From f61a0446e7f892a90592d6c55bce630be3e17857 Mon Sep 17 00:00:00 2001 From: Ryan Date: Mon, 10 Aug 2026 19:42:24 -0400 Subject: [PATCH] Story 9 --- config.yaml | 2 + src/crash_recovery.py | 137 +++++++++++ src/drift_detector.py | 219 +++++++++++++++++ src/health_check.py | 134 +++++++++++ src/metrics.py | 261 +++++++++++++++++++++ src/retry.py | 114 +++++++++ tests/test_story_09.py | 521 +++++++++++++++++++++++++++++++++++++++++ 7 files changed, 1388 insertions(+) create mode 100644 src/crash_recovery.py create mode 100644 src/drift_detector.py create mode 100644 src/health_check.py create mode 100644 src/metrics.py create mode 100644 src/retry.py create mode 100644 tests/test_story_09.py diff --git a/config.yaml b/config.yaml index f8eec6d..aef05d3 100644 --- a/config.yaml +++ b/config.yaml @@ -232,6 +232,8 @@ monitoring: enabled: true schedule: "0 2 * * 0" # cron: Sundays at 2 AM baseline_source: db + health_check: + port: 8080 crash_recovery: lock_timeout_minutes: 5 auto_requeue: true diff --git a/src/crash_recovery.py b/src/crash_recovery.py new file mode 100644 index 0000000..dff874e --- /dev/null +++ b/src/crash_recovery.py @@ -0,0 +1,137 @@ +""" +Crash recovery and idempotent processing guarantees. + +On startup: re-queues any video stuck in PROCESSING past the lock timeout. +During processing: saves per-video checkpoint JSON to scratch for resume. +""" + +import json +import logging +from datetime import datetime, timezone +from pathlib import Path +from typing import Dict, List, Optional + +logger = logging.getLogger(__name__) + +_CHECKPOINT_DIR = "checkpoints" + + +class CrashRecovery: + """Handle crash recovery, checkpointing, and idempotent re-queueing.""" + + def __init__(self, db_connector, config): + self._db = db_connector + scratch = config.get("storage.scratch_path", "/scratch") + self._checkpoint_dir = Path(scratch) / _CHECKPOINT_DIR + mon = config.get_section("monitoring") + cr = mon.get("crash_recovery", {}) + self._lock_timeout_minutes: int = int(cr.get("lock_timeout_minutes", 5)) + self._auto_requeue: bool = bool(cr.get("auto_requeue", True)) + + # ------------------------------------------------------------------ + # Startup recovery + # ------------------------------------------------------------------ + + def recover_on_startup(self) -> int: + """ + Re-queue videos stuck in PROCESSING past the lock timeout. + + Returns the count of videos re-queued. + """ + if not self._auto_requeue: + logger.info("auto_requeue disabled; skipping crash recovery") + return 0 + + try: + result = self._db.execute( + """UPDATE videos + SET status = 'PENDING', updated_at = NOW() + WHERE status = 'PROCESSING' + AND updated_at < NOW() - INTERVAL %s MINUTE""", + (self._lock_timeout_minutes,), + ) + count = result if isinstance(result, int) else 0 + if count: + logger.warning("Crash recovery: re-queued %d stuck PROCESSING videos", count) + else: + logger.info("Crash recovery: no stuck videos found") + return count + except Exception as exc: + logger.error("Crash recovery query failed: %s", exc) + return 0 + + def list_stuck_videos(self) -> List[Dict]: + """Return videos currently stuck past the lock timeout (read-only).""" + try: + rows = self._db.fetchall( + """SELECT id, file_path, updated_at + FROM videos + WHERE status = 'PROCESSING' + AND updated_at < NOW() - INTERVAL %s MINUTE""", + (self._lock_timeout_minutes,), + ) + return rows or [] + except Exception as exc: + logger.error("Failed to list stuck videos: %s", exc) + return [] + + # ------------------------------------------------------------------ + # Per-video checkpointing + # ------------------------------------------------------------------ + + def save_checkpoint(self, video_id: int, state: str, progress: Dict) -> None: + """Persist a checkpoint JSON for the given video.""" + self._checkpoint_dir.mkdir(parents=True, exist_ok=True) + checkpoint = { + "video_id": video_id, + "state": state, + "progress": progress, + "timestamp": datetime.now(timezone.utc).isoformat(), + } + path = self._checkpoint_dir / f"{video_id}.json" + try: + path.write_text(json.dumps(checkpoint, indent=2), encoding="utf-8") + logger.debug("Saved checkpoint for video %d (state=%s)", video_id, state) + except Exception as exc: + logger.warning("Failed to save checkpoint for video %d: %s", video_id, exc) + + def load_checkpoint(self, video_id: int) -> Optional[Dict]: + """Return the stored checkpoint for video_id, or None if absent.""" + path = self._checkpoint_dir / f"{video_id}.json" + if not path.exists(): + return None + try: + return json.loads(path.read_text(encoding="utf-8")) + except Exception as exc: + logger.warning("Failed to read checkpoint for video %d: %s", video_id, exc) + return None + + def delete_checkpoint(self, video_id: int) -> None: + """Remove the checkpoint file after successful processing.""" + path = self._checkpoint_dir / f"{video_id}.json" + if path.exists(): + try: + path.unlink() + logger.debug("Deleted checkpoint for video %d", video_id) + except Exception as exc: + logger.warning("Failed to delete checkpoint for video %d: %s", video_id, exc) + + # ------------------------------------------------------------------ + # Idempotency guard + # ------------------------------------------------------------------ + + def is_already_completed(self, video_id: int) -> bool: + """ + Return True if this video already has a COMPLETED status. + + Prevents duplicate processing if a job is accidentally re-submitted. + """ + try: + rows = self._db.fetchall( + "SELECT status FROM videos WHERE id = %s", + (video_id,), + ) + return bool(rows and rows[0].get("status") == "COMPLETED") + except Exception as exc: + logger.warning("Idempotency check failed for video %d: %s", video_id, exc) + return False diff --git a/src/drift_detector.py b/src/drift_detector.py new file mode 100644 index 0000000..9a98fd1 --- /dev/null +++ b/src/drift_detector.py @@ -0,0 +1,219 @@ +""" +Confidence distribution drift detection. + +Compares recent video confidence scores against a stored baseline and fires +alerts when the distribution shifts beyond configured thresholds. +""" + +import logging +from datetime import datetime, timedelta, timezone +from typing import Callable, Dict, List, Optional + +logger = logging.getLogger(__name__) + + +class DriftDetector: + """Detect and alert on confidence distribution drift.""" + + def __init__(self, db_connector, config, alert_fn: Optional[Callable[[str], None]] = None): + self._db = db_connector + mon = config.get_section("monitoring") + alerts = mon.get("alerts", {}) + self._drift_threshold: float = float(alerts.get("confidence_drift_threshold", 0.10)) + self._review_max_size: int = int(alerts.get("review_queue_max_size", 1000)) + self._review_max_age_h: int = int(alerts.get("review_queue_max_age_hours", 24)) + self._throughput_min: float = float(alerts.get("throughput_min_videos_per_hour", 20)) + self._throughput_window_h: int = int(alerts.get("throughput_min_duration_hours", 1)) + self._error_rate_threshold: float = float(alerts.get("error_rate_threshold", 0.05)) + self._error_rate_window_h: int = int(alerts.get("error_rate_window_hours", 1)) + self._alert_fn = alert_fn or _log_alert + + # ------------------------------------------------------------------ + # Public API + # ------------------------------------------------------------------ + + def run_all_checks(self) -> Dict[str, bool]: + """ + Run all alert checks and return a dict of {check_name: triggered}. + + Fires the alert function for each triggered condition. + """ + results = {} + results["confidence_drift"] = self.check_confidence_drift() + results["review_queue_growth"] = self.check_review_queue_growth() + results["low_throughput"] = self.check_low_throughput() + results["high_error_rate"] = self.check_error_rate() + return results + + def check_confidence_drift(self) -> bool: + """ + Compare p(C > 0.5) in recent processing logs to the stored baseline. + + Returns True if drift exceeds the threshold. + """ + current = self._fetch_recent_confidences(hours=24 * 7) + baseline = self._fetch_baseline_confidences() + if not current or not baseline: + logger.debug("Not enough data for drift check") + return False + + triggered = detect_drift(current, baseline, self._drift_threshold) + if triggered: + current_high = _high_conf_fraction(current) + baseline_high = _high_conf_fraction(baseline) + shift = abs(current_high - baseline_high) + self._alert_fn( + f"Confidence drift detected: {shift:.1%} shift in p(C>0.5) " + f"(current={current_high:.1%}, baseline={baseline_high:.1%}, " + f"threshold={self._drift_threshold:.1%})" + ) + return triggered + + def check_review_queue_growth(self) -> bool: + """Alert if the unannotated review queue has been above max_size for max_age hours.""" + try: + cutoff = datetime.now(timezone.utc) - timedelta(hours=self._review_max_age_h) + rows = self._db.fetchall( + """SELECT COUNT(*) AS cnt + FROM review_queue + WHERE annotated = FALSE + AND created_at < %s""", + (cutoff.replace(tzinfo=None),), + ) + count = int((rows or [{"cnt": 0}])[0]["cnt"]) + if count > self._review_max_size: + self._alert_fn( + f"Review queue growth alert: {count} unannotated items older than " + f"{self._review_max_age_h}h (threshold={self._review_max_size})" + ) + return True + return False + except Exception as exc: + logger.warning("Review queue growth check failed: %s", exc) + return False + + def check_low_throughput(self) -> bool: + """Alert if throughput dropped below the minimum for the last window hours.""" + try: + cutoff = datetime.now(timezone.utc) - timedelta(hours=self._throughput_window_h) + rows = self._db.fetchall( + """SELECT COUNT(*) AS cnt + FROM videos + WHERE status = 'COMPLETED' + AND last_processed_time >= %s""", + (cutoff.replace(tzinfo=None),), + ) + count = int((rows or [{"cnt": 0}])[0]["cnt"]) + videos_per_hour = count / max(self._throughput_window_h, 1) + if videos_per_hour < self._throughput_min: + self._alert_fn( + f"Low throughput alert: {videos_per_hour:.1f} videos/hour " + f"over last {self._throughput_window_h}h " + f"(minimum={self._throughput_min})" + ) + return True + return False + except Exception as exc: + logger.warning("Throughput check failed: %s", exc) + return False + + def check_error_rate(self) -> bool: + """Alert if the fraction of ERROR/UNSCANNABLE videos exceeds the threshold.""" + try: + cutoff = datetime.now(timezone.utc) - timedelta(hours=self._error_rate_window_h) + rows = self._db.fetchall( + """SELECT status, COUNT(*) AS cnt + FROM videos + WHERE last_processed_time >= %s + AND status IN ('COMPLETED', 'ERROR', 'UNSCANNABLE') + GROUP BY status""", + (cutoff.replace(tzinfo=None),), + ) + counts: Dict[str, int] = {r["status"]: int(r["cnt"]) for r in (rows or [])} + total = sum(counts.values()) + if total == 0: + return False + errors = counts.get("ERROR", 0) + counts.get("UNSCANNABLE", 0) + rate = errors / total + if rate > self._error_rate_threshold: + self._alert_fn( + f"High error rate alert: {rate:.1%} over last {self._error_rate_window_h}h " + f"({errors}/{total} videos, threshold={self._error_rate_threshold:.1%})" + ) + return True + return False + except Exception as exc: + logger.warning("Error rate check failed: %s", exc) + return False + + # ------------------------------------------------------------------ + # DB helpers + # ------------------------------------------------------------------ + + def _fetch_recent_confidences(self, hours: int = 168) -> List[float]: + """Return confidence scores from the last `hours` hours.""" + try: + cutoff = datetime.now(timezone.utc) - timedelta(hours=hours) + rows = self._db.fetchall( + """SELECT confidence_score + FROM processing_logs + WHERE processed_at >= %s + AND confidence_score IS NOT NULL""", + (cutoff.replace(tzinfo=None),), + ) + return [float(r["confidence_score"]) for r in (rows or [])] + except Exception as exc: + logger.warning("Failed to fetch recent confidences: %s", exc) + return [] + + def _fetch_baseline_confidences(self) -> List[float]: + """Return the baseline confidence distribution stored in the DB.""" + try: + rows = self._db.fetchall( + """SELECT confidence_score + FROM processing_logs + WHERE confidence_score IS NOT NULL + ORDER BY processed_at ASC + LIMIT 10000""" + ) + return [float(r["confidence_score"]) for r in (rows or [])] + except Exception as exc: + logger.warning("Failed to fetch baseline confidences: %s", exc) + return [] + + +# --------------------------------------------------------------------------- +# Module-level pure functions (testable without DB) +# --------------------------------------------------------------------------- + +def detect_drift( + current_confidences: List[float], + baseline_confidences: List[float], + threshold: float = 0.10, +) -> bool: + """ + Return True if |p(C>0.5)_current - p(C>0.5)_baseline| > threshold. + + Both lists must be non-empty. + """ + if not current_confidences or not baseline_confidences: + return False + current_high = _high_conf_fraction(current_confidences) + baseline_high = _high_conf_fraction(baseline_confidences) + shift = abs(current_high - baseline_high) + if shift > threshold: + logger.warning( + "Drift detected: shift=%.2f%% (current=%.2f%% baseline=%.2f%%)", + shift * 100, current_high * 100, baseline_high * 100, + ) + return True + return False + + +def _high_conf_fraction(confidences: List[float]) -> float: + """Return the fraction of scores > 0.5.""" + return sum(1 for c in confidences if c > 0.5) / len(confidences) + + +def _log_alert(message: str): + logger.warning("ALERT: %s", message) diff --git a/src/health_check.py b/src/health_check.py new file mode 100644 index 0000000..eca7abd --- /dev/null +++ b/src/health_check.py @@ -0,0 +1,134 @@ +""" +Worker health check HTTP endpoint. + +Exposes GET /health returning a JSON status document. +Designed to be started in a background thread alongside the main worker. +""" + +import json +import logging +import os +import threading +import time +from datetime import datetime, timezone +from http.server import BaseHTTPRequestHandler, HTTPServer +from pathlib import Path +from typing import Any, Dict, Optional + +logger = logging.getLogger(__name__) + +_START_TIME = time.monotonic() + + +class HealthStatus: + """Mutable singleton updated by the worker as it processes videos.""" + + def __init__(self): + self._lock = threading.Lock() + self._data: Dict[str, Any] = { + "status": "starting", + "gpu_available": False, + "gpu_memory_used_gb": 0.0, + "queue_depth": 0, + "videos_processed_today": 0, + "last_error": None, + } + + def update(self, **kwargs): + with self._lock: + self._data.update(kwargs) + + def snapshot(self) -> Dict[str, Any]: + with self._lock: + data = dict(self._data) + data["uptime_seconds"] = round(time.monotonic() - _START_TIME, 1) + return data + + +# Module-level singleton shared between the HTTP handler and the worker +health = HealthStatus() + + +class _HealthHandler(BaseHTTPRequestHandler): + """Minimal HTTP handler — only responds to GET /health.""" + + def do_GET(self): # noqa: N802 + if self.path.rstrip("/") == "/health": + body = json.dumps(health.snapshot(), default=str).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + else: + self.send_response(404) + self.end_headers() + + def log_message(self, fmt, *args): # silence default access log + logger.debug("health_check: " + fmt, *args) + + +def start_health_server(port: int = 8080) -> threading.Thread: + """ + Start the health HTTP server in a daemon thread. + + Returns the thread so callers can join if needed. + """ + server = HTTPServer(("0.0.0.0", port), _HealthHandler) + + def _serve(): + logger.info("Health check server listening on port %d", port) + server.serve_forever() + + t = threading.Thread(target=_serve, daemon=True, name="health-server") + t.start() + return t + + +def refresh_health_from_system(db_connector=None, scratch_path: str = "/scratch"): + """ + Update the health singleton with current GPU, queue, and disk state. + + Call this periodically from the worker main loop. + """ + updates: Dict[str, Any] = {"status": "healthy"} + + # GPU state + try: + import pynvml + pynvml.nvmlInit() + handle = pynvml.nvmlDeviceGetHandleByIndex(0) + mem = pynvml.nvmlDeviceGetMemoryInfo(handle) + updates["gpu_available"] = True + updates["gpu_memory_used_gb"] = round(mem.used / 1024**3, 2) + pynvml.nvmlShutdown() + except Exception: + updates["gpu_available"] = _torch_gpu_available() + + # Queue depth from DB + if db_connector is not None: + try: + rows = db_connector.fetchall( + "SELECT COUNT(*) AS cnt FROM videos WHERE status IN ('PENDING','PROCESSING')" + ) + updates["queue_depth"] = int((rows or [{"cnt": 0}])[0]["cnt"]) + + today = datetime.now(timezone.utc).date().isoformat() + rows2 = db_connector.fetchall( + "SELECT COUNT(*) AS cnt FROM videos " + "WHERE status = 'COMPLETED' AND DATE(last_processed_time) = %s", + (today,), + ) + updates["videos_processed_today"] = int((rows2 or [{"cnt": 0}])[0]["cnt"]) + except Exception as exc: + logger.debug("Health DB query failed: %s", exc) + + health.update(**updates) + + +def _torch_gpu_available() -> bool: + try: + import torch + return torch.cuda.is_available() + except ImportError: + return False diff --git a/src/metrics.py b/src/metrics.py new file mode 100644 index 0000000..6b6361b --- /dev/null +++ b/src/metrics.py @@ -0,0 +1,261 @@ +""" +Prometheus metrics definitions for VideoDetect. + +All metric objects are module-level singletons — import and call them directly. +Call start_metrics_server() once at worker startup to expose /metrics. +""" + +import logging +import threading +from typing import Optional + +logger = logging.getLogger(__name__) + +# Populated on first _ensure_prometheus() call; None if library unavailable. +_prometheus_available: Optional[bool] = None +_server_started = False +_server_lock = threading.Lock() + +# --------------------------------------------------------------------------- +# Lazy metric definitions — avoid import errors when prometheus_client absent +# --------------------------------------------------------------------------- + +def _ensure_prometheus(): + global _prometheus_available + if _prometheus_available is not None: + return _prometheus_available + try: + import prometheus_client # noqa: F401 + _prometheus_available = True + except ImportError: + logger.warning("prometheus_client not installed; metrics disabled") + _prometheus_available = False + return _prometheus_available + + +class _NoOpMetric: + """Drop-in replacement used when prometheus_client is unavailable.""" + def labels(self, **_): + return self + def inc(self, *_, **__): + pass + def set(self, *_, **__): + pass + def observe(self, *_, **__): + pass + + +def _counter(name, doc, labelnames=()): + if not _ensure_prometheus(): + return _NoOpMetric() + from prometheus_client import Counter + return Counter(name, doc, list(labelnames)) + + +def _gauge(name, doc, labelnames=()): + if not _ensure_prometheus(): + return _NoOpMetric() + from prometheus_client import Gauge + return Gauge(name, doc, list(labelnames)) + + +def _histogram(name, doc, labelnames=(), buckets=None): + if not _ensure_prometheus(): + return _NoOpMetric() + from prometheus_client import Histogram + kwargs = {"labelnames": list(labelnames)} + if buckets is not None: + kwargs["buckets"] = buckets + return Histogram(name, doc, **kwargs) + + +# --------------------------------------------------------------------------- +# Counters +# --------------------------------------------------------------------------- + +videos_processed_total = _counter( + "videos_processed_total", + "Total videos processed", + ["routing_decision"], +) + +videos_processed_by_model_total = _counter( + "videos_processed_by_model_total", + "Total videos processed per model version", + ["model_version"], +) + +frames_extracted_total = _counter( + "frames_extracted_total", + "Total frames extracted across all videos", +) + +faces_detected_total = _counter( + "faces_detected_total", + "Total face crops detected", +) + +inference_errors_total = _counter( + "inference_errors_total", + "Total inference errors by type", + ["error_type"], +) + +retry_attempts_total = _counter( + "retry_attempts_total", + "Total retry attempts by pipeline step", + ["step"], +) + +# --------------------------------------------------------------------------- +# Gauges +# --------------------------------------------------------------------------- + +gpu_utilization_percent = _gauge( + "gpu_utilization_percent", + "GPU utilization as a percentage", + ["gpu_id"], +) + +gpu_memory_used_bytes = _gauge( + "gpu_memory_used_bytes", + "GPU memory currently allocated in bytes", + ["gpu_id"], +) + +gpu_memory_free_bytes = _gauge( + "gpu_memory_free_bytes", + "GPU memory currently free in bytes", + ["gpu_id"], +) + +queue_depth_pending = _gauge( + "queue_depth_pending", + "Number of videos in PENDING state", +) + +queue_depth_processing = _gauge( + "queue_depth_processing", + "Number of videos currently in PROCESSING state", +) + +queue_depth_review = _gauge( + "queue_depth_review", + "Number of unannotated videos in the review queue", +) + +scratch_usage_bytes = _gauge( + "scratch_usage_bytes", + "Scratch space used in bytes", +) + +scratch_usage_percent = _gauge( + "scratch_usage_percent", + "Scratch space used as a percentage of total", +) + +# --------------------------------------------------------------------------- +# Histograms +# --------------------------------------------------------------------------- + +_CONFIDENCE_BUCKETS = [0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0] + +video_processing_duration_seconds = _histogram( + "video_processing_duration_seconds", + "End-to-end processing time per video in seconds", + ["routing_decision"], +) + +confidence_score_distribution = _histogram( + "confidence_score_distribution", + "Distribution of video-level confidence scores", + buckets=_CONFIDENCE_BUCKETS, +) + +frame_count_per_video = _histogram( + "frame_count_per_video", + "Number of frames sampled per video", + buckets=[1, 2, 5, 10, 20, 50, 100, 200, 500], +) + +face_count_per_video = _histogram( + "face_count_per_video", + "Number of face crops detected per video", + buckets=[0, 1, 2, 5, 10, 20, 50, 100], +) + +# --------------------------------------------------------------------------- +# Server startup +# --------------------------------------------------------------------------- + +def start_metrics_server(port: int = 9090): + """Start the Prometheus HTTP metrics endpoint (idempotent).""" + global _server_started + with _server_lock: + if _server_started: + return + if not _ensure_prometheus(): + logger.warning("Cannot start metrics server: prometheus_client unavailable") + return + try: + from prometheus_client import start_http_server + start_http_server(port) + _server_started = True + logger.info("Prometheus metrics server started on port %d", port) + except Exception as exc: + logger.error("Failed to start metrics server: %s", exc) + + +# --------------------------------------------------------------------------- +# Convenience update functions +# --------------------------------------------------------------------------- + +def update_queue_depths(db_connector): + """Query the DB and refresh all queue-depth gauges.""" + try: + rows = db_connector.fetchall( + """SELECT status, COUNT(*) AS cnt + FROM videos + WHERE status IN ('PENDING', 'PROCESSING') + GROUP BY status""" + ) + counts = {r["status"]: int(r["cnt"]) for r in (rows or [])} + queue_depth_pending.set(counts.get("PENDING", 0)) + queue_depth_processing.set(counts.get("PROCESSING", 0)) + + review_rows = db_connector.fetchall( + "SELECT COUNT(*) AS cnt FROM review_queue WHERE annotated = FALSE" + ) + queue_depth_review.set(int((review_rows or [{"cnt": 0}])[0]["cnt"])) + except Exception as exc: + logger.warning("Failed to update queue depth metrics: %s", exc) + + +def update_gpu_metrics(): + """Poll nvidia-smi (via pynvml) and refresh GPU gauges.""" + try: + import pynvml + pynvml.nvmlInit() + count = pynvml.nvmlDeviceGetCount() + for i in range(count): + handle = pynvml.nvmlDeviceGetHandleByIndex(i) + util = pynvml.nvmlDeviceGetUtilizationRates(handle) + mem = pynvml.nvmlDeviceGetMemoryInfo(handle) + gpu_id = str(i) + gpu_utilization_percent.labels(gpu_id=gpu_id).set(util.gpu) + gpu_memory_used_bytes.labels(gpu_id=gpu_id).set(mem.used) + gpu_memory_free_bytes.labels(gpu_id=gpu_id).set(mem.free) + pynvml.nvmlShutdown() + except Exception as exc: + logger.debug("GPU metrics unavailable: %s", exc) + + +def update_scratch_metrics(scratch_path: str = "/scratch"): + """Measure scratch disk usage and refresh gauges.""" + try: + import shutil + total, used, free = shutil.disk_usage(scratch_path) + scratch_usage_bytes.set(used) + scratch_usage_percent.set(round(used / max(total, 1) * 100, 2)) + except Exception as exc: + logger.debug("Scratch metrics unavailable: %s", exc) diff --git a/src/retry.py b/src/retry.py new file mode 100644 index 0000000..343d583 --- /dev/null +++ b/src/retry.py @@ -0,0 +1,114 @@ +""" +Retry decorator with exponential backoff for transient failures. + +Usage: + @retry(max_attempts=3, step="extract") + def extract_frames(...): ... + +Non-retryable error types are re-raised immediately without consuming attempts. +""" + +import functools +import logging +import time +from typing import Callable, Optional, Tuple, Type + +logger = logging.getLogger(__name__) + +# Error class names that should never be retried +_NON_RETRYABLE_NAMES = frozenset({ + "CodecUnsupportedError", + "FileCorruptError", + "InvalidPathError", + "FileNotFoundError", + "PermissionError", + "IsADirectoryError", +}) + + +class RetryExhaustedError(Exception): + """Raised when all retry attempts have been exhausted.""" + + +def retry( + max_attempts: int = 3, + initial_delay: float = 1.0, + backoff_factor: float = 2.0, + exceptions: Tuple[Type[Exception], ...] = (Exception,), + step: str = "unknown", +): + """ + Decorator: retry ``func`` up to ``max_attempts`` times on retryable exceptions. + + Non-retryable exceptions (see _NON_RETRYABLE_NAMES) propagate immediately. + Each retry waits ``initial_delay * backoff_factor ** attempt`` seconds. + """ + def decorator(func: Callable) -> Callable: + @functools.wraps(func) + def wrapper(*args, **kwargs): + last_exc: Optional[Exception] = None + + for attempt in range(max_attempts): + try: + return func(*args, **kwargs) + except exceptions as exc: + if _is_non_retryable(exc): + logger.error( + "Non-retryable error in step=%s (attempt %d/%d): %s: %s", + step, attempt + 1, max_attempts, + type(exc).__name__, exc, + ) + raise + + last_exc = exc + _increment_retry_counter(step) + + if attempt < max_attempts - 1: + delay = initial_delay * (backoff_factor ** attempt) + logger.warning( + "Transient error in step=%s (attempt %d/%d), " + "retrying in %.1fs: %s: %s", + step, attempt + 1, max_attempts, delay, + type(exc).__name__, exc, + ) + time.sleep(delay) + else: + logger.error( + "All %d attempts exhausted in step=%s: %s: %s", + max_attempts, step, type(exc).__name__, exc, + ) + + raise RetryExhaustedError( + f"step={step} failed after {max_attempts} attempts" + ) from last_exc + + return wrapper + return decorator + + +def retry_from_config(config, step: str = "unknown"): + """Build a ``@retry`` decorator from config.yaml monitoring.retry settings.""" + mon = config.get_section("monitoring") + cfg = mon.get("retry", {}) + return retry( + max_attempts=int(cfg.get("max_attempts", 3)), + initial_delay=float(cfg.get("initial_delay", 1.0)), + backoff_factor=float(cfg.get("backoff_factor", 2.0)), + step=step, + ) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _is_non_retryable(exc: Exception) -> bool: + return type(exc).__name__ in _NON_RETRYABLE_NAMES + + +def _increment_retry_counter(step: str): + try: + from metrics import retry_attempts_total + retry_attempts_total.labels(step=step).inc() + except Exception: + pass # metrics not available; don't break the retry logic diff --git a/tests/test_story_09.py b/tests/test_story_09.py new file mode 100644 index 0000000..b46feae --- /dev/null +++ b/tests/test_story_09.py @@ -0,0 +1,521 @@ +""" +Tests for Story 09: Observability, Monitoring & Hardening. + +Covers: + - metrics.py: NoOp fallback, update_queue_depths, update_scratch_metrics + - crash_recovery.py: recover_on_startup, checkpointing, idempotency guard + - retry.py: successful call, retry on transient error, non-retryable bypass, exhaustion + - drift_detector.py: detect_drift logic, all alert checks + - health_check.py: HealthStatus snapshot, HTTP /health endpoint +""" + +import json +import os +import sys +import tempfile +import threading +import time +import unittest +import urllib.request +from pathlib import Path +from unittest.mock import MagicMock, call, patch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + +from crash_recovery import CrashRecovery +from drift_detector import DriftDetector, detect_drift +from health_check import HealthStatus, start_health_server, health +from retry import RetryExhaustedError, retry, _is_non_retryable + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_config(overrides: dict = None): + cfg = MagicMock() + mon_defaults = { + "alerts": { + "confidence_drift_threshold": 0.10, + "review_queue_max_size": 1000, + "review_queue_max_age_hours": 24, + "throughput_min_videos_per_hour": 20, + "throughput_min_duration_hours": 1, + "error_rate_threshold": 0.05, + "error_rate_window_hours": 1, + }, + "drift_detection": {"enabled": True, "schedule": "0 2 * * 0", "baseline_source": "db"}, + "crash_recovery": {"lock_timeout_minutes": 5, "auto_requeue": True}, + "retry": {"max_attempts": 3, "initial_delay": 0.0, "backoff_factor": 2.0}, + } + if overrides: + mon_defaults.update(overrides) + + def _get_section(section): + if section == "monitoring": + return mon_defaults + return {} + + def _get(path, default=None): + mapping = { + "storage.scratch_path": "/tmp/test_scratch", + "storage.models_path": "/tmp/test_models", + } + return mapping.get(path, default) + + cfg.get_section.side_effect = _get_section + cfg.get.side_effect = _get + return cfg + + +def _make_db(): + return MagicMock() + + +# --------------------------------------------------------------------------- +# metrics.py tests +# --------------------------------------------------------------------------- + +class TestMetricsNoOp(unittest.TestCase): + """The _NoOpMetric must absorb all method calls without raising.""" + + def test_noop_labels_inc_does_not_raise(self): + from metrics import _NoOpMetric + m = _NoOpMetric() + m.labels(routing_decision="MATCH").inc() + + def test_noop_set_does_not_raise(self): + from metrics import _NoOpMetric + m = _NoOpMetric() + m.set(42) + + def test_noop_observe_does_not_raise(self): + from metrics import _NoOpMetric + m = _NoOpMetric() + m.observe(0.75) + + +class TestMetricsQueueDepths(unittest.TestCase): + def test_update_queue_depths_sets_gauges(self): + from metrics import update_queue_depths, queue_depth_pending, queue_depth_processing + + db = MagicMock() + db.fetchall.side_effect = [ + [{"status": "PENDING", "cnt": 10}, {"status": "PROCESSING", "cnt": 3}], + [{"cnt": 7}], + ] + # Should not raise even if prometheus is absent + update_queue_depths(db) + + def test_update_queue_depths_handles_db_error_gracefully(self): + from metrics import update_queue_depths + db = MagicMock() + db.fetchall.side_effect = Exception("DB down") + update_queue_depths(db) # must not raise + + +class TestMetricsScratch(unittest.TestCase): + def test_update_scratch_metrics_runs_without_error(self): + from metrics import update_scratch_metrics + with tempfile.TemporaryDirectory() as tmpdir: + update_scratch_metrics(tmpdir) # must not raise + + def test_update_scratch_metrics_handles_missing_path(self): + from metrics import update_scratch_metrics + update_scratch_metrics("/nonexistent_path_xyz") # must not raise + + +# --------------------------------------------------------------------------- +# crash_recovery.py tests +# --------------------------------------------------------------------------- + +class TestCrashRecoveryRequeue(unittest.TestCase): + def test_recover_on_startup_calls_update(self): + db = _make_db() + db.execute.return_value = 3 + cfg = _make_config() + cr = CrashRecovery(db, cfg) + count = cr.recover_on_startup() + self.assertEqual(count, 3) + db.execute.assert_called_once() + sql = db.execute.call_args[0][0] + self.assertIn("PENDING", sql) + self.assertIn("PROCESSING", sql) + + def test_recover_on_startup_skipped_when_disabled(self): + db = _make_db() + cfg = _make_config({"crash_recovery": {"lock_timeout_minutes": 5, "auto_requeue": False}}) + cr = CrashRecovery(db, cfg) + count = cr.recover_on_startup() + self.assertEqual(count, 0) + db.execute.assert_not_called() + + def test_recover_on_startup_handles_db_error(self): + db = _make_db() + db.execute.side_effect = Exception("connection refused") + cfg = _make_config() + cr = CrashRecovery(db, cfg) + count = cr.recover_on_startup() + self.assertEqual(count, 0) + + def test_list_stuck_videos_returns_rows(self): + db = _make_db() + db.fetchall.return_value = [{"id": 5, "file_path": "/data/vid.mp4", "updated_at": None}] + cfg = _make_config() + cr = CrashRecovery(db, cfg) + rows = cr.list_stuck_videos() + self.assertEqual(len(rows), 1) + self.assertEqual(rows[0]["id"], 5) + + +class TestCrashRecoveryCheckpoint(unittest.TestCase): + def test_save_and_load_checkpoint_roundtrip(self): + with tempfile.TemporaryDirectory() as tmpdir: + cfg = _make_config() + cfg.get.side_effect = lambda k, d=None: { + "storage.scratch_path": tmpdir, + }.get(k, d) + cr = CrashRecovery(_make_db(), cfg) + cr.save_checkpoint(42, "extracting", {"frames_done": 5}) + result = cr.load_checkpoint(42) + self.assertIsNotNone(result) + self.assertEqual(result["video_id"], 42) + self.assertEqual(result["state"], "extracting") + self.assertEqual(result["progress"]["frames_done"], 5) + + def test_load_checkpoint_returns_none_when_absent(self): + with tempfile.TemporaryDirectory() as tmpdir: + cfg = _make_config() + cfg.get.side_effect = lambda k, d=None: { + "storage.scratch_path": tmpdir, + }.get(k, d) + cr = CrashRecovery(_make_db(), cfg) + self.assertIsNone(cr.load_checkpoint(999)) + + def test_delete_checkpoint_removes_file(self): + with tempfile.TemporaryDirectory() as tmpdir: + cfg = _make_config() + cfg.get.side_effect = lambda k, d=None: { + "storage.scratch_path": tmpdir, + }.get(k, d) + cr = CrashRecovery(_make_db(), cfg) + cr.save_checkpoint(7, "classifying", {}) + cr.delete_checkpoint(7) + self.assertIsNone(cr.load_checkpoint(7)) + + def test_checkpoint_timestamp_is_iso_format(self): + with tempfile.TemporaryDirectory() as tmpdir: + cfg = _make_config() + cfg.get.side_effect = lambda k, d=None: { + "storage.scratch_path": tmpdir, + }.get(k, d) + cr = CrashRecovery(_make_db(), cfg) + cr.save_checkpoint(1, "detecting", {}) + ckpt = cr.load_checkpoint(1) + # Should parse without error + from datetime import datetime + datetime.fromisoformat(ckpt["timestamp"].replace("Z", "+00:00")) + + +class TestIdempotencyGuard(unittest.TestCase): + def test_returns_true_when_already_completed(self): + db = _make_db() + db.fetchall.return_value = [{"status": "COMPLETED"}] + cr = CrashRecovery(db, _make_config()) + self.assertTrue(cr.is_already_completed(1)) + + def test_returns_false_when_not_completed(self): + db = _make_db() + db.fetchall.return_value = [{"status": "PENDING"}] + cr = CrashRecovery(db, _make_config()) + self.assertFalse(cr.is_already_completed(1)) + + def test_returns_false_when_no_row(self): + db = _make_db() + db.fetchall.return_value = [] + cr = CrashRecovery(db, _make_config()) + self.assertFalse(cr.is_already_completed(99)) + + +# --------------------------------------------------------------------------- +# retry.py tests +# --------------------------------------------------------------------------- + +class TestRetryDecorator(unittest.TestCase): + def test_successful_call_returns_value(self): + @retry(max_attempts=3, initial_delay=0.0, step="test") + def always_succeeds(): + return 42 + + self.assertEqual(always_succeeds(), 42) + + def test_retries_on_transient_error(self): + call_count = {"n": 0} + + @retry(max_attempts=3, initial_delay=0.0, step="test") + def flaky(): + call_count["n"] += 1 + if call_count["n"] < 3: + raise ConnectionError("transient") + return "ok" + + result = flaky() + self.assertEqual(result, "ok") + self.assertEqual(call_count["n"], 3) + + def test_raises_retry_exhausted_after_max_attempts(self): + @retry(max_attempts=3, initial_delay=0.0, step="test") + def always_fails(): + raise ConnectionError("always fails") + + with self.assertRaises(RetryExhaustedError): + always_fails() + + def test_non_retryable_error_propagates_immediately(self): + call_count = {"n": 0} + + @retry(max_attempts=3, initial_delay=0.0, step="test") + def raises_non_retryable(): + call_count["n"] += 1 + raise FileNotFoundError("no such file") + + with self.assertRaises(FileNotFoundError): + raises_non_retryable() + + self.assertEqual(call_count["n"], 1) + + def test_only_specified_exception_types_are_retried(self): + @retry(max_attempts=3, initial_delay=0.0, exceptions=(ValueError,), step="test") + def raises_type_error(): + raise TypeError("wrong type") + + with self.assertRaises(TypeError): + raises_type_error() + + def test_preserves_return_value_on_first_try(self): + @retry(max_attempts=5, initial_delay=0.0, step="test") + def returns_dict(): + return {"key": "value"} + + self.assertEqual(returns_dict(), {"key": "value"}) + + +class TestIsNonRetryable(unittest.TestCase): + def test_file_not_found_is_non_retryable(self): + self.assertTrue(_is_non_retryable(FileNotFoundError("x"))) + + def test_permission_error_is_non_retryable(self): + self.assertTrue(_is_non_retryable(PermissionError("x"))) + + def test_connection_error_is_retryable(self): + self.assertFalse(_is_non_retryable(ConnectionError("x"))) + + def test_runtime_error_is_retryable(self): + self.assertFalse(_is_non_retryable(RuntimeError("x"))) + + +# --------------------------------------------------------------------------- +# drift_detector.py tests +# --------------------------------------------------------------------------- + +class TestDetectDrift(unittest.TestCase): + def test_no_drift_when_distributions_match(self): + base = [0.3] * 50 + [0.7] * 50 # 50% high confidence + curr = [0.3] * 50 + [0.7] * 50 + self.assertFalse(detect_drift(curr, base, threshold=0.10)) + + def test_drift_detected_when_shift_exceeds_threshold(self): + base = [0.3] * 80 + [0.8] * 20 # 20% high + curr = [0.8] * 70 + [0.3] * 30 # 70% high → shift = 0.50 + self.assertTrue(detect_drift(curr, base, threshold=0.10)) + + def test_no_drift_just_below_threshold(self): + base = [0.8] * 50 + [0.2] * 50 # 50% high + curr = [0.8] * 59 + [0.2] * 41 # 59% high → shift = 9% + self.assertFalse(detect_drift(curr, base, threshold=0.10)) + + def test_drift_at_boundary(self): + base = [0.8] * 50 + [0.2] * 50 # 50% + curr = [0.8] * 61 + [0.2] * 39 # 61% → shift = 11% + self.assertTrue(detect_drift(curr, base, threshold=0.10)) + + def test_empty_current_returns_false(self): + self.assertFalse(detect_drift([], [0.5] * 10, threshold=0.10)) + + def test_empty_baseline_returns_false(self): + self.assertFalse(detect_drift([0.5] * 10, [], threshold=0.10)) + + +class TestDriftDetectorAlerts(unittest.TestCase): + def _make_detector(self, db=None, overrides=None): + alerts = [] + cfg = _make_config(overrides or {}) + db = db or _make_db() + detector = DriftDetector(db, cfg, alert_fn=alerts.append) + return detector, alerts + + def test_check_review_queue_growth_triggers_alert(self): + db = _make_db() + db.fetchall.return_value = [{"cnt": 1500}] + detector, alerts = self._make_detector(db) + triggered = detector.check_review_queue_growth() + self.assertTrue(triggered) + self.assertEqual(len(alerts), 1) + self.assertIn("1500", alerts[0]) + + def test_check_review_queue_growth_no_alert_below_threshold(self): + db = _make_db() + db.fetchall.return_value = [{"cnt": 50}] + detector, alerts = self._make_detector(db) + triggered = detector.check_review_queue_growth() + self.assertFalse(triggered) + self.assertEqual(len(alerts), 0) + + def test_check_low_throughput_triggers_alert(self): + db = _make_db() + db.fetchall.return_value = [{"cnt": 5}] # 5 videos in last 1h < 20 min + detector, alerts = self._make_detector(db) + triggered = detector.check_low_throughput() + self.assertTrue(triggered) + self.assertEqual(len(alerts), 1) + + def test_check_low_throughput_no_alert_above_threshold(self): + db = _make_db() + db.fetchall.return_value = [{"cnt": 50}] # 50 > 20 + detector, alerts = self._make_detector(db) + triggered = detector.check_low_throughput() + self.assertFalse(triggered) + + def test_check_error_rate_triggers_alert(self): + db = _make_db() + db.fetchall.return_value = [ + {"status": "COMPLETED", "cnt": 80}, + {"status": "ERROR", "cnt": 10}, + {"status": "UNSCANNABLE", "cnt": 10}, + ] + detector, alerts = self._make_detector(db) + triggered = detector.check_error_rate() + self.assertTrue(triggered) # 20/100 = 20% > 5% + self.assertEqual(len(alerts), 1) + + def test_check_error_rate_no_alert_below_threshold(self): + db = _make_db() + db.fetchall.return_value = [ + {"status": "COMPLETED", "cnt": 98}, + {"status": "ERROR", "cnt": 2}, + ] + detector, alerts = self._make_detector(db) + triggered = detector.check_error_rate() + self.assertFalse(triggered) # 2% < 5% + + def test_check_error_rate_no_alert_zero_videos(self): + db = _make_db() + db.fetchall.return_value = [] + detector, alerts = self._make_detector(db) + triggered = detector.check_error_rate() + self.assertFalse(triggered) + + def test_run_all_checks_returns_dict_with_expected_keys(self): + db = _make_db() + db.fetchall.return_value = [{"cnt": 0}] + detector, _ = self._make_detector(db) + with patch.object(detector, "_fetch_recent_confidences", return_value=[]): + results = detector.run_all_checks() + self.assertIn("confidence_drift", results) + self.assertIn("review_queue_growth", results) + self.assertIn("low_throughput", results) + self.assertIn("high_error_rate", results) + + def test_check_handles_db_error_gracefully(self): + db = _make_db() + db.fetchall.side_effect = Exception("DB offline") + detector, alerts = self._make_detector(db) + # Should not raise + self.assertFalse(detector.check_review_queue_growth()) + self.assertFalse(detector.check_low_throughput()) + self.assertFalse(detector.check_error_rate()) + + +# --------------------------------------------------------------------------- +# health_check.py tests +# --------------------------------------------------------------------------- + +class TestHealthStatus(unittest.TestCase): + def test_snapshot_contains_required_keys(self): + hs = HealthStatus() + snap = hs.snapshot() + for key in ("status", "gpu_available", "gpu_memory_used_gb", + "queue_depth", "uptime_seconds", "videos_processed_today", "last_error"): + self.assertIn(key, snap, f"Missing key: {key}") + + def test_update_changes_values(self): + hs = HealthStatus() + hs.update(status="healthy", queue_depth=55) + snap = hs.snapshot() + self.assertEqual(snap["status"], "healthy") + self.assertEqual(snap["queue_depth"], 55) + + def test_uptime_increases_over_time(self): + hs = HealthStatus() + snap1 = hs.snapshot() + time.sleep(0.05) + snap2 = hs.snapshot() + self.assertGreaterEqual(snap2["uptime_seconds"], snap1["uptime_seconds"]) + + def test_update_is_thread_safe(self): + hs = HealthStatus() + errors = [] + + def writer(n): + try: + for _ in range(100): + hs.update(queue_depth=n) + except Exception as exc: + errors.append(exc) + + threads = [threading.Thread(target=writer, args=(i,)) for i in range(5)] + for t in threads: + t.start() + for t in threads: + t.join() + + self.assertEqual(errors, []) + + def test_http_health_endpoint_returns_200(self): + """Start a real health server and hit /health with urllib.""" + import socket + + # Find a free port + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + port = s.getsockname()[1] + + t = start_health_server(port) + time.sleep(0.1) # give the server a moment to bind + + url = f"http://127.0.0.1:{port}/health" + with urllib.request.urlopen(url, timeout=2) as resp: + self.assertEqual(resp.status, 200) + body = json.loads(resp.read()) + self.assertIn("status", body) + self.assertIn("uptime_seconds", body) + + def test_http_404_for_unknown_path(self): + """Non /health paths return 404.""" + import socket + from urllib.error import HTTPError + + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + port = s.getsockname()[1] + + start_health_server(port) + time.sleep(0.1) + + with self.assertRaises(HTTPError) as ctx: + urllib.request.urlopen(f"http://127.0.0.1:{port}/unknown", timeout=2) + self.assertEqual(ctx.exception.code, 404) + + +if __name__ == "__main__": + unittest.main()