feat: runtime feature flags, upload/download, RAW decoding

Adds Redis-backed feature flags for vision stages with admin UI toggles
and manual backfill trigger, photo upload and download routers with
frontend upload modal, and rawpy-based RAW decoding with JPEG fallback
for misnamed DNGs. Fixes pgvector serialization, is_trashed filter, and
naive-datetime bind in incremental duplicate regrouping; bumps Celery
time limits on regroup tasks beyond the 5-minute default.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
root
2026-04-14 21:31:52 +02:00
parent 800ee447ad
commit 5c531f11da
16 changed files with 2232 additions and 35 deletions

View File

@@ -11,7 +11,7 @@ import os
from app.config import settings
from app.database import init_db
from app.routers import photos, folders, heaps, tags, discard, library, search, auth, admin, sharing
from app.routers import photos, folders, heaps, tags, discard, library, search, auth, admin, sharing, upload, download, features
from app.services.scanner import start_initial_scan, bootstrap_default_source_root
from app.services.cleanup import cleanup_data_integrity
@@ -95,6 +95,9 @@ app.include_router(tags.router, prefix="/api/v1/tags", tags=["tags"])
app.include_router(discard.router, prefix="/api/v1/discard", tags=["discard"])
app.include_router(library.router, prefix="/api/v1/library", tags=["library"])
app.include_router(search.router, prefix="/api/v1/photos/search", tags=["search"])
app.include_router(upload.router, prefix="/api/v1/upload", tags=["upload"])
app.include_router(download.router, prefix="/api/v1/download", tags=["download"])
app.include_router(features.router, prefix="/api/v1/features", tags=["features"])
@app.get("/")
async def root():

View File

