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:
Rajat Ahuja 2025-10-16 11:46:55 -04:00 committed by GitHub
parent 3ba63edad2
commit 77a965e97f
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
15 changed files with 278 additions and 116 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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