Files
mule-image/backend/app/tasks/vision.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

185 lines
6.1 KiB
Python

"""
Celery tasks for the vision pipeline — embedding, OCR, object detection,
face recognition.
All tasks run on the dedicated `vision` queue with limited concurrency
(memory-bound CPU inference). They read thumbnails generated by
generate_thumbnails, so they MUST run after thumbs complete.
"""
import asyncio
import logging
from pathlib import Path
import numpy as np
from celery import shared_task
from sqlalchemy import select, delete, text
from PIL import Image
from app.database import AsyncSessionLocal
from app.models import Photo
from app.models.embeddings import Embedding
from app.config import settings
logger = logging.getLogger(__name__)
def _load_thumb(photo_id: str, size: str = "medium") -> np.ndarray | None:
"""Load a thumbnail as an RGB numpy array."""
thumb_path = Path(f"/data/thumbs/{photo_id}/{size}.webp")
if not thumb_path.exists():
logger.warning("Thumbnail not found: %s", thumb_path)
return None
img = Image.open(thumb_path).convert("RGB")
return np.array(img)
@shared_task(name='embed_photo', queue='vision')
def embed_photo(photo_id: str):
"""Generate CLIP embedding for a photo and store in pgvector."""
if not settings.vision.enabled:
return {'status': 'skipped', 'reason': 'vision disabled'}
return asyncio.run(_embed_photo_async(photo_id))
async def _embed_photo_async(photo_id: str):
image = _load_thumb(photo_id, "medium") # 640px
if image is None:
return {'status': 'error', 'message': 'thumbnail not found'}
from app.services.vision.registry import registry
embedder = registry.get_embedder()
vector = embedder.embed_image(image)
model_name = settings.vision.embedder.name
async with AsyncSessionLocal() as session:
# Upsert: delete existing then insert
await session.execute(
delete(Embedding).where(
Embedding.photo_id == photo_id,
Embedding.model == model_name,
)
)
emb = Embedding(
photo_id=photo_id,
model=model_name,
vector=vector.tolist(),
)
session.add(emb)
await session.commit()
logger.info("Embedded photo %s with %s", photo_id, model_name)
return {'status': 'success', 'photo_id': photo_id}
@shared_task(name='vision_fanout', queue='vision')
def vision_fanout(photo_id: str):
"""Dispatch all enabled vision tasks for a photo."""
if not settings.vision.enabled:
return {'status': 'skipped', 'reason': 'vision disabled'}
embed_photo.delay(photo_id)
if settings.vision.ocr.enabled:
ocr_photo.delay(photo_id)
if settings.vision.detector.enabled:
detect_objects.delay(photo_id)
if settings.vision.faces.enabled:
extract_faces.delay(photo_id)
return {'status': 'dispatched', 'photo_id': photo_id}
@shared_task(name='ocr_photo', queue='vision')
def ocr_photo(photo_id: str):
"""Run OCR on a photo and store text regions."""
if not settings.vision.enabled or not settings.vision.ocr.enabled:
return {'status': 'skipped', 'reason': 'OCR disabled'}
return asyncio.run(_ocr_photo_async(photo_id))
async def _ocr_photo_async(photo_id: str):
image = _load_thumb(photo_id, "large") # 1280px for better OCR accuracy
if image is None:
return {'status': 'error', 'message': 'thumbnail not found'}
from app.services.vision.registry import registry
ocr_engine = registry.get_ocr()
results = ocr_engine.run(image)
if not results:
logger.info("No OCR text found for photo %s", photo_id)
return {'status': 'success', 'photo_id': photo_id, 'regions': 0}
from app.models.ocr_text import OCRText
async with AsyncSessionLocal() as session:
# Delete existing OCR results for this photo (re-run safe)
await session.execute(
delete(OCRText).where(OCRText.photo_id == photo_id)
)
for r in results:
session.add(OCRText(
photo_id=photo_id,
text=r.text,
language=r.language,
confidence=r.confidence,
bbox=r.bbox,
))
await session.commit()
logger.info("OCR: %d text regions for photo %s", len(results), photo_id)
return {'status': 'success', 'photo_id': photo_id, 'regions': len(results)}
@shared_task(name='detect_objects', queue='vision')
def detect_objects(photo_id: str):
"""Detect objects in a photo — implemented in PR6."""
return {'status': 'not_implemented'}
@shared_task(name='extract_faces', queue='vision')
def extract_faces(photo_id: str):
"""Detect faces and extract embeddings — implemented in PR7."""
return {'status': 'not_implemented'}
@shared_task(name='backfill_vision')
def backfill_vision(task: str | None = None, limit: int | None = None):
"""Queue vision tasks for photos that haven't been processed yet."""
return asyncio.run(_backfill_vision_async(task, limit))
async def _backfill_vision_async(task: str | None, limit: int | None):
model_name = settings.vision.embedder.name
async with AsyncSessionLocal() as session:
# Find photos without embeddings
stmt = text("""
SELECT p.id FROM photos p
LEFT JOIN embeddings e ON e.photo_id = p.id AND e.model = :model
WHERE e.photo_id IS NULL
AND p.processing_status = 'completed'
ORDER BY p.created_at DESC
""")
if limit:
stmt = text(str(stmt) + f" LIMIT {limit}")
result = await session.execute(stmt, {"model": model_name})
photo_ids = [row[0] for row in result.fetchall()]
count = 0
for pid in photo_ids:
if task == 'embed' or task is None:
embed_photo.delay(pid)
if task == 'ocr' or task is None:
ocr_photo.delay(pid)
if task == 'detect' or task is None:
detect_objects.delay(pid)
if task == 'faces' or task is None:
extract_faces.delay(pid)
count += 1
logger.info("Backfill queued %d photos for vision processing", count)
return {'status': 'queued', 'count': count}