feat: fix race condition in message sequence batching (#235)
* feat: fix race condition in message sequence batching * fix: CodeRabbit comments; commit early to release the advisory lock before generating embeddings * fix: use index + rm unused method * fix: PR comments * fix: bug in lock timeout * fix: patch tracked_db for peers route within conftest.py
This commit is contained in:
parent
3ba63edad2
commit
77a965e97f
|
|
@ -0,0 +1,135 @@
|
|||
"""add seq_in_session column to messages table
|
||||
|
||||
Revision ID: bb6fb3a7a643
|
||||
Revises: 76ffba56fe8c
|
||||
Create Date: 2025-10-20 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
from migrations.utils import column_exists, constraint_exists, get_schema
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "bb6fb3a7a643"
|
||||
down_revision: str | None = "76ffba56fe8c"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
schema = get_schema()
|
||||
|
||||
BATCH_SIZE = 10_000
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("messages", "seq_in_session"):
|
||||
op.add_column(
|
||||
"messages",
|
||||
sa.Column("seq_in_session", sa.BigInteger(), nullable=True),
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
conn = op.get_bind()
|
||||
preparer = conn.dialect.identifier_preparer
|
||||
messages_table = sa.Table("messages", sa.MetaData(), schema=schema)
|
||||
qualified_messages = preparer.format_table(messages_table)
|
||||
id_col = preparer.quote("id")
|
||||
workspace_col = preparer.quote("workspace_name")
|
||||
session_col = preparer.quote("session_name")
|
||||
seq_col = preparer.quote("seq_in_session")
|
||||
|
||||
distinct_sessions = conn.execute(
|
||||
sa.text(
|
||||
f"""
|
||||
SELECT DISTINCT
|
||||
{workspace_col} AS workspace_name,
|
||||
{session_col} AS session_name
|
||||
FROM {qualified_messages}
|
||||
"""
|
||||
)
|
||||
)
|
||||
|
||||
update_stmt = sa.text(
|
||||
f"""
|
||||
WITH params AS (
|
||||
SELECT
|
||||
:workspace_name AS workspace_name,
|
||||
:session_name AS session_name,
|
||||
COALESCE(
|
||||
(
|
||||
SELECT MAX({seq_col})
|
||||
FROM {qualified_messages}
|
||||
WHERE {workspace_col} = :workspace_name
|
||||
AND {session_col} = :session_name
|
||||
),
|
||||
0
|
||||
) AS offset
|
||||
),
|
||||
batch AS (
|
||||
SELECT
|
||||
m.{id_col} AS id,
|
||||
ROW_NUMBER() OVER (ORDER BY m.{id_col}) AS rn,
|
||||
params.offset
|
||||
FROM {qualified_messages} AS m
|
||||
JOIN params ON TRUE
|
||||
WHERE m.{workspace_col} = params.workspace_name
|
||||
AND m.{session_col} = params.session_name
|
||||
AND m.{seq_col} IS NULL
|
||||
ORDER BY m.{id_col}
|
||||
LIMIT :batch_size
|
||||
)
|
||||
UPDATE {qualified_messages} AS m
|
||||
SET {seq_col} = batch.rn + batch.offset
|
||||
FROM batch
|
||||
WHERE m.{id_col} = batch.id
|
||||
"""
|
||||
)
|
||||
|
||||
for workspace_name, session_name in distinct_sessions:
|
||||
while True:
|
||||
result = conn.execute(
|
||||
update_stmt,
|
||||
{
|
||||
"workspace_name": workspace_name,
|
||||
"session_name": session_name,
|
||||
"batch_size": BATCH_SIZE,
|
||||
},
|
||||
)
|
||||
updated_rows = result.rowcount or 0
|
||||
result.close()
|
||||
if updated_rows == 0:
|
||||
break
|
||||
distinct_sessions.close()
|
||||
|
||||
op.alter_column(
|
||||
"messages",
|
||||
"seq_in_session",
|
||||
nullable=False,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if not constraint_exists("messages", "uq_messages_session_seq", "unique"):
|
||||
op.create_unique_constraint(
|
||||
"uq_messages_session_seq",
|
||||
"messages",
|
||||
["workspace_name", "session_name", "seq_in_session"],
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
schema = get_schema()
|
||||
|
||||
if constraint_exists("messages", "uq_messages_session_seq", "unique"):
|
||||
op.drop_constraint(
|
||||
"uq_messages_session_seq",
|
||||
"messages",
|
||||
type_="unique",
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
if column_exists("messages", "seq_in_session"):
|
||||
op.drop_column("messages", "seq_in_session", schema=schema)
|
||||
|
|
@ -9,7 +9,6 @@ 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,
|
||||
|
|
@ -67,7 +66,6 @@ __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",
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from logging import getLogger
|
|||
from typing import Any
|
||||
|
||||
from nanoid import generate as generate_nanoid
|
||||
from sqlalchemy import ColumnElement, Select, and_, func, select
|
||||
from sqlalchemy import ColumnElement, Select, and_, func, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src import models, schemas
|
||||
|
|
@ -78,9 +78,32 @@ async def create_messages(
|
|||
workspace_name=workspace_name,
|
||||
)
|
||||
|
||||
await db.execute(text("SET LOCAL lock_timeout = '5s'"))
|
||||
await db.execute(
|
||||
text(
|
||||
"SELECT pg_advisory_xact_lock(hashtext(:workspace_name), hashtext(:session_name))"
|
||||
),
|
||||
{"workspace_name": workspace_name, "session_name": session_name},
|
||||
)
|
||||
|
||||
# Get the last sequence number on a session - uses (workspace_name, session_name, seq_in_session) index
|
||||
last_seq = (
|
||||
await db.scalar(
|
||||
select(models.Message.seq_in_session)
|
||||
.where(
|
||||
models.Message.workspace_name == workspace_name,
|
||||
models.Message.session_name == session_name,
|
||||
)
|
||||
.order_by(models.Message.seq_in_session.desc())
|
||||
.limit(1)
|
||||
)
|
||||
or 0
|
||||
)
|
||||
|
||||
# Create list of message objects (this will trigger the before_insert event)
|
||||
message_objects: list[models.Message] = []
|
||||
for message in messages:
|
||||
for offset, message in enumerate(messages, start=1):
|
||||
message_seq_in_session = last_seq + offset
|
||||
message_obj = models.Message(
|
||||
session_name=session_name,
|
||||
peer_name=message.peer_name,
|
||||
|
|
@ -90,46 +113,55 @@ async def create_messages(
|
|||
public_id=generate_nanoid(),
|
||||
token_count=len(message.encoded_message),
|
||||
created_at=message.created_at, # Use provided created_at if available
|
||||
seq_in_session=message_seq_in_session,
|
||||
)
|
||||
message_objects.append(message_obj)
|
||||
|
||||
db.add_all(message_objects)
|
||||
await db.flush()
|
||||
|
||||
if settings.EMBED_MESSAGES:
|
||||
encoded_message_lookup = {
|
||||
msg.public_id: orig_msg.encoded_message
|
||||
for msg, orig_msg in zip(message_objects, messages, strict=True)
|
||||
}
|
||||
id_resource_dict = {
|
||||
message.public_id: (
|
||||
message.content,
|
||||
encoded_message_lookup[message.public_id],
|
||||
)
|
||||
for message in message_objects
|
||||
}
|
||||
embedding_dict = await embedding_client.batch_embed(id_resource_dict)
|
||||
|
||||
# Create MessageEmbedding entries for each embedded message
|
||||
embedding_objects: list[models.MessageEmbedding] = []
|
||||
for message_obj in message_objects:
|
||||
embeddings = embedding_dict.get(message_obj.public_id, [])
|
||||
for embedding in embeddings:
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
embedding=embedding,
|
||||
message_id=message_obj.public_id,
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=message_obj.peer_name,
|
||||
)
|
||||
embedding_objects.append(embedding_obj)
|
||||
|
||||
# Add all embedding objects to the session
|
||||
if embedding_objects:
|
||||
db.add_all(embedding_objects)
|
||||
|
||||
# Commit here to release the advisory lock before generating embeddings
|
||||
await db.commit()
|
||||
try:
|
||||
if settings.EMBED_MESSAGES:
|
||||
encoded_message_lookup = {
|
||||
msg.public_id: orig_msg.encoded_message
|
||||
for msg, orig_msg in zip(message_objects, messages, strict=True)
|
||||
}
|
||||
id_resource_dict = {
|
||||
message.public_id: (
|
||||
message.content,
|
||||
encoded_message_lookup[message.public_id],
|
||||
)
|
||||
for message in message_objects
|
||||
}
|
||||
embedding_dict = await embedding_client.batch_embed(id_resource_dict)
|
||||
|
||||
# Create MessageEmbedding entries for each embedded message
|
||||
embedding_objects: list[models.MessageEmbedding] = []
|
||||
for message_obj in message_objects:
|
||||
embeddings = embedding_dict.get(message_obj.public_id, [])
|
||||
for embedding in embeddings:
|
||||
embedding_obj = models.MessageEmbedding(
|
||||
content=message_obj.content,
|
||||
embedding=embedding,
|
||||
message_id=message_obj.public_id,
|
||||
workspace_name=workspace_name,
|
||||
session_name=session_name,
|
||||
peer_name=message_obj.peer_name,
|
||||
)
|
||||
embedding_objects.append(embedding_obj)
|
||||
|
||||
# Add all embedding objects to the session
|
||||
if embedding_objects:
|
||||
db.add_all(embedding_objects)
|
||||
await db.commit()
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to generate message embeddings for %s messages in workspace %s and session %s.",
|
||||
len(message_objects),
|
||||
workspace_name,
|
||||
session_name,
|
||||
)
|
||||
|
||||
return message_objects
|
||||
|
||||
|
|
@ -268,61 +300,13 @@ async def get_message_seq_in_session(
|
|||
The sequence number of the message (1-indexed)
|
||||
"""
|
||||
stmt = (
|
||||
select(func.count(models.Message.id))
|
||||
select(models.Message.seq_in_session)
|
||||
.where(models.Message.workspace_name == workspace_name)
|
||||
.where(models.Message.session_name == session_name)
|
||||
.where(models.Message.id < message_id)
|
||||
.where(models.Message.id == message_id)
|
||||
)
|
||||
result = await db.execute(stmt)
|
||||
count = result.scalar() or 0
|
||||
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}
|
||||
seq: int | None = await db.scalar(stmt)
|
||||
return int(seq) if seq is not None else 0
|
||||
|
||||
|
||||
async def get_message(
|
||||
|
|
|
|||
|
|
@ -336,6 +336,7 @@ async def clone_session(
|
|||
"h_metadata": message.h_metadata,
|
||||
"workspace_name": workspace_name,
|
||||
"peer_name": message.peer_name,
|
||||
"seq_in_session": message.seq_in_session,
|
||||
}
|
||||
for message in messages_to_clone
|
||||
]
|
||||
|
|
|
|||
|
|
@ -98,12 +98,6 @@ 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:
|
||||
|
|
@ -114,7 +108,6 @@ async def handle_session(
|
|||
peers_with_configuration,
|
||||
session.id,
|
||||
deriver_disabled=deriver_disabled,
|
||||
message_seq_map=message_seq_map,
|
||||
)
|
||||
)
|
||||
return queue_records
|
||||
|
|
@ -253,7 +246,6 @@ 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.
|
||||
|
|
@ -272,10 +264,9 @@ async def generate_queue_records(
|
|||
observed = message["peer_name"]
|
||||
message_id: int = message["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:
|
||||
# Prefer the sequence captured during message creation; fallback only if missing
|
||||
message_seq_in_session = int(message.get("seq_in_session") or 0)
|
||||
if message_seq_in_session <= 0:
|
||||
message_seq_in_session = await crud.get_message_seq_in_session(
|
||||
db_session,
|
||||
workspace_name=message["workspace_name"],
|
||||
|
|
|
|||
|
|
@ -186,6 +186,7 @@ class Message(Base):
|
|||
"internal_metadata", JSONB, default=dict
|
||||
)
|
||||
token_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False)
|
||||
seq_in_session: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True)
|
||||
|
||||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), index=True, default=func.now()
|
||||
|
|
@ -216,6 +217,12 @@ class Message(Base):
|
|||
"id",
|
||||
postgresql_include=["id", "created_at"],
|
||||
),
|
||||
UniqueConstraint(
|
||||
"workspace_name",
|
||||
"session_name",
|
||||
"seq_in_session",
|
||||
name="uq_messages_session_seq",
|
||||
),
|
||||
# Full text search index on content column
|
||||
Index(
|
||||
"idx_messages_content_gin",
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ async def create_messages_for_session(
|
|||
"peer_name": message.peer_name,
|
||||
"created_at": message.created_at,
|
||||
"message_public_id": message.public_id,
|
||||
"message_seq_in_session": message.seq_in_session,
|
||||
}
|
||||
for message in created_messages
|
||||
]
|
||||
|
|
@ -134,6 +135,7 @@ async def create_messages_with_file(
|
|||
"peer_name": message.peer_name,
|
||||
"created_at": message.created_at,
|
||||
"message_public_id": message.public_id,
|
||||
"message_seq_in_session": message.seq_in_session,
|
||||
}
|
||||
for message in created_messages
|
||||
]
|
||||
|
|
|
|||
|
|
@ -492,6 +492,7 @@ def mock_tracked_db(db_session: AsyncSession):
|
|||
patch("src.deriver.queue_manager.tracked_db", mock_tracked_db_context),
|
||||
patch("src.routers.sessions.tracked_db", mock_tracked_db_context),
|
||||
patch("src.crud.representation.tracked_db", mock_tracked_db_context),
|
||||
patch("src.routers.peers.tracked_db", mock_tracked_db_context),
|
||||
):
|
||||
yield
|
||||
|
||||
|
|
|
|||
|
|
@ -105,12 +105,14 @@ class TestWorkspaceCRUD:
|
|||
workspace_name=test_workspace.name,
|
||||
session_name=session.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
message2 = models.Message(
|
||||
content="Test message 2",
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=session.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=2,
|
||||
)
|
||||
db_session.add_all([message1, message2])
|
||||
await db_session.flush()
|
||||
|
|
@ -422,12 +424,14 @@ class TestWorkspaceCRUD:
|
|||
workspace_name=test_workspace.name,
|
||||
session_name=session1.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
message2 = models.Message(
|
||||
content="Test message 2",
|
||||
workspace_name=test_workspace.name,
|
||||
session_name=session2.name,
|
||||
peer_name=peer2.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add_all([message1, message2])
|
||||
|
||||
|
|
|
|||
|
|
@ -83,18 +83,21 @@ async def sample_messages(
|
|||
"content": "Hello, this is the first message from peer1",
|
||||
"peer_name": peer1.name,
|
||||
"workspace_name": session.workspace_name,
|
||||
"seq_in_session": 1,
|
||||
},
|
||||
{
|
||||
"session_name": session.name,
|
||||
"content": "Hi there! This is a response from peer2",
|
||||
"peer_name": peer2.name,
|
||||
"workspace_name": session.workspace_name,
|
||||
"seq_in_session": 2,
|
||||
},
|
||||
{
|
||||
"session_name": session.name,
|
||||
"content": "I'm just observing this conversation as peer3",
|
||||
"peer_name": peer3.name,
|
||||
"workspace_name": session.workspace_name,
|
||||
"seq_in_session": 3,
|
||||
},
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -169,6 +169,7 @@ class TestDeriverProcessing:
|
|||
session_name="test_session",
|
||||
peer_name="alice",
|
||||
content=f"message {message_id}",
|
||||
seq_in_session=i + 1,
|
||||
token_count=0,
|
||||
created_at=now - timedelta(minutes=7 - i),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -121,13 +121,14 @@ class TestQueueProcessing:
|
|||
|
||||
# Create and save messages to the database first
|
||||
messages: list[models.Message] = []
|
||||
for _ in range(3):
|
||||
for i in range(3):
|
||||
message = models.Message(
|
||||
session_name=session.name,
|
||||
workspace_name=session.workspace_name,
|
||||
peer_name=peer.name,
|
||||
content="hello",
|
||||
token_count=10,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
@ -291,6 +292,7 @@ class TestQueueProcessing:
|
|||
peer_name=peer.name,
|
||||
content=f"Test message {i}",
|
||||
token_count=token_count,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
@ -396,13 +398,14 @@ class TestQueueProcessing:
|
|||
]
|
||||
|
||||
messages: list[models.Message] = []
|
||||
for peer, token_count in messages_data:
|
||||
for i, (peer, token_count) in enumerate(messages_data):
|
||||
message = models.Message(
|
||||
session_name=session.name,
|
||||
workspace_name=session.workspace_name,
|
||||
peer_name=peer.name,
|
||||
content=f"Message from {peer.name}",
|
||||
token_count=token_count,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
@ -561,13 +564,14 @@ class TestQueueProcessing:
|
|||
]
|
||||
|
||||
messages: list[models.Message] = []
|
||||
for peer, token_count in messages_data:
|
||||
for i, (peer, token_count) in enumerate(messages_data):
|
||||
message = models.Message(
|
||||
session_name=session.name,
|
||||
workspace_name=session.workspace_name,
|
||||
peer_name=peer.name,
|
||||
content=f"Message from {peer.name}",
|
||||
token_count=token_count,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
@ -704,6 +708,7 @@ class TestQueueProcessing:
|
|||
peer_name=peer.name,
|
||||
content="First summary message",
|
||||
public_id=generate_nanoid(),
|
||||
seq_in_session=1,
|
||||
),
|
||||
models.Message(
|
||||
id=1000,
|
||||
|
|
@ -712,6 +717,7 @@ class TestQueueProcessing:
|
|||
peer_name=peer.name,
|
||||
content="Second summary message",
|
||||
public_id=generate_nanoid(),
|
||||
seq_in_session=2,
|
||||
),
|
||||
]
|
||||
|
||||
|
|
@ -820,6 +826,7 @@ class TestQueueProcessing:
|
|||
peer_name=peer.name,
|
||||
content=f"Test message {i}",
|
||||
token_count=token_count,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
@ -930,6 +937,7 @@ class TestQueueProcessing:
|
|||
peer_name=peer.name,
|
||||
content=f"Test message {i}",
|
||||
token_count=token_count,
|
||||
seq_in_session=i + 1,
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch
|
|||
|
||||
import pytest
|
||||
from nanoid import generate as generate_nanoid
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from src import crud, models, schemas
|
||||
|
|
@ -26,6 +26,15 @@ class TestEnqueueFunction:
|
|||
count: int = 1,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Create real messages in database and return payload with actual IDs"""
|
||||
# Get the current max sequence number for this session
|
||||
result = await db_session.execute(
|
||||
select(func.max(models.Message.seq_in_session)).where(
|
||||
models.Message.workspace_name == workspace_name,
|
||||
models.Message.session_name == session_name,
|
||||
)
|
||||
)
|
||||
current_max_seq = result.scalar() or 0
|
||||
|
||||
messages: list[models.Message] = []
|
||||
for i in range(count):
|
||||
message = models.Message(
|
||||
|
|
@ -34,6 +43,7 @@ class TestEnqueueFunction:
|
|||
peer_name=peer_name,
|
||||
content=f"Test message {i}",
|
||||
public_id=generate_nanoid(),
|
||||
seq_in_session=current_max_seq + i + 1,
|
||||
token_count=10,
|
||||
h_metadata={"test": f"value_{i}"},
|
||||
)
|
||||
|
|
@ -1085,6 +1095,15 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
count: int = 1,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Create real messages in database and return payload with actual IDs"""
|
||||
# Get the current max sequence number for this session
|
||||
result = await db_session.execute(
|
||||
select(func.max(models.Message.seq_in_session)).where(
|
||||
models.Message.workspace_name == workspace_name,
|
||||
models.Message.session_name == session_name,
|
||||
)
|
||||
)
|
||||
current_max_seq = result.scalar() or 0
|
||||
|
||||
messages: list[models.Message] = []
|
||||
for i in range(count):
|
||||
message = models.Message(
|
||||
|
|
@ -1093,6 +1112,7 @@ class TestAdvancedEnqueueEdgeCases:
|
|||
peer_name=peer_name,
|
||||
content=f"Test message {i}",
|
||||
public_id=generate_nanoid(),
|
||||
seq_in_session=current_max_seq + i + 1,
|
||||
token_count=10,
|
||||
h_metadata={"test": f"value_{i}"},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -173,6 +173,7 @@ async def test_get_messages(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -211,12 +212,14 @@ async def test_get_messages_with_reverse(
|
|||
content="First message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
test_message2 = models.Message(
|
||||
session_name=test_session.name,
|
||||
content="Second message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=2,
|
||||
)
|
||||
db_session.add(test_message1)
|
||||
db_session.add(test_message2)
|
||||
|
|
@ -266,6 +269,7 @@ async def test_get_messages_with_empty_filter(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -299,6 +303,7 @@ async def test_get_messages_with_null_filter(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -332,6 +337,7 @@ async def test_get_messages_no_body(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -364,6 +370,7 @@ async def test_get_filtered_messages(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"key": "value"},
|
||||
seq_in_session=1,
|
||||
)
|
||||
test_message2 = models.Message(
|
||||
session_name=test_session.name,
|
||||
|
|
@ -371,6 +378,7 @@ async def test_get_filtered_messages(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"key": "value2"},
|
||||
seq_in_session=2,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
db_session.add(test_message2)
|
||||
|
|
@ -410,6 +418,7 @@ async def test_get_filtered_messages_with_complex_filter(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"type": "question", "priority": "high", "category": "technical"},
|
||||
seq_in_session=1,
|
||||
)
|
||||
test_message2 = models.Message(
|
||||
session_name=test_session.name,
|
||||
|
|
@ -417,6 +426,7 @@ async def test_get_filtered_messages_with_complex_filter(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"type": "answer", "priority": "high", "category": "technical"},
|
||||
seq_in_session=2,
|
||||
)
|
||||
test_message3 = models.Message(
|
||||
session_name=test_session.name,
|
||||
|
|
@ -424,6 +434,7 @@ async def test_get_filtered_messages_with_complex_filter(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"type": "question", "priority": "low", "category": "general"},
|
||||
seq_in_session=3,
|
||||
)
|
||||
db_session.add(test_message1)
|
||||
db_session.add(test_message2)
|
||||
|
|
@ -494,6 +505,7 @@ async def test_update_message(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -526,6 +538,7 @@ async def test_update_message_with_complex_metadata(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -568,6 +581,7 @@ async def test_update_message_empty_metadata(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"test_key": "test_value"},
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -607,6 +621,7 @@ async def test_update_message_with_empty_dict_metadata(
|
|||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
h_metadata={"old_key": "old_value"},
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -638,6 +653,7 @@ async def test_get_single_message(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
@ -870,6 +886,7 @@ async def test_update_message_handles_crud_value_error(
|
|||
content="Test message",
|
||||
workspace_name=test_workspace.name,
|
||||
peer_name=test_peer.name,
|
||||
seq_in_session=1,
|
||||
)
|
||||
db_session.add(test_message)
|
||||
await db_session.commit()
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
from typing import Any
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
|
@ -309,15 +308,10 @@ def test_get_sessions_for_peer_with_empty_filter(
|
|||
assert isinstance(data["items"], list)
|
||||
|
||||
|
||||
@patch("src.routers.peers.tracked_db")
|
||||
def test_chat(
|
||||
mock_tracked_db: AsyncMock,
|
||||
client: TestClient,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
db_session: AsyncSession,
|
||||
):
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, test_peer = sample_data
|
||||
target_peer = str(generate_nanoid())
|
||||
|
||||
|
|
@ -335,15 +329,11 @@ def test_chat(
|
|||
assert "content" in data
|
||||
|
||||
|
||||
@patch("src.routers.peers.tracked_db")
|
||||
def test_chat_with_optional_params(
|
||||
mock_tracked_db: AsyncMock,
|
||||
client: TestClient,
|
||||
sample_data: tuple[Workspace, Peer],
|
||||
db_session: AsyncSession,
|
||||
):
|
||||
"""Test chat endpoint with optional parameters"""
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
|
||||
test_workspace, test_peer = sample_data
|
||||
session_id = str(generate_nanoid())
|
||||
|
|
|
|||
Loading…
Reference in New Issue