Merge pull request #675 from Routstr/feat/client-app-identification

feat: identify client app in request and error logs
This commit is contained in:
9qeklajc
2026-09-08 22:40:10 +02:00
committed by GitHub
3 changed files with 268 additions and 5 deletions
+22 -4
View File
@@ -184,6 +184,18 @@ class RequestIdFilter(logging.Filter):
return True
class ClientAppFilter(logging.Filter):
"""Attach request-local app attribution to log records."""
def filter(self, record: logging.LogRecord) -> bool:
# Import here to avoid circular imports
from .middleware import UNKNOWN_CLIENT_APP, client_app_context
client_app = client_app_context.get(None)
record.client_app = client_app if client_app else UNKNOWN_CLIENT_APP
return True
# Standard ``LogRecord`` attributes that are never user-supplied ``extra``
# fields; skipped when redacting structured extras (``msg``/``message`` are
# handled separately above).
@@ -323,7 +335,7 @@ def setup_logging() -> None:
"rich_tracebacks": True,
"markup": True,
"console": _console,
"filters": ["request_id_filter", "security_filter"],
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
}
else:
console_handler = {
@@ -331,7 +343,7 @@ def setup_logging() -> None:
"level": log_level,
"formatter": "plain",
"stream": "ext://sys.stdout",
"filters": ["request_id_filter", "security_filter"],
"filters": ["request_id_filter", "client_app_filter", "security_filter"],
}
LOGGING_CONFIG = {
@@ -340,7 +352,7 @@ def setup_logging() -> None:
"formatters": {
"json": {
"()": jsonlogger.JsonFormatter,
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s",
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s %(client_app)s",
"datefmt": "%Y-%m-%d %H:%M:%S",
},
"plain": {
@@ -351,6 +363,7 @@ def setup_logging() -> None:
"filters": {
"version_filter": {"()": VersionFilter},
"request_id_filter": {"()": RequestIdFilter},
"client_app_filter": {"()": ClientAppFilter},
"security_filter": {"()": SecurityFilter},
},
"handlers": {
@@ -364,7 +377,12 @@ def setup_logging() -> None:
"interval": 1, # Every 1 day
"backupCount": 30, # Keep 30 days of logs
"atTime": None, # Rotate at midnight (00:00)
"filters": ["version_filter", "request_id_filter", "security_filter"],
"filters": [
"version_filter",
"request_id_filter",
"client_app_filter",
"security_filter",
],
},
},
"loggers": {
+48 -1
View File
@@ -2,8 +2,10 @@ import time
import uuid
from contextvars import ContextVar
from typing import Callable
from urllib.parse import urlsplit
from fastapi import Request, Response
from starlette.datastructures import Headers
from starlette.middleware.base import BaseHTTPMiddleware
from .logging import get_logger
@@ -13,6 +15,41 @@ logger = get_logger(__name__)
# Context variable to store request ID across async context
request_id_context: ContextVar[str | None] = ContextVar("request_id")
client_app_context: ContextVar[str | None] = ContextVar("client_app")
UNKNOWN_CLIENT_APP = "unknown"
# Prefer OpenRouter app headers, then browser and SDK fallbacks.
_CLIENT_APP_HEADERS: tuple[str, ...] = (
"x-title",
"http-referer",
"referer",
"user-agent",
)
# Limit untrusted header data repeated in every log record.
_CLIENT_APP_MAX_LENGTH = 120
def client_app_from_headers(headers: Headers) -> str:
for header in _CLIENT_APP_HEADERS:
raw = headers.get(header)
if raw is None:
continue
cleaned = "".join(ch for ch in raw if ch.isprintable()).strip()
if header in ("http-referer", "referer"):
try:
url = urlsplit(cleaned)
if url.scheme not in ("http", "https") or not url.hostname:
continue
except ValueError:
continue
# Attribution needs the origin, not credentials or private page URLs.
cleaned = f"{url.scheme}://{url.netloc.rsplit('@', 1)[-1]}"
if cleaned:
return cleaned[:_CLIENT_APP_MAX_LENGTH]
return UNKNOWN_CLIENT_APP
# Methods that are never logged: HEAD requests are health probes from
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
@@ -71,6 +108,10 @@ class LoggingMiddleware(BaseHTTPMiddleware):
# Set request ID in context for logging
token = request_id_context.set(request_id)
client_app_token = client_app_context.set(
client_app_from_headers(request.headers)
)
path = request.url.path
should_log = _should_log(request.method, path)
@@ -130,6 +171,12 @@ class LoggingMiddleware(BaseHTTPMiddleware):
finally:
# Reset context
request_id_context.reset(token)
client_app_context.reset(client_app_token)
__all__ = ["LoggingMiddleware", "request_id_context"]
__all__ = [
"LoggingMiddleware",
"UNKNOWN_CLIENT_APP",
"client_app_context",
"request_id_context",
]
+198
View File
@@ -0,0 +1,198 @@
"""Tests for client-app identification in request logging."""
import asyncio
import logging
import pytest
from fastapi import FastAPI, Request, Response
from fastapi.testclient import TestClient
from starlette.datastructures import Headers
from routstr.core.logging import ClientAppFilter
from routstr.core.middleware import (
UNKNOWN_CLIENT_APP,
LoggingMiddleware,
client_app_context,
client_app_from_headers,
)
def _record() -> logging.LogRecord:
return logging.LogRecord(
name="routstr.test",
level=logging.INFO,
pathname=__file__,
lineno=1,
msg="test",
args=None,
exc_info=None,
)
@pytest.mark.parametrize(
("headers", "expected"),
[
(
{
"x-title": "Goose",
"http-referer": "https://myapp.example.com",
"user-agent": "python-httpx/0.27",
},
"Goose",
),
(
{"http-referer": "https://myapp.example.com", "user-agent": "curl/8.4.0"},
"https://myapp.example.com",
),
(
{"referer": "https://myapp.example.com", "user-agent": "curl/8.4.0"},
"https://myapp.example.com",
),
({"user-agent": "curl/8.4.0"}, "curl/8.4.0"),
({}, UNKNOWN_CLIENT_APP),
({"x-title": " ", "user-agent": "curl/8.4.0"}, "curl/8.4.0"),
({"x-title": " ", "user-agent": "\t"}, UNKNOWN_CLIENT_APP),
],
ids=[
"x-title-wins",
"http-referer",
"referer",
"user-agent-fallback",
"no-identity-headers",
"blank-falls-through",
"all-blank",
],
)
def test_client_app_from_headers(headers: dict[str, str], expected: str) -> None:
assert client_app_from_headers(Headers(headers)) == expected
@pytest.mark.parametrize("header", ["http-referer", "referer"])
@pytest.mark.parametrize(
("url", "expected"),
[
(
"https://alice:password@app.example:8443/private/chat?token=secret#access_token=secret",
"https://app.example:8443",
),
("http://[::1]:3000/chat?key=secret", "http://[::1]:3000"),
("https://app.example/" + "a" * 200, "https://app.example"),
("https://[invalid", "curl/8.4.0"),
("/private/chat?token=secret", "curl/8.4.0"),
("javascript:secret", "curl/8.4.0"),
("https:///private", "curl/8.4.0"),
],
)
def test_referrer_only_identifies_origin(header: str, url: str, expected: str) -> None:
headers = Headers({header: url, "user-agent": "curl/8.4.0"})
assert client_app_from_headers(headers) == expected
def test_value_is_truncated_to_120_chars() -> None:
assert client_app_from_headers(Headers({"x-title": "a" * 500})) == "a" * 120
def test_control_characters_are_stripped() -> None:
headers = Headers({"user-agent": "evil-app\x1b[0m fake INFO line"})
assert client_app_from_headers(headers) == "evil-app[0m fake INFO line"
def test_filter_reads_context_variable() -> None:
token = client_app_context.set("Goose")
try:
record = _record()
assert ClientAppFilter().filter(record) is True
assert record.client_app == "Goose" # type: ignore[attr-defined]
finally:
client_app_context.reset(token)
@pytest.mark.parametrize("fail", [False, True])
async def test_context_is_restored_after_request(fail: bool) -> None:
middleware = LoggingMiddleware(FastAPI())
request = Request(
{
"type": "http",
"method": "GET",
"path": "/test",
"query_string": b"",
"headers": [],
}
)
async def call_next(request: Request) -> Response:
assert client_app_context.get() == UNKNOWN_CLIENT_APP
if fail:
raise RuntimeError("handler failed")
return Response()
token = client_app_context.set("outer")
try:
if fail:
with pytest.raises(RuntimeError, match="handler failed"):
await middleware.dispatch(request, call_next)
else:
await middleware.dispatch(request, call_next)
assert client_app_context.get() == "outer"
finally:
client_app_context.reset(token)
async def test_concurrent_requests_keep_their_own_client_app() -> None:
middleware = LoggingMiddleware(FastAPI())
ready = asyncio.Event()
apps: list[str] = []
async def call_next(request: Request) -> Response:
apps.append(request.headers["x-title"])
if len(apps) == 2:
ready.set()
await asyncio.wait_for(ready.wait(), timeout=5)
assert client_app_context.get() == request.headers["x-title"]
return Response()
await asyncio.gather(
*(
middleware.dispatch(
Request(
{
"type": "http",
"method": "GET",
"path": "/test",
"query_string": b"",
"headers": [(b"x-title", app)],
}
),
call_next,
)
for app in (b"Goose", b"Pi")
)
)
def test_filter_defaults_to_unknown_outside_request_context() -> None:
record = _record()
assert ClientAppFilter().filter(record) is True
assert record.client_app == UNKNOWN_CLIENT_APP # type: ignore[attr-defined]
def test_handler_logs_carry_client_app(caplog: pytest.LogCaptureFixture) -> None:
app = FastAPI()
handler_logger = logging.getLogger("routstr.test.handler")
@app.get("/whoami")
async def whoami() -> dict[str, bool]:
handler_logger.warning("something went wrong")
return {"ok": True}
app.add_middleware(LoggingMiddleware)
caplog.handler.addFilter(ClientAppFilter())
handler_logger.addHandler(caplog.handler)
try:
TestClient(app).get("/whoami", headers={"X-Title": "Goose"})
finally:
handler_logger.removeHandler(caplog.handler)
record = next(r for r in caplog.records if r.name == "routstr.test.handler")
assert record.client_app == "Goose" # type: ignore[attr-defined]