This commit is contained in:
Ryan
2026-08-10 19:42:24 -04:00
parent 226da873d0
commit f61a0446e7
7 changed files with 1388 additions and 0 deletions
+521
View File
@@ -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()