fix: attribute log line to the upstream actually tried

This commit is contained in:
9qeklajc
2026-09-27 23:44:34 +02:00
parent a4328e4162
commit 17f1d875d7
3 changed files with 321 additions and 68 deletions
+19 -22
View File
@@ -94,6 +94,15 @@ def _should_log(method: str, path: str) -> bool:
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
def _attribution(request: Request) -> dict[str, object]:
"""Model/provider fields, omitted rather than null on routes that resolve none."""
return {
field: value
for field in ("model", "provider")
if (value := getattr(request.state, field, None))
}
class LoggingMiddleware(BaseHTTPMiddleware):
"""Middleware to log proxy interactions and page navigation.
@@ -143,14 +152,8 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"path": path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
**_attribution(request),
}
# Omitted rather than null on routes that resolve no model.
model = getattr(request.state, "model", None)
if model:
extra["model"] = model
provider = getattr(request.state, "provider", None)
if provider:
extra["provider"] = provider
if response.status_code >= 400:
error_detail = getattr(request.state, "error_detail", None)
if isinstance(error_detail, dict):
@@ -169,23 +172,17 @@ class LoggingMiddleware(BaseHTTPMiddleware):
except Exception as e:
# Always log failures, even for skipped paths, so we don't lose errors.
duration = time.time() - start_time
failure_extra: dict[str, object] = {
"request_id": request_id,
"method": request.method,
"path": path,
"duration_ms": round(duration * 1000, 2),
"error": str(e),
"error_type": type(e).__name__,
}
model = getattr(request.state, "model", None)
if model:
failure_extra["model"] = model
provider = getattr(request.state, "provider", None)
if provider:
failure_extra["provider"] = provider
logger.error(
"Request failed",
extra=failure_extra,
extra={
"request_id": request_id,
"method": request.method,
"path": path,
"duration_ms": round(duration * 1000, 2),
"error": str(e),
"error_type": type(e).__name__,
**_attribution(request),
},
exc_info=True,
)
raise
+23 -7
View File
@@ -405,6 +405,19 @@ _RETRYABLE_UPSTREAM_5XX = frozenset({502, 503, 504})
_UPSTREAM_5XX_RETRY_BACKOFF_SECONDS = 0.5
def _attribute_request(
request: Request, model_obj: Model, upstream: BaseUpstreamProvider
) -> None:
"""Attribute the completion log line to the candidate being tried.
Uses the provider's model id rather than the requested alias, so aliases
and cross-provider spellings resolve to the model that was forwarded.
"""
if model_obj.id:
request.state.model = model_obj.id
request.state.provider = upstream.provider_type
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
async def proxy(
request: Request, path: str, session: AsyncSession = Depends(get_session)
@@ -461,8 +474,10 @@ async def _proxy(
model_id = request_body_dict.get("model", "unknown")
# Set before routing so the completion log is attributed even when the
# request fails before an upstream is chosen (400/401/402).
request.state.model = model_id
# request fails before an upstream is chosen (400/401/402). "unknown" is
# the no-model sentinel, not a model.
if isinstance(model_id, str) and model_id and model_id != "unknown":
request.state.model = model_id
# Exact Tinfoil attestation GET routes don't map to models — forward
# without model/cost/auth lookups. Do not prefix-match here: paths such as
@@ -643,6 +658,7 @@ async def _proxy(
if x_cashu := headers.get("x-cashu", None):
last_error = None
for i, (model_obj, upstream) in enumerate(candidates):
_attribute_request(request, model_obj, upstream)
try:
if is_ehbp:
if not upstream.supports_ehbp:
@@ -722,7 +738,8 @@ async def _proxy(
logger.debug("Processing unauthenticated GET request", extra={"path": path})
last_error_response = None
for i, (_, upstream) in enumerate(candidates):
for i, (model_obj, upstream) in enumerate(candidates):
_attribute_request(request, model_obj, upstream)
try:
headers = upstream.prepare_headers(dict(request.headers))
response = await upstream.forward_get_request(request, path, headers)
@@ -785,10 +802,6 @@ async def _proxy(
already_stripped: set[str] = set()
for i, (model_obj, upstream) in enumerate(candidates):
# Served model id, not the requested alias, so the completion log
# matches the billing lines for this request.
request.state.model = getattr(model_obj, "id", None) or model_id
request.state.provider = upstream.provider_type
if i > 0 and request_body_dict:
# The reservation was sized to the previous candidate's envelope;
# settlement bills the serving candidate, so a pricier fallback
@@ -820,6 +833,9 @@ async def _proxy(
await _finish_read_transaction(session)
max_cost_for_model = candidate_max
# Only once the candidate is actually tried: a fallback skipped for its
# reservation must not take over the last attempted upstream's line.
_attribute_request(request, model_obj, upstream)
retries_left = settings.upstream_5xx_retry_attempts
retry_index = 0
headers = upstream.prepare_headers(dict(request.headers))
+279 -39
View File
@@ -3,21 +3,29 @@
import json
import logging
from collections.abc import Iterator
from contextlib import contextmanager
from pathlib import Path
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import FastAPI, Request
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import Response
from fastapi.testclient import TestClient
from httpx import ASGITransport, AsyncClient
from pythonjsonlogger import jsonlogger
from routstr import proxy as proxy_module
from routstr.core.db import get_session
from routstr.core.exceptions import UpstreamError
from routstr.core.logging import (
DailyRotatingFileHandler,
RequestIdFilter,
SecurityFilter,
VersionFilter,
)
from routstr.core.middleware import LoggingMiddleware
from routstr.core.middleware import LoggingMiddleware, _attribution
@pytest.fixture
@@ -43,20 +51,41 @@ def handler(tmp_path: Path) -> Iterator[DailyRotatingFileHandler]:
h.close()
def _records(handler: DailyRotatingFileHandler) -> list[dict[str, Any]]:
handler.flush()
text = Path(handler.baseFilename).read_text()
return [json.loads(line) for line in text.strip().splitlines() if line.strip()]
@contextmanager
def _middleware_logs_to(handler: DailyRotatingFileHandler) -> Iterator[None]:
middleware_logger = logging.getLogger("routstr.core.middleware")
saved = (
middleware_logger.handlers,
middleware_logger.level,
middleware_logger.propagate,
)
middleware_logger.handlers = [handler]
middleware_logger.setLevel(logging.INFO)
middleware_logger.propagate = False
try:
yield
finally:
(
middleware_logger.handlers,
middleware_logger.level,
middleware_logger.propagate,
) = saved
def _record(handler: DailyRotatingFileHandler, message: str) -> dict[str, Any]:
matches = [r for r in _records(handler) if r.get("message") == message]
handler.flush()
lines = Path(handler.baseFilename).read_text().strip().splitlines()
matches = [r for r in map(json.loads, lines) if r.get("message") == message]
assert matches, f"no {message!r} record was written"
return matches[-1]
def _build_app() -> FastAPI:
"""Middleware wired like ``main.py``, with routes that set attribution."""
# --------------------------------------------------------------------------- #
# Middleware: fields land on the log lines.
# --------------------------------------------------------------------------- #
def _middleware_app() -> FastAPI:
app = FastAPI()
app.add_middleware(LoggingMiddleware)
@@ -79,49 +108,29 @@ def _build_app() -> FastAPI:
return app
def _with_handler(handler: DailyRotatingFileHandler) -> Any:
middleware_logger = logging.getLogger("routstr.core.middleware")
middleware_logger.setLevel(logging.INFO)
middleware_logger.propagate = False
original_handlers = middleware_logger.handlers
middleware_logger.handlers = [handler]
class _Ctx:
def __enter__(self) -> None:
return None
def __exit__(self, *exc: object) -> None:
middleware_logger.handlers = original_handlers
return _Ctx()
def test_completion_log_carries_model_and_provider(
handler: DailyRotatingFileHandler,
) -> None:
app = _build_app()
with _with_handler(handler):
with TestClient(app, raise_server_exceptions=False) as client:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post(
"/v1/chat/completions", json={"model": "glm-5.3-flash"}
)
assert response.status_code == 200
assert response.status_code == 200
rec = _record(handler, "Request completed")
assert rec["model"] == "z-ai/glm-5.3-flash"
assert rec["provider"] == "openrouter"
assert rec["status_code"] == 200
assert isinstance(rec["duration_ms"], (int, float))
def test_completion_log_omits_attribution_when_route_sets_none(
handler: DailyRotatingFileHandler,
) -> None:
app = _build_app()
with _with_handler(handler):
with TestClient(app, raise_server_exceptions=False) as client:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post("/v1/models", json={})
assert response.status_code == 200
assert response.status_code == 200
rec = _record(handler, "Request completed")
assert "model" not in rec
@@ -131,13 +140,244 @@ def test_completion_log_omits_attribution_when_route_sets_none(
def test_failed_request_log_carries_attribution(
handler: DailyRotatingFileHandler,
) -> None:
app = _build_app()
with _with_handler(handler):
with TestClient(app, raise_server_exceptions=False) as client:
with _middleware_logs_to(handler):
with TestClient(_middleware_app(), raise_server_exceptions=False) as client:
response = client.post("/v1/broken", json={})
assert response.status_code == 500
assert response.status_code == 500
rec = _record(handler, "Request failed")
assert rec["model"] == "deepseek/deepseek-v4.1-flash"
assert rec["provider"] == "venice"
assert rec["error_type"] == "RuntimeError"
def test_middleware_logger_state_is_restored(
handler: DailyRotatingFileHandler,
) -> None:
middleware_logger = logging.getLogger("routstr.core.middleware")
before = (
list(middleware_logger.handlers),
middleware_logger.level,
middleware_logger.propagate,
)
with _middleware_logs_to(handler):
pass
after = (
list(middleware_logger.handlers),
middleware_logger.level,
middleware_logger.propagate,
)
assert after == before
# --------------------------------------------------------------------------- #
# Proxy: which model/provider each routing path attributes the request to.
# --------------------------------------------------------------------------- #
def _model(model_id: str) -> MagicMock:
return MagicMock(id=model_id)
def _upstream(provider_type: str) -> MagicMock:
upstream = MagicMock()
upstream.provider_type = provider_type
upstream.prepare_headers = MagicMock(return_value={})
upstream.on_upstream_error_redirect = AsyncMock()
return upstream
@pytest.fixture
def captured() -> dict[str, object]:
return {}
@pytest.fixture
def proxy_app(captured: dict[str, object]) -> FastAPI:
app = FastAPI()
app.include_router(proxy_module.proxy_router)
app.dependency_overrides[get_session] = lambda: AsyncMock()
@app.middleware("http")
async def capture(request: Request, call_next: Any) -> Response:
try:
return await call_next(request)
finally:
captured.update(_attribution(request))
return app
@pytest.fixture
def routing(monkeypatch: pytest.MonkeyPatch) -> dict[str, Any]:
"""Stub pricing/reservation so only candidate routing drives the test."""
max_costs: dict[str, int] = {}
async def max_cost(
model: str, session: object, model_obj: MagicMock | None = None
) -> int:
return max_costs.get(getattr(model_obj, "id", ""), 100)
async def discounted(cost: int, body: object, model_obj: object = None) -> int:
return cost
state: dict[str, Any] = {
"candidates": [],
"max_costs": max_costs,
"pay": AsyncMock(return_value=MagicMock()),
}
monkeypatch.setattr(
proxy_module, "get_candidates", lambda _model_id: state["candidates"]
)
monkeypatch.setattr(proxy_module, "get_max_cost_for_model", max_cost)
monkeypatch.setattr(proxy_module, "calculate_discounted_max_cost", discounted)
monkeypatch.setattr(proxy_module, "check_token_balance", lambda *_a: None)
monkeypatch.setattr(
proxy_module,
"get_bearer_token_key",
AsyncMock(return_value=MagicMock(hashed_key="abcdef123456", balance=0)),
)
monkeypatch.setattr(proxy_module, "pay_for_request", state["pay"])
monkeypatch.setattr(proxy_module, "revert_pay_for_request", AsyncMock())
monkeypatch.setattr(proxy_module, "_finish_read_transaction", AsyncMock())
return state
async def _send(app: FastAPI, method: str, path: str, **kwargs: Any) -> httpx.Response:
async with AsyncClient(
transport=ASGITransport(app=app), # type: ignore[arg-type]
base_url="http://test",
) as client:
return await client.request(method, path, **kwargs)
@pytest.mark.asyncio
async def test_unauthenticated_request_is_attributed_to_the_requested_model(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
routing["candidates"] = [(_model("prov/model-a"), _upstream("prov"))]
response = await _send(
proxy_app, "POST", "/v1/chat/completions", json={"model": "model-a"}
)
assert response.status_code == 401
assert captured == {"model": "model-a"}
@pytest.mark.parametrize("body", [{}, {"model": "unknown"}, {"model": 123}])
@pytest.mark.asyncio
async def test_request_without_a_model_is_not_attributed(
proxy_app: FastAPI,
routing: dict[str, Any],
captured: dict[str, object],
body: dict[str, object],
) -> None:
response = await _send(proxy_app, "POST", "/v1/chat/completions", json=body)
assert response.status_code == 400
assert captured == {}
@pytest.mark.asyncio
async def test_paid_fallback_is_attributed_to_the_serving_candidate(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_request = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.forward_request = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a-v2"), fallback),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert captured == {"model": "prov-b/model-a-v2", "provider": "prov-b"}
@pytest.mark.asyncio
async def test_fallback_rejected_at_reservation_keeps_last_attempted_attribution(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
"""A pricier fallback the key cannot reserve is never tried, so the line
stays with the upstream that actually handled (and failed) the request."""
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_request = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.forward_request = AsyncMock()
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
routing["max_costs"]["prov-b/model-a"] = 200
routing["pay"].side_effect = [
MagicMock(),
HTTPException(status_code=402, detail="Insufficient balance"),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"authorization": "Bearer sk-test"},
)
assert response.status_code == 402
fallback.forward_request.assert_not_awaited()
assert captured == {"model": "prov-a/model-a", "provider": "prov-a"}
@pytest.mark.asyncio
async def test_x_cashu_fallback_is_attributed_to_the_serving_candidate(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.handle_x_cashu = AsyncMock(
side_effect=UpstreamError("down", status_code=502)
)
fallback.handle_x_cashu = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
response = await _send(
proxy_app,
"POST",
"/v1/chat/completions",
json={"model": "model-a"},
headers={"x-cashu": "cashuAtoken"},
)
assert response.status_code == 200
assert captured == {"model": "prov-b/model-a", "provider": "prov-b"}
@pytest.mark.asyncio
async def test_unauthenticated_get_fallback_is_attributed_to_the_serving_upstream(
proxy_app: FastAPI, routing: dict[str, Any], captured: dict[str, object]
) -> None:
primary, fallback = _upstream("prov-a"), _upstream("prov-b")
primary.forward_get_request = AsyncMock(return_value=Response(status_code=502))
fallback.forward_get_request = AsyncMock(return_value=Response(status_code=200))
routing["candidates"] = [
(_model("prov-a/model-a"), primary),
(_model("prov-b/model-a"), fallback),
]
response = await _send(proxy_app, "GET", "/v1/models")
assert response.status_code == 200
assert captured["provider"] == "prov-b"