Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 807f478cac | |||
| 513019786e |
@@ -165,7 +165,9 @@ jobs:
|
||||
SECRET_KEY: ci_integration_placeholder
|
||||
services:
|
||||
postgres:
|
||||
image: postgres:16-alpine
|
||||
# pgvector image so `alembic upgrade head` can run migration 0067
|
||||
# (CREATE EXTENSION vector). PG17 — matches the prod/quickstart image.
|
||||
image: pgvector/pgvector:pg17
|
||||
env:
|
||||
POSTGRES_USER: scribe
|
||||
POSTGRES_PASSWORD: ci_integration
|
||||
@@ -189,7 +191,7 @@ jobs:
|
||||
set -eux
|
||||
echo "=== container landscape (diagnostic for the name filter) ==="
|
||||
docker ps -a --format '{{.ID}} {{.Image}} -> {{.Names}}'
|
||||
PG=$(docker ps --filter "name=integration" --filter "ancestor=postgres:16-alpine" -q | head -n1)
|
||||
PG=$(docker ps --filter "name=integration" --filter "ancestor=pgvector/pgvector:pg17" -q | head -n1)
|
||||
test -n "$PG"
|
||||
PG_IP=$(docker inspect -f '{{range .NetworkSettings.Networks}}{{.IPAddress}}{{end}}' "$PG")
|
||||
test -n "$PG_IP"
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
"""pgvector: note_embeddings.embedding JSONB -> vector(384) + HNSW index
|
||||
|
||||
Revision ID: 0067
|
||||
Revises: 0066
|
||||
Create Date: 2026-06-22
|
||||
|
||||
Moves semantic search off the full-table Python cosine scan onto a native
|
||||
pgvector column so ranking + top-k run as an indexed `ORDER BY embedding <=> :q
|
||||
LIMIT k` in Postgres (see services/embeddings.semantic_search_notes).
|
||||
|
||||
Requires a Postgres image that bundles the `vector` extension — the stack moved
|
||||
from postgres:16-alpine to pgvector/pgvector:pg16 in the same change (compose +
|
||||
CI). `CREATE EXTENSION IF NOT EXISTS vector` below is the in-db half.
|
||||
|
||||
Embeddings are DERIVED data (regenerated from note text by
|
||||
backfill_note_embeddings at startup), so this migration is free to drop any row
|
||||
it can't cleanly convert: only rows whose stored JSONB array is exactly 384-dim
|
||||
are carried over (guarding against stale vectors from an earlier model — the
|
||||
same mixed-dim hazard _cosine_similarity defended against). Dropped rows are
|
||||
re-embedded on next boot.
|
||||
"""
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision = "0067"
|
||||
down_revision = "0066"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
|
||||
|
||||
# New native-vector column, populated only from cleanly-convertible rows.
|
||||
# A JSONB array like [0.1, 0.2, ...] renders to text that is exactly
|
||||
# pgvector's input literal, so (embedding::text)::vector is a direct cast.
|
||||
op.execute("ALTER TABLE note_embeddings ADD COLUMN embedding_vec vector(384)")
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE note_embeddings
|
||||
SET embedding_vec = (embedding::text)::vector
|
||||
WHERE jsonb_array_length(embedding) = 384
|
||||
"""
|
||||
)
|
||||
# Stale-dim rows (couldn't convert) are derived data — drop and let the
|
||||
# startup backfill regenerate them at the current dimension.
|
||||
op.execute("DELETE FROM note_embeddings WHERE embedding_vec IS NULL")
|
||||
|
||||
op.execute("ALTER TABLE note_embeddings ALTER COLUMN embedding_vec SET NOT NULL")
|
||||
op.execute("ALTER TABLE note_embeddings DROP COLUMN embedding")
|
||||
op.execute("ALTER TABLE note_embeddings RENAME COLUMN embedding_vec TO embedding")
|
||||
|
||||
# HNSW index for cosine distance — matches Vector.cosine_distance (`<=>`).
|
||||
op.execute(
|
||||
"""
|
||||
CREATE INDEX ix_note_embeddings_embedding_hnsw
|
||||
ON note_embeddings
|
||||
USING hnsw (embedding vector_cosine_ops)
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Back to JSONB. pgvector renders a vector to a text literal that is a valid
|
||||
# JSON array, so the reverse cast is symmetric. The `vector` extension is
|
||||
# intentionally left installed (other objects may depend on it; dropping an
|
||||
# extension is the riskier, rarely-wanted direction).
|
||||
op.execute("DROP INDEX IF EXISTS ix_note_embeddings_embedding_hnsw")
|
||||
op.execute("ALTER TABLE note_embeddings ADD COLUMN embedding_json jsonb")
|
||||
op.execute("UPDATE note_embeddings SET embedding_json = (embedding::text)::jsonb")
|
||||
op.execute("ALTER TABLE note_embeddings ALTER COLUMN embedding_json SET NOT NULL")
|
||||
op.execute("ALTER TABLE note_embeddings DROP COLUMN embedding")
|
||||
op.execute("ALTER TABLE note_embeddings RENAME COLUMN embedding_json TO embedding")
|
||||
@@ -0,0 +1,60 @@
|
||||
"""retrieval_logs: per-call semantic-retrieval telemetry for KB-injection tuning
|
||||
|
||||
Revision ID: 0068
|
||||
Revises: 0067
|
||||
Create Date: 2026-06-22
|
||||
|
||||
One row per semantic-retrieval call (MCP search tool, REST search route, and —
|
||||
once it lands — the title-first auto-inject path). Captures the effective query
|
||||
params and the score distribution of the results so the similarity threshold
|
||||
and top-k can be tuned from real usage. FK-free on user_id (mirrors app_logs):
|
||||
telemetry should outlive the row it describes.
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
|
||||
revision = "0068"
|
||||
down_revision = "0067"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"retrieval_logs",
|
||||
sa.Column("id", sa.Integer(), primary_key=True),
|
||||
sa.Column(
|
||||
"created_at", sa.DateTime(timezone=True), nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.Column("user_id", sa.Integer(), nullable=True),
|
||||
sa.Column("source", sa.Text(), nullable=False),
|
||||
sa.Column("query", sa.Text(), nullable=True),
|
||||
sa.Column("threshold", sa.Float(), nullable=True),
|
||||
sa.Column("limit_n", sa.Integer(), nullable=True),
|
||||
sa.Column("project_id", sa.Integer(), nullable=True),
|
||||
sa.Column("is_task", sa.Boolean(), nullable=True),
|
||||
sa.Column("result_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("top_score", sa.Float(), nullable=True),
|
||||
sa.Column("min_score", sa.Float(), nullable=True),
|
||||
sa.Column("result_ids", JSONB(), nullable=False, server_default=sa.text("'[]'::jsonb")),
|
||||
sa.Column("duration_ms", sa.Float(), nullable=True),
|
||||
)
|
||||
op.create_index("ix_retrieval_logs_created_at", "retrieval_logs", ["created_at"])
|
||||
op.create_index("ix_retrieval_logs_user_id", "retrieval_logs", ["user_id"])
|
||||
op.create_index("ix_retrieval_logs_source", "retrieval_logs", ["source"])
|
||||
op.create_index(
|
||||
"ix_retrieval_logs_source_created_at",
|
||||
"retrieval_logs",
|
||||
["source", sa.text("created_at DESC")],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_retrieval_logs_source_created_at", table_name="retrieval_logs")
|
||||
op.drop_index("ix_retrieval_logs_source", table_name="retrieval_logs")
|
||||
op.drop_index("ix_retrieval_logs_user_id", table_name="retrieval_logs")
|
||||
op.drop_index("ix_retrieval_logs_created_at", table_name="retrieval_logs")
|
||||
op.drop_table("retrieval_logs")
|
||||
@@ -21,7 +21,11 @@ services:
|
||||
max_attempts: 5
|
||||
|
||||
db:
|
||||
image: postgres:16-alpine
|
||||
# pgvector image (Debian/glibc, PG17) — bundles the `vector` extension that
|
||||
# migration 0067 enables. Moved off postgres:16-alpine via logical
|
||||
# dump/restore (which doubles as the PG16->PG17 major upgrade); see the
|
||||
# TRANSITION runbook in the PR.
|
||||
image: pgvector/pgvector:pg17
|
||||
stop_grace_period: 120s
|
||||
volumes:
|
||||
- pgdata:/var/lib/postgresql/data
|
||||
|
||||
@@ -35,7 +35,8 @@ services:
|
||||
start_period: 30s
|
||||
|
||||
db:
|
||||
image: postgres:16-alpine
|
||||
# pgvector image (PG17) — bundles the `vector` extension (migration 0067).
|
||||
image: pgvector/pgvector:pg17
|
||||
stop_grace_period: 120s
|
||||
volumes:
|
||||
- pgdata:/var/lib/postgresql/data
|
||||
|
||||
@@ -21,6 +21,7 @@ dependencies = [
|
||||
"APScheduler>=3.10,<4.0",
|
||||
"mcp[cli]>=1.0",
|
||||
"fastembed>=0.4",
|
||||
"pgvector>=0.3",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
@@ -7,8 +7,11 @@ working. Differences from fable-mcp:
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from scribe.mcp._context import current_user_id
|
||||
from scribe.services.embeddings import semantic_search_notes
|
||||
from scribe.services.embeddings import DEFAULT_SIMILARITY_THRESHOLD, semantic_search_notes
|
||||
from scribe.services.retrieval_telemetry import record_retrieval
|
||||
|
||||
|
||||
async def search(
|
||||
@@ -43,10 +46,17 @@ async def search(
|
||||
uid = current_user_id()
|
||||
limit = max(1, min(limit, 50))
|
||||
is_task = {"note": False, "task": True}.get(content_type) # None => any
|
||||
t0 = time.perf_counter()
|
||||
raw = await semantic_search_notes(
|
||||
uid, q, limit=limit, is_task=is_task,
|
||||
project_id=project_id or None,
|
||||
)
|
||||
record_retrieval(
|
||||
user_id=uid, source="mcp_search", query=q,
|
||||
threshold=DEFAULT_SIMILARITY_THRESHOLD, limit=limit,
|
||||
project_id=project_id or None, is_task=is_task, results=raw,
|
||||
duration_ms=(time.perf_counter() - t0) * 1000.0,
|
||||
)
|
||||
return {
|
||||
"results": [
|
||||
{
|
||||
|
||||
@@ -26,6 +26,7 @@ from scribe.models.app_log import AppLog # noqa: E402, F401
|
||||
from scribe.models.password_reset import PasswordResetToken # noqa: E402, F401
|
||||
from scribe.models.invitation import InvitationToken # noqa: E402, F401
|
||||
from scribe.models.embedding import NoteEmbedding # noqa: E402, F401
|
||||
from scribe.models.retrieval_log import RetrievalLog # noqa: E402, F401
|
||||
from scribe.models.project import Project # noqa: E402, F401
|
||||
from scribe.models.event import Event # noqa: E402, F401
|
||||
from scribe.models.milestone import Milestone # noqa: E402, F401
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from pgvector.sqlalchemy import Vector
|
||||
from sqlalchemy import DateTime, ForeignKey, Integer
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from scribe.models import Base
|
||||
|
||||
# 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
|
||||
# as an indexed `ORDER BY embedding <=> :q LIMIT k` in Postgres rather than a
|
||||
# full-table Python cosine scan.
|
||||
EMBEDDING_DIM = 384
|
||||
|
||||
|
||||
class NoteEmbedding(Base):
|
||||
"""Stores the embedding vector for a note, used for semantic search."""
|
||||
@@ -18,7 +24,7 @@ class NoteEmbedding(Base):
|
||||
primary_key=True,
|
||||
)
|
||||
user_id: Mapped[int] = mapped_column(Integer, nullable=False, index=True)
|
||||
embedding: Mapped[list] = mapped_column(JSONB, nullable=False)
|
||||
embedding: Mapped[list] = mapped_column(Vector(EMBEDDING_DIM), nullable=False)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
default=lambda: datetime.now(timezone.utc),
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, Float, Index, Integer, Text
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from scribe.models import Base
|
||||
|
||||
|
||||
class RetrievalLog(Base):
|
||||
"""One row per semantic-retrieval call, for KB-injection tuning.
|
||||
|
||||
Captures what a query asked for, what came back, and the score
|
||||
distribution of the results — the empirical basis for tuning the
|
||||
similarity threshold and top-k per surface. `result_ids` holds the ranked
|
||||
hits (id + score + rank) so a later pass can correlate "what we surfaced"
|
||||
against "what the agent then fetched/referenced".
|
||||
|
||||
Deliberately FK-free on user_id (mirrors AppLog): telemetry should outlive
|
||||
the row it describes, and a deleted user shouldn't cascade away history.
|
||||
"""
|
||||
|
||||
__tablename__ = "retrieval_logs"
|
||||
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
||||
)
|
||||
user_id: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
# Retrieval surface: 'mcp_search' | 'rest_search' | 'auto_inject' | ...
|
||||
source: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
query: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
# Effective parameters actually used for this call.
|
||||
threshold: Mapped[float | None] = mapped_column(Float, nullable=True)
|
||||
limit_n: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
project_id: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
# The content-type filter as passed to semantic_search_notes: True=tasks,
|
||||
# False=notes, NULL=any.
|
||||
is_task: Mapped[bool | None] = mapped_column(Boolean, nullable=True)
|
||||
result_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
top_score: Mapped[float | None] = mapped_column(Float, nullable=True)
|
||||
min_score: Mapped[float | None] = mapped_column(Float, nullable=True)
|
||||
# [{"id": int, "score": float, "rank": int}, ...], highest-first.
|
||||
result_ids: Mapped[list] = mapped_column(JSONB, nullable=False, default=list)
|
||||
duration_ms: Mapped[float | None] = mapped_column(Float, nullable=True)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_retrieval_logs_created_at", "created_at"),
|
||||
Index("ix_retrieval_logs_user_id", "user_id"),
|
||||
Index("ix_retrieval_logs_source", "source"),
|
||||
Index("ix_retrieval_logs_source_created_at", "source", created_at.desc()),
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"id": self.id,
|
||||
"created_at": self.created_at.isoformat() if self.created_at else None,
|
||||
"user_id": self.user_id,
|
||||
"source": self.source,
|
||||
"query": self.query,
|
||||
"threshold": self.threshold,
|
||||
"limit_n": self.limit_n,
|
||||
"project_id": self.project_id,
|
||||
"is_task": self.is_task,
|
||||
"result_count": self.result_count,
|
||||
"top_score": self.top_score,
|
||||
"min_score": self.min_score,
|
||||
"result_ids": self.result_ids,
|
||||
"duration_ms": self.duration_ms,
|
||||
}
|
||||
@@ -1,7 +1,14 @@
|
||||
import time
|
||||
|
||||
from quart import Blueprint, jsonify, request
|
||||
|
||||
from scribe.auth import login_required, get_current_user_id
|
||||
from scribe.services.embeddings import semantic_search_notes
|
||||
from scribe.services.retrieval_telemetry import record_retrieval
|
||||
|
||||
# This route searches with a looser floor than the MCP tool default — it powers
|
||||
# an interactive feed where loosely-related hits still have value.
|
||||
_REST_SEARCH_THRESHOLD = 0.3
|
||||
|
||||
search_bp = Blueprint("search", __name__, url_prefix="/api/search")
|
||||
|
||||
@@ -27,8 +34,15 @@ async def search_route():
|
||||
limit = min(request.args.get("limit", 10, type=int), 50)
|
||||
is_task = _content_type_to_is_task(content_type)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
results = await semantic_search_notes(
|
||||
uid, q, limit=limit, is_task=is_task, threshold=0.3
|
||||
uid, q, limit=limit, is_task=is_task, threshold=_REST_SEARCH_THRESHOLD
|
||||
)
|
||||
record_retrieval(
|
||||
user_id=uid, source="rest_search", query=q,
|
||||
threshold=_REST_SEARCH_THRESHOLD, limit=limit,
|
||||
project_id=None, is_task=is_task, results=results,
|
||||
duration_ms=(time.perf_counter() - t0) * 1000.0,
|
||||
)
|
||||
return jsonify({
|
||||
"results": [
|
||||
|
||||
@@ -28,6 +28,10 @@ logger = logging.getLogger(__name__)
|
||||
# loosely-related results that pad the sidebar without adding real value.
|
||||
_SIMILARITY_THRESHOLD = 0.45
|
||||
|
||||
# Public alias so callers (and telemetry) can record the effective default
|
||||
# threshold without reaching for the underscored name.
|
||||
DEFAULT_SIMILARITY_THRESHOLD = _SIMILARITY_THRESHOLD
|
||||
|
||||
_MODEL_NAME = "BAAI/bge-small-en-v1.5"
|
||||
_CACHE_DIR = os.environ.get("FASTEMBED_CACHE_DIR", "/data/fastembed-cache")
|
||||
|
||||
@@ -115,6 +119,14 @@ async def semantic_search_notes(
|
||||
|
||||
Scores are cosine similarities in [-1, 1]; only notes at or above
|
||||
*threshold* are returned, sorted highest-first.
|
||||
|
||||
Ranking and the top-k cut happen in Postgres via pgvector's cosine-distance
|
||||
operator (`<=>`, exposed as ``Vector.cosine_distance``) backed by the HNSW
|
||||
index from migration 0067 — so this is an indexed ``ORDER BY ... LIMIT k``
|
||||
rather than a full-table scan. 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``.
|
||||
|
||||
Returns an empty list if the embedder is unavailable or on any error.
|
||||
"""
|
||||
if not query or not query.strip():
|
||||
@@ -125,10 +137,17 @@ async def semantic_search_notes(
|
||||
logger.debug("Semantic search skipped — embedder unavailable")
|
||||
return []
|
||||
|
||||
# Distance ceiling equivalent to the similarity floor. Clamp to the valid
|
||||
# cosine-distance range [0, 2] so a threshold of, say, -1 doesn't produce a
|
||||
# nonsensical ceiling.
|
||||
max_distance = min(2.0, max(0.0, 1.0 - threshold))
|
||||
distance = NoteEmbedding.embedding.cosine_distance(query_vec)
|
||||
|
||||
try:
|
||||
async with async_session() as session:
|
||||
stmt = (
|
||||
select(NoteEmbedding, Note)
|
||||
select(Note, distance.label("distance"))
|
||||
.select_from(NoteEmbedding)
|
||||
.join(Note, NoteEmbedding.note_id == Note.id)
|
||||
.where(NoteEmbedding.user_id == user_id, Note.deleted_at.is_(None))
|
||||
)
|
||||
@@ -142,30 +161,14 @@ async def semantic_search_notes(
|
||||
stmt = stmt.where(Note.status.is_(None))
|
||||
if exclude_ids:
|
||||
stmt = stmt.where(NoteEmbedding.note_id.notin_(exclude_ids))
|
||||
stmt = stmt.where(distance <= max_distance).order_by(distance.asc()).limit(limit)
|
||||
rows = list((await session.execute(stmt)).all())
|
||||
except Exception:
|
||||
logger.warning("Failed to query note embeddings", exc_info=True)
|
||||
return []
|
||||
|
||||
if not rows:
|
||||
return []
|
||||
|
||||
def _score() -> list[tuple[float, Note]]:
|
||||
out: list[tuple[float, Note]] = []
|
||||
for ne, note in rows:
|
||||
try:
|
||||
sim = _cosine_similarity(query_vec, ne.embedding)
|
||||
except Exception:
|
||||
continue
|
||||
if sim >= threshold:
|
||||
out.append((sim, note))
|
||||
out.sort(key=lambda x: x[0], reverse=True)
|
||||
return out[:limit]
|
||||
|
||||
# Offload the O(rows) cosine scoring off the event loop so a large corpus
|
||||
# doesn't stall other requests while ranking. Results are unchanged; the
|
||||
# real scaling fix (ORDER BY / LIMIT in pgvector) is a separate effort.
|
||||
return await asyncio.to_thread(_score)
|
||||
# Recover similarity (1 - distance) and preserve the highest-first contract.
|
||||
return [(1.0 - float(dist), note) for note, dist in rows]
|
||||
|
||||
|
||||
async def backfill_note_embeddings() -> None:
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
"""Retrieval telemetry — one RetrievalLog row per semantic-retrieval call.
|
||||
|
||||
This is the empirical basis for KB-injection tuning: it records what each query
|
||||
asked for, the score distribution of what came back, and the effective params,
|
||||
so the similarity threshold and top-k can be tuned from data rather than guessed.
|
||||
|
||||
Design notes:
|
||||
- Fire-and-forget, mirroring upsert_note_embedding: `record_retrieval` extracts
|
||||
the primitives it needs SYNCHRONOUSLY (while the caller's Note objects are
|
||||
still valid) and schedules the DB insert as a background task, so logging
|
||||
never adds latency to — or can break — the search response.
|
||||
- Result objects are reduced to {id, score, rank} before scheduling; the
|
||||
background writer touches only plain data, never a possibly-detached ORM row.
|
||||
- Every failure path is swallowed: telemetry must never take down retrieval.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
|
||||
from scribe.models import async_session
|
||||
from scribe.models.note import Note
|
||||
from scribe.models.retrieval_log import RetrievalLog
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _build_payload(
|
||||
*,
|
||||
user_id: int | None,
|
||||
source: str,
|
||||
query: str | None,
|
||||
threshold: float | None,
|
||||
limit: int | None,
|
||||
project_id: int | None,
|
||||
is_task: bool | None,
|
||||
results: list[tuple[float, Note]],
|
||||
duration_ms: float | None,
|
||||
) -> dict:
|
||||
"""Reduce a retrieval call to a flat, JSON-safe RetrievalLog payload.
|
||||
|
||||
Pure and synchronous (no DB, no event loop) so it is unit-testable and safe
|
||||
to run inline before scheduling the write. `results` is the
|
||||
`(score, Note)` list from semantic_search_notes, already highest-first.
|
||||
"""
|
||||
items = [
|
||||
{"id": int(note.id), "score": round(float(score), 5), "rank": rank}
|
||||
for rank, (score, note) in enumerate(results)
|
||||
]
|
||||
scores = [it["score"] for it in items]
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"source": source,
|
||||
"query": query,
|
||||
"threshold": threshold,
|
||||
"limit_n": limit,
|
||||
"project_id": project_id,
|
||||
"is_task": is_task,
|
||||
"result_count": len(items),
|
||||
"top_score": (scores[0] if scores else None),
|
||||
"min_score": (scores[-1] if scores else None),
|
||||
"result_ids": items,
|
||||
"duration_ms": (round(duration_ms, 2) if duration_ms is not None else None),
|
||||
}
|
||||
|
||||
|
||||
async def _insert_retrieval_log(payload: dict) -> None:
|
||||
"""Persist one RetrievalLog row. Best-effort: all errors are swallowed."""
|
||||
try:
|
||||
async with async_session() as session:
|
||||
session.add(RetrievalLog(**payload))
|
||||
await session.commit()
|
||||
except Exception:
|
||||
logger.debug("retrieval telemetry write skipped", exc_info=True)
|
||||
|
||||
|
||||
def record_retrieval(
|
||||
*,
|
||||
user_id: int | None,
|
||||
source: str,
|
||||
query: str | None,
|
||||
threshold: float | None,
|
||||
limit: int | None,
|
||||
project_id: int | None,
|
||||
is_task: bool | None,
|
||||
results: list[tuple[float, Note]],
|
||||
duration_ms: float | None = None,
|
||||
) -> None:
|
||||
"""Fire-and-forget: record one retrieval call.
|
||||
|
||||
Builds the payload inline (synchronously) then schedules the insert so the
|
||||
caller returns immediately. Never raises — telemetry must not affect search.
|
||||
"""
|
||||
try:
|
||||
payload = _build_payload(
|
||||
user_id=user_id,
|
||||
source=source,
|
||||
query=query,
|
||||
threshold=threshold,
|
||||
limit=limit,
|
||||
project_id=project_id,
|
||||
is_task=is_task,
|
||||
results=results,
|
||||
duration_ms=duration_ms,
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("retrieval telemetry payload build failed", exc_info=True)
|
||||
return
|
||||
|
||||
try:
|
||||
asyncio.get_running_loop().create_task(_insert_retrieval_log(payload))
|
||||
except RuntimeError:
|
||||
# No running loop (e.g. called from sync context outside the app) —
|
||||
# skip rather than block. The app paths always run on the loop.
|
||||
logger.debug("retrieval telemetry skipped — no running event loop")
|
||||
@@ -0,0 +1,101 @@
|
||||
"""Real-Postgres integration test for pgvector semantic search.
|
||||
|
||||
Runs only in the CI integration lane (real Postgres + `vector` extension +
|
||||
schema built by `alembic upgrade head`, which includes migration 0067). This
|
||||
exercises what the unit mocks cannot: the native `vector(384)` column, the
|
||||
`<=>` cosine-distance operator behind `Vector.cosine_distance`, the HNSW index,
|
||||
and the distance->similarity recovery in `semantic_search_notes`.
|
||||
|
||||
The embedder itself is stubbed (get_embedding is patched) so the test does not
|
||||
depend on downloading the fastembed model — only the Postgres/pgvector path is
|
||||
under test.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy import delete
|
||||
|
||||
from scribe.models import async_session, engine
|
||||
from scribe.models.embedding import EMBEDDING_DIM, NoteEmbedding
|
||||
from scribe.models.note import Note
|
||||
from scribe.models.user import User
|
||||
from scribe.services.embeddings import semantic_search_notes
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
|
||||
def _vec(*nonzero_first):
|
||||
"""A 384-dim vector with the given leading values, zero-padded."""
|
||||
v = list(nonzero_first) + [0.0] * (EMBEDDING_DIM - len(nonzero_first))
|
||||
return v[:EMBEDDING_DIM]
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(autouse=True)
|
||||
async def _dispose_engine():
|
||||
# Per-loop pool: dispose after each test (see test_integration_db_maintenance).
|
||||
yield
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def seeded():
|
||||
"""Insert a user + a near and a far note with hand-crafted embeddings.
|
||||
|
||||
Returns (user_id, near_note_id, far_note_id). Cleaned up after the test.
|
||||
"""
|
||||
async with async_session() as s:
|
||||
user = User(username="pgvec_itest")
|
||||
s.add(user)
|
||||
await s.flush()
|
||||
near = Note(user_id=user.id, title="near", body="near body")
|
||||
far = Note(user_id=user.id, title="far", body="far body")
|
||||
s.add_all([near, far])
|
||||
await s.flush()
|
||||
# query vector will be [1,0,0,...]; near ~ identical (sim≈1.0),
|
||||
# far is orthogonal (sim≈0.0 -> filtered by the default threshold).
|
||||
s.add(NoteEmbedding(note_id=near.id, user_id=user.id, embedding=_vec(1.0)))
|
||||
s.add(NoteEmbedding(note_id=far.id, user_id=user.id, embedding=_vec(0.0, 1.0)))
|
||||
await s.commit()
|
||||
ids = (user.id, near.id, far.id)
|
||||
yield ids
|
||||
user_id = ids[0]
|
||||
async with async_session() as s:
|
||||
await s.execute(delete(NoteEmbedding).where(NoteEmbedding.user_id == user_id))
|
||||
await s.execute(delete(Note).where(Note.user_id == user_id))
|
||||
await s.execute(delete(User).where(User.id == user_id))
|
||||
await s.commit()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_semantic_search_ranks_and_thresholds_via_pgvector(seeded):
|
||||
user_id, near_id, far_id = seeded
|
||||
with patch(
|
||||
"scribe.services.embeddings.get_embedding",
|
||||
AsyncMock(return_value=_vec(1.0)),
|
||||
):
|
||||
results = await semantic_search_notes(user_id=user_id, query="anything", limit=10)
|
||||
|
||||
ids = [note.id for _score, note in results]
|
||||
# Near note returned and ranked first; far (orthogonal, sim≈0) excluded by
|
||||
# the default 0.45 similarity threshold.
|
||||
assert near_id in ids
|
||||
assert far_id not in ids
|
||||
assert ids[0] == near_id
|
||||
top_score = results[0][0]
|
||||
assert top_score == pytest.approx(1.0, abs=1e-3)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_low_threshold_lets_orthogonal_through(seeded):
|
||||
user_id, near_id, far_id = seeded
|
||||
with patch(
|
||||
"scribe.services.embeddings.get_embedding",
|
||||
AsyncMock(return_value=_vec(1.0)),
|
||||
):
|
||||
results = await semantic_search_notes(
|
||||
user_id=user_id, query="anything", limit=10, threshold=-1.0,
|
||||
)
|
||||
ids = [note.id for _score, note in results]
|
||||
# With the floor dropped, both come back and near still ranks above far.
|
||||
assert ids.index(near_id) < ids.index(far_id)
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Tests for services.retrieval_telemetry.
|
||||
|
||||
_build_payload is pure (no DB, no loop) and gets unit coverage. The persistence
|
||||
path (_insert_retrieval_log + the RetrievalLog model / JSONB roundtrip) is an
|
||||
integration test against real Postgres.
|
||||
"""
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from scribe.services.retrieval_telemetry import (
|
||||
_build_payload,
|
||||
record_retrieval,
|
||||
)
|
||||
|
||||
|
||||
def _note(nid):
|
||||
"""Minimal stand-in — _build_payload only reads .id."""
|
||||
return SimpleNamespace(id=nid)
|
||||
|
||||
|
||||
# ─── _build_payload (pure) ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_build_payload_ranks_and_score_bounds():
|
||||
results = [(0.91, _note(11)), (0.72, _note(22)), (0.55, _note(33))]
|
||||
p = _build_payload(
|
||||
user_id=7, source="mcp_search", query="hello", threshold=0.45,
|
||||
limit=10, project_id=3, is_task=None, results=results, duration_ms=12.345,
|
||||
)
|
||||
assert p["result_count"] == 3
|
||||
assert p["top_score"] == 0.91
|
||||
assert p["min_score"] == 0.55
|
||||
assert [it["rank"] for it in p["result_ids"]] == [0, 1, 2]
|
||||
assert [it["id"] for it in p["result_ids"]] == [11, 22, 33]
|
||||
assert p["duration_ms"] == 12.35 # rounded to 2dp
|
||||
assert p["user_id"] == 7 and p["project_id"] == 3 and p["threshold"] == 0.45
|
||||
|
||||
|
||||
def test_build_payload_empty_results():
|
||||
p = _build_payload(
|
||||
user_id=1, source="rest_search", query="x", threshold=0.3,
|
||||
limit=5, project_id=None, is_task=False, results=[], duration_ms=None,
|
||||
)
|
||||
assert p["result_count"] == 0
|
||||
assert p["top_score"] is None and p["min_score"] is None
|
||||
assert p["result_ids"] == []
|
||||
assert p["duration_ms"] is None
|
||||
|
||||
|
||||
def test_build_payload_rounds_scores_to_5dp():
|
||||
p = _build_payload(
|
||||
user_id=1, source="mcp_search", query="q", threshold=0.45,
|
||||
limit=1, project_id=None, is_task=None,
|
||||
results=[(0.123456789, _note(1))], duration_ms=0.0,
|
||||
)
|
||||
assert p["result_ids"][0]["score"] == 0.12346
|
||||
|
||||
|
||||
def test_record_retrieval_without_event_loop_is_safe():
|
||||
"""Called from a sync context (no running loop) it must swallow and return,
|
||||
never raise — telemetry can't be allowed to break a caller."""
|
||||
# No event loop running in this plain sync test.
|
||||
assert record_retrieval(
|
||||
user_id=1, source="mcp_search", query="q", threshold=0.45,
|
||||
limit=10, project_id=None, is_task=None,
|
||||
results=[(0.9, _note(1))],
|
||||
) is None
|
||||
|
||||
|
||||
# ─── persistence (integration) ───────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def _dispose_engine():
|
||||
from scribe.models import engine
|
||||
yield
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_insert_retrieval_log_roundtrip(_dispose_engine):
|
||||
from sqlalchemy import delete, select
|
||||
|
||||
from scribe.models import async_session
|
||||
from scribe.models.retrieval_log import RetrievalLog
|
||||
from scribe.services.retrieval_telemetry import _insert_retrieval_log
|
||||
|
||||
payload = _build_payload(
|
||||
user_id=990001, source="mcp_search", query="pgvector tuning",
|
||||
threshold=0.45, limit=10, project_id=None, is_task=None,
|
||||
results=[(0.88, _note(501)), (0.61, _note(502))], duration_ms=9.9,
|
||||
)
|
||||
await _insert_retrieval_log(payload)
|
||||
|
||||
async with async_session() as s:
|
||||
row = (
|
||||
await s.execute(
|
||||
select(RetrievalLog).where(RetrievalLog.user_id == 990001)
|
||||
)
|
||||
).scalars().first()
|
||||
assert row is not None
|
||||
assert row.source == "mcp_search"
|
||||
assert row.result_count == 2
|
||||
assert row.top_score == 0.88
|
||||
# JSONB roundtrips as a list of dicts with the expected shape.
|
||||
assert row.result_ids[0] == {"id": 501, "score": 0.88, "rank": 0}
|
||||
assert row.created_at is not None # server_default now()
|
||||
await s.execute(delete(RetrievalLog).where(RetrievalLog.user_id == 990001))
|
||||
await s.commit()
|
||||
Reference in New Issue
Block a user