fix(retrieval): every semantic search scopes first, then ranks - the note, milestone and system searches join the rule search on one shared shape, _rank_scoped (#4961)
CI & Build / Python lint (push) Successful in 3s
CI & Build / Plugin hooks (push) Successful in 17s
CI & Build / TypeScript typecheck (push) Successful in 54s
CI & Build / integration (push) Successful in 1m0s
CI & Build / Python tests (push) Successful in 1m54s
CI & Build / Build & push image (push) Successful in 27s
CI & Build / Python lint (push) Successful in 3s
CI & Build / Plugin hooks (push) Successful in 17s
CI & Build / TypeScript typecheck (push) Successful in 54s
CI & Build / integration (push) Successful in 1m0s
CI & Build / Python tests (push) Successful in 1m54s
CI & Build / Build & push image (push) Successful in 27s
Ordered straight off a *_embeddings table, the planner walks the HNSW index, takes ~ef_search (40) nearest chunks across every owner and project, and only then applies the scope: an in-scope record behind 40 nearer ones the caller cannot see was silently dropped. #4958 fixed the rule search alone; the note, milestone and system searches kept the fault. _scoped_chunks builds the in-scope chunks with their distance; _rank_scoped materializes them as a CTE and ranks exactly. All four searches go through it, and the row shape is unchanged, so callers and mocks are untouched. Tests: a structural guard that every semantic_search_* ranks through _rank_scoped and orders nothing itself (with a replay of the old shape), a compiled-SQL check of the MATERIALIZED CTE, and integration crowd tests - 80 nearer out-of-scope records - for notes (another user; the reader own other project), milestones and systems. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,189 @@
|
||||
"""Real-Postgres crowd tests: an in-scope record is found however many nearer
|
||||
records sit outside the scope (#4961; the rule search's is #4958's, in
|
||||
test_integration_rule_scope).
|
||||
|
||||
Ordered straight off a `*_embeddings` table, the planner can walk the HNSW
|
||||
index, which takes ~`hnsw.ef_search` (40) nearest chunks from every owner and
|
||||
project and only then applies the scope. Each test here puts 80 out-of-scope
|
||||
records nearer the query than the reader's one in-scope record, so the old
|
||||
shape could not reach it when the index served the order. What a test cannot
|
||||
force is the plan — the shape guard in test_search_scope_shape is what fails
|
||||
on the old query whatever Postgres picks.
|
||||
|
||||
Rows are inserted directly with hand-made vectors and the embedder is stubbed,
|
||||
so the scope alone decides what comes back.
|
||||
"""
|
||||
import traceback
|
||||
import uuid
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import delete
|
||||
|
||||
from scribe.models import async_session
|
||||
from scribe.models.embedding import (
|
||||
EMBEDDING_DIM, MilestoneEmbedding, NoteEmbedding, SystemEmbedding,
|
||||
)
|
||||
from scribe.models.milestone import Milestone
|
||||
from scribe.models.note import Note
|
||||
from scribe.models.project import Project
|
||||
from scribe.models.system import System
|
||||
from scribe.models.user import User
|
||||
from scribe.services import embeddings as emb
|
||||
from scribe.services.embeddings import CHUNKER_VERSION, EMBEDDING_MODEL
|
||||
from tests.helpers import ensure_user
|
||||
|
||||
pytestmark = [pytest.mark.integration, pytest.mark.usefixtures("_dispose_engine", "_no_embedding")]
|
||||
|
||||
CROWD = 80
|
||||
QUERY_VEC = [1.0] + [0.0] * (EMBEDDING_DIM - 1)
|
||||
# Cosine 0.8 to the query: comfortably above the bar, and behind every crowd row.
|
||||
MINE_VEC = [0.8, 0.6] + [0.0] * (EMBEDDING_DIM - 2)
|
||||
|
||||
|
||||
def _near(i: int) -> list[float]:
|
||||
"""Cosine ~0.9996 to the query, each distinct so the graph is a graph."""
|
||||
vec = [1.0] + [0.0] * (EMBEDDING_DIM - 1)
|
||||
vec[1 + i] = 0.02
|
||||
return vec
|
||||
|
||||
|
||||
def _chunk(model, key: str, record_id: int, vec: list[float], **extra):
|
||||
return model(**{key: record_id}, chunk_index=0, embedding=vec,
|
||||
chunk_text=f"record {record_id}", chunker_version=CHUNKER_VERSION,
|
||||
embedding_model=EMBEDDING_MODEL, **extra)
|
||||
|
||||
|
||||
async def _found(search, user_id: int, **scope) -> set[int]:
|
||||
"""The ids a search finds. Every search fails open, so a query that raised
|
||||
reads as an empty result — the swallowed traceback is captured and named
|
||||
in the failure instead (#4958)."""
|
||||
raised: list[str] = []
|
||||
with patch("scribe.services.embeddings.get_embedding",
|
||||
AsyncMock(return_value=QUERY_VEC)), \
|
||||
patch("scribe.services.embeddings.logger") as log:
|
||||
log.warning.side_effect = lambda *_a, **_k: raised.append(traceback.format_exc())
|
||||
hits = await search(user_id, "anything", limit=10, threshold=0.5, **scope)
|
||||
assert not raised, "the search raised:\n" + "\n".join(raised)
|
||||
return {int(record.id) for _score, record in hits}
|
||||
|
||||
|
||||
async def _people(tag: str):
|
||||
"""A reader and a crowd, each with a project of their own."""
|
||||
async with async_session() as s:
|
||||
reader = await ensure_user(s, f"scope_reader_{tag}")
|
||||
crowd = await ensure_user(s, f"scope_crowd_{tag}")
|
||||
mine = Project(user_id=reader.id, title="Reader's project")
|
||||
other = Project(user_id=reader.id, title="Reader's other project")
|
||||
theirs = Project(user_id=crowd.id, title="Crowd's project")
|
||||
s.add_all([mine, other, theirs])
|
||||
await s.commit()
|
||||
return reader.id, crowd.id, mine.id, other.id, theirs.id
|
||||
|
||||
|
||||
async def _cleanup(*user_ids: int) -> None:
|
||||
# These vectors sit right beside the query every other vector test uses.
|
||||
async with async_session() as s:
|
||||
for model in (Note, Milestone, System):
|
||||
await s.execute(delete(model).where(model.user_id.in_(user_ids)))
|
||||
await s.execute(delete(Project).where(Project.user_id.in_(user_ids)))
|
||||
await s.execute(delete(User).where(User.id.in_(user_ids)))
|
||||
await s.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_note_is_found_behind_another_users_nearer_notes():
|
||||
reader, crowd, mine_p, _other, theirs_p = await _people(uuid.uuid4().hex[:8])
|
||||
try:
|
||||
async with async_session() as s:
|
||||
mine = Note(user_id=reader, project_id=mine_p, title="mine", body="mine")
|
||||
crowd_notes = [Note(user_id=crowd, project_id=theirs_p, title=f"crowd {i}",
|
||||
body="not the reader's") for i in range(CROWD)]
|
||||
s.add_all([mine, *crowd_notes])
|
||||
await s.flush()
|
||||
s.add(_chunk(NoteEmbedding, "note_id", mine.id, MINE_VEC, user_id=reader))
|
||||
s.add_all(_chunk(NoteEmbedding, "note_id", n.id, _near(i), user_id=crowd)
|
||||
for i, n in enumerate(crowd_notes))
|
||||
await s.commit()
|
||||
mine_id = mine.id
|
||||
|
||||
assert await _found(emb.semantic_search_notes, reader) == {mine_id}
|
||||
found = await _found(emb.semantic_search_notes, crowd)
|
||||
assert mine_id not in found and len(found) == 10
|
||||
finally:
|
||||
await _cleanup(reader, crowd)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_note_is_found_behind_the_readers_own_other_project():
|
||||
"""The single-user case, and the one the hook arms live in: a project-scoped
|
||||
search where the reader's OTHER projects hold the nearer chunks."""
|
||||
reader, crowd, mine_p, other_p, _theirs = await _people(uuid.uuid4().hex[:8])
|
||||
try:
|
||||
async with async_session() as s:
|
||||
mine = Note(user_id=reader, project_id=mine_p, title="mine", body="mine")
|
||||
elsewhere = [Note(user_id=reader, project_id=other_p, title=f"elsewhere {i}",
|
||||
body="another project") for i in range(CROWD)]
|
||||
s.add_all([mine, *elsewhere])
|
||||
await s.flush()
|
||||
s.add(_chunk(NoteEmbedding, "note_id", mine.id, MINE_VEC, user_id=reader))
|
||||
s.add_all(_chunk(NoteEmbedding, "note_id", n.id, _near(i), user_id=reader)
|
||||
for i, n in enumerate(elsewhere))
|
||||
await s.commit()
|
||||
mine_id = mine.id
|
||||
|
||||
assert await _found(emb.semantic_search_notes, reader,
|
||||
project_id=mine_p) == {mine_id}
|
||||
finally:
|
||||
await _cleanup(reader, crowd)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_milestone_is_found_behind_another_users_nearer_plans():
|
||||
reader, crowd, mine_p, _other, theirs_p = await _people(uuid.uuid4().hex[:8])
|
||||
try:
|
||||
async with async_session() as s:
|
||||
mine = Milestone(user_id=reader, project_id=mine_p, title="my plan")
|
||||
plans = [Milestone(user_id=crowd, project_id=theirs_p, title=f"plan {i}")
|
||||
for i in range(CROWD)]
|
||||
s.add_all([mine, *plans])
|
||||
await s.flush()
|
||||
s.add(_chunk(MilestoneEmbedding, "milestone_id", mine.id, MINE_VEC))
|
||||
s.add_all(_chunk(MilestoneEmbedding, "milestone_id", m.id, _near(i))
|
||||
for i, m in enumerate(plans))
|
||||
await s.commit()
|
||||
mine_id = mine.id
|
||||
|
||||
assert await _found(emb.semantic_search_milestones, reader) == {mine_id}
|
||||
assert await _found(emb.semantic_search_milestones, reader,
|
||||
project_id=mine_p) == {mine_id}
|
||||
found = await _found(emb.semantic_search_milestones, crowd)
|
||||
assert mine_id not in found and len(found) == 10
|
||||
finally:
|
||||
await _cleanup(reader, crowd)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_system_is_found_behind_another_users_nearer_charters():
|
||||
reader, crowd, mine_p, _other, theirs_p = await _people(uuid.uuid4().hex[:8])
|
||||
try:
|
||||
async with async_session() as s:
|
||||
mine = System(user_id=reader, project_id=mine_p, name="mine",
|
||||
description="the reader's area")
|
||||
areas = [System(user_id=crowd, project_id=theirs_p, name=f"area {i}",
|
||||
description="someone else's area") for i in range(CROWD)]
|
||||
s.add_all([mine, *areas])
|
||||
await s.flush()
|
||||
s.add(_chunk(SystemEmbedding, "system_id", mine.id, MINE_VEC))
|
||||
s.add_all(_chunk(SystemEmbedding, "system_id", a.id, _near(i))
|
||||
for i, a in enumerate(areas))
|
||||
await s.commit()
|
||||
mine_id = mine.id
|
||||
|
||||
assert await _found(emb.semantic_search_systems, reader) == {mine_id}
|
||||
assert await _found(emb.semantic_search_systems, reader,
|
||||
project_id=mine_p) == {mine_id}
|
||||
found = await _found(emb.semantic_search_systems, crowd)
|
||||
assert mine_id not in found and len(found) == 10
|
||||
finally:
|
||||
await _cleanup(reader, crowd)
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Every semantic search scopes first, then ranks (#4958, #4961).
|
||||
|
||||
Ordered straight off a `*_embeddings` table, the planner walks the HNSW index:
|
||||
it takes ~`hnsw.ef_search` (40) nearest chunks from every owner and project
|
||||
and filters them afterwards, so an in-scope record behind 40 nearer ones the
|
||||
caller cannot see is silently dropped. The fix is one shape, `_rank_scoped`,
|
||||
and these pin that every search goes through it — the rule search was fixed
|
||||
alone first, and the other three kept the fault because nothing said they
|
||||
were the same query.
|
||||
|
||||
The integration lane's crowd tests show the behaviour on a real index; these
|
||||
are the guard that fails on the old shape whatever plan Postgres picks.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from scribe.models.embedding import NoteEmbedding
|
||||
from scribe.models.note import Note
|
||||
from scribe.services import embeddings as emb
|
||||
from tests.helpers import compiled_sql
|
||||
|
||||
_SOURCE = Path("src/scribe/services/embeddings.py")
|
||||
|
||||
|
||||
def _searches() -> dict[str, ast.AsyncFunctionDef]:
|
||||
tree = ast.parse(_SOURCE.read_text())
|
||||
return {
|
||||
node.name: node for node in tree.body
|
||||
if isinstance(node, ast.AsyncFunctionDef)
|
||||
and node.name.startswith("semantic_search_")
|
||||
}
|
||||
|
||||
|
||||
def _calls(fn: ast.AST) -> list[str]:
|
||||
return [
|
||||
getattr(n.func, "attr", None) or getattr(n.func, "id", None)
|
||||
for n in ast.walk(fn) if isinstance(n, ast.Call)
|
||||
]
|
||||
|
||||
|
||||
def test_every_semantic_search_ranks_through_the_scoped_shape():
|
||||
searches = _searches()
|
||||
# The four corpora with an embedding table: if one is renamed or a fifth
|
||||
# is added, this has to be looked at rather than silently passing.
|
||||
assert set(searches) == {
|
||||
"semantic_search_notes", "semantic_search_rules",
|
||||
"semantic_search_milestones", "semantic_search_systems",
|
||||
}
|
||||
for name, fn in searches.items():
|
||||
calls = _calls(fn)
|
||||
assert "_rank_scoped" in calls, f"{name} does not rank through _rank_scoped"
|
||||
assert "_scoped_chunks" in calls, f"{name} does not build its scope as chunks"
|
||||
# Its own ORDER BY is the old shape: a distance ordered on the
|
||||
# embedding table, which the index serves before the scope applies.
|
||||
assert "order_by" not in calls, f"{name} orders by distance itself"
|
||||
|
||||
|
||||
def test_the_old_shape_would_be_caught():
|
||||
"""Rule 167: replay the guard against the query every search used to run."""
|
||||
old = ast.parse(
|
||||
"async def semantic_search_x():\n"
|
||||
" rows = await session.execute(select(Note, distance).select_from(E)"
|
||||
".join(Note, E.note_id == Note.id).where(scope)"
|
||||
".order_by(distance).limit(k))\n"
|
||||
).body[0]
|
||||
calls = _calls(old)
|
||||
assert "_rank_scoped" not in calls and "order_by" in calls
|
||||
|
||||
|
||||
def test_the_scope_is_materialized_and_the_order_is_over_it():
|
||||
# Column against column, so there is no vector literal to inline: the
|
||||
# shape is the subject here, not the query.
|
||||
distance = NoteEmbedding.embedding.cosine_distance(NoteEmbedding.embedding)
|
||||
scoped = (
|
||||
emb._scoped_chunks(NoteEmbedding, NoteEmbedding.note_id, distance)
|
||||
.join(Note, NoteEmbedding.note_id == Note.id)
|
||||
.where(Note.user_id == 1)
|
||||
)
|
||||
sql = compiled_sql(
|
||||
emb._rank_scoped(Note, scoped, name="scoped_x", limit=5),
|
||||
dialect=postgresql.dialect(),
|
||||
)
|
||||
# MATERIALIZED is what keeps the planner from inlining the CTE and walking
|
||||
# the index again; the ORDER BY must name the CTE's column, not the table's.
|
||||
assert "scoped_x AS MATERIALIZED" in sql
|
||||
assert "ORDER BY scoped_x.distance" in sql
|
||||
assert "note_embeddings.embedding <=>" in sql.split("ORDER BY")[0]
|
||||
# The scope sits inside the CTE, where the ranking draws from.
|
||||
cte_body = sql.split("AS MATERIALIZED", 1)[1].split("SELECT notes", 1)[0]
|
||||
assert "notes.user_id" in cte_body
|
||||
Reference in New Issue
Block a user