faisaltitu/Drift-Detection
0
1"""2Observability Middleware3- Request latency tracking (p50, p95, p99)4- Per-endpoint counters5- SLO breach alerting6"""7 8import logging9import time10from collections import defaultdict11from dataclasses import dataclass, field12from typing import Dict, List13 14import numpy as np15from starlette.middleware.base import BaseHTTPMiddleware16from starlette.requests import Request17from starlette.responses import Response18 19logger = logging.getLogger(__name__)20 21# ── SLO Targets ──────────────────────────────────────────────────────22PREDICT_P95_SLO_MS = 150 # /predict p95 must stay below 150 ms23HEALTH_P95_SLO_MS = 50 # /health p95 must stay below 50 ms24DEFAULT_P95_SLO_MS = 500 # catch-all25 26SLO_MAP = {27 "/predict": PREDICT_P95_SLO_MS,28 "/health": HEALTH_P95_SLO_MS,29}30 31 32@dataclass33class EndpointStats:34 """Accumulates latencies and counts for one endpoint."""35 latencies_ms: List[float] = field(default_factory=list)36 total_requests: int = 037 error_count: int = 038 slo_breaches: int = 039 _max_samples: int = 10_000 # rolling window40 41 def record(self, latency_ms: float, status_code: int, slo_ms: float) -> None:42 self.total_requests += 143 self.latencies_ms.append(latency_ms)44 if status_code >= 500:45 self.error_count += 146 if latency_ms > slo_ms:47 self.slo_breaches += 148 # Trim to rolling window49 if len(self.latencies_ms) > self._max_samples:50 self.latencies_ms = self.latencies_ms[-self._max_samples:]51 52 def percentile(self, pct: float) -> float:53 if not self.latencies_ms:54 return 0.055 return float(np.percentile(self.latencies_ms, pct))56 57 def to_dict(self, slo_ms: float) -> Dict:58 return {59 "total_requests": self.total_requests,60 "error_count": self.error_count,61 "error_rate": round(self.error_count / max(self.total_requests, 1), 4),62 "slo_target_ms": slo_ms,63 "slo_breaches": self.slo_breaches,64 "slo_breach_rate": round(self.slo_breaches / max(self.total_requests, 1), 4),65 "p50_ms": round(self.percentile(50), 2),66 "p95_ms": round(self.percentile(95), 2),67 "p99_ms": round(self.percentile(99), 2),68 "avg_ms": round(float(np.mean(self.latencies_ms)) if self.latencies_ms else 0, 2),69 "max_ms": round(max(self.latencies_ms) if self.latencies_ms else 0, 2),70 }71 72 73class MetricsCollector:74 """Singleton that aggregates per-endpoint metrics."""75 76 def __init__(self) -> None:77 self._stats: Dict[str, EndpointStats] = defaultdict(EndpointStats)78 79 def record(self, path: str, latency_ms: float, status_code: int) -> None:80 slo_ms = SLO_MAP.get(path, DEFAULT_P95_SLO_MS)81 self._stats[path].record(latency_ms, status_code, slo_ms)82 83 def snapshot(self) -> Dict:84 """Return metrics summary for every tracked endpoint."""85 result: Dict = {}86 total_reqs = 087 total_errs = 088 all_latencies: List[float] = []89 for path, stats in sorted(self._stats.items()):90 slo_ms = SLO_MAP.get(path, DEFAULT_P95_SLO_MS)91 result[path] = stats.to_dict(slo_ms)92 total_reqs += stats.total_requests93 total_errs += stats.error_count94 all_latencies.extend(stats.latencies_ms)95 96 # Global summary97 result["_global"] = {98 "total_requests": total_reqs,99 "error_count": total_errs,100 "error_rate": round(total_errs / max(total_reqs, 1), 4),101 "p50_ms": round(float(np.percentile(all_latencies, 50)) if all_latencies else 0, 2),102 "p95_ms": round(float(np.percentile(all_latencies, 95)) if all_latencies else 0, 2),103 "p99_ms": round(float(np.percentile(all_latencies, 99)) if all_latencies else 0, 2),104 }105 return result106 107 def reset(self) -> None:108 self._stats.clear()109 110 111# Module-level singleton112metrics_collector = MetricsCollector()113 114 115class LatencyMiddleware(BaseHTTPMiddleware):116 """Measures request latency and feeds MetricsCollector."""117 118 async def dispatch(self, request: Request, call_next) -> Response:119 start = time.perf_counter()120 response: Response = await call_next(request)121 elapsed_ms = (time.perf_counter() - start) * 1000.0122 123 path = request.url.path124 metrics_collector.record(path, elapsed_ms, response.status_code)125 126 # Attach latency header for client visibility127 response.headers["X-Response-Time-Ms"] = f"{elapsed_ms:.2f}"128 129 # Log slow requests130 slo = SLO_MAP.get(path, DEFAULT_P95_SLO_MS)131 if elapsed_ms > slo:132 logger.warning(133 "SLO breach: %s %s took %.1f ms (target: %d ms)",134 request.method, path, elapsed_ms, slo,135 )136 137 return response138 