fix: internal N+1 query in dialectic agent calls - DEV-1721 (#652)
* fix: internal N+1 query in dialectic agent calls * fix: comments
This commit is contained in:
parent
ad7c1b3040
commit
a4ae372932
|
|
@ -140,13 +140,12 @@ async def _build_merged_snippets(
|
|||
for msg in matched_messages:
|
||||
session_matches.setdefault(msg.session_name, []).append(msg)
|
||||
|
||||
snippets: list[tuple[list[models.Message], list[models.Message]]] = []
|
||||
|
||||
# Build merged ranges per session, then issue a single batched query
|
||||
session_ranges: dict[str, list[tuple[int, int, list[models.Message]]]] = {}
|
||||
for sess_name, matches in session_matches.items():
|
||||
matches.sort(key=lambda m: m.seq_in_session)
|
||||
|
||||
merged_ranges: list[tuple[int, int, list[models.Message]]] = []
|
||||
|
||||
for match in matches:
|
||||
start = match.seq_in_session - context_window
|
||||
end = match.seq_in_session + context_window
|
||||
|
|
@ -161,25 +160,42 @@ 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())
|
||||
session_ranges[sess_name] = merged_ranges
|
||||
|
||||
# One OR-of-ANDs predicate covers every (session, range) pair
|
||||
session_predicates = [
|
||||
and_(
|
||||
models.Message.session_name == sess_name,
|
||||
or_(
|
||||
*(
|
||||
models.Message.seq_in_session.between(start_seq, end_seq)
|
||||
for start_seq, end_seq, _ in merged_ranges
|
||||
)
|
||||
),
|
||||
)
|
||||
for sess_name, merged_ranges in session_ranges.items()
|
||||
]
|
||||
|
||||
context_result = await db.execute(context_stmt)
|
||||
all_context_messages = list(context_result.scalars().all())
|
||||
context_stmt = (
|
||||
select(models.Message)
|
||||
.where(models.Message.workspace_name == workspace_name)
|
||||
.where(or_(*session_predicates))
|
||||
.order_by(
|
||||
models.Message.session_name.asc(),
|
||||
models.Message.seq_in_session.asc(),
|
||||
)
|
||||
)
|
||||
|
||||
# Partition results back into their respective ranges
|
||||
context_result = await db.execute(context_stmt)
|
||||
by_session: dict[str, list[models.Message]] = {}
|
||||
for msg in context_result.scalars().all():
|
||||
by_session.setdefault(msg.session_name, []).append(msg)
|
||||
|
||||
snippets: list[
|
||||
tuple[list[models.Message], list[models.Message]]
|
||||
] = [] # list of tuples, each containing query matches and context messages
|
||||
for sess_name, merged_ranges in session_ranges.items():
|
||||
all_context_messages = by_session.get(sess_name, [])
|
||||
for start_seq, end_seq, range_matches in merged_ranges:
|
||||
context_messages = [
|
||||
msg
|
||||
|
|
|
|||
|
|
@ -17,11 +17,50 @@ from src import models
|
|||
from src.config import settings
|
||||
from src.crud import create_messages
|
||||
from src.crud import message as message_crud
|
||||
from src.models import Peer, Workspace
|
||||
from src.models import Message, Peer, Workspace
|
||||
from src.schemas import MessageCreate
|
||||
from src.utils.search import search
|
||||
|
||||
|
||||
class _FakeScalarResult:
|
||||
def __init__(self, rows: list[models.Message]):
|
||||
self._rows: list[Message] = rows
|
||||
|
||||
def all(self) -> list[models.Message]:
|
||||
return self._rows
|
||||
|
||||
|
||||
class _FakeResult:
|
||||
def __init__(self, rows: list[models.Message]):
|
||||
self._rows: list[Message] = rows
|
||||
|
||||
def scalars(self) -> _FakeScalarResult:
|
||||
return _FakeScalarResult(self._rows)
|
||||
|
||||
|
||||
class _CountingDb:
|
||||
def __init__(self, rows: list[models.Message]):
|
||||
self._rows: list[Message] = rows
|
||||
self.execute_count: int = 0
|
||||
|
||||
async def execute(self, _stmt: Any) -> _FakeResult:
|
||||
self.execute_count += 1
|
||||
return _FakeResult(self._rows)
|
||||
|
||||
|
||||
def _message(session_name: str, seq_in_session: int) -> models.Message:
|
||||
return models.Message(
|
||||
workspace_name="workspace",
|
||||
session_name=session_name,
|
||||
peer_name="peer",
|
||||
content=f"{session_name}:{seq_in_session}",
|
||||
public_id=generate_nanoid(),
|
||||
seq_in_session=seq_in_session,
|
||||
token_count=1,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_embedding_created_when_setting_enabled(
|
||||
db_session: AsyncSession,
|
||||
|
|
@ -260,6 +299,46 @@ async def test_semantic_search_when_embeddings_enabled(
|
|||
assert created_message.public_id in found_message_ids
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_merged_snippets_batches_context_query_across_sessions():
|
||||
"""Context expansion should not issue one DB query per matched session."""
|
||||
matched_messages = [
|
||||
_message("session_a", 10),
|
||||
_message("session_b", 20),
|
||||
_message("session_c", 30),
|
||||
]
|
||||
context_messages = [
|
||||
_message("session_a", 9),
|
||||
_message("session_a", 10),
|
||||
_message("session_a", 11),
|
||||
_message("session_a", 99),
|
||||
_message("session_b", 19),
|
||||
_message("session_b", 20),
|
||||
_message("session_b", 21),
|
||||
_message("session_c", 29),
|
||||
_message("session_c", 30),
|
||||
_message("session_c", 31),
|
||||
]
|
||||
db = _CountingDb(context_messages)
|
||||
|
||||
snippets = await message_crud._build_merged_snippets( # pyright: ignore[reportPrivateUsage]
|
||||
db, # pyright: ignore[reportArgumentType]
|
||||
workspace_name="workspace",
|
||||
matched_messages=matched_messages,
|
||||
context_window=1,
|
||||
)
|
||||
|
||||
assert db.execute_count == 1
|
||||
assert [len(matches) for matches, _ in snippets] == [1, 1, 1]
|
||||
assert [
|
||||
[msg.content for msg in context_messages] for _, context_messages in snippets
|
||||
] == [
|
||||
["session_a:9", "session_a:10", "session_a:11"],
|
||||
["session_b:19", "session_b:20", "session_b:21"],
|
||||
["session_c:29", "session_c:30", "session_c:31"],
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_search_messages_external_lookup_happens_before_tracked_db(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
|
|
|
|||
Loading…
Reference in New Issue