diff --git a/src/crud/message.py b/src/crud/message.py index 41e0053b..08c4c861 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -9,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from src import models, schemas from src.config import settings +from src.dependencies import tracked_db from src.embedding_client import embedding_client from src.utils.filter import apply_filter from src.utils.formatting import ILIKE_ESCAPE_CHAR, escape_ilike_pattern @@ -34,6 +35,40 @@ def _deduplicate_messages( return result +def _expunge_snippets( + db: AsyncSession, snippets: list[tuple[list[models.Message], list[models.Message]]] +) -> None: + """Detach snippet messages from the session, guarding against duplicates.""" + seen: set[int] = set() + for matches, context in snippets: + for msg in [*matches, *context]: + obj_id = id(msg) + if obj_id in seen: + continue + db.expunge(msg) + seen.add(obj_id) + + +async def get_peer_session_names( + db: AsyncSession, + workspace_name: str, + peer_name: str, +) -> list[str]: + """Get all session names where a peer has any membership record. + + Any membership record (regardless of joined_at/left_at) grants visibility + to all messages in that session. + """ + stmt = ( + select(models.session_peers_table.c.session_name) + .where(models.session_peers_table.c.workspace_name == workspace_name) + .where(models.session_peers_table.c.peer_name == peer_name) + .distinct() + ) + result = await db.execute(stmt) + return [row[0] for row in result.all()] + + def _apply_token_limit( base_conditions: list[ColumnElement[Any]], token_limit: int ) -> Select[tuple[models.Message]]: @@ -595,22 +630,19 @@ async def update_message( async def _search_messages_external( - db: AsyncSession, workspace_name: str, query_embedding: list[float], limit: int, *, session_name: str | None = None, + allowed_session_names: list[str] | None = None, after_date: datetime | None = None, before_date: datetime | None = None, -) -> list[models.Message]: - """Query the external vector store for messages and fetch them from the DB. +) -> list[str]: + """Query the external vector store and return ordered message IDs. Multiple vector records can map to the same message (chunked embeddings), so we oversample from the vector store and deduplicate by message_id. - - Date filters are applied at the DB level since external vector stores - don't support temporal filtering. """ external_vector_store = get_external_vector_store() if external_vector_store is None: @@ -621,6 +653,8 @@ async def _search_messages_external( vector_filters: dict[str, Any] = {} if session_name: vector_filters["session_name"] = session_name + elif allowed_session_names is not None: + vector_filters["session_name"] = {"in": allowed_session_names} # Oversample: chunks can map to the same message, and date filters are # applied post-fetch (vector stores don't support temporal filtering), @@ -648,7 +682,18 @@ async def _search_messages_external( if not message_ids: return [] - # Fetch from DB with optional date filtering + return message_ids + + +async def _fetch_messages_by_ids( + db: AsyncSession, + workspace_name: str, + message_ids: list[str], + *, + after_date: datetime | None = None, + before_date: datetime | None = None, +) -> list[models.Message]: + """Fetch messages by ID, preserving the supplied ordering.""" fetch_stmt = ( select(models.Message) .where(models.Message.public_id.in_(message_ids)) @@ -662,18 +707,139 @@ async def _search_messages_external( result = await db.execute(fetch_stmt) messages_by_id = {msg.public_id: msg for msg in result.scalars().all()} - # Preserve vector store similarity order, apply limit - return [messages_by_id[mid] for mid in message_ids if mid in messages_by_id][:limit] + return [messages_by_id[mid] for mid in message_ids if mid in messages_by_id] + + +async def _search_messages_pgvector( + db: AsyncSession, + workspace_name: str, + session_name: str | None, + *, + query_embedding: list[float], + allowed_session_names: list[str] | None = None, + after_date: datetime | None = None, + before_date: datetime | None = None, + limit: int = 10, + context_window: int = 2, +) -> list[tuple[list[models.Message], list[models.Message]]]: + """Run semantic message search against pgvector-backed embeddings.""" + # pgvector path: cosine distance in SQL + # Oversample because a message with multiple embedding chunks can + # produce duplicate rows; we deduplicate in Python to preserve HNSW + # index usage (a DISTINCT ON subquery would prevent the index scan). + match_stmt = ( + select(models.Message) + .join( + models.MessageEmbedding, + models.Message.public_id == models.MessageEmbedding.message_id, + ) + .where(models.MessageEmbedding.workspace_name == workspace_name) + .order_by(models.MessageEmbedding.embedding.cosine_distance(query_embedding)) + .limit(limit * 2) + ) + + if session_name: + match_stmt = match_stmt.where( + models.MessageEmbedding.session_name == session_name + ) + elif allowed_session_names is not None: + match_stmt = match_stmt.where( + models.MessageEmbedding.session_name.in_(allowed_session_names) + ) + + if after_date: + 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) + + result = await db.execute(match_stmt) + matched_messages = _deduplicate_messages(result.scalars().all(), limit) + + return await _build_merged_snippets( + db, workspace_name, matched_messages, context_window + ) + + +async def _semantic_search_messages( + workspace_name: str, + session_name: str | None, + *, + query_embedding: list[float], + limit: int = 10, + context_window: int = 2, + operation_name: str, + after_date: datetime | None = None, + before_date: datetime | None = None, + observer: 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. + """ + # Pre-fetch peer session scope if needed (short-lived DB session) + allowed_session_names: list[str] | None = None + if observer and not session_name: + async with tracked_db(f"{operation_name}.peer_scope") as db: + allowed_session_names = await get_peer_session_names( + db, workspace_name, observer + ) + if not allowed_session_names: + return [] + + if settings.VECTOR_STORE.TYPE != "pgvector" and settings.VECTOR_STORE.MIGRATED: + message_ids = await _search_messages_external( + workspace_name, + query_embedding, + limit, + session_name=session_name, + allowed_session_names=allowed_session_names, + after_date=after_date, + before_date=before_date, + ) + if not message_ids: + return [] + + async with tracked_db(operation_name) as db: + matched_messages = ( + await _fetch_messages_by_ids( + db, + workspace_name, + message_ids, + after_date=after_date, + before_date=before_date, + ) + )[:limit] + snippets = await _build_merged_snippets( + db, workspace_name, matched_messages, context_window + ) + _expunge_snippets(db, snippets) + return snippets + + async with tracked_db(operation_name) as db: + snippets = await _search_messages_pgvector( + db, + workspace_name, + session_name, + query_embedding=query_embedding, + allowed_session_names=allowed_session_names, + after_date=after_date, + before_date=before_date, + limit=limit, + context_window=context_window, + ) + _expunge_snippets(db, snippets) + return snippets async def search_messages( - db: AsyncSession, workspace_name: str, session_name: str | None, query: str, limit: int = 10, context_window: int = 2, embedding: list[float] | None = None, + observer: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: """ Search for messages using semantic similarity and return conversation snippets. @@ -682,86 +848,44 @@ async def search_messages( snippets within the same session are merged to avoid repetition. Args: - db: Database session workspace_name: Name of the workspace session_name: Name of the session (optional) query: Search query text limit: Maximum number of matching messages to return context_window: Number of messages before/after each match to include embedding: Optional pre-computed embedding + observer: When provided and session_name is None, scope results + to sessions this peer belongs to Returns: List of tuples: (matched_messages, context_messages) Each snippet may contain multiple matches if they were close together. Context messages are ordered chronologically and include the matched messages. """ - # Use provided embedding or generate one query_embedding = ( embedding if embedding is not None else await embedding_client.embed(query) ) - - if settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED: - # pgvector path: cosine distance in SQL - # Oversample because a message with multiple embedding chunks can - # produce duplicate rows; we deduplicate in Python to preserve HNSW - # index usage (a DISTINCT ON subquery would prevent the index scan). - match_stmt = ( - select(models.Message) - .join( - models.MessageEmbedding, - models.Message.public_id == models.MessageEmbedding.message_id, - ) - .where(models.MessageEmbedding.workspace_name == workspace_name) - .order_by( - models.MessageEmbedding.embedding.cosine_distance(query_embedding) - ) - .limit(limit * 2) - ) - - if session_name: - match_stmt = match_stmt.where( - models.MessageEmbedding.session_name == session_name - ) - - result = await db.execute(match_stmt) - matched_messages = _deduplicate_messages(result.scalars().all(), limit) - else: - # External vector store path - matched_messages = await _search_messages_external( - db, workspace_name, query_embedding, limit, session_name=session_name - ) - - return await _build_merged_snippets( - db, workspace_name, matched_messages, context_window + return await _semantic_search_messages( + workspace_name, + session_name, + query_embedding=query_embedding, + limit=limit, + context_window=context_window, + operation_name="message.search_messages", + observer=observer, ) -async def grep_messages( +async def _grep_messages_internal( db: AsyncSession, workspace_name: str, session_name: str | None, text: str, limit: int = 10, context_window: int = 2, + allowed_session_names: list[str] | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: - """ - Search for messages containing specific text (case-insensitive substring match). - - Unlike semantic search, this finds EXACT text matches. Useful for finding - specific names, dates, phrases, or keywords. - - Args: - db: Database session - workspace_name: Name of the workspace - session_name: Name of the session (optional - searches all sessions if None) - text: Text to search for (case-insensitive) - limit: Maximum number of matching messages to return - context_window: Number of messages before/after each match to include - - Returns: - List of tuples: (matched_messages, context_messages) - Each snippet may contain multiple matches if they were close together. - """ + """Internal implementation of exact-text message search.""" # Build the base query with ILIKE for case-insensitive text search escaped_text = escape_ilike_pattern(text) match_stmt = ( @@ -776,6 +900,10 @@ async def grep_messages( if session_name: match_stmt = match_stmt.where(models.Message.session_name == session_name) + elif allowed_session_names is not None: + match_stmt = match_stmt.where( + models.Message.session_name.in_(allowed_session_names) + ) result = await db.execute(match_stmt) matched_messages = list(result.scalars().all()) @@ -785,6 +913,56 @@ async def grep_messages( ) +async def grep_messages( + workspace_name: str, + session_name: str | None, + text: str, + limit: int = 10, + context_window: int = 2, + observer: str | None = None, +) -> list[tuple[list[models.Message], list[models.Message]]]: + """ + Search for messages containing specific text (case-insensitive substring match). + + Unlike semantic search, this finds EXACT text matches. Useful for finding + specific names, dates, phrases, or keywords. + + Args: + workspace_name: Name of the workspace + session_name: Name of the session (optional - searches all sessions if None) + text: Text to search for (case-insensitive) + limit: Maximum number of matching messages to return + context_window: Number of messages before/after each match to include + observer: When provided and session_name is None, scope results + to sessions this peer belongs to + + Returns: + List of tuples: (matched_messages, context_messages) + Each snippet may contain multiple matches if they were close together. + """ + async with tracked_db("message.grep_messages") as db: + # Pre-fetch peer session scope if needed + allowed_session_names = None + if observer and not session_name: + allowed_session_names = await get_peer_session_names( + db, workspace_name, observer + ) + if not allowed_session_names: + return [] + + snippets = await _grep_messages_internal( + db, + workspace_name, + session_name, + text, + limit, + context_window, + allowed_session_names=allowed_session_names, + ) + _expunge_snippets(db, snippets) + return snippets + + async def get_messages_by_date_range( db: AsyncSession, workspace_name: str, @@ -793,6 +971,7 @@ async def get_messages_by_date_range( before_date: datetime | None = None, limit: int = 20, order: str = "desc", + observer: str | None = None, ) -> list[models.Message]: """ Get messages within a date range. @@ -805,14 +984,27 @@ async def get_messages_by_date_range( before_date: Return messages before this datetime limit: Maximum messages to return order: Sort order - 'asc' for oldest first, 'desc' for newest first + observer: When provided and session_name is None, scope results + to sessions this peer belongs to Returns: List of messages within the date range """ + # Pre-fetch peer session scope if needed + allowed_session_names = None + if observer and not session_name: + allowed_session_names = await get_peer_session_names( + db, workspace_name, observer + ) + if not allowed_session_names: + return [] + stmt = select(models.Message).where(models.Message.workspace_name == workspace_name) if session_name: stmt = stmt.where(models.Message.session_name == session_name) + elif allowed_session_names is not None: + stmt = stmt.where(models.Message.session_name.in_(allowed_session_names)) if after_date: stmt = stmt.where(models.Message.created_at >= after_date) if before_date: @@ -830,7 +1022,6 @@ async def get_messages_by_date_range( async def search_messages_temporal( - db: AsyncSession, workspace_name: str, session_name: str | None, query: str, @@ -839,6 +1030,7 @@ async def search_messages_temporal( limit: int = 10, context_window: int = 2, embedding: list[float] | None = None, + observer: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: """ Search for messages using semantic similarity with optional date filtering. @@ -847,7 +1039,6 @@ async def search_messages_temporal( to find recent mentions, or before_date to find what was said before a certain point. Args: - db: Database session workspace_name: Name of the workspace session_name: Name of the session (optional) query: Search query text @@ -856,58 +1047,24 @@ async def search_messages_temporal( limit: Maximum number of matching messages to return context_window: Number of messages before/after each match to include embedding: Optional pre-computed embedding for the query + observer: When provided and session_name is None, scope results + to sessions this peer belongs to Returns: List of tuples: (matched_messages, context_messages) Each snippet may contain multiple matches if they were close together. """ - # Use provided embedding or generate one query_embedding = ( embedding if embedding is not None else await embedding_client.embed(query) ) - - if settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED: - # pgvector path: cosine distance in SQL with date filters - # Oversample to handle chunk duplicates (see search_messages comment) - match_stmt = ( - select(models.Message) - .join( - models.MessageEmbedding, - models.Message.public_id == models.MessageEmbedding.message_id, - ) - .where(models.MessageEmbedding.workspace_name == workspace_name) - ) - - if session_name: - match_stmt = match_stmt.where( - models.MessageEmbedding.session_name == session_name - ) - - # Apply date filters on the Message table - if after_date: - 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) - - # Order by similarity and limit - match_stmt = match_stmt.order_by( - models.MessageEmbedding.embedding.cosine_distance(query_embedding) - ).limit(limit * 2) - - result = await db.execute(match_stmt) - matched_messages = _deduplicate_messages(result.scalars().all(), limit) - else: - # External vector store path with post-fetch date filtering - matched_messages = await _search_messages_external( - db, - workspace_name, - query_embedding, - limit, - session_name=session_name, - after_date=after_date, - before_date=before_date, - ) - - return await _build_merged_snippets( - db, workspace_name, matched_messages, context_window + return await _semantic_search_messages( + workspace_name, + session_name, + query_embedding=query_embedding, + after_date=after_date, + before_date=before_date, + limit=limit, + context_window=context_window, + operation_name="message.search_messages_temporal", + observer=observer, ) diff --git a/src/crud/peer.py b/src/crud/peer.py index 0088ebeb..4c144b9b 100644 --- a/src/crud/peer.py +++ b/src/crud/peer.py @@ -223,7 +223,11 @@ async def update_peer( db: AsyncSession, workspace_name: str, peer_name: str, peer: schemas.PeerUpdate ) -> models.Peer: """ - Update a peer. + Get or create a peer, then apply metadata and configuration updates. + + If the peer does not exist, the workspace and peer are created first. + Provided metadata and configuration replace the existing values when + present. Args: db: Database session @@ -235,9 +239,8 @@ async def update_peer( The updated peer Raises: - ResourceNotFoundException: If the peer does not exist - ValidationException: If the update data is invalid - ConflictException: If the update violates a unique constraint + ConflictException: If concurrent creation prevents fetching or creating + the peer """ peers_result = await get_or_create_peers( db, workspace_name, [schemas.PeerCreate(name=peer_name)] @@ -269,7 +272,6 @@ async def update_peer( return honcho_peer await db.commit() - await db.refresh(honcho_peer) await peers_result.post_commit() cache_key = peer_cache_key(workspace_name, honcho_peer.name) diff --git a/src/crud/session.py b/src/crud/session.py index 6ce73db7..9580c16d 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -137,21 +137,30 @@ async def get_or_create_session( _retry: bool = False, ) -> GetOrCreateResult[models.Session]: """ - Get or create a session in a workspace with specified peers. - If the session already exists, the peers are added to the session. + Get an active session in a workspace or create it if it does not exist. + + If the session already exists, provided metadata replaces the current + metadata, provided configuration keys are merged into the existing + configuration, and any provided peers are ensured to be members of the + session. If the session does not exist, the workspace and peers are created + as needed before the session is created. Args: db: Database session - session: Session creation schema + session: Session creation payload, including optional metadata, + configuration, and session-peer configuration workspace_name: Name of the workspace - peer_names: List of peer names to add to the session - _retry: Whether to retry the operation + _retry: Whether to retry after a concurrent create conflict + Returns: GetOrCreateResult containing the session and whether it was created Raises: - ResourceNotFoundException: If the session does not exist and create is false - ConflictException: If we fail to get or create the session + ValueError: If session.name is empty + ResourceNotFoundException: If the named session exists but is inactive + ObserverException: If adding peers would exceed the observer limit + ConflictException: If concurrent creation prevents fetching or creating + the session """ if not session.name: @@ -247,10 +256,10 @@ async def get_or_create_session( workspace_name=workspace_name, session_name=session.name, peer_names=session.peer_names, + fetch_after_upsert=False, ) await db.commit() - await db.refresh(honcho_session) # Run deferred cache operations from workspace/peer creation if ws_result is not None: @@ -334,7 +343,11 @@ async def update_session( session_name: str, ) -> models.Session: """ - Update a session. + Get or create a session, then apply metadata and configuration updates. + + Provided metadata replaces the current metadata when present. Provided + configuration keys are merged into the existing configuration instead of + replacing it wholesale. Args: db: Database session @@ -346,7 +359,9 @@ async def update_session( The updated session Raises: - ResourceNotFoundException: If the session does not exist or peer is not in session + ResourceNotFoundException: If the named session exists but is inactive + ConflictException: If concurrent creation prevents fetching or creating + the session """ honcho_session: models.Session = ( await get_or_create_session( @@ -381,7 +396,6 @@ async def update_session( return honcho_session await db.commit() - await db.refresh(honcho_session) # Only invalidate if we actually updated cache_key = session_cache_key(workspace_name, session_name) @@ -729,7 +743,6 @@ async def clone_session( db.add(new_session_peer) await db.commit() - await db.refresh(new_session) logger.debug("Session %s cloned successfully", original_session_name) # Cache will be populated on next read - read-through pattern @@ -795,7 +808,13 @@ async def get_peers_from_session( # Get all active peers in the session (where left_at is NULL) return ( select(models.Peer) - .join(models.SessionPeer, models.Peer.name == models.SessionPeer.peer_name) + .join( + models.SessionPeer, + and_( + models.Peer.name == models.SessionPeer.peer_name, + models.Peer.workspace_name == models.SessionPeer.workspace_name, + ), + ) .where(models.SessionPeer.session_name == session_name) .where(models.Peer.workspace_name == workspace_name) .where(models.SessionPeer.left_at.is_(None)) # Only active peers @@ -825,7 +844,13 @@ async def get_session_peer_configuration( models.SessionPeer.configuration.label("session_peer_configuration"), (models.SessionPeer.left_at.is_(None)).label("is_active"), ) - .join(models.SessionPeer, models.Peer.name == models.SessionPeer.peer_name) + .join( + models.SessionPeer, + and_( + models.Peer.name == models.SessionPeer.peer_name, + models.Peer.workspace_name == models.SessionPeer.workspace_name, + ), + ) .where(models.SessionPeer.session_name == session_name) .where(models.Peer.workspace_name == workspace_name) .where(models.SessionPeer.workspace_name == workspace_name) @@ -912,24 +937,35 @@ async def _get_or_add_peers_to_session( workspace_name: str, session_name: str, peer_names: dict[str, schemas.SessionPeerConfig], + *, + fetch_after_upsert: bool = True, ) -> list[models.SessionPeer]: """ - Add multiple peers to an existing session. If a peer already exists in the session, - it will be skipped gracefully. + Upsert session-peer memberships for a session and optionally fetch the + active memberships afterward. + + New peers are inserted, peers that previously left the session are rejoined, + and already-active peers keep their existing session-level configuration. Args: db: Database session + workspace_name: Name of the workspace session_name: Name of the session - peer_names: Set of peer names to add to the session + peer_names: Mapping of peer names to session-level configuration + fetch_after_upsert: If True, query and return the active session peers + after the upsert. If False, skip that read and return an empty list. Returns: - List of all SessionPeer objects (both existing and newly created) + Active SessionPeer objects after the upsert, or an empty list when the + post-upsert fetch is skipped Raises: - ValueError: If adding peers would exceed the maximum limit + ObserverException: If adding peers would exceed the observer limit """ # If no peers to add, skip the insert and just return existing active session peers if not peer_names: + if not fetch_after_upsert: + return [] select_stmt = select(models.SessionPeer).where( models.SessionPeer.session_name == session_name, models.SessionPeer.workspace_name == workspace_name, @@ -994,6 +1030,9 @@ async def _get_or_add_peers_to_session( ) await db.execute(stmt) + if not fetch_after_upsert: + return [] + # Return all active session peers after the upsert select_stmt = select(models.SessionPeer).where( models.SessionPeer.session_name == session_name, diff --git a/src/crud/webhook.py b/src/crud/webhook.py index 7dc567eb..b607ed08 100644 --- a/src/crud/webhook.py +++ b/src/crud/webhook.py @@ -18,17 +18,20 @@ async def get_or_create_webhook_endpoint( webhook: schemas.WebhookEndpointCreate, ) -> GetOrCreateResult[schemas.WebhookEndpoint]: """ - Get or create a webhook endpoint, optionally for a workspace. + Get an existing webhook endpoint for a workspace or create it if missing. Args: db: Database session + workspace_name: Name of the workspace webhook: Webhook endpoint creation schema Returns: GetOrCreateResult containing the webhook endpoint and whether it was created Raises: - ResourceNotFoundException: If the workspace is specified and does not exist + ResourceNotFoundException: If the workspace does not exist + ValueError: If the workspace already has the maximum number of webhook + endpoints """ # Verify workspace exists await get_workspace(db, workspace_name=workspace_name) @@ -39,12 +42,6 @@ async def get_or_create_webhook_endpoint( result = await db.execute(stmt) endpoints = result.scalars().all() - # No more than WORKSPACE_LIMIT webhooks per workspace - if len(endpoints) >= settings.WEBHOOK.MAX_WORKSPACE_LIMIT: - raise ValueError( - f"Maximum number of webhook endpoints ({settings.WEBHOOK.MAX_WORKSPACE_LIMIT}) reached for this workspace." - ) - # Check if webhook already exists for this workspace for endpoint in endpoints: if endpoint.url == webhook.url: @@ -52,6 +49,12 @@ async def get_or_create_webhook_endpoint( schemas.WebhookEndpoint.model_validate(endpoint), created=False ) + # No more than WORKSPACE_LIMIT webhooks per workspace + if len(endpoints) >= settings.WEBHOOK.MAX_WORKSPACE_LIMIT: + raise ValueError( + f"Maximum number of webhook endpoints ({settings.WEBHOOK.MAX_WORKSPACE_LIMIT}) reached for this workspace." + ) + # Create new webhook endpoint webhook_endpoint = models.WebhookEndpoint( workspace_name=workspace_name, @@ -59,7 +62,6 @@ async def get_or_create_webhook_endpoint( ) db.add(webhook_endpoint) await db.commit() - await db.refresh(webhook_endpoint) logger.debug("Webhook endpoint created: %s", webhook.url) return GetOrCreateResult( diff --git a/src/crud/workspace.py b/src/crud/workspace.py index b59a99b1..3df2bb46 100644 --- a/src/crud/workspace.py +++ b/src/crud/workspace.py @@ -202,7 +202,11 @@ async def update_workspace( db: AsyncSession, workspace_name: str, workspace: schemas.WorkspaceUpdate ) -> models.Workspace: """ - Update a workspace. + Get or create a workspace, then apply metadata and configuration updates. + + Provided metadata replaces the current metadata when present. Provided + configuration keys are merged into the existing configuration instead of + replacing it wholesale. Args: db: Database session @@ -211,6 +215,10 @@ async def update_workspace( Returns: The updated workspace + + Raises: + ConflictException: If concurrent creation prevents fetching or creating + the workspace """ ws_result = await get_or_create_workspace( db, @@ -250,7 +258,6 @@ async def update_workspace( return honcho_workspace await db.commit() - await db.refresh(honcho_workspace) await ws_result.post_commit() # Only invalidate if we actually updated diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index fa2b9259..d4fd2a04 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -70,8 +70,7 @@ async def process_item(queue_item: models.QueueItem) -> None: queue_payload, ) raise ValueError(f"Invalid payload structure: {str(e)}") from e - async with tracked_db() as db: - await webhook_delivery.deliver_webhook(db, validated, workspace_name) + await webhook_delivery.deliver_webhook(validated, workspace_name) elif task_type == "summary": try: diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 5e885255..cda6decd 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -56,6 +56,48 @@ class WorkerOwnership(NamedTuple): aqs_id: str # The ID of the ActiveQueueSession that the worker is processing +def _detach_queue_batch_objects( + db: AsyncSession, + messages_context: list[models.Message], + items_to_process: list[QueueItem], +) -> None: + """Detach loaded batch objects so they remain usable after tracked_db exits.""" + seen: set[int] = set() + for obj in [*messages_context, *items_to_process]: + obj_id = id(obj) + if obj_id in seen: + continue + db.expunge(obj) + seen.add(obj_id) + + +def _resolve_batch_configuration( + items_to_process: list[QueueItem], +) -> tuple[list[QueueItem], ResolvedConfiguration | None]: + """Keep only the initial homogeneous configuration prefix for a batch.""" + if not items_to_process: + return [], None + + raw_config = items_to_process[0].payload.get("configuration") + resolved_config = ( + None if raw_config is None else ResolvedConfiguration.model_validate(raw_config) + ) + + valid_items: list[QueueItem] = [] + for item in items_to_process: + item_raw_config = item.payload.get("configuration") + item_config = ( + None + if item_raw_config is None + else ResolvedConfiguration.model_validate(item_raw_config) + ) + if item_config != resolved_config: + break + valid_items.append(item) + + return valid_items, resolved_config + + class QueueManager: def __init__(self): self.shutdown_event: asyncio.Event = asyncio.Event() @@ -608,21 +650,19 @@ class QueueManager: ) batch_max_tokens = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + parsed_key = parse_work_unit_key(work_unit_key) + messages_context: list[models.Message] = [] + items_to_process: list[QueueItem] = [] async with tracked_db("get_queue_item_batch") as db: # For batch tasks, get messages based on token limit. - # Step 1: Parse work_unit_key to get session context and focused sender - parsed_key = parse_work_unit_key(work_unit_key) - - # Verify worker still owns the work_unit_key + # Step 1: Verify worker still owns the work_unit_key. ownership_check = await db.execute( select(models.ActiveQueueSession.id) .where(models.ActiveQueueSession.work_unit_key == work_unit_key) .where(models.ActiveQueueSession.id == aqs_id) ) if not ownership_check.scalar_one_or_none(): - # Worker lost ownership, return empty - await db.commit() return [], [], None # Step 2: Build a single SQL query that: @@ -716,11 +756,8 @@ class QueueManager: result = await db.execute(query) rows = result.all() if not rows: - await db.commit() return [], [], None - messages_context: list[models.Message] = [] - items_to_process: list[QueueItem] = [] seen_messages: set[int] = set() for m, qi in rows: if m.id not in seen_messages: @@ -729,48 +766,21 @@ class QueueManager: if qi is not None: items_to_process.append(qi) - if items_to_process: - # Enforce homogeneous peer_card_config in the batch - # We stop collecting items as soon as we encounter a different configuration - payload = items_to_process[0].payload + _detach_queue_batch_objects(db, messages_context, items_to_process) - raw_config = payload.get("configuration") - if raw_config is None: - resolved_config = None - else: - resolved_config = ResolvedConfiguration.model_validate(raw_config) + items_to_process, resolved_config = _resolve_batch_configuration( + items_to_process + ) - valid_items: list[QueueItem] = [] - for item in items_to_process: - item_raw_config = item.payload.get("configuration") - if item_raw_config is None: - item_config = None - else: - item_config = ResolvedConfiguration.model_validate( - item_raw_config - ) - if item_config != resolved_config: - break - valid_items.append(item) - items_to_process = valid_items - else: - resolved_config = None + if items_to_process: + max_queue_item_message_id = max( + qi.message_id for qi in items_to_process if qi.message_id is not None + ) + messages_context = [ + m for m in messages_context if m.id <= max_queue_item_message_id + ] - if items_to_process: - max_queue_item_message_id = max( - [ - qi.message_id - for qi in items_to_process - if qi.message_id is not None - ] - ) - messages_context = [ # remove any messages that are after the last message_id from queue items - m for m in messages_context if m.id <= max_queue_item_message_id - ] - - await db.commit() - - return messages_context, items_to_process, resolved_config + return messages_context, items_to_process, resolved_config async def mark_queue_items_as_processed( self, items: list[QueueItem], work_unit_key: str diff --git a/src/routers/peers.py b/src/routers/peers.py index 01aa0f71..fb765737 100644 --- a/src/routers/peers.py +++ b/src/routers/peers.py @@ -446,11 +446,10 @@ async def search_peer( ..., description="Message search parameters. Use `limit` to control the number of results returned.", ), - db: AsyncSession = db, ): """Search a Peer's messages, optionally filtered by various criteria.""" # take user-provided filter and add workspace_id and peer_id to it filters = body.filters or {} filters["workspace_id"] = workspace_id filters["peer_id"] = peer_id - return await search(db, body.query, filters=filters, limit=body.limit) + return await search(body.query, filters=filters, limit=body.limit) diff --git a/src/routers/sessions.py b/src/routers/sessions.py index f68f93ff..9071aef1 100644 --- a/src/routers/sessions.py +++ b/src/routers/sessions.py @@ -794,7 +794,6 @@ async def search_session( body: schemas.MessageSearchOptions = Body( ..., description="Message search parameters" ), - db: AsyncSession = db, ): """ Search a Session with optional filters. Use `limit` to control the number of results returned. @@ -804,7 +803,6 @@ async def search_session( filters["workspace_id"] = workspace_id filters["session_id"] = session_id return await search( - db, body.query, filters=filters, limit=body.limit, diff --git a/src/routers/workspaces.py b/src/routers/workspaces.py index 3402c723..90530e92 100644 --- a/src/routers/workspaces.py +++ b/src/routers/workspaces.py @@ -142,7 +142,6 @@ async def search_workspace( body: schemas.MessageSearchOptions = Body( ..., description="Message search parameters" ), - db: AsyncSession = db, ): """ Search messages in a Workspace using optional filters. Use `limit` to control the number of @@ -151,7 +150,7 @@ async def search_workspace( # take user-provided filter and add workspace_id to it filters = body.filters or {} filters["workspace_id"] = workspace_id - return await search(db, body.query, filters=filters, limit=body.limit) + return await search(body.query, filters=filters, limit=body.limit) @router.get( diff --git a/src/utils/agent_tools.py b/src/utils/agent_tools.py index a6c3009d..21132397 100644 --- a/src/utils/agent_tools.py +++ b/src/utils/agent_tools.py @@ -854,6 +854,7 @@ async def get_observation_context( workspace_name: str, session_name: str | None, message_ids: list[str], + observer: str | None = None, ) -> list[models.Message]: """ Retrieve messages for given message IDs along with surrounding context. @@ -867,6 +868,8 @@ async def get_observation_context( workspace_name: Workspace identifier session_name: Session identifier (optional) message_ids: List of message IDs to retrieve + observer: When provided and session_name is None, scope results + to sessions this peer belongs to Returns: List of messages in chronological order, including the requested messages and surrounding context @@ -874,6 +877,17 @@ async def get_observation_context( if not message_ids: return [] + # Pre-fetch peer session scope if needed + allowed_session_names: list[str] | None = None + if observer and not session_name: + from src.crud.message import get_peer_session_names + + allowed_session_names = await get_peer_session_names( + db, workspace_name, observer + ) + if not allowed_session_names: + return [] + # Use a CTE to get seq_in_session values for target messages stmt = ( select(models.Message.seq_in_session) @@ -883,6 +897,8 @@ async def get_observation_context( if session_name: stmt = stmt.where(models.Message.session_name == session_name) + elif allowed_session_names is not None: + stmt = stmt.where(models.Message.session_name.in_(allowed_session_names)) target_seqs_cte = stmt.cte("target_seqs") @@ -905,6 +921,8 @@ async def get_observation_context( if session_name: stmt = stmt.where(models.Message.session_name == session_name) + elif allowed_session_names is not None: + stmt = stmt.where(models.Message.session_name.in_(allowed_session_names)) result = await db.execute(stmt) messages = list(result.scalars().all()) @@ -916,6 +934,7 @@ async def extract_preferences( workspace_name: str, session_name: str | None, observed: str, + observer: str | None = None, ) -> dict[str, list[str]]: """ Extract user preferences and standing instructions from conversation history. @@ -927,6 +946,8 @@ async def extract_preferences( workspace_name: Workspace identifier session_name: Session identifier (optional) observed: The peer whose preferences to extract + observer: When provided and session_name is None, scope results + to sessions this peer belongs to Returns: Dict with 'messages' list containing potentially relevant messages @@ -959,27 +980,26 @@ async def extract_preferences( for query in semantic_queries: try: - async with tracked_db("extract_preferences") as db: - snippets = await crud.search_messages( - db, - workspace_name=workspace_name, - session_name=session_name, - query=query, - limit=10, - context_window=0, - embedding=( - query_embeddings_by_query.get(query) - if query_embeddings_by_query is not None - else None - ), - ) - for matches, _ in snippets: - for msg in matches: - if msg.peer_name == observed: - content_key = msg.content[:100].lower() - if content_key not in seen_content: - seen_content.add(content_key) - messages.append(f"'{msg.content.strip()}'") + snippets = await crud.search_messages( + workspace_name=workspace_name, + session_name=session_name, + query=query, + limit=10, + context_window=0, + embedding=( + query_embeddings_by_query.get(query) + if query_embeddings_by_query is not None + else None + ), + observer=observer, + ) + for matches, _ in snippets: + for msg in matches: + if msg.peer_name == observed: + content_key = msg.content[:100].lower() + if content_key not in seen_content: + seen_content.add(content_key) + messages.append(f"'{msg.content.strip()}'") except Exception as e: logger.warning("Error in semantic search for '%s': %s", query, e) @@ -1265,20 +1285,19 @@ async def _handle_search_memory(ctx: ToolContext, tool_input: dict[str, Any]) -> if ctx.agent_type == "dialectic": limit = min(_safe_int(tool_input.get("top_k"), 20), 20) message_output = None - async with tracked_db("tool.search_memory.fallback") as db: - snippets = await crud.search_messages( - db, - workspace_name=ctx.workspace_name, - session_name=ctx.session_name, - query=query, - limit=limit, - context_window=0, - embedding=query_embedding, + snippets = await crud.search_messages( + workspace_name=ctx.workspace_name, + session_name=ctx.session_name, + query=query, + limit=limit, + context_window=0, + embedding=query_embedding, + observer=ctx.observer, + ) + if snippets: + message_output = _format_message_snippets( + snippets, f"for query '{query}'" ) - if snippets: - message_output = _format_message_snippets( - snippets, f"for query '{query}'" - ) if message_output: return ( f"No observations yet. Message search results:\n\n{message_output}" @@ -1302,6 +1321,7 @@ async def _handle_get_observation_context( workspace_name=ctx.workspace_name, session_name=ctx.session_name, message_ids=tool_input["message_ids"], + observer=ctx.observer, ) if not messages: return f"No messages found for IDs {tool_input['message_ids']}" @@ -1326,19 +1346,18 @@ async def _handle_search_messages(ctx: ToolContext, tool_input: dict[str, Any]) # Pre-compute embedding outside DB session to avoid holding a connection # during the external API call (same pattern as _handle_search_memory). query_embedding = await embedding_client.embed(query) - async with tracked_db("tool.search_messages") as db: - snippets = await crud.search_messages( - db, - workspace_name=ctx.workspace_name, - session_name=ctx.session_name, - query=query, - limit=limit, - context_window=2, - embedding=query_embedding, - ) - if not snippets: - return f"No messages found for query '{query}'" - formatted = _format_message_snippets(snippets, f"for query '{query}'") + snippets = await crud.search_messages( + workspace_name=ctx.workspace_name, + session_name=ctx.session_name, + query=query, + limit=limit, + context_window=2, + embedding=query_embedding, + observer=ctx.observer, + ) + if not snippets: + return f"No messages found for query '{query}'" + formatted = _format_message_snippets(snippets, f"for query '{query}'") return formatted @@ -1352,35 +1371,32 @@ async def _handle_grep_messages(ctx: ToolContext, tool_input: dict[str, Any]) -> _safe_int(tool_input.get("context_window"), 2), 2 ) # Cap context - async with tracked_db("tool.grep_messages") as db: - snippets = await crud.grep_messages( - db, - workspace_name=ctx.workspace_name, - session_name=ctx.session_name, - text=text, - limit=limit, - context_window=context_window, - ) - if not snippets: - return f"No messages found containing '{text}'" + snippets = await crud.grep_messages( + workspace_name=ctx.workspace_name, + session_name=ctx.session_name, + text=text, + limit=limit, + context_window=context_window, + observer=ctx.observer, + ) + if not snippets: + return f"No messages found containing '{text}'" - # Format with pattern-based snippet extraction - snippet_texts: list[str] = [] - total_matches = sum(len(matches) for matches, _ in snippets) - for i, (matches, context) in enumerate(snippets, 1): - lines: list[str] = [] - for msg in context: - truncated = _extract_pattern_snippet(msg.content, text) - lines.append( - format_new_turn_with_timestamp( - truncated, msg.created_at, msg.peer_name - ) - ) - sess = context[0].session_name if context else "unknown" - snippet_texts.append( - f"--- Snippet {i} (session: {sess}, {len(matches)} match(es)) ---\n" - + "\n".join(lines) + # Format with pattern-based snippet extraction + snippet_texts: list[str] = [] + total_matches = sum(len(matches) for matches, _ in snippets) + for i, (matches, context) in enumerate(snippets, 1): + lines: list[str] = [] + for msg in context: + truncated = _extract_pattern_snippet(msg.content, text) + lines.append( + format_new_turn_with_timestamp(truncated, msg.created_at, msg.peer_name) ) + sess = context[0].session_name if context else "unknown" + snippet_texts.append( + f"--- Snippet {i} (session: {sess}, {len(matches)} match(es)) ---\n" + + "\n".join(lines) + ) output = ( f"Found {total_matches} messages containing '{text}' in {len(snippets)} conversation snippets:\n\n" @@ -1425,6 +1441,7 @@ async def _handle_get_messages_by_date_range( before_date=before_date, limit=limit, order=order, + observer=ctx.observer, ) msg_count = len(messages) messages_text = ( @@ -1483,31 +1500,28 @@ async def _handle_search_messages_temporal( # Pre-compute embedding outside DB session to avoid holding a connection # during the external API call. query_embedding = await embedding_client.embed(query) - async with tracked_db("tool.search_messages_temporal") as db: - snippets = await crud.search_messages_temporal( - db, - workspace_name=ctx.workspace_name, - session_name=ctx.session_name, - query=query, - after_date=after_date, - before_date=before_date, - limit=limit, - context_window=context_window, - embedding=query_embedding, - ) - date_filter: list[str] = [] - if after_date_str: - date_filter.append(f"after {after_date_str}") - if before_date_str: - date_filter.append(f"before {before_date_str}") - filter_desc = f" ({' and '.join(date_filter)})" if date_filter else "" + snippets = await crud.search_messages_temporal( + workspace_name=ctx.workspace_name, + session_name=ctx.session_name, + query=query, + after_date=after_date, + before_date=before_date, + limit=limit, + context_window=context_window, + embedding=query_embedding, + observer=ctx.observer, + ) + date_filter: list[str] = [] + if after_date_str: + date_filter.append(f"after {after_date_str}") + if before_date_str: + date_filter.append(f"before {before_date_str}") + filter_desc = f" ({' and '.join(date_filter)})" if date_filter else "" - if not snippets: - return f"No messages found for query '{query}'{filter_desc}" + if not snippets: + return f"No messages found for query '{query}'{filter_desc}" - formatted = _format_message_snippets( - snippets, f"for query '{query}'{filter_desc}" - ) + formatted = _format_message_snippets(snippets, f"for query '{query}'{filter_desc}") return formatted @@ -1659,6 +1673,7 @@ async def _handle_extract_preferences( workspace_name=ctx.workspace_name, session_name=ctx.session_name, observed=ctx.observed, + observer=ctx.observer, ) messages = results.get("messages", []) diff --git a/src/utils/search.py b/src/utils/search.py index fcc77273..67a0d355 100644 --- a/src/utils/search.py +++ b/src/utils/search.py @@ -13,6 +13,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from src import models from src.config import settings +from src.dependencies import tracked_db from src.embedding_client import embedding_client from src.exceptions import ValidationException from src.models import session_peers_table @@ -23,6 +24,13 @@ from src.vector_store import get_external_vector_store T = TypeVar("T") +def _uses_pgvector_message_search() -> bool: + """Return True when semantic message search can stay entirely in Postgres.""" + return ( + settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED + ) + + def reciprocal_rank_fusion(*ranked_lists: list[T], k: int = 60, limit: int) -> list[T]: """ Combine multiple ranked lists using Reciprocal Rank Fusion (RRF). @@ -65,122 +73,115 @@ def reciprocal_rank_fusion(*ranked_lists: list[T], k: int = 60, limit: int) -> l return result[:limit] -async def _semantic_search( - db: AsyncSession, - query: str, +async def query_external_vector_message_ids( workspace_name: str, + embedding_query: list[float], limit: int, filters: dict[str, Any] | None = None, -) -> list[models.Message]: - """ - Perform semantic search using external vector store for message embeddings. - - Args: - db: Database session - query: Search query - workspace_name: Name of the workspace to search in - limit: Maximum number of results to return - filters: Optional filters to apply at vector store level (supports: session_id, peer_id) - - Returns: - list of messages ordered by semantic similarity - """ - try: - embedding_query = await embedding_client.embed(query) - except ValueError as e: - raise ValidationException( - f"Query exceeds maximum token limit of {settings.MAX_EMBEDDING_TOKENS}." - ) from e - - # Query Postgres / pgvector directly - if settings.EMBED_MESSAGES and ( - settings.VECTOR_STORE.TYPE == "pgvector" or not settings.VECTOR_STORE.MIGRATED - ): - # Join message_embeddings with messages to get full message objects - distance_expr = models.MessageEmbedding.embedding.cosine_distance( - embedding_query - ) - - stmt = ( - select(models.Message) - .join( - models.MessageEmbedding, - models.Message.public_id == models.MessageEmbedding.message_id, - ) - .where(models.MessageEmbedding.embedding.isnot(None)) - .where(models.MessageEmbedding.workspace_name == workspace_name) - ) - - # Apply all additional filters using the standard filter utility - # filters dict uses external names (session_id, peer_id) which apply_filter will map - # to internal column names (session_name, peer_name) - if filters: - # Create a copy with workspace added - internal_filters = filters.copy() - internal_filters["workspace_id"] = workspace_name - stmt = apply_filter(stmt, models.Message, internal_filters) - - # Order by cosine distance and limit - stmt = stmt.order_by(distance_expr).limit(limit) - - result = await db.execute(stmt) - return list(result.scalars().all()) - - # FALLBACK: Use external vector store (Turbopuffer, LanceDB) +) -> list[str]: + """Query the external vector store and return ordered message IDs.""" external_vector_store = get_external_vector_store() if external_vector_store is None: return [] namespace = external_vector_store.get_vector_namespace("message", workspace_name) - # Build vector store filters from the provided filters vector_filters: dict[str, Any] = {} if filters: - # Map external filter keys to vector store metadata keys if "session_id" in filters: vector_filters["session_name"] = filters["session_id"] if "peer_id" in filters: vector_filters["peer_name"] = filters["peer_id"] - # Query external vector store for similar message embeddings - # Since all filters are applied at the vector store level, we don't need to oversample + # Oversample: multiple chunk-level hits can map to the same message, + # so fetch extra to ensure enough unique messages after deduplication. vector_results = await external_vector_store.query( namespace, embedding_query, - top_k=limit, + top_k=limit * 3, filters=vector_filters if vector_filters else None, ) if not vector_results: return [] - # Extract message IDs from vector metadata - # Use dict to deduplicate while preserving order (dict keys maintain insertion order in Python 3.7+) seen_message_ids: dict[str, None] = {} - for result in vector_results: message_id = result.metadata.get("message_id") if message_id and message_id not in seen_message_ids: seen_message_ids[message_id] = None - message_ids = list(seen_message_ids.keys()) + return list(seen_message_ids.keys()) - # Fetch messages from database by the IDs from vector search and reapply filters - semantic_query = select(models.Message).where( - models.Message.public_id.in_(message_ids) - ) - semantic_query = apply_filter(semantic_query, models.Message, filters) - result = await db.execute(semantic_query) +async def fetch_messages_by_ids( + db: AsyncSession, + message_ids: list[str], + filters: dict[str, Any] | None = None, +) -> list[models.Message]: + """Fetch messages by ID and preserve the input ordering.""" + if not message_ids: + return [] + + stmt = select(models.Message).where(models.Message.public_id.in_(message_ids)) + stmt = apply_filter(stmt, models.Message, filters) + + result = await db.execute(stmt) messages = {msg.public_id: msg for msg in result.scalars().all()} - # Return messages in order of similarity (preserving vector store order) - ordered_messages: list[models.Message] = [] - for msg_id in message_ids: - if msg_id in messages: - ordered_messages.append(messages[msg_id]) + return [messages[msg_id] for msg_id in message_ids if msg_id in messages] - return ordered_messages + +async def _semantic_search_pgvector( + db: AsyncSession, + workspace_name: str, + embedding_query: list[float], + limit: int, + filters: dict[str, Any] | None = None, +) -> list[models.Message]: + """ + Perform semantic message search using pgvector in Postgres. + + Args: + db: Database session + workspace_name: Name of the workspace to search in + embedding_query: Pre-computed embedding for the search query + limit: Maximum number of results to return + filters: Optional filters to apply to the message query + + Returns: + list of messages ordered by semantic similarity + """ + distance_expr = models.MessageEmbedding.embedding.cosine_distance(embedding_query) + + stmt = ( + select(models.Message) + .join( + models.MessageEmbedding, + models.Message.public_id == models.MessageEmbedding.message_id, + ) + .where(models.MessageEmbedding.embedding.isnot(None)) + .where(models.MessageEmbedding.workspace_name == workspace_name) + ) + + if filters: + internal_filters = filters.copy() + internal_filters["workspace_id"] = workspace_name + stmt = apply_filter(stmt, models.Message, internal_filters) + + # Oversample because a message with multiple embedding chunks can + # produce duplicate rows; we deduplicate in Python to preserve HNSW + # index usage (a DISTINCT ON subquery would prevent the index scan). + stmt = stmt.order_by(distance_expr).limit(limit * 2) + + result = await db.execute(stmt) + seen: set[str] = set() + deduped: list[models.Message] = [] + for msg in result.scalars().all(): + if msg.public_id not in seen: + seen.add(msg.public_id) + deduped.append(msg) + return deduped[:limit] async def _filter_by_peer_perspective( @@ -308,7 +309,6 @@ async def _fulltext_search( async def search( - db: AsyncSession, query: str, *, filters: dict[str, Any] | None = None, @@ -321,7 +321,6 @@ async def search( are available, providing better search results than either method alone. Args: - db: Database session query: Search query to match against message content filters: Optional filters to scope search (must include workspace_id for semantic search). Special filter 'peer_perspective' will search across all messages from sessions that the peer is/was a member of, @@ -368,50 +367,81 @@ async def search( stmt = apply_filter(stmt, models.Message, filters) - search_results: list[list[models.Message]] = [] + workspace_name: str | None = None + if filters: + workspace_value = filters.get("workspace_id") or filters.get("workspace_name") + if isinstance(workspace_value, str): + workspace_name = workspace_value + + semantic_limit = limit * 4 if peer_perspective_name else limit * 2 + query_embedding: list[float] | None = None + semantic_message_ids: list[str] | None = None - # Perform semantic search if enabled and we have workspace context - # workspace_id is required for semantic search to determine the vector namespace - 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 (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, - workspace_name=workspace_name, - limit=semantic_limit, - filters=filters, - ) + try: + query_embedding = await embedding_client.embed(query) + except ValueError as e: + raise ValidationException( + f"Query exceeds maximum token limit of {settings.MAX_EMBEDDING_TOKENS}." + ) from e - # 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 + if not _uses_pgvector_message_search(): + semantic_message_ids = await query_external_vector_message_ids( + workspace_name=workspace_name, + embedding_query=query_embedding, + limit=semantic_limit, + filters=filters, ) - search_results.append(semantic_results) + async def _run_search(active_db: AsyncSession) -> list[models.Message]: + search_results: list[list[models.Message]] = [] - # Perform full-text search - # Get more results for fusion - fulltext_limit = limit * 2 - fulltext_results = await _fulltext_search( - db=db, query=query, stmt=stmt, limit=fulltext_limit - ) - search_results.append(fulltext_results) + if ( + settings.EMBED_MESSAGES + and isinstance(workspace_name, str) + and query_embedding is not None + ): + if _uses_pgvector_message_search(): + semantic_results = await _semantic_search_pgvector( + db=active_db, + workspace_name=workspace_name, + embedding_query=query_embedding, + limit=semantic_limit, + filters=filters, + ) + else: + semantic_results = await fetch_messages_by_ids( + db=active_db, + message_ids=semantic_message_ids or [], + filters=filters, + ) - # Combine results using RRF if we have multiple search methods - if len(search_results) > 1: - # Use RRF to combine semantic and full-text results - combined_results = reciprocal_rank_fusion(*search_results, limit=limit) - elif len(search_results) == 1: - # Single search method - apply limit directly - combined_results = search_results[0] - combined_results = combined_results[:limit] - else: - # No search results - combined_results = [] + if peer_perspective_name: + semantic_results = await _filter_by_peer_perspective( + active_db, + semantic_results, + workspace_name, + peer_perspective_name, + ) - return combined_results + search_results.append(semantic_results) + + fulltext_results = await _fulltext_search( + db=active_db, + query=query, + stmt=stmt, + limit=limit * 2, + ) + search_results.append(fulltext_results) + + if len(search_results) > 1: + return reciprocal_rank_fusion(*search_results, limit=limit) + if len(search_results) == 1: + return search_results[0][:limit] + return [] + + async with tracked_db("search.messages") as managed_db: + combined_results = await _run_search(managed_db) + for message in combined_results: + managed_db.expunge(message) + return combined_results diff --git a/src/webhooks/webhook_delivery.py b/src/webhooks/webhook_delivery.py index d26aa404..3d2df830 100644 --- a/src/webhooks/webhook_delivery.py +++ b/src/webhooks/webhook_delivery.py @@ -9,42 +9,41 @@ from sqlalchemy.ext.asyncio import AsyncSession from src.config import settings from src.crud.webhook import list_webhook_endpoints +from src.dependencies import tracked_db from src.utils.formatting import utc_now_iso from src.utils.queue_payload import WebhookPayload logger = logging.getLogger(__name__) -async def deliver_webhook( - db: AsyncSession, payload: WebhookPayload, workspace_name: str -) -> None: +async def deliver_webhook(payload: WebhookPayload, workspace_name: str) -> None: """ Deliver a single webhook event to its configured endpoints. """ - async with httpx.AsyncClient(timeout=30.0) as client: - try: + try: + async with tracked_db("webhook.deliver") as db: webhook_urls = await _get_webhook_urls(db, workspace_name) - if not webhook_urls: - logger.debug( - f"No webhook endpoints for workspace {workspace_name}, skipping." - ) - return - event_payload = { - "type": payload.event_type, - "data": payload.data, - "timestamp": utc_now_iso(), - } - event_json = json.dumps( - event_payload, separators=(",", ":"), sort_keys=True + if not webhook_urls: + logger.debug( + f"No webhook endpoints for workspace {workspace_name}, skipping." ) + return - try: - signature = _generate_webhook_signature(event_json) - except ValueError: - logger.exception("Failed to generate webhook signature") - return + event_payload = { + "type": payload.event_type, + "data": payload.data, + "timestamp": utc_now_iso(), + } + event_json = json.dumps(event_payload, separators=(",", ":"), sort_keys=True) + try: + signature = _generate_webhook_signature(event_json) + except ValueError: + logger.exception("Failed to generate webhook signature") + return + + async with httpx.AsyncClient(timeout=30.0) as client: tasks = [ client.post( url=url, @@ -73,10 +72,10 @@ async def deliver_webhook( f"Failed delivery for {payload.event_type} to {url}. Exception: {result}" ) - except httpx.RequestError: - logger.exception(f"Error sending webhook for {workspace_name}.") - except Exception: - logger.exception("Unexpected error delivering webhook.") + except httpx.RequestError: + logger.exception(f"Error sending webhook for {workspace_name}.") + except Exception: + logger.exception("Unexpected error delivering webhook.") async def _get_webhook_urls(db: AsyncSession, workspace_name: str) -> list[str]: diff --git a/tests/bench/runner_common.py b/tests/bench/runner_common.py index 093027df..0f9840e7 100644 --- a/tests/bench/runner_common.py +++ b/tests/bench/runner_common.py @@ -441,10 +441,6 @@ class BaseRunner(ABC, Generic[ResultT]): f"{self.get_metrics_prefix()}_{datetime.now().strftime('%Y%m%d_%H%M%S')}" ) self.logger: Logger = configure_logging() - # Semaphore for rate limiting concurrent item execution - self._concurrency_semaphore: asyncio.Semaphore | None = ( - asyncio.Semaphore(config.max_concurrent) if config.max_concurrent else None - ) # ------------------------------------------------------------------------- # Abstract methods - must be implemented by subclasses @@ -559,59 +555,64 @@ class BaseRunner(ABC, Generic[ResultT]): print(f"Limiting to {self.config.max_concurrent} concurrent item(s)") overall_start = time.time() - all_results: list[ResultT] = [] + all_results: list[ResultT | None] = [None] * len(items) - # Process in batches - batch_size = self.config.batch_size - for i in range(0, len(items), batch_size): - batch = items[i : i + batch_size] - batch_num = (i // batch_size) + 1 - total_batches = (len(items) + batch_size - 1) // batch_size + # Two-level concurrency: + # - inflight_sem limits how many items may be in the pipeline at once + # - active_sem limits how many items may actively hit Honcho at once + # Items release active_sem while waiting on queue polling so other work + # can progress, but inflight_sem prevents an unlimited thundering herd. + concurrency = self.config.max_concurrent or self.config.batch_size + inflight_sem = asyncio.Semaphore(concurrency) + active_sem = asyncio.Semaphore(concurrency) - print(f"\n{'=' * 60}") - print(f"Processing batch {batch_num}/{total_batches} ({len(batch)} items)") - print(f"{'=' * 60}") + async def _run_item(index: int, item: Any) -> None: + async with inflight_sem: + result = await self.execute_item( + item, + self._get_honcho_url(index), + active_sem=active_sem, + ) + all_results[index] = result - # Run items in batch concurrently (with optional rate limiting) - batch_results = await asyncio.gather( - *[ - self._execute_item_with_limit(item, self._get_honcho_url(i + idx)) - for idx, item in enumerate(batch) - ] - ) - - all_results.extend(batch_results) + tasks = [ + asyncio.create_task(_run_item(index, item)) + for index, item in enumerate(items) + ] + await asyncio.gather(*tasks) overall_duration = time.time() - overall_start # Finalize metrics self.metrics_collector.finalize_collection() - return all_results, overall_duration + missing_indexes = [ + index for index, result in enumerate(all_results) if result is None + ] + if missing_indexes: + raise RuntimeError( + f"Missing benchmark results for item indexes: {missing_indexes}" + ) - async def _execute_item_with_limit(self, item: Any, honcho_url: str) -> ResultT: - """Wrapper that applies concurrency limiting if configured.""" - if self._concurrency_semaphore: - async with self._concurrency_semaphore: - return await self.execute_item(item, honcho_url) - return await self.execute_item(item, honcho_url) + return [cast(ResultT, result) for result in all_results], overall_duration - async def execute_item(self, item: Any, honcho_url: str) -> ResultT: + async def execute_item( + self, + item: Any, + honcho_url: str, + active_sem: asyncio.Semaphore | None = None, + ) -> ResultT: """ Execute a single benchmark item. - This method orchestrates the standard flow: - 1. Create workspace and client - 2. Setup peers and session - 3. Ingest messages - 4. Wait for queue to empty - 5. Trigger dreams - 6. Execute questions - 7. Cleanup (if configured) + Active work (setup, ingest, dream scheduling, query execution) acquires + ``active_sem`` when provided. Idle queue polling releases that slot so + other items can continue making forward progress. Args: item: The item to process honcho_url: URL of the Honcho instance to use + active_sem: Optional semaphore limiting active I/O phases Returns: Result for this item @@ -635,21 +636,22 @@ class BaseRunner(ABC, Generic[ResultT]): start_time = time.time() try: - # Setup peers - await self.setup_peers(ctx, item) + # Setup peers/session and ingest under the active semaphore. + if active_sem: + await active_sem.acquire() + try: + await self.setup_peers(ctx, item) + await self.setup_session(ctx, item) - # Setup session - await self.setup_session(ctx, item) - - # Ingest messages - print(f"[{workspace_id}] Ingesting messages...") - message_count = await self.ingest_messages(ctx, item) - print(f"[{workspace_id}] Ingested {message_count} messages") + print(f"[{workspace_id}] Ingesting messages...") + message_count = await self.ingest_messages(ctx, item) + print(f"[{workspace_id}] Ingested {message_count} messages") + finally: + if active_sem: + active_sem.release() # Wait for deriver queue print(f"[{workspace_id}] Waiting for deriver queue to empty...") - await asyncio.sleep(1) # Give time for tasks to be queued - queue_empty = await self._wait_for_queue_empty(ctx.honcho_client) if not queue_empty: raise TimeoutError( @@ -670,20 +672,72 @@ class BaseRunner(ABC, Generic[ResultT]): + f"{len(dream_observers)} observer(s) across " + f"{len(dream_session_ids)} session(s)..." ) - for observer in dream_observers: - for dream_session_id in dream_session_ids: - success = await self._trigger_dream( - ctx.honcho_client, workspace_id, observer, dream_session_id - ) - if not success: + + if self.config.skip_dream: + print(f"[{workspace_id}] Skipping dreams (--skip-dream)") + else: + + async def _schedule_dream( + observer: str, + session_id: str, + ) -> bool: + try: + if active_sem: + await active_sem.acquire() + try: + await ctx.honcho_client.aio.schedule_dream( + observer=observer, + session=session_id, + observed=observer, + ) + finally: + if active_sem: + active_sem.release() print( - f"[{workspace_id}] Warning: Dream for {observer} in " - + f"session {dream_session_id} did not complete" + f"[{workspace_id}] Dream triggered for " + + f"{observer}/{observer} in {session_id}" ) + return True + except Exception as e: + print( + f"[{workspace_id}] ERROR: Dream trigger exception " + + f"for {observer} in {session_id}: {e}" + ) + return False + + dream_results = await asyncio.gather( + *[ + _schedule_dream(observer, dream_session_id) + for observer in dream_observers + for dream_session_id in dream_session_ids + ] + ) + + if all(dream_results): + success = await self._wait_for_queue_empty(ctx.honcho_client) + if success: + print(f"[{workspace_id}] All dreams completed") + else: + print(f"[{workspace_id}] Dreams timed out") + elif any(dream_results): + failed = [i for i, ok in enumerate(dream_results) if not ok] + print( + f"[{workspace_id}] Warning: {len(failed)} of " + + f"{len(dream_results)} dream schedules failed" + ) + await self._wait_for_queue_empty(ctx.honcho_client) + else: + print(f"[{workspace_id}] Warning: No dreams were scheduled") # Execute questions print(f"[{workspace_id}] Executing questions...") - result = await self.execute_questions(ctx, item) + if active_sem: + await active_sem.acquire() + try: + result = await self.execute_questions(ctx, item) + finally: + if active_sem: + active_sem.release() # Cleanup if self.config.cleanup_workspace: @@ -765,13 +819,15 @@ class BaseRunner(ABC, Generic[ResultT]): async def _wait_for_queue_empty( self, honcho_client: Honcho, session_id: str | None = None ) -> bool: - """Wait for the deriver queue to be empty.""" + """Wait for the deriver queue to be empty with exponential backoff.""" start_time = time.time() + delay = 0.2 while True: try: status = await honcho_client.aio.queue_status(session=session_id) except Exception: - await asyncio.sleep(1) + await asyncio.sleep(delay) + delay = min(delay * 1.5, 2.0) if time.time() - start_time >= self.config.timeout_seconds: return False continue @@ -781,7 +837,8 @@ class BaseRunner(ABC, Generic[ResultT]): if time.time() - start_time >= self.config.timeout_seconds: return False - await asyncio.sleep(1) + await asyncio.sleep(delay) + delay = min(delay * 1.5, 2.0) async def _trigger_dream( self, @@ -815,8 +872,6 @@ class BaseRunner(ABC, Generic[ResultT]): print(f"[{workspace_id}] Dream triggered for {observer}/{observed}") - # Wait for dream to complete - await asyncio.sleep(2) success = await self._wait_for_queue_empty(honcho_client) if success: print(f"[{workspace_id}] Dream for {observer} completed") diff --git a/tests/conftest.py b/tests/conftest.py index bd426e89..2ef7086b 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -752,8 +752,11 @@ def mock_tracked_db(db_engine: AsyncEngine, request: pytest.FixtureRequest): patch("src.dialectic.chat.tracked_db", mock_tracked_db_context), patch("src.utils.summarizer.tracked_db", mock_tracked_db_context), patch("src.webhooks.events.tracked_db", mock_tracked_db_context), + patch("src.webhooks.webhook_delivery.tracked_db", mock_tracked_db_context), patch("src.utils.agent_tools.tracked_db", mock_tracked_db_context), + patch("src.utils.search.tracked_db", mock_tracked_db_context), patch("src.crud.document.tracked_db", mock_tracked_db_context), + patch("src.crud.message.tracked_db", mock_tracked_db_context), patch("src.dialectic.core.tracked_db", mock_tracked_db_context), patch("src.dreamer.specialists.tracked_db", mock_tracked_db_context), patch("src.dreamer.surprisal.tracked_db", mock_tracked_db_context), diff --git a/tests/integration/test_message_embeddings.py b/tests/integration/test_message_embeddings.py index de544ee0..ef045049 100644 --- a/tests/integration/test_message_embeddings.py +++ b/tests/integration/test_message_embeddings.py @@ -4,6 +4,8 @@ Tests for message embedding functionality. These tests verify that message embeddings are created, stored, and can be searched. """ +from contextlib import asynccontextmanager +from datetime import datetime, timezone from typing import Any import pytest @@ -12,7 +14,9 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from src import models +from src.config import settings from src.crud import create_messages +from src.crud import message as message_crud from src.models import Peer, Workspace from src.schemas import MessageCreate from src.utils.search import search @@ -240,8 +244,7 @@ async def test_semantic_search_when_embeddings_enabled( initial_call_count: int = mock_openai_embeddings["embed"].call_count search_results = await search( - db=db_session, - query=search_query, + search_query, filters={ "workspace_id": test_workspace.name, "session_id": test_session.name, @@ -257,6 +260,212 @@ async def test_semantic_search_when_embeddings_enabled( assert created_message.public_id in found_message_ids +@pytest.mark.asyncio +async def test_search_messages_external_lookup_happens_before_tracked_db( + monkeypatch: pytest.MonkeyPatch, +): + """External semantic lookup should finish before opening tracked_db.""" + monkeypatch.setattr(settings.VECTOR_STORE, "MIGRATED", True) + monkeypatch.setattr(settings.VECTOR_STORE, "TYPE", "external") + + call_order: list[str] = [] + message = models.Message( + workspace_name="workspace", + session_name="session", + peer_name="peer", + content="Relevant external search result", + seq_in_session=1, + token_count=5, + created_at=datetime.now(timezone.utc), + ) + + class FakeDb: + def expunge(self, _obj: object) -> None: + call_order.append("expunge") + + fake_db = FakeDb() + + async def fake_search_messages_external( + workspace_name: str, + query_embedding: list[float], + limit: int, + *, + session_name: str | None = None, + allowed_session_names: list[str] | None = None, + after_date: datetime | None = None, + before_date: datetime | None = None, + ) -> list[str]: + _ = ( + workspace_name, + query_embedding, + limit, + session_name, + allowed_session_names, + after_date, + before_date, + ) + call_order.append("external") + return ["message-1"] + + async def fake_fetch_messages_by_ids( + db: FakeDb, + workspace_name: str, + message_ids: list[str], + *, + after_date: datetime | None = None, + before_date: datetime | None = None, + ) -> list[models.Message]: + _ = (workspace_name, message_ids, after_date, before_date) + assert db is fake_db + call_order.append("fetch") + return [message] + + async def fake_build_merged_snippets( + db: FakeDb, + workspace_name: str, + matched_messages: list[models.Message], + context_window: int, + ) -> list[tuple[list[models.Message], list[models.Message]]]: + _ = (workspace_name, context_window) + assert db is fake_db + assert matched_messages == [message] + call_order.append("build") + return [([message], [message])] + + @asynccontextmanager + async def fake_tracked_db(_operation_name: str | None = None): + call_order.append("enter") + yield fake_db + call_order.append("exit") + + monkeypatch.setattr( + message_crud, "_search_messages_external", fake_search_messages_external + ) + monkeypatch.setattr( + message_crud, "_fetch_messages_by_ids", fake_fetch_messages_by_ids + ) + monkeypatch.setattr( + message_crud, "_build_merged_snippets", fake_build_merged_snippets + ) + monkeypatch.setattr(message_crud, "tracked_db", fake_tracked_db) + + snippets = await message_crud.search_messages( + workspace_name="workspace", + session_name="session", + query="relevant query", + embedding=[0.1, 0.2, 0.3], + ) + + assert snippets == [([message], [message])] + assert call_order.index("external") < call_order.index("enter") + + +@pytest.mark.asyncio +async def test_search_messages_temporal_external_lookup_happens_before_tracked_db( + monkeypatch: pytest.MonkeyPatch, +): + """Temporal external semantic lookup should finish before opening tracked_db.""" + monkeypatch.setattr(settings.VECTOR_STORE, "MIGRATED", True) + monkeypatch.setattr(settings.VECTOR_STORE, "TYPE", "external") + + call_order: list[str] = [] + after_date = datetime(2024, 1, 1, tzinfo=timezone.utc) + before_date = datetime(2024, 12, 31, tzinfo=timezone.utc) + message = models.Message( + workspace_name="workspace", + session_name="session", + peer_name="peer", + content="Relevant temporal external search result", + seq_in_session=1, + token_count=5, + created_at=datetime.now(timezone.utc), + ) + + class FakeDb: + def expunge(self, _obj: object) -> None: + call_order.append("expunge") + + fake_db = FakeDb() + + async def fake_search_messages_external( + workspace_name: str, + query_embedding: list[float], + limit: int, + *, + session_name: str | None = None, + allowed_session_names: list[str] | None = None, + after_date: datetime | None = None, + before_date: datetime | None = None, + ) -> list[str]: + _ = ( + workspace_name, + query_embedding, + limit, + session_name, + allowed_session_names, + ) + assert after_date is not None + assert before_date is not None + call_order.append("external") + return ["message-1"] + + async def fake_fetch_messages_by_ids( + db: FakeDb, + workspace_name: str, + message_ids: list[str], + *, + after_date: datetime | None = None, + before_date: datetime | None = None, + ) -> list[models.Message]: + _ = (workspace_name, message_ids) + assert db is fake_db + assert after_date is not None + assert before_date is not None + call_order.append("fetch") + return [message] + + async def fake_build_merged_snippets( + db: FakeDb, + workspace_name: str, + matched_messages: list[models.Message], + context_window: int, + ) -> list[tuple[list[models.Message], list[models.Message]]]: + _ = (workspace_name, context_window) + assert db is fake_db + assert matched_messages == [message] + call_order.append("build") + return [([message], [message])] + + @asynccontextmanager + async def fake_tracked_db(_operation_name: str | None = None): + call_order.append("enter") + yield fake_db + call_order.append("exit") + + monkeypatch.setattr( + message_crud, "_search_messages_external", fake_search_messages_external + ) + monkeypatch.setattr( + message_crud, "_fetch_messages_by_ids", fake_fetch_messages_by_ids + ) + monkeypatch.setattr( + message_crud, "_build_merged_snippets", fake_build_merged_snippets + ) + monkeypatch.setattr(message_crud, "tracked_db", fake_tracked_db) + + snippets = await message_crud.search_messages_temporal( + workspace_name="workspace", + session_name="session", + query="relevant query", + after_date=after_date, + before_date=before_date, + embedding=[0.1, 0.2, 0.3], + ) + + assert snippets == [([message], [message])] + assert call_order.index("external") < call_order.index("enter") + + @pytest.mark.asyncio async def test_message_chunking_creates_multiple_embeddings( db_session: AsyncSession, diff --git a/tests/sdk_typescript/conftest.py b/tests/sdk_typescript/conftest.py index 15f2c98b..8abf6848 100644 --- a/tests/sdk_typescript/conftest.py +++ b/tests/sdk_typescript/conftest.py @@ -136,5 +136,8 @@ def mock_tracked_db(ts_db_session: async_sessionmaker[AsyncSession]): patch("src.dialectic.chat.tracked_db", ts_tracked_db), patch("src.utils.summarizer.tracked_db", ts_tracked_db), patch("src.webhooks.events.tracked_db", ts_tracked_db), + patch("src.webhooks.webhook_delivery.tracked_db", ts_tracked_db), + patch("src.utils.search.tracked_db", ts_tracked_db), + patch("src.crud.message.tracked_db", ts_tracked_db), ): yield diff --git a/tests/test_search.py b/tests/test_search.py index 2881cfd5..84f3ffa3 100644 --- a/tests/test_search.py +++ b/tests/test_search.py @@ -6,7 +6,7 @@ import pytest from nanoid import generate as generate_nanoid from sqlalchemy.ext.asyncio import AsyncSession -from src import models +from src import crud, models from src.utils.search import search @@ -62,11 +62,10 @@ async def test_peer_perspective_search_single_session( created_at=join_time + datetime.timedelta(seconds=2), ) db_session.add_all([msg1, msg2]) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -132,11 +131,10 @@ async def test_peer_perspective_search_multiple_sessions( created_at=join_time + datetime.timedelta(seconds=2), ) db_session.add_all([msg1, msg2]) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -212,11 +210,10 @@ async def test_peer_perspective_search_temporal_constraints( created_at=leave_time + datetime.timedelta(seconds=1), ) db_session.add_all([msg_before, msg_during, msg_after]) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -279,11 +276,10 @@ async def test_peer_perspective_search_active_member( created_at=join_time + datetime.timedelta(seconds=100), ) db_session.add_all([msg1, msg2]) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -339,11 +335,10 @@ async def test_peer_perspective_search_no_sessions( created_at=join_time + datetime.timedelta(seconds=1), ) db_session.add(msg) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter for peer1 (not in any sessions) results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -408,11 +403,10 @@ async def test_peer_perspective_search_boundary_timestamps( created_at=leave_time, # Exact leave time ) db_session.add_all([msg_at_join, msg_at_leave]) - await db_session.flush() + await db_session.commit() # Search with peer_perspective filter results = await search( - db_session, "Message", filters={"peer_perspective": peer1.name, "workspace_id": workspace.name}, limit=10, @@ -422,3 +416,291 @@ async def test_peer_perspective_search_boundary_timestamps( assert len(results) == 2 assert msg_at_join.public_id in [m.public_id for m in results] assert msg_at_leave.public_id in [m.public_id for m in results] + + +# ============================================================================= +# Tests for observer scoping in CRUD message functions +# ============================================================================= + + +async def _setup_multi_session_workspace(db_session: AsyncSession): + """Helper: create workspace with 2 sessions, 2 peers. peer1 only in session1.""" + workspace = models.Workspace(name=generate_nanoid()) + db_session.add(workspace) + await db_session.flush() + + peer1 = models.Peer(name="observer", workspace_name=workspace.name) + peer2 = models.Peer(name="other", workspace_name=workspace.name) + db_session.add_all([peer1, peer2]) + await db_session.flush() + + session1 = models.Session(name="session_visible", workspace_name=workspace.name) + session2 = models.Session(name="session_hidden", workspace_name=workspace.name) + db_session.add_all([session1, session2]) + await db_session.flush() + + join_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta( + minutes=10 + ) + + # peer1 is only in session1 + await db_session.execute( + models.session_peers_table.insert().values( + workspace_name=workspace.name, + session_name=session1.name, + peer_name=peer1.name, + joined_at=join_time, + left_at=None, + ) + ) + # peer2 is in both sessions + for s in [session1, session2]: + await db_session.execute( + models.session_peers_table.insert().values( + workspace_name=workspace.name, + session_name=s.name, + peer_name=peer2.name, + joined_at=join_time, + left_at=None, + ) + ) + await db_session.flush() + + msg_visible = models.Message( + content="visible message with keyword", + session_name=session1.name, + peer_name=peer2.name, + workspace_name=workspace.name, + seq_in_session=1, + created_at=join_time + datetime.timedelta(seconds=1), + ) + msg_hidden = models.Message( + content="hidden message with keyword", + session_name=session2.name, + peer_name=peer2.name, + workspace_name=workspace.name, + seq_in_session=1, + created_at=join_time + datetime.timedelta(seconds=2), + ) + db_session.add_all([msg_visible, msg_hidden]) + await db_session.commit() + + return workspace, peer1, peer2, session1, session2, msg_visible, msg_hidden + + +@pytest.mark.asyncio +async def test_grep_messages_observer_scoping_excludes_non_member_sessions( + db_session: AsyncSession, +): + """grep_messages with observer excludes messages from sessions the observer isn't in.""" + ( + workspace, + peer1, + _, + _, + _, + msg_visible, + msg_hidden, + ) = await _setup_multi_session_workspace(db_session) + + # Without scoping: both messages found + results_unscoped = await crud.grep_messages( + workspace_name=workspace.name, + session_name=None, + text="keyword", + ) + all_matched_ids = [m.public_id for matches, _ in results_unscoped for m in matches] + assert msg_visible.public_id in all_matched_ids + assert msg_hidden.public_id in all_matched_ids + + # With observer scoping: only visible message found + results_scoped = await crud.grep_messages( + workspace_name=workspace.name, + session_name=None, + text="keyword", + observer=peer1.name, + ) + scoped_ids = [m.public_id for matches, _ in results_scoped for m in matches] + assert msg_visible.public_id in scoped_ids + assert msg_hidden.public_id not in scoped_ids + + +@pytest.mark.asyncio +async def test_get_messages_by_date_range_observer_scoping( + db_session: AsyncSession, +): + """get_messages_by_date_range with observer excludes non-member sessions.""" + ( + workspace, + peer1, + _, + _, + _, + msg_visible, + msg_hidden, + ) = await _setup_multi_session_workspace(db_session) + + # Without scoping + results_unscoped = await crud.get_messages_by_date_range( + db_session, + workspace_name=workspace.name, + session_name=None, + ) + unscoped_ids = [m.public_id for m in results_unscoped] + assert msg_visible.public_id in unscoped_ids + assert msg_hidden.public_id in unscoped_ids + + # With observer scoping + results_scoped = await crud.get_messages_by_date_range( + db_session, + workspace_name=workspace.name, + session_name=None, + observer=peer1.name, + ) + scoped_ids = [m.public_id for m in results_scoped] + assert msg_visible.public_id in scoped_ids + assert msg_hidden.public_id not in scoped_ids + + +@pytest.mark.asyncio +async def test_grep_messages_observer_scoping_noop_when_session_provided( + db_session: AsyncSession, +): + """When session_name is provided, observer is ignored.""" + ( + workspace, + peer1, + _, + _, + session_hidden, + _, + msg_hidden, + ) = await _setup_multi_session_workspace(db_session) + + results = await crud.grep_messages( + workspace_name=workspace.name, + session_name=session_hidden.name, + text="keyword", + observer=peer1.name, + ) + matched_ids = [m.public_id for matches, _ in results for m in matches] + assert msg_hidden.public_id in matched_ids + + +@pytest.mark.asyncio +async def test_grep_messages_observer_scoping_empty_when_no_sessions( + db_session: AsyncSession, +): + """Observer not in any sessions returns empty results.""" + workspace = models.Workspace(name=generate_nanoid()) + db_session.add(workspace) + await db_session.flush() + + loner = models.Peer(name="loner", workspace_name=workspace.name) + other = models.Peer(name="other", workspace_name=workspace.name) + db_session.add_all([loner, other]) + await db_session.flush() + + session = models.Session(name="s1", workspace_name=workspace.name) + db_session.add(session) + await db_session.flush() + + await db_session.execute( + models.session_peers_table.insert().values( + workspace_name=workspace.name, + session_name=session.name, + peer_name=other.name, + joined_at=datetime.datetime.now(datetime.timezone.utc), + left_at=None, + ) + ) + await db_session.flush() + + msg = models.Message( + content="some keyword content", + session_name=session.name, + peer_name=other.name, + workspace_name=workspace.name, + seq_in_session=1, + created_at=datetime.datetime.now(datetime.timezone.utc), + ) + db_session.add(msg) + await db_session.commit() + + results = await crud.grep_messages( + workspace_name=workspace.name, + session_name=None, + text="keyword", + observer=loner.name, + ) + assert results == [] + + +@pytest.mark.asyncio +async def test_grep_messages_observer_scoping_left_session_still_visible( + db_session: AsyncSession, +): + """Observer who left a session still sees all messages in that session. + + Any membership record (regardless of left_at) grants full session visibility. + """ + workspace = models.Workspace(name=generate_nanoid()) + db_session.add(workspace) + await db_session.flush() + + observer = models.Peer(name="obs", workspace_name=workspace.name) + other = models.Peer(name="other", workspace_name=workspace.name) + db_session.add_all([observer, other]) + await db_session.flush() + + session = models.Session(name="s1", workspace_name=workspace.name) + db_session.add(session) + await db_session.flush() + + base_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta( + minutes=10 + ) + join_time = base_time + leave_time = base_time + datetime.timedelta(minutes=5) + + await db_session.execute( + models.session_peers_table.insert().values( + workspace_name=workspace.name, + session_name=session.name, + peer_name=observer.name, + joined_at=join_time, + left_at=leave_time, + ) + ) + await db_session.flush() + + # Message during membership + msg_during = models.Message( + content="keyword during", + session_name=session.name, + peer_name=other.name, + workspace_name=workspace.name, + seq_in_session=1, + created_at=join_time + datetime.timedelta(minutes=2), + ) + # Message after observer left — still visible because any membership grants full access + msg_after = models.Message( + content="keyword after", + session_name=session.name, + peer_name=other.name, + workspace_name=workspace.name, + seq_in_session=2, + created_at=leave_time + datetime.timedelta(minutes=1), + ) + db_session.add_all([msg_during, msg_after]) + await db_session.commit() + + results = await crud.grep_messages( + workspace_name=workspace.name, + session_name=None, + text="keyword", + observer=observer.name, + ) + matched_ids = [m.public_id for matches, _ in results for m in matches] + assert msg_during.public_id in matched_ids + assert msg_after.public_id in matched_ids diff --git a/tests/utils/test_agent_tools.py b/tests/utils/test_agent_tools.py index 1387f681..bb0ff900 100644 --- a/tests/utils/test_agent_tools.py +++ b/tests/utils/test_agent_tools.py @@ -29,6 +29,7 @@ from src.utils.agent_tools import ( _handle_grep_messages, # pyright: ignore[reportPrivateUsage] _handle_search_memory, # pyright: ignore[reportPrivateUsage] _handle_search_messages, # pyright: ignore[reportPrivateUsage] + _handle_search_messages_temporal, # pyright: ignore[reportPrivateUsage] _handle_update_peer_card, # pyright: ignore[reportPrivateUsage] create_observations, create_tool_executor, @@ -528,15 +529,15 @@ class TestSearchMemory: return [] async def fake_search_messages( - db: AsyncSession, workspace_name: str, session_name: str | None, query: str, limit: int = 10, context_window: int = 2, embedding: list[float] | None = None, + observer: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: - _ = (db, workspace_name, session_name, query, limit, context_window) + _ = (workspace_name, session_name, query, limit, context_window, observer) fallback_embeddings.append(embedding) msg = models.Message( workspace_name=ctx.workspace_name, @@ -610,6 +611,78 @@ class TestGrepMessages: assert "ERROR" in result +@pytest.mark.asyncio +class TestSearchMessagesTemporal: + """Tests for _handle_search_messages_temporal.""" + + async def test_reuses_precomputed_embedding( + self, + make_tool_context: Callable[..., ToolContext], + monkeypatch: pytest.MonkeyPatch, + ): + """Embeds once and forwards the precomputed embedding to CRUD search.""" + ctx = make_tool_context() + + embed_calls: list[str] = [] + forwarded_embeddings: list[list[float] | None] = [] + + async def fake_embed(query: str) -> list[float]: + embed_calls.append(query) + return [0.9, 0.1, 0.3] + + async def fake_search_messages_temporal( + workspace_name: str, + session_name: str | None, + query: str, + after_date: datetime | None = None, + before_date: datetime | None = None, + limit: int = 10, + context_window: int = 2, + embedding: list[float] | None = None, + observer: str | None = None, + ) -> list[tuple[list[models.Message], list[models.Message]]]: + _ = ( + workspace_name, + session_name, + query, + after_date, + before_date, + limit, + context_window, + observer, + ) + forwarded_embeddings.append(embedding) + msg = models.Message( + workspace_name=ctx.workspace_name, + session_name=ctx.session_name, + peer_name=ctx.observed, + content="Relevant temporal fallback message", + seq_in_session=1, + token_count=5, + created_at=datetime.now(timezone.utc), + ) + return [([msg], [msg])] + + monkeypatch.setattr("src.utils.agent_tools.embedding_client.embed", fake_embed) + monkeypatch.setattr( + "src.utils.agent_tools.crud.search_messages_temporal", + fake_search_messages_temporal, + ) + + result = await _handle_search_messages_temporal( + ctx, + { + "query": "when did this happen", + "after_date": "2024-01-01", + "before_date": "2024-12-31", + }, + ) + + assert "Found" in result + assert embed_calls == ["when did this happen"] + assert forwarded_embeddings == [[0.9, 0.1, 0.3]] + + @pytest.mark.asyncio class TestGetMessagesByDateRange: """Tests for _handle_get_messages_by_date_range.""" @@ -991,15 +1064,15 @@ class TestExtractPreferences: embedding_args: list[list[float] | None] = [] async def fake_search_messages( - _db: AsyncSession, workspace_name: str, session_name: str | None, query: str, limit: int, context_window: int, embedding: list[float] | None, + observer: str | None = None, ) -> list[tuple[list[models.Message], list[models.Message]]]: - _ = (limit, context_window) + _ = (limit, context_window, observer) embedding_args.append(embedding) msg = models.Message( workspace_name=workspace_name, @@ -1292,3 +1365,55 @@ class TestObservationLockRegistry: # All 100 entries should be cleaned up remaining = sum(1 for k in _observation_locks if k[0].startswith("ws_growth_")) assert remaining == 0 + + +@pytest.mark.asyncio +class TestObserverPeerNameWiring: + """Tests that tool handlers pass observer to CRUD functions.""" + + async def test_grep_messages_passes_observer( + self, + make_tool_context: Callable[..., ToolContext], + monkeypatch: pytest.MonkeyPatch, + ): + """_handle_grep_messages passes ctx.observer as observer.""" + ctx = make_tool_context() + captured_kwargs: dict[str, Any] = {} + + async def fake_grep_messages( + **kwargs: Any, + ) -> list[tuple[list[models.Message], list[models.Message]]]: + captured_kwargs.update(kwargs) + return [] + + monkeypatch.setattr( + "src.utils.agent_tools.crud.grep_messages", fake_grep_messages + ) + + await _handle_grep_messages(ctx, {"text": "hello"}) + + assert captured_kwargs["observer"] == ctx.observer + + async def test_get_messages_by_date_range_passes_observer( + self, + make_tool_context: Callable[..., ToolContext], + monkeypatch: pytest.MonkeyPatch, + ): + """_handle_get_messages_by_date_range passes ctx.observer as observer.""" + ctx = make_tool_context() + captured_kwargs: dict[str, Any] = {} + + async def fake_get_messages_by_date_range( + _db: Any, **kwargs: Any + ) -> list[models.Message]: + captured_kwargs.update(kwargs) + return [] + + monkeypatch.setattr( + "src.utils.agent_tools.crud.get_messages_by_date_range", + fake_get_messages_by_date_range, + ) + + await _handle_get_messages_by_date_range(ctx, {"after_date": "2024-01-01"}) + + assert captured_kwargs["observer"] == ctx.observer diff --git a/tests/webhooks/test_webhook_delivery.py b/tests/webhooks/test_webhook_delivery.py index 3c4235f2..56198f55 100644 --- a/tests/webhooks/test_webhook_delivery.py +++ b/tests/webhooks/test_webhook_delivery.py @@ -121,7 +121,7 @@ async def test_deliver_webhook_skips_when_no_urls( ) payload = WebhookPayload(event_type="peer.created", data={"id": "p_123"}) - await webhook_delivery.deliver_webhook(AsyncMock(), payload, "workspace-a") + await webhook_delivery.deliver_webhook(payload, "workspace-a") assert fake_client.calls == [] @@ -162,7 +162,7 @@ async def test_deliver_webhook_posts_signed_payload_to_each_endpoint( event_type="message.created", data={"id": "m_1", "workspace": "workspace-a"}, ) - await webhook_delivery.deliver_webhook(AsyncMock(), payload, "workspace-a") + await webhook_delivery.deliver_webhook(payload, "workspace-a") expected_event_json = json.dumps( { @@ -210,7 +210,7 @@ async def test_deliver_webhook_handles_signature_generation_failure( monkeypatch.setattr(httpx, "AsyncClient", async_client_factory) payload = WebhookPayload(event_type="workspace.updated", data={"id": "ws_1"}) - await webhook_delivery.deliver_webhook(AsyncMock(), payload, "workspace-a") + await webhook_delivery.deliver_webhook(payload, "workspace-a") assert fake_client.calls == [] @@ -233,4 +233,4 @@ async def test_deliver_webhook_catches_request_errors( monkeypatch.setattr(httpx, "AsyncClient", async_client_factory) payload = WebhookPayload(event_type="workspace.updated", data={"id": "ws_1"}) - await webhook_delivery.deliver_webhook(AsyncMock(), payload, "workspace-a") + await webhook_delivery.deliver_webhook(payload, "workspace-a")