fix: N+1 query in search_messages

This commit is contained in:
Rajat Ahuja 2026-02-02 14:07:16 -05:00
parent 8ae21bb4fa
commit 78593aadd3
1 changed files with 25 additions and 12 deletions

View File

@ -3,7 +3,7 @@ from logging import getLogger
from typing import Any
from nanoid import generate as generate_nanoid
from sqlalchemy import ColumnElement, Select, and_, func, select, text, update
from sqlalchemy import ColumnElement, Select, and_, func, or_, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession
from src import models, schemas
@ -109,18 +109,31 @@ async def _build_merged_snippets(
else:
merged_ranges.append((start, end, [match]))
# Batch all ranges into a single query using OR conditions.
# NOTE: If callers ever pass a very high limit (many disjoint ranges),
# consider chunking to avoid oversized SQL / planner issues.
range_conditions = [
models.Message.seq_in_session.between(start_seq, end_seq)
for start_seq, end_seq, _ in merged_ranges
]
context_stmt = (
select(models.Message)
.where(models.Message.workspace_name == workspace_name)
.where(models.Message.session_name == sess_name)
.where(or_(*range_conditions))
.order_by(models.Message.seq_in_session.asc())
)
context_result = await db.execute(context_stmt)
all_context_messages = list(context_result.scalars().all())
# Partition results back into their respective ranges
for start_seq, end_seq, range_matches in merged_ranges:
context_stmt = (
select(models.Message)
.where(models.Message.workspace_name == workspace_name)
.where(models.Message.session_name == sess_name)
.where(models.Message.seq_in_session.between(start_seq, end_seq))
.order_by(models.Message.seq_in_session.asc())
)
context_result = await db.execute(context_stmt)
context_messages = list(context_result.scalars().all())
context_messages = [
msg
for msg in all_context_messages
if start_seq <= msg.seq_in_session <= end_seq
]
snippets.append((range_matches, context_messages))
return snippets