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:
@@ -7,9 +7,10 @@ from sqlalchemy.orm import Mapped, mapped_column
|
|||||||
from scribe.models import Base
|
from scribe.models import Base
|
||||||
|
|
||||||
# bge-small-en-v1.5 produces 384-dim unit-normalized vectors. The column is a
|
# bge-small-en-v1.5 produces 384-dim unit-normalized vectors. The column is a
|
||||||
# native pgvector `vector(384)` (see migration 0067) so similarity search runs
|
# native pgvector `vector(384)` (see migration 0067) so similarity is ranked in
|
||||||
# as an indexed `ORDER BY embedding <=> :q LIMIT k` in Postgres rather than a
|
# Postgres by `<=>` rather than by a full-table Python cosine scan. The semantic
|
||||||
# full-table Python cosine scan.
|
# searches rank the CALLER'S in-scope chunks exactly rather than through the
|
||||||
|
# HNSW index, which ranks before it filters (#4961; embeddings._rank_scoped).
|
||||||
EMBEDDING_DIM = 384
|
EMBEDDING_DIM = 384
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -856,6 +856,62 @@ def record_best_chunk(report: dict | None, chunks: dict[int, dict]) -> None:
|
|||||||
report["best_chunk"] = chunks
|
report["best_chunk"] = chunks
|
||||||
|
|
||||||
|
|
||||||
|
# --- scope first, then rank (#4958, #4961) -----------------------------------
|
||||||
|
#
|
||||||
|
# Every semantic search here asks the same question of a shared table: the
|
||||||
|
# nearest chunks AMONG THE ONES THIS CALLER MAY SEE. Ordered straight off a
|
||||||
|
# `*_embeddings` table, the planner answers a different one. It walks the HNSW
|
||||||
|
# index, takes about `hnsw.ef_search` (40) nearest chunks from every owner and
|
||||||
|
# project, and only then applies the scope — so an in-scope record ranked past
|
||||||
|
# the 40th chunk overall is silently gone, and on a shared install other
|
||||||
|
# users' records are what fill those 40. The searches fail open, so nothing
|
||||||
|
# reports it: the result is just shorter.
|
||||||
|
#
|
||||||
|
# The fix is one shape, shared by every search so it cannot be fixed in one
|
||||||
|
# and missed in the next (which is how #4958 left three behind). The scope is
|
||||||
|
# a MATERIALIZED CTE with the distance computed inside it; the index cannot
|
||||||
|
# order a materialized CTE, so the ranking over it is exact.
|
||||||
|
#
|
||||||
|
# What exact costs is one distance per in-scope chunk. Rules, milestones and
|
||||||
|
# Systems are hundreds of chunks. Notes are the big corpus, but the hook arms
|
||||||
|
# that run on every prompt are project-scoped, so the pass is over one
|
||||||
|
# project's chunks rather than the whole table; `retrieval_logs.duration_ms`
|
||||||
|
# on those arms is where that claim is checked after the deploy. If it ever
|
||||||
|
# stops holding, pgvector 0.8's `hnsw.iterative_scan` keeps the index and is
|
||||||
|
# the next step — gated on the extension version, because an unknown `hnsw.*`
|
||||||
|
# setting errors and these searches would read that error as "nothing found".
|
||||||
|
|
||||||
|
|
||||||
|
def _scoped_chunks(embedding_model, record_key, distance):
|
||||||
|
"""A select over an embedding table's chunks, each with its distance — the
|
||||||
|
caller joins its record and adds its scope, and `_rank_scoped` orders it.
|
||||||
|
|
||||||
|
The columns are labelled so every search's CTE carries the same four.
|
||||||
|
"""
|
||||||
|
return select(
|
||||||
|
record_key.label("record_id"),
|
||||||
|
distance.label("distance"),
|
||||||
|
embedding_model.chunk_index.label("chunk_index"),
|
||||||
|
embedding_model.chunk_text.label("chunk_text"),
|
||||||
|
).select_from(embedding_model)
|
||||||
|
|
||||||
|
|
||||||
|
def _rank_scoped(record_model, scoped_chunks, *, name: str, limit: int):
|
||||||
|
"""(record, distance, chunk_index, chunk_text) rows, nearest first, drawn
|
||||||
|
ONLY from the in-scope chunks — exactly, never through the index.
|
||||||
|
|
||||||
|
The row shape is what each search's collapse already unpacks, so a search
|
||||||
|
moves onto this without its callers or its mocks noticing.
|
||||||
|
"""
|
||||||
|
scoped = scoped_chunks.cte(name).prefix_with("MATERIALIZED")
|
||||||
|
return (
|
||||||
|
select(record_model, scoped.c.distance, scoped.c.chunk_index, scoped.c.chunk_text)
|
||||||
|
.join(scoped, scoped.c.record_id == record_model.id)
|
||||||
|
.order_by(scoped.c.distance)
|
||||||
|
.limit(limit)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def semantic_search_notes(
|
async def semantic_search_notes(
|
||||||
user_id: int,
|
user_id: int,
|
||||||
query: str,
|
query: str,
|
||||||
@@ -924,11 +980,14 @@ async def semantic_search_notes(
|
|||||||
a caller that forgets is wrong in the safe direction.
|
a caller that forgets is wrong in the safe direction.
|
||||||
|
|
||||||
Ranking and the top-k cut happen in Postgres via pgvector's cosine-distance
|
Ranking and the top-k cut happen in Postgres via pgvector's cosine-distance
|
||||||
operator (`<=>`, exposed as ``Vector.cosine_distance``) backed by the HNSW
|
operator (`<=>`, exposed as ``Vector.cosine_distance``), EXACTLY and over
|
||||||
index from migration 0067 — so this is an indexed ``ORDER BY ... LIMIT k``
|
the in-scope chunks only (#4961, see `_rank_scoped`). It used to be an
|
||||||
rather than a full-table scan. Cosine distance is ``1 - cosine_similarity``,
|
indexed ``ORDER BY ... LIMIT k`` straight off the HNSW index from
|
||||||
so a similarity floor of *threshold* is a distance ceiling of
|
migration 0067, which ranks before it filters and so could drop an
|
||||||
``1 - threshold`` and similarity is recovered as ``1 - distance``.
|
in-scope record that sat behind ~40 nearer ones the caller cannot see.
|
||||||
|
Cosine distance is ``1 - cosine_similarity``, so a similarity floor of
|
||||||
|
*threshold* is a distance ceiling of ``1 - threshold`` and similarity is
|
||||||
|
recovered as ``1 - distance``.
|
||||||
|
|
||||||
`demote_superseded` applies the supersession penalty (#278): a record a
|
`demote_superseded` applies the supersession penalty (#278): a record a
|
||||||
later note claims to have overtaken ranks below its equals. Callers asking
|
later note claims to have overtaken ranks below its equals. Callers asking
|
||||||
@@ -964,18 +1023,16 @@ async def semantic_search_notes(
|
|||||||
# Scope on Note, not NoteEmbedding.user_id: the embedding row belongs
|
# Scope on Note, not NoteEmbedding.user_id: the embedding row belongs
|
||||||
# to the note's owner, so filtering it would pin every scope to "own"
|
# to the note's owner, so filtering it would pin every scope to "own"
|
||||||
# and leave shared records unreachable by meaning.
|
# and leave shared records unreachable by meaning.
|
||||||
|
#
|
||||||
|
# Every filter below is the SCOPE, built into the chunks the
|
||||||
|
# ranking draws from rather than applied to what it returns
|
||||||
|
# (#4961) — see `_rank_scoped`.
|
||||||
stmt = (
|
stmt = (
|
||||||
# chunk_index/chunk_text ride along so the collapse below can
|
# chunk_index/chunk_text ride along so the collapse below can
|
||||||
# say WHICH passage matched. Without them the caller is left
|
# say WHICH passage matched. Without them the caller is left
|
||||||
# previewing the head of the body — a span this query has
|
# previewing the head of the body — a span this query has
|
||||||
# already determined is not why the record ranked (#4243).
|
# already determined is not why the record ranked (#4243).
|
||||||
select(
|
_scoped_chunks(NoteEmbedding, NoteEmbedding.note_id, distance)
|
||||||
Note,
|
|
||||||
distance.label("distance"),
|
|
||||||
NoteEmbedding.chunk_index,
|
|
||||||
NoteEmbedding.chunk_text,
|
|
||||||
)
|
|
||||||
.select_from(NoteEmbedding)
|
|
||||||
.join(Note, NoteEmbedding.note_id == Note.id)
|
.join(Note, NoteEmbedding.note_id == Note.id)
|
||||||
.where(
|
.where(
|
||||||
notes_visibility_clause(user_id, scope),
|
notes_visibility_clause(user_id, scope),
|
||||||
@@ -1032,24 +1089,22 @@ async def semantic_search_notes(
|
|||||||
# the results and the live record that should have replaced it was
|
# the results and the live record that should have replaced it was
|
||||||
# never fetched.
|
# never fetched.
|
||||||
#
|
#
|
||||||
# Ordering stays on RAW distance so pgvector's HNSW index still
|
# Ordering is on RAW distance; the penalty is applied after the
|
||||||
# serves it (migration 0067). Ordering by `distance + penalty`
|
# cut, over the over-fetched window. The cost of that, stated
|
||||||
# instead would be exact, and would turn an indexed top-k into a
|
# plainly: a live record outside the window cannot be promoted
|
||||||
# scan-and-sort of every embedded note.
|
# into the results. With a penalty far smaller than the window's
|
||||||
#
|
# score spread, that case needs the true answer to be more than
|
||||||
# The cost of that trade, stated plainly: a live record outside the
|
# _SUPERSESSION_OVERFETCH ranks down, which no observed query
|
||||||
# over-fetch window cannot be promoted into the results. With a
|
# comes close to.
|
||||||
# penalty far smaller than the window's score spread, that case
|
|
||||||
# needs the true answer to be more than _SUPERSESSION_OVERFETCH
|
|
||||||
# ranks down, which no observed query comes close to.
|
|
||||||
fetch = limit * _CHUNK_OVERFETCH * (
|
fetch = limit * _CHUNK_OVERFETCH * (
|
||||||
_SUPERSESSION_OVERFETCH if demote_superseded else 1
|
_SUPERSESSION_OVERFETCH if demote_superseded else 1
|
||||||
)
|
)
|
||||||
# NO threshold predicate — see the note above this function. The
|
# NO threshold predicate — see the note above this function. The
|
||||||
# bar is applied after the collapse, where the rejected scores can
|
# bar is applied after the collapse, where the rejected scores can
|
||||||
# still be seen.
|
# still be seen.
|
||||||
stmt = stmt.order_by(distance.asc()).limit(fetch)
|
rows = list((await session.execute(
|
||||||
rows = list((await session.execute(stmt)).all())
|
_rank_scoped(Note, stmt, name="scoped_note_chunks", limit=fetch)
|
||||||
|
)).all())
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Failed to query note embeddings", exc_info=True)
|
logger.warning("Failed to query note embeddings", exc_info=True)
|
||||||
return []
|
return []
|
||||||
@@ -1550,24 +1605,14 @@ async def semantic_search_rules(
|
|||||||
# function fails open.
|
# function fails open.
|
||||||
home = await rule_home(user_id, project_id, everywhere=everywhere)
|
home = await rule_home(user_id, project_id, everywhere=everywhere)
|
||||||
|
|
||||||
# SCOPE FIRST, THEN RANK — exactly (#4958). Ordered straight off
|
# SCOPE FIRST, THEN RANK — exactly (#4958; the shape is
|
||||||
# rule_embeddings, the planner walks the HNSW index, which hands back
|
# `_rank_scoped`'s). The home filter is built into the chunks the
|
||||||
# about `hnsw.ef_search` (40) nearest chunks from EVERY owner and
|
# ranking draws from: applied to what an HNSW walk returned, it ran
|
||||||
# project and only then applies the home filter: an in-scope rule
|
# after ~40 nearest chunks from every owner and project had been
|
||||||
# ranked past the 40th chunk overall is silently gone, and on a
|
# taken, and an in-scope rule ranked past them was silently gone.
|
||||||
# shared install other users' rules fill those 40. A MATERIALIZED
|
|
||||||
# CTE cannot be ordered through the index, so the distance is
|
|
||||||
# computed for each in-scope chunk and the ranking is exact. A
|
|
||||||
# rulebook is hundreds of chunks, not millions — exact is cheap here.
|
|
||||||
scoped = (
|
scoped = (
|
||||||
joined_to_homes(
|
joined_to_homes(
|
||||||
select(
|
_scoped_chunks(RuleEmbedding, RuleEmbedding.rule_id, distance)
|
||||||
RuleEmbedding.rule_id.label("rule_id"),
|
|
||||||
distance.label("distance"),
|
|
||||||
RuleEmbedding.chunk_index.label("chunk_index"),
|
|
||||||
RuleEmbedding.chunk_text.label("chunk_text"),
|
|
||||||
)
|
|
||||||
.select_from(RuleEmbedding)
|
|
||||||
.join(Rule, RuleEmbedding.rule_id == Rule.id)
|
.join(Rule, RuleEmbedding.rule_id == Rule.id)
|
||||||
)
|
)
|
||||||
.where(
|
.where(
|
||||||
@@ -1577,18 +1622,14 @@ async def semantic_search_rules(
|
|||||||
home,
|
home,
|
||||||
*( [Rule.kind == kind] if kind else [] ),
|
*( [Rule.kind == kind] if kind else [] ),
|
||||||
)
|
)
|
||||||
.cte("scoped_rule_chunks")
|
|
||||||
.prefix_with("MATERIALIZED")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
async with async_session() as session:
|
async with async_session() as session:
|
||||||
rows = (await session.execute(
|
rows = (await session.execute(
|
||||||
select(Rule, scoped.c.distance, scoped.c.chunk_index, scoped.c.chunk_text)
|
|
||||||
.join(scoped, scoped.c.rule_id == Rule.id)
|
|
||||||
# Overfetch so collapsing chunks to their best row still fills
|
# Overfetch so collapsing chunks to their best row still fills
|
||||||
# the page — the same reason the note search overfetches.
|
# the page — the same reason the note search overfetches.
|
||||||
.order_by(scoped.c.distance)
|
_rank_scoped(Rule, scoped, name="scoped_rule_chunks",
|
||||||
.limit(limit * _CHUNK_OVERFETCH)
|
limit=limit * _CHUNK_OVERFETCH)
|
||||||
)).all()
|
)).all()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Rule semantic search failed", exc_info=True)
|
logger.warning("Rule semantic search failed", exc_info=True)
|
||||||
@@ -1778,23 +1819,20 @@ async def semantic_search_milestones(
|
|||||||
scope = Milestone.project_id == project_id
|
scope = Milestone.project_id == project_id
|
||||||
else:
|
else:
|
||||||
scope = Milestone.user_id == user_id
|
scope = Milestone.user_id == user_id
|
||||||
async with async_session() as session:
|
# Scope first, then rank (#4961): see `_rank_scoped`.
|
||||||
rows = (await session.execute(
|
scoped = (
|
||||||
select(
|
_scoped_chunks(MilestoneEmbedding, MilestoneEmbedding.milestone_id, distance)
|
||||||
Milestone,
|
|
||||||
distance.label("distance"),
|
|
||||||
MilestoneEmbedding.chunk_index,
|
|
||||||
MilestoneEmbedding.chunk_text,
|
|
||||||
)
|
|
||||||
.select_from(MilestoneEmbedding)
|
|
||||||
.join(Milestone, MilestoneEmbedding.milestone_id == Milestone.id)
|
.join(Milestone, MilestoneEmbedding.milestone_id == Milestone.id)
|
||||||
.where(
|
.where(
|
||||||
scope,
|
scope,
|
||||||
Milestone.deleted_at.is_(None),
|
Milestone.deleted_at.is_(None),
|
||||||
*([Milestone.status == status] if status else []),
|
*([Milestone.status == status] if status else []),
|
||||||
)
|
)
|
||||||
.order_by(distance)
|
)
|
||||||
.limit(limit * _CHUNK_OVERFETCH)
|
async with async_session() as session:
|
||||||
|
rows = (await session.execute(
|
||||||
|
_rank_scoped(Milestone, scoped, name="scoped_milestone_chunks",
|
||||||
|
limit=limit * _CHUNK_OVERFETCH)
|
||||||
)).all()
|
)).all()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Milestone semantic search failed", exc_info=True)
|
logger.warning("Milestone semantic search failed", exc_info=True)
|
||||||
@@ -1953,23 +1991,20 @@ async def semantic_search_systems(
|
|||||||
scope = System.project_id == project_id
|
scope = System.project_id == project_id
|
||||||
else:
|
else:
|
||||||
scope = System.user_id == user_id
|
scope = System.user_id == user_id
|
||||||
async with async_session() as session:
|
# Scope first, then rank (#4961): see `_rank_scoped`.
|
||||||
rows = (await session.execute(
|
scoped = (
|
||||||
select(
|
_scoped_chunks(SystemEmbedding, SystemEmbedding.system_id, distance)
|
||||||
System,
|
|
||||||
distance.label("distance"),
|
|
||||||
SystemEmbedding.chunk_index,
|
|
||||||
SystemEmbedding.chunk_text,
|
|
||||||
)
|
|
||||||
.select_from(SystemEmbedding)
|
|
||||||
.join(System, SystemEmbedding.system_id == System.id)
|
.join(System, SystemEmbedding.system_id == System.id)
|
||||||
.where(
|
.where(
|
||||||
scope,
|
scope,
|
||||||
System.deleted_at.is_(None),
|
System.deleted_at.is_(None),
|
||||||
System.status != "archived",
|
System.status != "archived",
|
||||||
)
|
)
|
||||||
.order_by(distance)
|
)
|
||||||
.limit(limit * _CHUNK_OVERFETCH)
|
async with async_session() as session:
|
||||||
|
rows = (await session.execute(
|
||||||
|
_rank_scoped(System, scoped, name="scoped_system_chunks",
|
||||||
|
limit=limit * _CHUNK_OVERFETCH)
|
||||||
)).all()
|
)).all()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("System semantic search failed", exc_info=True)
|
logger.warning("System semantic search failed", exc_info=True)
|
||||||
|
|||||||
@@ -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