522 lines
19 KiB
Python
522 lines
19 KiB
Python
"""
|
|
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()
|