""" 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: → "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