diff --git a/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py b/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py new file mode 100644 index 00000000..36627479 --- /dev/null +++ b/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py @@ -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) diff --git a/src/crud/__init__.py b/src/crud/__init__.py index 69a4799e..7e0db845 100644 --- a/src/crud/__init__.py +++ b/src/crud/__init__.py @@ -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", diff --git a/src/crud/message.py b/src/crud/message.py index 143f1aa9..c6d3476d 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -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( diff --git a/src/crud/session.py b/src/crud/session.py index b3adf591..cfb77660 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -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 ] diff --git a/src/deriver/enqueue.py b/src/deriver/enqueue.py index a9eb736e..9405f6fa 100644 --- a/src/deriver/enqueue.py +++ b/src/deriver/enqueue.py @@ -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"], diff --git a/src/models.py b/src/models.py index 9b4e6f1d..1cd3b00c 100644 --- a/src/models.py +++ b/src/models.py @@ -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", diff --git a/src/routers/messages.py b/src/routers/messages.py index 9c858c8e..259fc9d7 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -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 ] diff --git a/tests/conftest.py b/tests/conftest.py index 844a2769..9ba3036e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 diff --git a/tests/crud/test_workspace.py b/tests/crud/test_workspace.py index a9d32fb6..18601128 100644 --- a/tests/crud/test_workspace.py +++ b/tests/crud/test_workspace.py @@ -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]) diff --git a/tests/deriver/conftest.py b/tests/deriver/conftest.py index 9ade5ef7..89410656 100644 --- a/tests/deriver/conftest.py +++ b/tests/deriver/conftest.py @@ -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, }, ] diff --git a/tests/deriver/test_deriver_processing.py b/tests/deriver/test_deriver_processing.py index dc986ae2..002377df 100644 --- a/tests/deriver/test_deriver_processing.py +++ b/tests/deriver/test_deriver_processing.py @@ -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), ) diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 6f1527ad..d620392a 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -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) diff --git a/tests/integration/test_enqueue.py b/tests/integration/test_enqueue.py index 0f541704..5b1870e8 100644 --- a/tests/integration/test_enqueue.py +++ b/tests/integration/test_enqueue.py @@ -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}"}, ) diff --git a/tests/routes/test_messages.py b/tests/routes/test_messages.py index 2f222339..5bc3f1d1 100644 --- a/tests/routes/test_messages.py +++ b/tests/routes/test_messages.py @@ -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() diff --git a/tests/routes/test_peers.py b/tests/routes/test_peers.py index 2341e15e..f106abd9 100644 --- a/tests/routes/test_peers.py +++ b/tests/routes/test_peers.py @@ -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())