CoolFace
Apppublic

faisaltitu/Drift-Detection

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
middleware.py138 linesDownload Raw Back to api
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