"""Request timing context and the counters emitted with completion logs.""" import time import uuid from contextvars import ContextVar from dataclasses import dataclass from typing import cast import structlog from starlette.types import ASGIApp, Receive, Scope, Send from structlog.contextvars import bind_contextvars, reset_contextvars @dataclass class RequestCounters: """Mutable counters scoped to one ASGI request context.""" spotify_calls: int = 0 cache_hits: int = 0 llm_input_tokens: int = 0 llm_output_tokens: int = 0 _COUNTERS: ContextVar[RequestCounters | None] = ContextVar("request_counters", default=None) _REQUEST_ID: ContextVar[str | None] = ContextVar("request_id", default=None) def increment_spotify_calls() -> None: """Count one Spotify HTTP request when a request context is active.""" counters = _COUNTERS.get() if counters is not None: counters.spotify_calls += 1 def increment_cache_hits() -> None: """Count one in-process cache hit when a request context is active.""" counters = _COUNTERS.get() if counters is not None: counters.cache_hits += 1 def record_llm_tokens(input_tokens: int, output_tokens: int) -> None: """Accumulate model token usage when the provider reports it.""" counters = _COUNTERS.get() if counters is not None: counters.llm_input_tokens += input_tokens counters.llm_output_tokens += output_tokens def current_request_id() -> str | None: """Return the active request identifier when one exists.""" return _REQUEST_ID.get() class RequestTimingMiddleware: """Log duration and counters after the complete response body is sent.""" def __init__(self, app: ASGIApp) -> None: """Wrap an ASGI application.""" self.app = app async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: """Install request context and measure an HTTP exchange.""" if scope["type"] != "http": await self.app(scope, receive, send) return request_id = uuid.uuid4().hex state = scope.setdefault("state", {}) cast(dict[str, object], state)["request_id"] = request_id started_at = time.monotonic() counters = RequestCounters() counter_token = _COUNTERS.set(counters) request_token = _REQUEST_ID.set(request_id) logging_tokens = bind_contextvars(request_id=request_id) try: await self.app(scope, receive, send) finally: self._log_completion(scope, request_id, started_at, counters) reset_contextvars(**logging_tokens) _COUNTERS.reset(counter_token) _REQUEST_ID.reset(request_token) @staticmethod def _log_completion( scope: Scope, request_id: str, started_at: float, counters: RequestCounters, ) -> None: structlog.get_logger().info( "request_complete", request_id=request_id, method=scope.get("method"), path=scope.get("path"), duration_ms=round((time.monotonic() - started_at) * 1000), spotify_calls=counters.spotify_calls, cache_hits=counters.cache_hits, llm_input_tokens=counters.llm_input_tokens, llm_output_tokens=counters.llm_output_tokens, )