Story 9
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user