fix: search and add create_observations
This commit is contained in:
parent
f1ac894713
commit
2178fc91ce
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue