Tighten Transaction Scopes (#525)

* fix: further remove extraneous transactions

* fix: (search) use 2 phase function to reduce un-needed transaction

* fix: refactor agent search to perform external operations before making a transaction

* fix: reduce scope of queue manager transaction

* fix: (bench) add concurrency to test bench

* fix: address review findings for search dedup, webhook idempotency, and bench throttling

* Fix Leakage in non-session-scoped chat call (#526)

* fix: (search) reduce scope for peer based searches

* fix: tests

* fix: (test) address coderabbit comment

* fix: drop db param from deliver_webhook

---------

Co-authored-by: Rajat Ahuja <rahuja445@gmail.com>
This commit is contained in:
Vineeth Voruganti 2026-04-08 11:14:50 -04:00 committed by GitHub
parent ff116b0601
commit 5b6bd59030
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
20 changed files with 1468 additions and 535 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@ -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", [])

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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