From 17f1d875d71fdceff4b073fb18db3595e5cc649f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sun, 27 Sep 2026 23:44:34 +0200 Subject: [PATCH] fix: attribute log line to the upstream actually tried --- routstr/core/middleware.py | 41 ++- routstr/proxy.py | 30 +- .../test_log_model_provider_attribution.py | 318 +++++++++++++++--- 3 files changed, 321 insertions(+), 68 deletions(-) diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 517f4d2c..65708d54 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -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 diff --git a/routstr/proxy.py b/routstr/proxy.py index 9c9e1f6e..ae36d1d4 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -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)) diff --git a/tests/unit/test_log_model_provider_attribution.py b/tests/unit/test_log_model_provider_attribution.py index 916f0599..76529af1 100644 --- a/tests/unit/test_log_model_provider_attribution.py +++ b/tests/unit/test_log_model_provider_attribution.py @@ -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"