fix: search and add create_observations

This commit is contained in:
Rajat Ahuja 2025-12-04 17:13:41 -05:00
parent f1ac894713
commit 2178fc91ce
5 changed files with 233 additions and 71 deletions

View File

@ -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(

View File

@ -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

View File

@ -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"""

View File

@ -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],

View File

@ -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)