- Create ocr_text table for storing per-region OCR results - Add tsvector search_vector column to photos with GIN index and auto-update trigger on filename/user_title/user_notes - Implement ocr_photo Celery task using rapidocr-onnxruntime - Add FTS leg to hybrid search: queries photos.search_vector and ocr_text via UNION, fused with semantic results via RRF (k=60) Migration 0004 backfills search_vector for existing rows. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
137 lines
5.1 KiB
Python
137 lines
5.1 KiB
Python
"""
|
|
Unified search service — hybrid FTS + semantic (RRF) search.
|
|
|
|
Phase 1 (PR4): semantic-only via pgvector cosine similarity.
|
|
Phase 2 (PR5): adds FTS via tsvector, enables RRF fusion.
|
|
"""
|
|
import logging
|
|
from typing import Optional
|
|
|
|
import numpy as np
|
|
from sqlalchemy import select, text, func
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.models import Photo
|
|
from app.models.embeddings import Embedding
|
|
from app.config import settings
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
async def hybrid_search(
|
|
db: AsyncSession,
|
|
q: Optional[str] = None,
|
|
tag_ids: Optional[list[str]] = None,
|
|
date_from: Optional[str] = None,
|
|
date_to: Optional[str] = None,
|
|
limit: int = 50,
|
|
offset: int = 0,
|
|
) -> list[dict]:
|
|
"""Run hybrid search (FTS + semantic) with RRF fusion.
|
|
|
|
Currently semantic-only; FTS leg added in PR5.
|
|
"""
|
|
model_name = settings.vision.embedder.name
|
|
results = {}
|
|
|
|
# ── Semantic search (CLIP text → pgvector cosine) ─────────────────
|
|
if q:
|
|
try:
|
|
from app.services.vision.registry import registry
|
|
embedder = registry.get_embedder()
|
|
query_vec = embedder.embed_text(q)
|
|
|
|
# pgvector cosine distance: <=> returns distance (lower = closer)
|
|
vec_str = "[" + ",".join(str(float(v)) for v in query_vec) + "]"
|
|
stmt = text("""
|
|
SELECT e.photo_id,
|
|
(e.vector <=> :qvec::vector) AS distance
|
|
FROM embeddings e
|
|
WHERE e.model = :model
|
|
ORDER BY e.vector <=> :qvec::vector
|
|
LIMIT 200
|
|
""")
|
|
rows = (await db.execute(stmt, {"qvec": vec_str, "model": model_name})).fetchall()
|
|
|
|
for rank, (photo_id, distance) in enumerate(rows):
|
|
if photo_id not in results:
|
|
results[photo_id] = {"semantic_rank": rank, "fts_rank": None}
|
|
else:
|
|
results[photo_id]["semantic_rank"] = rank
|
|
|
|
except Exception as e:
|
|
logger.warning("Semantic search failed (models may not be loaded): %s", e)
|
|
|
|
# ── FTS search (photos.search_vector + ocr_text) ────────────────
|
|
if q:
|
|
try:
|
|
fts_stmt = text("""
|
|
SELECT id, ts_rank(search_vector, plainto_tsquery('english', :q)) AS rank
|
|
FROM photos
|
|
WHERE search_vector @@ plainto_tsquery('english', :q)
|
|
UNION
|
|
SELECT o.photo_id AS id,
|
|
MAX(o.confidence) AS rank
|
|
FROM ocr_text o
|
|
WHERE to_tsvector('english', o.text) @@ plainto_tsquery('english', :q)
|
|
GROUP BY o.photo_id
|
|
ORDER BY rank DESC
|
|
LIMIT 200
|
|
""")
|
|
fts_rows = (await db.execute(fts_stmt, {"q": q})).fetchall()
|
|
for rank, (photo_id, score) in enumerate(fts_rows):
|
|
if photo_id not in results:
|
|
results[photo_id] = {"semantic_rank": None, "fts_rank": rank}
|
|
else:
|
|
results[photo_id]["fts_rank"] = rank
|
|
except Exception as e:
|
|
logger.warning("FTS search failed: %s", e)
|
|
|
|
# ── RRF fusion ────────────────────────────────────────────────────
|
|
k = 60
|
|
scored = []
|
|
for photo_id, ranks in results.items():
|
|
score = 0.0
|
|
if ranks["semantic_rank"] is not None:
|
|
score += 1.0 / (k + ranks["semantic_rank"])
|
|
if ranks.get("fts_rank") is not None:
|
|
score += 1.0 / (k + ranks["fts_rank"])
|
|
scored.append((photo_id, score))
|
|
|
|
scored.sort(key=lambda x: -x[1])
|
|
|
|
# If no text query, fall back to recent photos
|
|
if not q:
|
|
stmt = select(Photo.id).order_by(Photo.created_at.desc())
|
|
if tag_ids:
|
|
from app.models.tags import photo_tags
|
|
stmt = stmt.join(photo_tags, Photo.id == photo_tags.c.photo_id).where(
|
|
photo_tags.c.tag_id.in_(tag_ids)
|
|
).distinct()
|
|
if date_from:
|
|
stmt = stmt.where(Photo.taken_at >= date_from)
|
|
if date_to:
|
|
stmt = stmt.where(Photo.taken_at <= date_to)
|
|
stmt = stmt.offset(offset).limit(limit)
|
|
rows = (await db.execute(stmt)).fetchall()
|
|
return [{"photo_id": row[0], "score": 0.0} for row in rows]
|
|
|
|
# Apply filters to scored results
|
|
photo_ids = [pid for pid, _ in scored]
|
|
if not photo_ids:
|
|
return []
|
|
|
|
# Filter by tags if requested
|
|
if tag_ids:
|
|
from app.models.tags import photo_tags
|
|
stmt = select(photo_tags.c.photo_id).where(
|
|
photo_tags.c.photo_id.in_(photo_ids),
|
|
photo_tags.c.tag_id.in_(tag_ids),
|
|
).distinct()
|
|
valid_ids = {row[0] for row in (await db.execute(stmt)).fetchall()}
|
|
scored = [(pid, s) for pid, s in scored if pid in valid_ids]
|
|
|
|
# Paginate
|
|
page = scored[offset : offset + limit]
|
|
return [{"photo_id": pid, "score": score} for pid, score in page]
|