diff --git a/agent/agent_init.py b/agent/agent_init.py index 649d5338a996..25728ba044fe 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -70,6 +70,25 @@ def _ra(): return run_agent +def _generate_session_id(now: Optional[datetime] = None) -> str: + """Mint a session id: ``YYYYMMDD_HHMMSS_<24-bit hex>``. + + The 6-hex suffix is only 24 bits, so the id space is ~33.5M — birthday + collisions grow as P ≈ N²/33.5M (~12% at 2,000 sessions). The collision + contract is enforced at the persistence layer: + + * ``create_session`` / ``_insert_session_row`` upsert + (``ON CONFLICT(id) DO UPDATE``) — never raises on collision; the + hermes_state collision warning covers visibility there. + * The compression-rotation child is inserted with a *plain* INSERT + (``publish_compression_child``) and raises ``sqlite3.IntegrityError``; + ``conversation_compression.py`` catches it and retries with a fresh id + (max 3 attempts) instead of aborting the boundary. + """ + base = now or datetime.now() + return f"{base.strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:6]}" + + def _moa_reference_output_allowed(agent: Any) -> bool: """Keep MoA display events off only the machine-readable ``-Q`` surface.""" return not ( @@ -1491,10 +1510,10 @@ def init_agent( # Use provided session ID (e.g., from CLI) agent.session_id = session_id else: - # Generate a new session ID - timestamp_str = agent.session_start.strftime("%Y%m%d_%H%M%S") - short_uuid = uuid.uuid4().hex[:6] - agent.session_id = f"{timestamp_str}_{short_uuid}" + # Generate a new session ID (24-bit hex suffix — collision contract + # documented in _generate_session_id; persistence-layer retries live + # in conversation_compression.py for the plain-INSERT child path). + agent.session_id = _generate_session_id(now=agent.session_start) # Expose session ID to tools (terminal, execute_code) so agents can # reference their own session for --resume commands, cross-session diff --git a/agent/conversation_compression.py b/agent/conversation_compression.py index 3256bca9ec00..b40a489eb8d6 100644 --- a/agent/conversation_compression.py +++ b/agent/conversation_compression.py @@ -58,6 +58,7 @@ import logging import math import os +import sqlite3 import tempfile import time import uuid @@ -3251,24 +3252,50 @@ def _release_lock() -> None: _profile_for_child = None old_title = agent._session_db.get_session_title(agent.session_id) old_session_id = agent.session_id - new_session_id = ( - f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_" - f"{uuid.uuid4().hex[:6]}" - ) - agent._session_db.publish_compression_child( - parent_session_id=old_session_id, - child_session_id=new_session_id, - source=agent.platform - or os.environ.get("HERMES_SESSION_SOURCE", "cli"), - model=agent.model, - model_config=agent._session_init_model_config, - system_prompt=new_system_prompt, - messages=compressed, - cwd=getattr(agent, "working_directory", None), - profile_name=_profile_for_child, - compression_lock_holder=_lock_holder, - require_compression_lease=_lock_holder is not None, - ) + # ── Collision-safe child id (N4) ───────────────────── + # publish_compression_child inserts the child row with a + # plain INSERT (see hermes_state.publish_compression_child), + # so a 24-bit hex suffix collision (P ≈ N²/33.5M — ~12% at + # 2,000 sessions) raises sqlite3.IntegrityError and would + # abort compression at the boundary. Retry with a + # regenerated id (fresh timestamp + fresh 24-bit random), + # max 3 attempts. Each attempt runs in its own transaction + # (BEGIN IMMEDIATE + rollback on error), so a failed insert + # leaves no partial child/parent state to clean up. + new_session_id = None + for _child_attempt in range(1, 4): + new_session_id = ( + f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_" + f"{uuid.uuid4().hex[:6]}" + ) + try: + agent._session_db.publish_compression_child( + parent_session_id=old_session_id, + child_session_id=new_session_id, + source=agent.platform + or os.environ.get("HERMES_SESSION_SOURCE", "cli"), + model=agent.model, + model_config=agent._session_init_model_config, + system_prompt=new_system_prompt, + messages=compressed, + cwd=getattr(agent, "working_directory", None), + profile_name=_profile_for_child, + compression_lock_holder=_lock_holder, + require_compression_lease=_lock_holder is not None, + ) + break # child published — id is durable + except sqlite3.IntegrityError: + if _child_attempt >= 3: + raise + logger.warning( + "Session id collision on compression child " + "(%s); regenerating (attempt %d/3)", + new_session_id, + _child_attempt + 1, + ) + # Tiny sleep so a same-second retry gets a fresh + # timestamp, shrinking re-collision odds further. + time.sleep(0.05) agent.session_id = new_session_id try: from gateway.session_context import set_current_session_id diff --git a/agent/skill_utils.py b/agent/skill_utils.py index a302c6981a47..712203764095 100644 --- a/agent/skill_utils.py +++ b/agent/skill_utils.py @@ -563,13 +563,43 @@ def get_external_skills_dirs() -> List[Path]: return result +def get_workspace_skills_dirs() -> List[Path]: + """Return workspace-local skill directories for the current session. + + Reads the active session context (platform + chat_id + thread_id) and + discovers any workspace-local ``skills/`` directories under the resolved + workspace folder. These are injected into ``get_all_skills_dirs()`` + so workspace-specific skills are available without global installation. + + Uses a 1-second stat cache so changes are picked up automatically — + no gateway restart required. + """ + try: + from gateway.session_context import get_session_env + except ImportError: + return [] + platform = get_session_env("HERMES_SESSION_PLATFORM", "") + chat_id = get_session_env("HERMES_SESSION_CHAT_ID", "") + if not platform or not chat_id: + return [] + thread_id = get_session_env("HERMES_SESSION_THREAD_ID", "") or None + try: + from agent.workspace_resolver import get_workspace_skill_dirs + except ImportError: + return [] + return get_workspace_skill_dirs(platform, chat_id, thread_id) + + def get_all_skills_dirs() -> List[Path]: - """Return all skill directories: local ``~/.hermes/skills/`` first, then external. + """Return all skill directories: local ``~/.hermes/skills/`` first, then + workspace-local, then external. The local dir is always first (and always included even if it doesn't exist - yet — callers handle that). External dirs follow in config order. + yet — callers handle that). Workspace-local dirs follow so per-topic skills + can override global/external ones. External dirs come last in config order. """ dirs = [get_skills_dir()] + dirs.extend(get_workspace_skills_dirs()) dirs.extend(get_external_skills_dirs()) return dirs diff --git a/agent/workspace_resolver.py b/agent/workspace_resolver.py new file mode 100644 index 000000000000..e574f73369ee --- /dev/null +++ b/agent/workspace_resolver.py @@ -0,0 +1,358 @@ +"""Workspace resolver for Hermes channels/topics. + +Provides stat-cached discovery of per-topic and per-channel prompts + skills +from a folder hierarchy under ~/.hermes/workspaces/ and ~/.hermes/platforms/. + +This module is import-safe: no heavy dependencies, no tool registry. + +Folder layout:: + + ~/.hermes/ + ├── workspaces/ + │ ├── news-feed/ + │ │ ├── SYSTEM.md # YAML frontmatter + body + │ │ └── skills/ # optional workspace-local skills + │ │ └── some-skill/ + │ │ └── SKILL.md + │ └── code-review/ + │ ├── SYSTEM.md + │ └── skills/ + └── platforms/ + └── telegram/ + └── -1003682109119/ # chat_id as folder name + └── topics.yaml # thread_id → workspace_name + +SYSTEM.md frontmatter format:: + + --- + skills: + - telegram-summary-bot + - conventional-commits + --- + Respond in Hebrew. Focus on regional news... + +Resolution (per platform message):: + + 1. Discover platform mapping → workspace name + 2. Resolve workspace SYSTEM.md → prompt + skill list + 3. Inject into MessageEvent → channel_prompt + auto_skill + +Fallback chain (most-specific wins):: + + topic workspace prompt > channel workspace prompt > channel_prompts dict + topic workspace skills > channel workspace skills > existing auto_skill + +All disk reads are stat-cached (1-second TTL) so changes take effect +automatically --- no gateway restart required. +""" + +import logging +import os +import time +from pathlib import Path +from typing import Dict, List, NamedTuple, Optional, Tuple, Any, Iterable + +from hermes_constants import get_hermes_home + +logger = logging.getLogger(__name__) + +# ── File cache (path → (mtime_ns, content)) ────────────────────────────────── +_STAT_CACHE: Dict[Tuple[str, int], Tuple[int, str]] = {} +_CACHE_TTL_SECS = 1.0 + +def _stat_cached(path: Path) -> Optional[Tuple[int, str]]: + """Read *path* returning (mtime_ns, content). Returns None if missing. + + Uses a per-second TTL keyed on absolute path. This means edits on disk + are picked up within one second --- cheap enough to call on every + incoming message without restart. + """ + abs_str = str(path.resolve()) + now = time.monotonic() + key = (abs_str, int(now // _CACHE_TTL_SECS)) + cached = _STAT_CACHE.get(key) + if cached is not None: + return cached + if not path.exists(): + _STAT_CACHE[key] = None # type: ignore[assignment] + return None + try: + stat = path.stat() + content = path.read_text(encoding="utf-8") + result = (stat.st_mtime_ns, content) + _STAT_CACHE[key] = result # type: ignore[assignment] + return result + except (OSError, UnicodeDecodeError): + _STAT_CACHE[key] = None # type: ignore[assignment] + return None + + +def _clear_stat_cache() -> None: + """Test hook.""" + _STAT_CACHE.clear() + + +# ── Frontmatter parsing (minimal, dependency-light) ────────────────────────── + +def _parse_frontmatter(text: str) -> Tuple[Dict[str, Any], str]: + """Extract YAML frontmatter and body from SYSTEM.md-style content. + + Supports the ---\n...\n--- pattern. If no frontmatter, returns ({}, text). + """ + lines = text.splitlines(keepends=False) + if not lines or lines[0].strip() != "---": + return {}, text + try: + end = lines.index("---", 1) + except ValueError: + return {}, text + fm_text = "\n".join(lines[1:end]) + body = "\n".join(lines[end + 1 :]).strip() + try: + import yaml + fm = yaml.safe_load(fm_text) or {} + if not isinstance(fm, dict): + fm = {} + except Exception: + fm = {} + return fm, body + + +# ── Named result type ──────────────────────────────────────────────────────── + +class WorkspaceResult(NamedTuple): + """Resolved workspace metadata for a single message.""" + + prompt: str | None + skills: List[str] | None + model: str | None + + +# ═══════════════════════════════════════════════════════════════════════════ +# Public API +# ═══════════════════════════════════════════════════════════════════════════ + + +def resolve_workspace( + platform: str, + chat_id: str, + thread_id: str | None = None, +) -> WorkspaceResult: + """Resolve the workspace prompt + skills for a given (platform, chat, topic). + + Args: + platform: Platform slug (telegram, discord, slack, …) + chat_id: Numeric channel / group / chat id + thread_id: Optional topic / thread id (None for channel-level) + + Returns: + WorkspaceResult(prompt=str|None, skills=list|None) + """ + home = get_hermes_home() + workspace_name = _resolve_workspace_name(home, platform, chat_id, thread_id) + if not workspace_name: + return WorkspaceResult(None, None, None) + return _resolve_workspace_content(home, workspace_name) + + +def get_workspace_skill_dirs( + platform: str, + chat_id: str, + thread_id: str | None = None, +) -> List[Path]: + """Return workspace-local skill directories for the resolved topic. + + These paths are meant to be appended to ``get_all_skills_dirs()`` so + workspace-specific skills are discoverable without global installation. + """ + home = get_hermes_home() + workspace_name = _resolve_workspace_name(home, platform, chat_id, thread_id) + if not workspace_name: + return [] + ws_dir = _workspace_dir(home, workspace_name) + if ws_dir is None: + return [] + skills_dir = ws_dir / "skills" + if skills_dir.exists() and skills_dir.is_dir(): + return [skills_dir] + return [] + + +def apply_workspace_to_event(event: Any) -> bool: + """Resolve workspace prompt/skills/model from ``event.source`` and apply. + + Shared resolution path (review fix): called from the gateway message + handler for EVERY platform (Telegram, Discord, Slack, CLI, TUI, ...), so + workspace context is not Telegram-only. Idempotent — events already + resolved (``_workspace_applied`` flag) are skipped. + + Priority (matching the folder-resolver contract): + workspace prompt > platform topic/channel prompt already on the event + workspace skills > (merged with) existing auto_skill + workspace model > (overrides config default; /model session override + still wins — enforced in run.py) + + Returns True when a workspace was applied. + """ + if getattr(event, "_workspace_applied", False): + return False + source = getattr(event, "source", None) + if source is None: + return False + platform = getattr(source, "platform", None) + if platform is None: + return False + platform_slug = platform.value if hasattr(platform, "value") else str(platform) + chat_id = str(getattr(source, "chat_id", "") or "") + if not chat_id: + return False + thread_id = getattr(source, "thread_id", None) + thread_id_str = str(thread_id) if thread_id is not None else None + + ws = resolve_workspace(platform_slug, chat_id, thread_id_str) + if not (ws.prompt or ws.skills or ws.model): + event._workspace_applied = True # type: ignore[attr-defined] + return False + + if ws.prompt: + event.channel_prompt = ws.prompt + if ws.skills: + current = getattr(event, "auto_skill", None) + merged = list(ws.skills) + if isinstance(current, str): + if current not in merged: + merged.append(current) + elif isinstance(current, list): + for s in current: + if s not in merged: + merged.append(s) + event.auto_skill = merged + if ws.model: + event.workspace_model = ws.model + event._workspace_applied = True # type: ignore[attr-defined] + return True + + +# ═══════════════════════════════════════════════════════════════════════════ +# Internal helpers +# ═══════════════════════════════════════════════════════════════════════════ + + +def _resolve_workspace_name( + home: Path, + platform: str, + chat_id: str, + thread_id: str | None = None, +) -> str | None: + """Read platforms///topics.yaml, resolve workspace name.""" + mapping_file = home / "platforms" / platform / _safe_dir_name(chat_id) / "topics.yaml" + stat_result = _stat_cached(mapping_file) + if stat_result is None: + return None + _, text = stat_result + try: + import yaml + data = yaml.safe_load(text) or {} + except Exception: + logger.debug("[workspace] Failed to parse %s", mapping_file) + return None + + topics = data.get("topics", {}) + if isinstance(topics, dict): + # Direct dict form: {"7695": "news-feed"} + if thread_id is not None and thread_id in topics: + return topics[thread_id] + elif isinstance(topics, list): + # List-of-dicts form (allows YAML comments / ordering) + for entry in topics: + if isinstance(entry, dict) and str(entry.get("thread_id", "")) == str(thread_id or ""): + return entry.get("workspace") + + # Fall back to channel-level workspace + channel_ws = data.get("workspace") if isinstance(data, dict) else None + return channel_ws + + +def _resolve_workspace_content( + home: Path, + workspace_name: str, +) -> WorkspaceResult: + """Read workspaces//SYSTEM.md, extract prompt + skills + model.""" + ws_dir = _workspace_dir(home, workspace_name) + if ws_dir is None: + return WorkspaceResult(None, None, None) + system_file = ws_dir / "SYSTEM.md" + stat_result = _stat_cached(system_file) + if stat_result is None: + return WorkspaceResult(None, None, None) + _, text = stat_result + fm, body = _parse_frontmatter(text) + prompt = body.strip() or None + skills = fm.get("skills") + if isinstance(skills, str): + skills = [skills.strip()] if skills.strip() else None + elif isinstance(skills, list): + parsed: List[str] = [] + for s in skills: + if isinstance(s, str) and s.strip(): + parsed.append(s.strip()) + skills = parsed if parsed else None + else: + skills = None + model = fm.get("model") + if not isinstance(model, str) or not model.strip(): + model = None + return WorkspaceResult(prompt, skills, model) + + +def _safe_dir_name(value: str) -> str: + """Escape/validate a chat id or workspace name for use as a directory name. + + Rejects path separators, traversal segments, and NUL (review fix: a name + coming from topics.yaml or user input must never escape the parent + directory). A leading minus becomes escaped (e.g. "-1003" → "_-1003") + so ``Path`` doesn't interpret it as a relative-segment trick. Returns + "_invalid" for names that cannot be used as a single directory segment. + """ + if not value: + return "_invalid" + if value in (".", ".."): + return "_invalid" + if "/" in value or "\\" in value or "\x00" in value: + return "_invalid" + if ".." in value: + return "_invalid" + if value.startswith("-"): + return "_-" + value[1:] + return value + + +def _workspace_dir(home: Path, workspace_name: str) -> Optional[Path]: + """Return the validated workspace dir under ``/workspaces`` or None. + + Double-checks (defense in depth) that the resolved path still lives under + ``/workspaces`` before callers read SYSTEM.md from it. + """ + safe = _safe_dir_name(workspace_name) + if safe == "_invalid": + return None + base = (home / "workspaces").resolve() + candidate = (home / "workspaces" / safe).resolve() + if candidate == base or not str(candidate).startswith(str(base) + os.sep): + return None + return candidate + + +def _write_system_md(path: Path, frontmatter: Dict[str, Any], body: str) -> None: + """Write a SYSTEM.md file with YAML frontmatter + body. + + Preserves frontmatter fields and writes them in a deterministic order. + If frontmatter is empty, writes body only (no frontmatter fence). + """ + import yaml + if frontmatter: + fm_text = yaml.safe_dump(frontmatter, default_flow_style=False).strip() + content = f"---\n{fm_text}\n---\n{body}" + else: + content = body + path.write_text(content, encoding="utf-8") diff --git a/cli.py b/cli.py index aed3992922b4..97cfc1732820 100644 --- a/cli.py +++ b/cli.py @@ -9832,6 +9832,164 @@ def _show_gateway_status(self): print(f" 2. Or configure settings in {display_hermes_home()}/config.yaml") print() + def _handle_workspace_command(self, cmd: str): + """Handle the /workspace command — manage workspace prompts and skill links.""" + from hermes_constants import get_hermes_home + from agent.workspace_resolver import ( + resolve_workspace, _safe_dir_name, _clear_stat_cache, + ) + + parts = cmd.split(maxsplit=2) + subcmd = parts[1].lower() if len(parts) > 1 else "list" + rest = parts[2].strip() if len(parts) > 2 else "" + + hermes_home = get_hermes_home() + workspaces_dir = hermes_home / "workspaces" + platforms_dir = hermes_home / "platforms" + + if subcmd in ("list", ""): + if not workspaces_dir.exists(): + self.console.print(" No workspaces found. Use /workspace create to create one.") + return + ws_dirs = sorted(d.name for d in workspaces_dir.iterdir() if d.is_dir()) + if not ws_dirs: + self.console.print(" No workspaces found. Use /workspace create to create one.") + return + self.console.print("\n[bold]Workspaces[/bold]\n") + for ws_name in ws_dirs: + system_file = workspaces_dir / ws_name / "SYSTEM.md" + skills_dir = workspaces_dir / ws_name / "skills" + flags = [] + if system_file.exists(): + flags.append("prompt") + if skills_dir.exists() and any(skills_dir.iterdir()): + n = len([d for d in skills_dir.iterdir() if d.is_dir()]) + flags.append(f"{n} skill{'s' if n != 1 else ''}") + desc = f" ({', '.join(flags)})" if flags else "" + self.console.print(f" • [cyan]{ws_name}[/cyan]{desc}") + self.console.print() + + elif subcmd == "create": + create_parts = rest.split(None, 1) + if not create_parts: + self.console.print("Usage: [cyan]/workspace create [prompt text][/cyan]") + return + ws_name = create_parts[0] + prompt_text = create_parts[1].strip() if len(create_parts) > 1 else "" + if not ws_name.replace("-", "").replace("_", "").isalnum(): + self.console.print(f"[red]✗ Invalid workspace name '{ws_name}'. Use letters, numbers, hyphens, underscores.[/red]") + return + ws_dir = workspaces_dir / ws_name + if ws_dir.exists(): + self.console.print(f"[red]✗ Workspace '{ws_name}' already exists. Use /workspace show {ws_name}[/red]") + return + ws_dir.mkdir(parents=True, exist_ok=True) + system_file = ws_dir / "SYSTEM.md" + system_file.write_text(prompt_text + "\n" if prompt_text else f"# Workspace: {ws_name}\n", encoding="utf-8") + _clear_stat_cache() + self.console.print(f"[green]✓ Created workspace '{ws_name}'[/green]") + self.console.print(f" Edit prompt: {system_file}") + + elif subcmd == "show": + ws_name = rest.strip() + if not ws_name: + self.console.print("Usage: [cyan]/workspace show [/cyan]") + return + ws_dir = workspaces_dir / ws_name + if not ws_dir.exists(): + self.console.print(f"[red]✗ Workspace '{ws_name}' not found.[/red]") + return + from rich.markdown import Markdown + system_file = ws_dir / "SYSTEM.md" + self.console.print(f"\n[bold]Workspace: {ws_name}[/bold]\n") + if system_file.exists(): + content = system_file.read_text(encoding="utf-8") + self.console.print(Markdown(content)) + skills_dir = ws_dir / "skills" + if skills_dir.exists(): + skill_list = sorted(d.name for d in skills_dir.iterdir() if d.is_dir()) + if skill_list: + self.console.print(f"\n[bold]Skills:[/bold] {', '.join(skill_list)}") + + elif subcmd == "link": + link_parts = rest.split() + if len(link_parts) < 4: + self.console.print("Usage: [cyan]/workspace link [/cyan]") + self.console.print(" or: [cyan]/workspace link default [/cyan]") + return + link_platform, link_chat_id, link_thread_id, link_ws = link_parts[:4] + ws_dir = workspaces_dir / link_ws + if not ws_dir.exists(): + self.console.print(f"[red]✗ Workspace '{link_ws}' not found. Create it first: /workspace create {link_ws}[/red]") + return + safe_id = _safe_dir_name(link_chat_id) + topic_dir = platforms_dir / link_platform / safe_id + topic_dir.mkdir(parents=True, exist_ok=True) + topics_file = topic_dir / "topics.yaml" + import yaml as _yaml + data = {} + if topics_file.exists(): + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + if isinstance(topics, list): + topics = {str(item.get("thread_id", "")): str(item.get("workspace", "")) for item in topics if isinstance(item, dict)} + topics[str(link_thread_id)] = link_ws + data["topics"] = topics + if "workspace" not in data: + data["workspace"] = "default" + topics_file.write_text(_yaml.safe_dump(data, default_flow_style=False), encoding="utf-8") + _clear_stat_cache() + self.console.print(f"[green]✓ Linked {link_platform}/{link_chat_id} topic {link_thread_id} → {link_ws}[/green]") + + elif subcmd in ("remove", "unlink"): + if subcmd == "unlink": + unlink_parts = rest.split() + if len(unlink_parts) < 3: + self.console.print("Usage: [cyan]/workspace unlink [/cyan]") + return + u_platform, u_chat_id, u_thread_id = unlink_parts[:3] + safe_id = _safe_dir_name(u_chat_id) + topics_file = platforms_dir / u_platform / safe_id / "topics.yaml" + if not topics_file.exists(): + self.console.print(f"[red]✗ No topics.yaml found for {u_platform}/{u_chat_id}[/red]") + return + import yaml as _yaml + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + if str(u_thread_id) not in topics: + self.console.print(f"[red]✗ Thread '{u_thread_id}' not linked in {u_platform}/{u_chat_id}[/red]") + return + del topics[str(u_thread_id)] + data["topics"] = topics + topics_file.write_text(_yaml.safe_dump(data, default_flow_style=False), encoding="utf-8") + _clear_stat_cache() + self.console.print(f"[green]✓ Unlinked thread {u_thread_id} from {u_platform}/{u_chat_id}[/green]") + else: + ws_name = rest.strip() + if not ws_name: + self.console.print("Usage: [cyan]/workspace remove [/cyan]") + return + ws_dir = workspaces_dir / ws_name + if not ws_dir.exists(): + self.console.print(f"[red]✗ Workspace '{ws_name}' not found.[/red]") + return + import shutil + shutil.rmtree(ws_dir) + _clear_stat_cache() + self.console.print(f"[green]✓ Removed workspace '{ws_name}'[/green]") + self.console.print(" [dim]Any topic links pointing to it are now broken — use /workspace unlink to clean up.[/dim]") + + else: + self.console.print("\n[bold]Workspace commands[/bold]\n") + self.console.print(" [cyan]/workspace[/cyan] or [cyan]/workspace list[/cyan] — List all workspaces") + self.console.print(" [cyan]/workspace create [prompt][/cyan] — Create a workspace") + self.console.print(" [cyan]/workspace show [/cyan] — Show workspace details") + self.console.print(" [cyan]/workspace link [/cyan] — Link a topic") + self.console.print(" [cyan]/workspace unlink [/cyan] — Unlink a topic") + self.console.print(" [cyan]/workspace remove [/cyan] — Delete a workspace") + self.console.print() + + def process_command(self, command: str) -> bool: """ Process a slash command. @@ -10044,6 +10202,8 @@ def process_command(self, command: str) -> bool: elif canonical == "personality": # Use original case (handler lowercases the personality name itself) self._handle_personality_command(cmd_original) + elif canonical == "workspace": + self._handle_workspace_command(cmd_original) elif canonical == "pet": self._handle_pet_command(cmd_original) diff --git a/gateway/config.py b/gateway/config.py index b39ae8cfcb40..afc18a19e0da 100644 --- a/gateway/config.py +++ b/gateway/config.py @@ -1566,6 +1566,13 @@ def _merge_platform_map(source_platforms: Any) -> None: bridged["exclusive_bot_mentions"] = platform_cfg["exclusive_bot_mentions"] if plat == Platform.TELEGRAM and "observe_unmentioned_group_messages" in platform_cfg: bridged["observe_unmentioned_group_messages"] = platform_cfg["observe_unmentioned_group_messages"] + if plat == Platform.TELEGRAM: + # Bridge Telegram topic config from top-level to extra + # so users can write: telegram: { group_topics: [...] } + # instead of: telegram: { extra: { group_topics: [...] } } + for _topic_key in ("group_topics", "dm_topics"): + if _topic_key in platform_cfg: + bridged[_topic_key] = platform_cfg[_topic_key] if "dm_policy" in platform_cfg: bridged["dm_policy"] = platform_cfg["dm_policy"] if "allow_from" in platform_cfg: diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index fa5ba936d574..8a954f5dad57 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -3371,6 +3371,16 @@ def _atomic(conn): (clean_title, session_id), ).fetchone() if conflict: + # Defensive purge (N9): this row was inserted moments + # ago in the same transaction, so it has no messages + # or on-disk files — but under session-id reuse a + # stale gateway_routing entry could still point at + # this id. Purge routing on the in-flight conn (same + # pattern as delete_session) BEFORE deleting the row. + # Do NOT call db.delete_session() here: it opens its + # own BEGIN IMMEDIATE via _execute_write, which fails + # inside this already-open transaction. + db._purge_gateway_routing_for_sessions(conn, [session_id]) conn.execute( "DELETE FROM sessions WHERE id = ?", (session_id,) ) @@ -3447,6 +3457,36 @@ async def _handle_delete_session(self, request: "web.Request") -> "web.Response" sessions_dir = Path(get_hermes_home()) / "sessions" deleted = await asyncio.to_thread(db.delete_session, session_id, sessions_dir) + + # If the row was actually deleted (not already absent), drop the + # in-memory entry from the live SessionStore so a running gateway + # does not resurrect the dead session on the next save. forget_sessions + # may not exist yet while gateway/session.py is being upgraded, so + # probe defensively and never propagate store failures to the caller. + if deleted: + store = getattr(self, "_session_store", None) + if store is not None: + try: + forget = getattr(store, "forget_sessions", None) + if callable(forget): + forgotten = await asyncio.to_thread(forget, [session_id]) + logger.info( + "forget_sessions removed %s in-memory entr%s for deleted session %s after API delete", + forgotten, + "y" if forgotten == 1 else "ies", + session_id, + ) + else: + logger.debug( + "SessionStore has no forget_sessions yet; skipped in-memory forget for deleted session %s", + session_id, + ) + except Exception: + logger.warning( + "Failed to forget session %s from SessionStore after API delete", + session_id, + exc_info=True, + ) return web.json_response({"object": "hermes.session.deleted", "id": session_id, "deleted": bool(deleted)}) async def _handle_session_messages(self, request: "web.Request") -> "web.Response": diff --git a/gateway/platforms/base.py b/gateway/platforms/base.py index c42b9160737d..18c1e96ce4af 100644 --- a/gateway/platforms/base.py +++ b/gateway/platforms/base.py @@ -2107,6 +2107,10 @@ class MessageEvent: # Applied at API call time and never persisted to transcript history. channel_prompt: Optional[str] = None + # Workspace-specified model override. Takes effect on every message + # (including after /new resets). Priority: /model command > workspace > config. + workspace_model: Optional[str] = None + # Channel context recovered by history backfill (e.g. messages between # bot turns that were missed due to require_mention). Kept separate # from ``text`` so the sender-prefix logic in run.py can operate on the @@ -3405,11 +3409,21 @@ def _is_sender_authorized( def set_session_store(self, session_store: Any) -> None: """ - Set the session store for checking active sessions. - - Used by adapters that need to check if a thread/conversation - has an active session before processing messages (e.g., Slack - thread replies without explicit mentions). + set_session_store(storage) — chamado pelo gateway runner no startup + (gateway/run.py: startup, reconnect e _configure_profile_adapter); + usado pelo api_server para podar entries em memória após deletes + (api_server lê ``self._session_store`` via getattr e chama + ``forget_sessions`` quando uma sessão é apagada pela API). + + Também usado por adapters que precisam checar se uma + thread/conversação tem sessão ativa antes de processar mensagens + (e.g., Slack thread replies sem menção explícita) e para resolver + o diretório de sessões a partir do SessionStore vivo. + + Garantia de ordem: o runner instancia ``self.session_store`` no + ``__init__`` (antes de conectar qualquer adapter), então todo + adapter — incluindo o APIServerAdapter — recebe o store antes de + qualquer request de delete chegar. """ self._session_store = session_store diff --git a/gateway/run.py b/gateway/run.py index 78a49342c2a6..abccb7f81631 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -4332,6 +4332,28 @@ def run_sync(self): "run_agent resolved: model=%s provider=%s session=%s", model, runtime_kwargs.get("provider"), ctx.session_key or "", ) + # Workspace model override: takes effect when no /model session + # override is active. On /new, session overrides are cleared, + # so workspace model kicks back in automatically. + if ctx.workspace_model: + resolved_session_key = ctx.session_key + if not resolved_session_key and ctx.source is not None: + try: + resolved_session_key = self._runner._session_key_for_source(ctx.source) + except Exception: + resolved_session_key = None + session_override = self._runner._session_model_overrides.get(resolved_session_key) if resolved_session_key else None + if not session_override: + logger.info( + "Workspace model override: session=%s config_model=%s -> workspace_model=%s", + resolved_session_key or "", model, ctx.workspace_model, + ) + model = ctx.workspace_model + else: + logger.debug( + "Session /model override takes priority over workspace model: session=%s override=%s workspace=%s", + resolved_session_key or "", session_override.get("model", "?"), ctx.workspace_model, + ) except Exception as exc: return { "final_response": f"⚠️ Provider authentication failed: {exc}", @@ -10981,6 +11003,16 @@ async def start(self) -> bool: adapter.set_message_handler(self._primary_message_handler()) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) adapter.set_session_store(self.session_store) + # GAP-2 defense: SessionStore is created in __init__ (L3439) + # before any adapter exists; set_session_store (base.py L3387) + # lands it on the adapter HERE, before connect() at L8489 starts + # serving requests. api_server handlers read it via + # getattr(self, "_session_store") (api_server.py ~L3274). This + # assert turns a future reorder into a loud boot failure instead + # of a silent session-lookup regression. + assert getattr(adapter, "_session_store", None) is self.session_store, ( + f"{platform.value}: SessionStore not wired to adapter before connect" + ) adapter.set_busy_session_handler(self._handle_active_session_busy_message) _set_reaction = getattr(adapter, "set_reaction_handler", None) if callable(_set_reaction): @@ -12353,6 +12385,10 @@ async def _platform_reconnect_watcher(self) -> None: adapter.set_message_handler(self._primary_message_handler()) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) adapter.set_session_store(self.session_store) + # GAP-2 defense (reconnect path — new adapter object). + assert getattr(adapter, "_session_store", None) is self.session_store, ( + f"{platform.value}: SessionStore not wired to adapter before connect" + ) adapter.set_busy_session_handler(self._handle_active_session_busy_message) _set_reaction = getattr(adapter, "set_reaction_handler", None) if callable(_set_reaction): @@ -13295,6 +13331,10 @@ def _configure_profile_adapter( self._make_profile_fatal_error_handler(profile_name, platform) ) adapter.set_session_store(self.session_store) + # GAP-2 defense (secondary-profile path — new adapter object). + assert getattr(adapter, "_session_store", None) is self.session_store, ( + f"{platform.value}: SessionStore not wired to adapter before connect" + ) adapter.set_busy_session_handler(self._handle_active_session_busy_message) _set_reaction = getattr(adapter, "set_reaction_handler", None) if callable(_set_reaction): @@ -13684,6 +13724,10 @@ def _create_adapter( return None adapter = APIServerAdapter(config) adapter.gateway_runner = self + # The SessionStore is NOT wired here: the caller does it via + # set_session_store() right after _create_adapter returns + # (startup L8471 / reconnect L9571 / profile L10505), always + # before connect(). api_server reads it as ``_session_store``. return adapter elif platform == Platform.WEBHOOK: @@ -14132,6 +14176,7 @@ async def _busy_queue_command(self, event: MessageEvent, quick_key: str, source) reply_to_is_own_message=event.reply_to_is_own_message, auto_skill=event.auto_skill, channel_prompt=event.channel_prompt, + workspace_model=getattr(event, 'workspace_model', None), channel_context=event.channel_context, internal=event.internal, timestamp=event.timestamp, @@ -14163,6 +14208,7 @@ async def _busy_steer_command(self, event: MessageEvent, quick_key: str, source) source=event.source, message_id=event.message_id, channel_prompt=event.channel_prompt, + workspace_model=getattr(event, 'workspace_model', None), channel_context=event.channel_context, ) self._enqueue_fifo(quick_key, queued_event, adapter) @@ -14186,6 +14232,7 @@ async def _busy_steer_command(self, event: MessageEvent, quick_key: str, source) source=event.source, message_id=event.message_id, channel_prompt=event.channel_prompt, + workspace_model=getattr(event, 'workspace_model', None), channel_context=event.channel_context, ) self._enqueue_fifo(quick_key, queued_event, adapter) @@ -14225,6 +14272,16 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: """ source = event.source + # Workspace resolution — SHARED path (review fix): folder-based + # workspaces apply for EVERY platform (Discord, Slack, CLI, TUI, ...), + # not just Telegram. Idempotent: adapters that already resolved + # (``_workspace_applied`` flag) are skipped. + try: + from agent.workspace_resolver import apply_workspace_to_event + apply_workspace_to_event(event) + except Exception: + logger.debug("Workspace resolution skipped for %s", source.platform if source else "?", exc_info=True) + # 🔴 Cross-session leak guard. This handler runs inside a per-message # asyncio task created via create_task(), which snapshots the spawning # context with copy_context(). If a *concurrent* message had already @@ -15090,6 +15147,9 @@ async def _do_reset(): if canonical == "personality": return await self._handle_personality_command(event) + if canonical == "workspace": + return await self._handle_workspace_command(event) + if canonical == "kanban": return await self._handle_kanban_command(event) @@ -17399,6 +17459,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str, run_g run_generation=run_generation, event_message_id=self._reply_anchor_for_event(event), channel_prompt=event.channel_prompt, + workspace_model=getattr(event, 'workspace_model', None), moa_config=getattr(event, "_moa_config", None), persist_user_message=persist_user_message, persist_user_timestamp=persist_user_timestamp, @@ -18151,21 +18212,44 @@ def _reset_notice_session_info(self, source: SessionSource) -> str: inside this method, so contextvars behave correctly in the worker thread. """ + workspace_model = None + try: + from agent.workspace_resolver import resolve_workspace + ws = resolve_workspace( + source.platform.value if source.platform else "", + str(source.chat_id or ""), + str(source.thread_id) if getattr(source, "thread_id", None) else None, + ) + workspace_model = ws.model + except Exception: + pass if getattr(getattr(self, "config", None), "multiplex_profiles", False): with _profile_runtime_scope(self._resolve_profile_home_for_source(source)): - return self._format_session_info() - return self._format_session_info() + return self._format_session_info(workspace_model=workspace_model) + return self._format_session_info(workspace_model=workspace_model) - def _format_session_info(self) -> str: + def _format_session_info(self, workspace_model: Optional[str] = None) -> str: """Resolve current model config and return a formatted info block. Surfaces model, provider, context length, and endpoint so gateway users can immediately see if context detection went wrong (e.g. local models falling to the 128K default). + + When *workspace_model* is provided, the display shows the effective + model (workspace override if no /model session override is active) + rather than the config default, so gateway users see which model + will actually be used. """ from agent.model_metadata import get_model_context_length, DEFAULT_FALLBACK_CONTEXT - model = _resolve_gateway_model() + config_model = _resolve_gateway_model() + effective_model = model = config_model + # If a workspace model is provided and differs from the config default, + # use it as the effective model so the user sees what will actually run. + model_source = "config" + if workspace_model and workspace_model != config_model: + effective_model = model = workspace_model + model_source = "workspace" config_context_length = None provider = None base_url = None @@ -18265,7 +18349,7 @@ def _format_session_info(self) -> str: ctx_display = str(context_length) lines = [ - f"◆ Model: `{model}`", + f"◆ Model: `{effective_model}`{' (workspace)' if model_source == 'workspace' else ''}", f"◆ Provider: {provider or 'openrouter'}", f"◆ Context: {ctx_display} tokens ({ctx_source})", ] @@ -23872,6 +23956,7 @@ async def _run_agent( _interrupt_depth: int = 0, event_message_id: Optional[str] = None, channel_prompt: Optional[str] = None, + workspace_model: Optional[str] = None, moa_config: Optional[dict] = None, persist_user_message: Optional[Any] = None, persist_user_timestamp: Optional[float] = None, @@ -23891,7 +23976,7 @@ async def _run_agent( message, context_prompt, history, source, session_id, session_key=session_key, run_generation=run_generation, _interrupt_depth=_interrupt_depth, event_message_id=event_message_id, - channel_prompt=channel_prompt, moa_config=moa_config, + channel_prompt=channel_prompt, workspace_model=workspace_model, moa_config=moa_config, persist_user_message=persist_user_message, persist_user_timestamp=persist_user_timestamp, message_type=message_type, @@ -24025,6 +24110,7 @@ async def _run_agent_inner( _interrupt_depth: int = 0, event_message_id: Optional[str] = None, channel_prompt: Optional[str] = None, + workspace_model: Optional[str] = None, moa_config: Optional[dict] = None, persist_user_message: Optional[Any] = None, persist_user_timestamp: Optional[float] = None, @@ -24311,6 +24397,7 @@ def _generic_status_phrase(kind: str, *, tool_name: str | None = None, preview: _interrupt_depth=_interrupt_depth, event_message_id=event_message_id, moa_config=moa_config, + workspace_model=workspace_model, persist_user_message=persist_user_message, persist_user_timestamp=persist_user_timestamp, ) @@ -24578,10 +24665,12 @@ async def write_tool_log(): except Exception as _stts_err: logger.debug("Could not set up streaming TTS consumer: %s", _stts_err) + # run_sync extracted to TurnRunner.run_sync (bound method; the # executor call below is unchanged). Its closed-over locals travel # on turn_ctx; `nonlocal message` rebinds became ctx.message writes. run_sync = turn_runner.run_sync + # Start progress message sender if enabled. Gate on needs_progress_queue # (tool_progress OR thinking_progress), not tool_progress alone: the @@ -25432,6 +25521,7 @@ def _run_sync_with_timeout_lifecycle(): next_message = pending next_message_id = None next_channel_prompt = None + next_workspace_model = None next_session_key = session_key # #60671 — carry the pending event's message_type into the # recursive call so queued voice turns can stream TTS and @@ -25468,6 +25558,7 @@ def _run_sync_with_timeout_lifecycle(): return result next_message_id = self._reply_anchor_for_event(pending_event) next_channel_prompt = getattr(pending_event, "channel_prompt", None) + next_workspace_model = getattr(pending_event, "workspace_model", None) next_message_type = getattr(pending_event, "message_type", None) # Clear the completed streaming marker from the prior logical @@ -25523,6 +25614,7 @@ def _run_sync_with_timeout_lifecycle(): _interrupt_depth=_interrupt_depth + 1, event_message_id=next_message_id, channel_prompt=next_channel_prompt, + workspace_model=next_workspace_model, message_type=next_message_type, ) return _preserve_queued_followup_history_offset(result, followup_result) diff --git a/gateway/session.py b/gateway/session.py index 8bcbe56c06db..8efde3f0c6f6 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -1501,6 +1501,14 @@ def _prune_stale_sessions_locked(self) -> None: def _save(self) -> None: """Persist the routing index while the caller holds ``_lock``.""" data, generation = self._snapshot_routing_locked() + # Multi-process backstop (#GAP-3): another process (TUI session + # delete, second gateway, `hermes sessions delete`) may have deleted + # a state.db row this process still tracks. Drop such dead entries + # from the snapshot and from ``_entries`` before persisting so the + # save never resurrects them. + data, dead = self._prune_dead_persisted_entries(data) + if dead: + self._drop_entries_locked(dead) self._persist_routing_data(data, generation) def _next_routing_generation_locked(self) -> int: @@ -1612,8 +1620,17 @@ def _save_entries(self) -> None: """Snapshot latest state under ``_lock`` and persist after releasing it.""" with self._lock: data, generation = self._snapshot_routing_locked() + # Multi-process backstop (#GAP-3): drop entries whose state.db row + # was deleted by another process before persisting, so the save + # never resurrects them. The in-memory drop re-acquires ``_lock`` + # (unlike ``_save``, the caller does not hold it here). + data, dead = self._prune_dead_persisted_entries(data) + if dead: + with self._lock: + self._drop_entries_locked(dead) self._persist_routing_data(data, generation) + def _save_entry(self, session_key: str) -> None: """Persist ONE routing entry via UPSERT — the per-turn fast path. @@ -1714,6 +1731,158 @@ def _forget_fast_persisted(self, session_key: str) -> None: if fast_persisted is not None: fast_persisted.pop(session_key, None) + + def forget_sessions(self, session_ids: List[str]) -> int: + """Drop in-memory routing entries whose ``session_id`` is in *session_ids*. + + Contract: called after a session row has been deleted from state.db + (e.g. the API ``DELETE /api/sessions/{id}`` flow) so a running + gateway does not resurrect the dead session on its next save — the + entry would otherwise keep being re-persisted by ``_save`` / + ``_save_entries``. + + - Matches entries by ``entry.session_id`` (the routing key is not + used for matching, since a key can be re-bound to a new session). + - Returns the number of entries actually removed. + - Persists the routing index (via ``_save``) only when at least one + entry was removed — a no-op match avoids a pointless whole-index + rewrite. + - Thread-safe: runs under ``self._lock`` and ``_ensure_loaded_locked``. + """ + targets = {sid for sid in session_ids if sid} + if not targets: + return 0 + with self._lock: + self._ensure_loaded_locked() + removed = 0 + for key in [ + k for k, e in self._entries.items() if e.session_id in targets + ]: + del self._entries[key] + removed += 1 + if removed: + self._save() + return removed + + def _drop_entries_locked(self, dead: Dict[str, str]) -> None: + """Remove dead ``{session_key: session_id}`` pairs from ``_entries``. + + Must be called with ``self._lock`` held. Each key is only dropped + when it still maps to the exact dead ``session_id`` observed in the + snapshot — a key re-bound to a newer live session (reset/switch) + since the snapshot was taken is left untouched. + """ + for key, session_id in dead.items(): + entry = self._entries.get(key) + if entry is not None and entry.session_id == session_id: + del self._entries[key] + + def _prune_dead_persisted_entries( + self, data: Dict[str, Any] + ) -> tuple[Dict[str, Any], Dict[str, str]]: + """Filter entries whose state.db row no longer exists from a snapshot. + + Multi-process backstop (#GAP-3): ``SessionEntry.db_persisted`` means + a matching row exists in the ``sessions`` table of state.db. When + another process deletes that row, this process's in-memory entry is + stale — persisting it again would resurrect the deleted session. + This method verifies every ``db_persisted=True`` entry's + ``session_id`` still exists via one batched query + (``SELECT id FROM sessions WHERE id IN (...)``, 500 ids per chunk) + and returns the snapshot with dead entries removed, plus the + ``{session_key: session_id}`` pairs that died so callers can also + drop them from ``_entries``. + + - Legacy entries (``db_persisted=False``) are never queried and are + always preserved — they predate the state.db routing migration and + legitimately have no DB row. + - If no entry is ``db_persisted``, no DB query is issued at all. + - A DB failure never breaks the save: on any exception the snapshot + is returned unchanged (dead entries are persisted as before, with + a warning logged) — the backstop degrades to a no-op, it never + raises. + """ + persisted_ids: Dict[str, str] = {} + for key, entry_data in data.items(): + if not isinstance(entry_data, dict): + continue + if entry_data.get("db_persisted") and entry_data.get("session_id"): + persisted_ids[entry_data["session_id"]] = key + if not persisted_ids: + return data, {} + + _db = getattr(self, "_db", None) + if _db is None: + return data, {} + + existing: set[str] = set() + try: + read_ctx = getattr(_db, "_read_ctx", None) + ids = list(persisted_ids.keys()) + if callable(read_ctx): + with read_ctx() as conn: + existing = self._query_existing_session_ids(conn, ids) + else: + # Fallback: shared writer connection under the SessionDB lock. + db_lock = getattr(_db, "_lock", None) + conn = getattr(_db, "_conn", None) + if conn is None: + return data, {} + if db_lock is not None: + with db_lock: + existing = self._query_existing_session_ids(conn, ids) + else: + existing = self._query_existing_session_ids(conn, ids) + except Exception as exc: + logger.warning( + "gateway.session: dead-entry backstop query failed; " + "persisting %d routing entries unfiltered: %s", + len(data), + exc, + ) + return data, {} + + dead = { + key: sid + for sid, key in persisted_ids.items() + if sid not in existing + } + if not dead: + return data, {} + filtered = {k: v for k, v in data.items() if k not in dead} + logger.info( + "gateway.session: backstop dropped %d dead persisted entr%s " + "from routing save (%d remaining)", + len(dead), + "y" if len(dead) == 1 else "ies", + len(filtered), + ) + return filtered, dead + + @staticmethod + def _query_existing_session_ids( + conn: Any, session_ids: List[str] + ) -> set: + """Return the subset of *session_ids* that still have rows. + + Runs ``SELECT id FROM sessions WHERE id IN (...)`` in chunks of 500 + ids to stay well under SQLite's variable-number limit. *conn* is + expected to use ``sqlite3.Row`` (both SessionDB read/write paths do). + """ + existing: set = set() + for i in range(0, len(session_ids), 500): + chunk = session_ids[i : i + 500] + placeholders = ",".join("?" * len(chunk)) + rows = conn.execute( + f"SELECT id FROM sessions WHERE id IN ({placeholders})", chunk + ).fetchall() + for row in rows: + try: + existing.add(row["id"]) + except (KeyError, IndexError, TypeError): + existing.add(row[0]) + return existing + def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Optional[str]: """Return the profile namespace for session keys, or None when off. diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index 7b87e435055c..54648b67fa84 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -2559,6 +2559,508 @@ def _resolve_prompt(value): available = "`none`, " + ", ".join(f"`{n}`" for n in personalities) return t("gateway.personality.unknown", name=args, available=available) + def _find_workspace_name_from_context( + self, platform: str, chat_id: str, thread_id: str | None + ) -> str | None: + """Resolve the workspace NAME from a platform/chat/thread context. + + Reads topics.yaml to find the mapping, then returns the workspace name + (not the resolved prompt/skills). Returns None if no workspace is linked. + """ + try: + from pathlib import Path + from agent.workspace_resolver import _safe_dir_name + + topics_path = ( + Path(_hermes_home) + / "platforms" + / platform + / _safe_dir_name(chat_id) + / "topics.yaml" + ) + if not topics_path.exists(): + return None + + import yaml as _yaml + data = _yaml.safe_load(topics_path.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + if isinstance(topics, list): + topics = { + str(item.get("thread_id", "")): str(item.get("workspace", "")) + for item in topics + if isinstance(item, dict) + } + thread_key = str(thread_id) if thread_id else "default" + return topics.get(thread_key) + except Exception: + return None + + async def _handle_workspace_command(self, event: MessageEvent) -> str: + """Handle /workspace command for workspace management.""" + from agent.workspace_resolver import ( + resolve_workspace, + _safe_dir_name, + _clear_stat_cache, + _parse_frontmatter, + _write_system_md, + ) + + raw_args = event.get_command_args().strip() + parts = raw_args.split(None, 1) + subcmd = parts[0].lower() if parts else "" + rest = parts[1].strip() if len(parts) > 1 else "" + + hermes_home = _hermes_home + workspaces_dir = hermes_home / "workspaces" + platforms_dir = hermes_home / "platforms" + + # Determine current platform/chat_id/thread_id for context-aware ops + source = event.source + platform = source.platform.value if source.platform else "" + chat_id = str(source.chat_id) if source.chat_id else "" + thread_id = source.thread_id or "" + + if not subcmd or subcmd == "list": + # /workspace list — list all workspaces and current mapping + lines = ["📁 **Workspaces**\n"] + + if not workspaces_dir.exists(): + lines.append(" No workspaces found.") + lines.append(" Use `/workspace create ` to create one.") + else: + ws_dirs = sorted(d.name for d in workspaces_dir.iterdir() if d.is_dir()) + if not ws_dirs: + lines.append(" No workspaces found.") + else: + for ws_name in ws_dirs: + system_file = workspaces_dir / ws_name / "SYSTEM.md" + skills_dir = workspaces_dir / ws_name / "skills" + has_prompt = system_file.exists() + has_skills = skills_dir.exists() and any(skills_dir.iterdir()) + # Check for model in frontmatter + has_model = False + model_name = "" + if has_prompt: + content = system_file.read_text(encoding="utf-8") + fm, _ = _parse_frontmatter(content) + model_name = fm.get("model", "") + has_model = bool(model_name) + flags = [] + if has_prompt: + flags.append("prompt") + if has_model: + flags.append(f"model: {model_name}") + if has_skills: + n_skills = len([d for d in skills_dir.iterdir() if d.is_dir()]) + flags.append(f"{n_skills} skill{'s' if n_skills != 1 else ''}") + desc = f" ({', '.join(flags)})" if flags else "" + lines.append(f" • **{ws_name}**{desc}") + + # Show current context mapping + if platform and chat_id: + safe_id = _safe_dir_name(chat_id) + topics_file = platforms_dir / platform / safe_id / "topics.yaml" + if topics_file.exists(): + import yaml as _yaml + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + ws_default = data.get("workspace", "") + if topics or ws_default: + lines.append(f"\n🔗 **{platform}/{chat_id}**") + if ws_default: + lines.append(f" default → **{ws_default}**") + for tid, ws in topics.items(): + marker = " ◀ you" if str(tid) == str(thread_id) else "" + lines.append(f" topic {tid} → **{ws}**{marker}") + else: + lines.append(f"\n🔗 **{platform}/{chat_id}** — no topics linked yet") + lines.append(f" Use `/workspace link {platform} {chat_id} `") + + # Show active workspace for current topic + if platform and chat_id: + result = resolve_workspace(platform, chat_id, thread_id or None) + if result.prompt or result.skills or result.model: + lines.append(f"\n✅ **Active workspace** for this topic:") + if result.model: + lines.append(f" Model: **{result.model}**") + if result.prompt: + preview = result.prompt[:80] + ("..." if len(result.prompt) > 80 else "") + lines.append(f" Prompt: {preview}") + if result.skills: + lines.append(f" Skills: {', '.join(result.skills)}") + + return "\n".join(lines) + + elif subcmd == "create": + # /workspace create [prompt text] + # Auto-links to current topic if in one. + # When prompt text IS provided → create immediately (batch/script-friendly). + # When prompt text is NOT provided → interactive flow (if platform supports it). + # When NO name or minimal args → auto-generate name from context. + if not rest: + # Auto-generate workspace name from session/topic context + import re as _re + suggest_name = "" + # 1. Try session title from DB + try: + session_key = self._session_key_for_source(source) if source else None + if session_key and getattr(self, "async_session_store", None): + entry = self.async_session_store._entries.get(session_key) + if entry: + title = await self.async_session_store.get_session_title(entry.session_id) + if title: + suggest_name = title + except Exception: + pass + # 2. Fall back to chat_topic from source + if not suggest_name: + ct = getattr(source, "chat_topic", None) if source else None + if ct: + suggest_name = ct + # 3. Fall back to workspace prompt first words + if not suggest_name: + try: + from agent.workspace_resolver import resolve_workspace + ws = resolve_workspace(platform, chat_id, thread_id or None) + if ws.prompt: + words = ws.prompt.strip().split() + suggest_name = " ".join(words[:3]) if words else "" + except Exception: + pass + # 4. Final fallback + if not suggest_name: + suggest_name = f"topic-{thread_id}" if thread_id else "workspace" + # Slugify: lowercase, replace spaces/separators with hyphens, strip non-alnum + ws_name = _re.sub(r'[^a-zA-Z0-9_-]', '', suggest_name.lower().replace(" ", "-").replace(".", "-"))[:50] + if not ws_name or not ws_name.replace("-", "").replace("_", "").isalnum(): + ws_name = "workspace" + prompt_text = "" + else: + create_parts = rest.split(None, 1) + ws_name = create_parts[0] + prompt_text = create_parts[1].strip() if len(create_parts) > 1 else "" + + # Validate name + if not ws_name.replace("-", "").replace("_", "").isalnum(): + return f"❌ Invalid workspace name `{ws_name}`. Use letters, numbers, hyphens, underscores." + + ws_dir = workspaces_dir / ws_name + if ws_dir.exists(): + return f"❌ Workspace `{ws_name}` already exists. Use `/workspace show {ws_name}` to view it." + + # Resolve current model vs default for model option + current_model = "" + config_default = "" + try: + resolved_model, _ = self._resolve_session_agent_runtime( + session_key=self._session_key_for_source(source) if source else None, + user_config=_load_gateway_config(), + ) + config_default = _resolve_gateway_model(_load_gateway_config()) + current_model = resolved_model or "" + except Exception: + pass + + # --- Interactive flow (no prompt provided, platform supports it) --- + if not prompt_text: + adapter = self.adapters.get(source.platform) + has_picker = ( + adapter is not None + and getattr(type(adapter), "send_workspace_create_picker", None) is not None + ) + + if has_picker: + # Build workspace creation callback closure + _self = self + _ws_name = ws_name + _platform = platform + _chat_id = chat_id + _thread_id = thread_id + _workspaces_dir = workspaces_dir + _platforms_dir = platforms_dir + + # --- Prompt suggestion: current workspace/channel prompt --- + suggest_prompt = "" + try: + # Use current event's channel_prompt (set by workspace resolution) + cp = getattr(event, "channel_prompt", None) + if cp: + suggest_prompt = cp.strip() + except Exception: + pass + + async def _on_workspace_create( + picker_chat_id: str, + ws_name: str = _ws_name, + model: str | None = None, + prompt: str | None = None, + ) -> str: + """Create the workspace and return confirmation text.""" + import yaml as _yaml + + ws_dir = _workspaces_dir / ws_name + ws_dir.mkdir(parents=True, exist_ok=True) + system_file = ws_dir / "SYSTEM.md" + + # Build SYSTEM.md with frontmatter + frontmatter = {} + if model: + frontmatter["model"] = model + body = prompt or f"# Workspace: {ws_name}\n" + + _write_system_md(system_file, frontmatter, body) + _clear_stat_cache() + + # Auto-link to current topic if in one + link_msg = "" + if _platform and _chat_id and _thread_id and str(_thread_id) != "1": + safe_id = _safe_dir_name(_chat_id) + topic_dir = _platforms_dir / _platform / safe_id + topic_dir.mkdir(parents=True, exist_ok=True) + topics_file = topic_dir / "topics.yaml" + data = {} + if topics_file.exists(): + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + topics[str(_thread_id)] = ws_name + data["topics"] = topics + if "workspace" not in data: + data["workspace"] = "default" + topics_file.write_text(_yaml.safe_dump(data, default_flow_style=False), encoding="utf-8") + _clear_stat_cache() + link_msg = f"\n🔗 Auto-linked to {_platform}/{_chat_id} topic **{_thread_id}**" + elif _platform and _chat_id and str(_thread_id) == "1": + link_msg = f"\n⚠️ Skipped auto-link — you're in the **general topic** (thread 1).\nUse `/workspace link {_platform} {_chat_id} default {ws_name}` to set a channel default." + + lines = [f"✅ Created workspace **{ws_name}**"] + lines.append(f"📄 `{system_file}`") + if model: + lines.append(f"🌐 Model: **{model}**") + if prompt: + preview = prompt[:80] + ("..." if len(prompt) > 80 else "") + lines.append(f"✏️ Prompt: {preview}") + lines.append(link_msg) + return "\n".join(lines) + + metadata = self._thread_metadata_for_source(source, self._reply_anchor_for_event(event)) + result = await adapter.send_workspace_create_picker( + chat_id=source.chat_id, + ws_name=ws_name, + current_model=current_model, + config_default_model=config_default, + on_create_workspace=_on_workspace_create, + suggested_prompt=suggest_prompt, + metadata=metadata, + owner_user_id=getattr(source, "user_id", None), + thread_id=getattr(source, "thread_id", None), + ) + if result.success: + return None # Picker sent — adapter handles the response + + # Fallback: no interactive support → create with defaults + # (current behavior — creates with no model, no prompt) + + # --- Direct creation (prompt provided or platform without interactive support) --- + model_hint = "" + if current_model and config_default and current_model != config_default: + model_hint = ( + f"\n\n💡 You're currently using model **{current_model}** (default is **{config_default}**).\n" + f"To set this as the workspace model, add `model: {current_model}` to the frontmatter.\n" + f"Or run: `/workspace model {ws_name} {current_model}`" + ) + + ws_dir.mkdir(parents=True, exist_ok=True) + system_file = ws_dir / "SYSTEM.md" + system_file.write_text(prompt_text + "\n" if prompt_text else f"# Workspace: {ws_name}\n", encoding="utf-8") + _clear_stat_cache() + + # Auto-link to current topic if in one (skip Telegram general topic thread_id=1) + link_msg = "" + if platform and chat_id and thread_id and str(thread_id) != "1": + safe_id = _safe_dir_name(chat_id) + topic_dir = platforms_dir / platform / safe_id + topic_dir.mkdir(parents=True, exist_ok=True) + topics_file = topic_dir / "topics.yaml" + import yaml as _yaml + data = {} + if topics_file.exists(): + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + topics[str(thread_id)] = ws_name + data["topics"] = topics + if "workspace" not in data: + data["workspace"] = "default" + topics_file.write_text(_yaml.safe_dump(data, default_flow_style=False), encoding="utf-8") + _clear_stat_cache() + link_msg = f"\n🔗 Auto-linked to {platform}/{chat_id} topic **{thread_id}**" + elif platform and chat_id and str(thread_id) == "1": + link_msg = f"\n⚠️ Skipped auto-link — you're in the **general topic** (thread 1).\nUse `/workspace link {platform} {chat_id} default {ws_name}` to set a channel default." + + return f"✅ Created workspace **{ws_name}**{link_msg}\nEdit: `{system_file}`{model_hint}" + + elif subcmd == "show": + # /workspace show + ws_name = rest.strip() + if not ws_name: + # Show current resolved workspace + if platform and chat_id: + result = resolve_workspace(platform, chat_id, thread_id or None) + if result.prompt or result.skills or result.model: + lines = ["📄 **Current workspace context:**\n"] + if result.model: + lines.append(f"**Model:** {result.model}") + if result.prompt: + lines.append(f"**Prompt:**\n```\n{result.prompt}\n```") + if result.skills: + lines.append(f"**Skills:** {', '.join(result.skills)}") + return "\n".join(lines) + return "Usage: `/workspace show ` or use in a topic to see active workspace." + + ws_dir = workspaces_dir / ws_name + if not ws_dir.exists(): + return f"❌ Workspace `{ws_name}` not found." + + lines = [f"📄 **Workspace: {ws_name}**\n"] + system_file = ws_dir / "SYSTEM.md" + if system_file.exists(): + lines.append(f"**SYSTEM.md:**\n```\n{system_file.read_text(encoding='utf-8')}\n```") + skills_dir = ws_dir / "skills" + if skills_dir.exists(): + skill_list = [d.name for d in skills_dir.iterdir() if d.is_dir()] + if skill_list: + lines.append(f"**Skills:** {', '.join(skill_list)}") + return "\n".join(lines) + + elif subcmd == "link": + # /workspace link ← current topic + # /workspace link ← current chat, specific topic + # /workspace link ← full spec + # /workspace link default ← channel default + link_parts = rest.split() + if not link_parts: + return ( + "📁 **Link a workspace:**\n\n" + "• `/workspace link ` — Link current topic\n" + "• `/workspace link ` — Link a topic in this chat\n" + "• `/workspace link ` — Full spec\n\n" + "Example: `/workspace link news-feed` (from inside a topic)" + ) + + # Determine how many args to infer from current context + if len(link_parts) == 1: + # /workspace link — use current context + if not platform or not chat_id: + return "❌ Can't determine current chat. Use the full form: `/workspace link `" + link_ws = link_parts[0] + if not thread_id: + return f"❌ No topic/thread in current context for {platform}/{chat_id}. Use `/workspace link {platform} {chat_id} {link_ws}`" + # Telegram general topic has thread_id == 1 — not a real topic + return ( + f"❌ You're in the **general topic** (thread 1), which can't be linked like a regular topic.\n" + f"Use channel default instead:\n" + f"`/workspace link {platform} {chat_id} default {link_ws}`" + ) + link_platform, link_chat_id, link_thread_id = platform, chat_id, thread_id + elif len(link_parts) == 2: + # /workspace link — current platform+chat + if not platform or not chat_id: + return "❌ Can't determine current chat. Use the full form: `/workspace link `" + link_thread_id, link_ws = link_parts + link_platform, link_chat_id = platform, chat_id + elif len(link_parts) >= 4: + # Full form: /workspace link + link_platform, link_chat_id, link_thread_id, link_ws = link_parts[:4] + else: + # 3 args — ambiguous, could be + # but also (default). Prefer explicit. + return ( + "❌ Ambiguous args. Use one of:\n" + "• `/workspace link ` — current topic\n" + "• `/workspace link ` — this chat, specific topic\n" + "• `/workspace link ` — full spec" + ) + + # Verify workspace exists + ws_dir = workspaces_dir / link_ws + if not ws_dir.exists(): + return f"❌ Workspace `{link_ws}` not found. Create it first with `/workspace create {link_ws}`" + + # Write/update topics.yaml + safe_id = _safe_dir_name(link_chat_id) + topic_dir = platforms_dir / link_platform / safe_id + topic_dir.mkdir(parents=True, exist_ok=True) + topics_file = topic_dir / "topics.yaml" + + import yaml as _yaml + data = {} + if topics_file.exists(): + data = _yaml.safe_load(topics_file.read_text(encoding="utf-8")) or {} + topics = data.get("topics", {}) + if isinstance(topics, list): + # Convert list form to dict form + topics = {str(item.get("thread_id", "")): str(item.get("workspace", "")) for item in topics if isinstance(item, dict)} + topics[str(link_thread_id)] = link_ws + data["topics"] = topics + if "workspace" not in data: + data["workspace"] = "default" + topics_file.write_text(_yaml.safe_dump(data, default_flow_style=False), encoding="utf-8") + _clear_stat_cache() + + if link_thread_id == "default": + return f"✅ Set **default workspace** for {link_platform}/{link_chat_id} → **{link_ws}**" + return f"✅ Linked {link_platform}/{link_chat_id} topic **{link_thread_id}** → **{link_ws}**" + + elif subcmd == "model": + # /workspace model [name] [model_name] — set, show, or clear workspace model + # No ws name → auto-detect from current topic context + # No model_name → auto-detect current session model + if not rest: + # Try to resolve workspace from current context + if platform and chat_id: + result = resolve_workspace(platform, chat_id, thread_id or None) + if result.prompt is not None: + # Find workspace name from topics.yaml + ws_name = self._find_workspace_name_from_context(platform, chat_id, thread_id or None) + if ws_name: + rest = ws_name # pretend user typed the workspace name + else: + return "❌ Can't determine workspace name from current topic. Use: `/workspace model `" + else: + return "Usage: `/workspace model [model_name]`\nOmit model to use current session model\nUse `clear` to remove: `/workspace model clear`" + else: + return "Usage: `/workspace model [model_name]`\nOmit model to use current session model\nUse `clear` to remove: `/workspace model clear`" + + model_parts = rest.split(None, 1) + ws_name = model_parts[0] + second_arg = model_parts[1].strip() if len(model_parts) > 1 else "" + + ws_dir = workspaces_dir / ws_name + if not ws_dir.exists(): + # First arg isn't a workspace — maybe it's a model name for current workspace? + if platform and chat_id: + current_ws = self._find_workspace_name_from_context(platform, chat_id, thread_id or None) + if current_ws: + # Treat the first arg as model name for the current workspace + current_ws_dir = workspaces_dir / current_ws + if current_ws_dir.exists(): + # Re-parse: ws_name=current_ws, model_value=original first arg + second_arg = ws_name # original first arg becomes the model + ws_name = current_ws + ws_dir = current_ws_dir + if not ws_dir.exists(): + return f"❌ Workspace `{ws_name}` not found. Create it first with `/workspace create {ws_name}`" + + system_file = ws_dir / "SYSTEM.md" + import yaml as _yaml + if system_file.exists(): + content = system_file.read_text(encoding="utf-8") + fm, body = _parse_frontmatter(content) + else: + fm, body = {}, "" + + # Handle "clear" + + async def _handle_retry_command(self, event: MessageEvent) -> str: """Handle /retry command - re-send the last user message.""" source = event.source @@ -2595,6 +3097,7 @@ async def _handle_retry_command(self, event: MessageEvent) -> str: source=source, raw_message=event.raw_message, channel_prompt=event.channel_prompt, + workspace_model=getattr(event, 'workspace_model', None), ) # Let the normal message handler process it diff --git a/hermes_cli/commands.py b/hermes_cli/commands.py index 34280abc2464..9fec8f2fd385 100644 --- a/hermes_cli/commands.py +++ b/hermes_cli/commands.py @@ -239,6 +239,9 @@ class CommandDef: subcommands=("queue", "steer", "interrupt", "status")), # Tools & Skills + CommandDef("workspace", "Manage workspaces: create, link, model, list, show, or remove", + "Configuration", args_hint="[create|link|model|list|show|remove] [args]", + subcommands=("create", "link", "model", "list", "show", "remove", "unlink")), CommandDef("tools", "Manage tools: /tools [list|disable|enable] [name...]", "Tools & Skills", args_hint="[list|disable|enable] [name...]", cli_only=True), CommandDef("toolsets", "List available toolsets", "Tools & Skills", diff --git a/hermes_cli/web_server.py b/hermes_cli/web_server.py index 1fb3e6131629..60d2db2e807c 100644 --- a/hermes_cli/web_server.py +++ b/hermes_cli/web_server.py @@ -11306,6 +11306,50 @@ def _sweep() -> None: +# Architecture note — multi-process delete contract (serve has no SessionStore): +# This process never hosts a SessionStore (grep SessionStore in this file: none); +# it writes state.db directly via SessionDB and only the gateway process +# (``hermes gateway run``, gateway/session.py) keeps in-memory routing entries +# keyed by messaging session_key. A desktop DELETE therefore cannot resurrect +# anything by itself — the only resurrection risk is the gateway re-persisting +# an entry it still holds. That case is covered by the #GAP-3 multi-process +# backstop inside SessionStore: every save runs _prune_dead_persisted_entries +# (gateway/session.py:_save), which drops any entry whose state.db row is gone +# (db_persisted=True + row missing) from the snapshot and from _entries before +# persisting, so the next gateway save never re-creates the deleted session_id +# in gateway_routing or the sessions.json mirror. Routing cleanup here is +# synchronous and transactional: SessionDB.delete_session purges gateway_routing +# rows whose embedded session_id matches inside the same write transaction +# (_purge_gateway_routing_for_sessions, hermes_state.py), so no RPC to the +# gateway (session.close etc.) is needed — the staleness detection is +# data-driven over the shared state.db. +@app.delete("/api/sessions/{session_id}") +async def delete_session_endpoint(session_id: str, profile: Optional[str] = None): + # ``profile`` deletes a session belonging to another (local) profile by + # opening its state.db directly. Remote profiles never reach here — the + # desktop routes their DELETE to the remote backend. Omit for current/default. + def _delete(): + db = _open_session_db_for_profile(profile) + try: + # Resolve exact ids / unique prefixes like every other session endpoint + # (detail, messages, rename, export all do). A session that no longer + # exists is an idempotent success: DELETE's contract is "ensure it's + # gone", and the desktop optimistically removes the row then RESTORES it + # on any error — so a 404 on an already-absent row resurrected a ghost + # row and surfaced "session not found". /goal + auto-compression churn + # leaves transient empty rows (reaped by empty-session hygiene) that + # race the sidebar snapshot, which is exactly when this fired. Mirrors + # the bulk-delete endpoint, which already treats ghost ids as success. + sid = db.resolve_session_id(session_id) + if not sid: + return {"ok": True, "already_absent": True} + db.delete_session(sid, sessions_dir=_profile_sessions_dir(profile)) + return {"ok": True} + finally: + db.close() + + return await asyncio.to_thread(_delete) + diff --git a/hermes_state.py b/hermes_state.py index 22298e59f1c7..21699f6f7983 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -19,6 +19,7 @@ import errno import hashlib import json +import tempfile import logging import os import random @@ -44,8 +45,10 @@ from hermes_cli.sqlite_runtime import ( is_sqlite_wal_reset_vulnerable as _is_sqlite_wal_reset_vulnerable, ) + from typing import Any, Callable, Dict, Iterable, List, Optional, Set, Tuple, TypeVar + from hermes_state_common import ( # noqa: F401 (re-exported for back-compat) _BRANCH_CHILD_SQL, _COMPRESSION_CHILD_SQL, @@ -240,6 +243,91 @@ def _delete_delegate_children(conn, parent_ids: List[str]) -> List[str]: conn.execute(f"DELETE FROM sessions WHERE id IN ({ph})", ids) return ids + +def _cleanup_async_delegations( + conn: sqlite3.Connection, + session_ids: Iterable[str], + session_keys: Iterable[str], +) -> int: + """Delete durable async-delegation rows referencing deleted sessions. + + ``async_delegations`` has no FK to ``sessions`` (schema in + ``hermes_state_common.py``; table is also created lazily by + ``tools/async_delegation.py``). A pending/undelivered completion left + behind after a delete is re-hydrated on the next boot by + ``restore_undelivered_completions``, which re-creates the deleted + session to deliver the delegation's result (N3). The session may be + referenced by any of the id-typed columns (``origin_ui_session_id``, + ``parent_session_id``, ``origin_session_id`` — the raw api_server + session id added in a later schema revision) or by ``origin_session`` + (the routing *key* of the origin session). Runs on the in-flight + write transaction's connection; missing table/column on legacy DBs is + tolerated (best-effort). Returns the number of rows deleted. + """ + ids = [sid for sid in session_ids if sid] + keys = [k for k in session_keys if k] + if not ids and not keys: + return 0 + try: + conn.execute("SELECT 1 FROM async_delegations LIMIT 1") + except sqlite3.DatabaseError: + return 0 + deleted = 0 + for col in ("origin_ui_session_id", "parent_session_id", "origin_session_id"): + if not ids: + break + ph = ",".join("?" * len(ids)) + try: + cur = conn.execute( + f"DELETE FROM async_delegations WHERE {col} IN ({ph})", ids + ) + deleted += cur.rowcount + except sqlite3.DatabaseError: + pass # column absent in legacy schema — nothing to clean there + if keys: + ph = ",".join("?" * len(keys)) + try: + cur = conn.execute( + f"DELETE FROM async_delegations WHERE origin_session IN ({ph})", + keys, + ) + deleted += cur.rowcount + except sqlite3.DatabaseError: + pass + return deleted + + +def _cleanup_delivery_obligations( + conn: sqlite3.Connection, session_keys: Iterable[str] +) -> int: + """Delete delivery-obligation rows for deleted sessions. + + ``delivery_obligations`` (``gateway/delivery_ledger.py`` — same + ``state.db``, no FK, table created lazily by the ledger module) + records outbound final responses keyed by ``session_key``. A row left + behind after a delete is claimed/redelivered on the next gateway + restart, delivering the deleted conversation's content into whatever + session now owns the key (N5). Runs on the in-flight write + transaction's connection; a missing table is tolerated (best-effort). + Returns the number of rows deleted. + """ + keys = [k for k in session_keys if k] + if not keys: + return 0 + try: + conn.execute("SELECT 1 FROM delivery_obligations LIMIT 1") + except sqlite3.DatabaseError: + return 0 + ph = ",".join("?" * len(keys)) + try: + cur = conn.execute( + f"DELETE FROM delivery_obligations WHERE session_key IN ({ph})", keys + ) + return cur.rowcount + except sqlite3.DatabaseError: + return 0 + + T = TypeVar("T") DEFAULT_DB_PATH = get_hermes_home() / "state.db" @@ -2966,6 +3054,22 @@ def _insert_session_row( without a recoverable routing mapping (#59527). """ def _do(conn): + # N4 visibility: the ON CONFLICT(id) DO UPDATE below silently + # merges an existing row. When the existing row was created by a + # different origin than this call, log a warning so a delete + + # re-create of the same id (session resurrection) or a + # source-mismatched merge becomes visible. (sqlite3 rowcount + # cannot distinguish insert vs update for ON CONFLICT DO UPDATE — + # both report 1 — so the pre-check reads the existing source.) + existing = conn.execute( + "SELECT source FROM sessions WHERE id = ?", (session_id,) + ).fetchone() + if existing is not None and existing["source"] != source: + logger.warning( + "session upsert merged an existing row: session_id=%s " + "incoming source=%r over existing source=%r", + session_id, source, existing["source"], + ) system_prompt_hash = self._store_system_prompt(conn, system_prompt) conn.execute( """INSERT INTO sessions ( @@ -4815,11 +4919,12 @@ def update_token_counts( the caller already holds cumulative totals (gateway path, where the cached agent accumulates across messages). """ - # Ensure the session row exists so the UPDATE doesn't silently affect - # 0 rows. Under concurrent load (cron + kanban + delegate_task) the - # initial create_session() may have failed due to SQLite locking. - # INSERT OR IGNORE is cheap and idempotent. - self._insert_session_row(session_id, "unknown", model=model) + # Sessions are born in create_session(); token accounting must NEVER + # create a row. A missing row here means the session was deleted (or + # never created) — recreating it would resurrect a deleted session + # with end_reason=NULL, defeating the GAP-3 backstop (N1). The + # UPDATE below simply affects 0 rows in that case and the per-model + # usage insert is skipped inside the transaction. if absolute: sql = """UPDATE sessions SET input_tokens = ?, @@ -4910,6 +5015,17 @@ def _do(conn): "SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,), ).fetchone() + if row is None: + # Session deleted (or never created) — never recreate it from + # token accounting (N1). Queued deltas that were enqueued + # before a delete land here after it; dropping them is the + # intended behavior, not an error. + logger.info( + "token accounting skipped for non-existent session %s " + "(deleted or never created)", + session_id, + ) + return existing_model = row["model"] if row is not None else None existing_provider = row["billing_provider"] if row is not None else None existing_api_calls = int((row["api_call_count"] if row is not None else 0) or 0) @@ -4994,6 +5110,16 @@ def _record_model_usage( "FROM sessions WHERE id = ?", (session_id,), ).fetchone() + if row is None: + # Session deleted (or never created) — never recreate it from + # usage accounting (N1); drop the delta instead of inserting an + # orphan session_model_usage row or failing an FK. + logger.info( + "model usage skipped for non-existent session %s " + "(deleted or never created)", + session_id, + ) + return sess_model = row["model"] if row is not None else None sess_provider = row["billing_provider"] if row is not None else None sess_base_url = row["billing_base_url"] if row is not None else None @@ -5100,10 +5226,13 @@ def record_auxiliary_usage( """ if not session_id or not task: return - # FK on session_model_usage.session_id → sessions.id: ensure the row - # exists (same INSERT OR IGNORE guard update_token_counts uses — the - # initial create_session() can fail under concurrent SQLite locking). - self._insert_session_row(session_id, "unknown") + # session_model_usage has an FK-style relationship to sessions.id + # (session_model_usage rows are meaningless without the session). We + # deliberately do NOT create the session row here: a missing row + # means the session was deleted, and recreating it would resurrect a + # deleted session (N1). _record_model_usage no-ops when the row is + # gone, so the accounting is simply dropped — matching the + # best-effort contract below. def _do(conn): self._record_model_usage( @@ -5144,12 +5273,14 @@ def _do(conn): """, (cutoff,)).fetchall() ids = [r[0] if isinstance(r, (tuple, list)) else r["id"] for r in rows] if ids: - placeholders = ",".join("?" * len(ids)) + placeholders = ", ".join("?" * len(ids)) conn.execute( f"DELETE FROM sessions WHERE id IN ({placeholders})", ids ) self._purge_gateway_routing_for_sessions(conn, ids) + self._delete_unreferenced_system_prompts(conn) + return ids removed_ids = self._execute_write(_do) or [] @@ -8001,11 +8132,14 @@ def delete_session( ) def _do(conn): - cursor = conn.execute( - "SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,) - ) - if cursor.fetchone() is None: + + row = conn.execute( + "SELECT session_key FROM sessions WHERE id = ?", (session_id,) + ).fetchone() + if row is None: + return False + session_key = row["session_key"] if row["session_key"] is not None else "" if expected_ids is not None: actual_ids = { session_id, @@ -8022,10 +8156,24 @@ def _do(conn): ) conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,)) conn.execute("DELETE FROM sessions WHERE id = ?", (session_id,)) + + # No-FK side tables must die with the session, in the same + # transaction: a pending async delegation would otherwise be + # re-hydrated on boot (restore_undelivered_completions) and an + # undelivered delivery obligation redelivered on restart — both + # resurrecting the deleted session (N3/N5). Both helpers are + # best-effort (missing table/column tolerated). + _cleanup_async_delegations( + conn, + [session_id, *removed_delegate_ids], + [session_key] if session_key else [], + ) + _cleanup_delivery_obligations(conn, [session_key] if session_key else []) self._purge_gateway_routing_for_sessions( conn, [session_id, *removed_delegate_ids] ) self._delete_unreferenced_system_prompts(conn) + return True deleted = self._execute_write(_do) @@ -8075,7 +8223,9 @@ def _do(conn): ) if cursor.rowcount > 0: self._purge_gateway_routing_for_sessions(conn, [session_id]) + self._delete_unreferenced_system_prompts(conn) + return cursor.rowcount > 0 deleted = self._execute_write(_do) @@ -8130,10 +8280,14 @@ def _do(conn): # First, filter to IDs that actually exist — we want to # return the real deleted count, not the input length. cursor = conn.execute( - f"SELECT id FROM sessions WHERE id IN ({placeholders})", + f"SELECT id, session_key FROM sessions WHERE id IN ({placeholders})", unique_ids, ) - existing = [row["id"] for row in cursor.fetchall()] + rows = cursor.fetchall() + existing = [row["id"] for row in rows] + existing_keys = [ + row["session_key"] for row in rows if row["session_key"] is not None + ] if not existing: return 0 @@ -8157,10 +8311,18 @@ def _do(conn): f"DELETE FROM sessions WHERE id IN ({existing_placeholders})", existing, ) + + # Same-transaction cleanup of the no-FK side tables (N3/N5) — + # pending async delegations / undelivered delivery obligations + # referencing a deleted session must not survive to be restored + # on the next boot/restart. Best-effort per helper. + _cleanup_async_delegations(conn, existing + removed_delegate_ids, existing_keys) + _cleanup_delivery_obligations(conn, existing_keys) self._purge_gateway_routing_for_sessions( conn, existing + removed_delegate_ids ) self._delete_unreferenced_system_prompts(conn) + removed_ids.extend(existing) return len(existing) @@ -8260,13 +8422,16 @@ def _do(conn): conn.execute("DELETE FROM sessions WHERE id = ?", (sid,)) removed_ids.append(sid) self._purge_gateway_routing_for_sessions(conn, removed_ids) + self._delete_unreferenced_system_prompts(conn) + return len(session_ids) count = self._execute_write(_do) self._remove_sessions_json_entries(sessions_dir, removed_ids) for sid in removed_ids: self._remove_session_files(sessions_dir, sid) + self._remove_sessions_json_entries(sessions_dir, removed_ids) return count @staticmethod @@ -8598,7 +8763,9 @@ def _do(conn): conn.execute("DELETE FROM sessions WHERE id = ?", (sid,)) removed_ids.append(sid) self._purge_gateway_routing_for_sessions(conn, removed_ids) + self._delete_unreferenced_system_prompts(conn) + return len(session_ids) count = self._execute_write(_do) @@ -8606,6 +8773,7 @@ def _do(conn): # Clean up on-disk files outside the DB transaction for sid in removed_ids: self._remove_session_files(sessions_dir, sid) + self._remove_sessions_json_entries(sessions_dir, removed_ids) return count def purge_stale_tool_call_markers( diff --git a/hermes_state_schema.py b/hermes_state_schema.py index 0761a0804ef7..5f22e0fb5d72 100644 --- a/hermes_state_schema.py +++ b/hermes_state_schema.py @@ -369,12 +369,24 @@ def _reconcile_columns(self, cursor: sqlite3.Cursor) -> None: ) except sqlite3.OperationalError as exc: # Expected: "duplicate column name" from a race or - # re-run. Unexpected: "Cannot add a NOT NULL column - # with default value NULL" from a schema mistake. - # Log at DEBUG so it's visible in agent.log. - logger.debug( - "reconcile %s.%s: %s", table_name, col_name, exc, - ) + # re-run — the column already exists, so the ADD is + # a no-op and idempotence is preserved. Keep this at + # DEBUG so normal startups stay quiet. + msg = str(exc).lower() + if "duplicate column" in msg or "already exists" in msg: + logger.debug( + "reconcile %s.%s: %s", table_name, col_name, exc, + ) + else: + # Unexpected: a real schema mistake (e.g. "Cannot + # add a NOT NULL column with default value NULL" + # from a bad SCHEMA_SQL entry). This must never be + # swallowed silently — log ERROR with the column + # and the full traceback. + logger.exception( + "reconcile %s.%s: failed to add column (type=%s): %s", + table_name, col_name, col_type, exc, + ) def _heal_gateway_routing_pk(self, cursor: sqlite3.Cursor) -> None: """Rebuild ``gateway_routing`` when its PRIMARY KEY predates scoping. diff --git a/plugins/platforms/telegram/adapter.py b/plugins/platforms/telegram/adapter.py index ea6258b9f2dc..e8104fa6aa44 100644 --- a/plugins/platforms/telegram/adapter.py +++ b/plugins/platforms/telegram/adapter.py @@ -860,6 +860,11 @@ def __init__(self, config: PlatformConfig): # Clarify button state: clarify_id → session_key (for the clarify tool's # multiple-choice prompts; see GatewayRunner clarify_callback wiring). self._clarify_state: Dict[str, str] = {} + # Workspace creation interactive state. + # Key = "::" (review fix: capture is + # scoped to the initiating user + topic, so another user/topic can + # neither hijack nor delete the pending workspace name/prompt). + self._workspace_create_state: Dict[str, dict] = {} # Notification mode for message sends. # "important" — only final responses, approvals, and slash confirmations # trigger notifications; tool progress, streaming, status @@ -5840,6 +5845,278 @@ async def _handle_choice_picker_callback( await query.answer() self._choice_picker_state.pop(chat_id, None) + # ── Workspace creation (interactive picker) ─────────────────────────── + + def _workspace_state_key(self, chat_id: str, user_id: str | None, thread_id: str | None) -> str: + """Composite state key scoping capture to user + topic (review fix).""" + return f"{chat_id}:{user_id or '?'}:{thread_id or ''}" + + def _find_workspace_state(self, chat_id: str, user_id: str | None, thread_id: str | None): + """Return the pending workspace-create state for this chat/user/topic.""" + store = getattr(self, "_workspace_create_state", None) + if store is None: + return None + return store.get( + self._workspace_state_key(chat_id, user_id, thread_id) + ) + + def _pop_workspace_state(self, chat_id: str, user_id: str | None, thread_id: str | None): + store = getattr(self, "_workspace_create_state", None) + if store is None: + return None + return store.pop( + self._workspace_state_key(chat_id, user_id, thread_id), None + ) + + async def send_workspace_create_picker( + self, + chat_id: str, + ws_name: str, + current_model: str, + config_default_model: str, + on_create_workspace, + suggested_prompt: str = "", + metadata: Optional[Dict[str, Any]] = None, + owner_user_id: Optional[str] = None, + thread_id: Optional[str] = None, + ) -> SendResult: + """Send an interactive inline-keyboard workspace creator. + + Prompts the user to configure name, model, and prompt before creating. + Supports toggling the workspace model (if current != default), + editing the workspace name, and entering text-capture mode for + a custom prompt. Accepts a suggested prompt for pre-fill. + + Capture state is keyed by chat + owner user + topic thread, and the + picker message id is validated on every callback, so a different user + or topic cannot hijack the flow (review fix). + """ + if not self._bot: + return SendResult(success=False, error="Not connected") + + if metadata is not None and thread_id is None: + thread_id = metadata.get("thread_id") + state_key = self._workspace_state_key(chat_id, owner_user_id, thread_id) + if not hasattr(self, "_workspace_create_state"): + self._workspace_create_state = {} + + # Build initial state + has_model_diff = current_model and config_default_model and current_model != config_default_model + state = { + "ws_name": ws_name, + "current_model": current_model, + "config_default_model": config_default_model, + "include_model": has_model_diff, # pre-check if user already switched away from default + "has_model_diff": has_model_diff, + "prompt": suggested_prompt, + "awaiting_prompt": False, + "awaiting_name": False, + "on_create": on_create_workspace, + "owner_user_id": owner_user_id, + "thread_id": str(thread_id) if thread_id is not None else None, + } + self._workspace_create_state[state_key] = state + + text, keyboard = self._build_workspace_create_message(state) + try: + reply_to_id = self._reply_to_message_id_for_send(None, metadata) + msg = await self._send_message_with_thread_fallback( + chat_id=int(chat_id), + text=text, + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=keyboard, + reply_to_message_id=reply_to_id, + **self._thread_kwargs_for_send( + chat_id, thread_id, metadata, reply_to_message_id=reply_to_id, + ), + **self._link_preview_kwargs(), + ) + state["msg_id"] = msg.message_id + return SendResult(success=True, message_id=str(msg.message_id)) + except Exception as e: + self._workspace_create_state.pop(state_key, None) + logger.warning("[%s] send_workspace_create_picker failed: %s", self.name, e) + return SendResult(success=False, error=str(e)) + + def _build_workspace_create_message(self, state: dict) -> tuple: + """Build text + InlineKeyboardMarkup for the workspace creation prompt.""" + ws_name = state["ws_name"] + lines = [ + f"📁 *Create Workspace*", + "", + f"📛 Name: `{ws_name}`", + ] + + if state["has_model_diff"]: + if state["include_model"]: + lines.append(f"🌐 Model: *{state['current_model']}*") + else: + lines.append(f"🌐 Model: {state['config_default_model']} (default)") + + if state.get("awaiting_prompt"): + lines.append("✏️ Prompt: _awaiting your text..._") + elif state["prompt"]: + preview = state["prompt"][:80] + ("…" if len(state["prompt"]) > 80 else "") + lines.append(f"✏️ Prompt: {preview}") + else: + lines.append("✏️ Prompt: _(none)_") + + if state.get("awaiting_name"): + lines.append("") + lines.append("⏳ _Please type a new workspace name now..._") + + lines.append("") + lines.append("Configure your workspace:") + + text = self.format_message("\n".join(lines)) + + # Build keyboard + buttons: list = [] + + # Name edit button + if state.get("awaiting_name"): + buttons.append(InlineKeyboardButton("⌛ Awaiting name…", callback_data="wc:noop")) + else: + buttons.append(InlineKeyboardButton("📛 Edit name", callback_data="wc:n:edit")) + + if state["has_model_diff"]: + label = "✅ Use model" if state["include_model"] else "☐ Use model" + buttons.append(InlineKeyboardButton(label, callback_data="wc:m:toggle")) + + if state.get("awaiting_prompt"): + buttons.append(InlineKeyboardButton("⌛ Awaiting prompt…", callback_data="wc:noop")) + elif state["prompt"]: + buttons.append(InlineKeyboardButton("✏️ Edit prompt", callback_data="wc:p:edit")) + else: + buttons.append(InlineKeyboardButton("✏️ Add prompt", callback_data="wc:p:edit")) + + # Action row + action_buttons = [ + InlineKeyboardButton("✅ Create", callback_data="wc:c:create"), + InlineKeyboardButton("✗ Cancel", callback_data="wc:x:cancel"), + ] + + rows = [buttons[i : i + 2] for i in range(0, len(buttons), 2)] + rows.append(action_buttons) + + return text, InlineKeyboardMarkup(rows) + + async def _handle_workspace_create_callback( + self, query, data: str, chat_id: str + ) -> None: + """Handle workspace creation inline keyboard callbacks (wc:*).""" + # Resolve the pending state scoped to THIS user + topic + picker message + # (review fix: reject callbacks from other users/topics/stale pickers). + owner_user_id = str(query.from_user.id) if query.from_user else None + thread_id = getattr(getattr(query, "message", None), "message_thread_id", None) + state = self._find_workspace_state(chat_id, owner_user_id, thread_id) + if not state: + await query.answer(text="Picker expired — use /workspace create again.") + return + picker_msg_id = getattr(getattr(query, "message", None), "message_id", None) + if state.get("msg_id") is not None and str(state["msg_id"]) != str(picker_msg_id): + await query.answer(text="Picker expired — use /workspace create again.") + return + + if data == "wc:x:cancel": + self._pop_workspace_state(chat_id, owner_user_id, thread_id) + await query.edit_message_text( + text="Workspace creation cancelled.", + reply_markup=None, + ) + await query.answer() + return + + if data == "wc:noop": + await query.answer() + return + + if data == "wc:m:toggle": + state["include_model"] = not state["include_model"] + text, keyboard = self._build_workspace_create_message(state) + await query.edit_message_text( + text=text, + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=keyboard, + ) + await query.answer() + return + + if data == "wc:n:edit": + # Enter text-capture mode for name + state["awaiting_name"] = True + await query.edit_message_text( + text=self.format_message( + f"📁 *Create Workspace*\n\n" + f"📛 Please type a new name for this workspace now.\n" + f"Your next message will become the workspace name.\n\n" + f"Use letters, numbers, hyphens, and underscores only." + ), + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=None, + ) + await query.answer() + return + + if data == "wc:p:edit": + # Enter text-capture mode for prompt + state["awaiting_prompt"] = True + await query.edit_message_text( + text=self.format_message( + f"📁 *Create Workspace: `{state['ws_name']}`*\n\n" + f"✏️ Please type the prompt text for this workspace now.\n" + f"Your next message will become the SYSTEM.md body." + ), + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=None, + ) + await query.answer() + return + + if data == "wc:c:create": + # Call the creation callback + create_cb = state.get("on_create") + if not create_cb: + self._pop_workspace_state(chat_id, owner_user_id, thread_id) + await query.edit_message_text( + text="Picker expired.", + reply_markup=None, + ) + await query.answer() + return + + try: + result_text = await create_cb( + chat_id, + ws_name=state["ws_name"], + model=state["current_model"] if state["include_model"] else None, + prompt=state["prompt"] or None, + ) + except Exception as exc: + logger.error("Workspace create callback failed: %s", exc, exc_info=True) + result_text = f"Error creating workspace: {exc}" + + self._pop_workspace_state(chat_id, owner_user_id, thread_id) + try: + await query.edit_message_text( + text=result_text, + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=None, + ) + except Exception: + try: + await query.edit_message_text( + text=result_text, + parse_mode=None, + reply_markup=None, + ) + except Exception: + pass + await query.answer(text="Workspace created!") + return + + await query.answer() + _MODEL_PAGE_SIZE = 8 def _build_provider_keyboard(self, providers: list, page: int = 0) -> tuple: @@ -6329,6 +6606,13 @@ async def _handle_callback_query( await self._handle_model_picker_callback(query, data, chat_id) return + # --- Workspace creation callbacks (wc:*) --- + if data.startswith("wc:"): + chat_id = str(query.message.chat_id) if query.message else None + if chat_id: + await self._handle_workspace_create_callback(query, data, chat_id) + return + # --- Generic choice picker callbacks (/reasoning, /fast) --- if data.startswith("cp:"): chat_id = str(query.message.chat_id) if query.message else None @@ -8807,6 +9091,74 @@ async def _handle_text_message(self, update: Update, context: ContextTypes.DEFAU return await self._ensure_forum_commands(update.message) + # --- Workspace create text-capture intercept --- + # Scoped to the initiating user + topic (review fix): only the user who + # started the picker, in the same topic thread, is captured; messages + # from anyone else fall through to the normal pipeline untouched. + _ws_chat_id = str( + getattr(msg, "chat_id", None) + or (getattr(getattr(msg, "chat", None), "id", None) or "") + ) + _ws_user_id = str(msg.from_user.id) if getattr(msg, "from_user", None) else None + _ws_thread_id = getattr(msg, "message_thread_id", None) + ws_state = self._find_workspace_state(_ws_chat_id, _ws_user_id, _ws_thread_id) + if ws_state and (ws_state.get("awaiting_name") or ws_state.get("awaiting_prompt")): + if ws_state.get("awaiting_name"): + # Capture user's text as the workspace name + new_name = msg.text.strip() + # Validate name + if not new_name.replace("-", "").replace("_", "").isalnum(): + # Invalid name — show error and keep awaiting + try: + wc_msg_id = ws_state.get("msg_id") + if wc_msg_id and self._bot: + await self._bot.edit_message_text( + chat_id=msg.chat_id, + message_id=wc_msg_id, + text=self.format_message( + f"📁 *Create Workspace*\n\n" + f"❌ Invalid name `{new_name}`. Use letters, numbers, hyphens, underscores.\n" + f"Please try again:" + ), + parse_mode=ParseMode.MARKDOWN_V2, + ) + except Exception: + pass + return + ws_state["ws_name"] = new_name + ws_state["awaiting_name"] = False + + elif ws_state.get("awaiting_prompt"): + # Capture user's text as the workspace prompt + prompt_text = msg.text.strip() + ws_state["prompt"] = prompt_text + ws_state["awaiting_prompt"] = False + + # Delete the captured message for cleanliness + try: + await self._bot.delete_message( + chat_id=msg.chat_id, + message_id=msg.message_id, + ) + except Exception: + pass # non-fatal + + # Edit the workspace creation message with updated state + keyboard + text, keyboard = self._build_workspace_create_message(ws_state) + try: + wc_msg_id = ws_state.get("msg_id") + if wc_msg_id and self._bot: + await self._bot.edit_message_text( + chat_id=msg.chat_id, + message_id=wc_msg_id, + text=text, + parse_mode=ParseMode.MARKDOWN_V2, + reply_markup=keyboard, + ) + except Exception as e: + logger.warning("[%s] Failed to edit workspace create message: %s", self.name, e) + return + event = self._build_message_event(msg, MessageType.TEXT, update_id=update.update_id) event.text = self._clean_bot_trigger_text(event.text) await self._cache_replied_media(msg, event) @@ -9666,12 +10018,14 @@ def _build_message_event( thread_id_str = self._effective_message_thread_id(message) chat_topic = None topic_skill = None + topic_prompt = None if chat_type == "dm" and thread_id_str: topic_info = self._get_dm_topic_info(str(chat.id), thread_id_str) if topic_info: chat_topic = topic_info.get("name") topic_skill = topic_info.get("skill") + topic_prompt = topic_info.get("prompt") # Also check forum_topic_created service message for topic discovery if hasattr(message, "forum_topic_created") and message.forum_topic_created: @@ -9711,6 +10065,7 @@ def _build_message_event( if tid is not None and str(tid) == thread_id_str: chat_topic = topic.get("name") topic_skill = topic.get("skill") + topic_prompt = topic.get("prompt") break break @@ -9776,15 +10131,17 @@ def _build_message_event( reply_to_text = None # Per-channel/topic ephemeral prompt + # Priority: topic-level prompt from group_topics/dm_topics config + # > channel_prompts dict > None from gateway.platforms.base import resolve_channel_prompt _chat_id_str = str(chat.id) - _channel_prompt = resolve_channel_prompt( + _channel_prompt = topic_prompt or resolve_channel_prompt( self.config.extra, thread_id_str or _chat_id_str, _chat_id_str if thread_id_str else None, ) - return MessageEvent( + event = MessageEvent( text=message.text or "", message_type=msg_type, source=source, @@ -9795,9 +10152,22 @@ def _build_message_event( reply_to_text=reply_to_text, auto_skill=topic_skill, channel_prompt=_channel_prompt, + workspace_model=None, timestamp=message.date, ) + # Workspace resolution — SHARED path (review fix): folder-based + # workspaces apply for every platform via the gateway handler; the + # Telegram adapter also applies here so the event carries workspace + # context immediately. Idempotent (``_workspace_applied`` flag). + try: + from agent.workspace_resolver import apply_workspace_to_event + apply_workspace_to_event(event) + except Exception as e: + logger.debug("[%s] workspace resolution skipped: %s", self.name, e) + + return event + # ── Message reactions (processing lifecycle) ────────────────────────── def _reactions_enabled(self) -> bool: diff --git a/tests/agent/test_workspace_resolver.py b/tests/agent/test_workspace_resolver.py new file mode 100644 index 000000000000..3445afc512cd --- /dev/null +++ b/tests/agent/test_workspace_resolver.py @@ -0,0 +1,461 @@ +"""Tests for the workspace_resolver module.""" + +import os +import tempfile +from pathlib import Path +from typing import cast + +import pytest + +from agent.workspace_resolver import ( + WorkspaceResult, + _clear_stat_cache, + _resolve_workspace_content, + _resolve_workspace_name, + _safe_dir_name, + get_workspace_skill_dirs, + resolve_workspace, +) +from hermes_constants import get_hermes_home + +# ── Fixtures ────────────────────────────────────────────────────────────────── + +@pytest.fixture(autouse=True) +def _clear_cache(): + _clear_stat_cache() + yield + _clear_stat_cache() + + +@pytest.fixture +def temp_hermes_home(tmp_path: Path): + """Provide a temporary hermes home with workspaces + platforms.""" + # Save/restore real HERMES_HOME + original = os.environ.get("HERMES_HOME", "") + os.environ["HERMES_HOME"] = str(tmp_path) + yield tmp_path + if original: + os.environ["HERMES_HOME"] = original + else: + os.environ.pop("HERMES_HOME", None) + + +# ── Safe dir name ────────────────────────────────────────────────────────────── + +class TestSafeDirName: + def test_leading_minus(self): + assert _safe_dir_name("-1003") == "_-1003" + + def test_no_leading_minus(self): + assert _safe_dir_name("1003") == "1003" + + def test_empty(self): + # Review fix: empty names are rejected (would escape the parent dir) + assert _safe_dir_name("") == "_invalid" + + def test_special_chars(self): + assert _safe_dir_name("abc-123") == "abc-123" + + def test_rejects_separators(self): + assert _safe_dir_name("../evil") == "_invalid" + assert _safe_dir_name("a/b") == "_invalid" + assert _safe_dir_name("a\\b") == "_invalid" + + def test_rejects_traversal_segments(self): + assert _safe_dir_name("..") == "_invalid" + assert _safe_dir_name(".") == "_invalid" + assert _safe_dir_name("a..b") == "_invalid" + assert _safe_dir_name("a/../../etc") == "_invalid" + + def test_rejects_absolute_escape(self): + assert _safe_dir_name("/etc/passwd") == "_invalid" + assert _safe_dir_name("C:/Windows") == "_invalid" + assert _safe_dir_name("C:\\Windows") == "_invalid" + + +# ── topics.yaml resolution ──────────────────────────────────────────────────── + +class TestResolveWorkspaceName: + def test_topic_match(self, temp_hermes_home: Path): + self._write_topics(temp_hermes_home, { + "topics": { + "7695": "news-feed", + "7696": "code-review", + }, + }) + result = _resolve_workspace_name(temp_hermes_home, "telegram", "-1003682109119", "7695") + assert result == "news-feed" + + def test_unmapped_thread_no_fallback(self, temp_hermes_home: Path): + self._write_topics(temp_hermes_home, { + "topics": { + "7695": "news-feed", + }, + "workspace": "default", + }) + result = _resolve_workspace_name( + temp_hermes_home, "telegram", "-1003682109119", "7696" + ) + # 7696 is not in topics, but "workspace" fallback exists + assert result == "default" + + def test_no_mapping_file(self, temp_hermes_home: Path): + result = _resolve_workspace_name(temp_hermes_home, "telegram", "-1003682109119", "7695") + assert result is None + + def test_list_form_topics(self, temp_hermes_home: Path): + file_path = ( + temp_hermes_home / "platforms" / "telegram" / "123" / "topics.yaml" + ) + file_path.parent.mkdir(parents=True) + file_path.write_text(""" +topics: + - thread_id: "7695" + workspace: news-feed + - thread_id: "7696" + workspace: code-review +""", encoding="utf-8") + result = _resolve_workspace_name(temp_hermes_home, "telegram", "123", "7696") + assert result == "code-review" + + def test_string_thread_id_vs_int(self, temp_hermes_home: Path): + """Thread IDs from Telegram come as strings; config may use quoted numbers.""" + self._write_topics(temp_hermes_home, { + "topics": { + "7695": "news-feed", + } + }) + # Pass int-like string + result = _resolve_workspace_name( + temp_hermes_home, "telegram", "-1003682109119", "7695" + ) + assert result == "news-feed" + + @staticmethod + def _write_topics(home: Path, data: dict): + file_path = ( + home / "platforms" / "telegram" / "_-1003682109119" / "topics.yaml" + ) + file_path.parent.mkdir(parents=True) + import yaml + file_path.write_text(yaml.safe_dump(data), encoding="utf-8") + + +# ── SYSTEM.md resolution ──────────────────────────────────────────────────── + +class TestResolveWorkspaceContent: + def test_prompt_with_skills(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "news-feed", + """--- +skills: + - telegram-summary-bot +--- +Respond in Hebrew. +""", + ) + result = _resolve_workspace_content(temp_hermes_home, "news-feed") + assert result == WorkspaceResult( + prompt="Respond in Hebrew.", + skills=["telegram-summary-bot"], + model=None, + ) + + def test_single_skill_string(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "code-review", + """--- +skills: conventional-commits +--- +Follow standards. +""", + ) + result = _resolve_workspace_content(temp_hermes_home, "code-review") + assert result == WorkspaceResult( + prompt="Follow standards.", + skills=["conventional-commits"], + model=None, + ) + + def test_no_frontmatter(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "general", "General discussion.\n" + ) + result = _resolve_workspace_content(temp_hermes_home, "general") + assert result == WorkspaceResult(prompt="General discussion.", skills=None, model=None) + + def test_empty_skills_list(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "empty-test", + """--- +skills: [] +--- +Just a prompt. +""", + ) + result = _resolve_workspace_content(temp_hermes_home, "empty-test") + assert result == WorkspaceResult(prompt="Just a prompt.", skills=None, model=None) + + def test_no_system_file(self, temp_hermes_home: Path): + result = _resolve_workspace_content(temp_hermes_home, "nonexistent") + assert result == WorkspaceResult(None, None, None) + + def test_model_from_frontmatter(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "custom-model", + """--- +skills: + - thinking-before-acting +model: google/gemma-4-31b-it +--- +Use step-by-step reasoning. +""", + ) + result = _resolve_workspace_content(temp_hermes_home, "custom-model") + assert result == WorkspaceResult( + prompt="Use step-by-step reasoning.", + skills=["thinking-before-acting"], + model="google/gemma-4-31b-it", + ) + + def test_model_only_frontmatter(self, temp_hermes_home: Path): + self._write_system( + temp_hermes_home, "model-only", + """--- +model: anthropic/claude-sonnet-4 +--- +Some prompt. +""", + ) + result = _resolve_workspace_content(temp_hermes_home, "model-only") + assert result.model == "anthropic/claude-sonnet-4" + assert result.prompt == "Some prompt." + + @staticmethod + def _write_system(home: Path, name: str, content: str): + path = home / "workspaces" / name / "SYSTEM.md" + path.parent.mkdir(parents=True) + path.write_text(content, encoding="utf-8") + + +# ── Full resolution ────────────────────────────────────────────────────────── + +class TestResolveWorkspace: + def test_end_to_end_topic_match(self, temp_hermes_home: Path): + self._write_full_workspace(temp_hermes_home, workspace="news-feed") + result = resolve_workspace("telegram", "-1003682109119", "7695") + assert result.prompt == "Hebrew news." + assert result.skills == ["telegram-summary-bot"] + + def test_end_to_end_no_match(self, temp_hermes_home: Path): + # No topics.yaml + result = resolve_workspace("telegram", "-1003682109119", "7695") + assert result == WorkspaceResult(None, None, None) + + def test_end_to_end_unknown_chat(self, temp_hermes_home: Path): + self._write_full_workspace(temp_hermes_home, workspace="news-feed") + result = resolve_workspace("telegram", "-999123", "7695") + assert result == WorkspaceResult(None, None, None) + + @staticmethod + def _write_full_workspace(home: Path, workspace: str): + # topics.yaml + tp = home / "platforms" / "telegram" / "_-1003682109119" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text( + yaml.safe_dump({"topics": {"7695": workspace}}), + encoding="utf-8", + ) + # SYSTEM.md + sp = home / "workspaces" / workspace / "SYSTEM.md" + sp.parent.mkdir(parents=True) + sp.write_text( + "---\nskills:\n - telegram-summary-bot\n---\nHebrew news.", + encoding="utf-8", + ) + + +# ── Skill dirs ──────────────────────────────────────────────────────────────── + +class TestGetWorkspaceSkillDirs: + def test_no_skill_dir(self, temp_hermes_home: Path): + # Create workspace but no skills/ + tp = temp_hermes_home / "platforms" / "telegram" / "_-1003682109119" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text( + yaml.safe_dump({"topics": {"7695": "news-feed"}}), + encoding="utf-8", + ) + sp = temp_hermes_home / "workspaces" / "news-feed" / "SYSTEM.md" + sp.parent.mkdir(parents=True) + sp.write_text("Prompt", encoding="utf-8") + + result = get_workspace_skill_dirs("telegram", "-1003682109119", "7695") + assert result == [] + + def test_skill_dir_exists(self, temp_hermes_home: Path): + tp = temp_hermes_home / "platforms" / "telegram" / "_-1003682109119" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text( + yaml.safe_dump({"topics": {"7695": "news-feed"}}), + encoding="utf-8", + ) + # Create skills/ dir + skills_dir = temp_hermes_home / "workspaces" / "news-feed" / "skills" + skills_dir.mkdir(parents=True) + # Drop a placeholder + (skills_dir / "my-skill").mkdir() + (skills_dir / "my-skill" / "SKILL.md").write_text("---\nname: my-skill\n---\n", encoding="utf-8") + + result = get_workspace_skill_dirs("telegram", "-1003682109119", "7695") + assert len(result) == 1 + assert result[0].name == "skills" + + +# ── Stat cache behavior ─────────────────────────────────────────────────────── + +class TestStatCache: + def test_cache_returns_same_content(self, temp_hermes_home: Path): + """Multiple reads within 1s should use cache.""" + from agent.workspace_resolver import _stat_cached + + path = temp_hermes_home / "test.txt" + path.write_text("v1", encoding="utf-8") + r1 = _stat_cached(path) + r2 = _stat_cached(path) + assert r1 is not None and r2 is not None + assert r1[1] == r2[1] # same content string (cache hit) + + def test_reads_new_content_after_ttl(self, temp_hermes_home: Path): + """After file change and ttl expiry, new content is read.""" + import time + from agent.workspace_resolver import _stat_cached + + path = temp_hermes_home / "test.txt" + path.write_text("old", encoding="utf-8") + r1 = _stat_cached(path) + assert r1 is not None and r1[1] == "old" + + # Wait for cache bucket to expire (must cross integer second boundary) + time.sleep(1.1) + path.write_text("new", encoding="utf-8") + _clear_stat_cache() # simulate passing time; in real use bucket naturally rolls + r2 = _stat_cached(path) + assert r2 is not None and r2[1] == "new" + + +# ── Review fix: workspace dir hardening (never escape /workspaces) ───── + +class TestWorkspaceDirHardening: + def test_valid_name(self, temp_hermes_home: Path): + from agent.workspace_resolver import _workspace_dir + + result = _workspace_dir(temp_hermes_home, "news-feed") + assert result is not None + assert result == (temp_hermes_home / "workspaces" / "news-feed") + + def test_traversal_rejected(self, temp_hermes_home: Path): + from agent.workspace_resolver import _workspace_dir + + assert _workspace_dir(temp_hermes_home, "../evil") is None + assert _workspace_dir(temp_hermes_home, "..") is None + assert _workspace_dir(temp_hermes_home, "a/b") is None + assert _workspace_dir(temp_hermes_home, "a\\b") is None + + def test_absolute_rejected(self, temp_hermes_home: Path): + from agent.workspace_resolver import _workspace_dir + + assert _workspace_dir(temp_hermes_home, "/etc/passwd") is None + assert _workspace_dir(temp_hermes_home, "C:/Windows") is None + + def test_evil_content_resolution_is_empty(self, temp_hermes_home: Path): + """Traversal names in topics.yaml must resolve to nothing, not read files.""" + from agent.workspace_resolver import _resolve_workspace_content + + result = _resolve_workspace_content(temp_hermes_home, "../../secret") + assert result == WorkspaceResult(None, None, None) + + +# ── Review fix: shared resolution via apply_workspace_to_event ──────────────── + +class TestApplyWorkspaceToEvent: + def _make_event(self, platform="telegram", chat_id="-1001", thread_id="42"): + from gateway.platforms.base import MessageEvent, MessageType + from gateway.platforms.base import SessionSource + + source = SessionSource( + platform=platform, + chat_id=chat_id, + thread_id=thread_id, + ) + return MessageEvent( + text="hi", + message_type=MessageType.TEXT, + source=source, + ) + + def test_applies_prompt_skills_model(self, temp_hermes_home: Path): + from agent.workspace_resolver import apply_workspace_to_event + + tp = temp_hermes_home / "platforms" / "telegram" / "_-1001" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text(yaml.safe_dump({"topics": {"42": "news-feed"}}), encoding="utf-8") + sp = temp_hermes_home / "workspaces" / "news-feed" / "SYSTEM.md" + sp.parent.mkdir(parents=True) + sp.write_text( + "---\nskills:\n - telegram-summary-bot\nmodel: google/gemma-4-31b-it\n---\nHebrew news.", + encoding="utf-8", + ) + + event = self._make_event() + applied = apply_workspace_to_event(event) + assert applied is True + assert event.channel_prompt == "Hebrew news." + assert event.auto_skill == ["telegram-summary-bot"] + assert event.workspace_model == "google/gemma-4-31b-it" + + def test_idempotent(self, temp_hermes_home: Path): + from agent.workspace_resolver import apply_workspace_to_event + + tp = temp_hermes_home / "platforms" / "telegram" / "_-1001" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text(yaml.safe_dump({"topics": {"42": "news-feed"}}), encoding="utf-8") + sp = temp_hermes_home / "workspaces" / "news-feed" / "SYSTEM.md" + sp.parent.mkdir(parents=True) + sp.write_text("Prompt", encoding="utf-8") + + event = self._make_event() + assert apply_workspace_to_event(event) is True + # Second call must be a no-op (already applied) + assert apply_workspace_to_event(event) is False + + def test_no_workspace_no_change(self, temp_hermes_home: Path): + from agent.workspace_resolver import apply_workspace_to_event + + event = self._make_event() + applied = apply_workspace_to_event(event) + assert applied is False + assert event.channel_prompt is None + assert event.workspace_model is None + + def test_merges_workspace_skills_with_existing(self, temp_hermes_home: Path): + from agent.workspace_resolver import apply_workspace_to_event + + tp = temp_hermes_home / "platforms" / "telegram" / "_-1001" / "topics.yaml" + tp.parent.mkdir(parents=True) + import yaml + tp.write_text(yaml.safe_dump({"topics": {"42": "news-feed"}}), encoding="utf-8") + sp = temp_hermes_home / "workspaces" / "news-feed" / "SYSTEM.md" + sp.parent.mkdir(parents=True) + sp.write_text("---\nskills:\n - ws-skill\n---\nPrompt", encoding="utf-8") + + event = self._make_event() + event.auto_skill = "config-skill" + apply_workspace_to_event(event) + assert "ws-skill" in event.auto_skill + assert "config-skill" in event.auto_skill diff --git a/tests/gateway/test_config_env_bridge_authority.py b/tests/gateway/test_config_env_bridge_authority.py index 35e277664dc9..c96411a00ee9 100644 --- a/tests/gateway/test_config_env_bridge_authority.py +++ b/tests/gateway/test_config_env_bridge_authority.py @@ -55,10 +55,26 @@ def _run_gateway_import(hermes_home: Path, initial_env: dict[str, str]) -> dict[ print(f"{{k}}={{v}}") """ ) - env = dict(initial_env) + # Start from the FULL parent environment so the subprocess gets a + # functional OS environment. A minimal {}-based env breaks on Windows: + # Winsock cannot initialize without SystemRoot/windir etc., so + # ``import gateway.run`` dies with WinError 10106 ("requested service + # provider could not be loaded or initialized"). + # + # Every HERMES_* key is scrubbed first: the bridge must source those + # values exclusively from the .env / config.yaml under hermes_home — + # never from the parent process — which is exactly what the tests below + # verify. + env = { + k: v + for k, v in os.environ.items() + if not k.startswith("HERMES_") + } + env.update(initial_env) env["HERMES_HOME"] = str(hermes_home) - # Keep PATH / PYTHONPATH so venv imports resolve. - for k in ("PATH", "PYTHONPATH", "VIRTUAL_ENV", "HOME"): + # Keep PATH / PYTHONPATH so venv imports resolve (already present in + # the copy above; kept explicit for clarity). + for k in ("PATH", "PYTHONPATH", "VIRTUAL_ENV"): if k in os.environ and k not in env: env[k] = os.environ[k] diff --git a/tests/gateway/test_discord_liveness.py b/tests/gateway/test_discord_liveness.py index 4cd87c6ddb6f..4b2063aa3b1c 100644 --- a/tests/gateway/test_discord_liveness.py +++ b/tests/gateway/test_discord_liveness.py @@ -194,6 +194,281 @@ async def _wait_until(predicate, message: str, timeout: float = 2.0) -> None: await asyncio.sleep(0.01) +@pytest.mark.asyncio +async def test_liveness_probe_disabled_when_interval_zero(monkeypatch): + """interval<=0 must skip the probe entirely so users can opt out.""" + adapter = _make_adapter(monkeypatch, interval=0) + + bot_holder: dict = {} + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + bot.fetch_user = AsyncMock() + bot_holder["bot"] = bot + return bot + + await _connect(adapter, monkeypatch, factory) + assert adapter._liveness_task is None + await asyncio.sleep(0.05) + bot_holder["bot"].fetch_user.assert_not_called() + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_probe_disabled_when_threshold_zero(monkeypatch): + """threshold<=0 must also skip the probe.""" + adapter = _make_adapter(monkeypatch, interval=0.01, threshold=0) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + bot.fetch_user = AsyncMock() + return bot + + await _connect(adapter, monkeypatch, factory) + assert adapter._liveness_task is None + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_probe_does_not_call_rest_while_websocket_is_healthy(monkeypatch): + """A fresh Gateway ACK is sufficient; REST is not a transport health probe.""" + adapter = _make_adapter(monkeypatch, interval=0.01, threshold=3) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot) + bot.fetch_user = AsyncMock(return_value=SimpleNamespace(id=999)) + return bot + + await _connect(adapter, monkeypatch, factory) + await asyncio.sleep(0.05) + adapter._client.fetch_user.assert_not_awaited() + assert adapter._running is True + assert adapter.has_fatal_error is False + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_probe_forces_reconnect_when_rest_succeeds_but_gateway_ack_is_stale(monkeypatch): + """A REST response must not hide a stale Gateway heartbeat failure.""" + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=2, max_ack_age=0.01) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot, ack_age=3600) + bot.fetch_user = AsyncMock(return_value=SimpleNamespace(id=999)) + return bot + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + wedged = adapter._client + + # The sampler schedules the close + supervisor callback in a sibling task + # so the fatal path cannot cancel/await itself through disconnect(). + await _wait_until( + lambda: handler.await_count, + "liveness recovery notification did not complete within 2s", + ) + + assert adapter._liveness_task and adapter._liveness_task.done() + assert wedged.is_closed() is True + assert adapter.has_fatal_error is True + assert adapter.fatal_error_code == "discord_websocket_health_stale" + assert adapter.fatal_error_retryable is True + wedged.fetch_user.assert_not_awaited() + handler.assert_awaited_once() + + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_fatal_queues_primary_runner_reconnect_without_self_cancellation(monkeypatch): + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=1, max_ack_age=0.01) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot, ack_age=3600) + return bot + + runner = GatewayRunner.__new__(GatewayRunner) + runner.adapters = {Platform.DISCORD: adapter} + runner._failed_platforms = {} + runner._running = True + runner.stop = AsyncMock() + runner.delivery_router = SimpleNamespace(adapters=runner.adapters) + runner.config = SimpleNamespace(platforms={Platform.DISCORD: adapter.config}) + runner._update_platform_runtime_status = lambda *args, **kwargs: None + runner._adapter_disconnect_timeout_secs = lambda: 0.1 + adapter.set_fatal_error_handler(runner._handle_adapter_fatal_error) + await _connect(adapter, monkeypatch, factory) + + await _wait_until( + lambda: Platform.DISCORD in runner._failed_platforms, + "liveness fatal did not reach the runner reconnect queue", + ) + + assert adapter._liveness_notification_task is None or adapter._liveness_notification_task.done() + assert runner._failed_platforms[Platform.DISCORD]["attempts"] == 0 + runner.stop.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("health", "expected_reason"), + [ + ({"ready": False}, "not_ready"), + ({"socket_open": False}, "socket_closed"), + ({"latency": float("inf")}, "latency_non_finite"), + ], +) +async def test_liveness_probe_reports_gateway_health_failure_reason(monkeypatch, health, expected_reason): + # max_ack_age is generous (60s) so the latency case cannot race the + # ack window: the ack is created at bot-construction time, and on a + # loaded machine connect() can take >1s, which made ack_stale win + # over latency_non_finite (flaky under the full suite). + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=1, max_ack_age=60.0) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot, **health) + bot.fetch_user = AsyncMock(return_value=SimpleNamespace(id=999)) + return bot + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + + await _wait_until( + lambda: handler.await_count, + "liveness loop did not surface a websocket health failure", + ) + + assert expected_reason in (adapter.fatal_error_message or "") + adapter._client.fetch_user.assert_not_awaited() + handler.assert_awaited_once() + await adapter.disconnect() + + + + +@pytest.mark.asyncio +async def test_liveness_probe_treats_websocket_state_read_error_as_unhealthy(monkeypatch): + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=1) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + bot.ws = _BrokenWebSocket() + return bot + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + + await _wait_until( + lambda: handler.await_count, + "liveness loop did not surface a WebSocket state read error", + ) + + assert "socket_state_unavailable" in (adapter.fatal_error_message or "") + handler.assert_awaited_once() + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_probe_recovers_when_health_reader_raises(monkeypatch): + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=1) + + def factory(**kwargs): + return _LiveBot( + intents=kwargs["intents"], + allowed_mentions=kwargs.get("allowed_mentions"), + ) + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + monkeypatch.setattr( + adapter, + "_read_websocket_health", + lambda _client: (_ for _ in ()).throw(RuntimeError("unexpected state")), + ) + + await _wait_until( + lambda: handler.await_count, + "liveness loop did not recover from health-reader failure", + ) + + assert "health_check_error" in (adapter.fatal_error_message or "") + handler.assert_awaited_once() + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_recovery_keeps_websocket_fatal_when_client_task_exits(monkeypatch): + """The close callback must not replace stale-ACK recovery with task-exited.""" + adapter = _make_adapter(monkeypatch, interval=0.005, threshold=1, max_ack_age=0.01) + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot, ack_age=3600) + return bot + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + + await _wait_until( + lambda: handler.await_count, + "closed client task did not finish within 2s", + ) + + assert adapter._bot_task and adapter._bot_task.done() + assert adapter.fatal_error_code == "discord_websocket_health_stale" + assert handler.await_count == 1 + await adapter.disconnect() + + +@pytest.mark.asyncio +async def test_liveness_recovery_not_blocked_by_hanging_client_close(monkeypatch): + """A wedged close must not prevent fatal notification/reconnect queueing.""" + adapter = _make_adapter(monkeypatch, interval=60, threshold=1, max_ack_age=1.0) + monkeypatch.setenv("HERMES_GATEWAY_ADAPTER_DISCONNECT_TIMEOUT", "0.02") + + def factory(**kwargs): + bot = _LiveBot(intents=kwargs["intents"], allowed_mentions=kwargs.get("allowed_mentions")) + _set_websocket_health(bot, ack_age=3600) + bot.fetch_user = AsyncMock(return_value=SimpleNamespace(id=999)) + return bot + + handler = AsyncMock() + adapter.set_fatal_error_handler(handler) + await _connect(adapter, monkeypatch, factory) + wedged = adapter._client + close_started = asyncio.Event() + + async def hanging_close(): + close_started.set() + await asyncio.Event().wait() + + wedged.close = hanging_close + adapter._set_fatal_error( + "discord_websocket_health_stale", + "Discord Gateway WebSocket health check failed: ack_stale", + retryable=True, + ) + notify_task = asyncio.create_task(adapter._notify_liveness_fatal_error(wedged)) + await asyncio.wait_for(close_started.wait(), timeout=0.5) + await asyncio.wait_for(notify_task, timeout=2.0) + assert close_started.is_set() is True + assert handler.await_count == 1 + assert adapter.fatal_error_code == "discord_websocket_health_stale" + + # Restore a cooperative fake close so the test can release the bot task. + wedged.close = _LiveBot.close.__get__(wedged, _LiveBot) + await adapter.disconnect() + + @pytest.mark.asyncio async def test_liveness_close_timeout_aborts_aiohttp_transport_before_fatal_notification( monkeypatch, diff --git a/tests/gateway/test_dm_topics.py b/tests/gateway/test_dm_topics.py index f889f95c66b0..bf0631e4cbcc 100644 --- a/tests/gateway/test_dm_topics.py +++ b/tests/gateway/test_dm_topics.py @@ -12,7 +12,7 @@ import os import sys from pathlib import Path -from types import SimpleNamespace +from types import ModuleType, SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -20,24 +20,47 @@ from gateway.config import PlatformConfig -# Use the shared, comprehensive telegram mock from conftest instead of a -# file-local one. The previous local installer differed from every other -# telegram test's stub in two ways — it registered a SEPARATE string-valued -# ``telegram.constants`` module (others register the root mock, so ParseMode -# members stay auto-generated MagicMock attributes) and its ``telegram.error`` -# was a bare MagicMock (conftest defines real exception subclasses with PTB's -# hierarchy) — and it installed UNCONDITIONALLY (no real-library guard). -# Because it also force-reimported the adapter, the divergent stub leaked into -# sys.modules for the rest of the session: every later telegram test that -# asserts ParseMode repr or isinstance against telegram.error classes failed -# order-dependently in full runs while passing in isolation. -from tests.gateway.conftest import _ensure_telegram_mock # noqa: E402 +def _ensure_telegram_mock(): + existing = sys.modules.get("telegram") + if existing is not None and isinstance(existing, ModuleType): + return # Real telegram installed — nothing to mock. + + if existing is None: + # Safety net for standalone runs: the gateway conftest normally + # installs the suite-wide telegram mock before any test module is + # imported, so this branch is rarely hit. + root = MagicMock() + root.ext.ContextTypes.DEFAULT_TYPE = type(None) + sys.modules["telegram"] = root + sys.modules["telegram.ext"] = root.ext + sys.modules["telegram.request"] = root.request + + # Do NOT clobber the suite-wide telegram mock (tests/gateway/conftest.py). + # Replacing it (as this module used to) also breaks ``telegram.error`` + # resolution for every later test file — NetworkError/TimedOut become + # auto-generated MagicMocks and the reconnect classifier tests in + # test_telegram_network_reconnect.py fail — and its string-valued + # ParseMode.MARKDOWN_V2 breaks the parse-mode assertions in + # test_telegram_approval_buttons.py / test_telegram_model_picker.py / + # test_telegram_slash_confirm.py. + # + # The only thing this module needs beyond the conftest mock is a + # dedicated constants module with string-valued ChatType members so + # ``from telegram.constants import ChatType`` comparisons work (the + # conftest mock leaves ChatType as auto-generated MagicMocks). + constants_mod = MagicMock() + constants_mod.ChatType.GROUP = "group" + constants_mod.ChatType.SUPERGROUP = "supergroup" + constants_mod.ChatType.CHANNEL = "channel" + constants_mod.ChatType.PRIVATE = "private" + + sys.modules["telegram.constants"] = constants_mod + + # Force reimport so the adapter picks up the string-valued ChatType. + sys.modules.pop("plugins.platforms.telegram.adapter", None) + _ensure_telegram_mock() -# Force reimport so the adapter binds to whatever sys.modules now holds -# (the shared mock, or the real library when it is installed) rather than a -# stub an earlier test file may have bound it to. -sys.modules.pop("plugins.platforms.telegram.adapter", None) from plugins.platforms.telegram.adapter import TelegramAdapter # noqa: E402 @@ -57,6 +80,29 @@ def _make_adapter(dm_topics_config=None, group_topics_config=None): # ── _setup_dm_topics: load persisted thread_ids ── +@pytest.mark.asyncio +async def test_setup_dm_topics_loads_persisted_thread_ids(): + """Topics with thread_id in config should be loaded into cache, not created.""" + adapter = _make_adapter([ + { + "chat_id": 111, + "topics": [ + {"name": "General", "thread_id": 100}, + {"name": "Work", "thread_id": 200}, + ], + } + ]) + adapter._bot = AsyncMock() + + await adapter._setup_dm_topics() + + # Both should be in cache + assert adapter._dm_topics["111:General"] == 100 + assert adapter._dm_topics["111:Work"] == 200 + # create_forum_topic should NOT have been called + adapter._bot.create_forum_topic.assert_not_called() + + @pytest.mark.asyncio async def test_setup_dm_topics_creates_when_no_thread_id(): """Topics without thread_id should be created via API.""" @@ -114,6 +160,29 @@ async def test_setup_dm_topics_mixed_persisted_and_new(): adapter._bot.create_forum_topic.assert_called_once() +@pytest.mark.asyncio +async def test_setup_dm_topics_skips_empty_config(): + """Empty dm_topics config should be a no-op.""" + adapter = _make_adapter([]) + adapter._bot = AsyncMock() + + await adapter._setup_dm_topics() + + adapter._bot.create_forum_topic.assert_not_called() + assert adapter._dm_topics == {} + + +@pytest.mark.asyncio +async def test_setup_dm_topics_no_config(): + """No dm_topics in config at all should be a no-op.""" + adapter = _make_adapter() + adapter._bot = AsyncMock() + + await adapter._setup_dm_topics() + + adapter._bot.create_forum_topic.assert_not_called() + + # ── _create_dm_topic: error handling ── @@ -141,6 +210,17 @@ async def test_create_dm_topic_handles_generic_error(): assert result is None +@pytest.mark.asyncio +async def test_create_dm_topic_returns_none_without_bot(): + """No bot instance should return None.""" + adapter = _make_adapter() + adapter._bot = None + + result = await adapter._create_dm_topic(chat_id=111, name="General") + + assert result is None + + @pytest.mark.asyncio async def test_ensure_dm_topic_creates_on_demand_and_persists(): """Named delivery targets should create missing private DM topics on demand.""" @@ -165,6 +245,30 @@ async def test_ensure_dm_topic_creates_on_demand_and_persists(): ) +@pytest.mark.asyncio +async def test_ensure_dm_topic_force_create_replaces_persisted_thread_id(): + """Refreshing a stale named topic should replace the cached persisted thread_id.""" + adapter = _make_adapter() + bot = AsyncMock() + bot.create_forum_topic.return_value = SimpleNamespace(message_thread_id=777) + adapter._bot = bot + adapter._persist_dm_topic_thread_id = MagicMock() + adapter._dm_topics = {"111:General": 500} + adapter._dm_topics_config = [ + {"chat_id": 111, "topics": [{"name": "General", "thread_id": 500}]} + ] + + result = await adapter.ensure_dm_topic("111", "General", force_create=True) + + assert result == "777" + bot.create_forum_topic.assert_called_once_with(chat_id=111, name="General") + assert adapter._dm_topics["111:General"] == 777 + assert adapter._dm_topics_config[0]["topics"][0]["thread_id"] == 777 + adapter._persist_dm_topic_thread_id.assert_called_once_with( + 111, "General", 777, replace_existing=True + ) + + # ── _persist_dm_topic_thread_id ── @@ -209,6 +313,83 @@ def test_persist_dm_topic_thread_id_writes_config(tmp_path): assert "thread_id" not in topics[1] # "Work" should be untouched +def test_persist_dm_topic_thread_id_skips_if_already_set(tmp_path): + """Should not overwrite an existing thread_id.""" + import yaml + + config_data = { + "platforms": { + "telegram": { + "extra": { + "dm_topics": [ + { + "chat_id": 111, + "topics": [ + {"name": "General", "icon_color": 123, "thread_id": 500}, + ], + } + ] + } + } + } + } + + config_file = tmp_path / ".hermes" / "config.yaml" + config_file.parent.mkdir(parents=True) + with open(config_file, "w") as f: + yaml.dump(config_data, f) + + adapter = _make_adapter() + + with patch.object(Path, "home", return_value=tmp_path): + adapter._persist_dm_topic_thread_id(111, "General", 999) + + with open(config_file) as f: + result = yaml.safe_load(f) + + topics = result["platforms"]["telegram"]["extra"]["dm_topics"][0]["topics"] + assert topics[0]["thread_id"] == 500 # unchanged + + +def test_persist_dm_topic_thread_id_replaces_existing_when_requested(tmp_path): + """Forced refresh should overwrite a stale persisted thread_id.""" + import yaml + + config_data = { + "platforms": { + "telegram": { + "extra": { + "dm_topics": [ + { + "chat_id": 111, + "topics": [ + {"name": "General", "icon_color": 123, "thread_id": 500}, + ], + } + ] + } + } + } + } + + config_file = tmp_path / ".hermes" / "config.yaml" + config_file.parent.mkdir(parents=True) + with open(config_file, "w") as f: + yaml.dump(config_data, f) + + adapter = _make_adapter() + + with patch.object(Path, "home", return_value=tmp_path), \ + patch.dict(os.environ, {"HERMES_HOME": str(tmp_path / ".hermes")}): + adapter._persist_dm_topic_thread_id(111, "General", 999, replace_existing=True) + + with open(config_file) as f: + result = yaml.safe_load(f) + + topics = result["platforms"]["telegram"]["extra"]["dm_topics"][0]["topics"] + assert topics[0]["thread_id"] == 999 + + # ── _get_dm_topic_info ── @@ -273,6 +454,43 @@ def test_get_dm_topic_info_finds_cached_topic(): assert result["skill"] == "my-skill" +def test_get_dm_topic_info_returns_none_for_unknown(): + """Should return None for unknown thread_id.""" + adapter = _make_adapter([ + { + "chat_id": 111, + "topics": [{"name": "General"}], + } + ]) + # Mock reload to avoid filesystem access + adapter._reload_dm_topics_from_config = lambda: None + + result = adapter._get_dm_topic_info("111", "999") + + assert result is None + + +def test_get_dm_topic_info_returns_none_without_config(): + """Should return None if no dm_topics config.""" + adapter = _make_adapter() + adapter._reload_dm_topics_from_config = lambda: None + + result = adapter._get_dm_topic_info("111", "100") + + assert result is None + + +def test_get_dm_topic_info_returns_none_for_none_thread(): + """Should return None if thread_id is None.""" + adapter = _make_adapter([ + {"chat_id": 111, "topics": [{"name": "General"}]} + ]) + + result = adapter._get_dm_topic_info("111", None) + + assert result is None + + def test_get_dm_topic_info_hot_reloads_from_config(tmp_path): """Should find a topic added to config after startup (hot-reload).""" import yaml @@ -317,6 +535,15 @@ def test_get_dm_topic_info_hot_reloads_from_config(tmp_path): # ── _cache_dm_topic_from_message ── +def test_cache_dm_topic_from_message(): + """Should cache a new topic mapping.""" + adapter = _make_adapter() + + adapter._cache_dm_topic_from_message("111", "100", "General") + + assert adapter._dm_topics["111:General"] == 100 + + def test_cache_dm_topic_from_message_no_overwrite(): """Should not overwrite an existing cached topic.""" adapter = _make_adapter() @@ -410,6 +637,51 @@ def test_build_message_event_no_auto_skill_without_binding(): assert event.source.chat_topic == "General" +def test_build_message_event_no_auto_skill_without_thread(): + """Regular DM messages (no thread_id) should have auto_skill=None.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() + msg = _make_mock_message(chat_id=111, thread_id=None) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + + +def test_build_message_event_filters_non_topic_dm_thread_id(): + """A DM reply-thread id should not be persisted unless Telegram marks it as a topic message.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() + msg = _make_mock_message(chat_id=111, thread_id=777, is_topic_message=False) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.source.thread_id is None + assert event.source.chat_topic is None + assert event.auto_skill is None + + +def test_build_message_event_preserves_true_dm_topic_thread_id(): + """True DM topic messages should keep their thread id for routing.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter([ + { + "chat_id": 111, + "topics": [ + {"name": "General", "thread_id": 200}, + ], + } + ]) + adapter._dm_topics["111:General"] = 200 + + msg = _make_mock_message(chat_id=111, thread_id=200, is_topic_message=True) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.source.thread_id == "200" + assert event.source.chat_topic == "General" + + # ── _build_message_event: group_topics skill binding ── # The telegram mock sets sys.modules["telegram.constants"] = telegram_mod (root mock), @@ -475,6 +747,411 @@ def test_group_topic_skill_binding_second_topic(): assert event.source.chat_topic == "Sales" +def test_group_topic_no_skill_binding(): + """Group topic without a skill key should have auto_skill=None but set chat_topic.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + {"name": "General", "thread_id": 1}, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=1, + text="hey", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic == "General" + + +# ── _build_message_event: topic prompt resolution ── + + +def test_group_topic_prompt_binding(): + """Group topic with prompt config should set channel_prompt on the event.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1003682109119, + "topics": [ + { + "name": "news feed", + "thread_id": 7695, + "prompt": "Respond in Hebrew. Focus on regional news.", + }, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1003682109119, chat_type=_ChatType.SUPERGROUP, thread_id=7695, text="latest news", + is_topic_message=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.channel_prompt == "Respond in Hebrew. Focus on regional news." + assert event.source.chat_topic == "news feed" + + +def test_group_topic_prompt_with_skill(): + """Group topic with both skill and prompt should set both on the event.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + { + "name": "Engineering", + "thread_id": 42, + "skill": "software-development", + "prompt": "Follow conventional-commits and thinking-before-acting.", + }, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, chat_type=_ChatType.SUPERGROUP, thread_id=42, text="refactor the api", + is_topic_message=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill == "software-development" + assert event.channel_prompt == "Follow conventional-commits and thinking-before-acting." + assert event.source.chat_topic == "Engineering" + + +def test_group_topic_prompt_takes_priority_over_channel_prompts(): + """Topic-level prompt should override channel_prompts dict for the same thread_id.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + { + "name": "Engineering", + "thread_id": 42, + "prompt": "Topic-level prompt takes priority", + }, + ], + } + ]) + # Also set a channel_prompts entry for the same thread_id + adapter.config.extra["channel_prompts"] = {"42": "Channel-level prompt"} + + msg = _make_mock_message( + chat_id=-1001234567890, chat_type=_ChatType.SUPERGROUP, thread_id=42, text="hello", + is_topic_message=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + # Topic-level prompt should take priority + assert event.channel_prompt == "Topic-level prompt takes priority" + + +def test_group_topic_no_prompt_falls_back_to_channel_prompts(): + """When topic has no prompt, channel_prompts dict should still be used.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + {"name": "General", "thread_id": 42}, + ], + } + ]) + adapter.config.extra["channel_prompts"] = {"42": "Channel-level prompt fallback"} + + msg = _make_mock_message( + chat_id=-1001234567890, chat_type=_ChatType.SUPERGROUP, thread_id=42, text="hello", + is_topic_message=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + # Falls back to channel_prompts since topic has no prompt field + assert event.channel_prompt == "Channel-level prompt fallback" + + +def test_dm_topic_prompt_binding(): + """DM topic with prompt config should set channel_prompt on the event.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(dm_topics_config=[ + { + "chat_id": 12345, + "topics": [ + { + "name": "projects", + "thread_id": 100, + "skill": "project-tracker", + "prompt": "Track project status and deadlines.", + }, + ], + } + ]) + adapter._dm_topics["12345:projects"] = 100 + + msg = _make_mock_message( + chat_id=12345, chat_type=_ChatType.PRIVATE, thread_id=100, + is_topic_message=True, text="status update" + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill == "project-tracker" + assert event.channel_prompt == "Track project status and deadlines." + assert event.source.chat_topic == "projects" + + +def test_group_topic_unmapped_thread_id(): + """Thread ID not in config should fall through — no skill, no topic name.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + {"name": "Engineering", "thread_id": 5, "skill": "software-development"}, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=999, + text="random", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic is None + + +def test_group_topic_unmapped_chat_id(): + """Chat ID not in group_topics config should fall through silently.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": -1001234567890, + "topics": [ + {"name": "Engineering", "thread_id": 5, "skill": "software-development"}, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1009999999999, + chat_type=_ChatType.SUPERGROUP, + thread_id=5, + text="wrong group", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic is None + + +def test_group_topic_no_config(): + """No group_topics config at all should be fine — no skill, no topic.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() # no group_topics_config + + msg = _make_mock_message( + chat_id=-1001234567890, chat_type=_ChatType.GROUP, thread_id=5, text="hi" + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic is None + + +def test_group_topic_chat_id_int_string_coercion(): + """chat_id as string in config should match integer chat.id via str() coercion.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + { + "chat_id": "-1001234567890", # string, not int + "topics": [ + {"name": "Dev", "thread_id": "7", "skill": "hermes-agent-dev"}, + ], + } + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=7, + text="test", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill == "hermes-agent-dev" + assert event.source.chat_topic == "Dev" + + +def test_group_topic_mapping_shape_config(): + """Operator-edited mapping shape {chat_id: [topics]} must resolve like the list shape.""" + from gateway.platforms.base import MessageType + + # Dict/mapping shape instead of the canonical list-of-entries shape. + adapter = _make_adapter(group_topics_config={ + "-1001234567890": [ + {"name": "Engineering", "thread_id": 5, "skill": "software-development"}, + {"name": "Sales", "thread_id": 12, "skill": "sales-framework"}, + ], + }) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=12, + text="deal update", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill == "sales-framework" + assert event.source.chat_topic == "Sales" + + +def test_group_topic_malformed_config_does_not_crash(): + """Non-dict entries / non-list topics must be skipped, not raise AttributeError.""" + from gateway.platforms.base import MessageType + + # Junk list entries (str) are filtered out; a matching entry with a good + # topic still resolves; non-dict topic entries within it are skipped. + adapter = _make_adapter(group_topics_config=[ + "not-a-dict", + {"chat_id": -1001234567890, "topics": ["also-not-a-dict", + {"name": "Good", "thread_id": 5}]}, + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=5, + text="hi", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic == "Good" + + +def test_group_topic_non_list_topics_does_not_crash(): + """A matched entry whose topics is not a list must fall through, not raise.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=[ + {"chat_id": -1001234567890, "topics": "oops-not-a-list"}, + ]) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=5, + text="hi", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic is None + + +def test_group_topic_scalar_config_falls_through(): + """A scalar (int/str) group_topics value must fall through cleanly, not raise.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter(group_topics_config=42) + + msg = _make_mock_message( + chat_id=-1001234567890, + chat_type=_ChatType.SUPERGROUP, + thread_id=5, + text="hi", + is_topic_message=True, + is_forum=True, + ) + event = adapter._build_message_event(msg, MessageType.TEXT) + + assert event.auto_skill is None + assert event.source.chat_topic is None + + # ── _build_message_event: from_user=None fallback in DMs ── +def test_build_message_event_dm_from_user_none_falls_back_to_chat_id(): + """When from_user is None in a DM, user_id should fall back to chat.id.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() + msg = _make_mock_message(chat_id=12345, user_id=42, user_name="Alice") + # Simulate from_user being None (edge case on fresh restart / forwarded msg) + msg.from_user = None + + event = adapter._build_message_event(msg, MessageType.TEXT) + + # Should fall back to chat.id since chat_type is "dm" + assert event.source.user_id == "12345" + assert event.source.user_name == "Alice" # falls back to chat.full_name + + +def test_build_message_event_group_from_user_none_stays_none(): + """When from_user is None in a group, user_id should remain None.""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() + msg = _make_mock_message( + chat_id=-1001234567890, chat_type=_ChatType.SUPERGROUP, + user_id=42, user_name="Alice" + ) + msg.from_user = None + + event = adapter._build_message_event(msg, MessageType.TEXT) + + # Groups should NOT fall back — anonymous senders stay None + assert event.source.user_id is None + assert event.source.user_name is None + + +def test_build_message_event_dm_from_user_present_uses_user(): + """When from_user is present in a DM, it should be used (no fallback).""" + from gateway.platforms.base import MessageType + + adapter = _make_adapter() + msg = _make_mock_message(chat_id=12345, user_id=99999, user_name="Bob") + + event = adapter._build_message_event(msg, MessageType.TEXT) + + # Normal case — from_user is used directly + assert event.source.user_id == "99999" + assert event.source.user_name == "Bob" diff --git a/tests/gateway/test_email.py b/tests/gateway/test_email.py index 8e46600b0434..85c4b4491f25 100644 --- a/tests/gateway/test_email.py +++ b/tests/gateway/test_email.py @@ -13,6 +13,22 @@ """ import os + + +# Windows-safe env for ``@patch.dict(..., clear=True)`` tests: the decorator +# wipes the whole environment, and on Windows ``pathlib.Path.home()`` raises +# "Could not determine home directory" when USERPROFILE / HOME / +# HOMEDRIVE+HOMEPATH are all absent (POSIX falls back to the ``pwd`` module, +# which is why these tests only broke on Windows). ``get_hermes_home()`` +# reads HERMES_HOME first, then LOCALAPPDATA, then Path.home(). Preserving +# these variables keeps the "no Feishu env vars" semantics of clear=True +# while letting home resolution work on every platform. +_HOME_ENV = { + k: v + for k, v in os.environ.items() + if k in ("HERMES_HOME", "USERPROFILE", "HOMEDRIVE", "HOMEPATH", "HOME", "LOCALAPPDATA") + and v +} import unittest from email.mime.text import MIMEText from email.mime.multipart import MIMEMultipart @@ -26,6 +42,19 @@ class TestConfigEnvOverrides(unittest.TestCase): """Verify email config is loaded from environment variables.""" + @patch.dict(os.environ, { + "EMAIL_ADDRESS": "hermes@test.com", + "EMAIL_PASSWORD": "secret", + "EMAIL_IMAP_HOST": "imap.test.com", + "EMAIL_SMTP_HOST": "smtp.test.com", + }, clear=False) + def test_email_config_loaded_from_env(self): + from gateway.config import GatewayConfig, Platform, _apply_env_overrides + config = GatewayConfig() + _apply_env_overrides(config) + self.assertIn(Platform.EMAIL, config.platforms) + self.assertTrue(config.platforms[Platform.EMAIL].enabled) + self.assertEqual(config.platforms[Platform.EMAIL].extra["address"], "hermes@test.com") @patch.dict(os.environ, { "EMAIL_ADDRESS": "hermes@test.com", @@ -42,6 +71,12 @@ def test_email_home_channel_loaded(self): self.assertIsNotNone(home) self.assertEqual(home.chat_id, "user@test.com") + @patch.dict(os.environ, {**_HOME_ENV}, clear=True) + def test_email_not_loaded_without_env(self): + from gateway.config import GatewayConfig, Platform, _apply_env_overrides + config = GatewayConfig() + _apply_env_overrides(config) + self.assertNotIn(Platform.EMAIL, config.platforms) class TestCheckRequirements(unittest.TestCase): """Verify check_email_requirements function.""" @@ -56,10 +91,25 @@ def test_requirements_met(self): from plugins.platforms.email.adapter import check_email_requirements self.assertTrue(check_email_requirements()) + @patch.dict(os.environ, { + "EMAIL_ADDRESS": "a@b.com", + }, clear=True) + def test_requirements_not_met(self): + from plugins.platforms.email.adapter import check_email_requirements + self.assertFalse(check_email_requirements()) + + @patch.dict(os.environ, {}, clear=True) + def test_requirements_empty_env(self): + from plugins.platforms.email.adapter import check_email_requirements + self.assertFalse(check_email_requirements()) + class TestHelperFunctions(unittest.TestCase): """Test email parsing helper functions.""" + def test_decode_header_plain(self): + from plugins.platforms.email.adapter import _decode_header_value + self.assertEqual(_decode_header_value("Hello World"), "Hello World") def test_decode_header_encoded(self): from plugins.platforms.email.adapter import _decode_header_value @@ -75,6 +125,19 @@ def test_extract_email_address_with_name(self): "john@example.com" ) + def test_extract_email_address_bare(self): + from plugins.platforms.email.adapter import _extract_email_address + self.assertEqual( + _extract_email_address("john@example.com"), + "john@example.com" + ) + + def test_extract_email_address_uppercase(self): + from plugins.platforms.email.adapter import _extract_email_address + self.assertEqual( + _extract_email_address("John@Example.COM"), + "john@example.com" + ) def test_strip_html_basic(self): from plugins.platforms.email.adapter import _strip_html @@ -85,6 +148,19 @@ def test_strip_html_basic(self): self.assertNotIn("

", result) self.assertNotIn("", result) + def test_strip_html_br_tags(self): + from plugins.platforms.email.adapter import _strip_html + html = "Line 1
Line 2
Line 3" + result = _strip_html(html) + self.assertIn("Line 1", result) + self.assertIn("Line 2", result) + + def test_strip_html_entities(self): + from plugins.platforms.email.adapter import _strip_html + html = "a & b < c > d" + result = _strip_html(html) + self.assertIn("a & b", result) + class TestExtractTextBody(unittest.TestCase): """Test email body extraction from different message formats.""" @@ -95,6 +171,12 @@ def test_plain_text_body(self): result = _extract_text_body(msg) self.assertEqual(result, "Hello, this is a test.") + def test_html_body_fallback(self): + from plugins.platforms.email.adapter import _extract_text_body + msg = MIMEText("

Hello from HTML

", "html", "utf-8") + result = _extract_text_body(msg) + self.assertIn("Hello from HTML", result) + self.assertNotIn("

", result) def test_multipart_prefers_plain(self): from plugins.platforms.email.adapter import _extract_text_body @@ -104,6 +186,19 @@ def test_multipart_prefers_plain(self): result = _extract_text_body(msg) self.assertEqual(result, "Plain version") + def test_multipart_html_only(self): + from plugins.platforms.email.adapter import _extract_text_body + msg = MIMEMultipart("alternative") + msg.attach(MIMEText("

Only HTML

", "html", "utf-8")) + result = _extract_text_body(msg) + self.assertIn("Only HTML", result) + + def test_empty_body(self): + from plugins.platforms.email.adapter import _extract_text_body + msg = MIMEText("", "plain", "utf-8") + result = _extract_text_body(msg) + self.assertEqual(result, "") + class TestExtractAttachments(unittest.TestCase): """Test attachment extraction and caching.""" @@ -114,6 +209,45 @@ def test_no_attachments(self): result = _extract_attachments(msg) self.assertEqual(result, []) + @patch("plugins.platforms.email.adapter.cache_document_from_bytes") + def test_document_attachment(self, mock_cache): + from plugins.platforms.email.adapter import _extract_attachments + mock_cache.return_value = "/tmp/cached_doc.pdf" + + msg = MIMEMultipart() + msg.attach(MIMEText("See attached.", "plain", "utf-8")) + + part = MIMEBase("application", "pdf") + part.set_payload(b"%PDF-1.4 fake pdf content") + encoders.encode_base64(part) + part.add_header("Content-Disposition", "attachment; filename=report.pdf") + msg.attach(part) + + result = _extract_attachments(msg) + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["type"], "document") + self.assertEqual(result[0]["filename"], "report.pdf") + mock_cache.assert_called_once() + + @patch("plugins.platforms.email.adapter.cache_image_from_bytes") + def test_image_attachment(self, mock_cache): + from plugins.platforms.email.adapter import _extract_attachments + mock_cache.return_value = "/tmp/cached_img.jpg" + + msg = MIMEMultipart() + msg.attach(MIMEText("See photo.", "plain", "utf-8")) + + part = MIMEBase("image", "jpeg") + part.set_payload(b"\xff\xd8\xff\xe0 fake jpg") + encoders.encode_base64(part) + part.add_header("Content-Disposition", "attachment; filename=photo.jpg") + msg.attach(part) + + result = _extract_attachments(msg) + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["type"], "image") + mock_cache.assert_called_once() + class TestDispatchMessage(unittest.TestCase): """Test email message dispatch logic.""" @@ -234,6 +368,32 @@ async def capture_handle(event): self.assertNotIn("[Subject:", captured_events[0].text) self.assertEqual(captured_events[0].text, "Thanks for the help!") + def test_empty_body_handled(self): + """Email with no body should dispatch '(empty email)'.""" + import asyncio + adapter = self._make_adapter() + captured_events = [] + + async def capture_handle(event): + captured_events.append(event) + + adapter.handle_message = capture_handle + + msg_data = { + "uid": b"4", + "sender_addr": "user@test.com", + "sender_name": "User", + "subject": "Re: test", + "message_id": "", + "in_reply_to": "", + "body": "", + "attachments": [], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + self.assertEqual(len(captured_events), 1) + self.assertIn("(empty email)", captured_events[0].text) def test_image_attachment_sets_photo_type(self): """Email with image attachment should set message type to PHOTO.""" @@ -264,6 +424,159 @@ async def capture_handle(event): self.assertEqual(captured_events[0].message_type, MessageType.PHOTO) self.assertEqual(captured_events[0].media_urls, ["/tmp/img.jpg"]) + def test_document_attachment_sets_document_type(self): + """Email with a document attachment must set DOCUMENT so run.py injects file context.""" + import asyncio + from gateway.platforms.base import MessageType + adapter = self._make_adapter() + captured_events = [] + + async def capture_handle(event): + captured_events.append(event) + + adapter.handle_message = capture_handle + + msg_data = { + "uid": b"6", + "sender_addr": "user@test.com", + "sender_name": "User", + "subject": "Re: report", + "message_id": "", + "in_reply_to": "", + "body": "See attached", + "attachments": [{"path": "/tmp/report.pdf", "filename": "report.pdf", "type": "document", "media_type": "application/pdf"}], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + self.assertEqual(len(captured_events), 1) + self.assertEqual(captured_events[0].message_type, MessageType.DOCUMENT) + self.assertEqual(captured_events[0].media_urls, ["/tmp/report.pdf"]) + + def test_mixed_image_and_document_prefers_document(self): + """DOCUMENT wins for mixed attachments — image handling keys off per-path + mime types, but document injection gates strictly on MessageType.DOCUMENT.""" + import asyncio + from gateway.platforms.base import MessageType + adapter = self._make_adapter() + captured_events = [] + + async def capture_handle(event): + captured_events.append(event) + + adapter.handle_message = capture_handle + + msg_data = { + "uid": b"7", + "sender_addr": "user@test.com", + "sender_name": "User", + "subject": "Re: both", + "message_id": "", + "in_reply_to": "", + "body": "Photo and PDF", + "attachments": [ + {"path": "/tmp/img.jpg", "filename": "img.jpg", "type": "image", "media_type": "image/jpeg"}, + {"path": "/tmp/report.pdf", "filename": "report.pdf", "type": "document", "media_type": "application/pdf"}, + ], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + self.assertEqual(len(captured_events), 1) + self.assertEqual(captured_events[0].message_type, MessageType.DOCUMENT) + self.assertEqual(len(captured_events[0].media_urls), 2) + + def test_source_built_correctly(self): + """Session source should have correct chat_id and user info.""" + import asyncio + adapter = self._make_adapter() + captured_events = [] + + async def capture_handle(event): + captured_events.append(event) + + adapter.handle_message = capture_handle + + msg_data = { + "uid": b"6", + "sender_addr": "john@example.com", + "sender_name": "John Doe", + "subject": "Re: hi", + "message_id": "", + "in_reply_to": "", + "body": "Hello", + "attachments": [], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + event = captured_events[0] + self.assertEqual(event.source.chat_id, "john@example.com") + self.assertEqual(event.source.user_id, "john@example.com") + self.assertEqual(event.source.user_name, "John Doe") + self.assertEqual(event.source.chat_type, "dm") + + def test_non_allowlisted_sender_dropped(self): + """Senders not in EMAIL_ALLOWED_USERS should be dropped before dispatch.""" + import asyncio + with patch.dict(os.environ, { + "EMAIL_ALLOWED_USERS": "hermes@test.com,admin@test.com", + }): + adapter = self._make_adapter() + adapter._message_handler = MagicMock() + + msg_data = { + "uid": b"99", + "sender_addr": "outsider@evil.com", + "sender_name": "Spammer", + "subject": "Buy now!!!", + "message_id": "", + "in_reply_to": "", + "body": "Cheap meds", + "attachments": [], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + # Handler should NOT be called for non-allowlisted sender + adapter._message_handler.assert_not_called() + # Thread context should NOT be created + self.assertNotIn("outsider@evil.com", adapter._thread_context) + + def test_allowlisted_sender_proceeds(self): + """Senders in EMAIL_ALLOWED_USERS should proceed to dispatch normally.""" + import asyncio + with patch.dict(os.environ, { + "EMAIL_ALLOWED_USERS": "hermes@test.com,admin@test.com", + }): + adapter = self._make_adapter() + captured_events = [] + + async def mock_handler(event): + captured_events.append(event) + return None + + adapter._message_handler = mock_handler + + msg_data = { + "uid": b"100", + "sender_addr": "admin@test.com", + "sender_name": "Admin", + "subject": "Important", + "message_id": "", + "in_reply_to": "", + "body": "Hello", + "attachments": [], + "date": "", + # Authenticated From: (SPF/DKIM/DMARC passed at the receiving + # server). Allowlisted senders must be authenticated to proceed. + "sender_authenticated": True, + "auth_reason": "dmarc=pass", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + self.assertEqual(len(captured_events), 1) + self.assertEqual(captured_events[0].source.chat_id, "admin@test.com") def test_empty_allowlist_denies_without_optin(self): """No allowlist and no allow-all opt-in → adapter fails closed (2.6).""" @@ -293,6 +606,124 @@ def test_empty_allowlist_denies_without_optin(self): # Fail closed: an unset allowlist without allow-all drops the sender. adapter._message_handler.assert_not_called() + def test_empty_allowlist_allows_all_with_optin(self): + """EMAIL_ALLOW_ALL_USERS=true with no allowlist → all senders proceed.""" + import asyncio + with patch.dict(os.environ, {"EMAIL_ALLOW_ALL_USERS": "true"}, clear=False): + os.environ.pop("EMAIL_ALLOWED_USERS", None) + + adapter = self._make_adapter() + adapter._message_handler = MagicMock() + + msg_data = { + "uid": b"101", + "sender_addr": "anyone@test.com", + "sender_name": "Anyone", + "subject": "Hey", + "message_id": "", + "in_reply_to": "", + "body": "Hi", + "attachments": [], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + # With explicit allow-all opt-in the handler is called. + adapter._message_handler.assert_called() + + def test_spoofed_from_rejected_when_allowlisted(self): + """A forged From: matching the allowlist is dropped when unauthenticated. + + Core of GHSA-rxqh-5572-8m77: an attacker forges From: an-allowlisted + address. With an allowlist in effect and no allow-all, an unauthenticated + From: must be rejected before it can be matched against the allowlist. + """ + import asyncio + with patch.dict(os.environ, { + "EMAIL_ALLOWED_USERS": "admin@test.com", + "EMAIL_ALLOW_ALL_USERS": "", + "GATEWAY_ALLOW_ALL_USERS": "", + }): + adapter = self._make_adapter() + adapter._message_handler = MagicMock() + + msg_data = { + "uid": b"200", + "sender_addr": "admin@test.com", # forged From: matching allowlist + "sender_name": "Admin", + "subject": "Spoofed", + "message_id": "", + "in_reply_to": "", + "body": "rm -rf /", + "attachments": [], + "date": "", + "sender_authenticated": False, # SPF/DKIM/DMARC did not pass + "auth_reason": "authentication failed (spf=fail)", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + adapter._message_handler.assert_not_called() + self.assertNotIn("admin@test.com", adapter._thread_context) + + def test_unauthenticated_denied_without_allowlist_optin(self): + """No allowlist, no allow-all → adapter fails closed regardless of From auth.""" + import asyncio + with patch.dict(os.environ, {}, clear=False): + for k in ("EMAIL_ALLOWED_USERS", "GATEWAY_ALLOWED_USERS", + "EMAIL_ALLOW_ALL_USERS", "GATEWAY_ALLOW_ALL_USERS"): + os.environ.pop(k, None) + adapter = self._make_adapter() + adapter._message_handler = MagicMock() + + msg_data = { + "uid": b"201", + "sender_addr": "anyone@test.com", + "sender_name": "Anyone", + "subject": "Hi", + "message_id": "", + "in_reply_to": "", + "body": "Hi", + "attachments": [], + "date": "", + "sender_authenticated": False, + "auth_reason": "no Authentication-Results header", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + # Fail closed at the adapter — no allowlist and no allow-all opt-in. + adapter._message_handler.assert_not_called() + + def test_unauthenticated_allowed_with_trust_from_header(self): + """EMAIL_TRUST_FROM_HEADER=true disables the gate even with an allowlist.""" + import asyncio + with patch.dict(os.environ, { + "EMAIL_ALLOWED_USERS": "admin@test.com", + "EMAIL_TRUST_FROM_HEADER": "true", + }): + adapter = self._make_adapter() + captured = [] + + async def capture_handle(event): + captured.append(event) + + adapter.handle_message = capture_handle + + msg_data = { + "uid": b"202", + "sender_addr": "admin@test.com", + "sender_name": "Admin", + "subject": "Trusted", + "message_id": "", + "in_reply_to": "", + "body": "Hello", + "attachments": [], + "date": "", + "sender_authenticated": False, + "auth_reason": "no Authentication-Results header", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + self.assertEqual(len(captured), 1) def test_unauthenticated_allowed_with_allow_all(self): """EMAIL_ALLOW_ALL_USERS=true makes sender identity moot — gate skipped. @@ -360,6 +791,33 @@ def _make_adapter(self): adapter = EmailAdapter(PlatformConfig(enabled=True)) return adapter + def test_thread_context_stored_after_dispatch(self): + """After dispatching a message, thread context should be stored.""" + import asyncio + adapter = self._make_adapter() + + async def noop_handle(event): + pass + + adapter.handle_message = noop_handle + + msg_data = { + "uid": b"10", + "sender_addr": "user@test.com", + "sender_name": "User", + "subject": "Project question", + "message_id": "", + "in_reply_to": "", + "body": "Hello", + "attachments": [], + "date": "", + } + + asyncio.run(adapter._dispatch_message(msg_data)) + ctx = adapter._thread_context.get("user@test.com") + self.assertIsNotNone(ctx) + self.assertEqual(ctx["subject"], "Project question") + self.assertEqual(ctx["message_id"], "") def test_reply_uses_re_prefix(self): """Reply subject should have Re: prefix.""" @@ -382,6 +840,38 @@ def test_reply_uses_re_prefix(self): self.assertEqual(send_call["References"], "") self.assertIn("Date", send_call) + def test_reply_does_not_double_re(self): + """If subject already has Re:, don't add another.""" + adapter = self._make_adapter() + adapter._thread_context["user@test.com"] = { + "subject": "Re: Project question", + "message_id": "", + } + + with patch("smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value = mock_server + + adapter._send_email("user@test.com", "Follow up.", None) + + send_call = mock_server.send_message.call_args[0][0] + self.assertEqual(send_call["Subject"], "Re: Project question") + self.assertFalse(send_call["Subject"].startswith("Re: Re:")) + + def test_no_thread_context_uses_default_subject(self): + """Without thread context, subject should be 'Re: Hermes Agent'.""" + adapter = self._make_adapter() + + with patch("smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value = mock_server + + adapter._send_email("newuser@test.com", "Hello!", None) + + send_call = mock_server.send_message.call_args[0][0] + self.assertEqual(send_call["Subject"], "Re: Hermes Agent") + self.assertIn("Date", send_call) + class TestSendMethods(unittest.TestCase): """Test email send methods.""" @@ -398,6 +888,55 @@ def _make_adapter(self): adapter = EmailAdapter(PlatformConfig(enabled=True)) return adapter + def test_send_calls_smtp(self): + """send() should use SMTP to deliver email.""" + import asyncio + adapter = self._make_adapter() + + with patch("smtplib.SMTP") as mock_smtp: + mock_server = MagicMock() + mock_smtp.return_value = mock_server + + result = asyncio.run( + adapter.send("user@test.com", "Hello from Hermes!") + ) + + self.assertTrue(result.success) + mock_server.starttls.assert_called_once() + mock_server.login.assert_called_once_with("hermes@test.com", "secret") + mock_server.send_message.assert_called_once() + mock_server.quit.assert_called_once() + + def test_send_failure_returns_error(self): + """SMTP failure should return SendResult with error.""" + import asyncio + adapter = self._make_adapter() + + with patch("smtplib.SMTP") as mock_smtp: + mock_smtp.side_effect = Exception("Connection refused") + + result = asyncio.run( + adapter.send("user@test.com", "Hello") + ) + + self.assertFalse(result.success) + self.assertIn("Connection refused", result.error) + + def test_send_image_includes_url(self): + """send_image should include image URL in email body.""" + import asyncio + adapter = self._make_adapter() + + adapter.send = AsyncMock(return_value=SendResult(success=True)) + + asyncio.run( + adapter.send_image("user@test.com", "https://img.com/photo.jpg", "My photo") + ) + + call_args = adapter.send.call_args + body = call_args[0][1] + self.assertIn("https://img.com/photo.jpg", body) + self.assertIn("My photo", body) def test_send_document_with_attachment(self): """send_document should send email with file attachment.""" @@ -431,6 +970,12 @@ def test_send_document_with_attachment(self): finally: os.unlink(tmp_path) + def test_send_typing_is_noop(self): + """send_typing should do nothing for email.""" + import asyncio + adapter = self._make_adapter() + # Should not raise + asyncio.run(adapter.send_typing("user@test.com")) def test_get_chat_info(self): """get_chat_info should return email address as chat info.""" @@ -486,6 +1031,44 @@ def test_connect_success(self): if adapter._poll_task: adapter._poll_task.cancel() + def test_connect_imap_failure(self): + """IMAP connection failure returns False.""" + import asyncio + adapter = self._make_adapter() + + with patch("imaplib.IMAP4_SSL", side_effect=Exception("IMAP down")): + result = asyncio.run(adapter.connect()) + self.assertFalse(result) + self.assertFalse(adapter._running) + + def test_connect_smtp_failure(self): + """SMTP connection failure returns False.""" + import asyncio + adapter = self._make_adapter() + + mock_imap = MagicMock() + mock_imap.uid.return_value = ("OK", [b""]) + + with patch("imaplib.IMAP4_SSL", return_value=mock_imap), \ + patch("smtplib.SMTP", side_effect=Exception("SMTP down")): + result = asyncio.run(adapter.connect()) + self.assertFalse(result) + + def test_disconnect_cancels_poll(self): + """disconnect() should cancel the polling task.""" + import asyncio + adapter = self._make_adapter() + adapter._running = True + + async def _exercise_disconnect(): + adapter._poll_task = asyncio.create_task(asyncio.sleep(100)) + await adapter.disconnect() + + asyncio.run(_exercise_disconnect()) + + self.assertFalse(adapter._running) + self.assertIsNone(adapter._poll_task) + class TestFetchNewMessages(unittest.TestCase): """Test IMAP message fetching logic.""" @@ -531,6 +1114,54 @@ def uid_handler(command, *args): self.assertEqual(results[0]["sender_addr"], "user@test.com") self.assertIn(b"3", adapter._seen_uids) + def test_fetch_no_unseen_messages(self): + """No unseen messages returns empty list.""" + adapter = self._make_adapter() + + mock_imap = MagicMock() + mock_imap.uid.return_value = ("OK", [b""]) + + with patch("imaplib.IMAP4_SSL", return_value=mock_imap): + results = adapter._fetch_new_messages() + + self.assertEqual(results, []) + + def test_fetch_handles_imap_error(self): + """IMAP errors should be caught and return empty list.""" + adapter = self._make_adapter() + + with patch("imaplib.IMAP4_SSL", side_effect=Exception("Network error")): + results = adapter._fetch_new_messages() + + self.assertEqual(results, []) + + def test_fetch_extracts_sender_name(self): + """Sender name should be extracted from 'Name ' format.""" + adapter = self._make_adapter() + + raw_email = MIMEText("Hello", "plain", "utf-8") + raw_email["From"] = '"John Doe" ' + raw_email["Subject"] = "Test" + raw_email["Message-ID"] = "" + + mock_imap = MagicMock() + + def uid_handler(command, *args): + if command == "search": + return ("OK", [b"1"]) + if command == "fetch": + return ("OK", [(b"1", raw_email.as_bytes())]) + return ("NO", []) + + mock_imap.uid.side_effect = uid_handler + + with patch("imaplib.IMAP4_SSL", return_value=mock_imap): + results = adapter._fetch_new_messages() + + self.assertEqual(len(results), 1) + self.assertEqual(results[0]["sender_addr"], "john@test.com") + self.assertEqual(results[0]["sender_name"], "John Doe") + class TestPollLoop(unittest.TestCase): """Test the async polling loop.""" @@ -618,6 +1249,43 @@ async def _send_email(extra, chat_id, message): self.assertEqual(send_call["To"], "user@test.com") self.assertEqual(send_call["From"], "hermes@test.com") + @patch.dict(os.environ, { + "EMAIL_ADDRESS": "hermes@test.com", + "EMAIL_PASSWORD": "secret", + "EMAIL_SMTP_HOST": "smtp.test.com", + }) + def test_send_email_tool_failure(self): + """SMTP failure should return error dict.""" + import asyncio + from plugins.platforms.email.adapter import _standalone_send as _email_send + from types import SimpleNamespace + async def _send_email(extra, chat_id, message): + return await _email_send(SimpleNamespace(token=None, api_key=None, extra=extra or {}), chat_id, message) + + with patch("smtplib.SMTP", side_effect=Exception("SMTP error")): + result = asyncio.run( + _send_email({"address": "hermes@test.com", "smtp_host": "smtp.test.com"}, "user@test.com", "Hello") + ) + + self.assertIn("error", result) + self.assertIn("SMTP error", result["error"]) + + @patch.dict(os.environ, {}, clear=True) + def test_send_email_tool_not_configured(self): + """Missing config should return error.""" + import asyncio + from plugins.platforms.email.adapter import _standalone_send as _email_send + from types import SimpleNamespace + async def _send_email(extra, chat_id, message): + return await _email_send(SimpleNamespace(token=None, api_key=None, extra=extra or {}), chat_id, message) + + result = asyncio.run( + _send_email({}, "user@test.com", "Hello") + ) + + self.assertIn("error", result) + self.assertIn("not configured", result["error"]) + class TestSmtpConnectionCleanup(unittest.TestCase): """Verify SMTP connections are closed even when send_message raises.""" @@ -634,6 +1302,24 @@ def _make_adapter(self): from plugins.platforms.email.adapter import EmailAdapter return EmailAdapter(PlatformConfig(enabled=True)) + @patch.dict(os.environ, { + "EMAIL_ADDRESS": "hermes@test.com", + "EMAIL_PASSWORD": "secret", + "EMAIL_IMAP_HOST": "imap.test.com", + "EMAIL_SMTP_HOST": "smtp.test.com", + "EMAIL_SMTP_PORT": "587", + }, clear=False) + def test_smtp_quit_called_on_send_message_failure(self): + """SMTP quit() must be called even when send_message() raises.""" + adapter = self._make_adapter() + mock_smtp = MagicMock() + mock_smtp.send_message.side_effect = Exception("send failed") + + with patch("smtplib.SMTP", return_value=mock_smtp): + with self.assertRaises(Exception): + adapter._send_email("user@test.com", "Hello") + + mock_smtp.quit.assert_called_once() @patch.dict(os.environ, { "EMAIL_ADDRESS": "hermes@test.com", @@ -698,6 +1384,25 @@ def uid_handler(command, *args): self.assertEqual(results, []) mock_imap.logout.assert_called_once() + @patch.dict(os.environ, { + "EMAIL_ADDRESS": "hermes@test.com", + "EMAIL_PASSWORD": "secret", + "EMAIL_IMAP_HOST": "imap.test.com", + "EMAIL_IMAP_PORT": "993", + "EMAIL_SMTP_HOST": "smtp.test.com", + }, clear=False) + def test_imap_logout_called_on_early_return(self): + """IMAP logout() must be called even when returning early (no unseen).""" + adapter = self._make_adapter() + mock_imap = MagicMock() + mock_imap.uid.return_value = ("OK", [b""]) + + with patch("imaplib.IMAP4_SSL", return_value=mock_imap): + results = adapter._fetch_new_messages() + + self.assertEqual(results, []) + mock_imap.logout.assert_called_once() + class TestImapIdExtensionForNetEase(unittest.TestCase): """Regression for #22271: 163/NetEase mailbox requires the RFC 2971 @@ -747,6 +1452,32 @@ def test_connect_sends_imap_id_after_login(self): self.assertIn("login", names) self.assertLess(names.index("login"), names.index("xatom")) + def test_fetch_new_messages_sends_imap_id_after_login(self): + """_fetch_new_messages must also send ID — it opens its own IMAP session.""" + adapter = self._make_adapter() + mock_imap = MagicMock() + mock_imap.uid.return_value = ("OK", [b""]) + + with patch("imaplib.IMAP4_SSL", return_value=mock_imap): + adapter._fetch_new_messages() + + id_calls = [c for c in mock_imap.xatom.call_args_list if c.args and c.args[0] == "ID"] + self.assertTrue( + id_calls, + "_fetch_new_messages() must call imap.xatom('ID', ...) after " + "LOGIN — the polling path opens a fresh IMAP connection.", + ) + + def test_send_imap_id_swallows_errors_for_non_supporting_servers(self): + """Servers that reject ID must not break the connection.""" + from plugins.platforms.email.adapter import _send_imap_id + + mock_imap = MagicMock() + mock_imap.xatom.side_effect = Exception("BAD command unknown: ID") + + _send_imap_id(mock_imap) + mock_imap.xatom.assert_called_once() + class TestConnectSmtp(unittest.TestCase): """Test _connect_smtp() helper: protocol selection and IPv6 fallback.""" @@ -763,6 +1494,36 @@ def _make_adapter(self, port="587"): from plugins.platforms.email.adapter import EmailAdapter return EmailAdapter(PlatformConfig(enabled=True)) + def test_port_587_uses_smtp_with_starttls(self): + """Port 587 should use smtplib.SMTP + STARTTLS.""" + adapter = self._make_adapter("587") + + with patch("smtplib.SMTP") as mock_smtp, \ + patch("smtplib.SMTP_SSL") as mock_smtp_ssl: + mock_server = MagicMock() + mock_smtp.return_value = mock_server + + result = adapter._connect_smtp() + + mock_smtp.assert_called_once() + mock_smtp_ssl.assert_not_called() + mock_server.starttls.assert_called_once() + self.assertIs(result, mock_server) + + def test_port_465_uses_smtp_ssl(self): + """Port 465 should use smtplib.SMTP_SSL (implicit TLS).""" + adapter = self._make_adapter("465") + + with patch("smtplib.SMTP") as mock_smtp, \ + patch("smtplib.SMTP_SSL") as mock_smtp_ssl: + mock_server = MagicMock() + mock_smtp_ssl.return_value = mock_server + + result = adapter._connect_smtp() + + mock_smtp_ssl.assert_called_once() + mock_smtp.assert_not_called() + self.assertIs(result, mock_server) def test_ipv6_timeout_falls_back_to_ipv4(self): """When default connection times out, retry with an IPv4-only SMTP path.""" @@ -801,10 +1562,83 @@ def test_port_465_ipv6_fallback(self): "smtp.test.com", 465, timeout=30, context=ANY, ) + def test_tls_verification_error_does_not_retry_ipv4(self): + """Certificate failures are security errors, not IPv6 reachability failures.""" + import ssl as _ssl + import plugins.platforms.email.adapter as email_mod + + adapter = self._make_adapter("465") + + with patch("smtplib.SMTP_SSL", side_effect=_ssl.SSLError("cert verify failed")), \ + patch.object(email_mod, "_IPv4SMTP_SSL") as mock_ipv4_smtp_ssl: + with self.assertRaises(_ssl.SSLError): + adapter._connect_smtp() + + mock_ipv4_smtp_ssl.assert_not_called() + + def test_ipv4_connection_does_not_mutate_global_resolver(self): + """IPv4 fallback must not monkeypatch process-global socket state.""" + import socket as _socket + from plugins.platforms.email.adapter import _create_ipv4_connection + + original_getaddrinfo = _socket.getaddrinfo + fake_sock = MagicMock() + + with patch( + "socket.getaddrinfo", + return_value=[(_socket.AF_INET, _socket.SOCK_STREAM, 6, "", ("192.0.2.1", 587))], + ) as mock_getaddrinfo, patch("socket.socket", return_value=fake_sock): + result = _create_ipv4_connection("smtp.test.com", 587, 30) + + self.assertIs(result, fake_sock) + mock_getaddrinfo.assert_called_once_with( + "smtp.test.com", 587, _socket.AF_INET, _socket.SOCK_STREAM, + ) + self.assertIs(_socket.getaddrinfo, original_getaddrinfo) + class TestConnectionConfigResolution(unittest.TestCase): """Host/address resolution and pre-connect validation (#49736).""" + def test_host_and_address_whitespace_stripped(self): + """A stray space/newline must not reach IMAP4_SSL as part of the host. + + Whitespace in the host produced the misleading + ``[Errno 8] nodename nor servname`` (unresolvable name) instead of a + successful connection. + """ + from gateway.config import PlatformConfig + from plugins.platforms.email.adapter import EmailAdapter + with patch.dict(os.environ, { + "EMAIL_ADDRESS": " hermes@test.com\n", + "EMAIL_PASSWORD": "secret", + "EMAIL_IMAP_HOST": " imap.test.com ", + "EMAIL_SMTP_HOST": "smtp.test.com\n", + }, clear=False): + adapter = EmailAdapter(PlatformConfig(enabled=True)) + self.assertEqual(adapter._imap_host, "imap.test.com") + self.assertEqual(adapter._smtp_host, "smtp.test.com") + self.assertEqual(adapter._address, "hermes@test.com") + + def test_falls_back_to_platform_config_extra(self): + """When env vars are absent, settings come from PlatformConfig.extra — + the same dict gateway.config populates and `hermes config show` reads.""" + from gateway.config import PlatformConfig + from plugins.platforms.email.adapter import EmailAdapter + cfg = PlatformConfig(enabled=True) + cfg.extra.update({ + "address": "hermes@test.com", + "imap_host": "imap.test.com", + "smtp_host": "smtp.test.com", + }) + with patch.dict(os.environ, { + "EMAIL_ADDRESS": "", "EMAIL_IMAP_HOST": "", "EMAIL_SMTP_HOST": "", + "EMAIL_PASSWORD": "secret", + }, clear=False): + adapter = EmailAdapter(cfg) + self.assertEqual(adapter._imap_host, "imap.test.com") + self.assertEqual(adapter._smtp_host, "smtp.test.com") + self.assertEqual(adapter._address, "hermes@test.com") def test_connect_aborts_without_attempting_imap_when_host_missing(self): """A missing host returns False without the cryptic DNS error, and marks @@ -843,6 +1677,15 @@ def test_blank_present_env_vars_are_not_required(self): }, clear=False): self.assertFalse(check_email_requirements()) + def test_all_settings_present_satisfies_requirements(self): + """The connected check passes only when all four settings are non-blank.""" + from plugins.platforms.email.adapter import check_email_requirements + with patch.dict(os.environ, { + "EMAIL_ADDRESS": "hermes@test.com", "EMAIL_PASSWORD": "secret", + "EMAIL_IMAP_HOST": "imap.test.com", "EMAIL_SMTP_HOST": "smtp.test.com", + }, clear=False): + self.assertTrue(check_email_requirements()) + class TestSenderAuthentication(unittest.TestCase): """Verify _verify_sender_authentication parses Authentication-Results @@ -873,6 +1716,12 @@ def test_dmarc_pass_authenticates(self): ) self.assertTrue(ok, reason) + def test_spf_pass_aligned_authenticates(self): + ok, reason = self._verify( + "admin@example.com", + ["mx.google.com; spf=pass smtp.mailfrom=admin@example.com"], + ) + self.assertTrue(ok, reason) def test_dkim_pass_aligned_authenticates(self): ok, reason = self._verify( @@ -889,6 +1738,32 @@ def test_spf_pass_misaligned_rejected(self): ) self.assertFalse(ok, reason) + def test_dkim_pass_misaligned_rejected(self): + ok, reason = self._verify( + "admin@example.com", + ["mx.google.com; dkim=pass header.d=evil.com"], + ) + self.assertFalse(ok, reason) + + def test_all_fail_rejected(self): + ok, reason = self._verify( + "admin@example.com", + ["mx.google.com; dmarc=fail; spf=fail; dkim=fail"], + ) + self.assertFalse(ok, reason) + + def test_no_authentication_results_rejected(self): + ok, reason = self._verify("admin@example.com", []) + self.assertFalse(ok) + self.assertIn("no Authentication-Results", reason) + + def test_relaxed_alignment_subdomain(self): + # mail.example.com (DKIM signer) aligns with example.com (From). + ok, reason = self._verify( + "admin@example.com", + ["mx.google.com; dkim=pass header.d=mail.example.com"], + ) + self.assertTrue(ok, reason) def test_injected_header_below_trusted_does_not_authenticate(self): """An attacker-injected Authentication-Results sorts BELOW the receiving @@ -906,6 +1781,15 @@ def test_injected_header_below_trusted_does_not_authenticate(self): ) self.assertFalse(ok, reason) + def test_authserv_id_mismatch_skips_untrusted_header(self): + """A header from an authserv-id we don't trust is skipped entirely.""" + ok, reason = self._verify( + "admin@example.com", + ["attacker.relay.com; dmarc=pass header.from=example.com"], + authserv_id="mx.ourserver.com", + ) + self.assertFalse(ok, reason) + if __name__ == "__main__": unittest.main() diff --git a/tests/gateway/test_feishu.py b/tests/gateway/test_feishu.py index a923ea5f5d46..e3e02059ee0b 100644 --- a/tests/gateway/test_feishu.py +++ b/tests/gateway/test_feishu.py @@ -15,6 +15,22 @@ from gateway.platforms.base import ProcessingOutcome +# Windows-safe env for ``@patch.dict(..., clear=True)`` tests: the decorator +# wipes the whole environment, and on Windows ``pathlib.Path.home()`` raises +# "Could not determine home directory" when USERPROFILE / HOME / +# HOMEDRIVE+HOMEPATH are all absent (POSIX falls back to the ``pwd`` module, +# which is why these tests only broke on Windows). ``get_hermes_home()`` +# reads HERMES_HOME first, then LOCALAPPDATA, then Path.home(). Preserving +# these variables keeps the "no Feishu env vars" semantics of clear=True +# while letting home resolution work on every platform. +_HOME_ENV = { + k: v + for k, v in os.environ.items() + if k in ("HERMES_HOME", "USERPROFILE", "HOMEDRIVE", "HOMEPATH", "HOME", "LOCALAPPDATA") + and v +} + + try: import lark_oapi _HAS_LARK_OAPI = True @@ -47,7 +63,7 @@ def _mock_event_dispatcher_builder(mock_handler_class): class TestConfigEnvOverrides(unittest.TestCase): - @patch.dict(os.environ, { + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_APP_ID": "cli_xxx", "FEISHU_APP_SECRET": "secret_xxx", "FEISHU_CONNECTION_MODE": "websocket", @@ -64,9 +80,82 @@ def test_feishu_config_loaded_from_env(self): self.assertEqual(config.platforms[Platform.FEISHU].extra["app_id"], "cli_xxx") self.assertEqual(config.platforms[Platform.FEISHU].extra["connection_mode"], "websocket") + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_xxx", + "FEISHU_APP_SECRET": "secret_xxx", + "FEISHU_HOME_CHANNEL": "oc_xxx", + }, clear=False) + def test_feishu_home_channel_loaded(self): + from gateway.config import GatewayConfig, Platform, _apply_env_overrides + + config = GatewayConfig() + _apply_env_overrides(config) + + home = config.platforms[Platform.FEISHU].home_channel + self.assertIsNotNone(home) + self.assertEqual(home.chat_id, "oc_xxx") + + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_xxx", + "FEISHU_APP_SECRET": "secret_xxx", + }, clear=False) + def test_feishu_in_connected_platforms(self): + from gateway.config import GatewayConfig, Platform, _apply_env_overrides + + config = GatewayConfig() + _apply_env_overrides(config) + + self.assertIn(Platform.FEISHU, config.get_connected_platforms()) + class TestFeishuMessageNormalization(unittest.TestCase): + def test_normalize_merge_forward_preserves_summary_lines(self): + from plugins.platforms.feishu.adapter import normalize_feishu_message + + normalized = normalize_feishu_message( + message_type="merge_forward", + raw_content=json.dumps( + { + "title": "Sprint recap", + "messages": [ + {"sender_name": "Alice", "text": "Please review PR-128"}, + { + "sender_name": "Bob", + "message_type": "post", + "content": { + "en_us": { + "content": [[{"tag": "text", "text": "Ship it"}]], + } + }, + }, + ], + } + ), + ) + + self.assertEqual(normalized.relation_kind, "merge_forward") + self.assertEqual( + normalized.text_content, + "Sprint recap\n- Alice: Please review PR-128\n- Bob: Ship it", + ) + + def test_normalize_share_chat_exposes_summary_and_metadata(self): + from plugins.platforms.feishu.adapter import normalize_feishu_message + + normalized = normalize_feishu_message( + message_type="share_chat", + raw_content=json.dumps( + { + "chat_id": "oc_chat_shared", + "chat_name": "Backend Guild", + } + ), + ) + self.assertEqual(normalized.relation_kind, "share_chat") + self.assertEqual(normalized.text_content, "Shared chat: Backend Guild\nChat ID: oc_chat_shared") + self.assertEqual(normalized.metadata["chat_id"], "oc_chat_shared") + self.assertEqual(normalized.metadata["chat_name"], "Backend Guild") def test_normalize_interactive_card_preserves_title_body_and_actions(self): from plugins.platforms.feishu.adapter import normalize_feishu_message @@ -106,7 +195,7 @@ def test_websocket_sdk_accepts_channel_ua_tag(self): """The shipped SDK must support the Channel signaling argument. Guarded on the pinned version: the repo pins lark-oapi==1.6.8 (the - first release with ``extra_ua_tags``). Dev machines can carry an + first release with extra_ua_tags). Dev machines can carry an older lazy-installed lark-oapi that predates the argument — that is an environment artifact, not a product regression, so skip rather than fail there. Environments installing the pin (the feishu extra) @@ -128,6 +217,96 @@ def test_websocket_sdk_accepts_channel_ua_tag(self): signature = inspect.signature(FeishuWSClient) self.assertIn("extra_ua_tags", signature.parameters) + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_app", + "FEISHU_APP_SECRET": "secret_app", + "FEISHU_CONNECTION_MODE": "webhook", + "FEISHU_WEBHOOK_HOST": "127.0.0.1", + "FEISHU_WEBHOOK_PORT": "9001", + "FEISHU_WEBHOOK_PATH": "/hook", + "FEISHU_VERIFICATION_TOKEN": "vtok", + }, clear=True) + def test_connect_webhook_mode_starts_local_server(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + runner = AsyncMock() + site = AsyncMock() + web_module = SimpleNamespace( + Application=lambda **_kwargs: SimpleNamespace(router=SimpleNamespace(add_post=lambda *_args, **_kwargs: None)), + AppRunner=lambda _app: runner, + TCPSite=lambda _runner, host, port: SimpleNamespace(start=site.start, host=host, port=port), + ) + + with ( + patch("plugins.platforms.feishu.adapter.FEISHU_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.FEISHU_WEBHOOK_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.EventDispatcherHandler") as mock_handler_class, + patch("plugins.platforms.feishu.adapter.acquire_scoped_lock", return_value=(True, None)), + patch("plugins.platforms.feishu.adapter.release_scoped_lock"), + patch.object(adapter, "_hydrate_bot_identity", new=AsyncMock()), + patch.object(adapter, "_build_lark_client", return_value=SimpleNamespace()), + patch("plugins.platforms.feishu.adapter.web", web_module), + ): + _mock_event_dispatcher_builder(mock_handler_class) + connected = asyncio.run(adapter.connect()) + + self.assertTrue(connected) + runner.setup.assert_awaited_once() + site.start.assert_awaited_once() + + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_app", + "FEISHU_APP_SECRET": "secret_app", + }, clear=True) + def test_connect_acquires_scoped_lock_and_disconnect_releases_it(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + ws_client = SimpleNamespace() + + with ( + patch("plugins.platforms.feishu.adapter.FEISHU_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.FEISHU_WEBSOCKET_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.lark", SimpleNamespace(LogLevel=SimpleNamespace(INFO="INFO", WARNING="WARNING"))), + patch("plugins.platforms.feishu.adapter.EventDispatcherHandler") as mock_handler_class, + patch("plugins.platforms.feishu.adapter.FeishuWSClient", return_value=ws_client), + patch("plugins.platforms.feishu.adapter._run_official_feishu_ws_client"), + patch("plugins.platforms.feishu.adapter.acquire_scoped_lock", return_value=(True, None)) as acquire_lock, + patch("plugins.platforms.feishu.adapter.release_scoped_lock") as release_lock, + patch.object(adapter, "_hydrate_bot_identity", new=AsyncMock()), + patch.object(adapter, "_build_lark_client", return_value=SimpleNamespace()), + ): + _mock_event_dispatcher_builder(mock_handler_class) + + loop = asyncio.new_event_loop() + future = loop.create_future() + future.set_result(None) + + class _Loop: + def run_in_executor(self, *_args, **_kwargs): + return future + + def is_closed(self): + return False + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.get_running_loop", return_value=_Loop()): + connected = asyncio.run(adapter.connect()) + asyncio.run(adapter.disconnect()) + finally: + loop.close() + + self.assertTrue(connected) + self.assertIsNone(adapter._event_handler) + acquire_lock.assert_called_once_with( + "feishu-app-id", + "cli_app", + metadata={"platform": "feishu"}, + ) + release_lock.assert_called_once_with("feishu-app-id", "cli_app") def test_disconnect_sends_websocket_close_frame(self): """Regression test for #10202: disconnect() must call the WSS @@ -181,8 +360,105 @@ async def _fake_disconnect() -> None: # _disable_websocket_auto_reconnect() must still run. self.assertIsNone(adapter._ws_client) + def test_disconnect_tolerates_missing_internal_disconnect(self): + """If the lark_oapi client layout changes and ``_disconnect`` + disappears, disconnect() must not raise — fall through to the + existing task-cancel path. + """ + from types import SimpleNamespace + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + # No ``_disconnect`` attribute — ``hasattr`` guard should skip. + adapter._ws_client = SimpleNamespace(_auto_reconnect=True) + adapter._ws_thread_loop = None + adapter._ws_future = None + + # Must not raise. + asyncio.run(adapter.disconnect()) + self.assertIsNone(adapter._ws_client) + + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_app", + "FEISHU_APP_SECRET": "secret_app", + }, clear=True) + def test_connect_rejects_existing_app_lock(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + with ( + patch("plugins.platforms.feishu.adapter.FEISHU_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.FEISHU_WEBSOCKET_AVAILABLE", True), + patch( + "plugins.platforms.feishu.adapter.acquire_scoped_lock", + return_value=(False, {"pid": 4321}), + ), + ): + connected = asyncio.run(adapter.connect()) + + self.assertFalse(connected) + self.assertEqual(adapter.fatal_error_code, "feishu_app_lock") + self.assertFalse(adapter.fatal_error_retryable) + self.assertIn("PID 4321", adapter.fatal_error_message) + + @patch.dict(os.environ, {**_HOME_ENV, + "FEISHU_APP_ID": "cli_app", + "FEISHU_APP_SECRET": "secret_app", + }, clear=True) + def test_connect_retries_transient_startup_failure(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + ws_client = SimpleNamespace() + sleeps = [] + + with ( + patch("plugins.platforms.feishu.adapter.FEISHU_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.FEISHU_WEBSOCKET_AVAILABLE", True), + patch("plugins.platforms.feishu.adapter.lark", SimpleNamespace(LogLevel=SimpleNamespace(INFO="INFO", WARNING="WARNING"))), + patch("plugins.platforms.feishu.adapter.EventDispatcherHandler") as mock_handler_class, + patch("plugins.platforms.feishu.adapter.FeishuWSClient", return_value=ws_client), + patch("plugins.platforms.feishu.adapter.acquire_scoped_lock", return_value=(True, None)), + patch("plugins.platforms.feishu.adapter.release_scoped_lock"), + patch.object(adapter, "_hydrate_bot_identity", new=AsyncMock()), + patch("plugins.platforms.feishu.adapter.asyncio.sleep", side_effect=lambda delay: sleeps.append(delay)), + patch.object(adapter, "_build_lark_client", return_value=SimpleNamespace()), + ): + _mock_event_dispatcher_builder(mock_handler_class) + + loop = asyncio.new_event_loop() + future = loop.create_future() + future.set_result(None) + + class _Loop: + def __init__(self): + self.calls = 0 + + def run_in_executor(self, *_args, **_kwargs): + self.calls += 1 + if self.calls == 1: + raise OSError("temporary websocket failure") + return future + + def is_closed(self): + return False + + fake_loop = _Loop() + try: + with patch("plugins.platforms.feishu.adapter.asyncio.get_running_loop", return_value=fake_loop): + connected = asyncio.run(adapter.connect()) + finally: + loop.close() + + self.assertTrue(connected) + self.assertEqual(sleeps, [1]) + self.assertEqual(fake_loop.calls, 2) - @patch.dict(os.environ, { + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_APP_ID": "cli_app", "FEISHU_APP_SECRET": "secret_app", }, clear=True) @@ -241,8 +517,49 @@ def is_closed(self): self.assertEqual(call_kwargs["extra_ua_tags"], ["channel"], "extra_ua_tags must be ['channel'] to enable group event routing") + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_edit_message_updates_existing_feishu_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _MessageAPI: + def update(self, request): + captured["request"] = request + return SimpleNamespace(success=lambda: True) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.edit_message( + chat_id="oc_chat", + message_id="om_progress", + content="📖 read_file: \"/tmp/image.png\"", + ) + ) + + self.assertTrue(result.success) + self.assertEqual(result.message_id, "om_progress") + self.assertEqual(captured["request"].message_id, "om_progress") + self.assertEqual(captured["request"].request_body.msg_type, "text") + self.assertEqual( + captured["request"].request_body.content, + json.dumps({"text": "📖 read_file: \"/tmp/image.png\""}, ensure_ascii=False), + ) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_edit_message_falls_back_to_text_when_post_update_is_rejected(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -285,6 +602,40 @@ async def _direct(func, *args, **kwargs): json.dumps({"text": "可以用 粗体 和 斜体。"}, ensure_ascii=False), ) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_get_chat_info_uses_real_feishu_chat_api(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + class _ChatAPI: + def get(self, request): + self.request = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(name="Hermes Group", chat_type="group"), + ) + + chat_api = _ChatAPI() + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + chat=chat_api, + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + info = asyncio.run(adapter.get_chat_info("oc_chat")) + + self.assertEqual(chat_api.request.chat_id, "oc_chat") + self.assertEqual(info["chat_id"], "oc_chat") + self.assertEqual(info["name"], "Hermes Group") + self.assertEqual(info["type"], "group") class TestAdapterModule(unittest.TestCase): def test_load_settings_uses_sdk_defaults_for_invalid_ws_reconnect_values(self): @@ -300,6 +651,44 @@ def test_load_settings_uses_sdk_defaults_for_invalid_ws_reconnect_values(self): self.assertEqual(settings.ws_reconnect_nonce, 30) self.assertEqual(settings.ws_reconnect_interval, 120) + def test_load_settings_accepts_custom_ws_reconnect_values(self): + from plugins.platforms.feishu.adapter import FeishuAdapter + + settings = FeishuAdapter._load_settings( + { + "ws_reconnect_nonce": 0, + "ws_reconnect_interval": 3, + } + ) + + self.assertEqual(settings.ws_reconnect_nonce, 0) + self.assertEqual(settings.ws_reconnect_interval, 3) + + def test_load_settings_accepts_custom_ws_ping_values(self): + from plugins.platforms.feishu.adapter import FeishuAdapter + + settings = FeishuAdapter._load_settings( + { + "ws_ping_interval": 10, + "ws_ping_timeout": 8, + } + ) + + self.assertEqual(settings.ws_ping_interval, 10) + self.assertEqual(settings.ws_ping_timeout, 8) + + def test_load_settings_ignores_invalid_ws_ping_values(self): + from plugins.platforms.feishu.adapter import FeishuAdapter + + settings = FeishuAdapter._load_settings( + { + "ws_ping_interval": 0, + "ws_ping_timeout": -1, + } + ) + + self.assertIsNone(settings.ws_ping_interval) + self.assertIsNone(settings.ws_ping_timeout) def test_runtime_ws_overrides_reapply_after_sdk_configure(self): import sys @@ -368,7 +757,7 @@ def _admits_group(adapter, message, sender_id, chat_id=""): class TestAdapterBehavior(unittest.TestCase): - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_build_event_handler_registers_reaction_and_card_processors(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -450,7 +839,7 @@ def builder(_encrypt_key, _verification_token): ], ) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_bot_origin_reactions_are_dropped_to_avoid_feedback_loops(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -471,7 +860,7 @@ def test_bot_origin_reactions_are_dropped_to_avoid_feedback_loops(self): adapter._on_reaction_event("im.message.reaction.created_v1", data) run_threadsafe.assert_not_called() - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_user_reaction_with_managed_emoji_is_still_routed(self): # Operator-origin filter is enough to prevent feedback loops; we must # not additionally swallow user-origin reactions just because their @@ -529,7 +918,7 @@ def _build_reaction_adapter(self, *, msg_sender_id: str): adapter.get_chat_info = AsyncMock(return_value={"name": "Test Chat"}) return adapter - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_reaction_on_peer_bot_message_is_not_routed(self): # GET im/v1/messages sender for bot messages carries id=app_id; a peer # bot's message has a different app_id than ours, so it must be dropped. @@ -546,15 +935,93 @@ def test_reaction_on_peer_bot_message_is_not_routed(self): ) adapter._handle_message_with_guards.assert_not_awaited() + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_reaction_on_our_own_bot_message_is_routed(self): + adapter = self._build_reaction_adapter(msg_sender_id="cli_self_app") - def test_per_group_allowlist_policy_gates_by_sender(self): + event = SimpleNamespace( + message_id="om_self_msg", + user_id=SimpleNamespace(open_id="ou_human", user_id=None, union_id=None), + reaction_type=SimpleNamespace(emoji_type="THUMBSUP"), + ) + data = SimpleNamespace(event=event) + asyncio.run( + adapter._handle_reaction_event("im.message.reaction.created_v1", data) + ) + adapter._handle_message_with_guards.assert_awaited_once() + + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "open"}, clear=True) + def test_group_message_requires_mentions_even_when_policy_open(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter - config = PlatformConfig( - extra={ - "group_rules": { - "oc_chat_a": { + adapter = FeishuAdapter(PlatformConfig()) + message = SimpleNamespace(mentions=[]) + sender_id = SimpleNamespace(open_id="ou_any", user_id=None) + self.assertFalse(_admits_group(adapter, message, sender_id, "")) + + message_with_mention = SimpleNamespace(mentions=[SimpleNamespace(key="@_user_1")]) + self.assertFalse(_admits_group(adapter, message_with_mention, sender_id, "")) + + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "open"}, clear=True) + def test_group_message_with_other_user_mention_is_rejected_when_bot_identity_unknown(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + sender_id = SimpleNamespace(open_id="ou_any", user_id=None) + other_mention = SimpleNamespace( + name="Other User", + id=SimpleNamespace(open_id="ou_other", user_id="u_other"), + ) + + self.assertFalse( + _admits_group(adapter, SimpleNamespace(mentions=[other_mention]), sender_id, "") + ) + + @patch.dict( + os.environ, + {**_HOME_ENV, "FEISHU_GROUP_POLICY": "allowlist", "FEISHU_ALLOWED_USERS": "ou_allowed", "FEISHU_BOT_NAME": "Hermes Bot"}, + clear=True, + ) + def test_group_message_allowlist_and_mention_both_required(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + # Mention without IDs — name fallback legitimately engages. + mentioned = SimpleNamespace( + mentions=[ + SimpleNamespace( + name="Hermes Bot", + id=SimpleNamespace(open_id=None, user_id=None), + ) + ] + ) + + self.assertTrue( + _admits_group(adapter, + mentioned, + SimpleNamespace(open_id="ou_allowed", user_id=None), + "", + ) + ) + self.assertFalse( + _admits_group(adapter, + mentioned, + SimpleNamespace(open_id="ou_blocked", user_id=None), + "", + ) + ) + + def test_per_group_allowlist_policy_gates_by_sender(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + config = PlatformConfig( + extra={ + "group_rules": { + "oc_chat_a": { "policy": "allowlist", "allowlist": ["ou_alice", "ou_bob"], } @@ -583,6 +1050,113 @@ def test_per_group_allowlist_policy_gates_by_sender(self): ) ) + def test_per_group_blacklist_policy_blocks_specific_users(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + config = PlatformConfig( + extra={ + "group_rules": { + "oc_chat_b": { + "policy": "blacklist", + "blacklist": ["ou_blocked"], + } + } + } + ) + adapter = FeishuAdapter(config) + adapter._bot_open_id = "ou_bot" + + message = SimpleNamespace( + mentions=[SimpleNamespace(name="Bot", id=SimpleNamespace(open_id="ou_bot", user_id=None))] + ) + + self.assertTrue( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_alice", user_id=None), + "oc_chat_b", + ) + ) + self.assertFalse( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_blocked", user_id=None), + "oc_chat_b", + ) + ) + + def test_per_group_admin_only_policy_requires_admin(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + config = PlatformConfig( + extra={ + "admins": ["ou_admin"], + "group_rules": { + "oc_chat_c": { + "policy": "admin_only", + } + }, + } + ) + adapter = FeishuAdapter(config) + adapter._bot_open_id = "ou_bot" + + message = SimpleNamespace( + mentions=[SimpleNamespace(name="Bot", id=SimpleNamespace(open_id="ou_bot", user_id=None))] + ) + + self.assertTrue( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_admin", user_id=None), + "oc_chat_c", + ) + ) + self.assertFalse( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_regular", user_id=None), + "oc_chat_c", + ) + ) + + def test_per_group_disabled_policy_blocks_all(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + config = PlatformConfig( + extra={ + "admins": ["ou_admin"], + "group_rules": { + "oc_chat_d": { + "policy": "disabled", + } + }, + } + ) + adapter = FeishuAdapter(config) + adapter._bot_open_id = "ou_bot" + + message = SimpleNamespace( + mentions=[SimpleNamespace(name="Bot", id=SimpleNamespace(open_id="ou_bot", user_id=None))] + ) + + self.assertTrue( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_admin", user_id=None), + "oc_chat_d", + ) + ) + self.assertFalse( + _admits_group(adapter, + message, + SimpleNamespace(open_id="ou_regular", user_id=None), + "oc_chat_d", + ) + ) def test_global_admins_bypass_all_group_rules(self): from gateway.config import PlatformConfig @@ -638,8 +1212,32 @@ def test_default_group_policy_fallback_for_chats_without_explicit_rule(self): ) ) + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "open"}, clear=True) + def test_group_message_matches_bot_open_id_when_configured(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._bot_open_id = "ou_bot" + sender_id = SimpleNamespace(open_id="ou_any", user_id=None) + + bot_mention = SimpleNamespace( + name="Hermes", + id=SimpleNamespace(open_id="ou_bot", user_id="u_bot"), + ) + other_mention = SimpleNamespace( + name="Other", + id=SimpleNamespace(open_id="ou_other", user_id="u_other"), + ) + + self.assertTrue( + _admits_group(adapter, SimpleNamespace(mentions=[bot_mention]), sender_id, "") + ) + self.assertFalse( + _admits_group(adapter, SimpleNamespace(mentions=[other_mention]), sender_id, "") + ) - @patch.dict(os.environ, {"FEISHU_GROUP_POLICY": "open"}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "open"}, clear=True) def test_group_message_matches_bot_name_when_only_name_available(self): """Name fallback engages when either side lacks an open_id. When BOTH the mention and the bot carry open_ids, IDs are authoritative — a @@ -696,8 +1294,71 @@ def test_group_message_matches_bot_name_when_only_name_available(self): _admits_group(adapter2, SimpleNamespace(mentions=[bot_mention]), sender_id, "") ) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_post_message_as_text(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + message = SimpleNamespace( + message_type="post", + content='{"zh_cn":{"title":"Title","content":[[{"tag":"text","text":"hello "}],[{"tag":"a","text":"doc","href":"https://example.com"}]]}}', + message_id="om_post", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "Title\nhello\n[doc](https://example.com)") + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_post_message_uses_first_available_language_block(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + message = SimpleNamespace( + message_type="post", + content='{"fr_fr":{"title":"Subject","content":[[{"tag":"text","text":"bonjour"}]]}}', + message_id="om_post_fr", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "Subject\nbonjour") + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_post_message_with_rich_elements_does_not_drop_content(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + message = SimpleNamespace( + message_type="post", + content=( + '{"en_us":{"title":"Rich message","content":[' + '[{"tag":"img","alt":"diagram"}],' + '[{"tag":"at","user_name":"Alice"},{"tag":"text","text":" please check the attachment"}],' + '[{"tag":"media","file_name":"spec.pdf"}],' + '[{"tag":"emotion","emoji_type":"smile"}]' + ']}}' + ), + message_id="om_post_rich", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "Rich message\n[Image: diagram]\n@Alice please check the attachment\n[Attachment: spec.pdf]\n:smile:") + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_extract_post_message_downloads_embedded_resources(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -733,91 +1394,281 @@ def test_extract_post_message_downloads_embedded_resources(self): fallback_filename="spec.pdf", ) - - @patch.dict(os.environ, {}, clear=True) - def test_extract_audio_message_downloads_and_caches(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_merge_forward_message_as_text_summary(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) - adapter._download_feishu_message_resource = AsyncMock( - return_value=("/tmp/feishu-audio.ogg", "audio/ogg") - ) message = SimpleNamespace( - message_type="audio", - content='{"file_key":"file_audio","file_name":"voice.ogg"}', - message_id="om_audio", + message_type="merge_forward", + content=json.dumps( + { + "title": "Forwarded updates", + "messages": [ + {"sender_name": "Alice", "text": "Investigating the incident"}, + {"sender_name": "Bob", "text": "ETA 10 minutes"}, + ], + } + ), + message_id="om_merge_forward", ) text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) - self.assertEqual(text, "") - # Lark "audio" msg_type is a native voice recording (the fixture is - # literally voice.ogg) — it must classify as VOICE so the gateway - # auto-transcribes it, not AUDIO (a non-transcribed file attachment). - # See the #28993 follow-up fix in _resolve_normalized_message_type. - self.assertEqual(msg_type.value, "voice") - self.assertEqual(media_urls, ["/tmp/feishu-audio.ogg"]) - self.assertEqual(media_types, ["audio/ogg"]) - + self.assertEqual( + text, + "Forwarded updates\n- Alice: Investigating the incident\n- Bob: ETA 10 minutes", + ) + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) - @patch.dict(os.environ, {}, clear=True) - def test_extract_text_message_starting_with_slash_becomes_command(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_share_chat_message_as_text_summary(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) - adapter._dispatch_inbound_event = AsyncMock() - adapter.get_chat_info = AsyncMock( - return_value={"chat_id": "oc_chat", "name": "Feishu DM", "type": "dm"} - ) - adapter._resolve_sender_profile = AsyncMock( - return_value={"user_id": "ou_user", "user_name": "张三", "user_id_alt": None} - ) message = SimpleNamespace( - chat_id="oc_chat", - thread_id=None, - parent_id=None, - upper_message_id=None, - message_type="text", - content='{"text":"/help test"}', - message_id="om_command", + message_type="share_chat", + content='{"chat_id":"oc_shared","chat_name":"Platform Ops"}', + message_id="om_share_chat", ) - asyncio.run( - adapter._process_inbound_message( - data=SimpleNamespace(event=SimpleNamespace(message=message)), - message=message, - sender_id=SimpleNamespace(open_id="ou_user", user_id=None, union_id=None), - is_bot=False, - chat_type="p2p", - message_id="om_command", - ) - ) + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) - event = adapter._dispatch_inbound_event.await_args.args[0] - self.assertEqual(event.message_type.value, "command") - self.assertEqual(event.text, "/help test") + self.assertEqual(text, "Shared chat: Platform Ops\nChat ID: oc_shared") + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) - @patch.dict(os.environ, {}, clear=True) - def test_extract_text_file_injects_content(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_interactive_message_as_text_summary(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) - with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as tmp: - tmp.write("hello from feishu") - path = tmp.name - - try: - text = asyncio.run(adapter._maybe_extract_text_document(path, "text/plain")) - finally: - os.unlink(path) - - self.assertIn("hello from feishu", text) + message = SimpleNamespace( + message_type="interactive", + content=json.dumps( + { + "card": { + "header": {"title": {"tag": "plain_text", "content": "Approval Request"}}, + "elements": [ + {"tag": "div", "text": {"tag": "plain_text", "content": "Requester: Alice"}}, + { + "tag": "action", + "actions": [ + {"tag": "button", "text": {"tag": "plain_text", "content": "Approve"}}, + ], + }, + ], + } + } + ), + message_id="om_interactive", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "Approval Request\nRequester: Alice\nApprove\nActions: Approve") + self.assertEqual(msg_type.value, "text") + self.assertEqual(media_urls, []) + self.assertEqual(media_types, []) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_image_message_downloads_and_caches(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._download_feishu_image = AsyncMock(return_value=("/tmp/feishu-image.png", "image/png")) + message = SimpleNamespace( + message_type="image", + content='{"image_key":"img_123"}', + message_id="om_image", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "") + self.assertEqual(msg_type.value, "photo") + self.assertEqual(media_urls, ["/tmp/feishu-image.png"]) + self.assertEqual(media_types, ["image/png"]) + adapter._download_feishu_image.assert_awaited_once_with( + message_id="om_image", + image_key="img_123", + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_audio_message_downloads_and_caches(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._download_feishu_message_resource = AsyncMock( + return_value=("/tmp/feishu-audio.ogg", "audio/ogg") + ) + message = SimpleNamespace( + message_type="audio", + content='{"file_key":"file_audio","file_name":"voice.ogg"}', + message_id="om_audio", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "") + # Lark "audio" msg_type is a native voice recording (the fixture is + # literally voice.ogg) — it must classify as VOICE so the gateway + # auto-transcribes it, not AUDIO (a non-transcribed file attachment). + # See the #28993 follow-up fix in _resolve_normalized_message_type. + self.assertEqual(msg_type.value, "voice") + self.assertEqual(media_urls, ["/tmp/feishu-audio.ogg"]) + self.assertEqual(media_types, ["audio/ogg"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_file_message_downloads_and_caches(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._download_feishu_message_resource = AsyncMock( + return_value=("/tmp/doc_123_report.pdf", "application/pdf") + ) + message = SimpleNamespace( + message_type="file", + content='{"file_key":"file_doc","file_name":"report.pdf"}', + message_id="om_file", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "") + self.assertEqual(msg_type.value, "document") + self.assertEqual(media_urls, ["/tmp/doc_123_report.pdf"]) + self.assertEqual(media_types, ["application/pdf"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_media_message_with_image_mime_becomes_photo(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._download_feishu_message_resource = AsyncMock( + return_value=("/tmp/feishu-media.jpg", "image/jpeg") + ) + message = SimpleNamespace( + message_type="media", + content='{"file_key":"file_media","file_name":"photo.jpg"}', + message_id="om_media", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "") + self.assertEqual(msg_type.value, "photo") + self.assertEqual(media_urls, ["/tmp/feishu-media.jpg"]) + self.assertEqual(media_types, ["image/jpeg"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_media_message_with_video_mime_becomes_video(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._download_feishu_message_resource = AsyncMock( + return_value=("/tmp/feishu-video.mp4", "video/mp4") + ) + message = SimpleNamespace( + message_type="media", + content='{"file_key":"file_video","file_name":"clip.mp4"}', + message_id="om_video", + ) + + text, msg_type, media_urls, media_types, _mentions = asyncio.run(adapter._extract_message_content(message)) + + self.assertEqual(text, "") + self.assertEqual(msg_type.value, "video") + self.assertEqual(media_urls, ["/tmp/feishu-video.mp4"]) + self.assertEqual(media_types, ["video/mp4"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_text_from_raw_content_uses_relation_message_fallbacks(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + shared = adapter._extract_text_from_raw_content( + msg_type="share_chat", + raw_content='{"chat_id":"oc_shared","chat_name":"Platform Ops"}', + ) + attachment = adapter._extract_text_from_raw_content( + msg_type="file", + raw_content='{"file_key":"file_1","file_name":"report.pdf"}', + ) + + self.assertEqual(shared, "Shared chat: Platform Ops\nChat ID: oc_shared") + self.assertEqual(attachment, "[Attachment: report.pdf]") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_text_message_starting_with_slash_becomes_command(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._dispatch_inbound_event = AsyncMock() + adapter.get_chat_info = AsyncMock( + return_value={"chat_id": "oc_chat", "name": "Feishu DM", "type": "dm"} + ) + adapter._resolve_sender_profile = AsyncMock( + return_value={"user_id": "ou_user", "user_name": "张三", "user_id_alt": None} + ) + message = SimpleNamespace( + chat_id="oc_chat", + thread_id=None, + parent_id=None, + upper_message_id=None, + message_type="text", + content='{"text":"/help test"}', + message_id="om_command", + ) + + asyncio.run( + adapter._process_inbound_message( + data=SimpleNamespace(event=SimpleNamespace(message=message)), + message=message, + sender_id=SimpleNamespace(open_id="ou_user", user_id=None, union_id=None), + is_bot=False, + chat_type="p2p", + message_id="om_command", + ) + ) + + event = adapter._dispatch_inbound_event.await_args.args[0] + self.assertEqual(event.message_type.value, "command") + self.assertEqual(event.text, "/help test") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_extract_text_file_injects_content(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as tmp: + tmp.write("hello from feishu") + path = tmp.name + + try: + text = asyncio.run(adapter._maybe_extract_text_document(path, "text/plain")) + finally: + os.unlink(path) + + self.assertIn("hello from feishu", text) self.assertIn("[Content of", text) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_message_event_submits_to_adapter_loop(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -852,7 +1703,7 @@ def _submit(coro, _loop): self.assertTrue(submit.called) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_webhook_request_uses_same_message_dispatch_path(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -876,7 +1727,7 @@ def test_webhook_request_uses_same_message_dispatch_path(self): self.assertEqual(response.status, 200) adapter._on_message_event.assert_called_once() - @patch.dict(os.environ, {"FEISHU_VERIFICATION_TOKEN": "expected-token"}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_VERIFICATION_TOKEN": "expected-token"}, clear=True) def test_url_verification_requires_configured_verification_token(self): """url_verification must be rejected when token is set but mismatched. @@ -904,7 +1755,7 @@ def test_url_verification_requires_configured_verification_token(self): self.assertEqual(response.status, 401) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_process_inbound_message_uses_event_sender_identity_only(self): from gateway.config import PlatformConfig from gateway.platforms.base import MessageType @@ -950,12 +1801,49 @@ def test_process_inbound_message_uses_event_sender_identity_only(self): self.assertEqual(event.source.user_id_alt, "on_union") self.assertEqual(event.source.chat_name, "Feishu DM") + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_text_batch_merges_rapid_messages_into_single_event(self): + from gateway.config import PlatformConfig + from gateway.platforms.base import MessageEvent, MessageType + from plugins.platforms.feishu.adapter import FeishuAdapter + from gateway.session import SessionSource + + adapter = FeishuAdapter(PlatformConfig()) + adapter.handle_message = AsyncMock() + source = SessionSource( + platform=adapter.platform, + chat_id="oc_chat", + chat_name="Feishu DM", + chat_type="dm", + user_id="ou_user", + user_name="张三", + ) + + async def _sleep(_delay): + return None + + async def _run() -> None: + with patch("plugins.platforms.feishu.adapter.asyncio.sleep", side_effect=_sleep): + await adapter._dispatch_inbound_event( + MessageEvent(text="A", message_type=MessageType.TEXT, source=source, message_id="om_1") + ) + await adapter._dispatch_inbound_event( + MessageEvent(text="B", message_type=MessageType.TEXT, source=source, message_id="om_2") + ) + pending = list(adapter._pending_text_batch_tasks.values()) + self.assertEqual(len(pending), 1) + await asyncio.gather(*pending, return_exceptions=True) + + asyncio.run(_run()) + + adapter.handle_message.assert_awaited_once() + event = adapter.handle_message.await_args.args[0] + self.assertEqual(event.text, "A\nB") + self.assertEqual(event.message_type, MessageType.TEXT) @patch.dict( os.environ, - { - "HERMES_FEISHU_TEXT_BATCH_MAX_MESSAGES": "2", - }, + {**_HOME_ENV, "HERMES_FEISHU_TEXT_BATCH_MAX_MESSAGES": "2"}, clear=True, ) def test_text_batch_flushes_when_message_count_limit_is_hit(self): @@ -1001,7 +1889,7 @@ async def _run() -> None: self.assertEqual(first.text, "A\nB") self.assertEqual(second.text, "C") - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_media_batch_merges_rapid_photo_messages(self): from gateway.config import PlatformConfig from gateway.platforms.base import MessageEvent, MessageType @@ -1056,6 +1944,47 @@ async def _run() -> None: self.assertIn("第一张", event.text) self.assertIn("第二张", event.text) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_image_downloads_then_uses_native_image_send(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter.send_image_file = AsyncMock(return_value=SimpleNamespace(success=True, message_id="om_img")) + + async def _run(): + with patch("plugins.platforms.feishu.adapter.cache_image_from_url", new=AsyncMock(return_value="/tmp/cached.png")): + return await adapter.send_image("oc_chat", "https://example.com/cat.png", caption="cat") + + result = asyncio.run(_run()) + + self.assertTrue(result.success) + adapter.send_image_file.assert_awaited_once() + self.assertEqual(adapter.send_image_file.await_args.kwargs["image_path"], "/tmp/cached.png") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_animation_degrades_to_document_send(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter.send_document = AsyncMock(return_value=SimpleNamespace(success=True, message_id="om_gif")) + + async def _run(): + with patch.object( + adapter, + "_download_remote_document", + new=AsyncMock(return_value=("/tmp/anim.gif", "anim.gif")), + ): + return await adapter.send_animation("oc_chat", "https://example.com/anim.gif", caption="look") + + result = asyncio.run(_run()) + + self.assertTrue(result.success) + adapter.send_document.assert_awaited_once() + caption = adapter.send_document.await_args.kwargs["caption"] + self.assertIn("GIF downgraded to file", caption) + self.assertIn("look", caption) def test_download_remote_document_reads_response_before_httpx_client_closes(self): """#18451 — snapshot Content-Type + body while the httpx.AsyncClient @@ -1095,7 +2024,7 @@ async def get(self, *_a: object, **_k: object) -> _FakeResponse: return _FakeResponse() with tempfile.TemporaryDirectory() as tmp: - with patch.dict(os.environ, {"HERMES_HOME": tmp}, clear=False): + with patch.dict(os.environ, {**_HOME_ENV, "HERMES_HOME": tmp}, clear=False): adapter = FeishuAdapter(PlatformConfig()) async def _run() -> tuple[str, str]: @@ -1139,85 +2068,1030 @@ def fake_getaddrinfo(_host, port, *_args, **_kwargs): (socket.AF_INET, socket.SOCK_STREAM, 6, "", (ip, port or 0)) ] - connect_attempts = [] + connect_attempts = [] + + async def fake_connect_tcp( + _self, + host, + port, + timeout=None, + local_address=None, + socket_options=None, + ): + connect_attempts.append((host, port)) + raise httpcore.ConnectError("stop before network") + + proxy_vars = { + name: "" + for name in ( + "HTTP_PROXY", + "HTTPS_PROXY", + "ALL_PROXY", + "http_proxy", + "https_proxy", + "all_proxy", + ) + } + with ( + patch.dict(os.environ, proxy_vars, clear=False), + patch("socket.getaddrinfo", side_effect=fake_getaddrinfo), + patch.object(AutoBackend, "connect_tcp", new=fake_connect_tcp), + self.assertRaises(SSRFConnectionBlocked), + ): + asyncio.run( + adapter._download_remote_document( + "http://rebind.example/doc.bin", + default_ext=".bin", + preferred_name="doc", + ) + ) + + self.assertEqual(connect_attempts, []) + + def test_dedup_state_persists_across_adapter_restart(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + with tempfile.TemporaryDirectory() as temp_home: + with patch.dict(os.environ, {**_HOME_ENV, "HERMES_HOME": temp_home}, clear=False): + first = FeishuAdapter(PlatformConfig()) + self.assertFalse(first._is_duplicate("om_same")) + second = FeishuAdapter(PlatformConfig()) + self.assertTrue(second._is_duplicate("om_same")) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_process_inbound_group_message_keeps_group_type_when_chat_lookup_falls_back(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._dispatch_inbound_event = AsyncMock() + adapter.get_chat_info = AsyncMock( + return_value={"chat_id": "oc_group", "name": "oc_group", "type": "dm"} + ) + adapter._resolve_sender_profile = AsyncMock( + return_value={"user_id": "ou_user", "user_name": "张三", "user_id_alt": None} + ) + message = SimpleNamespace( + chat_id="oc_group", + thread_id=None, + message_type="text", + content='{"text":"hello group"}', + message_id="om_group_text", + ) + sender_id = SimpleNamespace(open_id="ou_user", user_id=None, union_id=None) + sender = SimpleNamespace(sender_type="user", sender_id=sender_id) + data = SimpleNamespace(event=SimpleNamespace(message=message)) + + asyncio.run( + adapter._process_inbound_message( + data=data, + message=message, + sender_id=sender.sender_id, + chat_type="group", + message_id="om_group_text", + ) + ) + + event = adapter._dispatch_inbound_event.await_args.args[0] + self.assertEqual(event.source.chat_type, "group") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_process_inbound_message_fetches_reply_to_text(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._dispatch_inbound_event = AsyncMock() + adapter.get_chat_info = AsyncMock( + return_value={"chat_id": "oc_chat", "name": "Feishu DM", "type": "dm"} + ) + adapter._resolve_sender_profile = AsyncMock( + return_value={"user_id": "ou_user", "user_name": "张三", "user_id_alt": None} + ) + adapter._fetch_message_text = AsyncMock(return_value="父消息内容") + message = SimpleNamespace( + chat_id="oc_chat", + thread_id=None, + parent_id="om_parent", + upper_message_id=None, + message_type="text", + content='{"text":"reply"}', + message_id="om_reply", + ) + + asyncio.run( + adapter._process_inbound_message( + data=SimpleNamespace(event=SimpleNamespace(message=message)), + message=message, + sender_id=SimpleNamespace(open_id="ou_user", user_id=None, union_id=None), + is_bot=False, + chat_type="p2p", + message_id="om_reply", + ) + ) + + event = adapter._dispatch_inbound_event.await_args.args[0] + self.assertEqual(event.reply_to_message_id, "om_parent") + self.assertEqual(event.reply_to_text, "父消息内容") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_replies_in_thread_when_thread_metadata_present(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _ReplyAPI: + def reply(self, request): + captured["request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_reply"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_ReplyAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content="hello", + reply_to="om_parent", + metadata={"thread_id": "omt-thread"}, + ) + ) + + self.assertTrue(result.success) + self.assertEqual(result.message_id, "om_reply") + self.assertTrue(captured["request"].request_body.reply_in_thread) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_uses_metadata_reply_target_for_threaded_feishu_topic(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _MessageAPI: + def reply(self, request): + captured["request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_reply"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace(v1=SimpleNamespace(message=_MessageAPI())) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content="status update", + metadata={ + "thread_id": "omt-thread", + "reply_to_message_id": "om_trigger", + }, + ) + ) + + self.assertTrue(result.success) + self.assertEqual(captured["request"].message_id, "om_trigger") + self.assertTrue(captured["request"].request_body.reply_in_thread) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_retries_transient_failure(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {"attempts": 0} + sleeps = [] + + class _MessageAPI: + def create(self, request): + captured["attempts"] += 1 + captured["request"] = request + if captured["attempts"] == 1: + raise OSError("temporary send failure") + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_retry"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + async def _sleep(delay): + sleeps.append(delay) + + with ( + patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct), + patch("plugins.platforms.feishu.adapter.asyncio.sleep", side_effect=_sleep), + ): + result = asyncio.run(adapter.send(chat_id="oc_chat", content="hello retry")) + + self.assertTrue(result.success) + self.assertEqual(result.message_id, "om_retry") + self.assertEqual(captured["attempts"], 2) + self.assertEqual(sleeps, [1]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_does_not_retry_deterministic_api_failure(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {"attempts": 0} + sleeps = [] + + class _MessageAPI: + def create(self, request): + captured["attempts"] += 1 + return SimpleNamespace( + success=lambda: False, + code=400, + msg="bad request", + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + async def _sleep(delay): + sleeps.append(delay) + + with ( + patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct), + patch("plugins.platforms.feishu.adapter.asyncio.sleep", side_effect=_sleep), + ): + result = asyncio.run(adapter.send(chat_id="oc_chat", content="bad payload")) + + self.assertFalse(result.success) + self.assertEqual(result.error, "[400] bad request") + self.assertEqual(captured["attempts"], 1) + self.assertEqual(sleeps, []) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_document_reply_uses_thread_flag(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _FileAPI: + def create(self, request): + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(file_key="file_123"), + ) + + class _MessageAPI: + def reply(self, request): + captured["request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_file_reply"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + file=_FileAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".pdf", delete=False) as tmp: + tmp.write(b"%PDF-1.4 test") + file_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send_document( + chat_id="oc_chat", + file_path=file_path, + reply_to="om_parent", + metadata={"thread_id": "omt-thread"}, + ) + ) + finally: + os.unlink(file_path) + + self.assertTrue(result.success) + self.assertTrue(captured["request"].request_body.reply_in_thread) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_document_uploads_file_and_sends_file_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _FileAPI: + def create(self, request): + captured["upload_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(file_key="file_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_file_msg"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + file=_FileAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".pdf", delete=False) as tmp: + tmp.write(b"%PDF-1.4 test") + file_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter.send_document(chat_id="oc_chat", file_path=file_path)) + finally: + os.unlink(file_path) + + self.assertTrue(result.success) + self.assertEqual(result.message_id, "om_file_msg") + self.assertEqual(captured["upload_request"].request_body.file_type, "pdf") + self.assertEqual( + captured["message_request"].request_body.content, + '{"file_key": "file_123"}', + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_document_with_caption_uses_single_post_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _FileAPI: + def create(self, request): + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(file_key="file_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_post_msg"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + file=_FileAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".pdf", delete=False) as tmp: + tmp.write(b"%PDF-1.4 test") + file_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send_document(chat_id="oc_chat", file_path=file_path, caption="报告请看") + ) + finally: + os.unlink(file_path) + + self.assertTrue(result.success) + self.assertEqual(captured["message_request"].request_body.msg_type, "post") + self.assertIn('"tag": "media"', captured["message_request"].request_body.content) + self.assertIn('"file_key": "file_123"', captured["message_request"].request_body.content) + self.assertIn("报告请看", captured["message_request"].request_body.content) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_image_file_uploads_image_and_sends_image_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _ImageAPI: + def create(self, request): + captured["upload_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(image_key="img_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_image_msg"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + image=_ImageAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".png", delete=False) as tmp: + tmp.write(b"\x89PNG\r\n\x1a\n") + image_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter.send_image_file(chat_id="oc_chat", image_path=image_path)) + finally: + os.unlink(image_path) + + self.assertTrue(result.success) + self.assertEqual(result.message_id, "om_image_msg") + self.assertEqual(captured["upload_request"].request_body.image_type, "message") + self.assertEqual( + captured["message_request"].request_body.content, + '{"image_key": "img_123"}', + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_image_file_with_caption_uses_single_post_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _ImageAPI: + def create(self, request): + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(image_key="img_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_post_img"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + image=_ImageAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".png", delete=False) as tmp: + tmp.write(b"\x89PNG\r\n\x1a\n") + image_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send_image_file(chat_id="oc_chat", image_path=image_path, caption="截图说明") + ) + finally: + os.unlink(image_path) + + self.assertTrue(result.success) + self.assertEqual(captured["message_request"].request_body.msg_type, "post") + self.assertIn('"tag": "img"', captured["message_request"].request_body.content) + self.assertIn('"image_key": "img_123"', captured["message_request"].request_body.content) + self.assertIn("截图说明", captured["message_request"].request_body.content) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_video_uploads_file_and_sends_media_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _FileAPI: + def create(self, request): + captured["upload_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(file_key="file_video_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_video_msg"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + file=_FileAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".mp4", delete=False) as tmp: + tmp.write(b"\x00\x00\x00\x18ftypmp42") + video_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter.send_video(chat_id="oc_chat", video_path=video_path)) + finally: + os.unlink(video_path) + + self.assertTrue(result.success) + self.assertEqual(captured["upload_request"].request_body.file_type, "mp4") + self.assertEqual(captured["message_request"].request_body.msg_type, "media") + self.assertEqual(captured["message_request"].request_body.content, '{"file_key": "file_video_123"}') + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_voice_uploads_opus_and_sends_audio_message(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _FileAPI: + def create(self, request): + captured["upload_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(file_key="file_audio_123"), + ) + + class _MessageAPI: + def create(self, request): + captured["message_request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_audio_msg"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + file=_FileAPI(), + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with tempfile.NamedTemporaryFile("wb", suffix=".opus", delete=False) as tmp: + tmp.write(b"opus") + audio_path = tmp.name + + try: + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter.send_voice(chat_id="oc_chat", audio_path=audio_path)) + finally: + os.unlink(audio_path) + + self.assertTrue(result.success) + self.assertEqual(captured["upload_request"].request_body.file_type, "opus") + self.assertEqual(captured["message_request"].request_body.msg_type, "audio") + self.assertEqual(captured["message_request"].request_body.content, '{"file_key": "file_audio_123"}') + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_extracts_title_and_links(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads(adapter._build_post_payload("# 标题\n访问 [文档](https://example.com)")) + + elements = payload["zh_cn"]["content"][0] + self.assertEqual(elements, [{"tag": "md", "text": "# 标题\n访问 [文档](https://example.com)"}]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_wraps_markdown_in_md_tag(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads( + adapter._build_post_payload("支持 **粗体**、*斜体* 和 `代码`") + ) + + elements = payload["zh_cn"]["content"][0] + self.assertEqual( + elements, + [ + {"tag": "md", "text": "支持 **粗体**、*斜体* 和 `代码`"}, + ], + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_keeps_full_markdown_text(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads( + adapter._build_post_payload( + "---\n1. 第一项\n 2. 子项\n- 外层\n - 内层\n下划线 和 ~~删除线~~" + ) + ) + + rows = payload["zh_cn"]["content"] + self.assertEqual( + rows, + [[{"tag": "md", "text": "---\n1. 第一项\n 2. 子项\n- 外层\n - 内层\n下划线 和 ~~删除线~~"}]], + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_outbound_payload_uses_post_for_markdown_table(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + content = "| Name | Score |\n| --- | ---: |\n| Ada | 10 |" + + msg_type, raw_payload = adapter._build_outbound_payload(content) + + self.assertEqual(msg_type, "post") + payload = json.loads(raw_payload) + self.assertEqual( + payload["zh_cn"]["content"], + [[{"tag": "md", "text": content}]], + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_uses_post_for_every_chunk_of_multi_chunk_markdown(self): + """Regression for #26841: when a long Markdown message is split + across multiple chunks, every chunk must go out as + ``msg_type=post`` — including chunk 1. The bug was that the + first chunk often had only plain prose (the per-chunk regex + didn't match) and was sent as ``text``, so users saw literal + ``**bold``/``## heading``/code fences while later chunks + rendered correctly. + """ + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = [] + + class _MessageAPI: + def create(self, request): + captured.append(request) + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace( + message_id=f"om_chunk_{len(captured)}", + ), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + # Force a deterministic split so the test doesn't depend on the + # exact 8000-char limit. Chunk 1 is plain prose; chunk 2 has + # the markdown markers. Without the fix, chunk 1 went out as + # ``msg_type=text``. + first_chunk = "Here is a short intro that has no markdown markers at all." + second_chunk = "## Heading\nAnd then some **bold** text." + + with patch.object( + adapter, "truncate_message", return_value=[first_chunk, second_chunk], + ), patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content=first_chunk + "\n" + second_chunk, + ) + ) + + self.assertTrue(result.success) + self.assertEqual(len(captured), 2) + msg_types = [r.request_body.msg_type for r in captured] + self.assertEqual(msg_types, ["post", "post"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_plain_text_message_not_upgraded_by_prefer_post(self): + """A message with no markdown at all must still go out as plain + ``msg_type=text`` — the whole-message ``prefer_post`` decision + only flips on when the formatted message matches the hint regex. + """ + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = [] + + class _MessageAPI: + def create(self, request): + captured.append(request) + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace( + message_id=f"om_chunk_{len(captured)}", + ), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content="just a plain sentence", + ) + ) + + self.assertTrue(result.success) + self.assertEqual(len(captured), 1) + self.assertEqual(captured[0].request_body.msg_type, "text") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_uses_post_for_inline_markdown(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _MessageAPI: + def create(self, request): + captured["request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_markdown"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content="可以用 **粗体** 和 *斜体*。", + ) + ) + + self.assertTrue(result.success) + self.assertEqual(captured["request"].request_body.msg_type, "post") + payload = json.loads(captured["request"].request_body.content) + elements = payload["zh_cn"]["content"][0] + self.assertEqual(elements, [{"tag": "md", "text": "可以用 **粗体** 和 *斜体*。"}]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_splits_fenced_code_blocks_into_separate_post_rows(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + captured = {} + + class _MessageAPI: + def create(self, request): + captured["request"] = request + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(message_id="om_codeblock"), + ) + + adapter._client = SimpleNamespace( + im=SimpleNamespace( + v1=SimpleNamespace( + message=_MessageAPI(), + ) + ) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + content = ( + "确认已入库 ✓\n" + "文件路径:`/root/.hermes/profiles/agent_cto/cron/jobs.json`\n" + "**解码后的内容:**\n" + "```json\n" + '{"cron": "list"}\n' + "```\n" + "后续说明仍应保留。" + ) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content=content, + ) + ) + + self.assertTrue(result.success) + self.assertEqual(captured["request"].request_body.msg_type, "post") + payload = json.loads(captured["request"].request_body.content) + rows = payload["zh_cn"]["content"] + self.assertEqual( + rows, + [ + [ + { + "tag": "md", + "text": "确认已入库 ✓\n文件路径:`/root/.hermes/profiles/agent_cto/cron/jobs.json`\n**解码后的内容:**", + } + ], + [{"tag": "md", "text": "```json\n{\"cron\": \"list\"}\n```"}], + [{"tag": "md", "text": "后续说明仍应保留。"}], + ], + ) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_keeps_fence_like_code_lines_inside_code_block(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads( + adapter._build_post_payload( + "before\n```python\n```oops\n```\nafter" + ) + ) + + self.assertEqual( + payload["zh_cn"]["content"], + [ + [{"tag": "md", "text": "before"}], + [{"tag": "md", "text": "```python\n```oops\n```"}], + [{"tag": "md", "text": "after"}], + ], + ) - async def fake_connect_tcp( - _self, - host, - port, - timeout=None, - local_address=None, - socket_options=None, - ): - connect_attempts.append((host, port)) - raise httpcore.ConnectError("stop before network") + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_preserves_trailing_spaces_in_code_block(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter - proxy_vars = { - name: "" - for name in ( - "HTTP_PROXY", - "HTTPS_PROXY", - "ALL_PROXY", - "http_proxy", - "https_proxy", - "all_proxy", - ) - } - with ( - patch.dict(os.environ, proxy_vars, clear=False), - patch("socket.getaddrinfo", side_effect=fake_getaddrinfo), - patch.object(AutoBackend, "connect_tcp", new=fake_connect_tcp), - self.assertRaises(SSRFConnectionBlocked), - ): - asyncio.run( - adapter._download_remote_document( - "http://rebind.example/doc.bin", - default_ext=".bin", - preferred_name="doc", - ) + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads( + adapter._build_post_payload( + "before\n```python\nline with two spaces \n```\nafter" ) + ) - self.assertEqual(connect_attempts, []) + self.assertEqual( + payload["zh_cn"]["content"], + [ + [{"tag": "md", "text": "before"}], + [{"tag": "md", "text": "```python\nline with two spaces \n```"}], + [{"tag": "md", "text": "after"}], + ], + ) - def test_dedup_state_persists_across_adapter_restart(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_build_post_payload_splits_multiple_fenced_code_blocks(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter - with tempfile.TemporaryDirectory() as temp_home: - with patch.dict(os.environ, {"HERMES_HOME": temp_home}, clear=False): - first = FeishuAdapter(PlatformConfig()) - self.assertFalse(first._is_duplicate("om_same")) - second = FeishuAdapter(PlatformConfig()) - self.assertTrue(second._is_duplicate("om_same")) + adapter = FeishuAdapter(PlatformConfig()) + payload = json.loads( + adapter._build_post_payload( + "before\n```python\nprint(1)\n```\nmiddle\n```json\n{}\n```\nafter" + ) + ) + self.assertEqual( + payload["zh_cn"]["content"], + [ + [{"tag": "md", "text": "before"}], + [{"tag": "md", "text": "```python\nprint(1)\n```"}], + [{"tag": "md", "text": "middle"}], + [{"tag": "md", "text": "```json\n{}\n```"}], + [{"tag": "md", "text": "after"}], + ], + ) - @patch.dict(os.environ, {}, clear=True) - def test_send_document_reply_uses_thread_flag(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_falls_back_to_text_when_post_payload_is_rejected(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) - captured = {} - - class _FileAPI: - def create(self, request): - return SimpleNamespace( - success=lambda: True, - data=SimpleNamespace(file_key="file_123"), - ) + captured = {"calls": []} class _MessageAPI: - def reply(self, request): - captured["request"] = request + def create(self, request): + captured["calls"].append(request) + if len(captured["calls"]) == 1: + raise RuntimeError("content format of the post type is incorrect") return SimpleNamespace( success=lambda: True, - data=SimpleNamespace(message_id="om_file_reply"), + data=SimpleNamespace(message_id="om_plain"), ) adapter._client = SimpleNamespace( im=SimpleNamespace( v1=SimpleNamespace( - file=_FileAPI(), message=_MessageAPI(), ) ) @@ -1226,51 +3100,38 @@ def reply(self, request): async def _direct(func, *args, **kwargs): return func(*args, **kwargs) - with tempfile.NamedTemporaryFile("wb", suffix=".pdf", delete=False) as tmp: - tmp.write(b"%PDF-1.4 test") - file_path = tmp.name - - try: - with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): - result = asyncio.run( - adapter.send_document( - chat_id="oc_chat", - file_path=file_path, - reply_to="om_parent", - metadata={"thread_id": "omt-thread"}, - ) + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run( + adapter.send( + chat_id="oc_chat", + content="可以用 **粗体** 和 *斜体*。", ) - finally: - os.unlink(file_path) + ) self.assertTrue(result.success) - self.assertTrue(captured["request"].request_body.reply_in_thread) - + self.assertEqual(captured["calls"][0].request_body.msg_type, "post") + self.assertEqual(captured["calls"][1].request_body.msg_type, "text") + self.assertEqual( + captured["calls"][1].request_body.content, + json.dumps({"text": "可以用 粗体 和 斜体。"}, ensure_ascii=False), + ) - @patch.dict(os.environ, {}, clear=True) - def test_send_uses_post_for_every_chunk_of_multi_chunk_markdown(self): - """Regression for #26841: when a long Markdown message is split - across multiple chunks, every chunk must go out as - ``msg_type=post`` — including chunk 1. The bug was that the - first chunk often had only plain prose (the per-chunk regex - didn't match) and was sent as ``text``, so users saw literal - ``**bold``/``## heading``/code fences while later chunks - rendered correctly. - """ + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_falls_back_to_text_when_post_response_is_unsuccessful(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) - captured = [] + captured = {"calls": []} class _MessageAPI: def create(self, request): - captured.append(request) + captured["calls"].append(request) + if len(captured["calls"]) == 1: + return SimpleNamespace(success=lambda: False, code=230001, msg="content format of the post type is incorrect") return SimpleNamespace( success=lambda: True, - data=SimpleNamespace( - message_id=f"om_chunk_{len(captured)}", - ), + data=SimpleNamespace(message_id="om_plain_response"), ) adapter._client = SimpleNamespace( @@ -1284,31 +3145,24 @@ def create(self, request): async def _direct(func, *args, **kwargs): return func(*args, **kwargs) - # Force a deterministic split so the test doesn't depend on the - # exact 8000-char limit. Chunk 1 is plain prose; chunk 2 has - # the markdown markers. Without the fix, chunk 1 went out as - # ``msg_type=text``. - first_chunk = "Here is a short intro that has no markdown markers at all." - second_chunk = "## Heading\nAnd then some **bold** text." - - with patch.object( - adapter, "truncate_message", return_value=[first_chunk, second_chunk], - ), patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): result = asyncio.run( adapter.send( chat_id="oc_chat", - content=first_chunk + "\n" + second_chunk, + content="可以用 **粗体** 和 *斜体*。", ) ) self.assertTrue(result.success) - self.assertEqual(len(captured), 2) - msg_types = [r.request_body.msg_type for r in captured] - self.assertEqual(msg_types, ["post", "post"]) - + self.assertEqual(captured["calls"][0].request_body.msg_type, "post") + self.assertEqual(captured["calls"][1].request_body.msg_type, "text") + self.assertEqual( + captured["calls"][1].request_body.content, + json.dumps({"text": "可以用 粗体 和 斜体。"}, ensure_ascii=False), + ) - @patch.dict(os.environ, {}, clear=True) - def test_send_splits_fenced_code_blocks_into_separate_post_rows(self): + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_send_uses_post_for_advanced_markdown_lines(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1320,7 +3174,7 @@ def create(self, request): captured["request"] = request return SimpleNamespace( success=lambda: True, - data=SimpleNamespace(message_id="om_codeblock"), + data=SimpleNamespace(message_id="om_markdown_advanced"), ) adapter._client = SimpleNamespace( @@ -1334,21 +3188,11 @@ def create(self, request): async def _direct(func, *args, **kwargs): return func(*args, **kwargs) - content = ( - "确认已入库 ✓\n" - "文件路径:`/root/.hermes/profiles/agent_cto/cron/jobs.json`\n" - "**解码后的内容:**\n" - "```json\n" - '{"cron": "list"}\n' - "```\n" - "后续说明仍应保留。" - ) - with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): result = asyncio.run( adapter.send( chat_id="oc_chat", - content=content, + content="---\n1. 第一项\n下划线\n~~删除线~~", ) ) @@ -1358,16 +3202,7 @@ async def _direct(func, *args, **kwargs): rows = payload["zh_cn"]["content"] self.assertEqual( rows, - [ - [ - { - "tag": "md", - "text": "确认已入库 ✓\n文件路径:`/root/.hermes/profiles/agent_cto/cron/jobs.json`\n**解码后的内容:**", - } - ], - [{"tag": "md", "text": "```json\n{\"cron\": \"list\"}\n```"}], - [{"tag": "md", "text": "后续说明仍应保留。"}], - ], + [[{"tag": "md", "text": "---\n1. 第一项\n下划线\n~~删除线~~"}]], ) @@ -1387,7 +3222,7 @@ def _make_adapter(self): return FeishuAdapter(PlatformConfig()) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_hydration_populates_open_id_from_bot_info(self): adapter = self._make_adapter() adapter._client = Mock() @@ -1410,10 +3245,7 @@ def test_hydration_populates_open_id_from_bot_info(self): @patch.dict( os.environ, - { - "FEISHU_BOT_OPEN_ID": "ou_env", - "FEISHU_BOT_NAME": "Env Hermes", - }, + {**_HOME_ENV, "FEISHU_BOT_OPEN_ID": "ou_env", "FEISHU_BOT_NAME": "Env Hermes"}, clear=True, ) def test_hydration_refreshes_env_values_when_bot_info_available(self): @@ -1439,6 +3271,61 @@ def test_hydration_refreshes_env_values_when_bot_info_available(self): self.assertEqual(adapter._bot_open_id, "ou_hydrated") self.assertEqual(adapter._bot_name, "Hydrated Hermes") + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_BOT_OPEN_ID": "ou_env"}, clear=True) + def test_hydration_overwrites_stale_env_open_id(self): + """A stale env open_id should not break group mention gating after app migration.""" + adapter = self._make_adapter() + adapter._client = Mock() + payload = json.dumps( + { + "code": 0, + "bot": { + "bot_name": "Hermes Bot", + "open_id": "ou_probe_DIFFERENT", + }, + } + ).encode("utf-8") + adapter._client.request = Mock(return_value=SimpleNamespace(raw=SimpleNamespace(content=payload))) + + asyncio.run(adapter._hydrate_bot_identity()) + + self.assertEqual(adapter._bot_open_id, "ou_probe_DIFFERENT") + self.assertEqual(adapter._bot_name, "Hermes Bot") # filled in + + @patch.dict( + os.environ, + {**_HOME_ENV, "FEISHU_BOT_OPEN_ID": "ou_env", "FEISHU_BOT_NAME": "Env Hermes"}, + clear=True, + ) + def test_hydration_preserves_env_values_when_bot_info_probe_fails(self): + adapter = self._make_adapter() + adapter._client = Mock() + adapter._client.request = Mock(side_effect=RuntimeError("network down")) + + asyncio.run(adapter._hydrate_bot_identity()) + + self.assertEqual(adapter._bot_open_id, "ou_env") + self.assertEqual(adapter._bot_name, "Env Hermes") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_hydration_tolerates_probe_failure_and_falls_back_to_app_info(self): + adapter = self._make_adapter() + adapter._client = Mock() + adapter._client.request = Mock(side_effect=RuntimeError("network down")) + + # Make the application-info fallback succeed for _bot_name. + app_response = Mock() + app_response.success = Mock(return_value=True) + app_response.data = SimpleNamespace(app=SimpleNamespace(app_name="Fallback Bot")) + adapter._client.application.v6.application.get = Mock(return_value=app_response) + adapter._build_get_application_request = Mock(return_value=object()) + + asyncio.run(adapter._hydrate_bot_identity()) + + # Primary probe failed — open_id stays empty, but bot_name came from app-info. + self.assertEqual(adapter._bot_open_id, "") + self.assertEqual(adapter._bot_name, "Fallback Bot") + @unittest.skipUnless(_HAS_LARK_OAPI, "lark-oapi not installed") class TestPendingInboundQueue(unittest.TestCase): @@ -1446,7 +3333,7 @@ class TestPendingInboundQueue(unittest.TestCase): before or during adapter loop transitions must be queued for replay rather than silently dropped.""" - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_event_queued_when_loop_not_ready(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1466,7 +3353,7 @@ def test_event_queued_when_loop_not_ready(self): # Drain scheduled flag set. self.assertTrue(adapter._pending_drain_scheduled) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_drainer_replays_queued_events_when_loop_becomes_ready(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1487,30 +3374,103 @@ def is_closed(self): self.assertEqual(len(adapter._pending_inbound_events), 3) - # Now the loop becomes ready; run the drainer inline (not as a thread) - # to verify it replays the queue. + # Now the loop becomes ready; run the drainer inline (not as a thread) + # to verify it replays the queue. + adapter._loop = _ReadyLoop() + + future = SimpleNamespace(add_done_callback=lambda *_a, **_kw: None) + submitted: list = [] + + def _submit(coro, _loop): + submitted.append(coro) + coro.close() + return future + + with patch( + "plugins.platforms.feishu.adapter.asyncio.run_coroutine_threadsafe", + side_effect=_submit, + ) as submit: + adapter._drain_pending_inbound_events() + + # All three events dispatched to the loop. + self.assertEqual(submit.call_count, 3) + # Queue emptied. + self.assertEqual(len(adapter._pending_inbound_events), 0) + # Drain flag reset so a future race can schedule a new drainer. + self.assertFalse(adapter._pending_drain_scheduled) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_drainer_drops_queue_when_adapter_shuts_down(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._loop = None + adapter._running = False # Shutdown state + + with patch("plugins.platforms.feishu.adapter.threading.Thread"): + adapter._on_message_event(SimpleNamespace(tag="evt-lost")) + + self.assertEqual(len(adapter._pending_inbound_events), 1) + + # Drainer should drop the queue immediately since _running is False. + adapter._drain_pending_inbound_events() + + self.assertEqual(len(adapter._pending_inbound_events), 0) + self.assertFalse(adapter._pending_drain_scheduled) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_queue_cap_evicts_oldest_beyond_max_depth(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._loop = None + adapter._pending_inbound_max_depth = 3 # Shrink for test + + with patch("plugins.platforms.feishu.adapter.threading.Thread"): + for i in range(5): + adapter._on_message_event(SimpleNamespace(tag=f"evt-{i}")) + + # Only the last 3 should remain; evt-0 and evt-1 dropped. + self.assertEqual(len(adapter._pending_inbound_events), 3) + tags = [getattr(e, "tag", None) for e in adapter._pending_inbound_events] + self.assertEqual(tags, ["evt-2", "evt-3", "evt-4"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_normal_path_unchanged_when_loop_ready(self): + """When the loop is ready, events should dispatch directly without + ever touching the pending queue.""" + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + class _ReadyLoop: + def is_closed(self): + return False + adapter._loop = _ReadyLoop() future = SimpleNamespace(add_done_callback=lambda *_a, **_kw: None) - submitted: list = [] def _submit(coro, _loop): - submitted.append(coro) coro.close() return future with patch( "plugins.platforms.feishu.adapter.asyncio.run_coroutine_threadsafe", side_effect=_submit, - ) as submit: - adapter._drain_pending_inbound_events() + ) as submit, patch( + "plugins.platforms.feishu.adapter.threading.Thread" + ) as thread_cls: + adapter._on_message_event(SimpleNamespace(tag="evt")) - # All three events dispatched to the loop. - self.assertEqual(submit.call_count, 3) - # Queue emptied. + self.assertEqual(submit.call_count, 1) self.assertEqual(len(adapter._pending_inbound_events), 0) - # Drain flag reset so a future race can schedule a new drainer. self.assertFalse(adapter._pending_drain_scheduled) + # No drainer thread spawned when the happy path runs. + self.assertEqual(thread_cls.call_count, 0) @unittest.skipUnless(_HAS_LARK_OAPI, "lark-oapi not installed") @@ -1521,7 +3481,7 @@ def _make_adapter(self, encrypt_key: str = "") -> "FeishuAdapter": from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter - with patch.dict(os.environ, {"FEISHU_APP_ID": "cli", "FEISHU_APP_SECRET": "sec", "FEISHU_ENCRYPT_KEY": encrypt_key}, clear=True): + with patch.dict(os.environ, {**_HOME_ENV, "FEISHU_APP_ID": "cli", "FEISHU_APP_SECRET": "sec", "FEISHU_ENCRYPT_KEY": encrypt_key}, clear=True): return FeishuAdapter(PlatformConfig()) def test_signature_valid_passes(self): @@ -1537,6 +3497,30 @@ def test_signature_valid_passes(self): headers = {"x-lark-request-timestamp": timestamp, "x-lark-request-nonce": nonce, "x-lark-signature": sig} self.assertTrue(adapter._is_webhook_signature_valid(headers, body)) + def test_signature_invalid_rejected(self): + adapter = self._make_adapter("test_secret") + headers = { + "x-lark-request-timestamp": "1700000000", + "x-lark-request-nonce": "abc", + "x-lark-signature": "deadbeef" * 8, + } + self.assertFalse(adapter._is_webhook_signature_valid(headers, b'{"type":"event"}')) + + def test_signature_missing_headers_rejected(self): + adapter = self._make_adapter("test_secret") + self.assertFalse(adapter._is_webhook_signature_valid({}, b'{}')) + + def test_rate_limit_allows_requests_within_window(self): + adapter = self._make_adapter() + for _ in range(5): + self.assertTrue(adapter._check_webhook_rate_limit("10.0.0.1")) + + def test_rate_limit_blocks_after_exceeding_max(self): + from plugins.platforms.feishu.adapter import _FEISHU_WEBHOOK_RATE_LIMIT_MAX + adapter = self._make_adapter() + for _ in range(_FEISHU_WEBHOOK_RATE_LIMIT_MAX): + adapter._check_webhook_rate_limit("10.0.0.2") + self.assertFalse(adapter._check_webhook_rate_limit("10.0.0.2")) def test_rate_limit_resets_after_window_expires(self): from plugins.platforms.feishu.adapter import _FEISHU_WEBHOOK_RATE_LIMIT_MAX, _FEISHU_WEBHOOK_RATE_WINDOW_SECONDS @@ -1550,6 +3534,19 @@ def test_rate_limit_resets_after_window_expires(self): adapter._webhook_rate_counts[ip] = (count, window_start - _FEISHU_WEBHOOK_RATE_WINDOW_SECONDS - 1) self.assertTrue(adapter._check_webhook_rate_limit(ip)) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_webhook_request_rejects_oversized_body(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter, _FEISHU_WEBHOOK_MAX_BODY_BYTES + + adapter = FeishuAdapter(PlatformConfig()) + # Simulate a request whose Content-Length already signals oversize. + request = SimpleNamespace( + remote="127.0.0.1", + content_length=_FEISHU_WEBHOOK_MAX_BODY_BYTES + 1, + ) + response = asyncio.run(adapter._handle_webhook_request(request)) + self.assertEqual(response.status, 413) def test_webhook_request_rejects_oversized_chunked_body_while_reading(self): from gateway.config import PlatformConfig @@ -1575,8 +3572,37 @@ def test_webhook_request_rejects_oversized_chunked_body_while_reading(self): self.assertEqual(response.status, 413) self.assertEqual(content.read_sizes, [_FEISHU_WEBHOOK_MAX_BODY_BYTES + 1]) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_webhook_request_rejects_invalid_json(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + request = SimpleNamespace( + remote="127.0.0.1", + content_length=None, + content=_FakeRequestContent(b"not-json"), + ) + response = asyncio.run(adapter._handle_webhook_request(request)) + self.assertEqual(response.status, 400) + + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_ENCRYPT_KEY": "secret"}, clear=True) + def test_webhook_request_rejects_bad_signature(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + body = json.dumps({"header": {"event_type": "im.message.receive_v1"}}).encode() + request = SimpleNamespace( + remote="127.0.0.1", + content_length=None, + headers={"x-lark-request-timestamp": "123", "x-lark-request-nonce": "abc", "x-lark-signature": "bad"}, + content=_FakeRequestContent(body), + ) + response = asyncio.run(adapter._handle_webhook_request(request)) + self.assertEqual(response.status, 401) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_webhook_connect_requires_inbound_auth_secret(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1589,7 +3615,7 @@ def test_webhook_connect_requires_inbound_auth_secret(self): ) self.assertFalse(asyncio.run(adapter.connect())) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_webhook_loads_auth_secrets_from_platform_extra(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1609,11 +3635,28 @@ def test_webhook_loads_auth_secrets_from_platform_extra(self): self.assertEqual(adapter._verification_token, "token_from_extra") self.assertEqual(adapter._encrypt_key, "encrypt_from_extra") + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_webhook_url_verification_challenge_passes_without_signature(self): + """Challenge requests must succeed even when no encrypt_key is set.""" + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + body = json.dumps({"type": "url_verification", "challenge": "test_challenge_token"}).encode() + request = SimpleNamespace( + remote="127.0.0.1", + content_length=None, + content=_FakeRequestContent(body), + ) + response = asyncio.run(adapter._handle_webhook_request(request)) + self.assertEqual(response.status, 200) + self.assertIn(b"test_challenge_token", response.body) + class TestDedupTTL(unittest.TestCase): """Tests for TTL-aware deduplication.""" - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_duplicate_within_ttl_is_rejected(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1624,8 +3667,20 @@ def test_duplicate_within_ttl_is_rejected(self): adapter._seen_message_order = ["om_dup"] self.assertTrue(adapter._is_duplicate("om_dup")) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_expired_entry_is_not_considered_duplicate(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter, _FEISHU_DEDUP_TTL_SECONDS + + adapter = FeishuAdapter(PlatformConfig()) + # Plant an entry that expired well past the TTL. + stale_ts = time.time() - _FEISHU_DEDUP_TTL_SECONDS - 60 + adapter._seen_message_ids = {"om_old": stale_ts} + adapter._seen_message_order = ["om_old"] + with patch.object(adapter, "_persist_seen_message_ids"): + self.assertFalse(adapter._is_duplicate("om_old")) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_load_tolerates_malformed_timestamp_values(self): """Regression #13632 — a non-numeric timestamp in the persisted dedup state must not crash adapter startup. The bad key is @@ -1636,7 +3691,7 @@ def test_load_tolerates_malformed_timestamp_values(self): from plugins.platforms.feishu.adapter import FeishuAdapter with tempfile.TemporaryDirectory() as temp_home: - with patch.dict(os.environ, {"HERMES_HOME": temp_home}, clear=True): + with patch.dict(os.environ, {**_HOME_ENV, "HERMES_HOME": temp_home}, clear=True): adapter = FeishuAdapter(PlatformConfig()) adapter._dedup_state_path.parent.mkdir(parents=True, exist_ok=True) adapter._dedup_state_path.write_text( @@ -1656,12 +3711,54 @@ def test_load_tolerates_malformed_timestamp_values(self): assert "om_bad_str" not in adapter._seen_message_ids assert "om_bad_null" not in adapter._seen_message_ids + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_persist_saves_timestamps_as_dict(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + ts = time.time() + adapter._seen_message_ids = {"om_ts1": ts} + adapter._seen_message_order = ["om_ts1"] + with tempfile.TemporaryDirectory() as tmpdir: + adapter._dedup_state_path = Path(tmpdir) / "dedup.json" + adapter._persist_seen_message_ids() + saved = json.loads(adapter._dedup_state_path.read_text()) + self.assertIsInstance(saved["message_ids"], dict) + self.assertAlmostEqual(saved["message_ids"]["om_ts1"], ts, places=1) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_load_backward_compat_list_format(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + with tempfile.TemporaryDirectory() as tmpdir: + path = Path(tmpdir) / "dedup.json" + path.write_text(json.dumps({"message_ids": ["om_a", "om_b"]}), encoding="utf-8") + adapter._dedup_state_path = path + adapter._load_seen_message_ids() + self.assertIn("om_a", adapter._seen_message_ids) + self.assertIn("om_b", adapter._seen_message_ids) + class TestGroupMentionAtAll(unittest.TestCase): """Tests for @_all (Feishu @everyone) group mention routing.""" + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "open"}, clear=True) + def test_at_all_in_content_accepts_without_explicit_bot_mention(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + message = SimpleNamespace( + content='{"text":"@_all 请注意"}', + mentions=[], + ) + sender_id = SimpleNamespace(open_id="ou_any", user_id=None) + self.assertTrue(_admits_group(adapter, message, sender_id, "")) - @patch.dict(os.environ, {"FEISHU_GROUP_POLICY": "allowlist", "FEISHU_ALLOWED_USERS": "ou_allowed"}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_GROUP_POLICY": "allowlist", "FEISHU_ALLOWED_USERS": "ou_allowed"}, clear=True) def test_at_all_still_requires_policy_gate(self): """@_all bypasses mention gating but NOT the allowlist policy.""" from gateway.config import PlatformConfig @@ -1681,8 +3778,17 @@ def test_at_all_still_requires_policy_gate(self): class TestSenderNameResolution(unittest.TestCase): """Tests for _resolve_sender_name_from_api (contact API + cache).""" + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_returns_none_when_client_is_none(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + adapter._client = None + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_abc")) + self.assertIsNone(result) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_returns_cached_name_within_ttl(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1694,7 +3800,7 @@ def test_returns_cached_name_within_ttl(self): result = asyncio.run(adapter._resolve_sender_name_from_api("ou_cached")) self.assertEqual(result, "Alice") - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_fetches_and_caches_name_from_api(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1723,6 +3829,56 @@ def get(self, request): self.assertEqual(result, "Bob") self.assertIn("ou_bob", adapter._sender_name_cache) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_expired_cache_triggers_new_api_call(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + # Expired cache entry. + adapter._sender_name_cache["ou_expired"] = ("OldName", time.time() - 1) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + user_obj = SimpleNamespace(name="NewName", display_name=None, nickname=None, en_name=None) + + class _ContactAPI: + def get(self, request): + return SimpleNamespace(success=lambda: True, data=SimpleNamespace(user=user_obj)) + + adapter._client = SimpleNamespace( + contact=SimpleNamespace(v3=SimpleNamespace(user=_ContactAPI())) + ) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_expired")) + + self.assertEqual(result, "NewName") + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_api_failure_returns_none_without_raising(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + class _BrokenContactAPI: + def get(self, _request): + raise RuntimeError("API down") + + adapter._client = SimpleNamespace( + contact=SimpleNamespace(v3=SimpleNamespace(user=_BrokenContactAPI())) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_broken")) + + self.assertIsNone(result) + @unittest.skipUnless(_HAS_LARK_OAPI, "lark-oapi not installed") class TestBotNameResolution(unittest.TestCase): @@ -1751,7 +3907,7 @@ def _fake_request(request): adapter._client = SimpleNamespace(request=_fake_request) return adapter, calls - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_returns_cached_bot_name_without_api_call(self): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter @@ -1764,7 +3920,7 @@ def test_returns_cached_bot_name_without_api_call(self): result = asyncio.run(adapter._resolve_sender_name_from_api("ou_peer", is_bot=True)) self.assertEqual(result, "Peer Bot") - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_fetches_and_caches_bot_name(self): adapter, calls = self._build_adapter_with_bots({"ou_peer": "Peer Bot"}) @@ -1781,6 +3937,79 @@ async def _direct(func, *args, **kwargs): # Feishu expects repeated ?bot_ids= params, not comma-joined. self.assertEqual(calls[0].queries, [("bot_ids", "ou_peer")]) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_api_failure_returns_none_and_does_not_poison_cache(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + + def _broken_request(_req): + raise RuntimeError("API down") + + adapter._client = SimpleNamespace(request=_broken_request) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_peer", is_bot=True)) + + self.assertIsNone(result) + self.assertNotIn("ou_peer", adapter._sender_name_cache) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_bot_absent_from_response_is_not_cached(self): + """Bot not in ``data.bots`` (e.g. landed in ``failed_bots``) → no + cache entry, next lookup re-fetches.""" + adapter, _ = self._build_adapter_with_bots({"ou_other": "Other Bot"}) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_ghost", is_bot=True)) + + self.assertIsNone(result) + self.assertNotIn("ou_ghost", adapter._sender_name_cache) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_empty_name_in_response_is_negative_cached(self): + """API returns name="" → cache "" so repeat lookups short-circuit.""" + adapter, calls = self._build_adapter_with_bots({"ou_nameless": ""}) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + first = asyncio.run(adapter._resolve_sender_name_from_api("ou_nameless", is_bot=True)) + second = asyncio.run(adapter._resolve_sender_name_from_api("ou_nameless", is_bot=True)) + + self.assertIsNone(first) + self.assertIsNone(second) + self.assertEqual(adapter._sender_name_cache["ou_nameless"][0], "") + self.assertEqual(len(calls), 1) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_non_zero_code_returns_none(self): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + error_payload = b'{"code":99991663,"msg":"permission denied"}' + adapter._client = SimpleNamespace( + request=lambda _r: SimpleNamespace(raw=SimpleNamespace(content=error_payload)) + ) + + async def _direct(func, *args, **kwargs): + return func(*args, **kwargs) + + with patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct): + result = asyncio.run(adapter._resolve_sender_name_from_api("ou_peer", is_bot=True)) + + self.assertIsNone(result) + self.assertNotIn("ou_peer", adapter._sender_name_cache) + @unittest.skipUnless(_HAS_LARK_OAPI, "lark-oapi not installed") class TestProcessingReactions(unittest.TestCase): @@ -1850,7 +4079,7 @@ async def _direct(func, *args, **kwargs): return patch("plugins.platforms.feishu.adapter.asyncio.to_thread", side_effect=_direct) # ------------------------------------------------------------------ start - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_start_adds_typing_and_caches_reaction_id(self): adapter, tracker = self._build_adapter(next_reaction_id="r_typing") with self._patch_to_thread(): @@ -1858,9 +4087,24 @@ def test_start_adds_typing_and_caches_reaction_id(self): self.assertEqual(tracker.create_calls, ["Typing"]) self.assertEqual(adapter._pending_processing_reactions["om_msg"], "r_typing") + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_start_is_idempotent_for_same_message_id(self): + adapter, tracker = self._build_adapter(next_reaction_id="r_typing") + with self._patch_to_thread(): + self._run(adapter.on_processing_start(self._event())) + self._run(adapter.on_processing_start(self._event())) + self.assertEqual(tracker.create_calls, ["Typing"]) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_start_does_not_cache_when_create_fails(self): + adapter, tracker = self._build_adapter(create_success=False) + with self._patch_to_thread(): + self._run(adapter.on_processing_start(self._event())) + self.assertEqual(tracker.create_calls, ["Typing"]) + self.assertNotIn("om_msg", adapter._pending_processing_reactions) # --------------------------------------------------------------- complete - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_success_removes_typing_and_adds_nothing(self): adapter, tracker = self._build_adapter(next_reaction_id="r_typing") with self._patch_to_thread(): @@ -1872,7 +4116,7 @@ def test_success_removes_typing_and_adds_nothing(self): self.assertEqual(tracker.delete_calls, ["r_typing"]) self.assertNotIn("om_msg", adapter._pending_processing_reactions) - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_failure_removes_typing_then_adds_cross_mark(self): adapter, tracker = self._build_adapter(next_reaction_id="r_typing") with self._patch_to_thread(): @@ -1883,9 +4127,40 @@ def test_failure_removes_typing_then_adds_cross_mark(self): self.assertEqual(tracker.create_calls, ["Typing", "CrossMark"]) self.assertEqual(tracker.delete_calls, ["r_typing"]) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_cancelled_removes_typing_and_adds_nothing(self): + adapter, tracker = self._build_adapter(next_reaction_id="r_typing") + with self._patch_to_thread(): + self._run(adapter.on_processing_start(self._event())) + self._run( + adapter.on_processing_complete(self._event(), ProcessingOutcome.CANCELLED) + ) + self.assertEqual(tracker.create_calls, ["Typing"]) + self.assertEqual(tracker.delete_calls, ["r_typing"]) + self.assertNotIn("om_msg", adapter._pending_processing_reactions) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_failure_without_preceding_start_still_adds_cross_mark(self): + adapter, tracker = self._build_adapter() + with self._patch_to_thread(): + self._run( + adapter.on_processing_complete(self._event(), ProcessingOutcome.FAILURE) + ) + self.assertEqual(tracker.create_calls, ["CrossMark"]) + self.assertEqual(tracker.delete_calls, []) + + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_success_without_preceding_start_is_full_noop(self): + adapter, tracker = self._build_adapter() + with self._patch_to_thread(): + self._run( + adapter.on_processing_complete(self._event(), ProcessingOutcome.SUCCESS) + ) + self.assertEqual(tracker.create_calls, []) + self.assertEqual(tracker.delete_calls, []) # ------------------------- delete failure: don't stack badges ----------- - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) def test_delete_failure_on_failure_outcome_skips_cross_mark(self): # Removing Typing is best-effort — but if it fails, we must NOT # additionally add CrossMark, or the UI would show two contradictory @@ -1904,13 +4179,76 @@ def test_delete_failure_on_failure_outcome_skips_cross_mark(self): adapter._pending_processing_reactions["om_msg"], "r_typing", ) # handle retained + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_delete_failure_on_success_outcome_retains_handle(self): + adapter, tracker = self._build_adapter( + next_reaction_id="r_typing", delete_success=False, + ) + with self._patch_to_thread(): + self._run(adapter.on_processing_start(self._event())) + self._run( + adapter.on_processing_complete(self._event(), ProcessingOutcome.SUCCESS) + ) + self.assertEqual(tracker.create_calls, ["Typing"]) + self.assertEqual(tracker.delete_calls, ["r_typing"]) + self.assertEqual( + adapter._pending_processing_reactions["om_msg"], "r_typing", + ) # ------------------------------------------------------------- env toggle + @patch.dict(os.environ, {**_HOME_ENV, "FEISHU_REACTIONS": "false"}, clear=True) + def test_env_disable_short_circuits_both_hooks(self): + adapter, tracker = self._build_adapter() + with self._patch_to_thread(): + self._run(adapter.on_processing_start(self._event())) + self._run( + adapter.on_processing_complete(self._event(), ProcessingOutcome.FAILURE) + ) + self.assertEqual(tracker.create_calls, []) + self.assertEqual(tracker.delete_calls, []) # ------------------------------------------------------------- LRU bounds + @patch.dict(os.environ, {**_HOME_ENV, }, clear=True) + def test_cache_evicts_oldest_entry_beyond_size_limit(self): + from plugins.platforms.feishu.adapter import _FEISHU_PROCESSING_REACTION_CACHE_SIZE + + adapter, _ = self._build_adapter() + counter = {"n": 0} + + def _create(_request): + counter["n"] += 1 + return SimpleNamespace( + success=lambda: True, + data=SimpleNamespace(reaction_id=f"r{counter['n']}"), + ) + + adapter._client.im.v1.message_reaction.create = _create + + with self._patch_to_thread(): + for i in range(_FEISHU_PROCESSING_REACTION_CACHE_SIZE + 1): + self._run(adapter.on_processing_start(self._event(f"om_{i}"))) + + self.assertNotIn("om_0", adapter._pending_processing_reactions) + self.assertIn( + f"om_{_FEISHU_PROCESSING_REACTION_CACHE_SIZE}", + adapter._pending_processing_reactions, + ) + self.assertEqual( + len(adapter._pending_processing_reactions), + _FEISHU_PROCESSING_REACTION_CACHE_SIZE, + ) class TestFeishuMentionMap(unittest.TestCase): + def test_build_mentions_map_handles_at_all(self): + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity, FeishuMentionRef + + mention = SimpleNamespace(key="@_all", id=None, name="") + result = _build_mentions_map( + [mention], + _FeishuBotIdentity(open_id="ou_bot", name="Hermes"), + ) + self.assertEqual(result["@_all"], FeishuMentionRef(is_all=True)) def test_build_mentions_map_marks_self_by_open_id(self): from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity @@ -1925,6 +4263,16 @@ def test_build_mentions_map_marks_self_by_open_id(self): self.assertEqual(ref.open_id, "ou_bot") self.assertEqual(ref.name, "Hermes") + def test_build_mentions_map_marks_self_by_name_fallback(self): + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity + + mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="", user_id=""), + name="Hermes", + ) + result = _build_mentions_map([mention], _FeishuBotIdentity(name="Hermes")) + self.assertTrue(result["@_user_1"].is_self) def test_build_mentions_map_name_match_does_not_override_mismatching_open_id(self): """Regression: a human user whose display name matches the bot must @@ -1963,6 +4311,32 @@ def test_build_mentions_map_falls_back_to_name_when_bot_open_id_not_hydrated(sel ) self.assertTrue(result["@_user_1"].is_self) + def test_build_mentions_map_non_self_user(self): + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity + + mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_alice", user_id=""), + name="Alice", + ) + ref = _build_mentions_map([mention], _FeishuBotIdentity(open_id="ou_bot"))["@_user_1"] + self.assertFalse(ref.is_self) + self.assertEqual(ref.open_id, "ou_alice") + self.assertEqual(ref.name, "Alice") + + def test_build_mentions_map_returns_empty_for_none_input(self): + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity + + self.assertEqual(_build_mentions_map(None, _FeishuBotIdentity(open_id="ou_bot")), {}) + + def test_build_mentions_map_tolerates_missing_id_object(self): + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity + + mention = SimpleNamespace(key="@_user_9", id=None, name="") + ref = _build_mentions_map([mention], _FeishuBotIdentity(open_id="ou_bot"))["@_user_9"] + self.assertEqual(ref.open_id, "") + self.assertFalse(ref.is_self) + class TestFeishuMentionHint(unittest.TestCase): def test_hint_single_user(self): @@ -1986,6 +4360,11 @@ def test_hint_multiple_users(self): "[Mentioned: Alice (open_id=ou_alice), Bob (open_id=ou_bob)]", ) + def test_hint_at_all(self): + from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint + + refs = [FeishuMentionRef(is_all=True)] + self.assertEqual(_build_mention_hint(refs), "[Mentioned: @all]") def test_hint_filters_self_mentions(self): from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint @@ -1999,6 +4378,28 @@ def test_hint_filters_self_mentions(self): "[Mentioned: Alice (open_id=ou_alice)]", ) + def test_hint_returns_empty_when_only_self(self): + from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint + + refs = [FeishuMentionRef(name="Hermes", open_id="ou_bot", is_self=True)] + self.assertEqual(_build_mention_hint(refs), "") + + def test_hint_returns_empty_for_no_refs(self): + from plugins.platforms.feishu.adapter import _build_mention_hint + + self.assertEqual(_build_mention_hint([]), "") + + def test_hint_falls_back_when_open_id_missing(self): + from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint + + refs = [FeishuMentionRef(name="Alice", open_id="")] + self.assertEqual(_build_mention_hint(refs), "[Mentioned: Alice]") + + def test_hint_uses_unknown_placeholder_when_name_missing(self): + from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint + + refs = [FeishuMentionRef(name="", open_id="ou_xxx")] + self.assertEqual(_build_mention_hint(refs), "[Mentioned: unknown (open_id=ou_xxx)]") def test_hint_dedupes_repeated_user(self): from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint @@ -2013,6 +4414,12 @@ def test_hint_dedupes_repeated_user(self): "[Mentioned: Alice (open_id=ou_alice), Bob (open_id=ou_bob)]", ) + def test_hint_dedupes_repeated_at_all(self): + from plugins.platforms.feishu.adapter import FeishuMentionRef, _build_mention_hint + + refs = [FeishuMentionRef(is_all=True), FeishuMentionRef(is_all=True)] + self.assertEqual(_build_mention_hint(refs), "[Mentioned: @all]") + class TestFeishuStripLeadingSelf(unittest.TestCase): def _make_refs(self, *, self_name="Hermes", other_name=None): @@ -2023,6 +4430,17 @@ def _make_refs(self, *, self_name="Hermes", other_name=None): refs.append(FeishuMentionRef(name=other_name, open_id="ou_alice")) return refs + def test_strips_leading_self(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions + + result = _strip_edge_self_mentions("@Hermes /help", self._make_refs()) + self.assertEqual(result, "/help") + + def test_strips_consecutive_leading_self(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions + + result = _strip_edge_self_mentions("@Hermes @Hermes hi", self._make_refs()) + self.assertEqual(result, "hi") def test_stops_at_first_non_self_token(self): from plugins.platforms.feishu.adapter import _strip_edge_self_mentions @@ -2032,6 +4450,17 @@ def test_stops_at_first_non_self_token(self): ) self.assertEqual(result, "@Alice make a group") + def test_preserves_mid_text_self(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions + + result = _strip_edge_self_mentions("check @Hermes said yesterday", self._make_refs()) + self.assertEqual(result, "check @Hermes said yesterday") + + def test_strips_trailing_self_at_end_of_text(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions + + result = _strip_edge_self_mentions("look up docs @Hermes", self._make_refs()) + self.assertEqual(result, "look up docs") def test_strips_trailing_self_with_terminal_punct(self): from plugins.platforms.feishu.adapter import _strip_edge_self_mentions @@ -2049,6 +4478,10 @@ def test_preserves_trailing_self_before_non_terminal_char(self): ) self.assertEqual(result, "please don't @Hermes anymore") + def test_returns_input_when_refs_empty(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions + + self.assertEqual(_strip_edge_self_mentions("@Hermes /help", []), "@Hermes /help") def test_returns_input_when_no_self_refs(self): from plugins.platforms.feishu.adapter import _strip_edge_self_mentions, FeishuMentionRef @@ -2056,8 +4489,26 @@ def test_returns_input_when_no_self_refs(self): refs = [FeishuMentionRef(name="Alice", open_id="ou_alice")] self.assertEqual(_strip_edge_self_mentions("@Alice hi", refs), "@Alice hi") + def test_uses_open_id_fallback_when_name_missing(self): + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions, FeishuMentionRef + + refs = [FeishuMentionRef(name="", open_id="ou_bot", is_self=True)] + self.assertEqual(_strip_edge_self_mentions("@ou_bot hi", refs), "hi") + + def test_word_boundary_prevents_prefix_collision(self): + """A bot named 'Al' must not eat the leading '@Alice' of a different user.""" + from plugins.platforms.feishu.adapter import _strip_edge_self_mentions, FeishuMentionRef + + refs = [FeishuMentionRef(name="Al", open_id="ou_bot", is_self=True)] + self.assertEqual(_strip_edge_self_mentions("@Alice hi", refs), "@Alice hi") + class TestFeishuNormalizeText(unittest.TestCase): + def test_renders_mention_with_display_name(self): + from plugins.platforms.feishu.adapter import _normalize_feishu_text, FeishuMentionRef + + refs = {"@_user_1": FeishuMentionRef(name="Alice", open_id="ou_alice")} + self.assertEqual(_normalize_feishu_text("@_user_1 hello", refs), "@Alice hello") def test_renders_self_mention_with_name(self): from plugins.platforms.feishu.adapter import _normalize_feishu_text, FeishuMentionRef @@ -2073,6 +4524,27 @@ def test_at_all_rendered_as_english_literal(self): self.assertEqual(_normalize_feishu_text("@_all notice", None), "@all notice") + def test_unknown_placeholder_degrades_to_space(self): + from plugins.platforms.feishu.adapter import _normalize_feishu_text + + # No map: fall back to the old behavior (substitute with space, then collapse). + self.assertEqual(_normalize_feishu_text("@_user_9 hello", None), "hello") + + def test_backward_compatible_without_map(self): + from plugins.platforms.feishu.adapter import _normalize_feishu_text + + self.assertEqual(_normalize_feishu_text("hello world"), "hello world") + + def test_mention_for_missing_map_entry_degrades_to_space(self): + from plugins.platforms.feishu.adapter import _normalize_feishu_text, FeishuMentionRef + + refs = {"@_user_1": FeishuMentionRef(name="Alice")} + # @_user_2 has no entry — should degrade to a space (legacy behavior) + self.assertEqual( + _normalize_feishu_text("@_user_1 @_user_2 hi", refs), + "@Alice hi", + ) + class TestFeishuPostMentionParsing(unittest.TestCase): def test_post_at_tag_renders_via_mentions_map(self): @@ -2095,6 +4567,38 @@ def test_post_at_tag_renders_via_mentions_map(self): result = parse_feishu_post_payload(payload, mentions_map=mentions_map) self.assertEqual(result.text_content, "@Alice hello") + def test_post_at_tag_falls_back_to_inline_user_name_when_map_misses(self): + """When the mentions payload is missing a placeholder, fall back to the + inline user_name in the tag itself.""" + from plugins.platforms.feishu.adapter import parse_feishu_post_payload + + payload = { + "en_us": { + "content": [[ + {"tag": "at", "user_id": "@_user_7", "user_name": "Unknown"}, + {"tag": "text", "text": " hi"}, + ]] + } + } + result = parse_feishu_post_payload(payload, mentions_map={}) + self.assertEqual(result.text_content, "@Unknown hi") + + def test_post_at_all_tag_renders_as_at_all(self): + """Post-format @everyone has user_id == '@_all' (confirmed via live + im.v1.message.get). Rendered as literal '@all' regardless of map.""" + from plugins.platforms.feishu.adapter import parse_feishu_post_payload + + payload = { + "en_us": { + "content": [[ + {"tag": "at", "user_id": "@_all", "user_name": "everyone"}, + {"tag": "text", "text": " meeting"}, + ]] + } + } + result = parse_feishu_post_payload(payload) + self.assertIn("@all", result.text_content) + class TestFeishuNormalizeWithMentions(unittest.TestCase): def test_text_message_renders_mention_by_name(self): @@ -2116,6 +4620,72 @@ def test_text_message_renders_mention_by_name(self): self.assertEqual(normalized.mentions[0].open_id, "ou_alice") self.assertFalse(normalized.mentions[0].is_self) + def test_text_message_marks_bot_self_mention(self): + from plugins.platforms.feishu.adapter import normalize_feishu_message, _FeishuBotIdentity + + mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_bot", user_id=""), + name="Hermes", + ) + normalized = normalize_feishu_message( + message_type="text", + raw_content=json.dumps({"text": "@_user_1 /help"}), + mentions=[mention], + bot=_FeishuBotIdentity(open_id="ou_bot"), + ) + self.assertTrue(normalized.mentions[0].is_self) + # self mention is still rendered — strip is a separate adapter-level pass + self.assertEqual(normalized.text_content, "@Hermes /help") + + def test_text_message_at_all_surfaces_ref(self): + from plugins.platforms.feishu.adapter import normalize_feishu_message + + mention = SimpleNamespace(key="@_all", id=None, name="") + normalized = normalize_feishu_message( + message_type="text", + raw_content=json.dumps({"text": "@_all meeting"}), + mentions=[mention], + ) + self.assertEqual(normalized.text_content, "@all meeting") + self.assertEqual(len(normalized.mentions), 1) + self.assertTrue(normalized.mentions[0].is_all) + + def test_text_message_at_all_in_text_without_mentions_payload(self): + """Feishu SDK sometimes omits @_all from the mentions payload (confirmed + via im.v1.message.get). The fallback scan on raw text must still yield + an is_all ref so [Mentioned: @all] gets injected.""" + from plugins.platforms.feishu.adapter import normalize_feishu_message + + normalized = normalize_feishu_message( + message_type="text", + raw_content=json.dumps({"text": "@_all hello"}), + mentions=None, + ) + self.assertEqual(normalized.text_content, "@all hello") + self.assertEqual(len(normalized.mentions), 1) + self.assertTrue(normalized.mentions[0].is_all) + + def test_text_message_at_all_not_synthesized_if_absent_from_text(self): + """No @_all in text → no synthetic ref even if mentions_map is empty.""" + from plugins.platforms.feishu.adapter import normalize_feishu_message + + normalized = normalize_feishu_message( + message_type="text", + raw_content=json.dumps({"text": "plain hello"}), + mentions=None, + ) + self.assertEqual(normalized.mentions, []) + + def test_text_message_without_mentions_param_is_backward_compatible(self): + from plugins.platforms.feishu.adapter import normalize_feishu_message + + normalized = normalize_feishu_message( + message_type="text", + raw_content=json.dumps({"text": "hello world"}), + ) + self.assertEqual(normalized.text_content, "hello world") + self.assertEqual(normalized.mentions, []) def test_post_message_marks_self_via_mentions_map_lookup(self): """Real Feishu post: + top-level mentions array @@ -2173,6 +4743,10 @@ def test_post_mentions_bot_uses_is_self_flag(self): ) ) + def test_post_mentions_bot_empty_returns_false(self): + adapter = self._build_adapter() + self.assertFalse(adapter._post_mentions_bot([])) + class TestFeishuExtractMessageContent(unittest.TestCase): def _build_adapter(self): @@ -2207,6 +4781,19 @@ def test_returns_five_tuple_with_mentions(self): self.assertEqual(len(mentions), 1) self.assertEqual(mentions[0].open_id, "ou_alice") + def test_returns_empty_mentions_when_missing(self): + adapter = self._build_adapter() + message = SimpleNamespace( + content=json.dumps({"text": "plain hello"}), + message_type="text", + message_id="m2", + mentions=None, + ) + + text, _, _, _, mentions = asyncio.run(adapter._extract_message_content(message)) + self.assertEqual(text, "plain hello") + self.assertEqual(mentions, []) + class TestFeishuProcessInboundMessage(unittest.TestCase): def _build_adapter(self): @@ -2227,6 +4814,37 @@ def _build_adapter(self): adapter._dispatch_inbound_event = AsyncMock() return adapter + def test_leading_self_mention_stripped_for_command(self): + from gateway.platforms.base import MessageType + + adapter = self._build_adapter() + bot_mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_bot", user_id=""), + name="Hermes", + ) + message = SimpleNamespace( + content=json.dumps({"text": "@_user_1 /help"}), + message_type="text", + message_id="m1", + mentions=[bot_mention], + chat_id="oc_chat", + parent_id=None, + upper_message_id=None, + thread_id=None, + ) + asyncio.run( + adapter._process_inbound_message( + data=message, + message=message, + sender_id=None, + chat_type="group", + message_id="m1", + ) + ) + event = adapter._dispatch_inbound_event.call_args.args[0] + self.assertEqual(event.text, "/help") + self.assertEqual(event.message_type, MessageType.COMMAND) def test_non_command_message_with_mentions_injects_hint(self): from gateway.platforms.base import MessageType @@ -2301,6 +4919,65 @@ def test_command_message_never_injects_hint(self): self.assertNotIn("[Mentioned:", event.text) self.assertTrue(event.text.startswith("/model")) + def test_mid_text_self_mention_preserved(self): + adapter = self._build_adapter() + bot_mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_bot", user_id=""), + name="Hermes", + ) + message = SimpleNamespace( + content=json.dumps({"text": "stop pinging @_user_1 please"}), + message_type="text", + message_id="m4", + mentions=[bot_mention], + chat_id="oc_chat", + parent_id=None, + upper_message_id=None, + thread_id=None, + ) + asyncio.run( + adapter._process_inbound_message( + data=message, + message=message, + sender_id=None, + chat_type="group", + message_id="m4", + ) + ) + event = adapter._dispatch_inbound_event.call_args.args[0] + self.assertEqual(event.text, "stop pinging @Hermes please") + + def test_pure_self_mention_message_is_ignored(self): + """A message containing only '@Bot' (no body, no media) must not dispatch. + + Regression guard: the rendered '@Hermes' slips past the pre-strip empty + guard; the post-strip guard must catch it. + """ + adapter = self._build_adapter() + bot_mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_bot", user_id=""), + name="Hermes", + ) + message = SimpleNamespace( + content=json.dumps({"text": "@_user_1"}), + message_type="text", + message_id="m5", + mentions=[bot_mention], + chat_id="oc_chat", + parent_id=None, + upper_message_id=None, + thread_id=None, + ) + asyncio.run( + adapter._process_inbound_message( + data=message, message=message, sender_id=None, + chat_type="group", message_id="m5", + ) + ) + adapter._dispatch_inbound_event.assert_not_called() + class TestFeishuFetchMessageText(unittest.TestCase): def _build_adapter(self): @@ -2315,6 +4992,51 @@ def _build_adapter(self): adapter._build_get_message_request = Mock(return_value=object()) return adapter + def test_fetch_message_text_renders_mentions_without_hint_prefix(self): + adapter = self._build_adapter() + + alice_mention = SimpleNamespace( + key="@_user_1", + id="ou_alice", + id_type="open_id", + name="Alice", + ) + parent = SimpleNamespace( + body=SimpleNamespace(content=json.dumps({"text": "@_user_1 hi"})), + msg_type="text", + mentions=[alice_mention], + ) + response = Mock() + response.success = Mock(return_value=True) + response.data = SimpleNamespace(items=[parent]) + adapter._client.im.v1.message.get = Mock(return_value=response) + + result = asyncio.run(adapter._fetch_message_text("m_parent")) + self.assertEqual(result, "@Alice hi") + # No [Mentioned:] wrapper — reply-context path intentionally skips the hint. + self.assertNotIn("[Mentioned:", result) + + def test_extract_text_from_raw_content_accepts_mentions_kwarg(self): + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter.__new__(FeishuAdapter) + adapter._bot_open_id = "" + adapter._bot_user_id = "" + adapter._bot_name = "" + + alice_mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_alice", user_id=""), + name="Alice", + ) + self.assertEqual( + adapter._extract_text_from_raw_content( + msg_type="text", + raw_content=json.dumps({"text": "@_user_1 hello"}), + mentions=[alice_mention], + ), + "@Alice hello", + ) def test_fetch_message_text_marks_is_self_via_string_id_shape(self): """History-path Mention objects carry id as str + id_type; is_self must still work.""" @@ -2342,6 +5064,24 @@ def test_fetch_message_text_marks_is_self_via_string_id_shape(self): result = asyncio.run(adapter._fetch_message_text("m_parent")) self.assertEqual(result, "@Hermes hi") + def test_build_mentions_map_string_id_shape(self): + """_build_mentions_map accepts the reply-history shape (id as str + + id_type='open_id'). user_id id_type is not load-bearing for self + detection — inbound mention payloads always include an open_id.""" + from plugins.platforms.feishu.adapter import _build_mentions_map, _FeishuBotIdentity + + # open_id discriminator, non-self + alice = SimpleNamespace(key="@_user_1", id="ou_alice", id_type="open_id", name="Alice") + ref = _build_mentions_map([alice], _FeishuBotIdentity(open_id="ou_bot"))["@_user_1"] + self.assertEqual(ref.open_id, "ou_alice") + self.assertFalse(ref.is_self) + + # open_id discriminator, is_self matches via open_id + bot_oid = SimpleNamespace(key="@_user_3", id="ou_bot", id_type="open_id", name="Hermes") + self.assertTrue( + _build_mentions_map([bot_oid], _FeishuBotIdentity(open_id="ou_bot"))["@_user_3"].is_self + ) + class TestFeishuMentionEndToEnd(unittest.TestCase): """High-level scenarios from the design spec — verify the full pipeline.""" @@ -2390,6 +5130,52 @@ def _run(self, adapter, text, mentions): ) return adapter._dispatch_inbound_event.call_args.args[0] + def test_scenario_bot_plus_alice_plus_bob_build_group(self): + adapter = self._build_adapter() + event = self._run( + adapter, + "@_user_1 @_user_2 @_user_3 build me a group", + [ + {"key": "@_user_1", "open_id": "ou_bot", "name": "Hermes"}, + {"key": "@_user_2", "open_id": "ou_alice", "name": "Alice"}, + {"key": "@_user_3", "open_id": "ou_bob", "name": "Bob"}, + ], + ) + self.assertIn("[Mentioned: Alice (open_id=ou_alice), Bob (open_id=ou_bob)]", event.text) + self.assertIn("@Alice @Bob build me a group", event.text) + self.assertNotIn("@Hermes", event.text) + + def test_scenario_at_all_announcement(self): + adapter = self._build_adapter() + event = self._run( + adapter, + "@_all meeting at 3pm", + [{"key": "@_all"}], + ) + self.assertTrue(event.text.startswith("[Mentioned: @all]")) + self.assertIn("@all meeting at 3pm", event.text) + + def test_scenario_trailing_self_mention_stripped(self): + """Trailing @bot at the end of a message is routing noise, not content — + strip it so the agent sees a clean instruction body.""" + adapter = self._build_adapter() + event = self._run( + adapter, + "who are you @_user_1", + [{"key": "@_user_1", "open_id": "ou_bot", "name": "Hermes"}], + ) + self.assertEqual(event.text, "who are you") + + def test_scenario_mid_text_self_mention_preserved(self): + """Self mention in the middle of a sentence (followed by a non-terminal + character) is meaningful content — preserve it.""" + adapter = self._build_adapter() + event = self._run( + adapter, + "please don't @_user_1 anymore", + [{"key": "@_user_1", "open_id": "ou_bot", "name": "Hermes"}], + ) + self.assertEqual(event.text, "please don't @Hermes anymore") def test_scenario_no_mentions_zero_regression(self): adapter = self._build_adapter() @@ -2397,6 +5183,42 @@ def test_scenario_no_mentions_zero_regression(self): self.assertEqual(event.text, "plain message") self.assertNotIn("[Mentioned:", event.text) + def test_scenario_post_at_alice_exposes_open_id(self): + """Post-type @mention: placeholder resolves via top-level mentions, + agent gets real open_id in the hint (mirrors text-type behavior).""" + adapter = self._build_adapter() + alice_mention = SimpleNamespace( + key="@_user_1", + id=SimpleNamespace(open_id="ou_alice", user_id=""), + name="Alice", + ) + post_content = json.dumps({ + "zh_cn": { + "content": [[ + {"tag": "at", "user_id": "@_user_1", "user_name": "Alice"}, + {"tag": "text", "text": " lookup this doc"}, + ]] + } + }) + message = SimpleNamespace( + content=post_content, + message_type="post", + message_id="m_post", + mentions=[alice_mention], + chat_id="oc_chat", + parent_id=None, + upper_message_id=None, + thread_id=None, + ) + asyncio.run( + adapter._process_inbound_message( + data=message, message=message, sender_id=None, + chat_type="group", message_id="m_post", + ) + ) + event = adapter._dispatch_inbound_event.call_args.args[0] + self.assertIn("[Mentioned: Alice (open_id=ou_alice)]", event.text) + self.assertIn("@Alice lookup this doc", event.text) def test_scenario_post_bot_plus_alice_filters_self_from_hint(self): """Post-type message @-ing both the bot and Alice: leading bot is @@ -2466,4 +5288,41 @@ def test_chat_locks_is_ordered_dict(self): adapter = self._make_adapter() self.assertIsInstance(adapter._chat_locks, _collections.OrderedDict) + def test_same_id_returns_same_lock_and_stays_bounded(self): + adapter = self._make_adapter(max_size=5) + locks = [adapter._get_chat_lock(f"c{i}") for i in range(5)] + self.assertEqual(len(adapter._chat_locks), 5) + # Re-requesting an existing id returns the identical lock, no growth. + self.assertIs(adapter._get_chat_lock("c2"), locks[2]) + self.assertEqual(len(adapter._chat_locks), 5) + + def test_lru_eviction_respects_recent_access(self): + adapter = self._make_adapter(max_size=5) + for i in range(5): + adapter._get_chat_lock(f"c{i}") + # Touch c0 so it is no longer the LRU entry, then add a new chat. + adapter._get_chat_lock("c0") + adapter._get_chat_lock("c_new") + self.assertEqual(len(adapter._chat_locks), 5) + self.assertNotIn("c1", adapter._chat_locks) # c1 was the true LRU + self.assertIn("c0", adapter._chat_locks) + self.assertIn("c_new", adapter._chat_locks) + + def test_eviction_skips_held_locks(self): + adapter = self._make_adapter(max_size=3) + + async def _run(): + held = adapter._get_chat_lock("held") + await held.acquire() + try: + adapter._get_chat_lock("x") + adapter._get_chat_lock("y") + # At capacity; "held" is LRU but locked, so "x" should go instead. + adapter._get_chat_lock("z") + self.assertIn("held", adapter._chat_locks) + self.assertNotIn("x", adapter._chat_locks) + self.assertEqual(len(adapter._chat_locks), 3) + finally: + held.release() + asyncio.run(_run()) diff --git a/tests/gateway/test_gap3_forget_verify.py b/tests/gateway/test_gap3_forget_verify.py new file mode 100644 index 000000000000..6cf3a6ae0c58 --- /dev/null +++ b/tests/gateway/test_gap3_forget_verify.py @@ -0,0 +1,106 @@ +"""Ad-hoc verification for forget_sessions wiring in _handle_delete_session (gap3).""" +from unittest.mock import patch + +import pytest +from aiohttp import web +from aiohttp.test_utils import TestClient, TestServer + +from gateway.config import PlatformConfig +from gateway.platforms.api_server import APIServerAdapter +from hermes_state import SessionDB + + +class FakeStore: + def __init__(self, result=1, raise_on_forget=False, sessions_dir="C:/tmp/sessions"): + self.sessions_dir = sessions_dir + self.result = result + self.raise_on_forget = raise_on_forget + self.calls = [] + + def forget_sessions(self, ids): + self.calls.append(list(ids)) + if self.raise_on_forget: + raise RuntimeError("boom") + return self.result + + +def _make_app(adapter): + app = web.Application() + app.router.add_delete("/api/sessions/{session_id}", adapter._handle_delete_session) + return app + + +@pytest.fixture +def db(tmp_path): + d = SessionDB(tmp_path / "state.db") + try: + yield d + finally: + close = getattr(d, "close", None) + if callable(close): + close() + + +def _adapter(db, store): + a = APIServerAdapter(PlatformConfig(enabled=True)) + a._session_db = db + a._session_store = store + return a + + +@pytest.mark.asyncio +async def test_delete_calls_forget_sessions_and_logs_result(db, tmp_path): + db.create_session("sess-1", "api_server") + store = FakeStore(result=1, sessions_dir=str(tmp_path / "sessions")) + app = _make_app(_adapter(db, store)) + async with TestClient(TestServer(app)) as cli: + with patch("gateway.platforms.api_server.logger") as mock_logger: + resp = await cli.delete("/api/sessions/sess-1") + assert resp.status == 200 + body = await resp.json() + assert body == {"object": "hermes.session.deleted", "id": "sess-1", "deleted": True} + assert store.calls == [["sess-1"]] + info_args = [c.args[0] for c in mock_logger.info.call_args_list] + assert any("forget_sessions removed" in a for a in info_args) + assert mock_logger.warning.call_count == 0 + + +@pytest.mark.asyncio +async def test_delete_store_without_forget_sessions_is_skipped(db, tmp_path): + db.create_session("sess-2", "api_server") + + class NoForgetStore(FakeStore): + forget_sessions = None # attribute present but not callable + + store = NoForgetStore(sessions_dir=str(tmp_path / "sessions")) + app = _make_app(_adapter(db, store)) + async with TestClient(TestServer(app)) as cli: + with patch("gateway.platforms.api_server.logger") as mock_logger: + resp = await cli.delete("/api/sessions/sess-2") + assert resp.status == 200 + assert (await resp.json())["deleted"] is True + assert store.calls == [] + assert mock_logger.debug.call_count == 1 + assert mock_logger.warning.call_count == 0 + + +@pytest.mark.asyncio +async def test_delete_forget_exception_never_fails_delete(db, tmp_path): + db.create_session("sess-3", "api_server") + store = FakeStore(raise_on_forget=True, sessions_dir=str(tmp_path / "sessions")) + app = _make_app(_adapter(db, store)) + async with TestClient(TestServer(app)) as cli: + resp = await cli.delete("/api/sessions/sess-3") + assert resp.status == 200 + assert (await resp.json())["deleted"] is True + assert store.calls == [["sess-3"]] + + +@pytest.mark.asyncio +async def test_delete_unknown_session_404_does_not_call_forget(db, tmp_path): + store = FakeStore(sessions_dir=str(tmp_path / "sessions")) + app = _make_app(_adapter(db, store)) + async with TestClient(TestServer(app)) as cli: + resp = await cli.delete("/api/sessions/ghost-session") + assert resp.status == 404 + assert store.calls == [] diff --git a/tests/gateway/test_session_reset_notify.py b/tests/gateway/test_session_reset_notify.py index c5783f467397..2c3defd8d5ba 100644 --- a/tests/gateway/test_session_reset_notify.py +++ b/tests/gateway/test_session_reset_notify.py @@ -170,7 +170,12 @@ def test_was_auto_reset_persists_across_roundtrip(self, tmp_path): def _make_db_mock() -> MagicMock: """Return a SessionDB mock with safe defaults for all lookup methods.""" db = MagicMock() - db.get_session.return_value = None + db.get_session.return_value = { + "id": "sess", + "end_reason": None, + } # alive row: end_reason=None; a bare MagicMock (truthy) would make + # every session look "ended" and _is_session_ended_in_db would drop the + # routing entry even when the reset policy says keep it db.get_compression_tip.return_value = None # avoids MagicMock leaking into session_id db.find_latest_gateway_session_for_peer.return_value = None db.reopen_session.return_value = None diff --git a/tests/gateway/test_session_store_lock_io.py b/tests/gateway/test_session_store_lock_io.py index 61990c3f74a0..0b403e184bf2 100644 --- a/tests/gateway/test_session_store_lock_io.py +++ b/tests/gateway/test_session_store_lock_io.py @@ -56,15 +56,57 @@ def held(self) -> bool: return self._held +class _FakeCursor: + """Fake sqlite cursor returning pre-built rows.""" + + def __init__(self, rows): + self._rows = rows + + def fetchall(self): + return self._rows + + +class _FakeConn: + """Fake connection for the #GAP-3 backstop read path.""" + + def __init__(self, rows): + self._rows = rows + + def execute(self, sql, params=None): + # Mirrors _query_existing_session_ids: SELECT id FROM sessions + # WHERE id IN (...) — return only ids that exist in the mock store. + ids = params if isinstance(params, (list, tuple)) else [] + return _FakeCursor([{"id": sid} for sid in ids if sid in self._rows]) + + +class _FakeReadCtx: + """Context manager standing in for SessionDB._read_ctx.""" + + def __init__(self, rows): + self._rows = rows + + def __enter__(self): + return _FakeConn(self._rows) + + def __exit__(self, *exc): + return False + + def _db_with_rows(rows: dict) -> MagicMock: """Mock SessionDB where ``get_session`` maps session_id -> row dict.""" db = MagicMock() db.get_session.side_effect = lambda sid: rows.get(sid) db.find_latest_gateway_session_for_peer.return_value = None db.reopen_session.return_value = None - db.create_session.return_value = None + db.create_session.side_effect = lambda **kwargs: rows.setdefault( + kwargs["session_id"], {"id": kwargs["session_id"], "end_reason": None} + ) # Identity compression tip (no child session). db.get_compression_tip.side_effect = lambda sid: sid + # The #GAP-3 backstop (_prune_dead_persisted_entries) reads existence + # through _read_ctx; without a faithful fake it sees an empty store and + # drops the just-created entry as "dead". Mirror the real DB. + db._read_ctx.side_effect = lambda: _FakeReadCtx(rows) return db @@ -123,10 +165,10 @@ def test_is_session_ended_not_holding_lock(self, tmp_path): orig = store._is_session_ended_in_db - def tracking(sid, **kw): + def tracking(sid, **kwargs): if lock.held: calls_under_lock.append(sid) - return orig(sid, **kw) + return orig(sid, **kwargs) store._is_session_ended_in_db = tracking # type: ignore[method-assign] @@ -227,6 +269,33 @@ def synchronized_query(**kwargs): assert created_ids == {entries[0].session_id} +def test_concurrent_force_new_returns_one_published_session(tmp_path): + """Concurrent /new delivery must not create orphan SQLite sessions.""" + source = _source() + db = _db_with_rows({}) + store = _make_store(tmp_path, db) + owner_started = threading.Event() + release_owner = threading.Event() + original_impl = store._get_or_create_session_impl + + def synchronized_impl(*args, **kwargs): + owner_started.set() + assert release_owner.wait(timeout=10) + return original_impl(*args, **kwargs) + + store._get_or_create_session_impl = synchronized_impl # type: ignore[method-assign] + with ThreadPoolExecutor(max_workers=2) as pool: + owner = pool.submit(store.get_or_create_session, source, True) + assert owner_started.wait(timeout=10) + follower = pool.submit(store.get_or_create_session, source, True) + release_owner.set() + entries = [owner.result(timeout=10), follower.result(timeout=10)] + + assert entries[0] is entries[1] + created_ids = {call.kwargs["session_id"] for call in db.create_session.call_args_list} + assert created_ids == {entries[0].session_id} + + def test_auto_reset_does_not_recover_session_being_ended(tmp_path): source = _source() db = _db_with_rows({}) @@ -254,3 +323,109 @@ def test_auto_reset_does_not_recover_session_being_ended(tmp_path): db.end_session.assert_not_called() +def test_legacy_and_off_lock_saves_share_one_serialization_lock(tmp_path): + db = _db_with_rows({}) + persisted: dict[str, str] = {} + first_write_started = threading.Event() + release_first_write = threading.Event() + write_count = 0 + count_lock = threading.Lock() + + def replace(entries, *, scope): + nonlocal write_count, persisted + with count_lock: + write_count += 1 + call_number = write_count + if call_number == 1: + first_write_started.set() + assert release_first_write.wait(timeout=10) + persisted = dict(entries) + + db.replace_gateway_routing_entries.side_effect = replace + store = _make_store(tmp_path, db) + source_a = _source() + source_b = SessionSource( + platform=Platform.TELEGRAM, + chat_id="67890", + chat_type="dm", + user_id="67890", + ) + key_a = store._generate_session_key(source_a) + key_b = store._generate_session_key(source_b) + _seed_entry(store, key_a, "sid-a") + + with ThreadPoolExecutor(max_workers=2) as pool: + future_a = pool.submit(store._save_entries) + assert first_write_started.wait(timeout=10) + _seed_entry(store, key_b, "sid-b") + future_b = pool.submit(store._save) + release_first_write.set() + future_a.result(timeout=10) + future_b.result(timeout=10) + + assert set(persisted) == {key_a, key_b} + + +def test_save_serialization_snapshots_latest_routing_index(tmp_path): + """A delayed earlier writer must snapshot the state visible when it writes.""" + db = _db_with_rows({}) + persisted: dict[str, str] = {} + first_write_started = threading.Event() + release_first_write = threading.Event() + write_count = 0 + count_lock = threading.Lock() + + def replace(entries, *, scope): + nonlocal write_count, persisted + with count_lock: + write_count += 1 + call_number = write_count + if call_number == 1: + first_write_started.set() + assert release_first_write.wait(timeout=10) + persisted = dict(entries) + + db.replace_gateway_routing_entries.side_effect = replace + store = _make_store(tmp_path, db) + source_a = _source() + source_b = SessionSource( + platform=Platform.TELEGRAM, + chat_id="67890", + chat_type="dm", + user_id="67890", + ) + key_a = store._generate_session_key(source_a) + key_b = store._generate_session_key(source_b) + entry_a = _seed_entry(store, key_a, "sid-a") + + with ThreadPoolExecutor(max_workers=2) as pool: + future_a = pool.submit(store._save_entries) + assert first_write_started.wait(timeout=10) + entry_b = _seed_entry(store, key_b, "sid-b") + future_b = pool.submit(store._save_entries) + release_first_write.set() + future_a.result(timeout=10) + future_b.result(timeout=10) + + assert set(store._entries) == {key_a, key_b} + assert set(persisted) == {key_a, key_b} + assert json.loads(persisted[key_a])["session_id"] == entry_a.session_id + assert json.loads(persisted[key_b])["session_id"] == entry_b.session_id + + +def test_recovery_rejects_other_profile_row(tmp_path, monkeypatch): + """The lock-free recovery path must retain the canonical profile guard.""" + source = _source() + db = _db_with_rows({}) + db.find_latest_gateway_session_for_peer.return_value = { + "id": "foreign-session", + "session_key": "agent:other:telegram:dm:12345", + "started_at": datetime.now().timestamp(), + } + store = _make_store(tmp_path, db) + monkeypatch.setattr(store, "_active_profile_name", lambda: "default") + + entry = store.get_or_create_session(source) + + assert entry.session_id != "foreign-session" + db.reopen_session.assert_not_called() diff --git a/tests/gateway/test_setup_feishu.py b/tests/gateway/test_setup_feishu.py index 6ae9fe228e91..d982b155e4cf 100644 --- a/tests/gateway/test_setup_feishu.py +++ b/tests/gateway/test_setup_feishu.py @@ -7,6 +7,22 @@ import os from unittest.mock import patch +# Windows-safe env for ``@patch.dict(..., clear=True)`` tests: the decorator +# wipes the whole environment, and on Windows ``pathlib.Path.home()`` raises +# "Could not determine home directory" when USERPROFILE / HOME / +# HOMEDRIVE+HOMEPATH are all absent (POSIX falls back to the ``pwd`` module, +# which is why these tests only broke on Windows). ``get_hermes_home()`` +# reads HERMES_HOME first, then LOCALAPPDATA, then Path.home(). Preserving +# these variables keeps the "no Feishu env vars" semantics of clear=True +# while letting home resolution work on every platform. +_HOME_ENV = { + k: v + for k, v in os.environ.items() + if k in ("HERMES_HOME", "USERPROFILE", "HOMEDRIVE", "HOMEPATH", "HOME", "LOCALAPPDATA") + and v +} + + # --------------------------------------------------------------------------- # Helpers @@ -76,6 +92,23 @@ def mock_remove(name): class TestSetupFeishuQrPath: """Tests for the QR scan-to-create happy path.""" + def test_qr_success_saves_core_credentials(self): + env, _ = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", + "app_secret": "secret_test", + "domain": "feishu", + "open_id": "ou_owner", + "bot_name": "TestBot", + "bot_open_id": "ou_bot", + }, + prompt_yes_no_responses=[True], # Start QR + prompt_choice_responses=[0, 0, 0], # method=QR, dm=pairing, group=open + prompt_responses=[""], # home channel: skip + ) + assert env["FEISHU_APP_ID"] == "cli_test" + assert env["FEISHU_APP_SECRET"] == "secret_test" + assert env["FEISHU_DOMAIN"] == "feishu" def test_qr_success_does_not_persist_bot_identity(self): """Bot identity is discovered at runtime by _hydrate_bot_identity — not persisted @@ -104,6 +137,16 @@ def test_qr_success_does_not_persist_bot_identity(self): class TestSetupFeishuConnectionMode: """Connection mode: QR always websocket, manual path lets user choose.""" + def test_qr_path_defaults_to_websocket(self): + env, _ = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", "app_secret": "s", "domain": "feishu", + "open_id": None, "bot_name": None, "bot_open_id": None, + }, + prompt_choice_responses=[0, 0, 0], # method=QR, dm=pairing, group=open + prompt_responses=[""], + ) + assert env["FEISHU_CONNECTION_MODE"] == "websocket" @patch("plugins.platforms.feishu.adapter.probe_bot", return_value=None) def test_manual_path_websocket(self, _mock_probe): @@ -114,6 +157,15 @@ def test_manual_path_websocket(self, _mock_probe): ) assert env["FEISHU_CONNECTION_MODE"] == "websocket" + @patch("plugins.platforms.feishu.adapter.probe_bot", return_value=None) + def test_manual_path_webhook(self, _mock_probe): + env, _ = _run_setup_feishu( + qr_result=None, + prompt_choice_responses=[1, 0, 1, 0, 0], # method=manual, domain=feishu, connection=webhook, dm=pairing, group=open + prompt_responses=["cli_manual", "secret_manual", ""], # app_id, app_secret, home_channel + ) + assert env["FEISHU_CONNECTION_MODE"] == "webhook" + # --------------------------------------------------------------------------- # DM security policy @@ -134,6 +186,17 @@ def _run_with_dm_choice(self, dm_choice_idx, prompt_responses=None): ) return env + def test_pairing_sets_feishu_allow_all_false(self): + env = self._run_with_dm_choice(0) + assert env["FEISHU_ALLOW_ALL_USERS"] == "false" + assert env["FEISHU_ALLOWED_USERS"] == "" + assert "GATEWAY_ALLOW_ALL_USERS" not in env + + def test_allow_all_sets_feishu_allow_all_true(self): + env = self._run_with_dm_choice(1) + assert env["FEISHU_ALLOW_ALL_USERS"] == "true" + assert env["FEISHU_ALLOWED_USERS"] == "" + assert "GATEWAY_ALLOW_ALL_USERS" not in env def test_allowlist_sets_feishu_allow_all_false_with_list(self): env = self._run_with_dm_choice(2, prompt_responses=["ou_user1,ou_user2", ""]) @@ -141,6 +204,13 @@ def test_allowlist_sets_feishu_allow_all_false_with_list(self): assert env["FEISHU_ALLOWED_USERS"] == "ou_user1,ou_user2" assert "GATEWAY_ALLOW_ALL_USERS" not in env + def test_allowlist_prepopulates_with_scan_owner_open_id(self): + """When open_id is available from QR scan, it should be the default allowlist value.""" + # We return the owner's open_id from prompt (+ empty home channel). + env = self._run_with_dm_choice(2, prompt_responses=["ou_owner", ""]) + assert env["FEISHU_ALLOWED_USERS"] == "ou_owner" + + # --------------------------------------------------------------------------- # Group policy @@ -160,6 +230,18 @@ def test_open_with_mention(self): ) assert env["FEISHU_GROUP_POLICY"] == "open" + def test_disabled(self): + env, _ = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", "app_secret": "s", "domain": "feishu", + "open_id": None, "bot_name": None, "bot_open_id": None, + }, + prompt_yes_no_responses=[True], + prompt_choice_responses=[0, 0, 1], # method=QR, dm=pairing, group=disabled + prompt_responses=[""], + ) + assert env["FEISHU_GROUP_POLICY"] == "disabled" + # --------------------------------------------------------------------------- # Home channel (optional clear — Issue #12423) @@ -182,6 +264,48 @@ def test_blank_removes_existing_home_channel(self): assert "FEISHU_HOME_CHANNEL" in removed assert "FEISHU_HOME_CHANNEL" not in env + def test_blank_without_prior_home_still_attempts_remove(self): + _, removed = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", "app_secret": "s", "domain": "feishu", + "open_id": None, "bot_name": None, "bot_open_id": None, + }, + prompt_yes_no_responses=[True], + prompt_choice_responses=[0, 0, 0], + prompt_responses=[""], + existing_env={}, + ) + assert removed.count("FEISHU_HOME_CHANNEL") == 1 + + def test_nonempty_saves_home_channel(self): + env, removed = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", "app_secret": "s", "domain": "feishu", + "open_id": None, "bot_name": None, "bot_open_id": None, + }, + prompt_yes_no_responses=[True], + prompt_choice_responses=[0, 0, 0], + prompt_responses=["oc_chat123"], + existing_env={}, + ) + assert env["FEISHU_HOME_CHANNEL"] == "oc_chat123" + assert "FEISHU_HOME_CHANNEL" not in removed + + def test_whitespace_only_clears_home_channel(self): + """Whitespace-only input should clear, not save.""" + env, removed = _run_setup_feishu( + qr_result={ + "app_id": "cli_test", "app_secret": "s", "domain": "feishu", + "open_id": None, "bot_name": None, "bot_open_id": None, + }, + prompt_yes_no_responses=[True], + prompt_choice_responses=[0, 0, 0], + prompt_responses=[" "], + existing_env={"FEISHU_HOME_CHANNEL": "chat_old"}, + ) + assert "FEISHU_HOME_CHANNEL" in removed + assert "FEISHU_HOME_CHANNEL" not in env + # --------------------------------------------------------------------------- # Adapter integration: env vars → FeishuAdapterSettings @@ -211,12 +335,12 @@ def _make_env_from_setup(self, dm_idx=0, group_idx=0): ) return env - @patch.dict(os.environ, {}, clear=True) + @patch.dict(os.environ, {**_HOME_ENV}, clear=True) def test_qr_env_produces_valid_adapter_settings(self): """QR setup → adapter initializes with websocket mode.""" env = self._make_env_from_setup() - with patch.dict(os.environ, env, clear=True): + with patch.dict(os.environ, {**_HOME_ENV, **env}, clear=True): from gateway.config import PlatformConfig from plugins.platforms.feishu.adapter import FeishuAdapter adapter = FeishuAdapter(PlatformConfig()) @@ -225,4 +349,25 @@ def test_qr_env_produces_valid_adapter_settings(self): assert adapter._domain_name == "feishu" assert adapter._connection_mode == "websocket" + @patch.dict(os.environ, {**_HOME_ENV}, clear=True) + def test_open_dm_env_sets_correct_adapter_state(self): + """Setup with 'allow all DMs' → adapter sees allow-all flag.""" + env = self._make_env_from_setup(dm_idx=1) + + with patch.dict(os.environ, {**_HOME_ENV, **env}, clear=True): + from plugins.platforms.feishu.adapter import FeishuAdapter + from gateway.config import PlatformConfig + # Verify adapter initializes without error and env var is correct. + FeishuAdapter(PlatformConfig()) + assert os.getenv("FEISHU_ALLOW_ALL_USERS") == "true" + + @patch.dict(os.environ, {**_HOME_ENV}, clear=True) + def test_group_open_env_sets_adapter_group_policy(self): + """Setup with 'open groups' → adapter group_policy is 'open'.""" + env = self._make_env_from_setup(group_idx=0) + with patch.dict(os.environ, {**_HOME_ENV, **env}, clear=True): + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + adapter = FeishuAdapter(PlatformConfig()) + assert adapter._group_policy == "open" diff --git a/tests/gateway/test_slack_block_kit_adapter.py b/tests/gateway/test_slack_block_kit_adapter.py index f77a1362b81e..47505bc0c10f 100644 --- a/tests/gateway/test_slack_block_kit_adapter.py +++ b/tests/gateway/test_slack_block_kit_adapter.py @@ -7,6 +7,7 @@ * multi-chunk (>39k) messages fall back to plain text """ +from types import ModuleType from unittest.mock import AsyncMock, MagicMock, call import pytest @@ -66,6 +67,17 @@ async def test_disabled_by_default_no_blocks(self): assert "blocks" not in kwargs assert kwargs["text"] # plain text still sent + @pytest.mark.asyncio + async def test_enabled_sends_blocks_with_text_fallback(self): + adapter, client = _make_adapter({"rich_blocks": True}) + await adapter.send("C1", RICH_MD) + kwargs = client.chat_postMessage.await_args.kwargs + assert "blocks" in kwargs and kwargs["blocks"] + # text fallback is ALWAYS present alongside blocks (notifications/a11y) + assert kwargs["text"] + types = [b["type"] for b in kwargs["blocks"]] + assert "header" in types + assert "divider" in types @pytest.mark.asyncio async def test_enabled_but_unrenderable_falls_back_to_text(self): @@ -76,6 +88,21 @@ async def test_enabled_but_unrenderable_falls_back_to_text(self): assert "blocks" not in kwargs assert kwargs["text"] + @pytest.mark.asyncio + async def test_string_true_coerced(self): + adapter, client = _make_adapter({"rich_blocks": "true"}) + await adapter.send("C1", RICH_MD) + assert "blocks" in client.chat_postMessage.await_args.kwargs + + @pytest.mark.asyncio + async def test_multichunk_message_no_blocks(self): + adapter, client = _make_adapter({"rich_blocks": True}) + huge = "word " * 20000 # well over MAX_MESSAGE_LENGTH -> chunked + await adapter.send("C1", huge) + # every posted chunk is plain text, none carry blocks + for c in client.chat_postMessage.await_args_list: + assert "blocks" not in c.kwargs + assert c.kwargs["text"] @pytest.mark.asyncio async def test_feedback_buttons_opt_in_appended_to_blocks(self): @@ -89,6 +116,38 @@ async def test_feedback_buttons_opt_in_appended_to_blocks(self): assert feedback["elements"][0]["type"] == "feedback_buttons" assert feedback["elements"][0]["action_id"] == "hermes_feedback" + @pytest.mark.asyncio + async def test_feedback_buttons_require_rich_blocks(self): + """feedback_buttons alone must not implicitly enable Block Kit rendering.""" + adapter, client = _make_adapter({"feedback_buttons": True}) + + await adapter.send("C1", "final answer") + + assert "blocks" not in client.chat_postMessage.await_args.kwargs + + @pytest.mark.asyncio + async def test_block_rejection_retries_send_without_blocks_using_workspace_client(self): + adapter, client = _make_adapter({"rich_blocks": True}) + client.chat_postMessage = AsyncMock( + side_effect=[SlackRejectedBlocks("invalid_blocks"), {"ts": "111.333"}] + ) + + result = await adapter.send( + "C1", RICH_TABLE_MD, metadata={"team_id": "T_SECONDARY"} + ) + + assert result.success is True + assert adapter._get_client.call_args_list == [ + call("C1", team_id="T_SECONDARY"), + call("C1", team_id="T_SECONDARY"), + ] + assert client.chat_postMessage.await_count == 2 + first = client.chat_postMessage.await_args_list[0].kwargs + second = client.chat_postMessage.await_args_list[1].kwargs + assert "blocks" in first and first["blocks"] + assert "blocks" not in second + assert second["text"] + class TestEditMessageBlocks: @pytest.mark.asyncio @@ -107,6 +166,18 @@ async def test_finalize_edit_gets_blocks(self): assert "blocks" in kwargs and kwargs["blocks"] assert kwargs["text"] + @pytest.mark.asyncio + async def test_finalize_edit_gets_feedback_buttons_when_enabled(self): + adapter, client = _make_adapter({"rich_blocks": True, "feedback_buttons": True}) + await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + blocks = client.chat_update.await_args.kwargs["blocks"] + assert blocks[-1]["elements"][0]["type"] == "feedback_buttons" + + @pytest.mark.asyncio + async def test_finalize_edit_disabled_no_blocks(self): + adapter, client = _make_adapter() # rich_blocks off + await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + assert "blocks" not in client.chat_update.await_args.kwargs @pytest.mark.asyncio async def test_block_rejection_retries_edit_without_blocks_using_workspace_client(self): @@ -146,6 +217,144 @@ async def test_timeout_error_on_edit_is_retryable_transient(self): assert result.retryable is True assert result.error_kind == "transient" + @pytest.mark.asyncio + async def test_dns_connection_error_on_edit_is_retryable_transient(self): + from aiohttp import ClientConnectorDNSError + + adapter, client = _make_adapter() + client.chat_update = AsyncMock( + side_effect=ClientConnectorDNSError( + _slack_connection_key(), + OSError(8, "nodename nor servname provided, or not known"), + ) + ) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is True + assert result.error_kind == "transient" + + @pytest.mark.asyncio + async def test_slack_api_error_on_edit_is_not_retryable(self): + # Real slack_sdk required: the test pins that a genuine SlackApiError + # is never misclassified as transient. CI shards without the slack + # extras skip (adapter classification is still covered by the + # OSError/timeout tests above, which use stdlib exceptions). + errors_mod = pytest.importorskip("slack_sdk.errors") + # The gateway suite registers a MagicMock ``slack_sdk`` in sys.modules + # at collection time (see test_send_multiple_images.py / test_slack.py) + # whenever the real package is not installed. ``importorskip`` cannot + # detect that mock — it imports fine — so a mock SlackApiError would + # silently run the test with wrong semantics (AsyncMock side_effect + # CALLS a Mock instead of raising). Skip unless the real module is + # present. + if not isinstance(errors_mod, ModuleType): + pytest.skip( + "real slack_sdk required — only a collection-time mock is " + "registered in sys.modules" + ) + SlackApiError = errors_mod.SlackApiError + + adapter, client = _make_adapter() + client.chat_update = AsyncMock( + side_effect=SlackApiError( + "message_not_found", + {"ok": False, "error": "message_not_found"}, + ) + ) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is not True + assert result.error_kind != "transient" + + @pytest.mark.asyncio + async def test_certificate_error_on_edit_is_not_retryable(self): + import ssl + + from aiohttp import ClientConnectorCertificateError + + adapter, client = _make_adapter() + client.chat_update = AsyncMock( + side_effect=ClientConnectorCertificateError( + _slack_connection_key(), + ssl.SSLCertVerificationError("certificate verify failed"), + ) + ) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is not True + assert result.error_kind != "transient" + + @pytest.mark.asyncio + async def test_tls_integrity_errors_on_edit_are_not_retryable(self): + import ssl + + from aiohttp import ClientConnectorSSLError, ServerFingerprintMismatch + + errors = ( + ClientConnectorSSLError( + _slack_connection_key(), ssl.SSLError("handshake failed") + ), + ServerFingerprintMismatch(b"expected", b"got", "slack.com", 443), + ) + for error in errors: + adapter, client = _make_adapter() + client.chat_update = AsyncMock(side_effect=error) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is not True + assert result.error_kind != "transient" + + @pytest.mark.asyncio + async def test_plain_os_error_on_edit_is_not_retryable(self): + adapter, client = _make_adapter() + client.chat_update = AsyncMock(side_effect=OSError("invalid local socket state")) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is not True + assert result.error_kind != "transient" + + @pytest.mark.asyncio + async def test_lazy_rebound_aiohttp_connection_error_is_retryable( + self, monkeypatch + ): + # Exercises the REAL lazy-import rebind path in + # check_slack_requirements — requires slack_bolt/slack_sdk installed. + pytest.importorskip("slack_bolt") + pytest.importorskip("slack_sdk") + import tools.lazy_deps as lazy_deps + + monkeypatch.setattr(slack_module, "SLACK_AVAILABLE", False) + monkeypatch.delattr(slack_module, "aiohttp", raising=False) + + def ensure_and_bind(_group, import_fn, target_globals, *, prompt): + assert prompt is False + target_globals.update(import_fn()) + return True + + monkeypatch.setattr(lazy_deps, "ensure_and_bind", ensure_and_bind) + + assert slack_module.check_slack_requirements() is True + adapter, client = _make_adapter() + client.chat_update = AsyncMock( + side_effect=slack_module.aiohttp.ClientConnectionError("connection dropped") + ) + + result = await adapter.edit_message("C1", "111.222", RICH_MD, finalize=True) + + assert result.success is False + assert result.retryable is True + assert result.error_kind == "transient" + # --------------------------------------------------------------------------- # markdown_blocks mode — Slack's native ``markdown`` Block Kit block (#8552) @@ -175,6 +384,44 @@ async def test_enabled_sends_markdown_block_with_raw_content(self): # mrkdwn fallback text is still present for notifications/search assert kwargs["text"] + @pytest.mark.asyncio + async def test_text_fallback_is_mrkdwn_converted(self): + adapter, client = _make_adapter({"markdown_blocks": True}) + await adapter.send("C1", "**bold**") + kwargs = client.chat_postMessage.await_args.kwargs + assert kwargs["blocks"][0]["text"] == "**bold**" + assert kwargs["text"] == "*bold*" # mrkdwn conversion for fallback + + @pytest.mark.asyncio + async def test_markdown_block_preferred_over_rich_blocks(self): + adapter, client = _make_adapter( + {"markdown_blocks": True, "rich_blocks": True} + ) + await adapter.send("C1", RICH_TABLE_MD) + blocks = client.chat_postMessage.await_args.kwargs["blocks"] + assert blocks[0]["type"] == "markdown" + + @pytest.mark.asyncio + async def test_over_cap_falls_back_to_rich_or_text(self): + adapter, client = _make_adapter({"markdown_blocks": True}) + big = "x" * (SlackAdapter._MARKDOWN_BLOCK_MAX + 1) + payload = adapter._markdown_block_payload(big) + assert payload is None # declines >12k cumulative markdown cap + + @pytest.mark.asyncio + async def test_rejection_retries_without_blocks(self): + """Workspaces/surfaces without markdown-block support degrade to + the plain mrkdwn text payload instead of dropping the message.""" + adapter, client = _make_adapter({"markdown_blocks": True}) + client.chat_postMessage = AsyncMock( + side_effect=[SlackRejectedBlocks(), {"ts": "111.222"}] + ) + result = await adapter.send("C1", RICH_TABLE_MD) + assert result.success is True + assert client.chat_postMessage.await_count == 2 + retry_kwargs = client.chat_postMessage.await_args_list[1].kwargs + assert "blocks" not in retry_kwargs + assert retry_kwargs["text"] @pytest.mark.asyncio async def test_edit_finalize_uses_markdown_block(self): @@ -184,4 +431,14 @@ async def test_edit_finalize_uses_markdown_block(self): assert kwargs["blocks"][0]["type"] == "markdown" assert kwargs["blocks"][0]["text"] == RICH_TABLE_MD + @pytest.mark.asyncio + async def test_edit_streaming_stays_plain(self): + adapter, client = _make_adapter({"markdown_blocks": True}) + await adapter.edit_message("C1", "111.222", RICH_TABLE_MD, finalize=False) + kwargs = client.chat_update.await_args.kwargs + assert "blocks" not in kwargs + def test_empty_content_declines(self): + adapter, _ = _make_adapter({"markdown_blocks": True}) + assert adapter._markdown_block_payload("") is None + assert adapter._markdown_block_payload(" ") is None diff --git a/tests/hermes_state/test_aux_usage_accounting.py b/tests/hermes_state/test_aux_usage_accounting.py index dd1600bd1894..04d0e9922eeb 100644 --- a/tests/hermes_state/test_aux_usage_accounting.py +++ b/tests/hermes_state/test_aux_usage_accounting.py @@ -67,7 +67,28 @@ def test_accumulates_same_task_and_model(self, db): assert rows[0]["input_tokens"] == 3000 assert rows[0]["api_call_count"] == 3 - + def test_task_rows_do_not_touch_session_counters(self, db): + """Aux usage must NOT increment sessions.input_tokens — the gateway + overwrites those with absolute main-loop totals.""" + db.create_session("s1", source="cli") + db.record_auxiliary_usage("s1", "vision", model="m", input_tokens=999) + sess = db.get_session("s1") + assert (sess.get("input_tokens") or 0) == 0 + + def test_task_row_does_not_inherit_session_route(self, db): + """An aux call on a different provider must not borrow the session's + main-loop model/provider.""" + db.create_session("s1", source="cli", model="anthropic/claude-opus-4.6") + db.update_token_counts( + "s1", input_tokens=10, model="anthropic/claude-opus-4.6", + billing_provider="anthropic", api_call_count=1, + ) + db.record_auxiliary_usage("s1", "vision", input_tokens=5) # no model given + rows = {r["task"]: r for r in _usage_rows(db, "s1")} + assert rows["vision"]["model"] == "unknown" + assert rows["vision"]["billing_provider"] == "" + # main-loop row unaffected + assert rows[""]["model"] == "anthropic/claude-opus-4.6" def test_main_loop_and_aux_rows_coexist(self, db): db.create_session("s1", source="cli") @@ -83,7 +104,20 @@ def test_main_loop_and_aux_rows_coexist(self, db): tasks = sorted(r["task"] for r in rows) assert tasks == ["", "title_generation"] + def test_noop_without_session_or_task(self, db): + db.record_auxiliary_usage("", "vision", input_tokens=5) + db.create_session("s1", source="cli") + db.record_auxiliary_usage("s1", "", input_tokens=5) + assert _usage_rows(db, "s1") == [] + def test_usage_against_missing_session_is_safe_noop(self, db): + """Anti-resurrection contract (N1): recording usage for a session that does + not exist must NOT fail (FK safety) and must NOT create a session row + (a deleted session must stay deleted — token accounting must never + resurrect it).""" + db.record_auxiliary_usage("ghost", "vision", model="m", input_tokens=5) + assert _usage_rows(db, "ghost") == [] + assert db.get_session("ghost") is None class TestSchemaMigrationV22: @@ -176,6 +210,12 @@ def test_record_aux_usage_writes_through_context(self, db): assert rows[0]["input_tokens"] == 100 assert rows[0]["output_tokens"] == 20 + def test_noop_outside_context(self, db): + from agent.aux_accounting import record_aux_usage + + db.create_session("s1", source="cli") + record_aux_usage(_mk_response(), "vision") + assert _usage_rows(db, "s1") == [] def test_moa_tasks_excluded(self, db): """MoA advisor usage is already folded into the main-loop delta by @@ -195,7 +235,38 @@ def test_moa_tasks_excluded(self, db): reset_accounting_context(token) assert _usage_rows(db, "s1") == [] + def test_no_usage_object_is_noop(self, db): + from agent.aux_accounting import ( + record_aux_usage, + reset_accounting_context, + set_accounting_context, + ) + + db.create_session("s1", source="cli") + resp = SimpleNamespace(model="m", choices=[]) + token = set_accounting_context(db, "s1") + try: + record_aux_usage(resp, "vision") + finally: + reset_accounting_context(token) + assert _usage_rows(db, "s1") == [] + + def test_recording_failure_never_raises(self, db): + from agent.aux_accounting import ( + record_aux_usage, + reset_accounting_context, + set_accounting_context, + ) + + class ExplodingDB: + def record_auxiliary_usage(self, *a, **kw): + raise RuntimeError("disk full") + token = set_accounting_context(ExplodingDB(), "s1") + try: + record_aux_usage(_mk_response(), "vision") # must not raise + finally: + reset_accounting_context(token) def test_validate_llm_response_records(self, db): """The aux client's validation chokepoint feeds the recorder.""" @@ -217,6 +288,19 @@ def test_validate_llm_response_records(self, db): assert rows[0]["task"] == "web_extract" assert rows[0]["billing_provider"] == "openrouter" + def test_context_isolated_between_copied_contexts(self, db): + import contextvars + + from agent.aux_accounting import get_accounting_context, set_accounting_context + + def _set_and_get(sid): + set_accounting_context(db, sid) + return get_accounting_context()[1] + + a = contextvars.copy_context().run(_set_and_get, "agent-a") + b = contextvars.copy_context().run(_set_and_get, "agent-b") + assert (a, b) == ("agent-a", "agent-b") + assert get_accounting_context() is None class TestAnalyticsAuxRows: @@ -283,3 +367,25 @@ def test_overview_totals_include_aux_usage(self, db): models = {m["model"] for m in report["models"]} assert {"main-model", "glm-5"} <= models + def test_overview_totals_not_double_counted_with_absolute_updates(self, db): + """Gateway absolute overwrites + aux rows must not inflate totals.""" + from agent.insights import InsightsEngine + + db.create_session("s2", source="telegram") + db.update_token_counts( + "s2", input_tokens=2000, output_tokens=200, + model="main-model", billing_provider="nous", api_call_count=1, + ) + db.update_token_counts( + "s2", input_tokens=2000, output_tokens=200, + model="main-model", billing_provider="nous", + absolute=True, api_call_count=1, + ) + db.record_auxiliary_usage( + "s2", "title_generation", model="main-model", + billing_provider="nous", input_tokens=40, output_tokens=8, + ) + report = InsightsEngine(db).generate(days=30) + ov = report["overview"] + assert ov["total_input_tokens"] == 2040 + assert ov["total_output_tokens"] == 208 diff --git a/tests/test_gap3_round2_n1_n3_n5.py b/tests/test_gap3_round2_n1_n3_n5.py new file mode 100644 index 000000000000..33a859b470a2 --- /dev/null +++ b/tests/test_gap3_round2_n1_n3_n5.py @@ -0,0 +1,310 @@ +"""TDD round-2 tests for gap-3 fixes N1/N3/N5/N9 (hermes-state campaign). + +Hermetic tests: every ``SessionDB`` is built on ``tmp_path`` — never the +real ``~/.hermes`` state.db. + +* N1 — token/usage accounting must NEVER resurrect a deleted session row: + ``update_token_counts`` / ``_apply_token_batch`` on a deleted session must + leave ``sessions`` with COUNT(*) == 0 and must not raise. +* N1b — control: on a LIVE session the same calls still update counters. +* N3 — ``delete_session`` must also remove ``async_delegations`` rows that + belong to the deleted session (``origin_session``), while leaving rows of + other sessions untouched. +* N5 — ``delete_session`` must also remove ``delivery_obligations`` rows for + the deleted session (matched through the session's ``session_key``), while + leaving obligations of other session keys untouched. +* N9 — the api_server title-conflict create path must not leave an orphaned + ``gateway_routing`` entry behind when it rolls back the freshly inserted + session row. +""" +import json +import time + +import pytest + +from hermes_state import SessionDB + + +@pytest.fixture +def db(tmp_path): + return SessionDB(tmp_path / "state.db") + + +def _count_rows(db, table, where, params): + with db._lock: + row = db._conn.execute( + f"SELECT COUNT(*) FROM {table} WHERE {where}", params + ).fetchone() + return row[0] + + +# Real DDL mirrored from gateway/delivery_ledger.py::_initialize_schema — +# the delivery_obligations table is NOT part of the hermes_state schema, so +# tests create it with the production column set. +_DELIVERY_OBLIGATIONS_DDL = """ +CREATE TABLE IF NOT EXISTS delivery_obligations ( + obligation_id TEXT PRIMARY KEY, + session_key TEXT NOT NULL, + platform TEXT NOT NULL, + chat_id TEXT NOT NULL, + thread_id TEXT, + content TEXT NOT NULL, + state TEXT NOT NULL, + attempts INTEGER NOT NULL DEFAULT 0, + created_at REAL NOT NULL, + updated_at REAL NOT NULL, + owner_pid INTEGER, + owner_started_at INTEGER, + last_error TEXT +) +""" + + +def _insert_async_delegation( + db, + delegation_id, + origin_ui_session_id, + *, + parent_session_id=None, + origin_session="", + state="pending", +): + """Mirror tools/async_delegation.py's real row shape: the session id + lives in ``origin_ui_session_id`` / ``parent_session_id``, while + ``origin_session`` carries the routing key.""" + now = time.time() + with db._lock: + db._conn.execute( + """INSERT INTO async_delegations ( + delegation_id, origin_session, origin_ui_session_id, + parent_session_id, state, dispatched_at, completed_at, + updated_at, event_json, result_json, delivery_state, + delivery_attempts, delivered_at, owner_pid, owner_started_at, + task_json, delivery_claim, delivery_claimed_at + ) VALUES (?, ?, ?, ?, ?, ?, NULL, ?, NULL, NULL, + 'pending', 0, NULL, NULL, NULL, NULL, NULL, NULL)""", + ( + delegation_id, + origin_session, + origin_ui_session_id, + parent_session_id, + state, + now, + now, + ), + ) + + +def _insert_delivery_obligation( + db, obligation_id, session_key, *, state="pending" +): + now = time.time() + with db._lock: + db._conn.execute( + """INSERT INTO delivery_obligations ( + obligation_id, session_key, platform, chat_id, thread_id, + content, state, attempts, created_at, updated_at, + owner_pid, owner_started_at, last_error + ) VALUES (?, ?, 'telegram', 'chat-1', NULL, 'final reply', + ?, 0, ?, ?, NULL, NULL, NULL)""", + (obligation_id, session_key, state, now, now), + ) + + +# --------------------------------------------------------------------------- +# N1 — token/usage accounting must never resurrect a deleted session +# --------------------------------------------------------------------------- + +class TestN1TokenAccountingNeverResurrectsDeletedSession: + def test_update_token_counts_after_delete_leaves_no_row(self, db): + """A queued token delta arriving AFTER delete_session() must not + recreate the session row (and must not raise).""" + db.create_session("s1", source="cli") + db.append_message("s1", "user", content="hello") + assert db.delete_session("s1") is True + assert _count_rows(db, "sessions", "id = ?", ("s1",)) == 0 + + db.update_token_counts( + "s1", input_tokens=100, output_tokens=50, model="m1" + ) + + assert _count_rows(db, "sessions", "id = ?", ("s1",)) == 0 + + def test_apply_token_batch_after_delete_leaves_no_row(self, db): + """The async writer's _apply_token_batch path has the same contract: + deltas for a deleted session are dropped, never re-inserted.""" + db.create_session("s2", source="cli") + db.append_message("s2", "user", content="hello") + assert db.delete_session("s2") is True + + db._apply_token_batch( + [("s2", {"input_tokens": 10, "output_tokens": 5, "model": "m1"})] + ) + + assert _count_rows(db, "sessions", "id = ?", ("s2",)) == 0 + + def test_absolute_batch_after_delete_leaves_no_row(self, db): + """Gateway-style absolute (cumulative) deltas share the guard.""" + db.create_session("s3", source="gateway") + assert db.delete_session("s3") is True + + db.update_token_counts( + "s3", input_tokens=999, output_tokens=999, absolute=True + ) + + assert _count_rows(db, "sessions", "id = ?", ("s3",)) == 0 + + def test_live_session_token_update_still_works(self, db): + """Control (N1b): the guard must not break normal accounting on a + session that still exists.""" + db.create_session("live", source="cli") + db.update_token_counts( + "live", input_tokens=100, output_tokens=50, model="m1" + ) + sess = db.get_session("live") + assert sess is not None + assert sess["input_tokens"] == 100 + assert sess["output_tokens"] == 50 + + +# --------------------------------------------------------------------------- +# N3 — delete_session cascades into async_delegations +# --------------------------------------------------------------------------- + +class TestN3DeleteSessionRemovesAsyncDelegations: + def test_delegation_rows_for_deleted_session_are_removed(self, db): + db.create_session("s1", source="cli") + # Real row shapes: id in origin_ui_session_id + parent_session_id. + _insert_async_delegation( + db, "del-1", "s1", parent_session_id="s1", + origin_session="telegram:551199999999:default", + ) + _insert_async_delegation( + db, "del-2", "s1", parent_session_id="s1", + origin_session="telegram:551199999999:default", state="running", + ) + assert _count_rows( + db, "async_delegations", "origin_ui_session_id = ?", ("s1",) + ) == 2 + + assert db.delete_session("s1") is True + + assert _count_rows( + db, "async_delegations", "origin_ui_session_id = ?", ("s1",) + ) == 0 + assert _count_rows( + db, "async_delegations", "parent_session_id = ?", ("s1",) + ) == 0 + + def test_delegations_of_other_sessions_survive(self, db): + db.create_session("victim", source="cli") + db.create_session("other", source="cli") + _insert_async_delegation(db, "del-victim", "victim") + _insert_async_delegation(db, "del-other", "other") + + assert db.delete_session("victim") is True + + assert _count_rows( + db, "async_delegations", "delegation_id = ?", ("del-victim",) + ) == 0 + assert _count_rows( + db, "async_delegations", "delegation_id = ?", ("del-other",) + ) == 1 + + +# --------------------------------------------------------------------------- +# N5 — delete_session cascades into delivery_obligations +# --------------------------------------------------------------------------- + +class TestN5DeleteSessionRemovesDeliveryObligations: + def test_obligations_for_deleted_session_key_are_removed(self, db): + with db._lock: + db._conn.execute(_DELIVERY_OBLIGATIONS_DDL) + session_key = "telegram:551199999999:default" + db.create_session("s1", source="gateway", session_key=session_key) + _insert_delivery_obligation(db, "obl-1", session_key, state="pending") + _insert_delivery_obligation(db, "obl-2", session_key, state="delivered") + assert _count_rows( + db, "delivery_obligations", "session_key = ?", (session_key,) + ) == 2 + + assert db.delete_session("s1") is True + + assert _count_rows( + db, "delivery_obligations", "session_key = ?", (session_key,) + ) == 0 + + def test_obligations_of_other_session_keys_survive(self, db): + with db._lock: + db._conn.execute(_DELIVERY_OBLIGATIONS_DDL) + victim_key = "telegram:551199999999:default" + other_key = "telegram:558899999999:default" + db.create_session("victim", source="gateway", session_key=victim_key) + db.create_session("other", source="gateway", session_key=other_key) + _insert_delivery_obligation(db, "obl-victim", victim_key) + _insert_delivery_obligation(db, "obl-other", other_key) + + assert db.delete_session("victim") is True + + assert _count_rows( + db, "delivery_obligations", "obligation_id = ?", ("obl-victim",) + ) == 0 + assert _count_rows( + db, "delivery_obligations", "obligation_id = ?", ("obl-other",) + ) == 1 + + +# --------------------------------------------------------------------------- +# N9 — api_server title-conflict create leaves no orphaned routing +# --------------------------------------------------------------------------- + +class TestN9TitleConflictCreateLeavesNoOrphanRouting: + """POST /api/sessions with a title already owned by another session must + roll back the freshly inserted row AND purge any stale gateway_routing + entry still pointing at the reused session id.""" + + @pytest.mark.asyncio + async def test_title_conflict_purges_stale_routing(self, tmp_path, monkeypatch): + from gateway.config import PlatformConfig + from gateway.platforms.api_server import APIServerAdapter + + db = SessionDB(tmp_path / "state.db") + # Another session already owns the title the client wants to use. + db.create_session("existing", source="api_server") + db.set_session_title("existing", "Shared Title") + # Stale routing entry that still points at the id the client reuses. + db.save_gateway_routing_entry( + "telegram:551199999999:default", + json.dumps({"session_id": "reused-id", "source": "api_server"}), + scope="", + ) + assert _count_rows( + db, "gateway_routing", "session_key = ?", + ("telegram:551199999999:default",), + ) == 1 + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_check_auth", lambda request: None) + + async def _session_db(): + return db + + monkeypatch.setattr(adapter, "_ensure_session_db_async", _session_db) + + async def _read_body(request): + return {"id": "reused-id", "title": "Shared Title"}, None + + monkeypatch.setattr(adapter, "_read_json_body", _read_body) + + resp = await adapter._handle_create_session(object()) + + assert resp.status == 400 + # The rolled-back row must not exist... + assert _count_rows(db, "sessions", "id = ?", ("reused-id",)) == 0 + # ...and the stale routing entry must not be orphaned behind it. + assert _count_rows( + db, "gateway_routing", "session_key = ?", + ("telegram:551199999999:default",), + ) == 0 + # The rightful title owner is untouched. + assert _count_rows(db, "sessions", "id = ?", ("existing",)) == 1 diff --git a/tests/test_gateway_backstop.py b/tests/test_gateway_backstop.py new file mode 100644 index 000000000000..0bcc979895c8 --- /dev/null +++ b/tests/test_gateway_backstop.py @@ -0,0 +1,238 @@ +"""Multi-process backstop tests for the gateway routing index (GAP-3). + +When several hermes processes share one ``state.db`` (gateway + desktop + +CLI + MCP workers), any of them can delete a session row through +``SessionDB.delete_session`` while the gateway still holds the matching +entry in its in-memory routing index (``SessionStore._entries``). A naive +``_save_entries()`` then re-persists that ghost entry into the +``gateway_routing`` table and the sessions.json mirror, resurrecting a +route to a session that no longer exists. + +Contract under test (backstop implemented in ``gateway/session.py``, in +``_save_entries`` / ``_persist_routing_data``): + +* Entries with ``db_persisted=True`` whose row no longer exists in + ``state.db`` (deleted by ANOTHER process) are DROPPED before persistence + — the implementation checks existence with a batch query + (``SELECT id FROM sessions WHERE id IN (...)``). +* Legacy entries with ``db_persisted=False`` (pre-SQLite routing entries + that never had a state.db row) are ALWAYS preserved — absence from the + DB is expected for them, not evidence of deletion. +* A DB failure during the existence check must not break the save + (try/except, fail-safe: entries are preserved, nothing is dropped). +""" + +import json +from datetime import datetime, timezone +from pathlib import Path + +import pytest + +from gateway.config import GatewayConfig, Platform +from gateway.session import SessionEntry, SessionSource, SessionStore +from hermes_state import SessionDB + + +@pytest.fixture() +def _isolated_db(tmp_path, monkeypatch): + """Point state.db at tmp_path so tests never touch ~/.hermes.""" + import hermes_state + + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + return tmp_path + + +def _make_store(tmp_path): + return SessionStore(sessions_dir=tmp_path / "sessions", config=GatewayConfig()) + + +def _slack_source(chat_id): + return SessionSource( + platform=Platform.SLACK, + chat_id=chat_id, + chat_type="channel", + user_id="U1", + ) + + +def _persisted_routing(tmp_path, store): + """Read the routing index as a *fresh* process would (new SessionDB).""" + db = SessionDB(db_path=tmp_path / "state.db") + return db.load_gateway_routing_entries(scope=store._routing_scope()) + + +def _read_sessions_json(tmp_path): + path = Path(tmp_path) / "sessions" / "sessions.json" + assert path.exists(), "sessions.json mirror should have been written" + return json.loads(path.read_text(encoding="utf-8")) + + +# --------------------------------------------------------------------------- +# (a) db_persisted=True + row deleted by another process -> dropped +# --------------------------------------------------------------------------- + +def test_persisted_entry_with_row_deleted_by_other_process_is_dropped( + _isolated_db, tmp_path +): + """A persisted routing entry whose state.db row was deleted by another + process must NOT be re-persisted on the next _save_entries().""" + store = _make_store(tmp_path) + entry = store.get_or_create_session(_slack_source("C111")) + assert entry.db_persisted is True, "freshly created session must be DB-persisted" + + # Simulate the other process (desktop/CLI/MCP) deleting the session. + assert store._db.delete_session(entry.session_id) is True + + # Gateway's next whole-index save must not resurrect the ghost route. + store._save_entries() + + persisted = _persisted_routing(tmp_path, store) + assert entry.session_key not in persisted, ( + "ghost routing entry for a session deleted by another process must be " + "dropped before persistence (backstop)" + ) + sessions_json = _read_sessions_json(tmp_path) + assert entry.session_key not in sessions_json, ( + "sessions.json mirror must not contain the ghost routing entry either" + ) + + +# --------------------------------------------------------------------------- +# (b) db_persisted=True + row EXISTS -> preserved +# --------------------------------------------------------------------------- + +def test_persisted_entry_with_existing_row_is_preserved(_isolated_db, tmp_path): + """A persisted entry whose row still exists in state.db is untouched.""" + store = _make_store(tmp_path) + entry = store.get_or_create_session(_slack_source("C222")) + assert entry.db_persisted is True + + store._save_entries() + + persisted = _persisted_routing(tmp_path, store) + assert entry.session_key in persisted, "live persisted entry must be preserved" + assert json.loads(persisted[entry.session_key])["session_id"] == entry.session_id + assert entry.session_key in _read_sessions_json(tmp_path) + + +# --------------------------------------------------------------------------- +# (c) legacy db_persisted=False + row absent -> PRESERVED (legacy compat) +# --------------------------------------------------------------------------- + +def test_legacy_entry_never_persisted_is_preserved_when_row_absent( + _isolated_db, tmp_path +): + """Pre-SQLite legacy entries (db_persisted=False) are always preserved, + even when they have no row in state.db — absence is expected for them.""" + store = _make_store(tmp_path) + key = "agent:main:slack:channel:C333" + now = datetime.now(timezone.utc) + legacy = SessionEntry( + session_key=key, + session_id="legacy-session-no-db-row", + created_at=now, + updated_at=now, + origin=_slack_source("C333"), + db_persisted=False, + ) + store._entries[key] = legacy + store._loaded = True + + # Sanity: no such row exists in state.db. + assert store._db.get_session(legacy.session_id) is None + + store._save_entries() + + persisted = _persisted_routing(tmp_path, store) + assert key in persisted, ( + "legacy entry (db_persisted=False) must be preserved even with no DB row" + ) + assert json.loads(persisted[key])["session_id"] == legacy.session_id + assert key in _read_sessions_json(tmp_path) + + +# --------------------------------------------------------------------------- +# (d) 3 entries, 1 dead -> only the dead one falls +# --------------------------------------------------------------------------- + +def test_only_dead_entry_dropped_among_mixed_entries(_isolated_db, tmp_path): + """With several persisted entries, only the one whose row was deleted by + another process is dropped; the live ones survive the save.""" + store = _make_store(tmp_path) + e1 = store.get_or_create_session(_slack_source("C441")) + e2 = store.get_or_create_session(_slack_source("C442")) + e3 = store.get_or_create_session(_slack_source("C443")) + for e in (e1, e2, e3): + assert e.db_persisted is True + + assert store._db.delete_session(e2.session_id) is True + + store._save_entries() + + persisted = _persisted_routing(tmp_path, store) + assert e1.session_key in persisted + assert e3.session_key in persisted + assert e2.session_key not in persisted, ( + "only the entry whose row was deleted may be dropped" + ) + + sessions_json = _read_sessions_json(tmp_path) + assert e1.session_key in sessions_json + assert e3.session_key in sessions_json + assert e2.session_key not in sessions_json + + +# --------------------------------------------------------------------------- +# (e) DB exception during the existence check -> save does not break +# --------------------------------------------------------------------------- + +def test_db_error_during_existence_check_does_not_break_save( + _isolated_db, tmp_path, monkeypatch +): + """A DB exception raised by the existence check must be caught: the save + must not raise, and the entries must be preserved (fail-safe — nothing + is dropped on uncertainty).""" + store = _make_store(tmp_path) + entry = store.get_or_create_session(_slack_source("C555")) + assert entry.db_persisted is True + + # Force the backstop's DB read to fail mid-check (monkeypatch do db). + check_attempted = [] + real_db = store._db + real_read_ctx = getattr(real_db, "_read_ctx", None) + + def _boom_read_ctx(*args, **kwargs): + check_attempted.append("read_ctx") + raise RuntimeError("simulated DB failure during existence check") + + def _boom_query(*args, **kwargs): + check_attempted.append("query") + raise RuntimeError("simulated DB failure during existence check") + + if real_read_ctx is not None: + monkeypatch.setattr(real_db, "_read_ctx", _boom_read_ctx) + else: + # Alternative wiring: the batch query itself is the check. + monkeypatch.setattr(store, "_query_existing_session_ids", _boom_query) + + # Must not raise, and the check must actually have been attempted (and + # failed) — without this, a save that skips the check would false-pass. + store._save_entries() + assert check_attempted, "backstop existence check must have been exercised" + + # Fail-safe: nothing dropped on DB error; entries survive in memory. + assert entry.session_key in store._entries + data, _generation = store._snapshot_routing_locked() + assert entry.session_key in data, "entry must survive a failed existence check" + + # Once the DB is healthy again, a normal save persists the entry intact. + if real_read_ctx is not None: + monkeypatch.setattr(real_db, "_read_ctx", real_read_ctx) + else: + monkeypatch.undo() + store._save_entries() + persisted = _persisted_routing(tmp_path, store) + assert entry.session_key in persisted, ( + "entry preserved during the DB failure must be persistable afterwards" + ) diff --git a/tests/test_gateway_forget_sessions.py b/tests/test_gateway_forget_sessions.py new file mode 100644 index 000000000000..714d2821f01c --- /dev/null +++ b/tests/test_gateway_forget_sessions.py @@ -0,0 +1,242 @@ +"""Tests for SessionStore.forget_sessions (GAP-3, TDD). + +Contract under test (implemented in parallel in gateway/session.py): + + SessionStore.forget_sessions(session_ids: List[str]) -> int + +* removes from ``_entries`` every entry whose ``entry.session_id`` is in + ``session_ids``; +* re-persists the routing index (state.db ``gateway_routing`` table + the + legacy sessions.json mirror) WITHOUT the dead entries; +* returns the number of entries removed; +* 0 removed => no save (durable files untouched). + +Isolation: every test gets its own ``state.db`` (``DEFAULT_DB_PATH`` is +module-level and shared by every ``SessionDB()`` in the process), and its own +``sessions_dir`` under ``tmp_path``. Never touches ``~/.hermes``. +""" + +import json +import threading + +import hermes_state +import pytest + +from gateway.config import GatewayConfig, Platform, SessionResetPolicy +from gateway.session import SessionSource, SessionStore + + +@pytest.fixture(autouse=True) +def _isolated_db(tmp_path, monkeypatch): + """Each test gets its own state.db — rows would otherwise leak between + tests because DEFAULT_DB_PATH is module-level.""" + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + + +def _make_store(tmp_path) -> SessionStore: + config = GatewayConfig(default_reset_policy=SessionResetPolicy(mode="none")) + return SessionStore(sessions_dir=tmp_path / "sessions", config=config) + + +def _source(chat_id: str) -> SessionSource: + return SessionSource( + platform=Platform.TELEGRAM, + chat_id=chat_id, + chat_name=f"chat {chat_id}", + chat_type="dm", + user_id=chat_id, + ) + + +def _routing_rows(store) -> dict: + """Read back the gateway_routing table for this store's scope.""" + return store._db.load_gateway_routing_entries(scope=store._routing_scope()) + + +def _sessions_json(store) -> dict: + return json.loads((store.sessions_dir / "sessions.json").read_text(encoding="utf-8")) + + +class TestForgetSessions: + def test_forget_removes_entry_and_repersists_both_layers(self, tmp_path): + """(a) A real persisted entry is dropped from memory AND from both + durable layers (gateway_routing table + sessions.json).""" + store = _make_store(tmp_path) + try: + entry = store.get_or_create_session(_source("chat-1")) + key, sid = entry.session_key, entry.session_id + assert key in _routing_rows(store) + assert key in _sessions_json(store) + + removed = store.forget_sessions([sid]) + + assert removed == 1 + assert key not in store._entries + assert key not in _routing_rows(store) + data = _sessions_json(store) + assert key not in data + # The legacy-mirror sentinel survives the rewrite. + assert "_README" in data + finally: + store._db.close() + + def test_forget_unknown_id_returns_zero_and_does_not_rewrite( + self, tmp_path, monkeypatch + ): + """(b) Unknown id -> returns 0 and NO persistence happens at all: + no _persist_routing_data call, sessions.json bytes+mtime untouched, + routing table rows unchanged.""" + store = _make_store(tmp_path) + try: + entry = store.get_or_create_session(_source("chat-1")) + sessions_file = store.sessions_dir / "sessions.json" + json_before = sessions_file.read_bytes() + mtime_before = sessions_file.stat().st_mtime_ns + rows_before = _routing_rows(store) + + save_calls = [] + monkeypatch.setattr( + store, + "_persist_routing_data", + lambda data, generation: save_calls.append(generation), + ) + + removed = store.forget_sessions(["20990101_000000_no_such_session"]) + + assert removed == 0 + assert save_calls == [], "0 removals must not trigger a save" + assert sessions_file.read_bytes() == json_before + assert sessions_file.stat().st_mtime_ns == mtime_before + assert _routing_rows(store) == rows_before + assert store._entries[entry.session_key].session_id == entry.session_id + finally: + store._db.close() + + def test_forget_removes_only_target_among_multiple(self, tmp_path): + """(c) Multiple entries, only one targeted -> only it falls, in memory + and in both durable layers.""" + store = _make_store(tmp_path) + try: + entries = {} + for chat in ("chat-1", "chat-2", "chat-3"): + entry = store.get_or_create_session(_source(chat)) + entries[entry.session_key] = entry + keys = list(entries) + target_key = keys[1] + + removed = store.forget_sessions([entries[target_key].session_id]) + + assert removed == 1 + assert target_key not in store._entries + rows = _routing_rows(store) + data = _sessions_json(store) + for key, entry in entries.items(): + if key == target_key: + assert key not in rows + assert key not in data + else: + assert key in rows # survives in DB routing table + assert data[key]["session_id"] == entry.session_id # and JSON mirror + finally: + store._db.close() + + def test_forget_concurrent_threads_consistent(self, tmp_path): + """(d) Two threads calling forget_sessions concurrently on the same + store with disjoint targets -> no exception, each returns its own + count, both targets dropped from memory and both durable layers, + survivor intact.""" + store = _make_store(tmp_path) + try: + entries = {} + for chat in ("chat-1", "chat-2", "chat-3"): + entry = store.get_or_create_session(_source(chat)) + entries[entry.session_key] = entry + keys = list(entries) + targets = [entries[keys[0]].session_id, entries[keys[1]].session_id] + survivor_key = keys[2] + + barrier = threading.Barrier(2) + results = {} + results_lock = threading.Lock() + + def worker(i): + barrier.wait() # maximize contention + try: + n = store.forget_sessions([targets[i]]) + with results_lock: + results[i] = n + except Exception as exc: # pragma: no cover - failure path + with results_lock: + results[i] = f"ERR:{exc}" + + threads = [threading.Thread(target=worker, args=(i,)) for i in range(2)] + for t in threads: + t.start() + for t in threads: + t.join(timeout=30) + + assert results == {0: 1, 1: 1}, results + assert keys[0] not in store._entries + assert keys[1] not in store._entries + assert survivor_key in store._entries + rows = _routing_rows(store) + assert keys[0] not in rows + assert keys[1] not in rows + assert survivor_key in rows + data = _sessions_json(store) + assert keys[0] not in data + assert keys[1] not in data + assert survivor_key in data + finally: + store._db.close() + + def test_forget_works_on_not_yet_loaded_store(self, tmp_path): + """(e) Lazy store: forget_sessions on a store that has never loaded + must load the index first, then remove and persist.""" + store = _make_store(tmp_path) + entry = store.get_or_create_session(_source("chat-1")) + key, sid = entry.session_key, entry.session_id + store._db.close() + + restarted = _make_store(tmp_path) + try: + assert restarted._loaded is False + + removed = restarted.forget_sessions([sid]) + + assert removed == 1 + assert key not in restarted._entries + assert key not in _routing_rows(restarted) + assert key not in _sessions_json(restarted) + finally: + restarted._db.close() + + def test_forget_returns_removed_count_and_ignores_unknown_ids(self, tmp_path): + """(f) Return value counts only real removals: mixed list with unknown + ids counts only the matched ones; a second forget of an already-removed + id returns 0.""" + store = _make_store(tmp_path) + try: + entries = {} + for chat in ("chat-1", "chat-2", "chat-3"): + entry = store.get_or_create_session(_source(chat)) + entries[entry.session_key] = entry + keys = list(entries) + + mixed = store.forget_sessions( + [ + entries[keys[0]].session_id, + "20990101_000000_ghost", + entries[keys[1]].session_id, + ] + ) + assert mixed == 2 + + assert store.forget_sessions([entries[keys[2]].session_id]) == 1 + # Already removed -> 0, and the index is empty everywhere. + assert store.forget_sessions([entries[keys[2]].session_id]) == 0 + assert store._entries == {} + assert _routing_rows(store) == {} + assert set(_sessions_json(store)) == {"_README"} + finally: + store._db.close()