2.3.0 Release & N+1 Query Optimization (#190)
* chore (docs): Initial Changelog for 2.3.0 * optimize batch enqueue calculation * fix (embed): Truncate context during _query_documents_for_level * chore (docs): Update Changelog with patches * fix: Code Rabbit Suggestions * fix: Code Rabbit Suggestions
This commit is contained in:
parent
6efb950ef5
commit
f6e01a6f72
16
CHANGELOG.md
16
CHANGELOG.md
|
|
@ -10,16 +10,30 @@ and this project adheres to [Semantic Versioning](http://semver.org/).
|
|||
### Added
|
||||
|
||||
- `getSummaries` endpoint to get all available summaries for a session directly
|
||||
- Peer Card feature to improve context for deriver and dialectic
|
||||
|
||||
### Changed
|
||||
|
||||
- Session Peer limit to be based on observers instead, renamed config value to
|
||||
`SESSION_OBSERVERS_LIMIT`
|
||||
- Deriver uses `get_context` internally to prevent context window limit errors
|
||||
- `Messages` can take a custom timestamp for the `created_at` field, defaulting
|
||||
to the current time
|
||||
- `get_context` endpoint returns detailed `Summary` object rather than just
|
||||
summary content
|
||||
- Working representations use a FIFO queue structure to maintain facts rather
|
||||
than a full rewrite
|
||||
- Optimized deriver enqueue by prefetching message sequence numbers (eliminates N+1 queries)
|
||||
|
||||
### Fixed
|
||||
|
||||
- Deriver uses `get_context` internally to prevent context window limit errors
|
||||
- Embedding store will truncate context when querying documents to prevent embedding
|
||||
token limit errors
|
||||
- Queue manager to schedule work based on available works rather than total
|
||||
number of workers
|
||||
- Queue manager to use atomic db transactions rather than long lived transaction
|
||||
for the worker lifecycle
|
||||
- Timestamp formats unified to ISO 8601 across the codebase
|
||||
|
||||
## [2.2.0] — 2025-08-07
|
||||
|
||||
|
|
|
|||
|
|
@ -31,18 +31,31 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
|
|||
### Added
|
||||
|
||||
- `getSummaries` endpoint to get all available summaries for a session directly
|
||||
- Peer Card feature to improve context for deriver and dialectic
|
||||
|
||||
### Changed
|
||||
|
||||
- Session Peer limit to be based on observers instead, renamed config value to
|
||||
`SESSION_OBSERVERS_LIMIT`
|
||||
- Deriver uses `get_context` internally to prevent context window limit errors
|
||||
- `Messages` can take a custom timestamp for the `created_at` field, defaulting
|
||||
to the current time
|
||||
- `get_context` endpoint returns detailed `Summary` object rather than just
|
||||
summary content
|
||||
</Update>
|
||||
- Working representations use a FIFO queue structure to maintain facts rather
|
||||
than a full rewrite
|
||||
- Optimized deriver enqueue by prefetching message sequence numbers (eliminates N+1 queries)
|
||||
|
||||
### Fixed
|
||||
|
||||
- Deriver uses `get_context` internally to prevent context window limit errors
|
||||
- Embedding store will truncate context when querying documents to prevent embedding
|
||||
token limit errors
|
||||
- Queue manager to schedule work based on available works rather than total
|
||||
number of workers
|
||||
- Queue manager to use atomic db transactions rather than long lived transaction
|
||||
for the worker lifecycle
|
||||
- Timestamp formats unified to ISO 8601 across the codebase
|
||||
</Update>
|
||||
<Update label="v2.2.0">
|
||||
### Added
|
||||
|
||||
|
|
@ -259,6 +272,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
|
|||
### Added
|
||||
|
||||
- getSummaries API returning structured summaries
|
||||
- Webhook support
|
||||
|
||||
### Changed
|
||||
|
||||
|
|
@ -305,6 +319,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
|
|||
### Added
|
||||
|
||||
- getSummaries API returning structured summaries
|
||||
- Webhook support
|
||||
|
||||
### Changed
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ and this project adheres to [Semantic Versioning](http://semver.org/).
|
|||
### Added
|
||||
|
||||
- getSummaries API returning structured summaries
|
||||
- Webhook support
|
||||
|
||||
### Changed
|
||||
|
||||
|
|
@ -26,6 +27,7 @@ and this project adheres to [Semantic Versioning](http://semver.org/).
|
|||
|
||||
- Summaries are now included in `toOpenAI` and `toAnthropic` functions
|
||||
- `SessionContext.__len__` now counts the summary in its total
|
||||
- `filter` keyword changed to `filters`
|
||||
|
||||
### Fixed
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ and this project adheres to [Semantic Versioning](http://semver.org/).
|
|||
### Added
|
||||
|
||||
- getSummaries API returning structured summaries
|
||||
- Webhook support
|
||||
|
||||
### Changed
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from .message import (
|
|||
create_messages,
|
||||
get_message,
|
||||
get_message_seq_in_session,
|
||||
get_message_seqs_in_session_batch,
|
||||
get_messages,
|
||||
get_messages_id_range,
|
||||
update_message,
|
||||
|
|
@ -61,6 +62,7 @@ __all__ = [
|
|||
"get_messages_id_range",
|
||||
"get_message",
|
||||
"get_message_seq_in_session",
|
||||
"get_message_seqs_in_session_batch",
|
||||
"update_message",
|
||||
# Peer
|
||||
"get_or_create_peers",
|
||||
|
|
|
|||
|
|
@ -276,6 +276,53 @@ async def get_message_seq_in_session(
|
|||
return count + 1
|
||||
|
||||
|
||||
async def get_message_seqs_in_session_batch(
|
||||
db: AsyncSession,
|
||||
workspace_name: str,
|
||||
session_name: str,
|
||||
message_ids: list[int],
|
||||
) -> dict[int, int]:
|
||||
"""
|
||||
Get the sequence numbers for multiple messages within a session in a single query.
|
||||
|
||||
Args:
|
||||
db: Database session
|
||||
workspace_name: Name of the workspace
|
||||
session_name: Name of the session
|
||||
message_ids: List of message primary key IDs
|
||||
|
||||
Returns:
|
||||
Dictionary mapping message_id to sequence number (1-indexed).
|
||||
If a given ID does not exist in the specified session, its value will be 0.
|
||||
Note: duplicate IDs in the input are de-duplicated in the query and the
|
||||
resulting mapping will contain a single entry per unique message_id.
|
||||
"""
|
||||
if not message_ids:
|
||||
return {}
|
||||
|
||||
unique_ids = list(set(message_ids))
|
||||
|
||||
# Rank all messages in the session, then select only the ones we care about
|
||||
ranked = (
|
||||
select(
|
||||
models.Message.id.label("id"),
|
||||
func.row_number().over(order_by=models.Message.id).label("seq"),
|
||||
)
|
||||
.where(
|
||||
models.Message.workspace_name == workspace_name,
|
||||
models.Message.session_name == session_name,
|
||||
)
|
||||
.subquery()
|
||||
)
|
||||
stmt = select(ranked.c.id, ranked.c.seq).where(ranked.c.id.in_(unique_ids))
|
||||
result = await db.execute(stmt)
|
||||
rows = result.all()
|
||||
id_to_position = {row[0]: int(row[1]) for row in rows}
|
||||
|
||||
# Return positions for requested IDs; 0 for any missing/out-of-session IDs
|
||||
return {msg_id: id_to_position.get(msg_id, 0) for msg_id in message_ids}
|
||||
|
||||
|
||||
async def get_message(
|
||||
db: AsyncSession,
|
||||
workspace_name: str,
|
||||
|
|
|
|||
|
|
@ -83,6 +83,12 @@ async def handle_session(
|
|||
db_session, workspace_name, session_name
|
||||
)
|
||||
|
||||
# Get all message IDs to fetch sequences in batch
|
||||
message_ids = [msg["message_id"] for msg in payload]
|
||||
message_seq_map = await crud.get_message_seqs_in_session_batch(
|
||||
db_session, workspace_name, session_name, message_ids
|
||||
)
|
||||
|
||||
queue_records: list[dict[str, Any]] = []
|
||||
|
||||
for message in payload:
|
||||
|
|
@ -93,6 +99,7 @@ async def handle_session(
|
|||
peers_with_configuration,
|
||||
session.id,
|
||||
deriver_disabled=deriver_disabled,
|
||||
message_seq_map=message_seq_map,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -235,6 +242,7 @@ async def generate_queue_records(
|
|||
session_id: str,
|
||||
*,
|
||||
deriver_disabled: bool,
|
||||
message_seq_map: dict[int, int] | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
Process a single message and generate queue records based on configurations.
|
||||
|
|
@ -245,18 +253,24 @@ async def generate_queue_records(
|
|||
deriver_disabled: Whether deriver is disabled for the session
|
||||
peers_with_configuration: Dictionary of peer configurations
|
||||
session_id: Session ID
|
||||
message_seq_map: Optional pre-fetched mapping of message_id to sequence number
|
||||
|
||||
Returns:
|
||||
List of queue records for this message
|
||||
"""
|
||||
sender_name = message["peer_name"]
|
||||
message_id: int = message["message_id"]
|
||||
message_seq_in_session: int = await crud.get_message_seq_in_session(
|
||||
db_session,
|
||||
workspace_name=message["workspace_name"],
|
||||
session_name=message["session_name"],
|
||||
message_id=message_id,
|
||||
)
|
||||
|
||||
# Use pre-fetched sequence if available, otherwise fall back to individual query
|
||||
if message_seq_map and message_id in message_seq_map:
|
||||
message_seq_in_session = message_seq_map[message_id]
|
||||
else:
|
||||
message_seq_in_session = await crud.get_message_seq_in_session(
|
||||
db_session,
|
||||
workspace_name=message["workspace_name"],
|
||||
session_name=message["session_name"],
|
||||
message_id=message_id,
|
||||
)
|
||||
|
||||
records: list[dict[str, Any]] = []
|
||||
|
||||
|
|
|
|||
|
|
@ -242,7 +242,7 @@ class EmbeddingStore:
|
|||
workspace_name=self.workspace_name,
|
||||
peer_name=self.peer_name,
|
||||
collection_name=self.collection_name,
|
||||
query=query,
|
||||
query=self._build_truncated_query(query, ""),
|
||||
max_distance=max_distance,
|
||||
top_k=top_k,
|
||||
)
|
||||
|
|
@ -293,6 +293,69 @@ class EmbeddingStore:
|
|||
|
||||
return context
|
||||
|
||||
def _build_truncated_query(
|
||||
self,
|
||||
query: str,
|
||||
conversation_context: str = "",
|
||||
max_tokens: int | None = None,
|
||||
) -> str:
|
||||
"""Build a query that fits within token limits with clear priorities.
|
||||
|
||||
Args:
|
||||
query: The search query
|
||||
conversation_context: Optional conversation context to include
|
||||
max_tokens: Maximum tokens allowed (defaults to setting with buffer)
|
||||
|
||||
Returns:
|
||||
Truncated query string that fits within token limits
|
||||
"""
|
||||
max_tokens = max_tokens or (settings.MAX_EMBEDDING_TOKENS - 100)
|
||||
encoding = embedding_client.encoding
|
||||
|
||||
# Pre-calculate all token counts once
|
||||
query_prefix = "Current message: "
|
||||
context_prefix = "\nContext: "
|
||||
|
||||
prefix_tokens = len(encoding.encode(query_prefix))
|
||||
context_prefix_tokens = len(encoding.encode(context_prefix))
|
||||
query_tokens = encoding.encode(query)
|
||||
|
||||
# Simple case: query alone fits
|
||||
if prefix_tokens + len(query_tokens) <= max_tokens:
|
||||
if not conversation_context:
|
||||
return f"{query_prefix}{query}"
|
||||
|
||||
# Try to add context
|
||||
context_tokens = encoding.encode(conversation_context)
|
||||
total_without_context = (
|
||||
prefix_tokens + len(query_tokens) + context_prefix_tokens
|
||||
)
|
||||
|
||||
if total_without_context + len(context_tokens) <= max_tokens:
|
||||
return f"{query_prefix}{query}{context_prefix}{conversation_context}"
|
||||
|
||||
# Truncate context to fit
|
||||
available_context_tokens = max_tokens - total_without_context
|
||||
if available_context_tokens > 0:
|
||||
truncated_context = encoding.decode(
|
||||
context_tokens[-available_context_tokens:]
|
||||
)
|
||||
return f"{query_prefix}{query}{context_prefix}{truncated_context}"
|
||||
else:
|
||||
# No room left for context; keep full query intact
|
||||
return f"{query_prefix}{query}"
|
||||
|
||||
# Query itself is too long - truncate it
|
||||
available_query_tokens = max_tokens - prefix_tokens
|
||||
if available_query_tokens > 0:
|
||||
# Keep the end (recency) of the query
|
||||
truncated_query = encoding.decode(query_tokens[-available_query_tokens:])
|
||||
return f"{query_prefix}{truncated_query}"
|
||||
|
||||
# Pathological case - just return what we can
|
||||
logger.warning("Token limit too restrictive: %s", max_tokens)
|
||||
return encoding.decode(query_tokens[:max_tokens])
|
||||
|
||||
async def _query_documents_for_level(
|
||||
self,
|
||||
db: AsyncSession,
|
||||
|
|
@ -303,11 +366,8 @@ class EmbeddingStore:
|
|||
count: int,
|
||||
) -> list[models.Document]:
|
||||
"""Query documents for a specific level."""
|
||||
combined_query: str = (
|
||||
f"Current message: {query}\nContext: {conversation_context}"
|
||||
if conversation_context
|
||||
else query
|
||||
)
|
||||
# Construct the combined query with truncation to prevent token limit errors
|
||||
combined_query = self._build_truncated_query(query, conversation_context)
|
||||
|
||||
documents = await crud.query_documents(
|
||||
db,
|
||||
|
|
|
|||
|
|
@ -17,25 +17,43 @@ class TestEnqueueFunction:
|
|||
"""Test suite for the enqueue function's internal logic"""
|
||||
|
||||
# Helper methods
|
||||
def create_sample_payload(
|
||||
async def create_sample_payload(
|
||||
self,
|
||||
db_session: AsyncSession,
|
||||
workspace_name: str = "test_workspace",
|
||||
session_name: str | None = "test_session",
|
||||
peer_name: str = "test_peer",
|
||||
count: int = 1,
|
||||
):
|
||||
"""Create sample payload for testing"""
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Create real messages in database and return payload with actual IDs"""
|
||||
messages: list[models.Message] = []
|
||||
for i in range(count):
|
||||
message = models.Message(
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=peer_name,
|
||||
content=f"Test message {i}",
|
||||
public_id=generate_nanoid(),
|
||||
token_count=10,
|
||||
h_metadata={"test": f"value_{i}"},
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
||||
await db_session.commit()
|
||||
|
||||
# Return payload with real message IDs
|
||||
return [
|
||||
{
|
||||
"workspace_name": workspace_name,
|
||||
"session_name": session_name,
|
||||
"message_id": i + 1,
|
||||
"content": f"Test message {i}",
|
||||
"metadata": {"test": f"value_{i}"},
|
||||
"message_id": msg.id,
|
||||
"content": msg.content,
|
||||
"metadata": msg.h_metadata,
|
||||
"peer_name": peer_name,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"created_at": msg.created_at,
|
||||
}
|
||||
for i in range(count)
|
||||
for msg in messages
|
||||
]
|
||||
|
||||
async def count_queue_items(self, db_session: AsyncSession):
|
||||
|
|
@ -67,7 +85,7 @@ class TestEnqueueFunction:
|
|||
db_session: AsyncSession,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
):
|
||||
"""Test that deriver disabled sessions skip enqueue"""
|
||||
"""Test that deriver disabled sessions skip representation but allows summary"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, test_peer = sample_data
|
||||
|
|
@ -81,7 +99,8 @@ class TestEnqueueFunction:
|
|||
db_session.add(test_session)
|
||||
await db_session.commit()
|
||||
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer.name,
|
||||
|
|
@ -91,7 +110,12 @@ class TestEnqueueFunction:
|
|||
await enqueue(payload)
|
||||
final_count = await self.count_queue_items(db_session)
|
||||
|
||||
assert final_count == initial_count
|
||||
# When deriver is disabled, only summary records should be created (if applicable)
|
||||
# Since this is message 1, and 1 % 20 != 0 and 1 % 60 != 0, no summary should be created
|
||||
# No representation records should be created either (deriver disabled)
|
||||
assert (
|
||||
final_count == initial_count
|
||||
), f"Expected no queue items, but got {final_count - initial_count}"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch("src.deriver.enqueue.tracked_db")
|
||||
|
|
@ -115,7 +139,8 @@ class TestEnqueueFunction:
|
|||
)
|
||||
await db_session.commit()
|
||||
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer.name,
|
||||
|
|
@ -171,7 +196,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
NUM_MESSAGES = 3
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer1.name,
|
||||
|
|
@ -249,7 +275,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
NUM_MESSAGES = 3
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer1.name,
|
||||
|
|
@ -342,7 +369,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
NUM_MESSAGES = 3
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer1.name,
|
||||
|
|
@ -426,7 +454,8 @@ class TestEnqueueFunction:
|
|||
|
||||
await db_session.commit()
|
||||
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer.name,
|
||||
|
|
@ -477,13 +506,15 @@ class TestEnqueueFunction:
|
|||
)
|
||||
await db_session.commit()
|
||||
|
||||
payload1 = self.create_sample_payload(
|
||||
payload1 = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=test_peer1.name,
|
||||
)
|
||||
|
||||
payload2 = self.create_sample_payload(
|
||||
payload2 = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=additional_sender_peer.name,
|
||||
|
|
@ -578,7 +609,6 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Simulate sender leaving the session by setting left_at
|
||||
from datetime import datetime, timezone
|
||||
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
|
|
@ -592,7 +622,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message payload from the peer who left
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
|
|
@ -679,7 +710,6 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Simulate one observer leaving the session
|
||||
from datetime import datetime, timezone
|
||||
|
||||
session_peer_result = await db_session.execute(
|
||||
select(models.SessionPeer).where(
|
||||
|
|
@ -693,7 +723,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message payload
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
|
|
@ -730,7 +761,7 @@ class TestEnqueueFunction:
|
|||
"""Test get_effective_observe_me handles missing sender configuration gracefully"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, _existing_peer = sample_data
|
||||
test_workspace, existing_peer = sample_data
|
||||
|
||||
# Create observer peer
|
||||
observer_peer = models.Peer(
|
||||
|
|
@ -753,11 +784,11 @@ class TestEnqueueFunction:
|
|||
|
||||
# Create message from peer NOT in the session configuration
|
||||
# This simulates the race condition where a peer left after sending
|
||||
unknown_sender = str(generate_nanoid())
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=unknown_sender,
|
||||
peer_name=existing_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
|
|
@ -776,12 +807,12 @@ class TestEnqueueFunction:
|
|||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": unknown_sender,
|
||||
"target_name": unknown_sender,
|
||||
"sender_name": existing_peer.name,
|
||||
"target_name": existing_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": unknown_sender,
|
||||
"sender_name": existing_peer.name,
|
||||
"target_name": observer_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
|
|
@ -861,7 +892,6 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Mark some peers as having left the session
|
||||
from datetime import datetime, timezone
|
||||
|
||||
for peer_name in [inactive_observer.name, inactive_non_observer.name]:
|
||||
session_peer_result = await db_session.execute(
|
||||
|
|
@ -876,7 +906,8 @@ class TestEnqueueFunction:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message payload from sender
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
|
|
@ -1045,25 +1076,43 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
"""Test advanced edge cases for the enqueue system with race conditions"""
|
||||
|
||||
# Helper methods
|
||||
def create_sample_payload(
|
||||
async def create_sample_payload(
|
||||
self,
|
||||
db_session: AsyncSession,
|
||||
workspace_name: str = "test_workspace",
|
||||
session_name: str | None = "test_session",
|
||||
peer_name: str = "test_peer",
|
||||
count: int = 1,
|
||||
):
|
||||
"""Create sample payload for testing"""
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Create real messages in database and return payload with actual IDs"""
|
||||
messages: list[models.Message] = []
|
||||
for i in range(count):
|
||||
message = models.Message(
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=peer_name,
|
||||
content=f"Test message {i}",
|
||||
public_id=generate_nanoid(),
|
||||
token_count=10,
|
||||
h_metadata={"test": f"value_{i}"},
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
||||
await db_session.commit()
|
||||
|
||||
# Return payload with real message IDs
|
||||
return [
|
||||
{
|
||||
"workspace_name": workspace_name,
|
||||
"session_name": session_name,
|
||||
"message_id": i + 1,
|
||||
"content": f"Test message {i}",
|
||||
"metadata": {"test": f"value_{i}"},
|
||||
"message_id": msg.id,
|
||||
"content": msg.content,
|
||||
"metadata": msg.h_metadata,
|
||||
"peer_name": peer_name,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"created_at": msg.created_at,
|
||||
}
|
||||
for i in range(count)
|
||||
for msg in messages
|
||||
]
|
||||
|
||||
async def count_queue_items(self, db_session: AsyncSession):
|
||||
|
|
@ -1109,7 +1158,6 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
await db_session.commit()
|
||||
|
||||
# Mark all observers as having left
|
||||
from datetime import datetime, timezone
|
||||
|
||||
for peer_name in [observer1.name, observer2.name]:
|
||||
session_peer_result = await db_session.execute(
|
||||
|
|
@ -1124,7 +1172,8 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message payload
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
|
|
@ -1180,7 +1229,6 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
await db_session.commit()
|
||||
|
||||
# Mark both as having left (observer left first, then sender)
|
||||
from datetime import datetime, timezone
|
||||
|
||||
base_time = datetime.now(timezone.utc)
|
||||
|
||||
|
|
@ -1209,7 +1257,8 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message payload from sender who left
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=sender_peer.name,
|
||||
|
|
@ -1244,7 +1293,7 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
"""Test handling message from peer who was never in the session"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, _existing_peer = sample_data
|
||||
test_workspace, existing_peer = sample_data
|
||||
|
||||
observer_peer = models.Peer(
|
||||
workspace_name=test_workspace.name, name=str(generate_nanoid())
|
||||
|
|
@ -1265,11 +1314,11 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
await db_session.commit()
|
||||
|
||||
# Create message from peer who was NEVER in the session
|
||||
never_joined_peer = str(generate_nanoid())
|
||||
payload = self.create_sample_payload(
|
||||
payload = await self.create_sample_payload(
|
||||
db_session,
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=test_session.name,
|
||||
peer_name=never_joined_peer,
|
||||
peer_name=existing_peer.name,
|
||||
)
|
||||
|
||||
initial_count = await self.count_queue_items(db_session)
|
||||
|
|
@ -1288,12 +1337,12 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
|
||||
expected_payloads = [
|
||||
{
|
||||
"sender_name": never_joined_peer,
|
||||
"target_name": never_joined_peer,
|
||||
"sender_name": existing_peer.name,
|
||||
"target_name": existing_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
{
|
||||
"sender_name": never_joined_peer,
|
||||
"sender_name": existing_peer.name,
|
||||
"target_name": observer_peer.name,
|
||||
"task_type": "representation",
|
||||
},
|
||||
|
|
|
|||
Loading…
Reference in New Issue