@@ -18,6 +18,14 @@ from app.models.user import User
from app.models.photos import Photo
from app.models.folders import SourceRoot
from app.config import settings
from app.services.feature_flags import (
ALL_FLAGS,
snapshot as flags_snapshot,
set_flag,
reset_flag,
is_enabled,
FLAG_VISION_ENABLED,
)
logger = logging.getLogger(__name__)
@@ -258,3 +266,141 @@ async def delete_user(
logger.info(f"Admin '{admin.username}' deactivated user '{user.username}'")
return {"status": "ok", "detail": f"User '{user.username}' deactivated"}
# ---------------------------------------------------------------------------
# AI / vision feature flags + manual triggers
# ---------------------------------------------------------------------------
class FeatureFlagUpdate(BaseModel):
"""PATCH body for toggling a feature flag.
``value`` sets an explicit override (true/false); omitting it clears
the override and reverts the flag to its YAML default.
"""
value: Optional[bool] = None
@router.get("/feature-flags")
async def get_feature_flags(admin: User = Depends(require_admin)):
"""Return every tunable feature flag with its current effective
value, YAML default, and whether an admin override is in effect."""
return {"flags": flags_snapshot()}
@router.patch("/feature-flags/{flag_name}")
async def update_feature_flag(
flag_name: str,
body: FeatureFlagUpdate,
admin: User = Depends(require_admin),
):
"""Set or clear an override for one flag. With ``value`` set, the
flag is pinned to that boolean; without it, the override is deleted
and the YAML default takes over again.
New value is observed by vision tasks on their next invocation —
there's no worker restart required.
"""
if flag_name not in ALL_FLAGS:
raise HTTPException(status_code=404, detail=f"Unknown flag: {flag_name}")
try:
if body.value is None:
reset_flag(flag_name)
action = "cleared override"
else:
set_flag(flag_name, bool(body.value))
action = f"set to {body.value}"
except RuntimeError as e:
# Redis unreachable — surface as 503 so the UI doesn't think it
# succeeded silently.
raise HTTPException(status_code=503, detail=str(e))
logger.info(f"Admin '{admin.username}' {action} for flag '{flag_name}'")
return {"flags": flags_snapshot()}
class BackfillVisionBody(BaseModel):
"""POST body for triggering a vision backfill. ``task`` picks a
specific stage (``embed`` / ``ocr`` / ``detect`` / ``faces`` /
``classify``); leaving it null runs every enabled stage. ``limit``
caps how many photos per stage are queued — useful for smoke-
testing a newly-enabled feature before committing a full run.
"""
task: Optional[str] = None
limit: Optional[int] = None
@router.post("/ai/backfill")
async def trigger_ai_backfill(
body: BackfillVisionBody,
admin: User = Depends(require_admin),
):
"""Queue a vision backfill pass. Identical code path as the automatic
post-scan backfill — just triggered manually from the UI."""
if not is_enabled(FLAG_VISION_ENABLED):
raise HTTPException(
status_code=400,
detail="Vision is currently disabled; enable it before running a backfill.",
)
valid_tasks = {'embed', 'ocr', 'detect', 'faces', 'classify'}
if body.task is not None and body.task not in valid_tasks:
raise HTTPException(
status_code=400,
detail=f"task must be one of {sorted(valid_tasks)} or null",
)
if body.limit is not None and body.limit <= 0:
raise HTTPException(status_code=400, detail="limit must be positive")
# Import lazily so importing admin.py doesn't pull in the whole
# vision stack on startup (Celery task module loads numpy etc.).
from app.tasks.vision import backfill_vision
result = backfill_vision.apply_async(
kwargs={'task': body.task, 'limit': body.limit}
)
logger.info(
f"Admin '{admin.username}' queued vision backfill "
f"(task={body.task}, limit={body.limit}, celery_id={result.id})"
)
return {
"status": "queued",
"task_id": result.id,
"task": body.task,
"limit": body.limit,
}
@router.post("/ai/recluster-faces")
async def trigger_face_recluster(admin: User = Depends(require_admin)):
"""Kick off face recluster. Normally auto-fires after a scan via a
debounced scheduler; this endpoint is for admins who want to force
a fresh clustering pass (e.g. after tweaking ``cluster_eps`` in
the YAML config)."""
if not is_enabled(FLAG_VISION_ENABLED):
raise HTTPException(
status_code=400,
detail="Vision is currently disabled; enable it before reclustering.",
)
from app.tasks.vision import recluster_faces
result = recluster_faces.apply_async()
logger.info(
f"Admin '{admin.username}' queued face recluster (celery_id={result.id})"
)
return {"status": "queued", "task_id": result.id}
@router.post("/ai/rescan")
async def trigger_full_rescan(admin: User = Depends(require_admin)):
"""Dispatch the same scan_all_source_roots job the backend runs at
startup. Picks up any new files on disk and, through the
post-scan hook, queues a vision backfill for whatever still lacks
embeddings / OCR / etc.
"""
from app.tasks.scan import scan_all_source_roots
result = scan_all_source_roots.apply_async()
logger.info(
f"Admin '{admin.username}' queued full rescan (celery_id={result.id})"
)
return {"status": "queued", "task_id": result.id}

View File

@@ -0,0 +1,247 @@
"""
Download router — streams a .zip of every photo in a folder (recursively)
or a heap back to the browser.
Auth: both endpoints accept the regular Authorization header *or* a
``?token=JWT`` query string, mirroring the media endpoints. That lets the
frontend trigger a download with a plain ``<a href>`` (which can't set a
header), keeping the client side a one-liner.
Implementation: we build the zip into a ``NamedTemporaryFile`` and then
stream its bytes back, deleting the temp file on the way out. Stored
(uncompressed) mode because photos and videos are already compressed —
deflating them again just burns CPU for a fraction of a percent. For
very large libraries the temp-file route is mildly wasteful vs. a true
streaming zip (zipstream-ng etc), but it avoids a new dependency and
handles arbitrary folder sizes without blowing out RAM.
"""
import logging
import os
import re
import tempfile
import zipfile
from typing import List
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import StreamingResponse
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.dependencies import get_current_user_media
from app.models import Folder, Heap, Photo, SourceRoot
from app.models.heaps import heap_photos
from app.models.user import User
logger = logging.getLogger(__name__)
router = APIRouter()
def _safe_filename(name: str) -> str:
"""Strip characters that Content-Disposition or Windows filesystems
would choke on. Keeps the download's filename readable without
needing any escaping on the client side."""
cleaned = re.sub(r'[\\/:*?"<>|\r\n\t]', '_', name).strip().strip('.')
return cleaned or 'download'
async def _collect_folder_photos(
folder_id: str,
user: User,
db: AsyncSession,
) -> tuple[str, str, List[Photo]]:
"""Resolve a folder or source-root id → (base_path, display_name,
photos). ``base_path`` is the prefix we strip off each photo's
filepath when naming zip entries, so the archive mirrors the user's
on-disk structure under that folder.
"""
folder = (await db.execute(
select(Folder).where(Folder.id == folder_id, Folder.user_id == user.id)
)).scalar_one_or_none()
base_path: str
display_name: str
if folder is not None:
base_path = os.path.normpath(folder.path)
display_name = folder.name or os.path.basename(base_path)
else:
sr = (await db.execute(
select(SourceRoot).where(
SourceRoot.id == folder_id,
SourceRoot.user_id == user.id,
)
)).scalar_one_or_none()
if sr is None:
raise HTTPException(status_code=404, detail="Folder not found")
base_path = os.path.normpath(sr.path)
display_name = sr.name or os.path.basename(base_path)
# Every photo whose filepath is at or below the base path — matches
# the same prefix convention folders.py uses for recursive deletes.
descendant_prefix = base_path.rstrip(os.sep) + os.sep
result = await db.execute(
select(Photo).where(
Photo.user_id == user.id,
Photo.is_discarded == False, # noqa: E712
(Photo.filepath == base_path) | (Photo.filepath.like(descendant_prefix + '%')),
)
)
photos = list(result.scalars().all())
return base_path, display_name, photos
def _build_zip(
photos: List[Photo],
arcname_fn,
) -> tempfile.NamedTemporaryFile:
"""Write ``photos`` into a fresh ZIP_STORED temp file.
``arcname_fn(photo, used_names)`` returns the entry name to use for
the given photo; the caller supplies it because folder downloads
want path-preserving names while heap downloads flatten to bare
filenames (with a collision suffix).
"""
tmp = tempfile.NamedTemporaryFile(delete=False, suffix='.zip')
try:
used: set[str] = set()
with zipfile.ZipFile(tmp, 'w', zipfile.ZIP_STORED, allowZip64=True) as zf:
for p in photos:
if not p.filepath or not os.path.exists(p.filepath):
# Silent skip: the scanner may have indexed files
# that have since been moved / unlinked by a shell.
continue
name = arcname_fn(p, used)
used.add(name)
try:
zf.write(p.filepath, name)
except OSError as e:
logger.warning(f"Skipping {p.filepath} in zip: {e}")
tmp.close()
return tmp
except Exception:
tmp.close()
try:
os.unlink(tmp.name)
except OSError:
pass
raise
def _stream_and_cleanup(path: str):
"""Yield the temp zip in 1 MiB chunks and unlink it when the
iterator is exhausted (or GC'd, if the client disconnects early)."""
try:
with open(path, 'rb') as f:
while True:
chunk = f.read(1024 * 1024)
if not chunk:
break
yield chunk
finally:
try:
os.unlink(path)
except OSError as e:
logger.debug(f"Temp zip cleanup failed for {path}: {e}")
def _dedupe(name: str, used: set[str]) -> str:
"""Return ``name`` (or ``name (2)``, ``name (3)`` ...) such that the
result doesn't collide with anything in ``used``. Needed for heap
downloads where two members can have identical filenames from
different folders."""
if name not in used:
return name
stem, ext = os.path.splitext(name)
n = 2
while True:
cand = f"{stem} ({n}){ext}"
if cand not in used:
return cand
n += 1
@router.get("/folders/{folder_id}")
async def download_folder(
folder_id: str,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user_media),
):
"""Zip every (non-discarded) photo under a folder/source-root and
stream it back. Entries preserve the folder structure relative to
the downloaded root so the resulting archive is a faithful snapshot.
"""
base_path, display_name, photos = await _collect_folder_photos(
folder_id, current_user, db
)
if not photos:
raise HTTPException(status_code=404, detail="No photos to download")
def arcname(p: Photo, _used: set[str]) -> str:
# Relative path from the download root, falling back to the
# bare filename if the photo somehow lives outside base_path.
abs_path = os.path.normpath(p.filepath)
if abs_path.startswith(base_path + os.sep):
rel = abs_path[len(base_path) + 1:]
elif abs_path == base_path:
rel = os.path.basename(abs_path)
else:
rel = p.filename or os.path.basename(abs_path)
# Nest everything under display_name so users see one top-level
# folder inside the zip rather than loose files.
return os.path.join(_safe_filename(display_name), rel)
tmp = _build_zip(photos, arcname)
filename = _safe_filename(display_name) + '.zip'
return StreamingResponse(
_stream_and_cleanup(tmp.name),
media_type='application/zip',
headers={
'Content-Disposition': f'attachment; filename="{filename}"',
'Content-Length': str(os.path.getsize(tmp.name)),
},
)
@router.get("/heaps/{heap_id}")
async def download_heap(
heap_id: str,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user_media),
):
"""Zip every photo in a heap. Heaps are flat collections, so entries
use the original filename (with a ``(2)`` collision suffix when
two members share a name)."""
heap = (await db.execute(
select(Heap).where(Heap.id == heap_id, Heap.user_id == current_user.id)
)).scalar_one_or_none()
if heap is None:
raise HTTPException(status_code=404, detail="Heap not found")
result = await db.execute(
select(Photo)
.join(heap_photos, heap_photos.c.photo_id == Photo.id)
.where(
heap_photos.c.heap_id == heap_id,
Photo.is_discarded == False, # noqa: E712
)
)
photos = list(result.scalars().all())
if not photos:
raise HTTPException(status_code=404, detail="Heap is empty")
def arcname(p: Photo, used: set[str]) -> str:
bare = p.filename or os.path.basename(p.filepath or 'photo')
entry = os.path.join(_safe_filename(heap.name), _dedupe(bare, used))
return entry
tmp = _build_zip(photos, arcname)
filename = _safe_filename(heap.name) + '.zip'
return StreamingResponse(
_stream_and_cleanup(tmp.name),
media_type='application/zip',
headers={
'Content-Disposition': f'attachment; filename="{filename}"',
'Content-Length': str(os.path.getsize(tmp.name)),
},
)

