mirror of
https://github.com/hansjone/oclaw.git
synced 2026-10-09 00:40:45 +08:00
fix(whatsapp): isolate group sessions and inject quoted context safely
Default WhatsApp groups to per-user sessions, skip redundant quoted-text injection when already in recent history, and disable group long-term memory writes unless explicitly enabled. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
d90fca7293
commit
fc10560052
8 changed files with 450 additions and 11 deletions
|
|
@ -191,6 +191,8 @@ def run_attempt(*, store: Any, data: AttemptRunnerInput) -> AttemptRunnerOutput:
|
|||
user_text=str(data.persisted_user_text if data.persisted_user_text is not None else data.msg.text or ""),
|
||||
assistant_text=outcome.final_text,
|
||||
turn_uuid=outcome.turn_uuid,
|
||||
channel=str(data.msg.channel or ""),
|
||||
metadata=data.msg.metadata if isinstance(data.msg.metadata, dict) else {},
|
||||
)
|
||||
if data.trace_id:
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -422,6 +422,26 @@ def _session_title_user_label(*, is_group: bool, external_user_id: str, group_na
|
|||
return "group"
|
||||
|
||||
|
||||
def _group_session_title_user_label(
|
||||
*,
|
||||
is_group: bool,
|
||||
session_scope: str,
|
||||
external_user_id: str,
|
||||
group_name: str,
|
||||
external_chat_id: str,
|
||||
) -> str:
|
||||
if not is_group:
|
||||
return str(external_user_id or "").strip() or "unknown"
|
||||
if str(session_scope or "").strip().lower() == "chat":
|
||||
return _session_title_user_label(
|
||||
is_group=is_group,
|
||||
external_user_id=external_user_id,
|
||||
group_name=group_name,
|
||||
external_chat_id=external_chat_id,
|
||||
)
|
||||
return str(external_user_id or "").strip() or "unknown"
|
||||
|
||||
|
||||
def _parse_generic_inbound(channel_name: str, payload: dict[str, Any]) -> InboundMessage:
|
||||
meta = payload.get("metadata") if isinstance(payload.get("metadata"), dict) else {}
|
||||
user_id = str(payload.get("user_id") or payload.get("external_user_id") or "").strip()
|
||||
|
|
@ -781,13 +801,17 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
|||
|
||||
account = store.find_user_by_channel_account(channel=inbound.channel, account_id=account_id) or {}
|
||||
from runtime.orchestration.group_ingest import (
|
||||
build_group_focus_instruction,
|
||||
build_group_quoted_context_block,
|
||||
build_group_sender_context,
|
||||
enrich_alert_group_question,
|
||||
extract_group_quoted_message,
|
||||
extract_quoted_ume_alert_text,
|
||||
mentions_include_bot,
|
||||
metadata_mentions_bot,
|
||||
resolve_group_policy,
|
||||
session_user_key,
|
||||
should_inject_quoted_context,
|
||||
should_process_group_inbound,
|
||||
text_mentions_bot,
|
||||
)
|
||||
|
|
@ -893,9 +917,11 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
|||
session_external_user_id = session_user_key(
|
||||
is_group=inbound.is_group,
|
||||
external_user_id=inbound.external_user_id,
|
||||
session_scope=group_policy.session_scope,
|
||||
)
|
||||
title_user_label = _session_title_user_label(
|
||||
title_user_label = _group_session_title_user_label(
|
||||
is_group=inbound.is_group,
|
||||
session_scope=group_policy.session_scope,
|
||||
external_user_id=inbound.external_user_id,
|
||||
group_name=group_name,
|
||||
external_chat_id=inbound.external_chat_id,
|
||||
|
|
@ -990,7 +1016,20 @@ def process_inbound_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
|||
metadata=meta_for_group,
|
||||
external_user_id=inbound.external_user_id,
|
||||
)
|
||||
user_text = f"{sender_ctx}\n{user_text}" if user_text else sender_ctx
|
||||
group_rule = build_group_focus_instruction()
|
||||
quoted_ctx = ""
|
||||
quoted_info = extract_group_quoted_message(metadata=meta_for_group)
|
||||
quoted_text = str(quoted_info.get("quoted_text") or "").strip()
|
||||
if quoted_text:
|
||||
recent_messages = store.get_messages(str(session_id), limit=6)
|
||||
if should_inject_quoted_context(
|
||||
quoted_text=quoted_text,
|
||||
recent_messages=recent_messages,
|
||||
):
|
||||
quoted_ctx = build_group_quoted_context_block(metadata=meta_for_group)
|
||||
prefix_parts = [sender_ctx, group_rule, quoted_ctx]
|
||||
prefix = "\n".join(part for part in prefix_parts if str(part).strip())
|
||||
user_text = f"{prefix}\n{user_text}" if user_text else prefix
|
||||
gw_attachments = _channel_attachments_for_gateway(
|
||||
list(inbound.attachments or [])
|
||||
)
|
||||
|
|
|
|||
|
|
@ -135,6 +135,8 @@ def after_turn_memory(
|
|||
user_text: str,
|
||||
assistant_text: str,
|
||||
turn_uuid: str = "",
|
||||
channel: str = "",
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
ingest_after_turn(
|
||||
store=store,
|
||||
|
|
@ -156,6 +158,8 @@ def after_turn_memory(
|
|||
session_id=str(session_id or ""),
|
||||
user_text=str(user_text or ""),
|
||||
assistant_text=str(assistant_text or ""),
|
||||
channel=str(channel or ""),
|
||||
metadata=metadata if isinstance(metadata, dict) else {},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -154,6 +154,49 @@ def extract_quoted_ume_alert_text(*, metadata: dict[str, Any] | None) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def extract_group_quoted_message(*, metadata: dict[str, Any] | None) -> dict[str, str]:
|
||||
raw = _metadata_raw(metadata)
|
||||
quoted_text = str(raw.get("quotedText") or raw.get("quoted_text") or "").strip()
|
||||
quoted_participant = str(raw.get("quotedParticipant") or raw.get("quoted_participant") or "").strip()
|
||||
push_name = str(raw.get("quotedPushName") or raw.get("quoted_push_name") or "").strip()
|
||||
stanza_id = str(raw.get("quotedStanzaId") or raw.get("quoted_stanza_id") or "").strip()
|
||||
return {
|
||||
"quoted_text": quoted_text,
|
||||
"quoted_participant": quoted_participant,
|
||||
"quoted_push_name": push_name,
|
||||
"quoted_stanza_id": stanza_id,
|
||||
}
|
||||
|
||||
|
||||
def _normalize_quoted_compare_text(text: str) -> str:
|
||||
s = re.sub(r"\s+", " ", str(text or "").strip())
|
||||
return s[:280]
|
||||
|
||||
|
||||
def should_inject_quoted_context(*, quoted_text: str, recent_messages: list[Any]) -> bool:
|
||||
token = _normalize_quoted_compare_text(quoted_text)
|
||||
if not token:
|
||||
return False
|
||||
for row in recent_messages or []:
|
||||
content = _normalize_quoted_compare_text(getattr(row, "content", "") or "")
|
||||
if not content:
|
||||
continue
|
||||
if token == content or token in content or content in token:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def build_group_quoted_context_block(*, metadata: dict[str, Any] | None) -> str:
|
||||
info = extract_group_quoted_message(metadata=metadata)
|
||||
quoted_text = str(info.get("quoted_text") or "").strip()
|
||||
if not quoted_text:
|
||||
return ""
|
||||
quoted_participant = str(info.get("quoted_participant") or "").strip()
|
||||
quoted_push_name = str(info.get("quoted_push_name") or "").strip()
|
||||
speaker = quoted_push_name or quoted_participant or "unknown"
|
||||
return f"[被引用消息]\n{speaker}: {quoted_text}"
|
||||
|
||||
|
||||
def enrich_alert_group_question(*, user_text: str, quoted_alert: str) -> str:
|
||||
body = str(user_text or "").strip()
|
||||
quote = str(quoted_alert or "").strip()
|
||||
|
|
@ -173,8 +216,21 @@ def normalize_jids(jids: list[str]) -> set[str]:
|
|||
return out
|
||||
|
||||
|
||||
def session_user_key(*, is_group: bool, external_user_id: str) -> str:
|
||||
return GROUP_SESSION_USER_SENTINEL if is_group else str(external_user_id or "").strip()
|
||||
def normalize_group_session_scope(raw: Any) -> str:
|
||||
scope = str(raw or "").strip().lower()
|
||||
if scope in {"chat", "shared", "shared_chat"}:
|
||||
return "chat"
|
||||
if scope in {"user", "user_in_chat", "per_user", "member"}:
|
||||
return "user_in_chat"
|
||||
return "user_in_chat"
|
||||
|
||||
|
||||
def session_user_key(*, is_group: bool, external_user_id: str, session_scope: str = "user_in_chat") -> str:
|
||||
if not is_group:
|
||||
return str(external_user_id or "").strip()
|
||||
if normalize_group_session_scope(session_scope) == "chat":
|
||||
return GROUP_SESSION_USER_SENTINEL
|
||||
return str(external_user_id or "").strip()
|
||||
|
||||
|
||||
def infer_is_group_from_chat_id(chat_id: str) -> bool:
|
||||
|
|
@ -206,7 +262,7 @@ def should_send_channel_reply_text(text: str) -> bool:
|
|||
class GroupPolicyConfig:
|
||||
require_mention: bool = True
|
||||
triggers: tuple[str, ...] = ("/oclaw",)
|
||||
session_scope: str = "chat"
|
||||
session_scope: str = "user_in_chat"
|
||||
|
||||
|
||||
def _parse_bool_env(name: str, default: bool) -> bool:
|
||||
|
|
@ -236,7 +292,7 @@ def _parse_group_policy_dict(raw: Any) -> GroupPolicyConfig | None:
|
|||
return GroupPolicyConfig(
|
||||
require_mention=bool(require_mention) if require_mention is not None else True,
|
||||
triggers=triggers if triggers is not None else ("/oclaw",),
|
||||
session_scope=str(session_scope or "chat").strip() or "chat",
|
||||
session_scope=normalize_group_session_scope(session_scope),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -252,6 +308,7 @@ def resolve_group_policy(*, account: dict[str, Any] | None = None) -> GroupPolic
|
|||
return GroupPolicyConfig(
|
||||
require_mention=_parse_bool_env("AIA_WHATSAPP_GROUP_REQUIRE_MENTION", True),
|
||||
triggers=_parse_triggers_env("AIA_WHATSAPP_GROUP_TRIGGERS", ("/oclaw", "|oclaw")),
|
||||
session_scope=normalize_group_session_scope(os.environ.get("AIA_WHATSAPP_GROUP_SESSION_SCOPE")),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -304,6 +361,15 @@ def build_group_sender_context(*, metadata: dict[str, Any] | None, external_user
|
|||
return f"[群成员: {label}]"
|
||||
|
||||
|
||||
def build_group_focus_instruction(*, lang: str = "zh") -> str:
|
||||
if str(lang or "").strip().lower().startswith("en"):
|
||||
return (
|
||||
"[Group chat rule: answer only the current sender's request. "
|
||||
"Do not assume context from other members unless this message explicitly quotes or references it.]"
|
||||
)
|
||||
return "[群聊规则:只回答当前发言人的问题;除非本条消息明确引用或承接前文,否则不要默认继承其他群成员的上下文。]"
|
||||
|
||||
|
||||
def build_whatsapp_group_reply_metadata(
|
||||
*,
|
||||
inbound: Any,
|
||||
|
|
@ -335,18 +401,23 @@ __all__ = [
|
|||
"GROUP_SESSION_USER_SENTINEL",
|
||||
"GroupPolicyConfig",
|
||||
"build_group_sender_context",
|
||||
"build_group_focus_instruction",
|
||||
"build_group_quoted_context_block",
|
||||
"build_whatsapp_group_reply_metadata",
|
||||
"enrich_alert_group_question",
|
||||
"extract_group_quoted_message",
|
||||
"extract_quoted_ume_alert_text",
|
||||
"mentions_include_bot",
|
||||
"metadata_mentions_bot",
|
||||
"normalize_jid",
|
||||
"normalize_group_session_scope",
|
||||
"normalize_jids",
|
||||
"infer_is_group_from_chat_id",
|
||||
"is_nonsend_channel_reply_text",
|
||||
"resolve_is_group",
|
||||
"resolve_group_policy",
|
||||
"session_user_key",
|
||||
"should_inject_quoted_context",
|
||||
"should_process_group_inbound",
|
||||
"should_send_channel_reply_text",
|
||||
"text_mentions_bot",
|
||||
|
|
|
|||
|
|
@ -192,10 +192,22 @@ def maybe_write_turn_memory(
|
|||
session_id: str,
|
||||
user_text: str,
|
||||
assistant_text: str,
|
||||
channel: str = "",
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
runtime = read_vector_memory_runtime(store)
|
||||
if not runtime.writer_enabled:
|
||||
return {"ok": True, "written": 0, "reason": "writer_disabled"}
|
||||
md = metadata if isinstance(metadata, dict) else {}
|
||||
is_group = bool(md.get("is_group"))
|
||||
if not is_group:
|
||||
raw = md.get("raw")
|
||||
if isinstance(raw, dict):
|
||||
is_group = bool(raw.get("is_group")) or str(raw.get("remoteJid") or "").strip().lower().endswith("@g.us")
|
||||
if str(channel or "").strip().lower() == "whatsapp" and is_group:
|
||||
allow_group = str(store.get_setting("MEMORY_WRITE_GROUP_WHATSAPP") or "").strip().lower()
|
||||
if allow_group not in {"1", "true", "yes", "on"}:
|
||||
return {"ok": True, "written": 0, "reason": "group_whatsapp_disabled"}
|
||||
user_norm = _normalize_memory_text(user_text)
|
||||
assistant_norm = _normalize_memory_text(assistant_text)
|
||||
if not user_norm or not assistant_norm:
|
||||
|
|
|
|||
|
|
@ -4,13 +4,18 @@ import pytest
|
|||
|
||||
from runtime.orchestration.group_ingest import (
|
||||
GROUP_SESSION_USER_SENTINEL,
|
||||
build_group_focus_instruction,
|
||||
build_group_quoted_context_block,
|
||||
build_group_sender_context,
|
||||
build_whatsapp_group_reply_metadata,
|
||||
extract_group_quoted_message,
|
||||
infer_is_group_from_chat_id,
|
||||
is_nonsend_channel_reply_text,
|
||||
normalize_jid,
|
||||
normalize_group_session_scope,
|
||||
resolve_group_policy,
|
||||
session_user_key,
|
||||
should_inject_quoted_context,
|
||||
should_process_group_inbound,
|
||||
text_mentions_bot,
|
||||
)
|
||||
|
|
@ -237,11 +242,19 @@ def test_enrich_quoted_ume_alarm_when_mentioned() -> None:
|
|||
assert "what happened?" in enriched
|
||||
|
||||
|
||||
def test_session_user_key_group_sentinel() -> None:
|
||||
assert session_user_key(is_group=True, external_user_id="111@s.whatsapp.net") == GROUP_SESSION_USER_SENTINEL
|
||||
def test_session_user_key_group_scope() -> None:
|
||||
assert session_user_key(is_group=True, external_user_id="111@s.whatsapp.net") == "111@s.whatsapp.net"
|
||||
assert session_user_key(is_group=True, external_user_id="111@s.whatsapp.net", session_scope="chat") == GROUP_SESSION_USER_SENTINEL
|
||||
assert session_user_key(is_group=False, external_user_id="111@s.whatsapp.net") == "111@s.whatsapp.net"
|
||||
|
||||
|
||||
def test_normalize_group_session_scope() -> None:
|
||||
assert normalize_group_session_scope("chat") == "chat"
|
||||
assert normalize_group_session_scope("shared") == "chat"
|
||||
assert normalize_group_session_scope("user_in_chat") == "user_in_chat"
|
||||
assert normalize_group_session_scope("") == "user_in_chat"
|
||||
|
||||
|
||||
def test_resolve_group_policy_from_account_config() -> None:
|
||||
policy = resolve_group_policy(
|
||||
account={
|
||||
|
|
@ -249,12 +262,14 @@ def test_resolve_group_policy_from_account_config() -> None:
|
|||
"group_policy": {
|
||||
"require_mention": False,
|
||||
"triggers": ["!ask"],
|
||||
"session_scope": "chat",
|
||||
}
|
||||
}
|
||||
}
|
||||
)
|
||||
assert policy.require_mention is False
|
||||
assert policy.triggers == ("!ask",)
|
||||
assert policy.session_scope == "chat"
|
||||
|
||||
|
||||
def test_infer_is_group_from_chat_id() -> None:
|
||||
|
|
@ -345,6 +360,47 @@ def test_build_group_sender_context() -> None:
|
|||
assert "111@s.whatsapp.net" in ctx
|
||||
|
||||
|
||||
def test_build_group_focus_instruction() -> None:
|
||||
zh = build_group_focus_instruction()
|
||||
en = build_group_focus_instruction(lang="en")
|
||||
assert "群聊规则" in zh
|
||||
assert "current sender" in en
|
||||
|
||||
|
||||
def test_extract_group_quoted_message_and_build_context_block() -> None:
|
||||
meta = {
|
||||
"raw": {
|
||||
"quotedText": "服务器刚刚 502 了",
|
||||
"quotedParticipant": "111@s.whatsapp.net",
|
||||
"quotedPushName": "Alice",
|
||||
"quotedStanzaId": "Q1",
|
||||
}
|
||||
}
|
||||
info = extract_group_quoted_message(metadata=meta)
|
||||
assert info["quoted_text"] == "服务器刚刚 502 了"
|
||||
assert info["quoted_push_name"] == "Alice"
|
||||
block = build_group_quoted_context_block(metadata=meta)
|
||||
assert "[被引用消息]" in block
|
||||
assert "Alice" in block
|
||||
|
||||
|
||||
def test_should_inject_quoted_context_dedupes_recent_message() -> None:
|
||||
from svc.persistence.sqlite_store import ChatMessage
|
||||
|
||||
recent = [
|
||||
ChatMessage(
|
||||
id=1,
|
||||
session_id="s1",
|
||||
role="assistant",
|
||||
content="服务器刚刚 502 了",
|
||||
tool_calls=None,
|
||||
timestamp="",
|
||||
)
|
||||
]
|
||||
assert should_inject_quoted_context(quoted_text="服务器刚刚 502 了", recent_messages=recent) is False
|
||||
assert should_inject_quoted_context(quoted_text="另一条内容", recent_messages=recent) is True
|
||||
|
||||
|
||||
def test_build_whatsapp_group_reply_metadata() -> None:
|
||||
inbound = InboundMessage(
|
||||
channel="whatsapp",
|
||||
|
|
@ -367,7 +423,7 @@ def test_build_whatsapp_group_reply_metadata() -> None:
|
|||
assert meta["quote_participant"] == "111:12@s.whatsapp.net"
|
||||
|
||||
|
||||
def test_shared_group_session_for_multiple_senders(fresh_sqlite_store: SqliteStore) -> None:
|
||||
def test_default_group_session_is_per_user(fresh_sqlite_store: SqliteStore) -> None:
|
||||
store = fresh_sqlite_store
|
||||
tenant = store.create_tenant("WA")
|
||||
chat_id = "120363012345678@g.us"
|
||||
|
|
@ -387,6 +443,37 @@ def test_shared_group_session_for_multiple_senders(fresh_sqlite_store: SqliteSto
|
|||
external_user_id=session_user_key(is_group=True, external_user_id="222@s.whatsapp.net"),
|
||||
session_title="whatsapp|test+Family",
|
||||
)
|
||||
assert sid_a != sid_b
|
||||
|
||||
|
||||
def test_shared_group_session_for_multiple_senders_when_scope_chat(fresh_sqlite_store: SqliteStore) -> None:
|
||||
store = fresh_sqlite_store
|
||||
tenant = store.create_tenant("WA")
|
||||
chat_id = "120363012345678@g.us"
|
||||
sid_a = store.get_or_create_channel_session_v2(
|
||||
tenant_id=str(tenant["id"]),
|
||||
channel="whatsapp",
|
||||
account_id="wa-default",
|
||||
external_chat_id=chat_id,
|
||||
external_user_id=session_user_key(
|
||||
is_group=True,
|
||||
external_user_id="111@s.whatsapp.net",
|
||||
session_scope="chat",
|
||||
),
|
||||
session_title="whatsapp|test+Family",
|
||||
)
|
||||
sid_b = store.get_or_create_channel_session_v2(
|
||||
tenant_id=str(tenant["id"]),
|
||||
channel="whatsapp",
|
||||
account_id="wa-default",
|
||||
external_chat_id=chat_id,
|
||||
external_user_id=session_user_key(
|
||||
is_group=True,
|
||||
external_user_id="222@s.whatsapp.net",
|
||||
session_scope="chat",
|
||||
),
|
||||
session_title="whatsapp|test+Family",
|
||||
)
|
||||
assert sid_a == sid_b
|
||||
|
||||
|
||||
|
|
@ -412,6 +499,13 @@ def _setup_whatsapp_identity(store: SqliteStore, *, extra_user_ids: list[str] |
|
|||
external_user_id=ext_uid,
|
||||
user_id=user_id,
|
||||
)
|
||||
store.upsert_whatsapp_contact(
|
||||
tenant_id=tenant_id,
|
||||
account_id="wa-default",
|
||||
external_user_id=ext_uid,
|
||||
phone="".join(ch for ch in ext_uid.split("@", 1)[0] if ch.isdigit()),
|
||||
list_type="whitelist",
|
||||
)
|
||||
return tenant_id, user_id
|
||||
|
||||
|
||||
|
|
@ -475,7 +569,7 @@ def test_inbound_dm_still_processes_without_mention(monkeypatch: pytest.MonkeyPa
|
|||
assert "[群成员:" not in captured.get("text", "")
|
||||
|
||||
|
||||
def test_inbound_group_mention_uses_shared_session_and_sender_prefix(
|
||||
def test_inbound_group_mention_uses_per_user_session_and_sender_prefix(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh_sqlite_store: SqliteStore
|
||||
) -> None:
|
||||
store = fresh_sqlite_store
|
||||
|
|
@ -530,10 +624,92 @@ def test_inbound_group_mention_uses_shared_session_and_sender_prefix(
|
|||
)
|
||||
|
||||
assert len(session_ids) == 2
|
||||
assert session_ids[0] == session_ids[1]
|
||||
assert session_ids[0] != session_ids[1]
|
||||
assert "[群成员:" in captured["text"]
|
||||
assert "群聊规则" in captured["text"]
|
||||
assert "Bob" in captured["text"]
|
||||
|
||||
sid = store.get_or_create_channel_session_v2(
|
||||
tenant_id=tenant_id,
|
||||
channel="whatsapp",
|
||||
account_id="wa-default",
|
||||
external_chat_id=chat_id,
|
||||
external_user_id="111@s.whatsapp.net",
|
||||
session_title="whatsapp|wa-default+Family",
|
||||
)
|
||||
assert sid == session_ids[0]
|
||||
|
||||
_ = tenant_id, user_id
|
||||
|
||||
|
||||
def test_inbound_group_scope_chat_uses_shared_session(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh_sqlite_store: SqliteStore
|
||||
) -> None:
|
||||
store = fresh_sqlite_store
|
||||
tenant_id, user_id = _setup_whatsapp_identity(store, extra_user_ids=["222@s.whatsapp.net"])
|
||||
store.upsert_user_channel_account(
|
||||
tenant_id=tenant_id,
|
||||
user_id=user_id,
|
||||
channel="whatsapp",
|
||||
account_id="wa-default",
|
||||
name="wa-default",
|
||||
config={
|
||||
"group_policy": {
|
||||
"require_mention": True,
|
||||
"triggers": ["/oclaw"],
|
||||
"session_scope": "chat",
|
||||
}
|
||||
},
|
||||
is_active=True,
|
||||
)
|
||||
monkeypatch.setattr("svc.persistence.assistant_store.get_assistant_store", lambda: store)
|
||||
|
||||
session_ids: list[str] = []
|
||||
|
||||
class _Turn:
|
||||
turn_uuid = "turn-g"
|
||||
reply_text = "group-ok"
|
||||
|
||||
class _Gw:
|
||||
def __init__(self, *, store: object) -> None:
|
||||
_ = store
|
||||
|
||||
def handle_turn(self, **kwargs: object) -> _Turn:
|
||||
msg = kwargs.get("msg")
|
||||
session_ids.append(str(getattr(msg, "session_id", "") or ""))
|
||||
return _Turn()
|
||||
|
||||
monkeypatch.setattr("runtime.gateway.OclawGateway", _Gw)
|
||||
|
||||
chat_id = "120363012345678@g.us"
|
||||
base = {
|
||||
"channel": "whatsapp",
|
||||
"account_id": "wa-default",
|
||||
"chat_id": chat_id,
|
||||
"is_group": True,
|
||||
"metadata": {"bot_jid": "999@s.whatsapp.net", "source": "test"},
|
||||
}
|
||||
process_inbound_payload(
|
||||
{
|
||||
**base,
|
||||
"user_id": "111@s.whatsapp.net",
|
||||
"text": "@bot hi",
|
||||
"mentions": ["999@s.whatsapp.net"],
|
||||
"metadata": {**base["metadata"], "raw": {"pushName": "Alice"}},
|
||||
}
|
||||
)
|
||||
process_inbound_payload(
|
||||
{
|
||||
**base,
|
||||
"user_id": "222@s.whatsapp.net",
|
||||
"text": "@bot again",
|
||||
"mentions": ["999@s.whatsapp.net"],
|
||||
"metadata": {**base["metadata"], "raw": {"pushName": "Bob"}},
|
||||
}
|
||||
)
|
||||
assert len(session_ids) == 2
|
||||
assert session_ids[0] == session_ids[1]
|
||||
|
||||
sid = store.get_or_create_channel_session_v2(
|
||||
tenant_id=tenant_id,
|
||||
channel="whatsapp",
|
||||
|
|
@ -543,7 +719,124 @@ def test_inbound_group_mention_uses_shared_session_and_sender_prefix(
|
|||
session_title="whatsapp|wa-default+Family",
|
||||
)
|
||||
assert sid == session_ids[0]
|
||||
_ = user_id
|
||||
|
||||
|
||||
def test_inbound_group_injects_quoted_context_when_not_in_current_session(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh_sqlite_store: SqliteStore
|
||||
) -> None:
|
||||
store = fresh_sqlite_store
|
||||
tenant_id, user_id = _setup_whatsapp_identity(store, extra_user_ids=["222@s.whatsapp.net"])
|
||||
monkeypatch.setattr("svc.persistence.assistant_store.get_assistant_store", lambda: store)
|
||||
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class _Turn:
|
||||
turn_uuid = "turn-q"
|
||||
reply_text = "ok"
|
||||
|
||||
class _Gw:
|
||||
def __init__(self, *, store: object) -> None:
|
||||
_ = store
|
||||
|
||||
def handle_turn(self, **kwargs: object) -> _Turn:
|
||||
msg = kwargs.get("msg")
|
||||
captured["text"] = str(getattr(msg, "text", "") or "")
|
||||
return _Turn()
|
||||
|
||||
monkeypatch.setattr("runtime.gateway.OclawGateway", _Gw)
|
||||
|
||||
out = process_inbound_payload(
|
||||
{
|
||||
"channel": "whatsapp",
|
||||
"account_id": "wa-default",
|
||||
"user_id": "222@s.whatsapp.net",
|
||||
"chat_id": "120363012345678@g.us",
|
||||
"text": "@bot 这是不是和刚才发布有关?",
|
||||
"is_group": True,
|
||||
"mentions": ["999@s.whatsapp.net"],
|
||||
"metadata": {
|
||||
"bot_jid": "999@s.whatsapp.net",
|
||||
"raw": {
|
||||
"pushName": "Bob",
|
||||
"quotedText": "看起来像 OSPF 邻居抖动",
|
||||
"quotedParticipant": "999@s.whatsapp.net",
|
||||
"quotedPushName": "oclaw",
|
||||
"quotedStanzaId": "Q2",
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
assert out.get("ok") is True
|
||||
assert "[被引用消息]" in captured["text"]
|
||||
assert "看起来像 OSPF 邻居抖动" in captured["text"]
|
||||
_ = tenant_id, user_id
|
||||
|
||||
|
||||
def test_inbound_group_skips_quoted_context_when_already_in_current_session(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh_sqlite_store: SqliteStore
|
||||
) -> None:
|
||||
store = fresh_sqlite_store
|
||||
tenant_id, user_id = _setup_whatsapp_identity(store)
|
||||
monkeypatch.setattr("svc.persistence.assistant_store.get_assistant_store", lambda: store)
|
||||
|
||||
call_no = {"n": 0}
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
class _Turn:
|
||||
turn_uuid = "turn-q2"
|
||||
|
||||
@property
|
||||
def reply_text(self) -> str:
|
||||
return "看起来像 OSPF 邻居抖动" if call_no["n"] == 1 else "继续分析"
|
||||
|
||||
class _Gw:
|
||||
def __init__(self, *, store: object) -> None:
|
||||
_ = store
|
||||
|
||||
def handle_turn(self, **kwargs: object) -> _Turn:
|
||||
call_no["n"] += 1
|
||||
msg = kwargs.get("msg")
|
||||
captured["text"] = str(getattr(msg, "text", "") or "")
|
||||
return _Turn()
|
||||
|
||||
monkeypatch.setattr("runtime.gateway.OclawGateway", _Gw)
|
||||
|
||||
process_inbound_payload(
|
||||
{
|
||||
"channel": "whatsapp",
|
||||
"account_id": "wa-default",
|
||||
"user_id": "111@s.whatsapp.net",
|
||||
"chat_id": "120363012345678@g.us",
|
||||
"text": "@bot 先帮我判断原因",
|
||||
"is_group": True,
|
||||
"mentions": ["999@s.whatsapp.net"],
|
||||
"metadata": {"bot_jid": "999@s.whatsapp.net", "raw": {"pushName": "Alice"}},
|
||||
}
|
||||
)
|
||||
out = process_inbound_payload(
|
||||
{
|
||||
"channel": "whatsapp",
|
||||
"account_id": "wa-default",
|
||||
"user_id": "111@s.whatsapp.net",
|
||||
"chat_id": "120363012345678@g.us",
|
||||
"text": "@bot 那下一步怎么排查?",
|
||||
"is_group": True,
|
||||
"mentions": ["999@s.whatsapp.net"],
|
||||
"metadata": {
|
||||
"bot_jid": "999@s.whatsapp.net",
|
||||
"raw": {
|
||||
"pushName": "Alice",
|
||||
"quotedText": "看起来像 OSPF 邻居抖动",
|
||||
"quotedParticipant": "999@s.whatsapp.net",
|
||||
"quotedPushName": "oclaw",
|
||||
"quotedStanzaId": "Q3",
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
assert out.get("ok") is True
|
||||
assert "[被引用消息]" not in captured["text"]
|
||||
_ = tenant_id, user_id
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -71,6 +71,21 @@ class MemoryVectorTests(unittest.TestCase):
|
|||
)
|
||||
self.assertEqual(hits, [])
|
||||
|
||||
def test_group_whatsapp_memory_write_disabled_by_default(self) -> None:
|
||||
res = maybe_write_turn_memory(
|
||||
self.store,
|
||||
tenant_id="t1",
|
||||
user_id="u1",
|
||||
session_id="s1",
|
||||
user_text="记住我喜欢喝黑咖啡,不加糖。",
|
||||
assistant_text="好的,我记住了你的偏好。",
|
||||
channel="whatsapp",
|
||||
metadata={"is_group": True},
|
||||
)
|
||||
self.assertTrue(res.get("ok"))
|
||||
self.assertEqual(int(res.get("written") or 0), 0)
|
||||
self.assertEqual(str(res.get("reason") or ""), "group_whatsapp_disabled")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -39,8 +39,11 @@ def test_after_turn_memory_invokes_maybe_write_turn_memory(tmp_path: Path, monke
|
|||
user_id="u1",
|
||||
user_text="hello",
|
||||
assistant_text="world",
|
||||
channel="whatsapp",
|
||||
metadata={"is_group": True},
|
||||
)
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["tenant_id"] == "t1"
|
||||
assert calls[0]["user_id"] == "u1"
|
||||
assert calls[0]["channel"] == "whatsapp"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue