Files
mule-image/backend/app/services/search.py
dtoro 842a4fc864 feat: add OCR text extraction and Postgres full-text search
- 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>
2026-04-10 09:10:14 +02:00

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]