diff --git a/src/crud/message.py b/src/crud/message.py index 76f9cab7..a4136d00 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -722,6 +722,7 @@ async def _fetch_messages_by_ids( *, after_date: datetime | None = None, before_date: datetime | None = None, + peer_name: str | None = None, ) -> list[models.Message]: """Fetch messages by ID, preserving the supplied ordering.""" fetch_stmt = ( @@ -733,6 +734,8 @@ async def _fetch_messages_by_ids( fetch_stmt = fetch_stmt.where(models.Message.created_at >= after_date) if before_date: fetch_stmt = fetch_stmt.where(models.Message.created_at <= before_date) + if peer_name: + fetch_stmt = fetch_stmt.where(models.Message.peer_name == peer_name) result = await db.execute(fetch_stmt) messages_by_id = {msg.public_id: msg for msg in result.scalars().all()} @@ -751,6 +754,7 @@ async def _search_messages_pgvector( before_date: datetime | None = None, limit: int = 10, context_window: int = 2, + peer_name: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: """Run semantic message search against pgvector-backed embeddings.""" # pgvector path: cosine distance in SQL @@ -781,6 +785,8 @@ async def _search_messages_pgvector( match_stmt = match_stmt.where(models.Message.created_at >= after_date) if before_date: match_stmt = match_stmt.where(models.Message.created_at <= before_date) + if peer_name: + match_stmt = match_stmt.where(models.Message.peer_name == peer_name) result = await db.execute(match_stmt) matched_messages = _deduplicate_messages(result.scalars().all(), limit) @@ -801,11 +807,15 @@ async def _semantic_search_messages( after_date: datetime | None = None, before_date: datetime | None = None, observer: str | None = None, + peer_name: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: """Run semantic message search with optional temporal filters. When observer is provided and session_name is None, results are scoped to sessions the observer has any membership record in. + + When peer_name is provided, matched messages are restricted to those + authored by that peer. Context messages around the matches are unaffected. """ # Pre-fetch peer session scope if needed (short-lived DB session) allowed_session_names: list[str] | None = None @@ -838,6 +848,7 @@ async def _semantic_search_messages( message_ids, after_date=after_date, before_date=before_date, + peer_name=peer_name, ) )[:limit] snippets = await _build_merged_snippets( @@ -857,6 +868,7 @@ async def _semantic_search_messages( before_date=before_date, limit=limit, context_window=context_window, + peer_name=peer_name, ) _expunge_snippets(db, snippets) return snippets @@ -870,6 +882,7 @@ async def search_messages( context_window: int = 2, embedding: list[float] | None = None, observer: str | None = None, + peer_name: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: """ Search for messages using semantic similarity and return conversation snippets. @@ -886,6 +899,11 @@ async def search_messages( embedding: Optional pre-computed embedding observer: When provided and session_name is None, scope results to sessions this peer belongs to + peer_name: When provided, restrict matched messages to those authored + by this peer. Context (surrounding messages) is unaffected so the + conversation flow remains readable. Used by agent tools to scope + results to the observed peer's own statements (prevents bio info + from one peer leaking into another peer's representation). Returns: List of tuples: (matched_messages, context_messages) @@ -903,6 +921,7 @@ async def search_messages( context_window=context_window, operation_name="message.search_messages", observer=observer, + peer_name=peer_name, ) diff --git a/src/utils/agent_tools.py b/src/utils/agent_tools.py index 0d95cf6d..5d853e8b 100644 --- a/src/utils/agent_tools.py +++ b/src/utils/agent_tools.py @@ -1558,6 +1558,7 @@ async def _handle_search_messages(ctx: ToolContext, tool_input: dict[str, Any]) context_window=2, embedding=query_embedding, observer=ctx.observer, + peer_name=ctx.observed, ) if not snippets: return f"No messages found for query '{query}'" diff --git a/tests/integration/test_message_embeddings.py b/tests/integration/test_message_embeddings.py index c5e3f100..578271af 100644 --- a/tests/integration/test_message_embeddings.py +++ b/tests/integration/test_message_embeddings.py @@ -393,8 +393,9 @@ async def test_search_messages_external_lookup_happens_before_tracked_db( *, after_date: datetime | None = None, before_date: datetime | None = None, + peer_name: str | None = None, ) -> list[models.Message]: - _ = (workspace_name, message_ids, after_date, before_date) + _ = (workspace_name, message_ids, after_date, before_date, peer_name) assert db is fake_db call_order.append("fetch") return [message] @@ -495,8 +496,9 @@ async def test_search_messages_temporal_external_lookup_happens_before_tracked_d *, after_date: datetime | None = None, before_date: datetime | None = None, + peer_name: str | None = None, ) -> list[models.Message]: - _ = (workspace_name, message_ids) + _ = (workspace_name, message_ids, peer_name) assert db is fake_db assert after_date is not None assert before_date is not None diff --git a/tests/utils/test_agent_tools.py b/tests/utils/test_agent_tools.py index 95406587..cecbd7bb 100644 --- a/tests/utils/test_agent_tools.py +++ b/tests/utils/test_agent_tools.py @@ -756,6 +756,54 @@ class TestSearchMessages: # Should return some result (may be empty if semantic search doesn't match) assert isinstance(result, str) + async def test_filters_by_observed_peer_author( + self, + make_tool_context: Callable[..., ToolContext], + monkeypatch: pytest.MonkeyPatch, + ): + """search_messages must scope by ctx.observed (peer author). + + Bug: the handler passes `observer` and `session_name` to + `crud.search_messages` but does not pass `peer_name` (or equivalent + author filter). When the dreamer's deduction specialist runs for + observed=assistant, it asks for "bio facts about assistant" but the + un-scoped search returns user-authored messages too — leading to user + bio info leaking into the assistant's representation/peer card. + + Invariant: the kwargs passed to `crud.search_messages` must include a + peer-author filter equal to `ctx.observed`. This test fails until the + handler is updated to pass it (and `crud.search_messages` accepts and + applies it at the storage layer). + """ + ctx = make_tool_context() + captured_kwargs: dict[str, Any] = {} + + async def fake_search_messages( + **kwargs: Any, + ) -> list[tuple[list[models.Message], list[models.Message]]]: + captured_kwargs.update(kwargs) + return [] + + async def fake_embed(query: str) -> list[float]: + _ = query + return [0.1, 0.2, 0.3] + + monkeypatch.setattr("src.utils.agent_tools.embedding_client.embed", fake_embed) + monkeypatch.setattr( + "src.utils.agent_tools.crud.search_messages", fake_search_messages + ) + + await _handle_search_messages(ctx, {"query": "anything"}) + + assert "peer_name" in captured_kwargs, ( + "search_messages was called without a peer_name filter; " + "results would include messages from other peers in the session" + ) + assert captured_kwargs["peer_name"] == ctx.observed, ( + f"peer_name filter must be ctx.observed={ctx.observed!r}, " + f"got {captured_kwargs.get('peer_name')!r}" + ) + @pytest.mark.asyncio class TestGrepMessages: