diff --git a/routstr/core/logging.py b/routstr/core/logging.py index 61c96607..0fba407b 100644 --- a/routstr/core/logging.py +++ b/routstr/core/logging.py @@ -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": { diff --git a/routstr/core/middleware.py b/routstr/core/middleware.py index 63d43888..d4ddfb8f 100644 --- a/routstr/core/middleware.py +++ b/routstr/core/middleware.py @@ -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", +] diff --git a/tests/unit/test_client_app_logging.py b/tests/unit/test_client_app_logging.py new file mode 100644 index 00000000..107f8c2c --- /dev/null +++ b/tests/unit/test_client_app_logging.py @@ -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]