"""M7: the knowledge read paths must not grow a query per source or per passage. The same discipline `test_context_performance.py` holds for M6, applied to the four paths M7 adds. Each of them lists or joins over rows that a real library has many of, and each could plausibly have been written one query at a time: source list a chunk count and an embedded count per row source detail the source, and its passages retrieval lexical candidates, semantic candidates, their rows context build all of the above, inside a prompt assembly The assertions are on **growth**, not on an exact count: a fixed number breaks on any unrelated query and teaches the next person to raise it. What matters is that four times the library does not cost four times the queries. Also asserted here: candidates are bounded *in the database* before the Python reranking runs. "Do not load every chunk in the campaign merely to find the top few" is a statement about the SQL, so it is tested against the SQL. python -m pytest tests/test_knowledge_performance.py -v """ import asyncio import pytest from fastapi import Depends from fastapi.testclient import TestClient from sqlalchemy import event, select from app import auth, limits, memorybank, models from app.database import Base, SessionLocal, engine, get_db from app.knowledge import embeddings, retrieval from app.main import app from app.routers import adventures from fakes import ScriptedProvider, state_block class StubEmbedder: async def embed(self, texts): return [[1.0, float(len(t) % 5), 0.5] for t in texts] @pytest.fixture() def sql_log(): statements: list[str] = [] def record(conn, cursor, statement, parameters, context, executemany): statements.append(statement) event.listen(engine, "before_cursor_execute", record) try: yield statements finally: event.remove(engine, "before_cursor_execute", record) @pytest.fixture() def client(monkeypatch): Base.metadata.create_all(bind=engine) memorybank._vector_cache.clear() embeddings._cache.clear() setup = SessionLocal() user = models.User(is_guest=False, email="m7perf@example.com") setup.add(user) setup.flush() setup.add(models.Settings( user_id=user.id, model="test-model", embedding_model="nomic-embed-text", context_token_budget=8000, max_output_tokens=400, )) adventure = models.Adventure(user_id=user.id, title="Performance") setup.add(adventure) setup.flush() setup.add(models.Action(adventure_id=adventure.id, type="start", text="Aldric stands in the crypt beneath the Old Abbey.")) setup.commit() adv_id, user_id = adventure.id, user.id setup.close() monkeypatch.setattr(limits, "check_row_cap", lambda *a, **k: None) monkeypatch.setattr(adventures.turns, "OpenAICompatibleProvider", ScriptedProvider) monkeypatch.setattr(memorybank, "embedding_provider", lambda s: StubEmbedder()) monkeypatch.setattr(memorybank, "summary_provider", lambda s: StubEmbedder()) app.dependency_overrides[auth.get_current_user] = ( lambda db=Depends(get_db): db.get(models.User, user_id) ) test_client = TestClient(app) test_client.adv_id = adv_id test_client.user_id = user_id try: yield test_client finally: app.dependency_overrides.clear() memorybank._vector_cache.clear() embeddings._cache.clear() Base.metadata.drop_all(bind=engine) def add_sources(client, count, paragraphs=6, prefix="lore"): """Imports `count` sources, each with several passages of crypt-ish prose.""" for n in range(count): body = "\n\n".join( f"## {prefix} {n} section {p}\n\n" + ("The crypt beneath the Old Abbey at Westhaven is vaulted in " "stone, and the stair descends past niches cut for the dead. ") * 8 for p in range(paragraphs) ) response = client.post( f"/api/adventures/{client.adv_id}/knowledge", files={"file": (f"{prefix}-{n}.md", body.encode(), "text/markdown")}, data={"classification": ["canon", "reference", "inspiration"][n % 3], "allow_duplicate": "true"}, ) assert response.status_code == 201, response.text[:200] def counts(client): with SessionLocal() as db: sources = len(db.execute(select(models.KnowledgeSource)).scalars().all()) chunks = len(db.execute(select(models.KnowledgeChunk)).scalars().all()) return sources, chunks def measure(sql_log, call): sql_log.clear() result = call() return len(sql_log), result def retrieve(client): with SessionLocal() as db: adventure = db.get(models.Adventure, client.adv_id) settings = db.execute(select(models.Settings).where( models.Settings.user_id == client.user_id)).scalars().first() return asyncio.run(retrieval.retrieve(adventure, settings)) # --------------------------------------------------------------------- tests def test_the_source_list_does_not_cost_a_query_per_source(client, sql_log): add_sources(client, 4) small, _ = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/knowledge").json()) add_sources(client, 12, prefix="more") large, rows = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/knowledge").json()) assert len(rows) == 16 assert large == small, f"{small} queries for 4 sources, {large} for 16" # ...and the counts it shows are real, so the fixed query count is not # because the counts were dropped. assert all(row["chunk_count"] > 0 for row in rows) def test_source_detail_does_not_cost_a_query_per_passage(client, sql_log): add_sources(client, 1, paragraphs=3) small_id = client.get(f"/api/adventures/{client.adv_id}/knowledge").json()[0]["id"] small, _ = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/knowledge/{small_id}/chunks").json()) add_sources(client, 1, paragraphs=24, prefix="big") big_id = client.get(f"/api/adventures/{client.adv_id}/knowledge").json()[-1]["id"] large, chunks = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/knowledge/{big_id}/chunks").json()) assert len(chunks) > 3 assert large == small, f"{small} queries for a small source, {large} for a big one" def test_retrieval_does_not_grow_with_the_library(client, sql_log): add_sources(client, 4) embed_pending(client) small, small_result = measure(sql_log, lambda: retrieve(client)) add_sources(client, 16, prefix="more") embed_pending(client) embeddings.forget_cached(client.adv_id) large, large_result = measure(sql_log, lambda: retrieve(client)) sources, chunks = counts(client) assert sources == 20 and chunks > 40 assert small_result.candidates and large_result.candidates assert large <= small + 1, f"{small} queries at 4 sources, {large} at 20" def test_the_context_build_does_not_grow_with_the_library(client, sql_log): add_sources(client, 4) embed_pending(client) ScriptedProvider.replies = [f"The crypt is cold.\n{state_block([])}"] client.post(f"/api/adventures/{client.adv_id}/actions", json={"type": "do", "text": "Aldric descends into the crypt."}) small, _ = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/context").json()) add_sources(client, 16, prefix="more") embed_pending(client) embeddings.forget_cached(client.adv_id) large, report = measure(sql_log, lambda: client.get( f"/api/adventures/{client.adv_id}/context").json()) assert report["knowledge"]["used"] assert large <= small + 1, f"{small} queries at 4 sources, {large} at 20" def test_candidates_are_bounded_in_sql_before_the_python_ranking(client, sql_log): """"Do not load every chunk merely to find the top few", asserted on the SQL.""" add_sources(client, 20, paragraphs=8) embed_pending(client) embeddings.forget_cached(client.adv_id) _sources, chunks = counts(client) assert chunks > retrieval.LEXICAL_CANDIDATES * 2, chunks sql_log.clear() result = retrieve(client) # The lexical query names a LIMIT, and the merged candidate set is bounded # by the two per-path caps rather than by the size of the library. lexical = [s for s in sql_log if "knowledge_fts" in s and "MATCH" in s] assert lexical, sql_log assert all("LIMIT" in s for s in lexical) assert result.considered <= ( retrieval.LEXICAL_CANDIDATES + retrieval.SEMANTIC_CANDIDATES ) assert result.considered < chunks, (result.considered, chunks) # The row fetch for those candidates is one query, not one per candidate. loads = [s for s in sql_log if "knowledge_chunks" in s and "knowledge_sources" in s and " IN " in s.upper()] assert len(loads) <= 2, loads def test_the_semantic_scan_reads_only_narrow_columns(client, sql_log): """A vector is 6 kB; the catalogue read must not fetch passage text.""" add_sources(client, 6) embed_pending(client) embeddings.forget_cached(client.adv_id) sql_log.clear() retrieve(client) catalogue = [s for s in sql_log if "knowledge_embeddings.chunk_id" in s and "knowledge_embeddings.vector" not in s] assert catalogue, "the semantic catalogue read was not found" assert all("knowledge_chunks.text" not in s for s in catalogue) def embed_pending(client): with SessionLocal() as db: adventure = db.get(models.Adventure, client.adv_id) settings = db.execute(select(models.Settings).where( models.Settings.user_id == client.user_id)).scalars().first() asyncio.run(embeddings.embed_pending(db, adventure, settings)) db.commit()