"""AnalysisService — content analysis scoped to confirmed-SFW assets. The privacy invariant is enforced here and nowhere else: the vision provider is called **only** for canonical assets whose latest safety decision is ``sfw``. Confirmed ``nsfw`` and undecided assets never reach the provider — ``run`` records them as ``skipped_nsfw``/leaves them pending without constructing a request. The provider is an injected adapter so tests assert this with a call-recording fake and no network/key. Extracted from photo_analyzer.analyze_image (the OpenAI-compatible Gemini call, result fields, and album-hint prompt) and its ``photos`` schema, re-keyed to ``asset_id`` (donor_ledger.yaml: pa-analyze, pa-schema). Retry/rate-limit/throttle bookkeeping from the donor is out of scope for this story. ponytail: port the donor's retry+RPD throttling when analysis runs at real volume. """ from __future__ import annotations import json import os import uuid from datetime import datetime, timezone from typing import Protocol from sqlalchemy import func, select from sqlalchemy.orm import sessionmaker from photo_pipeline.integrations import exiftool from photo_pipeline.models import AnalysisResult, Asset, SafetyReview from photo_pipeline.services import hashing from photo_pipeline.services.safety import SFW MODEL = "gemini-2.5-flash" PROMPT_VERSION = "1" RESULT_FIELDS = ( "description", "tags", "people_count", "setting", "time_of_day", "season", "mood", "location_hint", "approx_year", ) class VisionProvider(Protocol): def analyze(self, path: str, *, album_hint: str) -> dict: """Return the analysis fields for one image. Raises on unrecoverable error.""" class AnalysisError(Exception): pass def _now() -> datetime: return datetime.now(timezone.utc) class AnalysisService: def __init__(self, session_factory: sessionmaker, *, provider: VisionProvider | None = None) -> None: self._session_factory = session_factory self._provider = provider def _sfw_asset_ids(self, session) -> set[str]: """Asset ids whose latest safety decision is ``sfw`` — the ONLY assets that may reach the provider.""" latest: dict[str, str | None] = {} for review in session.scalars(select(SafetyReview).order_by(SafetyReview.created_at)): latest[review.asset_id] = review.decision return {aid for aid, decision in latest.items() if decision == SFW} def eligible_asset_ids(self) -> list[str]: """Confirmed-SFW canonical active assets without a completed analysis.""" with self._session_factory() as session: sfw = self._sfw_asset_ids(session) if not sfw: return [] assets = session.scalars( select(Asset).where( Asset.id.in_(sfw), Asset.canonical_asset_id.is_(None), Asset.availability_state == "active", Asset.current_path.is_not(None), ) ) done = set( session.scalars( select(AnalysisResult.asset_id).where(AnalysisResult.status == "analyzed") ) ) return [a.id for a in assets if a.id not in done] def counts(self) -> dict[str, int]: with self._session_factory() as session: sfw = self._sfw_asset_ids(session) rows = dict( session.execute( select(AnalysisResult.status, func.count()).group_by(AnalysisResult.status) ).all() ) analyzed = int(rows.get("analyzed", 0)) errored = int(rows.get("error", 0)) return { "eligible": len(sfw), "analyzed": analyzed, "error": errored, "pending": max(len(sfw) - analyzed - errored, 0), } def run(self, asset_ids: list[str] | None = None) -> dict: """Analyze the given assets (default: all eligible). Enforces the gate. Returns ``{analyzed, skipped, errors}``. ``skipped`` counts assets that were requested but are not confirmed SFW — the provider is never called for them. """ with self._session_factory() as session: sfw = self._sfw_asset_ids(session) paths = { a.id: a.current_path for a in session.scalars(select(Asset).where(Asset.canonical_asset_id.is_(None))) if a.current_path } requested = asset_ids if asset_ids is not None else self.eligible_asset_ids() provider = self._provider or _default_provider() analyzed = skipped = errors = 0 for asset_id in requested: if asset_id not in sfw: # Gate: not confirmed SFW → never construct a provider request. self._store(asset_id, status="skipped_nsfw", result=None, error=None, tokens=0, raw="") skipped += 1 continue path = paths.get(asset_id) if not path: skipped += 1 continue try: result = provider.analyze(path, album_hint=_album_hint(path)) except Exception as error: # provider/validation failure is per-asset self._store(asset_id, status="error", result=None, error=str(error), tokens=0, raw="") errors += 1 continue self._store( asset_id, status="analyzed", result=result, error=None, tokens=int(result.get("_tokens", 0)) if isinstance(result, dict) else 0, raw=json.dumps(result, ensure_ascii=False), ) _write_analysis_exif(path, result) analyzed += 1 return {"analyzed": analyzed, "skipped": skipped, "errors": errors} def _store(self, asset_id, *, status, result, error, tokens, raw) -> None: now = _now() with self._session_factory() as session: row = session.get(AnalysisResult, asset_id) or AnalysisResult(asset_id=asset_id) row.status = status row.error_message = error row.model = MODEL row.prompt_version = PROMPT_VERSION row.tokens_total = tokens row.raw_response = raw or None if status == "analyzed" and isinstance(result, dict): row.description = result.get("description") row.tags = json.dumps(result.get("tags", []), ensure_ascii=False) row.people_count = result.get("people_count") row.setting = result.get("setting") row.time_of_day = result.get("time_of_day") row.season = result.get("season") row.mood = result.get("mood") row.location_hint = result.get("location_hint") row.approx_year = result.get("approx_year") row.analyzed_at = now row.exif_written_at = now session.add(row) session.commit() def get(self, asset_id: str) -> dict | None: with self._session_factory() as session: row = session.get(AnalysisResult, asset_id) return _result_dict(row) if row else None def _album_hint(path: str) -> str: from pathlib import Path return Path(path).parent.name def _write_analysis_exif(path: str, result: dict) -> None: """Additive analysis keywords into EXIF (Keywords/Subject), preserving safety and user keywords. The donor also wrote a managed caption; only keywords are written here via the shared adapter. ponytail: add the managed ``AI:`` caption segment + read-back preservation check when the analysis EXIF checkpoint is hardened.""" tags = result.get("tags") if isinstance(result, dict) else None if tags: exiftool.apply_keywords(path, add=[str(t) for t in tags]) def _result_dict(row: AnalysisResult) -> dict: data = {field: getattr(row, field) for field in RESULT_FIELDS} data["tags"] = json.loads(row.tags) if row.tags else [] data.update( asset_id=row.asset_id, status=row.status, model=row.model, prompt_version=row.prompt_version, tokens_total=row.tokens_total, error_message=row.error_message, ) return data def _default_provider() -> VisionProvider: # Test seam (concept §18: deterministic fakes replace the vision edge, enabled # only by test configuration). When this env var names a writable log file, the # worker/API use a fake that records every analyzed path — so the SFW-only gate # can be asserted end-to-end through the real integration layer — instead of the # real OpenAI-compatible call. Never set in production. log_path = os.environ.get("PHOTO_PIPELINE_FAKE_VISION_LOG") if log_path: return _RecordingFakeVision(log_path) return OpenAIVisionProvider() class _RecordingFakeVision: """Deterministic vision fake for end-to-end tests. Appends every analyzed file path to its log so a test can prove NSFW assets never reach the provider, and raises for a path whose stem contains ``boom`` to exercise per-asset error handling. Constructed only when ``PHOTO_PIPELINE_FAKE_VISION_LOG`` is set.""" def __init__(self, log_path: str) -> None: self._log_path = log_path def analyze(self, path: str, *, album_hint: str) -> dict: with open(self._log_path, "a", encoding="utf-8") as handle: handle.write(path + "\n") if "boom" in os.path.splitext(os.path.basename(path))[0]: raise AnalysisError("fake vision failure") return { "description": f"a deterministic scene in {album_hint}", "tags": ["fixture", "deterministic"], "people_count": 1, "setting": "outdoor", "time_of_day": "day", "season": "summer", "mood": "calm", "location_hint": None, "approx_year": None, "_tokens": 7, } class OpenAIVisionProvider: """The real provider: an OpenAI-compatible vision call (Gemini by default). Extracted from photo_analyzer.analyze_image. Constructed lazily from env (``OPENAI_API_KEY`` / ``OPENAI_BASE_URL``); never used in tests, which inject a fake. Kept intentionally thin — no retry/throttle bookkeeping (see module note). """ def __init__(self, *, model: str = MODEL, client=None) -> None: self._model = model self._client = client def _ensure_client(self): if self._client is None: from openai import OpenAI self._client = OpenAI() return self._client def analyze(self, path: str, *, album_hint: str) -> dict: b64, mime = _prepare_image(path) prompt = ANALYSIS_PROMPT + ( f'\n\nAlbum hint: this photo is filed in a folder named "{album_hint}". ' f'Folder names often contain the place and/or year — use it to inform ' f'"location_hint" and "approx_year", but trust the image if they conflict.' ) response = self._ensure_client().chat.completions.create( model=self._model, max_tokens=4096, messages=[ { "role": "user", "content": [ {"type": "image_url", "image_url": {"url": f"data:{mime};base64,{b64}"}}, {"type": "text", "text": prompt}, ], } ], ) raw = (response.choices[0].message.content or "").strip() if raw.startswith("```"): raw = raw.split("```")[1] raw = raw[4:] if raw.startswith("json") else raw raw = raw.strip() result = json.loads(raw) usage = getattr(response, "usage", None) result["_tokens"] = usage.total_tokens if usage else 0 return result ANALYSIS_PROMPT = ( "Analyze this photograph and return ONLY a JSON object with keys: description " "(one clear sentence), tags (8-12 specific keywords), people_count (integer), " "setting, time_of_day, season, mood, location_hint (or null), approx_year " "(integer or null)." ) def _prepare_image(path: str) -> tuple[str, str]: import base64 from io import BytesIO from PIL import Image with Image.open(path) as image: image = image.convert("RGB") image.thumbnail((1024, 1024)) buffer = BytesIO() image.save(buffer, format="JPEG", quality=85) return base64.b64encode(buffer.getvalue()).decode(), "image/jpeg"