diff --git a/README.md b/README.md index 961cf9b..fa036d6 100644 --- a/README.md +++ b/README.md @@ -38,13 +38,21 @@ replayed via `POST /api/v1/sync` on the next connect. uv pip install 'git+https://github.com/openmaxai/hermes-openmax.git@v0.1.5' # Or from a checkout: uv pip install -e /path/to/hermes-openmax -# The plugin is discovered through the Hermes entry point. For a directory -# checkout, this symlink is also supported: +# The plugin is discovered through the Hermes entry point. For a trusted +# directory checkout, this symlink is also supported; keep the sibling +# cws_agent_sdk/ directory in the checkout and install httpx + websockets in +# the Hermes environment: ln -s /path/to/hermes-openmax/hermes_openmax ~/.hermes/plugins/hermes-openmax # enable hermes-openmax in ~/.hermes/config.yaml, then restart: hermes gateway restart ``` +Directory loading bootstraps only the checkout's sibling `cws_agent_sdk` +package. It never installs dependencies at runtime. If that sibling package is +missing, startup logs an explicit installation error instead of leaving the +OpenMax platform silently offline. Prefer `hermes plugins install` or the pip +installation above for managed environments. + Put the secret in `~/.hermes/.env`: ```dotenv @@ -184,8 +192,9 @@ your local environment; never paste tokens into chat or commit them. See and deduplicated by conversation, work item, action, and resulting status. Missing source context, read-only actions, and notification failures do not break the underlying task operation. -- `send(metadata=...)` passes causation/interaction metadata through unchanged; - the server currently treats that metadata as opaque. +- `send(metadata=...)` preserves caller causation/interaction fields and adds + `agent_hop_count`, `agent_origin_member_id`, and `agent_trace_id` for the loop + guard; the server currently treats that metadata as opaque. - Core is authoritative for the Agent owner. The bridge reconciles the local `policy.json` cache at startup, every five minutes, after an internal WebSocket reconnect, and when an `agent.config.owner_changed` refresh hint arrives. A @@ -205,15 +214,40 @@ attachments. `workspace_members` provides directory, DM policy, and organization management. Connection remains explicitly unsupported: hermes-openmax does not register a Connection/`conn` tool and must not request credentials or simulate that surface. +Human access policy follows the Zylos safety contract: DM defaults to `owner`, +group scope defaults to `allowlist`, `disabled` blocks every group sender, and +an owner mention may bypass only a missing group registration. Owners remain +exempt from a configured group's `allowFrom`. Plain-text mentions use the Core +display name plus optional `CWS_SELF_ALIASES`, with a boundary check that keeps +`@Name` from matching `@NameSuffix`. + +Agent-to-Agent traffic intentionally has an additional Hermes loop guard. It is +fail-closed unless the relevant DM/group switch is enabled and the sender is in +`CWS_ALLOWED_AGENT_SENDERS`; Agent group traffic also requires a structured +mention and still passes group scope, registration, and `allowFrom`. The bridge +enforces propagated hop limits plus local duplicate and per-sender/conversation +turn budgets. Sender identity and structured mentions are authenticated by CWS; +hop metadata is defense in depth, not an authentication boundary. + +`silent` is bridge-only observation: admitted text updates a bounded in-memory +group history and the sync/read watermarks, but does not download attachments, +resolve work references, add an ack reaction, invoke billing, create a Hermes +turn, call the model, or reply. A later admitted message may receive the recent +bounded history as context. The vendored v1 conformance corpus still represents +this admission as `handle:true`; `contract/PROVENANCE.md` documents the Hermes +runtime overlay that consumes it before the host callback. + OpenMax WebSocket connectivity is transport health, not an installed IM channel. This adapter therefore does not call the channel-liveness snapshot endpoint. A runtime that owns IM channel processes must report its complete catalog-backed snapshot (for example, Feishu and Telegram) from their actual process health. -OpenMax group ingress is upstream-authorized by CWS. Group messages therefore +OpenMax group ingress is authenticated by CWS and then checked by the SDK bridge. +Group messages therefore intentionally have no per-member Hermes `user_id`; this preserves one shared -session per group while OpenMax enforces group scope, group allowlist, -allow-from, mention, smart, and silent policy before delivery. DM sessions +session per group while the bridge enforces group scope, group allowlist, +allow-from, mention, smart, Agent loop guards, and silent policy before runtime +delivery. DM sessions remain user/conversation-scoped. The bundled `hermes_openmax/skills/` docs preserve role boundaries, diff --git a/contract/PROVENANCE.md b/contract/PROVENANCE.md index 42528cb..425f6c1 100644 --- a/contract/PROVENANCE.md +++ b/contract/PROVENANCE.md @@ -8,3 +8,14 @@ Vendored from https://github.com/openmaxai/openmax-agent-sdk Per that repo's CONTRACT.md, passing fixtures/v1 against schemas/v1 is the definition of protocol conformance for any SDK in any language. Do not edit these files here — re-vendor from upstream and note the new commit. + +## Hermes runtime overlay + +The vendored v1 corpus classifies `silent` as an admitted `handle:true` policy +decision. Hermes preserves that normalization result for conformance, but its +production bridge consumes the admitted message into bounded bridge-owned +history and advances watermarks before the host callback. It does not deliver +the message into a Hermes session/model turn. This intentional runtime overlay +implements bridge-only observation without modifying the vendored contract; +changing the cross-runtime schema or fixtures requires an upstream SDK change +and a later re-vendor. diff --git a/cws_agent_sdk/access_policy.py b/cws_agent_sdk/access_policy.py index 057fc39..2bacf30 100644 --- a/cws_agent_sdk/access_policy.py +++ b/cws_agent_sdk/access_policy.py @@ -7,15 +7,16 @@ default False to prevent agent-to-agent chat loops. - Group: handle only when the agent is mentioned (@agent / all / all_agents), unless group_require_mention is disabled. -- Messages from SYSTEM senders: surfaced as handle=False by default - (delivered separately if the adapter wants lifecycle events). +- Messages from SYSTEM senders: delivered by default for scheduler/lifecycle + work and never participate in owner auto-binding. - Own messages: never handled (also enforced upstream in the bridge). """ from __future__ import annotations +import re from dataclasses import dataclass, field -from typing import Optional +from typing import Iterable, Optional from .types import InboundMessage @@ -24,19 +25,25 @@ class AccessPolicyConfig: # DM admission: "open" (any org member), "allowlist" (dm_allowlist + # owner), "owner" (bound owner only — zylos's default private model). - dm_policy: str = "open" + dm_policy: str = "owner" group_require_mention: bool = True allow_agent_senders: bool = False # let other agents' messages trigger us allow_sibling_dm: bool = False # same-owner agent DMs + agent_allowlist: list[str] = field(default_factory=list) + max_agent_hops: int = 4 + agent_turn_budget: int = 4 + agent_turn_window_s: float = 60.0 + agent_duplicate_window_s: float = 60.0 dm_allowlist: list[str] = field( default_factory=list ) # member_ids (dm_policy=allowlist) # System Member DMs (scheduler "dependencies ready", issue.activated, ...) # DRIVE the task flow — zylos lets them straight through. Default True. handle_system: bool = True - group_policy: str = "open" # open | allowlist | disabled + group_policy: str = "allowlist" # open | allowlist | disabled group_configs: dict[str, dict] = field(default_factory=dict) self_display_name: str = "" + self_aliases: list[str] = field(default_factory=list) @dataclass @@ -45,7 +52,19 @@ class AccessDecision: reason: str -def _is_mentioned(msg: InboundMessage, self_member_id: str) -> bool: +def _text_mentions_name(text: str, name: str) -> bool: + if not text or not name: + return False + # Do not let @Name match @NameSuffix or @Name-team. Python's Unicode-aware + # \w also keeps the boundary correct for non-ASCII display names. + return bool(re.search(r"@" + re.escape(name) + r"(?![\w-])", text, re.IGNORECASE)) + + +def _is_mentioned( + msg: InboundMessage, + self_member_id: str, + self_names: Iterable[str] = (), +) -> bool: for m in msg.mentions or []: if isinstance(m, str) and m == self_member_id: return True @@ -61,6 +80,24 @@ def _is_mentioned(msg: InboundMessage, self_member_id: str) -> bool: ) if str(target or "") == self_member_id: return True + return any(_text_mentions_name(msg.text or "", name) for name in self_names if name) + + +def _is_directly_mentioned(msg: InboundMessage, self_member_id: str) -> bool: + """Structured member mention only; broadcast mentions do not trigger Agents.""" + for mention in msg.mentions or []: + if isinstance(mention, str) and mention == self_member_id: + return True + if not isinstance(mention, dict): + continue + target = ( + mention.get("member_id") + or mention.get("entity_id") + or mention.get("mentioned_id") + or mention.get("id") + ) + if str(target or "") == self_member_id: + return True return False @@ -84,21 +121,25 @@ def decide_inbound( return AccessDecision(cfg.handle_system, "system_sender") if msg.sender_type == "agent": + agent_allowed = "*" in cfg.agent_allowlist or msg.sender_id in cfg.agent_allowlist if conv_type == "dm": if ( cfg.allow_sibling_dm + and agent_allowed and owner_member_id and sender_owner_member_id == owner_member_id ): return AccessDecision(True, "sibling_dm_allowed") return AccessDecision(False, "agent_dm_blocked") - # Group: an agent sender only triggers us when explicitly allowed AND - # we are mentioned — both gates guard against agent-to-agent loops. - if cfg.allow_agent_senders and _is_mentioned(msg, self_member_id): - return AccessDecision(True, "agent_mention") - return AccessDecision(False, "agent_sender_blocked") + # Agent group traffic has an extra loop guard. Once admitted it still + # passes through the ordinary group scope/allowlist/allowFrom gates. + if ( + not cfg.allow_agent_senders + or not agent_allowed + or not _is_directly_mentioned(msg, self_member_id) + ): + return AccessDecision(False, "agent_sender_blocked") - # Human sender. if conv_type == "dm": if is_owner: return AccessDecision(True, "dm_owner") # owner always exempt @@ -112,18 +153,18 @@ def decide_inbound( return AccessDecision(True, "dm") # Group / broadcast / bridge conversations. - mentioned = _is_mentioned(msg, self_member_id) - if not mentioned and cfg.self_display_name: - mentioned = ( - f"@{cfg.self_display_name}".casefold() in (msg.text or "").casefold() - ) - if is_owner and mentioned: - return AccessDecision(True, "group_owner_mention") # owner @-bypass + self_names = ( + (cfg.self_display_name, *cfg.self_aliases) + if msg.sender_type == "human" + else () + ) + mentioned = _is_mentioned(msg, self_member_id, self_names) group = cfg.group_configs.get(msg.conversation_id) - policy = (cfg.group_policy or "open").lower() + policy = (cfg.group_policy or "allowlist").lower() if policy == "disabled": return AccessDecision(False, "group_disabled") - if policy == "allowlist" and group is None: + owner_mention_bypass = policy == "allowlist" and group is None and is_owner and mentioned + if policy == "allowlist" and group is None and not owner_mention_bypass: return AccessDecision(False, "group_not_allowlisted") group = group or {} allow_from = [ @@ -135,7 +176,12 @@ def decide_inbound( ) or [] ] - if allow_from and "*" not in allow_from and msg.sender_id not in allow_from: + if ( + allow_from + and "*" not in allow_from + and msg.sender_id not in allow_from + and not is_owner + ): return AccessDecision(False, "group_sender_not_allowed") mode = str(group.get("mode") or "").lower() if mode == "silent": @@ -143,5 +189,8 @@ def decide_inbound( if mode == "smart" or not cfg.group_require_mention: return AccessDecision(True, "group_open") if mentioned: - return AccessDecision(True, "group_mention") + return AccessDecision( + True, + "group_owner_mention" if owner_mention_bypass else "group_mention", + ) return AccessDecision(False, "group_no_mention") diff --git a/cws_agent_sdk/bridge.py b/cws_agent_sdk/bridge.py index ffb732c..0f42172 100644 --- a/cws_agent_sdk/bridge.py +++ b/cws_agent_sdk/bridge.py @@ -15,10 +15,12 @@ from __future__ import annotations import asyncio -from collections import OrderedDict +import hashlib +import time +from collections import OrderedDict, deque from typing import Awaitable, Callable, Optional -from .access_policy import AccessPolicyConfig, decide_inbound +from .access_policy import AccessPolicyConfig, _is_mentioned, decide_inbound from .codec import FRAME_MESSAGE, FRAME_SYNC, FRAME_SYSTEM, Frame, encode_typing from .config import CwsConfig from .errors import CwsApiError @@ -50,10 +52,21 @@ _DEDUP_PERSIST_EVERY = 25 _GROUP_HISTORY_LEN = 10 _GROUP_CONTEXT_N = 5 +_AGENT_TURN_KEYS_MAX = 2048 DM_REJECT_NOTICE = ( "你好,我暂时无法处理这条私信(访问策略限制)。请联系我的 owner 开通权限。" ) +GROUP_REJECT_NOTICES = { + "group_disabled": "Sorry, group chat is currently disabled.", + "group_not_allowlisted": "Sorry, this group is not enabled for this agent.", + "group_sender_not_allowed": "Sorry, you are not allowed to trigger this agent in this group.", +} + +_VALID_DM_POLICIES = {"open", "allowlist", "owner"} +_VALID_GROUP_POLICIES = {"open", "allowlist", "disabled"} +_VALID_GROUP_MODES = {"mention", "smart", "silent"} +_VALID_LIST_ACTIONS = {"add", "remove", "set"} SMART_MODE_HINT = ( "You were not mentioned. Decide whether to respond. Do NOT " @@ -76,7 +89,7 @@ def __init__( policy: Optional[AccessPolicyConfig] = None, version: str = "", runtime_state: Optional[RuntimeStateProvider] = None, - billing_gate_enabled: bool = True, + billing_gate_enabled: bool = False, metrics_interval_s: float = 300.0, control_sync_interval_s: float = 300.0, on_config_event: Optional[Callable[[str, dict], Awaitable[None]]] = None, @@ -134,6 +147,7 @@ def __init__( ) self._marks_since_persist = 0 self._inflight: set[str] = set() + self._inflight_done: dict[str, asyncio.Event] = {} self._own_client_msg_ids: "OrderedDict[str, None]" = OrderedDict() self._group_history: dict[str, list[str]] = {} # conv_id -> recent "name: text" self._last_reject_notice: dict[str, float] = {} @@ -146,29 +160,39 @@ def __init__( self._participants: "OrderedDict[str, set]" = ( OrderedDict() ) # conv_id -> display names seen + self._agent_turns: "OrderedDict[str, deque[float]]" = OrderedDict() + self._agent_fingerprints: "OrderedDict[str, float]" = OrderedDict() + self._agent_loop_admitted: "OrderedDict[str, None]" = OrderedDict() self._sync_seq: int = int( (storage.read_json(_SYNC_SEQ_KEY) or {}).get("seq", 0) ) self._sync_lock = asyncio.Lock() self._owner_sync_lock = asyncio.Lock() + self._loop: Optional[asyncio.AbstractEventLoop] = None self._running = False # -- lifecycle --------------------------------------------------------- async def start(self) -> None: # Fail fast on bad credentials before going async. + self._loop = asyncio.get_running_loop() await self._tokens.get_access_token() await self._resolve_identity() self._running = True + self._ws.start() + await self._ws.wait_until_connected() if self._cfg.member_id: - await self._online.report(self._cfg.member_id) + # This is a boot/onboarding report, not proof of a model turn. Do + # not emit it until the transport has completed a real handshake. + self._spawn_bg( + self._online.report(self._cfg.member_id), "cws-online-report" + ) self._metrics_task = asyncio.create_task( self._metrics_loop(), name="cws-metrics" ) self._control_sync_task = asyncio.create_task( self._control_sync_loop(), name="cws-control-sync" ) - self._ws.start() # First install seeks to the inbox end; later starts replay from the # persisted cursor. Both run off the connect path. self._spawn_bg(self._initialize_or_sync(), "cws-initial-sync") @@ -192,6 +216,7 @@ async def stop(self) -> None: await self._ws.stop() await self._http.aclose() await self._tokens.aclose() + self._loop = None async def _resolve_identity(self) -> None: """Fill member_id from /me when unset; pull authoritative owner.""" @@ -319,6 +344,17 @@ def is_running(self) -> bool: # -- outbound ------------------------------------------------------------ + def _outbound_agent_metadata(self, metadata: Optional[dict], cmid: str) -> dict: + result = dict(metadata or {}) + try: + prior_hop = int(result.get("agent_hop_count") or 0) + except (TypeError, ValueError): + prior_hop = 0 + result["agent_hop_count"] = prior_hop + 1 + result.setdefault("agent_origin_member_id", self._cfg.member_id) + result.setdefault("agent_trace_id", cmid) + return result + async def send( self, conversation_id: str, @@ -331,11 +367,12 @@ async def send( cmid = new_client_msg_id() self._remember_own(cmid) + outbound_metadata = self._outbound_agent_metadata(metadata, cmid) receipt = await self.comm.send_message( conversation_id, self._canonicalize_mentions(conversation_id, content), reply_to=reply_to, - metadata=metadata, + metadata=outbound_metadata, client_msg_id=cmid, ) # Replying resolves the pending received-ack (zylos parity) and any @@ -354,6 +391,7 @@ async def send_image_file( *, caption: str = "", reply_to: Optional[str] = None, + metadata: Optional[dict] = None, ) -> SendReceipt: """Upload a local image (presigned two-phase) and send it as a native IMAGE message with a proper attachment.""" @@ -378,6 +416,10 @@ async def send_image_file( ) resp.raise_for_status() fin = await self.artifacts.finalize_conversation_upload(prep["upload_token"]) + from .codec import new_client_msg_id + + cmid = new_client_msg_id() + self._remember_own(cmid) receipt = await self.comm.send_image_message( conversation_id, artifact_id=str(fin.get("artifact_id", "")), @@ -386,6 +428,8 @@ async def send_image_file( size_bytes=size, caption=caption, reply_to=reply_to, + client_msg_id=cmid, + metadata=self._outbound_agent_metadata(metadata, cmid), ) await self._clear_ack(conversation_id) return receipt @@ -412,12 +456,13 @@ async def _handle_frame(self, frame: Frame) -> None: if frame.type == FRAME_MESSAGE: await self._handle_message_frame(frame) elif frame.type == FRAME_SYNC: - for ev in frame.payload.get("events") or []: - await self._deliver_by_id( - str(ev.get("conversation_id", "")), - ev.get("message_id"), - int(ev.get("seq") or 0), - ) + async with self._sync_lock: + for ev in frame.payload.get("events") or []: + await self._deliver_by_id( + str(ev.get("conversation_id", "")), + ev.get("message_id"), + int(ev.get("seq") or 0), + ) elif frame.type == FRAME_SYSTEM: await self._handle_system_frame(frame) # typing / read_state / presence / acks: no runtime delivery needed. @@ -441,7 +486,11 @@ async def _handle_system_frame(self, frame: Frame) -> None: self._log.log("config event:", event) if event == "agent.config.dm_allowlist_changed": action = str(data.get("action", "")).lower() - ids = [str(i) for i in data.get("member_ids") or []] + raw_ids = data.get("member_ids") + if action not in _VALID_LIST_ACTIONS or not isinstance(raw_ids, list): + self._log.warn("invalid dm allowlist event:", action) + return + ids = [str(i) for i in raw_ids if str(i)] if action == "add": for i in ids: if i not in self._policy.dm_allowlist: @@ -450,50 +499,63 @@ async def _handle_system_frame(self, frame: Frame) -> None: self._policy.dm_allowlist = [ i for i in self._policy.dm_allowlist if i not in ids ] - elif action == "set": + else: # set self._policy.dm_allowlist = ids self._save_policy_state() elif event == "agent.config.dm_policy_changed": policy = str(data.get("policy", "")).lower() - if policy in ("open", "allowlist", "owner"): - self._policy.dm_policy = policy - self._save_policy_state() + if policy not in _VALID_DM_POLICIES: + self._log.warn("invalid dm policy event:", policy) + return + self._policy.dm_policy = policy + self._save_policy_state() elif event == "agent.config.group_mode_changed": conv = str(data.get("conversation_id", "")) - if conv: - mode = str(data.get("mode", "")) - self._group_mode_overrides[conv] = mode - if mode == "silent": - self._policy.group_configs.pop(conv, None) - else: - group = self._policy.group_configs.setdefault( - conv, {"mode": "mention", "allow_from": ["*"]} - ) - group["mode"] = mode - self._save_policy_state() + mode = str(data.get("mode", "")).lower() + if not conv or mode not in _VALID_GROUP_MODES: + self._log.warn("invalid group mode event:", conv, mode) + return + self._group_mode_overrides[conv] = mode + group = self._policy.group_configs.setdefault( + conv, {"mode": "mention", "allow_from": ["*"]} + ) + group["mode"] = mode + self._save_policy_state() elif event == "agent.config.group_scope_changed": scope = str(data.get("scope", "")).lower() - if scope in ("open", "allowlist", "disabled"): - self._policy.group_policy = scope - self._save_policy_state() + if scope not in _VALID_GROUP_POLICIES: + self._log.warn("invalid group scope event:", scope) + return + self._policy.group_policy = scope + self._save_policy_state() elif event == "agent.config.group_allowlist_changed": + action = str(data.get("action", "")).lower() + raw_ids = data.get("conversation_ids") + if action not in _VALID_LIST_ACTIONS or not isinstance(raw_ids, list): + self._log.warn("invalid group allowlist event:", action) + return self._update_group_allowlist( - str(data.get("action", "")).lower(), - [str(v) for v in data.get("conversation_ids") or []], + action, + [str(v) for v in raw_ids if str(v)], ) elif event == "agent.config.group_allowfrom_changed": conv = str(data.get("conversation_id") or "") - if conv and isinstance(data.get("allow_from"), list): - group = self._policy.group_configs.setdefault( - conv, {"mode": "mention", "allow_from": ["*"]} - ) - group["allow_from"] = [str(v) for v in data["allow_from"]] - self._save_policy_state() + raw_allow = data.get("allow_from") + if not conv or not isinstance(raw_allow, list): + self._log.warn("invalid group allowFrom event:", conv) + return + group = self._policy.group_configs.setdefault( + conv, {"mode": "mention", "allow_from": ["*"]} + ) + group["allow_from"] = [str(v) for v in raw_allow if str(v)] + self._save_policy_state() elif event == "agent.config.owner_changed": # The pushed owner is only a refresh hint. Core's authenticated # member record remains authoritative. await self._sync_owner_from_core(notify=True) return + else: + return # Other events (dm_policy / group_scope / group_allowlist / allowfrom) # are forwarded to the adapter callback; interpretation is host policy. await self._report_policy() @@ -586,11 +648,16 @@ def _update_group_allowlist(self, action: str, conversation_ids: list[str]) -> N elif action == "remove": for conv in conversation_ids: groups.pop(conv, None) + self._group_mode_overrides.pop(conv, None) elif action == "set": old = dict(groups) + old_modes = dict(self._group_mode_overrides) groups.clear() + self._group_mode_overrides.clear() for conv in conversation_ids: groups[conv] = old.get(conv, {"mode": "mention", "allow_from": ["*"]}) + if conv in old_modes: + self._group_mode_overrides[conv] = old_modes[conv] else: return self._save_policy_state() @@ -604,6 +671,81 @@ async def _sender_owner(self, msg: InboundMessage) -> str: except Exception: # noqa: BLE001 return "" + def _agent_loop_rejection(self, msg: InboundMessage) -> str: + """Local fail-closed circuit breakers after Agent policy admission.""" + if msg.sender_type != "agent": + return "" + message_key = f"{msg.conversation_id}:{msg.message_id}" + if message_key in self._agent_loop_admitted: + return "" # delivery retry for the same un-acked message + if "agent_hop_count" not in msg.metadata: + hop = 1 + else: + raw_hop = msg.metadata["agent_hop_count"] + if isinstance(raw_hop, bool): + return "agent_hop_invalid" + if isinstance(raw_hop, int): + hop = raw_hop + elif isinstance(raw_hop, str): + try: + hop = int(raw_hop.strip()) + except (TypeError, ValueError): + return "agent_hop_invalid" + else: + return "agent_hop_invalid" + if hop < 1 or hop > self._policy.max_agent_hops: + return "agent_hop_limit" + + now = time.monotonic() + fingerprint_input = "\0".join( + ( + msg.conversation_id, + msg.sender_id, + msg.reply_to_message_id or "", + " ".join((msg.text or "").split()).casefold(), + ) + ) + fingerprint = hashlib.sha256(fingerprint_input.encode("utf-8")).hexdigest() + duplicate_window = self._policy.agent_duplicate_window_s + previous = self._agent_fingerprints.get(fingerprint) + if previous is not None and now - previous < duplicate_window: + return "agent_duplicate" + + turn_key = f"{msg.conversation_id}:{msg.sender_id}" + cutoff = now - self._policy.agent_turn_window_s + while self._agent_turns: + oldest_key, oldest_turns = next(iter(self._agent_turns.items())) + while oldest_turns and oldest_turns[0] <= cutoff: + oldest_turns.popleft() + if oldest_turns and len(self._agent_turns) <= _AGENT_TURN_KEYS_MAX: + break + self._agent_turns.pop(oldest_key, None) + turns = self._agent_turns.get(turn_key) + if turns is None: + while len(self._agent_turns) >= _AGENT_TURN_KEYS_MAX: + self._agent_turns.popitem(last=False) + turns = deque() + self._agent_turns[turn_key] = turns + else: + self._agent_turns.move_to_end(turn_key) + while turns and turns[0] <= cutoff: + turns.popleft() + if len(turns) >= self._policy.agent_turn_budget: + return "agent_turn_budget" + + turns.append(now) + self._agent_fingerprints[fingerprint] = now + self._agent_fingerprints.move_to_end(fingerprint) + while self._agent_fingerprints: + _, oldest = next(iter(self._agent_fingerprints.items())) + if len(self._agent_fingerprints) <= 2048 and now - oldest < duplicate_window: + break + self._agent_fingerprints.popitem(last=False) + self._agent_loop_admitted[message_key] = None + while len(self._agent_loop_admitted) > _DEDUP_MAX: + self._agent_loop_admitted.popitem(last=False) + return "" + async def _handle_message_frame(self, frame: Frame) -> None: p = frame.payload message_id = p.get("id") @@ -620,11 +762,28 @@ async def _deliver_by_id( self, conversation_id: str, message_id, seq: int, frame: Optional[Frame] = None ) -> None: key = f"{conversation_id}:{message_id}" - # The in-flight guard closes the WS-frame vs /sync-replay race: both - # can observe the same undelivered message concurrently. - if key in self._seen or key in self._inflight: + if key in self._seen: + # Realtime delivery never commits the global cursor. Serialized + # /sync later advances through already-seen messages in server order. + if frame is None and seq > 0: + await self._advance(conversation_id, 0, inbox_seq=seq) + return + # If /sync catches a realtime attempt still in progress, wait for its + # outcome. Success becomes an ordered cursor-only commit; failure is + # retried here and stops this serialized sync batch if it fails again. + if key in self._inflight: + if frame is not None: + return + done = self._inflight_done[key] + await done.wait() + if key in self._seen: + await self._advance(conversation_id, 0, inbox_seq=seq) + return + await self._deliver_by_id(conversation_id, message_id, seq) return self._inflight.add(key) + done = asyncio.Event() + self._inflight_done[key] = done try: detail = await self.comm.get_message(conversation_id, message_id) # /sync carries the org-level inbox watermark in event.seq. A @@ -636,29 +795,31 @@ async def _deliver_by_id( if isinstance(detail.get("message"), dict) else None ) - inbox_seq = int(seq or 0) if frame is None else int(detail_inbox or 0) - if inbox_seq > 0: + observed_inbox_seq = int(detail_inbox or 0) + inbox_seq = int(seq or 0) if frame is None else 0 + metadata_inbox_seq = observed_inbox_seq or inbox_seq + raw_message = detail.get("message") or detail + conversation_seq = int(raw_message.get("seq") or seq or 0) + if metadata_inbox_seq > 0: detail = dict(detail) - detail["_inbox_seq"] = inbox_seq + detail["_inbox_seq"] = metadata_inbox_seq msg = self._normalize(detail, conversation_id, seq) if msg is None: self._mark_seen(key) - await self._advance(conversation_id, seq, inbox_seq=seq) + await self._advance( + conversation_id, conversation_seq, inbox_seq=inbox_seq + ) return conv_info = await self._conversation_info(conversation_id) msg.conversation_type = conv_info["type"] if conv_info.get("name"): msg.metadata["conversation_name"] = conv_info["name"] - if not msg.sender_name and msg.sender_id: - msg.sender_name = await self._member_name(msg.sender_id) - if msg.sender_name: - self._record_participant(conversation_id, msg.sender_name) is_group = msg.conversation_type not in ("dm",) - if is_group: - # zylos group-history parity: record EVERY group message - # (handled or not) so later turns get conversation context. - self._record_group_history(conversation_id, msg) - mode = self._group_mode_overrides.get(conversation_id, "").lower() + effective_policy = self._effective_policy(conversation_id) + group_cfg = effective_policy.group_configs.get(conversation_id) or {} + mode = str(group_cfg.get("mode") or "").lower() + if not msg.sender_name and msg.sender_id and mode != "silent": + msg.sender_name = await self._member_name(msg.sender_id) # zylos parity: dm_policy=owner with no owner bound — the first # human DM sender becomes the owner (persisted; platform data # overrides on next identity resolve if it disagrees). @@ -672,20 +833,68 @@ async def _deliver_by_id( self.owner_member_id = msg.sender_id self._save_policy_state() self._log.log("owner auto-bound to first DM sender:", msg.sender_id) - await self._hydrate_media(msg) - await self._expand_work_references(msg) - await self._hydrate_reply_context(msg) decision = decide_inbound( msg, self_member_id=self._cfg.member_id, - cfg=self._effective_policy(conversation_id), + cfg=effective_policy, owner_member_id=self.owner_member_id, - sender_owner_member_id=await self._sender_owner(msg), + sender_owner_member_id=( + await self._sender_owner(msg) if not is_group else "" + ), ) - if decision.handle and is_group and ( - mode == "silent" or decision.reason == "group_silent" - ): - msg.metadata["group_silent"] = True + if decision.handle: + loop_rejection = self._agent_loop_rejection(msg) + if loop_rejection: + self._log.log( + f"policy skip [{loop_rejection}] conv={conversation_id} msg={message_id}" + ) + self._mark_seen(key) + await self._advance( + conversation_id, + msg.seq or conversation_seq, + inbox_seq=inbox_seq, + ) + return + context_eligible = decision.handle or decision.reason in ( + "group_no_mention", + "group_silent", + ) + if msg.sender_name and (not is_group or context_eligible): + self._record_participant(conversation_id, msg.sender_name) + if is_group and context_eligible: + # Allowed background traffic is cached in a small bridge-owned + # history window. Rejected groups/senders never enter context. + self._record_group_history(conversation_id, msg) + if not decision.handle: + self._log.log( + f"policy skip [{decision.reason}] conv={conversation_id} msg={message_id}" + ) + await self._maybe_send_reject_notice( + conversation_id, + msg, + decision.reason, + is_sync_replay=frame is None, + ) + self._mark_seen(key) + await self._advance( + conversation_id, + msg.seq or conversation_seq, + inbox_seq=inbox_seq, + ) + return + if is_group and (mode == "silent" or decision.reason == "group_silent"): + # Observe at the bridge only: no attachment/work hydration, + # billing lookup, ack reaction, Hermes session, or model turn. + self._mark_seen(key) + await self._advance( + conversation_id, + msg.seq or conversation_seq, + inbox_seq=inbox_seq, + ) + return + await self._hydrate_media(msg) + await self._expand_work_references(msg) + await self._hydrate_reply_context(msg) if decision.handle and is_group: history = self._group_history.get(conversation_id, []) recent = [h for h in history[:-1]][-_GROUP_CONTEXT_N:] @@ -693,9 +902,11 @@ async def _deliver_by_id( msg.metadata["group_context"] = ( "\n" + "\n".join(recent) + "\n" ) - from .access_policy import _is_mentioned - - if mode == "smart" and not _is_mentioned(msg, self._cfg.member_id): + if mode == "smart" and not _is_mentioned( + msg, + self._cfg.member_id, + (self._policy.self_display_name, *self._policy.self_aliases), + ): msg.metadata["smart_mode_hint"] = SMART_MODE_HINT if ( decision.handle @@ -703,31 +914,34 @@ async def _deliver_by_id( and await self._billing.is_suspended() ): self._log.warn("billing suspended — skipping delivery", conversation_id) - if self._billing.should_send_overdue_notice(conversation_id): + if ( + frame is not None + and self._billing.should_send_overdue_notice(conversation_id) + ): try: await self.comm.send_message(conversation_id, OVERDUE_NOTICE) except CwsApiError as exc: self._log.warn("overdue notice failed:", exc) self._mark_seen(key) - await self._advance(conversation_id, msg.seq or seq, inbox_seq=seq) - return - if not decision.handle: - self._log.log( - f"policy skip [{decision.reason}] conv={conversation_id} msg={message_id}" - ) - await self._maybe_send_reject_notice( - conversation_id, msg, decision.reason + await self._advance( + conversation_id, + msg.seq or conversation_seq, + inbox_seq=inbox_seq, ) - self._mark_seen(key) - await self._advance(conversation_id, msg.seq or seq, inbox_seq=seq) return await self._ack_received(conversation_id, str(message_id)) # Delivery point — exceptions propagate, watermark stays put, /sync replays. await self._on_message(msg) self._mark_seen(key) - await self._advance(conversation_id, msg.seq or seq, inbox_seq=seq) + await self._advance( + conversation_id, + msg.seq or conversation_seq, + inbox_seq=inbox_seq, + ) finally: self._inflight.discard(key) + self._inflight_done.pop(key, None) + done.set() # -- received-ack reaction (zylos parity: 👀 on receipt, cleared on reply) -- @@ -968,13 +1182,11 @@ def _effective_policy(self, conversation_id: str) -> AccessPolicyConfig: from dataclasses import replace low = mode.lower() - if "smart" in low: + if low == "smart": # smart: receive everything, model decides ([SKIP] to stay silent). return replace(self._policy, group_require_mention=False) - if "mention" in low: + if low == "mention": return replace(self._policy, group_require_mention=True) - if low in ("open", "all") or "open" in low: - return replace(self._policy, group_require_mention=False) if low == "silent": from copy import deepcopy @@ -992,13 +1204,18 @@ async def _conversation_info(self, conversation_id: str) -> dict: cached = self._conv_types.get(conversation_id) if cached: return cached - info_out = {"type": "dm", "name": ""} try: info = await self.comm.get_conversation(conversation_id) - info_out["type"] = str(info.get("type", "dm")).lower() or "dm" - info_out["name"] = str(info.get("name") or "") - except Exception as exc: # noqa: BLE001 — assume dm on failure - self._log.warn("get_conversation failed, assuming dm:", exc) + except Exception as exc: # noqa: BLE001 — unknown type must fail closed + self._log.warn("get_conversation failed; delivery remains pending:", exc) + raise + conversation_type = str(info.get("type") or "").lower() + if not conversation_type: + raise ValueError("conversation response is missing type") + info_out = { + "type": conversation_type, + "name": str(info.get("name") or ""), + } self._conv_types[conversation_id] = info_out while len(self._conv_types) > 512: self._conv_types.popitem(last=False) @@ -1108,22 +1325,36 @@ def _record_group_history(self, conversation_id: str, msg: InboundMessage) -> No del hist[:-_GROUP_HISTORY_LEN] async def _maybe_send_reject_notice( - self, conversation_id: str, msg: InboundMessage, reason: str + self, + conversation_id: str, + msg: InboundMessage, + reason: str, + *, + is_sync_replay: bool = False, ) -> None: - """zylos parity: polite rejection for human DMs blocked by policy — - throttled, never for agents/system, never for group-no-mention.""" - if reason not in ("dm_owner_only", "dm_not_allowlisted"): + """Send a throttled notice only for a live, actionable human rejection.""" + if msg.sender_type != "human" or is_sync_replay: return - if msg.sender_type != "human": + notice = "" + if reason in ("dm_owner_only", "dm_not_allowlisted"): + notice = DM_REJECT_NOTICE + elif reason in GROUP_REJECT_NOTICES and _is_mentioned( + msg, + self._cfg.member_id, + (self._policy.self_display_name, *self._policy.self_aliases), + ): + notice = GROUP_REJECT_NOTICES[reason] + if not notice: return import time as _time now = _time.time() - if now - self._last_reject_notice.get(conversation_id, 0) < 3600: + throttle_key = f"{conversation_id}:{reason}" + if now - self._last_reject_notice.get(throttle_key, 0) < 3600: return - self._last_reject_notice[conversation_id] = now + self._last_reject_notice[throttle_key] = now try: - await self.comm.send_message(conversation_id, DM_REJECT_NOTICE) + await self.comm.send_message(conversation_id, notice) except Exception as exc: # noqa: BLE001 — notice is best-effort self._log.warn("reject notice failed:", exc) @@ -1133,20 +1364,86 @@ def _load_policy_state(self) -> None: saved = self._storage.read_json(_POLICY_KEY) if not isinstance(saved, dict): return - if saved.get("dm_policy"): - self._policy.dm_policy = str(saved["dm_policy"]) + dm_policy = str(saved.get("dm_policy") or "").lower() + if dm_policy in _VALID_DM_POLICIES: + self._policy.dm_policy = dm_policy if saved.get("owner_member_id"): self.owner_member_id = str(saved["owner_member_id"]) if isinstance(saved.get("dm_allowlist"), list): - self._policy.dm_allowlist = [str(i) for i in saved["dm_allowlist"]] - if isinstance(saved.get("group_modes"), dict): - self._group_mode_overrides.update( - {str(k): str(v) for k, v in saved["group_modes"].items()} - ) - if saved.get("group_policy"): - self._policy.group_policy = str(saved["group_policy"]) + self._policy.dm_allowlist = [ + str(i) for i in saved["dm_allowlist"] if str(i) + ] + group_policy = str(saved.get("group_policy") or "").lower() + if group_policy in _VALID_GROUP_POLICIES: + self._policy.group_policy = group_policy + groups: dict[str, dict] = {} if isinstance(saved.get("group_configs"), dict): - self._policy.group_configs = dict(saved["group_configs"]) + for raw_conv, raw_config in saved["group_configs"].items(): + conv = str(raw_conv) + if not conv or not isinstance(raw_config, dict): + continue + mode = str(raw_config.get("mode") or "mention").lower() + if mode not in _VALID_GROUP_MODES: + mode = "mention" + raw_allow = ( + raw_config.get("allow_from") + if "allow_from" in raw_config + else raw_config.get("allowFrom") + ) + allow_from = ( + [str(v) for v in raw_allow if str(v)] + if isinstance(raw_allow, list) + else ["*"] + ) + groups[conv] = {"mode": mode, "allow_from": allow_from} + self._policy.group_configs = groups + if isinstance(saved.get("group_modes"), dict): + for raw_conv, raw_mode in saved["group_modes"].items(): + conv, mode = str(raw_conv), str(raw_mode).lower() + if not conv or mode not in _VALID_GROUP_MODES: + continue + self._group_mode_overrides[conv] = mode + group = self._policy.group_configs.setdefault( + conv, {"mode": "mention", "allow_from": ["*"]} + ) + group["mode"] = mode + + def get_dm_access(self) -> dict: + return { + "dm_policy": self._policy.dm_policy, + "dm_allowlist": list(self._policy.dm_allowlist), + } + + async def apply_local_dm_access( + self, dm_policy: str, dm_allowlist: list[str] + ) -> dict: + policy = str(dm_policy or "").lower() + if policy not in _VALID_DM_POLICIES: + raise ValueError("policy must be one of: open, allowlist, owner") + self._policy.dm_policy = policy + self._policy.dm_allowlist = list( + dict.fromkeys(str(v) for v in dm_allowlist if str(v)) + ) + self._save_policy_state() + self._spawn_bg(self._report_policy(), "cws-local-policy-report") + return self.get_dm_access() + + def apply_local_dm_access_threadsafe( + self, dm_policy: str, dm_allowlist: list[str] + ) -> dict: + loop = self._loop + if not loop or not loop.is_running(): + raise RuntimeError("CWS bridge event loop is not running") + try: + current_loop = asyncio.get_running_loop() + except RuntimeError: + current_loop = None + if current_loop is loop: + raise RuntimeError("cannot block the CWS bridge event loop") + future = asyncio.run_coroutine_threadsafe( + self.apply_local_dm_access(dm_policy, dm_allowlist), loop + ) + return future.result(timeout=10) def _save_policy_state(self) -> None: self._storage.write_json( diff --git a/cws_agent_sdk/services.py b/cws_agent_sdk/services.py index 49bbb42..63ecfdb 100644 --- a/cws_agent_sdk/services.py +++ b/cws_agent_sdk/services.py @@ -25,16 +25,24 @@ class AccessPolicyService: _KEY = "policy.json" _POLICIES = ("open", "allowlist", "owner") - def __init__(self, storage): + def __init__(self, storage, *, get_live_state=None, apply_live_state=None): self._storage = storage + self._get_live_state = get_live_state + self._apply_live_state = apply_live_state def _state(self) -> dict: + if self._get_live_state: + return dict(self._get_live_state()) state = self._storage.read_json(self._KEY) return dict(state) if isinstance(state, dict) else {} def _write(self, state: dict) -> dict: state["dm_policy"] = str(state.get("dm_policy") or "owner") state["dm_allowlist"] = [str(v) for v in state.get("dm_allowlist") or []] + if self._apply_live_state: + return self._apply_live_state( + state["dm_policy"], state["dm_allowlist"] + ) self._storage.write_json(self._KEY, state) return {"dm_policy": state["dm_policy"], "dm_allowlist": state["dm_allowlist"]} @@ -153,6 +161,7 @@ async def send_image_message( caption: str = "", reply_to: Optional[str] = None, client_msg_id: Optional[str] = None, + metadata: Optional[dict] = None, ) -> SendReceipt: """Send a native IMAGE message referencing a finalized upload.""" # zylos-openmax hard-won contract: body MUST carry file_name (an empty @@ -180,6 +189,8 @@ async def send_image_message( } if reply_to: body["parent_id"] = str(reply_to) + if metadata: + body["metadata"] = metadata data = await self._http.post( f"/api/v1/conversations/{conversation_id}/messages", json=body ) diff --git a/cws_agent_sdk/ws.py b/cws_agent_sdk/ws.py index b45f545..63e5a59 100644 --- a/cws_agent_sdk/ws.py +++ b/cws_agent_sdk/ws.py @@ -89,6 +89,29 @@ async def stop(self) -> None: def is_open(self) -> bool: return self._connected.is_set() + async def wait_until_connected(self) -> None: + """Wait for the first real handshake or propagate a fatal WS exit. + + Transient connection errors remain owned by the reconnect loop. The + host's platform-connect timeout bounds this wait during adapter startup. + """ + if self._connected.is_set(): + return + task = self._task + if task is None: + raise RuntimeError("cws ws client has not been started") + connected = asyncio.create_task(self._connected.wait()) + try: + done, _ = await asyncio.wait( + {connected, task}, return_when=asyncio.FIRST_COMPLETED + ) + if connected in done and connected.result(): + return + await task + raise ConnectionError("cws ws client stopped before connecting") + finally: + connected.cancel() + async def send_text(self, text: str) -> None: if not self._conn: raise ConnectionError("cws ws not connected") diff --git a/hermes_openmax/__init__.py b/hermes_openmax/__init__.py index a3bf9de..4af6c35 100644 --- a/hermes_openmax/__init__.py +++ b/hermes_openmax/__init__.py @@ -2,14 +2,55 @@ from __future__ import annotations +import importlib.util +import logging import os +import sys +from pathlib import Path + + +logger = logging.getLogger(__name__) + + +def _ensure_bundled_sdk_importable() -> None: + """Expose the sibling SDK for trusted directory-plugin installations.""" + try: + installed_spec = importlib.util.find_spec("cws_agent_sdk") + except (ImportError, ValueError): + installed_spec = None + if installed_spec is not None: + return + plugin_root = Path(__file__).resolve().parent.parent + sdk_init = plugin_root / "cws_agent_sdk" / "__init__.py" + if not sdk_init.is_file(): + message = ( + "hermes-openmax failed to load: cws_agent_sdk is not importable. " + "Install with `hermes plugins install openmaxai/hermes-openmax --enable` " + "or `uv pip install` the repository into the Hermes environment." + ) + logger.error(message) + raise ModuleNotFoundError(message) + root = str(plugin_root) + if root not in sys.path: + sys.path.insert(0, root) + importlib.invalidate_caches() + try: + bundled_spec = importlib.util.find_spec("cws_agent_sdk") + except (ImportError, ValueError): + bundled_spec = None + if bundled_spec is None: + message = f"hermes-openmax found but could not import bundled SDK at {sdk_init}" + logger.error(message) + raise ModuleNotFoundError(message) def check_requirements() -> bool: try: + _ensure_bundled_sdk_importable() import httpx # noqa: F401 import websockets # noqa: F401 - except ImportError: + except ImportError as exc: + logger.error("hermes-openmax dependency check failed: %s", exc) return False return True @@ -93,7 +134,7 @@ async def _standalone_send( def register(ctx): """Hermes plugin entry point.""" - from pathlib import Path + _ensure_bundled_sdk_importable() from .adapter import CwsAdapter from .tools import ALL_TOOLS diff --git a/hermes_openmax/adapter.py b/hermes_openmax/adapter.py index 0137fe1..a94a557 100644 --- a/hermes_openmax/adapter.py +++ b/hermes_openmax/adapter.py @@ -14,7 +14,9 @@ from __future__ import annotations +import asyncio import logging +import time from collections import OrderedDict from typing import Any, Dict, Optional @@ -43,6 +45,13 @@ def flag(name: str, default: bool) -> bool: return default return raw in ("1", "true", "yes", "on") + def positive_int(name: str, default: int) -> int: + try: + value = int(os.getenv(name, str(default))) + except ValueError: + return default + return value if value > 0 else default + allow = [ s.strip() for s in os.getenv("CWS_ALLOWED_USERS", "").split(",") if s.strip() ] @@ -57,12 +66,34 @@ def flag(name: str, default: bool) -> bool: dm_policy = "allowlist" else: dm_policy = "owner" + group_policy = os.getenv("CWS_GROUP_POLICY", "").strip().lower() + if group_policy not in ("open", "allowlist", "disabled"): + group_policy = "allowlist" + aliases = [ + value.strip() + for value in os.getenv("CWS_SELF_ALIASES", "").split(",") + if value.strip() + ] + agent_allowlist = [ + value.strip() + for value in os.getenv("CWS_ALLOWED_AGENT_SENDERS", "").split(",") + if value.strip() + ] return AccessPolicyConfig( dm_policy=dm_policy, + group_policy=group_policy, group_require_mention=flag("CWS_GROUP_REQUIRE_MENTION", True), allow_agent_senders=flag("CWS_ALLOW_AGENT_SENDERS", False), allow_sibling_dm=flag("CWS_ALLOW_SIBLING_DM", False), + agent_allowlist=agent_allowlist, + max_agent_hops=positive_int("CWS_MAX_AGENT_HOPS", 4), + agent_turn_budget=positive_int("CWS_AGENT_TURN_BUDGET", 4), + agent_turn_window_s=float(positive_int("CWS_AGENT_TURN_WINDOW_S", 60)), + agent_duplicate_window_s=float( + positive_int("CWS_AGENT_DUPLICATE_WINDOW_S", 60) + ), dm_allowlist=allow, + self_aliases=aliases, ) @@ -104,7 +135,9 @@ def __init__(self, config, **kwargs): self._bridge: Optional[CwsBridge] = None self._orientation: str = "" self._readonly_message_ids: OrderedDict[str, None] = OrderedDict() - self._silent_groups: set[str] = set() + self._agent_causation: "OrderedDict[tuple[str, str], tuple[Dict[str, Any], float]]" = ( + OrderedDict() + ) CwsAdapter._last_instance = self # -- lifecycle ----------------------------------------------------- @@ -126,14 +159,24 @@ async def connect(self, *, is_reconnect: bool = False) -> bool: policy=_policy_from_env(), version=cfg.client_version, on_config_event=self._on_config_event, + # Hermes owns its model/provider credentials. OpenMax hosted-LLM + # organization billing state must never gate external delivery. + billing_gate_enabled=False, ack_reaction="" if ack.lower() in ("off", "false", "none") else ack, ) - await self._bridge.start() - logger.info("[cws] bridge started (org=%s)", cfg.org_id or "") + try: + await self._bridge.start() + except (asyncio.CancelledError, Exception): + bridge = self._bridge + self._bridge = None + try: + await bridge.stop() + except Exception as cleanup_error: # noqa: BLE001 + logger.warning("[cws] failed-connect cleanup error: %s", cleanup_error) + raise + logger.info("[cws] bridge connected (org=%s)", cfg.org_id or "") # Orientation needs several REST calls — build it off the connect path # so slow starts don't trip the gateway's connect timeout. - import asyncio - asyncio.create_task(self._build_orientation()) return True @@ -180,6 +223,26 @@ def last_instance_connected(cls) -> bool: # -- outbound: gateway -> CWS --------------------------------------- + def _with_agent_causation( + self, + chat_id: str, + reply_to: Optional[str], + metadata: Optional[Dict[str, Any]], + ) -> Dict[str, Any]: + result = dict(metadata or {}) + causation = getattr(self, "_agent_causation", {}) + now = time.monotonic() + while causation: + _, (_, created_at) = next(iter(causation.items())) + if len(causation) <= 2048 and now - created_at < 600: + break + causation.popitem(last=False) + entry = causation.get((chat_id, str(reply_to))) if reply_to else None + if entry: + for key, value in entry[0].items(): + result.setdefault(key, value) + return result + async def send( self, chat_id: str, @@ -199,9 +262,7 @@ async def send( if content.strip().upper() == "[SKIP]": logger.info("[cws] reply intentionally skipped for %s", chat_id) return SendResult(success=True, message_id="") - if chat_id in getattr(self, "_silent_groups", set()) or ( - metadata and metadata.get("group_silent") - ): + if metadata and metadata.get("group_silent"): logger.info("[cws] group silent: suppressing reply for %s", chat_id) return SendResult(success=True, message_id="") try: @@ -209,7 +270,7 @@ async def send( conversation_id=chat_id, content=content, reply_to=reply_to, - metadata=metadata, + metadata=self._with_agent_causation(chat_id, reply_to, metadata), ) return SendResult(success=True, message_id=receipt.message_id) except Exception as exc: # noqa: BLE001 — surface any send failure to gateway @@ -258,10 +319,6 @@ async def _on_config_event(self, event: str, data: Dict[str, Any]) -> None: # has already updated owner_member_id, so rebuild the per-turn # workspace context before the next message arrives. await self._build_orientation() - elif event == "agent.config.group_mode_changed": - conv = str(data.get("conversation_id") or "") - if conv and str(data.get("mode", "")).lower() != "silent": - self._silent_groups.discard(conv) logger.info( "[cws] config event %s: %s", event, {k: data.get(k) for k in list(data)[:6]} ) @@ -271,6 +328,23 @@ async def _on_config_event(self, event: str, data: Dict[str, Any]) -> None: async def _on_inbound(self, msg: InboundMessage) -> None: """SDK delivery callback. Raising here prevents the ack watermark from advancing, so the message is replayed via /sync later.""" + if msg.sender_type == "agent": + if not hasattr(self, "_agent_causation"): + self._agent_causation = OrderedDict() + causation = { + key: msg.metadata[key] + for key in ( + "agent_hop_count", + "agent_origin_member_id", + "agent_trace_id", + ) + if msg.metadata.get(key) is not None + } + key = (msg.conversation_id, msg.message_id) + self._agent_causation[key] = (causation, time.monotonic()) + self._agent_causation.move_to_end(key) + while len(self._agent_causation) > 2048: + self._agent_causation.popitem(last=False) if msg.sender_type == "system": previous = getattr(self, "_readonly_message_ids", ()) if not isinstance(previous, OrderedDict): @@ -281,8 +355,6 @@ async def _on_inbound(self, msg: InboundMessage) -> None: self._readonly_message_ids.move_to_end(msg.message_id) while len(self._readonly_message_ids) > 1024: self._readonly_message_ids.popitem(last=False) - if msg.metadata.get("group_silent"): - self._silent_groups.add(msg.conversation_id) # OpenMax groups are conversation-scoped. Do not pass the sender as the # session participant, otherwise Hermes creates one session per member # instead of one shared session per group. @@ -351,7 +423,11 @@ async def send_image_file( return SendResult(success=True, message_id="") try: receipt = await self._bridge.send_image_file( - chat_id, image_path, caption=caption or "", reply_to=reply_to + chat_id, + image_path, + caption=caption or "", + reply_to=reply_to, + metadata=self._with_agent_causation(chat_id, reply_to, metadata), ) return SendResult(success=True, message_id=receipt.message_id) except Exception as exc: # noqa: BLE001 — fall back to caption-only text @@ -361,6 +437,7 @@ async def send_image_file( chat_id, f"{caption or ''}\n⚠️ 图片发送失败".strip(), reply_to=reply_to, + metadata=self._with_agent_causation(chat_id, reply_to, metadata), ) return SendResult(success=True, message_id=receipt.message_id) except Exception as exc2: # noqa: BLE001 @@ -406,18 +483,33 @@ async def send_image( ) text = caption or f"[image] {fname}" receipt = await self._bridge.send( - chat_id, text, reply_to=reply_to, metadata={"attachment": node} + chat_id, + text, + reply_to=reply_to, + metadata=self._with_agent_causation( + chat_id, + reply_to, + {**(metadata or {}), "attachment": node}, + ), ) return SendResult(success=True, message_id=receipt.message_id) # Remote URL: no re-hosting — send as markdown image link. text = f"![{caption or 'image'}]({image_url})" - receipt = await self._bridge.send(chat_id, text, reply_to=reply_to) + receipt = await self._bridge.send( + chat_id, + text, + reply_to=reply_to, + metadata=self._with_agent_causation(chat_id, reply_to, metadata), + ) return SendResult(success=True, message_id=receipt.message_id) except Exception as exc: # noqa: BLE001 logger.warning("[cws] send_image failed, falling back to text: %s", exc) try: receipt = await self._bridge.send( - chat_id, f"{caption or ''} {image_url}".strip(), reply_to=reply_to + chat_id, + f"{caption or ''} {image_url}".strip(), + reply_to=reply_to, + metadata=self._with_agent_causation(chat_id, reply_to, metadata), ) return SendResult(success=True, message_id=receipt.message_id) except Exception as exc2: # noqa: BLE001 diff --git a/hermes_openmax/plugin.yaml b/hermes_openmax/plugin.yaml index 7809f7b..4c2858d 100644 --- a/hermes_openmax/plugin.yaml +++ b/hermes_openmax/plugin.yaml @@ -31,7 +31,7 @@ optional_env: prompt: "Allowed member ids (optional)" password: false - name: CWS_ALLOW_ALL_USERS - description: "Allow all workspace members (default: true — CWS already scopes access by org)" + description: "Allow all workspace members to DM the agent (default: false)" prompt: "Allow all users? (true/false)" password: false - name: CWS_HOME_CHANNEL @@ -50,10 +50,38 @@ optional_env: description: "Allow DMs from sibling agents (default: false, prevents agent chat loops)" prompt: "Allow sibling agent DMs? (true/false)" password: false + - name: CWS_ALLOWED_AGENT_SENDERS + description: "Comma-separated Agent member ids allowed by the Agent loop guard; '*' allows any Agent after the remaining gates (default: empty/fail-closed)" + prompt: "Allowed Agent member ids (optional)" + password: false + - name: CWS_MAX_AGENT_HOPS + description: "Reject Agent messages whose propagated hop count exceeds this value (default: 4)" + prompt: "Maximum Agent hops" + password: false + - name: CWS_AGENT_TURN_BUDGET + description: "Maximum admitted Agent messages per sender/conversation window (default: 4)" + prompt: "Agent turn budget" + password: false + - name: CWS_AGENT_TURN_WINDOW_S + description: "Agent turn-budget window in seconds (default: 60)" + prompt: "Agent turn window seconds" + password: false + - name: CWS_AGENT_DUPLICATE_WINDOW_S + description: "Reject repeated Agent content within this many seconds (default: 60)" + prompt: "Agent duplicate window seconds" + password: false - name: CWS_DM_POLICY - description: "DM admission: open | allowlist | owner (default open; platform can change live via agent.config)" + description: "Human DM admission: open | allowlist | owner (default: owner; platform can change live via agent.config)" prompt: "DM policy (open/allowlist/owner)" password: false + - name: CWS_GROUP_POLICY + description: "Human group admission: open | allowlist | disabled (default: allowlist; platform can change live via agent.config)" + prompt: "Group policy (open/allowlist/disabled)" + password: false + - name: CWS_SELF_ALIASES + description: "Comma-separated stable aliases accepted for plain-text @mentions (optional)" + prompt: "Agent mention aliases (optional)" + password: false - name: CWS_PERSONA description: "Workspace-specific persona text, appended to the per-turn orientation (base personality stays in SOUL.md)" prompt: "Workspace persona (optional)" diff --git a/hermes_openmax/skills/workspace.md b/hermes_openmax/skills/workspace.md index 4676de7..2f3ff2e 100644 --- a/hermes_openmax/skills/workspace.md +++ b/hermes_openmax/skills/workspace.md @@ -189,8 +189,9 @@ API 查询(`list_projects`/`workspace_members list`)→ 默认 Inbox → 问人 派活后绝大多数协调走 **bot 对 bot DM**(不是发给人类): 1. 查 worker 的 member_id(`workspace_members list kind=agent search=<名>`;常用的记进记忆); -2. **确认双向 DM 权限已开**(对方能 DM 你、你也能 DM 对方 —— dm 策略由平台经 - agent.config 事件管理;拿不准先发条测试 DM,长时间无响应就报告人类,不要干等); +2. **确认双向 Agent DM 权限已开**(双方 runtime 都需启用 sibling DM、把对方 + member_id 放进 Agent sender allowlist,并由 Core 确认为 same-owner;拿不准先发条 + 测试 DM,长时间无响应就报告人类,不要干等); 3. `workspace_members create_dm(peer_member_id)` 拿 conversationId(幂等,存记忆复用); 4. 发目标:markdown 写清 **目标、所属 Issue ID、KB 产出位置、返回触发词、判定标准**; **Task 由被派的 worker 自己在该 Issue 下创建并认领**(谁执行谁建); @@ -203,12 +204,16 @@ API 查询(`list_projects`/`workspace_members list`)→ 默认 Inbox → 问人 - **DM 策略**(平台经 agent.config 热更):`open`(org 内任何成员)/ `allowlist`(白名单+owner)/ `owner`(仅 owner)。**owner 永远豁免**。 被策略拒绝的人类 DM 会收到一条礼貌提示(每会话每小时至多一条)。 -- **群策略**:默认需 @ 你才处理;owner 被 @ 永远放行;群可被平台设为 +- **群策略**:默认 scope=`allowlist` 且需 @ 你才处理;`disabled` 对 owner 也生效; + allowlist 中未登记的群仅允许 owner 明确 @ 时绕过,已登记群的 owner 不受 + `allowFrom` 限制。群可被平台设为 `mention`(默认)/ `smart`(全收,自行判断,不值得回就输出 [SKIP])/ - `silent`(整群静默丢弃)。 + `silent`(bridge 仅缓存有界文本上下文,不创建 Agent turn、不调用模型、不回复)。 - **System Member(调度器)不受任何策略约束**,直接进你的会话。 -- 其他 agent 的 DM 默认不触发你(防 agent 互聊死循环);群里 agent 消息 - 仅当平台开启且 @ 你时触发。 +- 其他 agent 的 DM 默认不触发你(防 agent 互聊死循环);开启后仍要求 same-owner + 且 sender 在 Agent allowlist。群里 agent 消息需显式开启、sender 在 Agent + allowlist、结构化 @ 你,并继续通过 group scope/群登记/`allowFrom`。bridge 还会 + 执行 hop 上限、重复消息和单位时间 turn budget 熔断。 ## 记忆触发点 diff --git a/hermes_openmax/tools.py b/hermes_openmax/tools.py index d5a2e93..306def6 100644 --- a/hermes_openmax/tools.py +++ b/hermes_openmax/tools.py @@ -96,13 +96,29 @@ async def go(): storage = FileStorage(_STATE_DIR) tokens = TokenManager(cfg, storage=storage) http = CwsHttpClient(cfg, tokens) + from .adapter import CwsAdapter + + instance = CwsAdapter._last_instance + live_bridge = ( + instance._bridge + if instance and instance._bridge and instance._bridge._running + else None + ) try: svc = { "tm": TmService(http), "kb": KbService(http), "core": CoreService(http), "comm": CommService(http), - "policy": AccessPolicyService(storage), + "policy": AccessPolicyService( + storage, + get_live_state=(live_bridge.get_dm_access if live_bridge else None), + apply_live_state=( + live_bridge.apply_local_dm_access_threadsafe + if live_bridge + else None + ), + ), } return await coro_factory(svc) finally: diff --git a/tests/integration/test_group_pipeline.py b/tests/integration/test_group_pipeline.py index d259446..1e4c510 100644 --- a/tests/integration/test_group_pipeline.py +++ b/tests/integration/test_group_pipeline.py @@ -6,7 +6,7 @@ from cws_agent_sdk.access_policy import AccessPolicyConfig from cws_agent_sdk.bridge import CwsBridge -from cws_agent_sdk.codec import FRAME_MESSAGE, Frame +from cws_agent_sdk.codec import FRAME_MESSAGE, FRAME_SYNC, Frame from cws_agent_sdk.config import CwsConfig from cws_agent_sdk.providers import FileStorage from cws_agent_sdk.types import InboundMessage @@ -128,6 +128,24 @@ async def capture_event(event): assert event.source.user_id is None assert event.text == "@agent hello" assert comm.read_marks == [(conversation_id, 7)] + assert comm.sync_acks == [] + + await bridge._handle_frame( + Frame( + type=FRAME_SYNC, + org_id="org-1", + payload={ + "events": [ + { + "conversation_id": conversation_id, + "message_id": "1001", + "seq": 501, + } + ] + }, + ) + ) + assert comm.sync_acks == [501] diff --git a/tests/test_access_policy.py b/tests/test_access_policy.py index ff6186c..062f8e4 100644 --- a/tests/test_access_policy.py +++ b/tests/test_access_policy.py @@ -4,12 +4,18 @@ ME = "me-1" -def msg(sender_type="human", sender_id="u-7", conv_type="dm", mentions=None): +def msg( + sender_type="human", + sender_id="u-7", + conv_type="dm", + mentions=None, + text="hi", +): return InboundMessage( message_id="1", conversation_id="c-1", org_id="o-1", - text="hi", + text=text, sender_id=sender_id, sender_type=sender_type, conversation_type=conv_type, @@ -17,8 +23,17 @@ def msg(sender_type="human", sender_id="u-7", conv_type="dm", mentions=None): ) +def test_safe_defaults_match_zylos_human_policy(): + cfg = AccessPolicyConfig() + assert cfg.dm_policy == "owner" + assert cfg.group_policy == "allowlist" + assert cfg.allow_sibling_dm is False + assert cfg.allow_agent_senders is False + assert cfg.agent_allowlist == [] + + def test_human_dm_handled(): - d = decide_inbound(msg(), self_member_id=ME) + d = decide_inbound(msg(), self_member_id=ME, cfg=AccessPolicyConfig(dm_policy="open")) assert d.handle and d.reason == "dm" @@ -44,7 +59,7 @@ def test_dm_owner_policy_and_exemption(): def test_group_owner_mention_bypass(): - cfg = AccessPolicyConfig(group_require_mention=True) + cfg = AccessPolicyConfig(group_policy="allowlist", group_require_mention=True) m = msg( sender_id="boss-1", conv_type="group", @@ -55,28 +70,33 @@ def test_group_owner_mention_bypass(): def test_group_requires_mention_by_default(): - d = decide_inbound(msg(conv_type="group"), self_member_id=ME) + cfg = AccessPolicyConfig(group_policy="open") + d = decide_inbound(msg(conv_type="group"), self_member_id=ME, cfg=cfg) assert not d.handle and d.reason == "group_no_mention" def test_group_mention_by_member_id(): + cfg = AccessPolicyConfig(group_policy="open") d = decide_inbound( msg(conv_type="group", mentions=[{"type": "member", "member_id": ME}]), self_member_id=ME, + cfg=cfg, ) assert d.handle and d.reason == "group_mention" def test_group_mention_all_agents(): + cfg = AccessPolicyConfig(group_policy="open") d = decide_inbound( msg(conv_type="group", mentions=[{"type": "all_agents"}]), self_member_id=ME, + cfg=cfg, ) assert d.handle def test_group_open_mode(): - cfg = AccessPolicyConfig(group_require_mention=False) + cfg = AccessPolicyConfig(group_policy="open", group_require_mention=False) assert decide_inbound(msg(conv_type="group"), self_member_id=ME, cfg=cfg).handle @@ -92,10 +112,82 @@ def test_agent_group_blocked_even_with_mention_unless_allowed(): mentions=[{"type": "member", "member_id": ME}], ) assert not decide_inbound(m, self_member_id=ME).handle - cfg = AccessPolicyConfig(allow_agent_senders=True) + cfg = AccessPolicyConfig( + group_policy="open", + allow_agent_senders=True, + agent_allowlist=["u-7"], + ) assert decide_inbound(m, self_member_id=ME, cfg=cfg).handle +def test_agent_group_still_obeys_group_scope_and_allow_from(): + message = msg( + sender_type="agent", + sender_id="agent-1", + conv_type="group", + mentions=[{"type": "member", "member_id": ME}], + ) + base = dict(allow_agent_senders=True, agent_allowlist=["agent-1"]) + assert not decide_inbound( + message, + self_member_id=ME, + cfg=AccessPolicyConfig(group_policy="disabled", **base), + ).handle + assert not decide_inbound( + message, + self_member_id=ME, + cfg=AccessPolicyConfig(group_policy="allowlist", **base), + ).handle + allowed_group = {"c-1": {"mode": "mention", "allow_from": ["someone-else"]}} + assert not decide_inbound( + message, + self_member_id=ME, + cfg=AccessPolicyConfig( + group_policy="allowlist", group_configs=allowed_group, **base + ), + ).handle + + +def test_agent_plain_text_mention_does_not_bypass_structured_loop_guard(): + cfg = AccessPolicyConfig( + group_policy="open", + allow_agent_senders=True, + agent_allowlist=["agent-1"], + self_display_name="COCO", + ) + decision = decide_inbound( + msg( + sender_type="agent", + sender_id="agent-1", + conv_type="group", + text="@COCO hello", + ), + self_member_id=ME, + cfg=cfg, + ) + assert not decision.handle and decision.reason == "agent_sender_blocked" + + +def test_agent_broadcast_mentions_do_not_bypass_direct_mention_loop_guard(): + cfg = AccessPolicyConfig( + group_policy="open", + allow_agent_senders=True, + agent_allowlist=["agent-1"], + ) + for mention_type in ("all", "all_agents"): + decision = decide_inbound( + msg( + sender_type="agent", + sender_id="agent-1", + conv_type="group", + mentions=[{"type": mention_type}], + ), + self_member_id=ME, + cfg=cfg, + ) + assert not decision.handle and decision.reason == "agent_sender_blocked" + + def test_system_sender_delivered_by_default(): # Scheduler DMs drive the task flow (dependency-ready, issue.activated). assert decide_inbound(msg(sender_type="system"), self_member_id=ME).handle diff --git a/tests/test_adapter.py b/tests/test_adapter.py index ce3e8fd..b8bf0e2 100644 --- a/tests/test_adapter.py +++ b/tests/test_adapter.py @@ -1,4 +1,10 @@ -"""Hermes OpenMax prompt and media behavior regressions.""" +"""Hermes OpenMax prompt, install, and media behavior regressions.""" + +import shutil +import subprocess +import sys +import textwrap +from pathlib import Path from hermes_openmax.behavior import ( build_workspace_orientation, @@ -7,6 +13,53 @@ ) +def test_directory_plugin_bootstraps_sibling_sdk_in_isolated_python(tmp_path): + source_root = Path(__file__).resolve().parents[1] + plugin_root = tmp_path / "checkout" + shutil.copytree(source_root / "hermes_openmax", plugin_root / "hermes_openmax") + shutil.copytree(source_root / "cws_agent_sdk", plugin_root / "cws_agent_sdk") + script = textwrap.dedent( + f""" + import importlib.util + import sys + from pathlib import Path + + plugin_dir = Path({str(plugin_root / 'hermes_openmax')!r}) + real_find_spec = importlib.util.find_spec + + def directory_install_find_spec(name, *args, **kwargs): + if name == "cws_agent_sdk" and str(plugin_dir.parent) not in sys.path: + return None + return real_find_spec(name, *args, **kwargs) + + importlib.util.find_spec = directory_install_find_spec + module_name = "hermes_plugins.hermes_openmax_directory_test" + spec = importlib.util.spec_from_file_location( + module_name, + plugin_dir / "__init__.py", + submodule_search_locations=[str(plugin_dir)], + ) + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + spec.loader.exec_module(module) + module._ensure_bundled_sdk_importable() + import cws_agent_sdk + assert Path(cws_agent_sdk.__file__).resolve().parent == ( + plugin_dir.parent / "cws_agent_sdk" + ).resolve() + """ + ) + + result = subprocess.run( + [sys.executable, "-I", "-c", script], + text=True, + capture_output=True, + check=False, + ) + + assert result.returncode == 0, result.stderr + + def test_extract_local_markdown_images_accepts_existing_file_uri_with_caption( tmp_path, monkeypatch ): diff --git a/tests/test_adapter_runtime_parity.py b/tests/test_adapter_runtime_parity.py index 9d7484e..d2fcd9a 100644 --- a/tests/test_adapter_runtime_parity.py +++ b/tests/test_adapter_runtime_parity.py @@ -7,7 +7,84 @@ from gateway.config import Platform from cws_agent_sdk.types import InboundMessage -from hermes_openmax.adapter import CwsAdapter + +import hermes_openmax.adapter as adapter_module +from hermes_openmax.adapter import CwsAdapter, _policy_from_env + + +def test_policy_env_defaults_are_safe_and_agent_controls_are_explicit(monkeypatch): + for name in ( + "CWS_DM_POLICY", + "CWS_GROUP_POLICY", + "CWS_ALLOW_ALL_USERS", + "CWS_ALLOWED_USERS", + "CWS_ALLOW_AGENT_SENDERS", + "CWS_ALLOW_SIBLING_DM", + "CWS_ALLOWED_AGENT_SENDERS", + "CWS_SELF_ALIASES", + ): + monkeypatch.delenv(name, raising=False) + + policy = _policy_from_env() + + assert policy.dm_policy == "owner" + assert policy.group_policy == "allowlist" + assert policy.allow_agent_senders is False + assert policy.allow_sibling_dm is False + assert policy.agent_allowlist == [] + + +def test_policy_env_loads_agent_allowlist_aliases_and_budgets(monkeypatch): + monkeypatch.setenv("CWS_ALLOWED_AGENT_SENDERS", "agent-1, agent-2") + monkeypatch.setenv("CWS_SELF_ALIASES", "COCO, helper.bot") + monkeypatch.setenv("CWS_MAX_AGENT_HOPS", "3") + monkeypatch.setenv("CWS_AGENT_TURN_BUDGET", "2") + + policy = _policy_from_env() + + assert policy.agent_allowlist == ["agent-1", "agent-2"] + assert policy.self_aliases == ["COCO", "helper.bot"] + assert policy.max_agent_hops == 3 + assert policy.agent_turn_budget == 2 + + +@pytest.mark.asyncio +async def test_adapter_connect_explicitly_disables_hosted_billing_gate(monkeypatch): + captured = {} + + class Config: + org_id = "org-1" + client_version = "test" + + def validate(self): + return [] + + class Bridge: + def __init__(self, cfg, **kwargs): + captured.update(kwargs) + self._cfg = SimpleNamespace(member_id="agent-1") + self.owner_member_id = "" + self.core = SimpleNamespace() + + async def start(self): + pass + + async def stop(self): + pass + + monkeypatch.setattr( + adapter_module.CwsConfig, "from_env", classmethod(lambda _cls: Config()) + ) + monkeypatch.setattr(adapter_module, "CwsBridge", Bridge) + adapter = CwsAdapter(SimpleNamespace()) + + async def no_orientation(): + pass + + adapter._build_orientation = no_orientation + + assert await adapter.connect() is True + assert captured["billing_gate_enabled"] is False class _Bridge: @@ -70,6 +147,63 @@ async def test_system_member_inbound_marks_only_that_message_read_only_for_repli assert result.success and result.message_id == "out-1" +@pytest.mark.asyncio +async def test_agent_inbound_causation_is_forwarded_to_outbound_reply(): + adapter = _adapter() + first = InboundMessage( + message_id="agent-msg-1", + conversation_id="agent-dm-1", + org_id="org-1", + text="Please continue", + sender_id="agent-2", + sender_type="agent", + conversation_type="dm", + metadata={ + "agent_hop_count": 2, + "agent_origin_member_id": "agent-1", + "agent_trace_id": "trace-1", + }, + ) + second = InboundMessage( + message_id="agent-msg-2", + conversation_id="agent-dm-1", + org_id="org-1", + text="Another request", + sender_id="agent-2", + sender_type="agent", + conversation_type="dm", + metadata={"agent_hop_count": 7, "agent_trace_id": "trace-2"}, + ) + human = InboundMessage( + message_id="human-msg-1", + conversation_id="agent-dm-1", + org_id="org-1", + text="Human interjection", + sender_id="human-1", + sender_type="human", + conversation_type="dm", + ) + + await adapter._on_inbound(first) + await adapter._on_inbound(second) + await adapter._on_inbound(human) + await adapter.send("agent-dm-1", "Done", reply_to="agent-msg-1") + assert adapter._bridge.sent[-1]["metadata"] == { + "agent_hop_count": 2, + "agent_origin_member_id": "agent-1", + "agent_trace_id": "trace-1", + } + + await adapter.send("agent-dm-1", "Second", reply_to="agent-msg-2") + assert adapter._bridge.sent[-1]["metadata"] == { + "agent_hop_count": 7, + "agent_trace_id": "trace-2", + } + + await adapter.send("agent-dm-1", "Proactive") + assert adapter._bridge.sent[-1]["metadata"] == {} + + @pytest.mark.asyncio async def test_system_member_conversation_suppresses_media_and_edits(): adapter = _adapter() diff --git a/tests/test_bridge.py b/tests/test_bridge.py index 9e1a973..674088c 100644 --- a/tests/test_bridge.py +++ b/tests/test_bridge.py @@ -3,9 +3,10 @@ import pytest from cws_agent_sdk.bridge import CwsBridge -from cws_agent_sdk.codec import FRAME_MESSAGE, Frame +from cws_agent_sdk.codec import FRAME_MESSAGE, FRAME_SYNC, Frame from cws_agent_sdk.config import CwsConfig from cws_agent_sdk.providers import FileStorage +from cws_agent_sdk.types import InboundMessage class FakeComm: @@ -13,6 +14,7 @@ def __init__(self): self.messages = {} self.read_marks = [] self.sync_acks = [] + self.sent_messages = [] async def get_message(self, conv_id, msg_id): return self.messages[f"{conv_id}:{msg_id}"] @@ -41,6 +43,7 @@ async def remove_reaction(self, message_id, code): async def send_message(self, conv_id, text, **kw): from cws_agent_sdk.types import SendReceipt + self.sent_messages.append((conv_id, text, kw)) return SendReceipt(message_id="out-1", conversation_id=conv_id) @@ -75,16 +78,31 @@ def msg_frame(msg_id=1, conv="conv-1", seq=10, sender="user-7"): ) -def detail(msg_id=1, conv="conv-1", seq=10, sender="user-7", text="hello"): +def detail( + msg_id=1, + conv="conv-1", + seq=10, + sender="user-7", + text="hello", + sender_type="HUMAN", + mentions=None, + metadata=None, + include_inbox_seq=True, +): + message = { + "id": msg_id, + "conversation_id": conv, + "seq": seq, + "sender_id": sender, + "sender_type": sender_type, + "client_msg_id": f"cm-{msg_id}", + "mentions": mentions or [], + "metadata": metadata or {}, + } + if include_inbox_seq: + message["inbox_seq"] = seq return { - "message": { - "id": msg_id, - "conversation_id": conv, - "seq": seq, - "sender_id": sender, - "sender_type": "HUMAN", - "client_msg_id": f"cm-{msg_id}", - }, + "message": message, "content": {"content_type": "text", "body": {"text": text}}, } @@ -124,9 +142,35 @@ async def on_message(m): assert got[0].text == "hello" assert got[0].sender_type == "human" assert b.comm.read_marks == [("conv-1", 10)] + assert b.comm.sync_acks == [] + + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={"events": [{"conversation_id": "conv-1", "message_id": 1, "seq": 10}]}, + ) + ) assert b.comm.sync_acks == [10] +@pytest.mark.asyncio +async def test_realtime_without_inbox_seq_does_not_advance_global_cursor(tmp_path): + got = [] + + async def on_message(message): + got.append(message) + + b = make_bridge(tmp_path, on_message) + b.comm.messages["conv-1:1"] = detail(seq=7, include_inbox_seq=False) + + await b._handle_frame(msg_frame(seq=7)) + + assert len(got) == 1 + assert b.comm.read_marks == [("conv-1", 7)] + assert b.comm.sync_acks == [] + assert b._sync_seq == 0 + + @pytest.mark.asyncio async def test_sync_replay_acks_org_inbox_seq_not_conversation_seq(tmp_path): """The global /sync cursor must never be advanced with a per-chat seq.""" @@ -164,6 +208,159 @@ async def on_message(m): assert "conv-1:1" not in b._seen +@pytest.mark.asyncio +async def test_external_bridge_default_does_not_gate_on_hosted_llm_billing(tmp_path): + got = [] + + async def on_message(message): + got.append(message) + + cfg = CwsConfig( + bff_url="https://bff.test", + ws_url="wss://comm.test", + api_key="cwsk_x", + org_id="org-1", + member_id="me-1", + ) + b = CwsBridge(cfg, storage=FileStorage(tmp_path), on_message=on_message) + b.comm = FakeComm() + b.comm.messages["conv-1:1"] = detail() + + assert b._billing is None + await b._handle_frame(msg_frame()) + + assert [message.message_id for message in got] == ["1"] + assert "conv-1:1" in b._seen + + +@pytest.mark.asyncio +async def test_sync_replay_never_emits_delayed_billing_notice(tmp_path): + got = [] + + async def on_message(message): + got.append(message) + + class SuspendedBilling: + async def is_suspended(self): + return True + + def should_send_overdue_notice(self, _conversation_id): + return True + + b = make_bridge(tmp_path, on_message) + b._billing = SuspendedBilling() + b.comm.messages["conv-1:1"] = detail() + + await b._deliver_by_id("conv-1", 1, 10) + + assert got == [] + assert b.comm.sent_messages == [] + assert "conv-1:1" in b._seen + assert b.comm.sync_acks == [10] + + +@pytest.mark.asyncio +async def test_sync_replay_never_emits_policy_rejection_notice(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + b._policy.dm_policy = "allowlist" + b.comm.messages["conv-1:1"] = detail() + + await b._deliver_by_id("conv-1", 1, 10) + + assert b.comm.sent_messages == [] + assert "conv-1:1" in b._seen + assert b.comm.sync_acks == [10] + + +@pytest.mark.asyncio +async def test_realtime_success_cannot_commit_past_lower_failed_inbox_seq(tmp_path): + import asyncio + + low_started = asyncio.Event() + release_low = asyncio.Event() + attempts = [] + + async def on_message(message): + attempts.append(message.message_id) + if message.message_id == "1" and attempts.count("1") == 1: + low_started.set() + await release_low.wait() + raise RuntimeError("low seq failed") + + b = make_bridge(tmp_path, on_message) + b.comm.messages["conv-1:1"] = detail(msg_id=1, seq=1) + b.comm.messages["conv-1:2"] = detail(msg_id=2, seq=2, text="higher") + + low = asyncio.create_task(b._handle_frame(msg_frame(msg_id=1, seq=1))) + await low_started.wait() + await b._handle_frame(msg_frame(msg_id=2, seq=2)) + release_low.set() + result = await asyncio.gather(low, return_exceptions=True) + + assert isinstance(result[0], RuntimeError) + assert b._sync_seq == 0 + assert b.comm.sync_acks == [] + assert "conv-1:1" not in b._seen + assert "conv-1:2" in b._seen + + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={ + "events": [ + {"conversation_id": "conv-1", "message_id": 1, "seq": 10}, + {"conversation_id": "conv-1", "message_id": 2, "seq": 11}, + ] + }, + ) + ) + + assert attempts == ["1", "2", "1"] + assert b._sync_seq == 11 + assert b.comm.sync_acks == [10, 11] + + +@pytest.mark.parametrize("policy_case", ["silent", "disabled", "unregistered"]) +@pytest.mark.asyncio +async def test_unknown_conversation_type_fails_closed_and_remains_retryable( + tmp_path, policy_case +): + got = [] + + async def on_message(message): + got.append(message) + + b = make_bridge(tmp_path, on_message) + if policy_case == "silent": + b._policy.group_configs["conv-1"] = { + "mode": "silent", + "allow_from": ["*"], + } + b._group_mode_overrides["conv-1"] = "silent" + elif policy_case == "disabled": + b._policy.group_policy = "disabled" + else: + b._policy.group_policy = "allowlist" + + async def fail_conversation_lookup(_conversation_id): + raise RuntimeError("conversation service unavailable") + + b.comm.get_conversation = fail_conversation_lookup + b.comm.messages["conv-1:1"] = detail( + mentions=[{"type": "member", "member_id": "me-1"}] + ) + + with pytest.raises(RuntimeError, match="conversation service unavailable"): + await b._handle_frame(msg_frame()) + + assert got == [] + assert b.owner_member_id == "" + assert "conv-1:1" not in b._seen + assert "conv-1" not in b._conv_types + assert b.comm.read_marks == [] + assert b.comm.sync_acks == [] + + @pytest.mark.asyncio async def test_own_echo_suppressed_by_sender(tmp_path): got = [] @@ -224,7 +421,7 @@ async def slow_on_message(m): b._handle_frame(msg_frame()), ) assert len(got) == 1 - assert b.comm.sync_acks == [10] + assert b.comm.sync_acks == [] @pytest.mark.asyncio @@ -255,6 +452,10 @@ async def on_message(m): b = make_bridge(tmp_path, on_message) b.comm.conv_type = "group" b._group_mode_overrides["conv-1"] = "smart" + b._policy.group_configs["conv-1"] = { + "mode": "smart", + "allow_from": ["*"], + } b.comm.messages["conv-1:1"] = detail(msg_id=1, seq=10, text="earlier chatter") b.comm.messages["conv-1:2"] = detail(msg_id=2, seq=11, text="hello smart") await b._handle_frame(msg_frame(msg_id=1, seq=10)) @@ -266,7 +467,7 @@ async def on_message(m): @pytest.mark.asyncio -async def test_group_silent_mode_observes_without_reply(tmp_path): +async def test_group_silent_mode_caches_without_model_delivery(tmp_path): got = [] async def on_message(m): @@ -275,13 +476,255 @@ async def on_message(m): b = make_bridge(tmp_path, on_message) b.comm.conv_type = "group" b._group_mode_overrides["conv-1"] = "silent" + b._policy.group_configs["conv-1"] = { + "mode": "silent", + "allow_from": ["*"], + } + + async def unexpected_member_lookup(_member_id): + raise AssertionError("silent must not resolve sender identity") + + b.core.get_member = unexpected_member_lookup b.comm.messages["conv-1:1"] = detail() await b._handle_frame(msg_frame()) - assert len(got) == 1 - assert got[0].metadata["group_silent"] is True + assert got == [] + assert b._group_history["conv-1"] == ["user-7: hello"] + assert getattr(b.comm, "reactions_added", []) == [] + assert b.comm.sync_acks == [] + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={"events": [{"conversation_id": "conv-1", "message_id": 1, "seq": 10}]}, + ) + ) assert b.comm.sync_acks == [10] # consumed silently +@pytest.mark.asyncio +async def test_agent_duplicate_and_turn_budget_breakers_consume_without_delivery(tmp_path): + got = [] + + async def on_message(message): + got.append(message) + + b = make_bridge(tmp_path, on_message) + b.comm.conv_type = "group" + b._policy.group_policy = "open" + b._policy.allow_agent_senders = True + b._policy.agent_allowlist = ["agent-1"] + b._policy.agent_turn_budget = 2 + mention = [{"type": "member", "member_id": "me-1"}] + b.comm.messages["conv-1:1"] = detail( + msg_id=1, + sender="agent-1", + sender_type="AGENT", + mentions=mention, + text="same task", + ) + b.comm.messages["conv-1:2"] = detail( + msg_id=2, + seq=11, + sender="agent-1", + sender_type="AGENT", + mentions=mention, + text="same task", + ) + b.comm.messages["conv-1:3"] = detail( + msg_id=3, + seq=12, + sender="agent-1", + sender_type="AGENT", + mentions=mention, + text="different task", + ) + b.comm.messages["conv-1:4"] = detail( + msg_id=4, + seq=13, + sender="agent-1", + sender_type="AGENT", + mentions=mention, + text="third task", + ) + + await b._handle_frame(msg_frame(msg_id=1, sender="agent-1")) + await b._handle_frame(msg_frame(msg_id=2, seq=11, sender="agent-1")) + await b._handle_frame(msg_frame(msg_id=3, seq=12, sender="agent-1")) + await b._handle_frame(msg_frame(msg_id=4, seq=13, sender="agent-1")) + + assert [message.text for message in got] == ["same task", "different task"] + assert b.comm.sync_acks == [] + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={ + "events": [ + {"conversation_id": "conv-1", "message_id": i, "seq": seq} + for i, seq in enumerate(range(10, 14), start=1) + ] + }, + ) + ) + assert b.comm.sync_acks == [10, 11, 12, 13] + + +@pytest.mark.asyncio +async def test_agent_hop_limit_is_fail_closed(tmp_path): + got = [] + + async def on_message(message): + got.append(message) + + b = make_bridge(tmp_path, on_message) + b.comm.conv_type = "group" + b._policy.group_policy = "open" + b._policy.allow_agent_senders = True + b._policy.agent_allowlist = ["agent-1"] + b._policy.max_agent_hops = 2 + b.comm.messages["conv-1:1"] = detail( + sender="agent-1", + sender_type="AGENT", + mentions=[{"type": "member", "member_id": "me-1"}], + metadata={"agent_hop_count": 3}, + ) + + await b._handle_frame(msg_frame(sender="agent-1")) + + assert got == [] + assert b.comm.sync_acks == [] + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={"events": [{"conversation_id": "conv-1", "message_id": 1, "seq": 10}]}, + ) + ) + assert b.comm.sync_acks == [10] + + +@pytest.mark.parametrize( + ("raw_hop", "expected"), + [ + (0, "agent_hop_limit"), + (-1, "agent_hop_limit"), + (True, "agent_hop_invalid"), + (1.5, "agent_hop_invalid"), + (None, "agent_hop_invalid"), + ("", "agent_hop_invalid"), + ("bad", "agent_hop_invalid"), + ], +) +def test_agent_hop_metadata_is_strictly_validated(tmp_path, raw_hop, expected): + b = make_bridge(tmp_path, lambda _message: None) + b._policy.max_agent_hops = 4 + message = InboundMessage( + message_id=f"m-{raw_hop!r}", + conversation_id="conv-1", + org_id="org-1", + text="hello", + sender_id="agent-1", + sender_type="agent", + metadata={"agent_hop_count": raw_hop}, + ) + + assert b._agent_loop_rejection(message) == expected + + +def test_missing_agent_hop_metadata_defaults_to_one(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + message = InboundMessage( + message_id="m-default-hop", + conversation_id="conv-1", + org_id="org-1", + text="hello", + sender_id="agent-1", + sender_type="agent", + ) + + assert b._agent_loop_rejection(message) == "" + + +def test_agent_turn_budget_state_expires_and_is_capacity_bounded(tmp_path, monkeypatch): + from types import SimpleNamespace + + import cws_agent_sdk.bridge as bridge_module + + clock = [0.0] + monkeypatch.setattr( + bridge_module, "time", SimpleNamespace(monotonic=lambda: clock[0]) + ) + b = make_bridge(tmp_path, lambda _message: None) + b._policy.agent_turn_window_s = 1 + b._policy.agent_turn_budget = 10 + + def message(index): + return InboundMessage( + message_id=f"m-{index}", + conversation_id=f"conv-{index}", + org_id="org-1", + text=f"task {index}", + sender_id=f"agent-{index}", + sender_type="agent", + ) + + for index in range(10): + assert b._agent_loop_rejection(message(index)) == "" + assert len(b._agent_turns) == 10 + + clock[0] = 10.0 + assert b._agent_loop_rejection(message(10)) == "" + assert list(b._agent_turns) == ["conv-10:agent-10"] + + for index in range(11, 2110): + assert b._agent_loop_rejection(message(index)) == "" + assert len(b._agent_turns) <= 2048 + + +@pytest.mark.asyncio +async def test_agent_delivery_failure_remains_retryable(tmp_path): + attempts = [] + + async def flaky_delivery(message): + attempts.append(message.message_id) + if len(attempts) == 1: + raise RuntimeError("gateway unavailable") + + b = make_bridge(tmp_path, flaky_delivery) + b.comm.conv_type = "group" + b._policy.group_policy = "open" + b._policy.allow_agent_senders = True + b._policy.agent_allowlist = ["agent-1"] + b.comm.messages["conv-1:1"] = detail( + sender="agent-1", + sender_type="AGENT", + mentions=[{"type": "member", "member_id": "me-1"}], + ) + + with pytest.raises(RuntimeError, match="gateway unavailable"): + await b._handle_frame(msg_frame(sender="agent-1")) + await b._handle_frame(msg_frame(sender="agent-1")) + + assert attempts == ["1", "1"] + assert b.comm.sync_acks == [] + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={"events": [{"conversation_id": "conv-1", "message_id": 1, "seq": 10}]}, + ) + ) + assert b.comm.sync_acks == [10] + + +@pytest.mark.asyncio +async def test_outbound_messages_propagate_agent_loop_metadata(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + + await b.send("conv-1", "reply", metadata={"agent_hop_count": 2}) + + metadata = b.comm.sent_messages[-1][2]["metadata"] + assert metadata["agent_hop_count"] == 3 + assert metadata["agent_origin_member_id"] == "me-1" + assert metadata["agent_trace_id"] + + @pytest.mark.asyncio async def test_dm_reject_notice_throttled(tmp_path): async def on_message(m): @@ -306,6 +749,43 @@ async def send_message(conv_id, text, **kw): assert len(b.comm.sent) == 1 # one notice, throttled +@pytest.mark.asyncio +async def test_group_reject_notice_requires_live_human_mention(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + b.comm.conv_type = "group" + b._policy.group_policy = "disabled" + b._group_mode_overrides["conv-1"] = "silent" + b._policy.group_configs["conv-1"] = { + "mode": "silent", + "allow_from": ["*"], + } + mention = [{"type": "member", "member_id": "me-1"}] + b.comm.messages["conv-1:1"] = detail(mentions=mention) + b.comm.messages["conv-1:2"] = detail(msg_id=2, seq=11, text="background") + b.comm.messages["conv-1:3"] = detail( + msg_id=3, seq=12, mentions=mention, text="replayed mention" + ) + + await b._handle_frame(msg_frame()) + await b._handle_frame(msg_frame(msg_id=2, seq=11)) + await b._handle_frame( + Frame( + type=FRAME_SYNC, + payload={ + "events": [ + {"conversation_id": "conv-1", "message_id": 1, "seq": 10}, + {"conversation_id": "conv-1", "message_id": 2, "seq": 11}, + {"conversation_id": "conv-1", "message_id": 3, "seq": 12}, + ] + }, + ) + ) + + assert len(b.comm.sent_messages) == 1 + assert "disabled" in b.comm.sent_messages[0][1] + assert b.comm.sync_acks == [10, 11, 12] + + @pytest.mark.asyncio async def test_dedup_persistence_across_restart(tmp_path): got = [] diff --git a/tests/test_native_surface_parity.py b/tests/test_native_surface_parity.py index 2c697d7..59e936c 100644 --- a/tests/test_native_surface_parity.py +++ b/tests/test_native_surface_parity.py @@ -342,6 +342,38 @@ def test_access_policy_service_matches_zylos_dm_contract_and_preserves_other_pol service.allow_dm_members([]) +def test_access_policy_service_uses_live_bridge_as_single_writer(): + storage = MemoryStorage( + {"dm_policy": "owner", "dm_allowlist": [], "group_policy": "allowlist"} + ) + live = {"dm_policy": "owner", "dm_allowlist": ["live-user"]} + applied = [] + + def apply(policy, allowlist): + live["dm_policy"] = policy + live["dm_allowlist"] = list(allowlist) + applied.append((policy, list(allowlist))) + return dict(live) + + service = AccessPolicyService( + storage, + get_live_state=lambda: dict(live), + apply_live_state=apply, + ) + + assert service.get_dm_access() == live + assert service.set_dm_policy("allowlist") == { + "dm_policy": "allowlist", + "dm_allowlist": ["live-user"], + } + assert service.allow_dm_members(["agent-owner"]) == { + "dm_policy": "allowlist", + "dm_allowlist": ["live-user", "agent-owner"], + } + assert applied[-1] == ("allowlist", ["live-user", "agent-owner"]) + assert storage.value["dm_allowlist"] == [] + + def test_comm_send_local_attachment_closes_upload_finalize_send_without_returning_url( tmp_path, ): diff --git a/tests/test_reporters.py b/tests/test_reporters.py index 7661bc2..d361d36 100644 --- a/tests/test_reporters.py +++ b/tests/test_reporters.py @@ -59,6 +59,60 @@ def handler(request): assert ("POST", "/api/v1/agents/me-1/online-report") in calls +@pytest.mark.asyncio +async def test_bridge_reports_online_only_after_first_ws_handshake(tmp_path): + import asyncio + + started = asyncio.Event() + release_handshake = asyncio.Event() + online_reports = [] + + class WaitingWs: + def start(self): + started.set() + + async def wait_until_connected(self): + await release_handshake.wait() + + async def stop(self): + pass + + def is_open(self): + return release_handshake.is_set() + + async def no_op(*_args, **_kwargs): + pass + + bridge = CwsBridge( + _cfg(), storage=FileStorage(tmp_path), on_message=no_op + ) + bridge._ws = WaitingWs() + bridge._tokens.get_access_token = no_op + bridge._resolve_identity = no_op + bridge._initialize_or_sync = no_op + bridge._metrics_loop = no_op + bridge._control_sync_loop = no_op + + async def report_online(member_id): + online_reports.append(member_id) + + bridge._online.report = report_online + + starting = asyncio.create_task(bridge.start()) + await started.wait() + await asyncio.sleep(0) + assert online_reports == [] + assert not starting.done() + + release_handshake.set() + await starting + await asyncio.sleep(0) + + assert online_reports == ["me-1"] + assert bridge.is_running() + await bridge.stop() + + @pytest.mark.asyncio async def test_metrics_report_version_only(): bodies = [] @@ -181,7 +235,7 @@ def sysframe(event, data): await b._handle_frame( sysframe( "agent.config.group_mode_changed", - {"conversation_id": "c-9", "mode": "open"}, + {"conversation_id": "c-9", "mode": "smart"}, ) ) assert b._effective_policy("c-9").group_require_mention is False diff --git a/tests/test_runtime_parity.py b/tests/test_runtime_parity.py index f963169..e47e24d 100644 --- a/tests/test_runtime_parity.py +++ b/tests/test_runtime_parity.py @@ -4,6 +4,7 @@ from cws_agent_sdk.access_policy import AccessPolicyConfig, decide_inbound from cws_agent_sdk.codec import FRAME_SYSTEM, Frame +from cws_agent_sdk.providers import FileStorage from cws_agent_sdk.types import InboundMessage from test_bridge import detail, make_bridge @@ -48,29 +49,96 @@ def test_group_scope_allowlist_and_allow_from(): ).handle -def test_group_disabled_but_owner_mention_bypasses(): +def test_group_disabled_blocks_owner_mention(): cfg = AccessPolicyConfig(group_policy="disabled") owner = _msg( sender_id="owner-1", mentions=[{"type": "member", "member_id": "me-1"}], ) - assert decide_inbound( + assert not decide_inbound( owner, self_member_id="me-1", cfg=cfg, owner_member_id="owner-1" ).handle assert not decide_inbound(_msg(), self_member_id="me-1", cfg=cfg).handle def test_plain_text_display_name_mention(): - cfg = AccessPolicyConfig(self_display_name="COCO") + cfg = AccessPolicyConfig(group_policy="open", self_display_name="COCO") decision = decide_inbound( _msg(text="@COCO 请看一下"), self_member_id="me-1", cfg=cfg ) assert decision.handle and decision.reason == "group_mention" +def test_plain_text_mention_uses_boundaries_and_aliases(): + cfg = AccessPolicyConfig( + group_policy="open", + self_display_name="COCO", + self_aliases=["helper.bot"], + ) + assert not decide_inbound( + _msg(text="@COCO-Suffix hello"), self_member_id="me-1", cfg=cfg + ).handle + assert not decide_inbound( + _msg(text="@COCOX hello"), self_member_id="me-1", cfg=cfg + ).handle + assert decide_inbound( + _msg(text="@helper.bot, hello"), self_member_id="me-1", cfg=cfg + ).handle + + +def test_owner_group_exemptions_match_zylos_ordering(): + owner = "owner-1" + mentioned = _msg( + sender_id=owner, + conversation_id="unlisted", + mentions=[{"type": "member", "member_id": "me-1"}], + ) + cfg = AccessPolicyConfig(group_policy="allowlist") + assert decide_inbound( + mentioned, self_member_id="me-1", cfg=cfg, owner_member_id=owner + ).handle + + smart = _msg(sender_id=owner) + cfg = AccessPolicyConfig( + group_policy="allowlist", + group_configs={ + "conv-1": {"mode": "smart", "allow_from": ["someone-else"]} + }, + ) + assert decide_inbound( + smart, self_member_id="me-1", cfg=cfg, owner_member_id=owner + ).handle + + +def test_corrupt_persisted_policy_falls_back_to_safe_normalized_state(tmp_path): + FileStorage(tmp_path).write_json( + "policy.json", + { + "dm_policy": "closed", + "group_policy": "everything", + "dm_allowlist": "not-a-list", + "group_modes": {"c-1": "open", "c-2": "silent"}, + "group_configs": { + "c-2": {"mode": "unexpected", "allow_from": "u-1"}, + "c-3": "bad-shape", + }, + }, + ) + + bridge = make_bridge(tmp_path, lambda _message: None) + + assert bridge._policy.dm_policy == "owner" + assert bridge._policy.group_policy == "allowlist" + assert bridge._policy.dm_allowlist == [] + assert bridge._group_mode_overrides == {"c-2": "silent"} + assert bridge._policy.group_configs == { + "c-2": {"mode": "silent", "allow_from": ["*"]} + } + + def test_sibling_dm_requires_same_owner(): same_owner = _msg(sender_id="agent-1", sender_type="agent", conversation_type="dm") - cfg = AccessPolicyConfig(allow_sibling_dm=True) + cfg = AccessPolicyConfig(allow_sibling_dm=True, agent_allowlist=["agent-1"]) assert decide_inbound( same_owner, self_member_id="me-1", @@ -103,7 +171,7 @@ def config(event, data): {"agent_member_id": "someone-else", "policy": "allowlist"}, ) ) - assert b._policy.dm_policy == "open" + assert b._policy.dm_policy == "owner" await b._handle_frame( config("agent.config.group_scope_changed", {"scope": "allowlist"}) @@ -124,6 +192,116 @@ def config(event, data): assert b._policy.group_configs["conv-1"]["allow_from"] == ["user-1"] +@pytest.mark.asyncio +async def test_invalid_config_events_do_not_mutate_persist_report_or_callback(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + b._cfg.member_id = "me-1" + reports = [] + callbacks = [] + + async def request(method, path, json=None, **_kwargs): + reports.append((method, path, json)) + return {} + + async def on_config(event, data): + callbacks.append((event, data)) + + b._http.request = request + b._on_config_event = on_config + + invalid = [ + ("agent.config.dm_policy_changed", {"policy": "closed"}), + ("agent.config.dm_allowlist_changed", {"action": "append", "member_ids": []}), + ("agent.config.group_mode_changed", {"conversation_id": "c-1", "mode": "open"}), + ("agent.config.group_scope_changed", {"scope": "closed"}), + ("agent.config.group_allowlist_changed", {"action": "set", "conversation_ids": "c-1"}), + ("agent.config.group_allowfrom_changed", {"conversation_id": "c-1", "allow_from": "u-1"}), + ] + for event, data in invalid: + await b._handle_frame( + Frame(type=FRAME_SYSTEM, payload={"event": event, "data": data}) + ) + + assert b._policy.dm_policy == "owner" + assert b._policy.group_policy == "allowlist" + assert b._policy.group_configs == {} + assert b._storage.read_json("policy.json") is None + assert reports == [] + assert callbacks == [] + + +@pytest.mark.asyncio +async def test_silent_config_is_persisted_and_reported_as_effective_group(tmp_path): + b = make_bridge(tmp_path, lambda _message: None) + b._cfg.member_id = "me-1" + reports = [] + + async def request(method, path, json=None, **_kwargs): + reports.append(json) + return {} + + b._http.request = request + await b._handle_frame( + Frame( + type=FRAME_SYSTEM, + payload={ + "event": "agent.config.group_mode_changed", + "data": {"conversation_id": "c-1", "mode": "silent"}, + }, + ) + ) + + assert b._policy.group_configs["c-1"]["mode"] == "silent" + assert b._storage.read_json("policy.json")["group_configs"]["c-1"]["mode"] == "silent" + assert reports[-1]["groups"] == [ + {"conversation_id": "c-1", "mode": "silent", "allow_from": ["*"]} + ] + + +@pytest.mark.asyncio +async def test_local_dm_policy_update_changes_live_state_persistence_and_report(tmp_path): + import asyncio + + b = make_bridge(tmp_path, lambda _message: None) + b._cfg.member_id = "me-1" + reports = [] + + async def request(method, path, json=None, **_kwargs): + reports.append(json) + return {} + + b._http.request = request + result = await b.apply_local_dm_access("allowlist", ["u-1", "u-1"]) + await asyncio.sleep(0) + + assert result == {"dm_policy": "allowlist", "dm_allowlist": ["u-1"]} + assert b.get_dm_access() == result + assert b._storage.read_json("policy.json")["dm_policy"] == "allowlist" + assert reports[-1]["dm_policy"] == "allowlist" + assert reports[-1]["dm_allowlist"] == ["u-1"] + + +@pytest.mark.asyncio +async def test_local_dm_policy_thread_bridge_runs_on_owner_loop(tmp_path): + import asyncio + + b = make_bridge(tmp_path, lambda _message: None) + b._cfg.member_id = "me-1" + b._loop = asyncio.get_running_loop() + + async def request(*_args, **_kwargs): + return {} + + b._http.request = request + result = await asyncio.to_thread( + b.apply_local_dm_access_threadsafe, "open", ["u-1"] + ) + await asyncio.sleep(0) + + assert result == {"dm_policy": "open", "dm_allowlist": ["u-1"]} + assert b.get_dm_access() == result + + @pytest.mark.asyncio async def test_edit_and_recall_are_delivered_as_runtime_messages(tmp_path): got = [] @@ -132,6 +310,7 @@ async def on_message(message): got.append(message) b = make_bridge(tmp_path, on_message) + b._policy.dm_policy = "open" b.comm.messages["conv-1:7"] = detail(msg_id=7, text="latest text") await b._handle_frame( diff --git a/tests/test_ws_runtime_parity.py b/tests/test_ws_runtime_parity.py index fb62e98..c2905a8 100644 --- a/tests/test_ws_runtime_parity.py +++ b/tests/test_ws_runtime_parity.py @@ -1,5 +1,7 @@ """Runtime-only parity regressions for websocket close handling.""" +import asyncio + import pytest from cws_agent_sdk.errors import CwsWsFatal @@ -54,3 +56,17 @@ def test_close_4003_resets_auth_and_remains_recoverable(): assert auth_resets == [True] assert fatals == [] + + +@pytest.mark.asyncio +async def test_initial_connect_wait_propagates_fatal_ws_exit(): + async def fail_fatally(): + raise CwsWsFatal(4002, "invalid credentials") + + client = _client() + client._task = asyncio.create_task(fail_fatally()) + + with pytest.raises(CwsWsFatal) as exc_info: + await client.wait_until_connected() + + assert exc_info.value.code == 4002