mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
fix: attribute log line to the upstream actually tried
This commit is contained in:
+19
-22
@@ -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
@@ -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))
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user