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:
parent
ff116b0601
commit
5b6bd59030
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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", [])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Reference in New Issue