View File

@@ -0,0 +1,23 @@
"""
Public feature-flag read API — lets the authenticated frontend know
which AI-powered sections to render.
This is NOT the admin mutation endpoint (that's in ``admin.py`` and
gated by ``require_admin``). Here we only expose the effective boolean
state so the UI can hide things like the People view, Tags view, or
text-search affordances when the underlying pipeline stage is off.
"""
from fastapi import APIRouter, Depends
from app.dependencies import get_current_user
from app.models.user import User
from app.services.feature_flags import ALL_FLAGS, is_enabled
router = APIRouter()
@router.get("")
async def get_enabled_features(_: User = Depends(get_current_user)):
"""Return ``{flag_name: bool}`` for every known flag, reflecting
the currently effective value (admin override or YAML default)."""
return {name: is_enabled(name) for name in ALL_FLAGS}

View File

@@ -0,0 +1,303 @@
"""
Upload router — lets users drop files (or whole folders) from their
desktop into a destination Folder, preserving any sub-folder structure
they bring with them.
Each POST handles one file. The frontend fans out many parallel requests
per drop, giving it per-file progress without the server having to
invent a chunking protocol. For folder uploads, the browser passes
`webkitRelativePath` under the `relative_path` field; any leading
sub-directories there are materialised on disk (and as Folder rows)
under the destination.
Uploaded files are placed under the destination folder on the owner's
media mount, indexed immediately (Photo row created), and queued for
the same thumb + metadata pipeline that the scanner uses. An optional
`heap_id` also drops them into a heap in the same request.
"""
import hashlib
import logging
import os
from pathlib import Path
from datetime import datetime
from typing import Optional
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile
from sqlalchemy import insert, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.dependencies import get_current_user
from app.models import Folder, Heap, Photo, SourceRoot
from app.models.heaps import heap_photos
from app.models.user import User
from app.services.date_guess import has_date_warning
from app.tasks.scan import SUPPORTED_EXTENSIONS, get_media_type
from app.tasks.thumbs import generate_thumbnails
from app.services.metadata import extract_metadata
logger = logging.getLogger(__name__)
router = APIRouter()
MAX_UPLOAD_BYTES = 500 * 1024 * 1024 # 500 MB per file cap.
def _validate_segment(segment: str) -> str:
"""Reject path segments that would escape the destination directory."""
segment = segment.strip()
if not segment or segment in ('.', '..') or '/' in segment or '\\' in segment:
raise HTTPException(status_code=400, detail=f"Invalid path segment: {segment!r}")
return segment
def _sanitize_relative_path(rel: Optional[str]) -> list[str]:
"""Split `relative_path` into safe segments (dirs + filename).
Empty or missing → []. Any absolute path, backslash, or `..` segment
raises 400 — we never want an upload to escape the destination.
"""
if not rel:
return []
# Normalise backslashes to forward slashes; browsers on Windows send
# webkitRelativePath with forward slashes anyway, but defend in depth.
rel = rel.replace('\\', '/').strip('/')
if not rel:
return []
segs = [_validate_segment(s) for s in rel.split('/') if s]
return segs
async def _resolve_destination(
folder_id: str,
user: User,
db: AsyncSession,
) -> Folder:
"""Resolve `folder_id` to a concrete Folder row the user owns.
Accepts both Folder ids and SourceRoot ids (for source roots, we
return the Folder row at the mount path — the scanner creates one
for every source root it walks). Raises 404 if neither matches.
"""
folder = (await db.execute(
select(Folder).where(Folder.id == folder_id, Folder.user_id == user.id)
)).scalar_one_or_none()
if folder is not None:
return folder
sr = (await db.execute(
select(SourceRoot).where(SourceRoot.id == folder_id, SourceRoot.user_id == user.id)
)).scalar_one_or_none()
if sr is None:
raise HTTPException(status_code=404, detail="Destination folder not found")
root_folder = (await db.execute(
select(Folder).where(
Folder.source_root_id == sr.id,
Folder.user_id == user.id,
Folder.path == os.path.normpath(sr.path),
)
)).scalar_one_or_none()
if root_folder is None:
# First-time source root with no walk yet — create the row now so
# uploads work even before the initial scan has run.
root_folder = Folder(
name=sr.name or os.path.basename(sr.path),
path=os.path.normpath(sr.path),
source_root_id=sr.id,
user_id=user.id,
)
os.makedirs(root_folder.path, exist_ok=True)
db.add(root_folder)
await db.flush()
return root_folder
async def _ensure_subfolder(
parent: Folder,
name: str,
user: User,
db: AsyncSession,
) -> Folder:
"""Return (or create) a Folder row named `name` under `parent`.
Also mkdirs the directory on disk. Idempotent — safe to call for a
path segment that already exists as a Folder row or directory.
"""
child_path = os.path.normpath(os.path.join(parent.path, name))
existing = (await db.execute(
select(Folder).where(
Folder.path == child_path,
Folder.user_id == user.id,
)
)).scalar_one_or_none()
if existing is not None:
os.makedirs(child_path, exist_ok=True)
return existing
os.makedirs(child_path, exist_ok=True)
child = Folder(
name=name,
path=child_path,
parent_id=parent.id,
source_root_id=parent.source_root_id,
user_id=user.id,
is_hidden=parent.is_hidden,
)
db.add(child)
await db.flush()
return child
def _unique_path(target_dir: str, filename: str) -> tuple[str, str]:
"""Return a (filepath, filename) that doesn't collide with an
existing file on disk. Suffixes " (2)", " (3)", ... until a free
slot is found. Prevents upload-over-existing and keeps the user's
original file intact.
"""
base, ext = os.path.splitext(filename)
candidate = os.path.join(target_dir, filename)
n = 2
while os.path.exists(candidate):
new_name = f"{base} ({n}){ext}"
candidate = os.path.join(target_dir, new_name)
n += 1
return candidate, os.path.basename(candidate)
@router.post("")
async def upload_file(
file: UploadFile = File(...),
destination_folder_id: str = Form(...),
relative_path: Optional[str] = Form(None),
heap_id: Optional[str] = Form(None),
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""Upload a single file into a destination folder (and optionally a
heap). For folder uploads, `relative_path` carries the sub-folder
chain from the browser's `webkitRelativePath`, and we materialise
it under the destination on disk + as Folder rows.
Returns the created photo's id on success. 4xx on unsupported file
type, bad path, missing destination, or too-large file.
"""
# --- validate inputs -------------------------------------------------
raw_name = file.filename or ''
if not raw_name:
raise HTTPException(status_code=400, detail="Missing filename")
# Prefer the leaf of relative_path when present (it contains the
# original filename as the browser saw it inside the picked folder).
segs = _sanitize_relative_path(relative_path)
if segs:
leaf = segs[-1]
subdirs = segs[:-1]
else:
leaf = _validate_segment(os.path.basename(raw_name))
subdirs = []
ext = Path(leaf).suffix.lower()
if ext not in SUPPORTED_EXTENSIONS:
raise HTTPException(
status_code=400,
detail=f"Unsupported file type: {ext or '(none)'}",
)
dest_folder = await _resolve_destination(destination_folder_id, current_user, db)
target_folder = dest_folder
for seg in subdirs:
target_folder = await _ensure_subfolder(target_folder, seg, current_user, db)
target_dir = target_folder.path
os.makedirs(target_dir, exist_ok=True)
filepath, final_name = _unique_path(target_dir, leaf)
# --- stream to disk, hash as we go ----------------------------------
hasher = hashlib.sha256()
total = 0
try:
with open(filepath, 'wb') as out:
while True:
chunk = await file.read(1024 * 1024)
if not chunk:
break
total += len(chunk)
if total > MAX_UPLOAD_BYTES:
out.close()
os.unlink(filepath)
raise HTTPException(
status_code=413,
detail=f"File exceeds {MAX_UPLOAD_BYTES // (1024*1024)}MB limit",
)
hasher.update(chunk)
out.write(chunk)
except HTTPException:
raise
except Exception as e:
logger.error(f"Upload write failed for {filepath}: {e}")
if os.path.exists(filepath):
try:
os.unlink(filepath)
except OSError:
pass
raise HTTPException(status_code=500, detail=f"Upload failed: {e}")
file_hash = hasher.hexdigest()
# --- validate heap before committing the DB row ---------------------
if heap_id:
heap = (await db.execute(
select(Heap).where(Heap.id == heap_id, Heap.user_id == current_user.id)
)).scalar_one_or_none()
if heap is None:
# Destination heap vanished — still keep the file + photo row,
# but tell the caller so the UI can surface the mismatch.
heap_id = None
# --- create Photo row ------------------------------------------------
mtime_dt = datetime.fromtimestamp(os.stat(filepath).st_mtime)
photo = Photo(
filepath=filepath,
filename=final_name,
folder_id=target_folder.id,
user_id=current_user.id,
file_hash=file_hash,
media_type=get_media_type(filepath),
original_format=Path(filepath).suffix.upper()[1:],
file_size=total,
taken_at=mtime_dt,
taken_at_source='filesystem',
has_date_warning=has_date_warning(filepath, mtime_dt),
is_hidden=bool(target_folder.is_hidden),
processing_status='pending',
)
db.add(photo)
await db.flush()
if heap_id:
await db.execute(
insert(heap_photos),
[{"heap_id": heap_id, "photo_id": photo.id}],
)
await db.commit()
# Queue the same background work the scanner does so thumbnails +
# EXIF show up without the user having to trigger a rescan.
try:
generate_thumbnails.delay(photo.id)
extract_metadata.delay(photo.id)
except Exception as e:
logger.warning(f"Failed to queue post-upload tasks for {photo.id}: {e}")
return {
"photo_id": photo.id,
"filename": final_name,
"folder_id": target_folder.id,
"folder_path": target_folder.path,
"heap_id": heap_id,
}

View File

@@ -221,6 +221,13 @@ async def incremental_regroup(
from datetime import timedelta
since = datetime.now(timezone.utc) - timedelta(hours=1)
# Photo.added_at is stored as TIMESTAMP WITHOUT TIME ZONE, so
# asyncpg rejects aware datetimes with "can't subtract offset-naive
# and offset-aware". Normalise: if `since` has a tzinfo, convert
# it to UTC and drop the tzinfo so the bind parameter is naive.
if since.tzinfo is not None:
since = since.astimezone(timezone.utc).replace(tzinfo=None)
# Get newly added photos (the "new" set).
new_rows = (
await session.execute(
@@ -349,6 +356,14 @@ async def _clip_neighbor_scan(
for photo_id, vector in target_embeddings:
# pgvector cosine distance: <=> operator
# Find top 20 nearest neighbors within threshold.
# Serialize the vector as "[a,b,c,...]" — pgvector's text
# format uses commas; numpy's default str() joins with spaces
# which Postgres rejects with "invalid input syntax for vector".
if hasattr(vector, 'tolist'):
vec_seq = vector.tolist()
else:
vec_seq = list(vector)
vec_text = '[' + ','.join(f'{float(x):.8f}' for x in vec_seq) + ']'
result = await session.execute(
text("""
SELECT e.photo_id, (e.vector <=> :vec) AS distance
@@ -356,14 +371,14 @@ async def _clip_neighbor_scan(
JOIN photos p ON p.id = e.photo_id
WHERE e.model = :model
AND e.photo_id != :pid
AND p.is_discarded = false
AND p.is_trashed = false
AND p.is_hidden = false
AND (e.vector <=> :vec) < :threshold
ORDER BY e.vector <=> :vec
LIMIT 20
"""),
{
'vec': str(vector),
'vec': vec_text,
'pid': photo_id,
'model': embedder_model,
'threshold': threshold,

View File

@@ -0,0 +1,301 @@
"""
Runtime feature flags for expensive pipeline stages.
The YAML config (``mulita.yml``) ships reasonable defaults. Admins can
toggle these at runtime from the Settings → AI Features tab without
rebuilding the image or editing the bind-mounted YAML; the overrides
live in Redis so both the FastAPI backend and the Celery workers see
the same value within ~1s of the write.
The key namespace is:
mulita:flags:<name> → "true" | "false"
An unset key means "fall back to the YAML default" — so an admin who
has never touched the tab sees exactly the config-file behaviour.
Only bool flags live here. Thresholds, batch sizes, model names etc.
stay in the YAML file because flipping them safely requires restarting
the vision workers (model reload, ONNX session re-init); that's not
something a single admin click should do.
"""
from __future__ import annotations
import logging
from typing import Optional
import redis
from app.config import settings
logger = logging.getLogger(__name__)
# Feature identifiers. The public name is what the admin UI sends; the
# ``yaml_default`` getter returns the value the YAML would have set.
# Keep these in sync with the VisionSettings fields in ``config.py``.
FLAG_VISION_ENABLED = 'vision.enabled'
FLAG_OCR_ENABLED = 'vision.ocr.enabled'
FLAG_DETECTOR_ENABLED = 'vision.detector.enabled'
FLAG_FACES_ENABLED = 'vision.faces.enabled'
FLAG_CLASSIFIER_ENABLED = 'vision.classifier.enabled'
ALL_FLAGS = (
FLAG_VISION_ENABLED,
FLAG_OCR_ENABLED,
FLAG_DETECTOR_ENABLED,
FLAG_FACES_ENABLED,
FLAG_CLASSIFIER_ENABLED,
)
_REDIS: Optional[redis.Redis] = None
def _redis() -> Optional[redis.Redis]:
"""Lazy Redis client. Returns None if the broker is unreachable so
callers can fall back to YAML defaults instead of crashing."""
global _REDIS
if _REDIS is None:
try:
_REDIS = redis.Redis.from_url(
settings.celery_broker_url, decode_responses=True
)
_REDIS.ping()
except Exception as e:
logger.warning(f"feature_flags: Redis unavailable, using YAML defaults ({e})")
_REDIS = None
return _REDIS
def _yaml_default(name: str) -> bool:
"""Return the YAML-configured default for a flag. Used when Redis
has no value set (fresh install or admin never touched the tab)."""
v = settings.vision
if name == FLAG_VISION_ENABLED:
return bool(v.enabled)
if name == FLAG_OCR_ENABLED:
return bool(v.ocr.enabled)
if name == FLAG_DETECTOR_ENABLED:
return bool(v.detector.enabled)
if name == FLAG_FACES_ENABLED:
return bool(v.faces.enabled)
if name == FLAG_CLASSIFIER_ENABLED:
return bool(v.classifier.enabled)
raise ValueError(f"Unknown feature flag: {name!r}")
def _redis_key(name: str) -> str:
return f"mulita:flags:{name}"
def is_enabled(name: str) -> bool:
"""Return True if feature ``name`` is currently enabled.
Order of precedence:
1. Redis override (set by PATCH /admin/feature-flags)
2. YAML default
Reads are cheap (~ms) and we intentionally do NOT add a local
process cache — the whole point of runtime flags is that a toggle
takes effect on the next task without a worker restart.
"""
r = _redis()
if r is not None:
try:
raw = r.get(_redis_key(name))
if raw is not None:
return raw.lower() == 'true'
except Exception as e:
logger.warning(f"feature_flags: Redis read failed for {name} ({e})")
return _yaml_default(name)
def set_flag(name: str, value: bool) -> None:
"""Persist a flag override to Redis. No-op if Redis is unreachable
(we don't silently pretend to have written; raise so the admin
request returns a 500 instead of misleading success)."""
if name not in ALL_FLAGS:
raise ValueError(f"Unknown feature flag: {name!r}")
r = _redis()
if r is None:
raise RuntimeError("Redis unavailable; cannot update feature flags")
r.set(_redis_key(name), 'true' if value else 'false')
_apply_worker_side_effects(name)
def reset_flag(name: str) -> None:
"""Delete the Redis override so the flag falls back to its YAML
default. Useful if an admin wants a clean slate without guessing
what the config defaults are."""
if name not in ALL_FLAGS:
raise ValueError(f"Unknown feature flag: {name!r}")
r = _redis()
if r is None:
raise RuntimeError("Redis unavailable; cannot reset feature flags")
r.delete(_redis_key(name))
_apply_worker_side_effects(name)
# ---------------------------------------------------------------------------
# Worker-level side effects: when the admin flips a flag we don't just want
# gating at task-start (which still executes the message, it just returns
# 'skipped'). We also want queued work gone and the vision worker genuinely
# idle when the master switch is off.
# ---------------------------------------------------------------------------
_VISION_QUEUE = 'vision'
# Flag → celery task name(s) whose queued messages should be dropped when
# the flag goes off. Keeps the queue from replaying yesterday's work the
# moment someone re-enables the stage.
_TASKS_BY_FLAG: dict[str, tuple[str, ...]] = {
FLAG_VISION_ENABLED: (
'embed_photo', 'ocr_photo', 'detect_objects', 'extract_faces',
'classify_content', 'vision_fanout', 'recluster_faces',
),
FLAG_OCR_ENABLED: ('ocr_photo',),
FLAG_DETECTOR_ENABLED: ('detect_objects',),
FLAG_FACES_ENABLED: ('extract_faces', 'recluster_faces'),
FLAG_CLASSIFIER_ENABLED: ('classify_content',),
}
def _apply_worker_side_effects(name: str) -> None:
"""Bring the live workers in line with the new flag value.
For the master ``vision.enabled`` flag we go beyond task gating and
actually stop consumption from the ``vision`` queue — flipping it
off puts the vision worker to sleep (no CPU, no model memory
churn) until it's flipped back on. For per-feature flags, the
running tasks already skip via ``is_enabled``; we just purge any
messages already sitting in the queue so the admin doesn't pay for
a backlog on re-enable.
All operations are best-effort — if control messaging or a Redis
op fails, we log and return; the flag state itself is already
persisted so the gating path continues to work.
"""
try:
# Lazy import: avoids a circular dependency between the services
# module (imported from tasks.vision) and the celery app config.
from app.tasks.celery import celery_app
except Exception as e:
logger.warning(f"feature_flags: celery app unavailable for side effects ({e})")
return
try:
if name == FLAG_VISION_ENABLED:
if is_enabled(FLAG_VISION_ENABLED):
# Re-attach the vision consumer so workers pick up tasks
# again. broadcast=True ensures every running worker
# receives the command.
celery_app.control.add_consumer(_VISION_QUEUE, reply=False)
logger.info("feature_flags: vision re-enabled; consumer added")
else:
celery_app.control.cancel_consumer(_VISION_QUEUE, reply=False)
_purge_queue(_VISION_QUEUE)
logger.info(
"feature_flags: vision disabled; consumer cancelled "
"and queue purged"
)
return
# Per-feature flag going off → drop pending tasks of its types.
if not is_enabled(name):
targets = _TASKS_BY_FLAG.get(name, ())
if targets:
removed = _purge_queue_by_task_names(_VISION_QUEUE, targets)
logger.info(
f"feature_flags: {name} disabled; removed {removed} "
f"pending messages from {_VISION_QUEUE}"
)
except Exception as e:
logger.warning(f"feature_flags: worker side effects failed for {name}: {e}")
def _purge_queue(queue: str) -> int:
"""Drop every pending message from ``queue``. Returns the count
deleted. Celery's control.purge() purges the default queue only,
so we delete the Redis key directly (the broker's queue list)."""
r = _redis()
if r is None:
return 0
try:
removed = r.delete(queue)
return int(removed or 0)
except Exception as e:
logger.warning(f"feature_flags: purge {queue} failed: {e}")
return 0
def _purge_queue_by_task_names(queue: str, task_names: tuple[str, ...]) -> int:
"""Walk ``queue`` and drop any message whose Celery task name is in
``task_names``. Other messages are preserved (pushed back in order)
so we don't flush embed tasks when the admin disabled only OCR.
Celery stores each message as a JSON blob in a Redis list; the
task name lives at ``headers.task``.
"""
import json
r = _redis()
if r is None:
return 0
try:
# Snapshot the queue, then rebuild it without the filtered names.
# Done inside a Redis transaction so a concurrent enqueue doesn't
# race with us (worst case it gets re-delivered after we release,
# which is the normal enqueue path anyway).
pipe = r.pipeline()
pipe.lrange(queue, 0, -1)
pipe.delete(queue)
raw_items, _ = pipe.execute()
kept: list[bytes | str] = []
removed = 0
for raw in raw_items or []:
try:
# Messages can be bytes or str depending on decode_responses.
payload = raw.decode() if isinstance(raw, bytes) else raw
msg = json.loads(payload)
task = (
msg.get('headers', {}).get('task')
or msg.get('task')
)
if task in task_names:
removed += 1
continue
except Exception:
# Unparseable message — keep it, better to leak than
# to silently drop a message we can't identify.
pass
kept.append(raw)
if kept:
r.rpush(queue, *kept)
return removed
except Exception as e:
logger.warning(f"feature_flags: selective purge failed on {queue}: {e}")
return 0
def snapshot() -> dict[str, dict[str, object]]:
"""Return every flag's current effective value, YAML default, and
whether it's overridden. Powers the admin UI tab.
"""
r = _redis()
out: dict[str, dict[str, object]] = {}
for name in ALL_FLAGS:
default = _yaml_default(name)
override = None
if r is not None:
try:
raw = r.get(_redis_key(name))
if raw is not None:
override = raw.lower() == 'true'
except Exception:
pass
out[name] = {
'effective': override if override is not None else default,
'default': default,
'overridden': override is not None,
}
return out

View File

@@ -312,7 +312,25 @@ async def _generate_thumbnails_async(photo_id: str, task):
else:
logger.error(f"Unsupported media type: {photo.media_type}")
image = create_placeholder_thumbnail(photo.media_type)
# Fallback: some files wear a RAW/HEIC extension but are actually
# plain JPEGs — e.g. iPhones that write ProRAW-style .DNG for
# images where no RAW sensor data was captured, or re-exports
# that kept the original suffix. Pillow can open them directly,
# so before giving up, try reading the file as a standard image.
if not image and photo.media_type in ('raw', 'heic'):
try:
image = process_standard_image(photo.filepath)
if image is not None:
logger.info(
f"{photo.filepath}: {photo.media_type} decode failed "
f"but file opens as a standard image — using fallback"
)
except Exception as e:
logger.debug(
f"Standard-image fallback failed for {photo.filepath}: {e}"
)
if not image:
raise Exception("Failed to process image")
@@ -493,7 +511,16 @@ async def _backfill_phashes_async():
}
@shared_task(name='regroup_duplicates')
@shared_task(
name='regroup_duplicates',
# Full regroup scales with O(N²) on phash plus one pgvector query per
# embedded photo. On a 16k-photo library that's comfortably past the
# default 5-minute soft limit — bump to 2h / 2h30m. (Passing None here
# does NOT disable limits; Celery falls back to the worker default
# of 300s/600s. An explicit number overrides.)
soft_time_limit=7200,
time_limit=9000,
)
def regroup_duplicates_task():
"""Full recompute of duplicate groups (pHash + CLIP similarity).
@@ -502,7 +529,15 @@ def regroup_duplicates_task():
return asyncio.run(regroup_duplicates())
@shared_task(name='incremental_regroup_duplicates')
@shared_task(
name='incremental_regroup_duplicates',
# O(new × N); still cheaper than a full regroup but can easily exceed
# the 5-minute default after a big batch import. Same caveat as
# regroup_duplicates above — None would just re-inherit the worker
# default, so we pass explicit values.
soft_time_limit=3600,
time_limit=4200,
)
def incremental_regroup_duplicates_task(since_iso: str | None = None):
"""Incremental duplicate detection for newly added photos.

View File

@@ -20,6 +20,14 @@ from PIL import Image
from app.models.embeddings import Embedding
from app.config import settings
from app.services.feature_flags import (
is_enabled,
FLAG_VISION_ENABLED,
FLAG_OCR_ENABLED,
FLAG_DETECTOR_ENABLED,
FLAG_FACES_ENABLED,
FLAG_CLASSIFIER_ENABLED,
)
logger = logging.getLogger(__name__)
@@ -83,7 +91,7 @@ def _load_thumb(photo_id: str, size: str = "medium") -> np.ndarray | None:
@shared_task(name='embed_photo', queue='vision', bind=True, max_retries=3)
def embed_photo(self, photo_id: str):
"""Generate CLIP embedding for a photo and store in pgvector."""
if not settings.vision.enabled:
if not is_enabled(FLAG_VISION_ENABLED):
return {'status': 'skipped', 'reason': 'vision disabled'}
image = _load_thumb(photo_id, "medium") # 640px
@@ -128,18 +136,18 @@ def embed_photo(self, photo_id: str):
@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:
if not is_enabled(FLAG_VISION_ENABLED):
return {'status': 'skipped', 'reason': 'vision disabled'}
embed_photo.delay(photo_id)
if settings.vision.ocr.enabled:
if is_enabled(FLAG_OCR_ENABLED):
ocr_photo.delay(photo_id)
if settings.vision.detector.enabled:
if is_enabled(FLAG_DETECTOR_ENABLED):
detect_objects.delay(photo_id)
if settings.vision.faces.enabled:
if is_enabled(FLAG_FACES_ENABLED):
extract_faces.delay(photo_id)
if settings.vision.classifier.enabled:
if is_enabled(FLAG_CLASSIFIER_ENABLED):
classify_content.delay(photo_id)
return {'status': 'dispatched', 'photo_id': photo_id}
@@ -148,7 +156,7 @@ def vision_fanout(photo_id: str):
@shared_task(name='ocr_photo', queue='vision', bind=True, max_retries=3)
def ocr_photo(self, photo_id: str):
"""Run OCR on a photo and store text regions."""
if not settings.vision.enabled or not settings.vision.ocr.enabled:
if not is_enabled(FLAG_VISION_ENABLED) or not is_enabled(FLAG_OCR_ENABLED):
return {'status': 'skipped', 'reason': 'OCR disabled'}
image = _load_thumb(photo_id, "large") # 1280px for better OCR accuracy
@@ -195,7 +203,7 @@ def ocr_photo(self, photo_id: str):
def detect_objects(self, photo_id: str):
"""Detect objects in a photo, create Tag(kind=object) rows, and
link via photo_tags with confidence/bbox/source."""
if not settings.vision.enabled or not settings.vision.detector.enabled:
if not is_enabled(FLAG_VISION_ENABLED) or not is_enabled(FLAG_DETECTOR_ENABLED):
return {'status': 'skipped', 'reason': 'detection disabled'}
image = _load_thumb(photo_id, "medium") # 640px
@@ -279,7 +287,7 @@ def detect_objects(self, photo_id: str):
def classify_content(self, photo_id: str):
"""Classify image content type (screenshot, document, artwork, etc.)
using CLIP zero-shot classification. Writes Tag(kind=content_type)."""
if not settings.vision.enabled or not settings.vision.classifier.enabled:
if not is_enabled(FLAG_VISION_ENABLED) or not is_enabled(FLAG_CLASSIFIER_ENABLED):
return {'status': 'skipped', 'reason': 'classifier disabled'}
image = _load_thumb(photo_id, "medium")
@@ -394,7 +402,7 @@ def extract_faces(self, photo_id: str):
"""Detect faces and store recognition embeddings using InsightFace
(RetinaFace + ArcFace). No YOLO workaround needed — RetinaFace has
strong human-vs-non-human precision on its own."""
if not settings.vision.enabled or not settings.vision.faces.enabled:
if not is_enabled(FLAG_VISION_ENABLED) or not is_enabled(FLAG_FACES_ENABLED):
return {'status': 'skipped', 'reason': 'faces disabled'}
image = _load_original(photo_id)
@@ -478,7 +486,7 @@ def recluster_faces(self):
logger.info("Vision worker not ready yet — retrying in 30s")
raise self.retry(countdown=30)
if not settings.vision.enabled or not settings.vision.faces.enabled:
if not is_enabled(FLAG_VISION_ENABLED) or not is_enabled(FLAG_FACES_ENABLED):
return {'status': 'skipped', 'reason': 'faces disabled'}
from app.models import Photo
@@ -603,7 +611,7 @@ def backfill_vision(self, task: str | None = None, limit: int | None = None):
embed_ids = [r[0] for r in session.execute(sa_text(sql), params).fetchall()]
ocr_ids = []
if task in ('ocr', None) and settings.vision.ocr.enabled:
if task in ('ocr', None) and is_enabled(FLAG_OCR_ENABLED):
sql = f"""
SELECT p.id FROM photos p
LEFT JOIN ocr_text o ON o.photo_id = p.id
@@ -613,7 +621,7 @@ def backfill_vision(self, task: str | None = None, limit: int | None = None):
ocr_ids = [r[0] for r in session.execute(sa_text(sql), params).fetchall()]
detect_ids = []
if task in ('detect', None) and settings.vision.detector.enabled:
if task in ('detect', None) and is_enabled(FLAG_DETECTOR_ENABLED):
sql = f"""
SELECT p.id FROM photos p
WHERE p.processing_status = 'completed'
@@ -626,7 +634,7 @@ def backfill_vision(self, task: str | None = None, limit: int | None = None):
detect_ids = [r[0] for r in session.execute(sa_text(sql), params).fetchall()]
face_ids = []
if task in ('faces', None) and settings.vision.faces.enabled:
if task in ('faces', None) and is_enabled(FLAG_FACES_ENABLED):
sql = f"""
SELECT p.id FROM photos p
LEFT JOIN face_embeddings fe ON fe.photo_id = p.id
@@ -636,7 +644,7 @@ def backfill_vision(self, task: str | None = None, limit: int | None = None):
face_ids = [r[0] for r in session.execute(sa_text(sql), params).fetchall()]
classify_ids = []
if task in ('classify', None) and settings.vision.classifier.enabled:
if task in ('classify', None) and is_enabled(FLAG_CLASSIFIER_ENABLED):
sql = f"""
SELECT p.id FROM photos p
WHERE p.processing_status = 'completed'