101 lines
3.3 KiB
Python
101 lines
3.3 KiB
Python
"""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,
|
|
)
|