From 97add7b0aabd36d7123980671e65b5b0347cf107 Mon Sep 17 00:00:00 2001 From: Tosko4 <1294707+Tosko4@users.noreply.github.com> Date: Sun, 28 Jun 2026 23:43:58 +0200 Subject: [PATCH] fix: preserve Discord lane metadata in LCM --- __init__.py | 39 ++++++++-- db_bootstrap.py | 23 +++++- engine.py | 1 + schemas.py | 7 ++ store.py | 65 +++++++++++++--- tests/test_lcm_core.py | 63 ++++++++++++++-- tests/test_lcm_engine.py | 72 ++++++++++++++++++ tests/test_packaging_install.py | 128 ++++++++++++++++++++++++++++++++ tools.py | 12 ++- 9 files changed, 387 insertions(+), 23 deletions(-) diff --git a/__init__.py b/__init__.py index d6f7b72dc..43ccc4bd4 100644 --- a/__init__.py +++ b/__init__.py @@ -156,11 +156,40 @@ def register(ctx): def _on_post_llm_call(**kwargs): history = kwargs.get("conversation_history") - if history: - try: - engine.ingest(history) - except Exception as exc: - logger.debug("LCM post_llm_call ingest error: %s", exc) + if not history: + return + active_engine = kwargs.get("context_compressor") + if not ( + active_engine is not None + and getattr(active_engine, "name", None) == "lcm" + and hasattr(active_engine, "ingest") + ): + active_engine = engine + + session_id = str(kwargs.get("session_id") or "") + conversation_id = str( + kwargs.get("conversation_id") + or kwargs.get("gateway_session_key") + or "" + ) + platform = str(kwargs.get("platform") or "") + + try: + if session_id and ( + str(getattr(active_engine, "current_session_id", "") or "") != session_id + or ( + conversation_id + and str(getattr(active_engine, "current_conversation_id", "") or "") != conversation_id + ) + ): + active_engine.on_session_start( + session_id, + platform=platform, + conversation_id=conversation_id or None, + ) + active_engine.ingest(history) + except Exception as exc: + logger.debug("LCM post_llm_call ingest error: %s", exc) _mgr._hooks.setdefault("post_llm_call", []).append(_on_post_llm_call) logger.debug("LCM registered post_llm_call hook for per-turn ingest") diff --git a/db_bootstrap.py b/db_bootstrap.py index 7cd337f8e..f3f9596b6 100644 --- a/db_bootstrap.py +++ b/db_bootstrap.py @@ -17,7 +17,7 @@ logger = logging.getLogger(__name__) -SCHEMA_VERSION = 4 +SCHEMA_VERSION = 5 SQLITE_BUSY_TIMEOUT_MS = 30_000 _MIN_DISK_SPACE_BYTES = 50 * 1024 * 1024 REQUIRED_CORE_TABLES = ( @@ -175,6 +175,22 @@ def ensure_lifecycle_state_columns(conn: sqlite3.Connection) -> None: conn.execute("ALTER TABLE lcm_lifecycle_state ADD COLUMN last_reset_at REAL") +def ensure_message_origin_columns(conn: sqlite3.Connection) -> None: + table_row = conn.execute( + "SELECT name FROM sqlite_master WHERE type='table' AND name='messages'" + ).fetchone() + if not table_row: + return + columns = { + row[1] for row in conn.execute("PRAGMA table_info(messages)").fetchall() + } + if "conversation_id" not in columns: + conn.execute("ALTER TABLE messages ADD COLUMN conversation_id TEXT DEFAULT ''") + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_msg_conversation_session ON messages(conversation_id, session_id, store_id)" + ) + + def mark_migration_step_complete(conn: sqlite3.Connection, step_name: str) -> None: ensure_migration_state_table(conn) conn.execute( @@ -598,4 +614,9 @@ def run_versioned_migrations(conn: sqlite3.Connection) -> None: mark_migration_step_complete(conn, "v4_lifecycle_debt_columns") current_version = 4 + ensure_message_origin_columns(conn) + if current_version < 5: + mark_migration_step_complete(conn, "v5_message_conversation_id") + current_version = 5 + set_schema_version(conn, current_version) diff --git a/engine.py b/engine.py index dbce49796..f5125cf3a 100644 --- a/engine.py +++ b/engine.py @@ -3585,6 +3585,7 @@ def _ingest_messages(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any protected_messages, estimates, source=self._session_platform, + conversation_id=self._conversation_id, ) self._ingest_cursor = n logger.debug("Ingested %d messages into LCM store", len(messages_to_store_with_index)) diff --git a/schemas.py b/schemas.py index fe4b3f5fb..214052ed0 100644 --- a/schemas.py +++ b/schemas.py @@ -67,6 +67,13 @@ "Use 'unknown' for explicit unknown-source content." ), }, + "conversation_id": { + "type": "string", + "description": ( + "Optional gateway conversation/session key filter for lane-scoped retrieval. " + "Use this to restrict Discord searches to one channel/thread/forum topic lane when rows carry metadata." + ), + }, "role": { "type": "string", "enum": ["system", "user", "assistant", "tool", "unknown"], diff --git a/store.py b/store.py index 7f7538c42..6b47fdb2a 100644 --- a/store.py +++ b/store.py @@ -53,7 +53,7 @@ _MESSAGE_ROLE_BIAS_SQL = "CASE m.role WHEN 'user' THEN 0 WHEN 'assistant' THEN 1 WHEN 'tool' THEN 2 ELSE 1 END" _MESSAGE_SELECT_COLUMNS = ( "store_id, session_id, source, role, content, tool_call_id, " - "tool_calls, tool_name, timestamp, token_estimate, pinned" + "tool_calls, tool_name, timestamp, token_estimate, pinned, conversation_id" ) _UNKNOWN_SOURCE = "unknown" @@ -71,6 +71,10 @@ def _normalize_source_value(source: str | None) -> str: return normalized or _UNKNOWN_SOURCE +def _normalize_conversation_id_value(conversation_id: str | None) -> str: + return (conversation_id or "").strip() + + def _source_filter_clause(column: str, source: str | None) -> tuple[str | None, list[str]]: normalized = _normalize_source_value(source) if source is not None else "" if not normalized: @@ -80,6 +84,13 @@ def _source_filter_clause(column: str, source: str | None) -> tuple[str | None, return f"{column} = ?", [normalized] +def _conversation_filter_clause(column: str, conversation_id: str | None) -> tuple[str | None, list[str]]: + normalized = _normalize_conversation_id_value(conversation_id) + if not normalized: + return None, [] + return f"{column} = ?", [normalized] + + def _message_role_bias(role: str | None) -> float: if role == "user": return 0.0 @@ -240,6 +251,7 @@ def _init_db(self): store_id INTEGER PRIMARY KEY AUTOINCREMENT, session_id TEXT NOT NULL, source TEXT DEFAULT '', + conversation_id TEXT DEFAULT '', role TEXT NOT NULL, content TEXT, tool_call_id TEXT, @@ -265,6 +277,7 @@ def _init_db(self): ) run_versioned_migrations(self._conn) self._ensure_source_column() + self._ensure_conversation_id_column() self._conn.commit() def _ensure_source_column(self) -> None: @@ -277,10 +290,21 @@ def _ensure_source_column(self) -> None: "CREATE INDEX IF NOT EXISTS idx_msg_source_session ON messages(source, session_id, store_id)" ) + def _ensure_conversation_id_column(self) -> None: + columns = { + row[1] for row in self._conn.execute("PRAGMA table_info(messages)").fetchall() + } + if "conversation_id" not in columns: + self._conn.execute("ALTER TABLE messages ADD COLUMN conversation_id TEXT DEFAULT ''") + self._conn.execute( + "CREATE INDEX IF NOT EXISTS idx_msg_conversation_session ON messages(conversation_id, session_id, store_id)" + ) + # -- Write operations --------------------------------------------------- def append(self, session_id: str, msg: Dict[str, Any], - token_estimate: int = 0, source: str = "") -> int: + token_estimate: int = 0, source: str = "", + conversation_id: str = "") -> int: """Persist a message and return its store_id.""" msg = protect_message_for_ingest( msg, @@ -294,12 +318,13 @@ def append(self, session_id: str, msg: Dict[str, Any], with self._write_lock: cur = self._conn.execute( """INSERT INTO messages - (session_id, source, role, content, tool_call_id, tool_calls, + (session_id, source, conversation_id, role, content, tool_call_id, tool_calls, tool_name, timestamp, token_estimate, pinned) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ( session_id, _normalize_source_value(source), + _normalize_conversation_id_value(conversation_id), msg.get("role", "unknown"), _normalize_content_value(msg.get("content")), msg.get("tool_call_id"), @@ -316,7 +341,8 @@ def append(self, session_id: str, msg: Dict[str, Any], def append_batch(self, session_id: str, messages: List[Dict[str, Any]], token_estimates: List[int] | None = None, - source: str = "") -> List[int]: + source: str = "", + conversation_id: str = "") -> List[int]: """Persist multiple messages in one transaction. Returns store_ids.""" protected_messages = protect_messages_for_ingest( messages, @@ -329,12 +355,14 @@ def append_batch(self, session_id: str, protected_messages, token_estimates, source=source, + conversation_id=conversation_id, ) def _append_protected_batch(self, session_id: str, messages: List[Dict[str, Any]], token_estimates: List[int] | None = None, - source: str = "") -> List[int]: + source: str = "", + conversation_id: str = "") -> List[int]: """Persist messages that already passed ingest protection. This is an internal fast path for callers that need the protected form @@ -353,12 +381,13 @@ def _append_protected_batch(self, session_id: str, ts = time.time() cur = self._conn.execute( """INSERT INTO messages - (session_id, source, role, content, tool_call_id, tool_calls, + (session_id, source, conversation_id, role, content, tool_call_id, tool_calls, tool_name, timestamp, token_estimate, pinned) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ( session_id, _normalize_source_value(source), + _normalize_conversation_id_value(conversation_id), msg.get("role", "unknown"), _normalize_content_value(msg.get("content")), msg.get("tool_call_id"), @@ -700,6 +729,7 @@ def get_time_bounds(self, store_ids: List[int]) -> tuple[float | None, float | N def search(self, query: str, session_id: str | None = None, limit: int = 20, sort: str | None = None, source: str | None = None, + conversation_id: str | None = None, role: str | None = None, time_from: float | None = None, time_to: float | None = None) -> List[Dict[str, Any]]: @@ -712,6 +742,7 @@ def search(self, query: str, session_id: str | None = None, - ``source`` limits which raw rows inside those sessions are eligible - ``source='unknown'`` means the explicit unknown-source bucket, with legacy blank-source rows treated as equivalent for back-compat + - ``conversation_id`` limits rows to one gateway conversation/session key """ safe_query = sanitize_fts5_query(query) terms = extract_search_terms(safe_query) @@ -723,6 +754,7 @@ def search(self, query: str, session_id: str | None = None, limit=limit, sort=sort, source=source, + conversation_id=conversation_id, role=role, time_from=time_from, time_to=time_to, @@ -738,6 +770,7 @@ def search(self, query: str, session_id: str | None = None, apply_directness_adjustment = should_apply_directness_rank_adjustment(terms, phrases) max_rank_bonus = compute_directness_rank_bonus_upper_bound(terms, phrases) * 3e-7 source_clause, source_args = _source_filter_clause("m.source", source) + conversation_clause, conversation_args = _conversation_filter_clause("m.conversation_id", conversation_id) offset = 0 scanned_rows = 0 results: list[Dict[str, Any]] = [] @@ -751,6 +784,9 @@ def search(self, query: str, session_id: str | None = None, if source_clause: where.append(source_clause) args.extend(source_args) + if conversation_clause: + where.append(conversation_clause) + args.extend(conversation_args) if role is not None: where.append("m.role = ?") args.append(role) @@ -763,7 +799,7 @@ def search(self, query: str, session_id: str | None = None, args.extend([fetch_limit, offset]) rows = self._conn.execute( f"""SELECT m.store_id, m.session_id, m.source, m.role, m.content, m.tool_call_id, - m.tool_calls, m.tool_name, m.timestamp, m.token_estimate, m.pinned, + m.tool_calls, m.tool_name, m.timestamp, m.token_estimate, m.pinned, m.conversation_id, rank as search_rank, snippet(messages_fts, 0, '>>>', '<<<', '...', 40) as snippet FROM messages_fts fts @@ -781,6 +817,7 @@ def search(self, query: str, session_id: str | None = None, limit=limit, sort=sort, source=source, + conversation_id=conversation_id, role=role, time_from=time_from, time_to=time_to, @@ -789,7 +826,7 @@ def search(self, query: str, session_id: str | None = None, raw_primary_values: list[float] = [] for r in rows: d = self._row_to_dict(r) - base_columns = 11 + base_columns = 12 d["search_rank"] = r[base_columns] if len(r) > base_columns else None d["snippet"] = r[base_columns + 1] if len(r) > (base_columns + 1) else "" d["_directness_score"] = _message_directness_score(d.get("role"), d.get("content"), terms, phrases) @@ -821,6 +858,7 @@ def search(self, query: str, session_id: str | None = None, def _search_like(self, query: str, session_id: str | None = None, limit: int = 20, sort: str | None = None, source: str | None = None, + conversation_id: str | None = None, role: str | None = None, time_from: float | None = None, time_to: float | None = None) -> List[Dict[str, Any]]: @@ -840,6 +878,10 @@ def _search_like(self, query: str, session_id: str | None = None, if source_clause: where.append(source_clause) args.extend(source_args) + conversation_clause, conversation_args = _conversation_filter_clause("conversation_id", conversation_id) + if conversation_clause: + where.append(conversation_clause) + args.extend(conversation_args) if role is not None: where.append("role = ?") args.append(role) @@ -1017,10 +1059,11 @@ def _row_to_dict(self, row) -> Dict[str, Any]: return {} cols = [ "store_id", "session_id", "source", "role", "content", "tool_call_id", - "tool_calls", "tool_name", "timestamp", "token_estimate", "pinned", + "tool_calls", "tool_name", "timestamp", "token_estimate", "pinned", "conversation_id", ] d = dict(zip(cols, row[:len(cols)])) d["source"] = _normalize_source_value(d.get("source")) + d["conversation_id"] = _normalize_conversation_id_value(d.get("conversation_id")) # Deserialize tool_calls JSON if d.get("tool_calls"): try: diff --git a/tests/test_lcm_core.py b/tests/test_lcm_core.py index 4befde5bc..e8e9f65e1 100644 --- a/tests/test_lcm_core.py +++ b/tests/test_lcm_core.py @@ -1128,6 +1128,59 @@ def test_source_stored_and_filterable(self, store): assert discord_results[0]["source"] == "discord" assert discord_results[0]["session_id"] == "sess2" + def test_conversation_id_stored_and_filterable_for_discord_lanes(self, store): + main_id = store.append( + "sess-main", + {"role": "user", "content": "docker in discord main lane"}, + source="discord", + conversation_id="agent:main:discord:group:main:user", + ) + thread_id = store.append( + "sess-thread", + {"role": "user", "content": "docker in discord forum topic"}, + source="discord", + conversation_id="agent:main:discord:thread:topic:topic", + ) + + main_results = store.search( + "docker", + source="discord", + conversation_id="agent:main:discord:group:main:user", + ) + thread_results = store.search( + "docker", + source="discord", + conversation_id="agent:main:discord:thread:topic:topic", + ) + + assert [result["store_id"] for result in main_results] == [main_id] + assert main_results[0]["conversation_id"] == "agent:main:discord:group:main:user" + assert [result["store_id"] for result in thread_results] == [thread_id] + assert thread_results[0]["conversation_id"] == "agent:main:discord:thread:topic:topic" + assert store.get(main_id)["conversation_id"] == "agent:main:discord:group:main:user" + + def test_like_fallback_filters_by_conversation_id(self, store): + store.append( + "sess-main", + {"role": "user", "content": "foo bar lane main"}, + source="discord", + conversation_id="agent:main:discord:group:main:user", + ) + thread_id = store.append( + "sess-thread", + {"role": "user", "content": "foo bar lane topic"}, + source="discord", + conversation_id="agent:main:discord:thread:topic:topic", + ) + + results = store.search( + 'foo"bar', + source="discord", + conversation_id="agent:main:discord:thread:topic:topic", + ) + + assert [result["store_id"] for result in results] == [thread_id] + def test_missing_source_is_normalized_to_unknown_and_filterable(self, store): store_id = store.append("sess-unknown", {"role": "user", "content": "docker with unknown source"}) @@ -1337,7 +1390,7 @@ def test_init_repairs_malformed_message_fts_and_sets_schema_version(self, tmp_pa version = store._conn.execute( "SELECT value FROM metadata WHERE key = 'schema_version'" ).fetchone() - assert version == ("4",) + assert version == ("5",) results = store.search("docker", session_id="sess1") assert len(results) == 1 @@ -1386,7 +1439,7 @@ def test_init_recreates_missing_message_fts_trigger(self, tmp_path): version = store._conn.execute( "SELECT value FROM metadata WHERE key = 'schema_version'" ).fetchone() - assert version == ("4",) + assert version == ("5",) migration_state = store._conn.execute( "SELECT step_name FROM lcm_migration_state ORDER BY step_name" @@ -2295,7 +2348,7 @@ def test_init_upgrades_legacy_db_and_keeps_missing_state_safe(self, tmp_path): version = state._conn.execute( "SELECT value FROM metadata WHERE key = 'schema_version'" ).fetchone()[0] - assert version == "4" + assert version == "5" tables = { row[0] @@ -2875,7 +2928,7 @@ def test_init_repairs_malformed_nodes_fts_and_sets_schema_version(self, tmp_path version = dag._conn.execute( "SELECT value FROM metadata WHERE key = 'schema_version'" ).fetchone() - assert version == ("4",) + assert version == ("5",) results = dag.search("docker", session_id="s1") assert len(results) == 1 @@ -2927,7 +2980,7 @@ def test_init_recreates_missing_nodes_fts_trigger(self, tmp_path): version = dag._conn.execute( "SELECT value FROM metadata WHERE key = 'schema_version'" ).fetchone() - assert version == ("4",) + assert version == ("5",) migration_state = dag._conn.execute( "SELECT step_name FROM lcm_migration_state ORDER BY step_name" diff --git a/tests/test_lcm_engine.py b/tests/test_lcm_engine.py index 28159cee2..0d49b5496 100644 --- a/tests/test_lcm_engine.py +++ b/tests/test_lcm_engine.py @@ -45,6 +45,36 @@ def test_shutdown_closes_lifecycle_store(tmp_path): assert engine._lifecycle._conn is None +def test_discord_short_turn_ingest_preserves_conversation_id(tmp_path): + config = LCMConfig(database_path=str(tmp_path / "discord-lanes.db")) + engine = LCMEngine(config=config) + try: + conversation_id = "agent:main:discord:thread:1520890589762031776:1520890589762031776" + engine.on_session_start( + "discord-session-1", + platform="discord", + conversation_id=conversation_id, + context_length=200_000, + ) + + engine.ingest([ + {"role": "user", "content": "needle from discord topic"}, + {"role": "assistant", "content": "topic answer"}, + ]) + + rows = engine._store.search( + "needle", + source="discord", + conversation_id=conversation_id, + ) + assert len(rows) == 1 + assert rows[0]["session_id"] == "discord-session-1" + assert rows[0]["source"] == "discord" + assert rows[0]["conversation_id"] == conversation_id + finally: + engine.shutdown() + + def test_engine_deallocation_releases_sqlite_fds_without_gc(tmp_path): fd_dir = Path("/proc/self/fd") if not fd_dir.exists(): @@ -1299,6 +1329,8 @@ def test_tool_schemas(self, engine): assert "source" in grep_props assert "descendant source lineage" in grep_props["source"]["description"] assert "unknown" in grep_props["source"]["description"] + assert "conversation_id" in grep_props + assert "Discord" in grep_props["conversation_id"]["description"] # The default scope still steers callers to the active session. description_lower = grep_schema["description"].lower() assert ( @@ -2867,6 +2899,46 @@ def test_lcm_grep_ingests_live_history_before_search(self, engine): assert result["total_results"] >= 1 assert any("needle phrase" in item["snippet"] for item in result["results"]) + def test_lcm_grep_filters_live_discord_history_by_conversation_id(self, engine): + target_conversation = "agent:main:discord:thread:topic-a:topic-a" + other_conversation = "agent:main:discord:thread:topic-b:topic-b" + engine.on_session_start( + "discord-topic-a", + platform="discord", + conversation_id=target_conversation, + context_length=200000, + ) + engine.ingest([ + {"role": "user", "content": "multichannel canary from topic a"}, + ]) + engine.on_session_start( + "discord-topic-b", + platform="discord", + conversation_id=other_conversation, + context_length=200000, + ) + engine.ingest([ + {"role": "user", "content": "multichannel canary from topic b"}, + ]) + + result = json.loads( + engine.handle_tool_call( + "lcm_grep", + { + "query": "multichannel canary", + "session_scope": "all", + "source": "discord", + "conversation_id": target_conversation, + "limit": 5, + }, + ) + ) + + assert result["conversation_id"] == target_conversation + assert result["summary_results_omitted"] is True + assert [item["conversation_id"] for item in result["results"]] == [target_conversation] + assert "topic a" in result["results"][0]["snippet"] + def test_compress_accepts_focus_topic(self, engine, monkeypatch): import importlib diff --git a/tests/test_packaging_install.py b/tests/test_packaging_install.py index 550a1f140..a4e90754c 100644 --- a/tests/test_packaging_install.py +++ b/tests/test_packaging_install.py @@ -5,6 +5,7 @@ import shutil import subprocess import sys +import types EXPECTED_LCM_TOOLS = { @@ -600,3 +601,130 @@ def spy_handle(name, args, **kwargs): assert "messages" in kwargs, f"{name}: messages kwarg not forwarded" # Depending on whether engine passes it through, at minimum verify it arrived assert kwargs["messages"] == test_messages, f"{name}: messages content mismatch" + + +def test_post_llm_hook_prefers_active_lcm_clone(monkeypatch, tmp_path): + module = _load_plugin_entrypoint_module("hermes_lcm_post_hook_active_clone") + manager = types.SimpleNamespace(_hooks={}) + fake_plugins = types.SimpleNamespace(get_plugin_manager=lambda: manager) + fake_hermes_cli = types.SimpleNamespace(plugins=fake_plugins) + monkeypatch.setitem(sys.modules, "hermes_cli", fake_hermes_cli) + monkeypatch.setitem(sys.modules, "hermes_cli.plugins", fake_plugins) + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes_home")) + + class _CtxNoTool: + def __init__(self): + self.engine = None + + def register_context_engine(self, engine): + self.engine = engine + + ctx = _CtxNoTool() + module.register(ctx) + assert ctx.engine is not None + + class _ActiveClone: + name = "lcm" + + def __init__(self): + self.current_session_id = "" + self.current_conversation_id = "" + self.starts = [] + self.ingested = [] + + def on_session_start(self, session_id, **kwargs): + self.current_session_id = session_id + self.current_conversation_id = kwargs.get("conversation_id") or session_id + self.starts.append((session_id, kwargs)) + + def ingest(self, history): + self.ingested.append(list(history)) + + active = _ActiveClone() + hook = manager._hooks["post_llm_call"][-1] + history = [{"role": "user", "content": "discord lane canary"}] + + hook( + context_compressor=active, + session_id="discord-session", + conversation_id="agent:main:discord:thread:t:t", + platform="discord", + conversation_history=history, + ) + + assert active.starts == [ + ( + "discord-session", + { + "platform": "discord", + "conversation_id": "agent:main:discord:thread:t:t", + }, + ) + ] + assert active.ingested == [history] + assert ctx.engine.current_session_id == "" + ctx.engine.shutdown() + + +def test_post_llm_hook_rebinds_legacy_singleton_between_gateway_lanes(monkeypatch, tmp_path): + module = _load_plugin_entrypoint_module("hermes_lcm_post_hook_singleton_rebind") + manager = types.SimpleNamespace(_hooks={}) + fake_plugins = types.SimpleNamespace(get_plugin_manager=lambda: manager) + fake_hermes_cli = types.SimpleNamespace(plugins=fake_plugins) + monkeypatch.setitem(sys.modules, "hermes_cli", fake_hermes_cli) + monkeypatch.setitem(sys.modules, "hermes_cli.plugins", fake_plugins) + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes_home")) + + class _CtxNoTool: + def __init__(self): + self.engine = None + + def register_context_engine(self, engine): + self.engine = engine + + ctx = _CtxNoTool() + module.register(ctx) + assert ctx.engine is not None + hook = manager._hooks["post_llm_call"][-1] + + ingests = [] + + def spy_ingest(history): + ingests.append( + ( + ctx.engine.current_session_id, + ctx.engine.current_conversation_id, + ctx.engine.current_session_platform, + list(history), + ) + ) + + monkeypatch.setattr(ctx.engine, "ingest", spy_ingest) + hook( + session_id="discord-topic-a", + conversation_id="agent:main:discord:thread:a:a", + platform="discord", + conversation_history=[{"role": "user", "content": "topic a"}], + ) + hook( + session_id="telegram-dm", + conversation_id="agent:main:telegram:private:1782862480", + platform="telegram", + conversation_history=[{"role": "user", "content": "telegram dm"}], + ) + + assert ingests == [ + ( + "discord-topic-a", + "agent:main:discord:thread:a:a", + "discord", + [{"role": "user", "content": "topic a"}], + ), + ( + "telegram-dm", + "agent:main:telegram:private:1782862480", + "telegram", + [{"role": "user", "content": "telegram dm"}], + ), + ] + ctx.engine.shutdown() diff --git a/tools.py b/tools.py index 00ba227e4..d089296e9 100644 --- a/tools.py +++ b/tools.py @@ -1067,6 +1067,7 @@ def lcm_grep(args: Dict[str, Any], **kwargs) -> str: str(raw_session_id_arg).strip() if raw_session_id_arg is not None else "" ) source = str(args.get("source") or "").strip() or None + conversation_id = str(args.get("conversation_id") or "").strip() or None role, role_error = _parse_grep_role(args.get("role")) if role_error: return json.dumps({"error": role_error}) @@ -1078,7 +1079,12 @@ def lcm_grep(args: Dict[str, Any], **kwargs) -> str: return json.dumps({"error": time_to_error}) if time_from is not None and time_to is not None and time_to < time_from: return json.dumps({"error": "time_to must be greater than or equal to time_from"}) - raw_message_filter_active = role is not None or time_from is not None or time_to is not None + raw_message_filter_active = ( + role is not None + or time_from is not None + or time_to is not None + or conversation_id is not None + ) if requested_session_scope == "current": if explicit_session_id: @@ -1130,6 +1136,7 @@ def lcm_grep(args: Dict[str, Any], **kwargs) -> str: limit=source_limit, sort=sort, source=source, + conversation_id=conversation_id, role=role, time_from=time_from, time_to=time_to, @@ -1143,6 +1150,7 @@ def lcm_grep(args: Dict[str, Any], **kwargs) -> str: "store_id": hit["store_id"], "session_id": hit["session_id"], "source": hit.get("source") or "", + "conversation_id": hit.get("conversation_id") or "", "role": hit["role"], "timestamp": timestamp_value, "snippet": hit.get("snippet", hit.get("content", "")[:200]), @@ -1211,6 +1219,7 @@ def lcm_grep(args: Dict[str, Any], **kwargs) -> str: "sort": sort, "session_scope": session_scope, "source": source, + "conversation_id": conversation_id, "limit": limit, "total_results": len(results), "results": results[:limit], @@ -1392,6 +1401,7 @@ def lcm_expand(args: Dict[str, Any], **kwargs) -> str: "source_type": "raw_message", "session_id": stored_session_id, "source": stored.get("source") or "", + "conversation_id": stored.get("conversation_id") or "", "role": stored.get("role"), "timestamp": stored.get("timestamp", 0), "tool_call_id": stored.get("tool_call_id") or "",