From 2178fc91cef27963ed56ef756f505371701799c9 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Thu, 4 Dec 2025 17:13:41 -0500 Subject: [PATCH] fix: search and add create_observations --- src/crud/document.py | 61 +++++++++++++++++++------ src/utils/search.py | 75 +++++++++++++++++++++++++++++-- tests/conftest.py | 68 ++++++++++++++++++++++++++++ tests/crud/test_document.py | 46 ++++++++++--------- tests/routes/test_observations.py | 54 +++++++++------------- 5 files changed, 233 insertions(+), 71 deletions(-) diff --git a/src/crud/document.py b/src/crud/document.py index 6d245579..2fc02205 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -423,22 +423,31 @@ async def create_observations( except ValueError as e: raise ValidationException(str(e)) from e - # Create document objects + # Create document objects and track embeddings for vector store honcho_documents: list[models.Document] = [] + # Group observations by collection (observer, observed) for vector store upserts + collection_embeddings: dict[ + tuple[str, str], list[tuple[models.Document, list[float]]] + ] = {} + for obs, embedding in zip(observations, embeddings, strict=True): - honcho_documents.append( - models.Document( - workspace_name=workspace_name, - observer=obs.observer_id, - observed=obs.observed_id, - content=obs.content, - level="explicit", # Manually created observations are always explicit - times_derived=1, - internal_metadata={}, # No message_ids since not derived from messages - embedding=embedding, - session_name=obs.session_id, - ) + doc = models.Document( + workspace_name=workspace_name, + observer=obs.observer_id, + observed=obs.observed_id, + content=obs.content, + level="explicit", # Manually created observations are always explicit + times_derived=1, + internal_metadata={}, # No message_ids since not derived from messages + session_name=obs.session_id, ) + honcho_documents.append(doc) + + # Track embedding for vector store (grouped by collection) + collection_key = (obs.observer_id, obs.observed_id) + if collection_key not in collection_embeddings: + collection_embeddings[collection_key] = [] + collection_embeddings[collection_key].append((doc, embedding)) try: db.add_all(honcho_documents) @@ -446,6 +455,32 @@ async def create_observations( # Refresh all documents to get generated IDs and timestamps for doc in honcho_documents: await db.refresh(doc) + + # Store embeddings in vector store after documents are committed (IDs now available) + vector_store = get_vector_store() + for (observer, observed), docs_with_embeddings in collection_embeddings.items(): + namespace = vector_store.get_document_namespace( + workspace_name, observer, observed + ) + + # Build vector records with metadata for filtering + vector_records: list[VectorRecord] = [] + for doc, embedding in docs_with_embeddings: + vector_records.append( + VectorRecord( + id=doc.id, + embedding=embedding, + metadata={ + "workspace_name": workspace_name, + "observer": observer, + "observed": observed, + "session_name": doc.session_name, + "level": doc.level, + }, + ) + ) + await vector_store.upsert_many(namespace, vector_records) + except IntegrityError as e: await db.rollback() raise ValidationException( diff --git a/src/utils/search.py b/src/utils/search.py index 3017a025..015fd4d2 100644 --- a/src/utils/search.py +++ b/src/utils/search.py @@ -148,6 +148,64 @@ async def _semantic_search( return ordered_messages +async def _filter_by_peer_perspective( + db: AsyncSession, + messages: list[models.Message], + workspace_name: str, + peer_name: str, +) -> list[models.Message]: + """ + Filter messages by peer perspective (temporal session membership). + + Only keeps messages from sessions where the peer was a member at the time + the message was created (between joined_at and left_at). + + Args: + db: Database session + messages: List of messages to filter + workspace_name: Name of the workspace + peer_name: Name of the peer whose perspective to use + + Returns: + Filtered list of messages + """ + if not messages: + return [] + + # Get all session memberships for this peer in this workspace + session_memberships_query = ( + select(session_peers_table) + .where(session_peers_table.c.workspace_name == workspace_name) + .where(session_peers_table.c.peer_name == peer_name) + ) + result = await db.execute(session_memberships_query) + memberships = result.all() + + # Build a lookup of session -> time windows + session_windows: dict[str, list[tuple[Any, Any]]] = {} + for membership in memberships: + session_name = membership.session_name + if session_name not in session_windows: + session_windows[session_name] = [] + session_windows[session_name].append((membership.joined_at, membership.left_at)) + + # Filter messages + filtered_messages: list[models.Message] = [] + for msg in messages: + if msg.session_name not in session_windows: + continue + + # Check if message was created during any of the peer's active windows in this session + for joined_at, left_at in session_windows[msg.session_name]: + if msg.created_at >= joined_at and ( + left_at is None or msg.created_at <= left_at + ): + filtered_messages.append(msg) + break # Don't add the same message twice + + return filtered_messages + + async def _fulltext_search( db: AsyncSession, query: str, @@ -237,8 +295,9 @@ async def search( stmt = select(models.Message) # Handle special peer_perspective filter + peer_perspective_name: str | None = None if filters and "peer_perspective" in filters: - peer_name = filters["peer_perspective"] + peer_perspective_name = filters["peer_perspective"] # Remove from filters dict so apply_filter doesn't try to handle it filters = {k: v for k, v in filters.items() if k != "peer_perspective"} # Safety: peer_perspective must be scoped to a workspace @@ -262,7 +321,7 @@ async def search( models.Message.created_at <= session_peers_table.c.left_at, ), ), - ).where(session_peers_table.c.peer_name == peer_name) + ).where(session_peers_table.c.peer_name == peer_perspective_name) stmt = apply_filter(stmt, models.Message, filters) @@ -273,8 +332,8 @@ async def search( workspace_name: str | None = filters.get("workspace_id") if filters else None if settings.EMBED_MESSAGES and isinstance(workspace_name, str): # Type narrowing: workspace_name is guaranteed to be str in this block - # Get more results for fusion - semantic_limit = limit * 2 + # Get more results for fusion (increase if peer_perspective filtering is applied post-search) + semantic_limit = limit * 4 if peer_perspective_name else limit * 2 semantic_results = await _semantic_search( db=db, query=query, @@ -282,6 +341,14 @@ async def search( limit=semantic_limit, filters=filters, ) + + # Apply peer_perspective filtering to semantic results if needed + # Vector store can't handle temporal filtering (joined_at/left_at), so filter post-search + if peer_perspective_name: + semantic_results = await _filter_by_peer_perspective( + db, semantic_results, workspace_name, peer_perspective_name + ) + search_results.append(semantic_results) # Perform full-text search diff --git a/tests/conftest.py b/tests/conftest.py index 9c11a570..50cc5b85 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -376,6 +376,74 @@ def mock_openai_embeddings(): yield {"embed": mock_embed, "batch_embed": mock_batch_embed} +@pytest.fixture(autouse=True) +def mock_vector_store(): + """Mock vector store operations for testing""" + from unittest.mock import AsyncMock, MagicMock + + from src.vector_store import QueryResult, VectorRecord + + # Create a mock vector store that stores vectors in memory + vector_storage: dict[str, dict[str, tuple[list[float], dict[str, Any]]]] = {} + + async def mock_upsert(namespace: str, vector: VectorRecord) -> None: + if namespace not in vector_storage: + vector_storage[namespace] = {} + vector_storage[namespace][vector.id] = (vector.embedding, vector.metadata) + + async def mock_upsert_many(namespace: str, vectors: list[VectorRecord]) -> None: + if namespace not in vector_storage: + vector_storage[namespace] = {} + for vector in vectors: + vector_storage[namespace][vector.id] = (vector.embedding, vector.metadata) + + async def mock_query( + namespace: str, embedding: list[float], **kwargs: Any + ) -> list[QueryResult]: + _ = embedding # unused in mock + if namespace not in vector_storage: + return [] + + # Simple mock: return all vectors in the namespace as results + results: list[QueryResult] = [] + for vec_id, (_vec_embedding, metadata) in vector_storage[namespace].items(): + results.append( + QueryResult( + id=vec_id, + score=0.1, # Mock score + metadata=metadata, + ) + ) + top_k: int = kwargs.get("top_k", 10) + return results[:top_k] + + async def mock_delete_many(namespace: str, ids: list[str]) -> None: + if namespace in vector_storage: + for vec_id in ids: + vector_storage[namespace].pop(vec_id, None) + + async def mock_delete_namespace(namespace: str) -> None: + vector_storage.pop(namespace, None) + + with ( + patch("src.vector_store.get_vector_store") as mock_get_vs, + ): + mock_vs = MagicMock() + mock_vs.upsert = AsyncMock(side_effect=mock_upsert) + mock_vs.upsert_many = AsyncMock(side_effect=mock_upsert_many) + mock_vs.query = AsyncMock(side_effect=mock_query) + mock_vs.delete_many = AsyncMock(side_effect=mock_delete_many) + mock_vs.delete_namespace = AsyncMock(side_effect=mock_delete_namespace) + mock_vs.get_document_namespace = ( + lambda ws, obs, obd: f"honcho:{ws}:{obs}:{obd}" # pyright: ignore[reportUnknownLambdaType] + ) + mock_vs.get_message_namespace = lambda ws: f"honcho:{ws}:messages" # pyright: ignore[reportUnknownLambdaType] + + mock_get_vs.return_value = mock_vs + + yield mock_vs + + @pytest.fixture(autouse=True) def mock_llm_call_functions(): """Mock LLM functions to avoid needing API keys during tests""" diff --git a/tests/crud/test_document.py b/tests/crud/test_document.py index a53f8df6..46533da1 100644 --- a/tests/crud/test_document.py +++ b/tests/crud/test_document.py @@ -60,7 +60,6 @@ class TestDocumentCRUD: observer=test_peer.name, observed=test_peer2.name, content="Test observation 1", - embedding=[0.1] * 1536, session_name=test_session.name, ) doc2 = models.Document( @@ -68,7 +67,6 @@ class TestDocumentCRUD: observer=test_peer.name, observed=test_peer2.name, content="Test observation 2", - embedding=[0.2] * 1536, session_name=test_session.name, ) db_session.add_all([doc1, doc2]) @@ -100,25 +98,34 @@ class TestDocumentCRUD: db_session, test_workspace, test_peer ) - # Create test documents with different embeddings - doc1 = models.Document( + # Create test documents using create_documents to ensure they're in vector store + doc_schemas = [ + schemas.DocumentCreate( + content="User likes pizza", + embedding=[0.9] * 1536, + session_name=test_session.name, + metadata=schemas.DocumentMetadata( + message_ids=[1], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + schemas.DocumentCreate( + content="User dislikes vegetables", + embedding=[0.1] * 1536, + session_name=test_session.name, + metadata=schemas.DocumentMetadata( + message_ids=[2], + message_created_at="2025-01-01T00:00:00Z", + ), + ), + ] + await crud.create_documents( + db_session, + doc_schemas, workspace_name=test_workspace.name, observer=test_peer.name, observed=test_peer2.name, - content="User likes pizza", - embedding=[0.9] * 1536, - session_name=test_session.name, ) - doc2 = models.Document( - workspace_name=test_workspace.name, - observer=test_peer.name, - observed=test_peer2.name, - content="User dislikes vegetables", - embedding=[0.1] * 1536, - session_name=test_session.name, - ) - db_session.add_all([doc1, doc2]) - await db_session.flush() # Query documents results = await crud.query_documents( @@ -150,7 +157,6 @@ class TestDocumentCRUD: observer=test_peer.name, observed=test_peer2.name, content="Test observation", - embedding=[0.1] * 1536, session_name=test_session.name, ) db_session.add(doc) @@ -214,8 +220,8 @@ class TestDocumentCRUD: doc_schemas = [ schemas.DocumentCreate( content="Observation 1", - session_name=test_session.name, embedding=[0.1] * 1536, + session_name=test_session.name, level="explicit", metadata=schemas.DocumentMetadata( message_ids=[1, 2, 3, 4, 5], @@ -224,8 +230,8 @@ class TestDocumentCRUD: ), schemas.DocumentCreate( content="Observation 2", - session_name=test_session.name, embedding=[0.2] * 1536, + session_name=test_session.name, level="deductive", metadata=schemas.DocumentMetadata( message_ids=[6, 7, 8, 9, 10], diff --git a/tests/routes/test_observations.py b/tests/routes/test_observations.py index 9f1bba94..d32801f9 100644 --- a/tests/routes/test_observations.py +++ b/tests/routes/test_observations.py @@ -62,7 +62,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="User prefers dark mode", - embedding=[0.1] * 1536, session_name=test_session.name, ) doc2 = models.Document( @@ -70,7 +69,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="User works late at night", - embedding=[0.2] * 1536, session_name=test_session.name, ) db_session.add_all([doc1, doc2]) @@ -170,7 +168,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="Peer1 observes Peer2", - embedding=[0.1] * 1536, session_name=test_session.name, ) doc2 = models.Document( @@ -178,7 +175,6 @@ class TestObservationRoutes: observer=test_peer2.name, observed=test_peer3.name, content="Peer2 observes Peer3", - embedding=[0.2] * 1536, session_name=test_session.name, ) db_session.add_all([doc1, doc2]) @@ -233,7 +229,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="First observation", - embedding=[0.1] * 1536, session_name=test_session.name, ) db_session.add(doc1) @@ -244,7 +239,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="Second observation", - embedding=[0.2] * 1536, session_name=test_session.name, ) db_session.add(doc2) @@ -298,7 +292,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content=f"Observation {i}", - embedding=[0.1 * i] * 1536, session_name=test_session.name, ) db_session.add(doc) @@ -350,30 +343,27 @@ class TestObservationRoutes: db_session.add(test_session) await db_session.commit() - # Create collection - await self._create_collection( - db_session, test_workspace.name, test_peer.name, test_peer2.name + # Create test observations via API (this populates the vector store) + create_response = client.post( + f"/v2/workspaces/{test_workspace.name}/observations", + json={ + "observations": [ + { + "content": "User loves pizza and pasta", + "observer_id": test_peer.name, + "observed_id": test_peer2.name, + "session_id": test_session.name, + }, + { + "content": "User dislikes vegetables", + "observer_id": test_peer.name, + "observed_id": test_peer2.name, + "session_id": test_session.name, + }, + ] + }, ) - - # Create test observations - doc1 = models.Document( - workspace_name=test_workspace.name, - observer=test_peer.name, - observed=test_peer2.name, - content="User loves pizza and pasta", - embedding=[0.9] * 1536, - session_name=test_session.name, - ) - doc2 = models.Document( - workspace_name=test_workspace.name, - observer=test_peer.name, - observed=test_peer2.name, - content="User dislikes vegetables", - embedding=[0.5] * 1536, - session_name=test_session.name, - ) - db_session.add_all([doc1, doc2]) - await db_session.commit() + assert create_response.status_code == 200 # Query observations response = client.post( @@ -436,7 +426,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content=f"Observation about topic {i}", - embedding=[0.1 * i] * 1536, session_name=test_session.name, ) db_session.add(doc) @@ -496,7 +485,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="Test observation", - embedding=[0.5] * 1536, session_name=test_session.name, ) db_session.add(doc) @@ -620,7 +608,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="Test observation to delete", - embedding=[0.1] * 1536, session_name=test_session.name, ) db_session.add(doc) @@ -725,7 +712,6 @@ class TestObservationRoutes: observer=test_peer.name, observed=test_peer2.name, content="Test observation content", - embedding=[0.1] * 1536, session_name=test_session.name, ) db_session.add(doc)