diff --git a/server/src/agent_control_server/auth_framework/providers/http_upstream.py b/server/src/agent_control_server/auth_framework/providers/http_upstream.py index 464ed2cb..79dcf483 100644 --- a/server/src/agent_control_server/auth_framework/providers/http_upstream.py +++ b/server/src/agent_control_server/auth_framework/providers/http_upstream.py @@ -41,11 +41,12 @@ from __future__ import annotations +import json import ssl from dataclasses import dataclass from datetime import datetime from time import perf_counter -from typing import Any +from typing import Any, Literal, TypedDict import httpx from agent_control_models import JSONObject @@ -84,6 +85,106 @@ ) _JSON_OBJECT_ADAPTER: TypeAdapter[JSONObject] = TypeAdapter(JSONObject) +# Upstream rejection bodies can be arbitrarily long. Bound decoding as well as +# logged output, and report the total when a parsed list is truncated. +_MAX_VALIDATION_RESPONSE_BYTES = 64 * 1024 +_MAX_LOGGED_VALIDATION_ERRORS = 5 +_MAX_LOGGED_LOCATION_PARTS = 4 +_UNKNOWN_VALIDATION_TOTAL: Literal["unknown"] = "unknown" +_SAFE_VALIDATION_LOCATION_PARTS = frozenset( + {"body", "context", "operation", "target_type", "target_id"} +) +_SAFE_VALIDATION_TYPES = frozenset( + {"enum", "literal_error", "missing", "string_type", "uuid_parsing", "uuid_type", "uuid_version"} +) +_MISSING = object() + + +class _ValidationDiagnostic(TypedDict): + type: str + loc: str + + +class _ValidationSummary(TypedDict): + total: int | Literal["unknown"] + errors: list[_ValidationDiagnostic] + status: Literal["parsed", "oversized", "unusable"] + + +def _field_shape(value: Any) -> str: + """Classify a target-context field without revealing its value. + + Target identifiers are caller data, so diagnostics carry only the kind of + value supplied. The length is included because whether an id was + UUID-shaped is the question these diagnostics exist to answer. + """ + if value is _MISSING: + return "missing" + if value is None: + return "null" + if not isinstance(value, str): + return f"non_string:{type(value).__name__}" + if not value: + return "empty" + return f"string:len={len(value)}" + + +def _target_context_shape(context: dict[str, Any] | None) -> dict[str, str]: + """Describe the target context passed to the upstream, values omitted.""" + if not context: + return {"present": "false"} + return { + "present": "true", + "target_type": _field_shape(context.get("target_type", _MISSING)), + "target_id": _field_shape(context.get("target_id", _MISSING)), + } + + +def _sanitized_validation_errors(response: httpx.Response) -> _ValidationSummary: + """Summarize an upstream validation body by field path and error kind. + + Keeps only known validation kinds and request fields. Dynamic location + segments, unknown kinds, ``input``, ``ctx``, and ``msg`` may contain + caller-supplied values and are dropped. Oversized or undecodable bodies + yield an empty summary with an unknown total and a reason rather than + changing the upstream rejection's 502. + """ + if len(response.content) > _MAX_VALIDATION_RESPONSE_BYTES: + return {"total": _UNKNOWN_VALIDATION_TOTAL, "errors": [], "status": "oversized"} + try: + body = response.json() + except (ValueError, RecursionError): + return {"total": _UNKNOWN_VALIDATION_TOTAL, "errors": [], "status": "unusable"} + + detail = body.get("detail") if isinstance(body, dict) else None + if not isinstance(detail, list): + return {"total": _UNKNOWN_VALIDATION_TOTAL, "errors": [], "status": "unusable"} + + errors: list[_ValidationDiagnostic] = [] + for item in detail[:_MAX_LOGGED_VALIDATION_ERRORS]: + if not isinstance(item, dict): + continue + loc = item.get("loc") + error_type = item.get("type") + errors.append( + { + "type": ( + error_type + if isinstance(error_type, str) and error_type in _SAFE_VALIDATION_TYPES + else "other" + ), + "loc": ".".join( + part + if isinstance(part, str) and part in _SAFE_VALIDATION_LOCATION_PARTS + else "" + for part in loc[:_MAX_LOGGED_LOCATION_PARTS] + ) + if isinstance(loc, list) and loc + else "unknown", + } + ) + return {"total": len(detail), "errors": errors, "status": "parsed"} + class _UpstreamGrant(BaseModel): """Strict schema for the upstream authorization-service response. @@ -397,10 +498,31 @@ def _handle_response( hint=hint, ) if 400 <= status < 500: + validation = _sanitized_validation_errors(response) + target_context = _target_context_shape(context) + # Orbit's JSON formatter drops list-valued extras. Index the bounded + # validation entries so its log record retains every type and loc. + upstream_validation = { + str(index): error for index, error in enumerate(validation["errors"]) + } _logger.warning( - "Authorization upstream rejected operation %s with status %d", + "Authorization upstream rejected operation %s with status %d " + "target_context=%s upstream_validation=%s upstream_validation_total=%s " + "upstream_validation_status=%s", operation, status, + json.dumps(target_context, separators=(",", ":")), + json.dumps(upstream_validation, separators=(",", ":")), + validation["total"], + validation["status"], + extra={ + "operation": operation, + "status_code": status, + "target_context": target_context, + "upstream_validation": upstream_validation, + "upstream_validation_total": validation["total"], + "upstream_validation_status": validation["status"], + }, ) raise APIError( status_code=502, diff --git a/server/src/agent_control_server/logging_utils.py b/server/src/agent_control_server/logging_utils.py index d9eeb9ce..093e7c94 100644 --- a/server/src/agent_control_server/logging_utils.py +++ b/server/src/agent_control_server/logging_utils.py @@ -1,3 +1,4 @@ +import json import logging from .config import LoggingSettings @@ -11,6 +12,34 @@ "NOTSET": logging.NOTSET, } _UVICORN_LEVELS = {"CRITICAL", "ERROR", "WARNING", "INFO", "DEBUG", "TRACE"} +_UPSTREAM_DIAGNOSTIC_FIELDS = ( + "operation", + "status_code", + "target_context", + "upstream_validation", + "upstream_validation_total", + "upstream_validation_status", +) + + +class _JsonLogFormatter(logging.Formatter): + """Render log messages and approved upstream diagnostics as valid JSON.""" + + def format(self, record: logging.LogRecord) -> str: + payload: dict[str, object] = { + "time": self.formatTime(record, self.datefmt), + "level": record.levelname, + "name": record.name, + "msg": record.getMessage(), + } + for field in _UPSTREAM_DIAGNOSTIC_FIELDS: + if field in record.__dict__: + payload[field] = record.__dict__[field] + if record.exc_info: + payload["exc_info"] = self.formatException(record.exc_info) + if record.stack_info: + payload["stack_info"] = self.formatStack(record.stack_info) + return json.dumps(payload, default=str) def _normalize_level_name(level: str | None) -> str | None: @@ -88,11 +117,6 @@ def configure_logging( resolved_level = level if level is not None else get_log_level_name(default_level) lvl = _parse_level(resolved_level) as_json = _parse_json(json) - fmt = ( - '{"time":"%(asctime)s","level":"%(levelname)s","name":"%(name)s","msg":"%(message)s"}' - if as_json - else "%(asctime)s %(levelname)s [%(name)s] %(message)s" - ) datefmt = "%Y-%m-%dT%H:%M:%S%z" root = logging.getLogger() @@ -100,7 +124,14 @@ def configure_logging( for h in list(root.handlers): root.removeHandler(h) handler = logging.StreamHandler() - handler.setFormatter(logging.Formatter(fmt=fmt, datefmt=datefmt)) + formatter = ( + _JsonLogFormatter(datefmt=datefmt) + if as_json + else logging.Formatter( + fmt="%(asctime)s %(levelname)s [%(name)s] %(message)s", datefmt=datefmt + ) + ) + handler.setFormatter(formatter) root.addHandler(handler) for name in ("uvicorn", "uvicorn.error", "uvicorn.access"): diff --git a/server/tests/test_auth_framework.py b/server/tests/test_auth_framework.py index 80400c25..9e717b14 100644 --- a/server/tests/test_auth_framework.py +++ b/server/tests/test_auth_framework.py @@ -2,6 +2,9 @@ from __future__ import annotations +import io +import json +import logging from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -36,6 +39,7 @@ ForbiddenError, NotFoundError, ) +from agent_control_server.logging_utils import configure_logging from agent_control_server.models import DEFAULT_NAMESPACE_KEY @@ -375,6 +379,64 @@ async def test_http_upstream_identity_preserves_401(): assert exc_info.value.status_code == 401 +@pytest.mark.asyncio +@pytest.mark.parametrize("status", [400, 422]) +async def test_http_upstream_identity_rejection_reports_sanitized_diagnostics( + caplog: pytest.LogCaptureFixture, + status: int, +) -> None: + # Given: identity resolution rejects a credential, echoing sensitive input. + sentinel = "SENTINEL-IDENTITY-DO-NOT-LOG" + provider = _build_upstream( + lambda request: httpx.Response( + status, + json={ + "detail": [ + { + "type": "missing", + "loc": ["body", sentinel], + "input": sentinel, + "msg": sentinel, + "ctx": {"credential": sentinel}, + } + ] + }, + ), + config_overrides={ + "identity_url": "https://identity.example/resolve", + "service_token": sentinel, + }, + ) + + # When: resolving identity through the same response handler as authorization. + with caplog.at_level(logging.WARNING), pytest.raises(APIError) as exc_info: + await provider.resolve_identity( + _build_request( + headers={ + "X-API-Key": sentinel, + "Authorization": f"Bearer {sentinel}", + "Cookie": f"session={sentinel}", + } + ), + Operation.CONTROL_BINDINGS_READ, + ) + + # Then: the rejection stays a 502 and logs only safe identity diagnostics. + assert exc_info.value.status_code == 502 + assert exc_info.value.error_code == "AUTH_UPSTREAM_REJECTED" + record = _rejection_record(caplog) + assert record.__dict__["operation"] == "identity.resolve" + assert record.__dict__["status_code"] == status + assert record.__dict__["target_context"] == {"present": "false"} + assert record.__dict__["upstream_validation"] == { + "0": {"type": "missing", "loc": "body."} + } + assert record.__dict__["upstream_validation_total"] == 1 + assert record.__dict__["upstream_validation_status"] == "parsed" + assert sentinel not in record.getMessage() + assert sentinel not in str(record.__dict__) + + @pytest.mark.asyncio async def test_http_upstream_identity_treats_missing_route_as_upstream_error(): provider = _build_upstream( @@ -634,6 +696,377 @@ async def test_http_upstream_unexpected_4xx_reports_upstream_rejection(status): assert "request shape" in exc_info.value.hint +def _rejection_record(caplog: pytest.LogCaptureFixture) -> logging.LogRecord: + """Return the upstream-rejection warning emitted by the 4xx branch.""" + records = [ + record + for record in caplog.records + if record.getMessage().startswith("Authorization upstream rejected operation") + ] + assert records + return records[-1] + + +@pytest.mark.asyncio +async def test_http_upstream_4xx_diagnostics_name_the_rejected_field( + caplog: pytest.LogCaptureFixture, +): + """A 422 names the upstream field path and error kind, not the value. + + This is the case the multitenant incident could not diagnose: the log + recorded only the operation and status, so the rejected field was unknown. + """ + provider = _build_upstream( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + { + "type": "uuid_parsing", + "loc": ["body", "context", "target_id"], + "msg": "Input should be a valid UUID", + "input": "not-a-uuid", + } + ] + }, + ) + ) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize( + _build_request(), + Operation.RUNTIME_TOKEN_EXCHANGE, + context={"target_type": "log_stream", "target_id": "not-a-uuid"}, + ) + + record = _rejection_record(caplog) + assert record.__dict__["operation"] == Operation.RUNTIME_TOKEN_EXCHANGE.value + assert record.__dict__["status_code"] == 422 + assert record.__dict__["upstream_validation"] == { + "0": {"type": "uuid_parsing", "loc": "body.context.target_id"} + } + assert record.__dict__["upstream_validation_total"] == 1 + assert record.__dict__["upstream_validation_status"] == "parsed" + assert record.__dict__["target_context"] == { + "present": "true", + "target_type": "string:len=10", + "target_id": "string:len=10", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("context", "expected_shape"), + [ + (None, {"present": "false"}), + ({}, {"present": "false"}), + ( + {"target_type": "log_stream"}, + {"present": "true", "target_type": "string:len=10", "target_id": "missing"}, + ), + ( + {"target_type": None, "target_id": None}, + {"present": "true", "target_type": "null", "target_id": "null"}, + ), + ( + {"target_type": "log_stream", "target_id": ""}, + {"present": "true", "target_type": "string:len=10", "target_id": "empty"}, + ), + ( + {"target_type": "log_stream", "target_id": 42}, + { + "present": "true", + "target_type": "string:len=10", + "target_id": "non_string:int", + }, + ), + ], +) +async def test_http_upstream_4xx_diagnostics_report_target_context_shape( + caplog: pytest.LogCaptureFixture, + context: dict[str, object] | None, + expected_shape: dict[str, str], +): + """Target context is described by shape, so absent and null stay distinct.""" + sent_payloads: list[dict[str, object]] = [] + + def reject(request: httpx.Request) -> httpx.Response: + sent_payloads.append(json.loads(request.content)) + return httpx.Response(422, text="rejected") + + provider = _build_upstream(reject) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize( + _build_request(), + Operation.RUNTIME_TOKEN_EXCHANGE, + context=context, + ) + + assert _rejection_record(caplog).__dict__["target_context"] == expected_shape + assert ("context" in sent_payloads[0]) is bool(context) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("as_json", [False, True], ids=["text", "json"]) +async def test_http_upstream_4xx_diagnostics_reach_configured_log_output( + monkeypatch: pytest.MonkeyPatch, + as_json: bool, +) -> None: + """The configured handler must emit the fields, not only retain them in LogRecord.""" + sentinel = "SENTINEL-DO-NOT-LOG" + provider = _build_upstream( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + { + "type": "uuid_parsing", + "loc": ["body", "context", "target_id"], + "msg": sentinel, + "input": sentinel, + } + ] + }, + ) + ) + monkeypatch.setenv("AGENT_CONTROL_CONFIGURE_LOGGING", "true") + root = logging.getLogger() + original_root = (list(root.handlers), root.level) + uvicorn_loggers = [ + logging.getLogger(name) for name in ("uvicorn", "uvicorn.error", "uvicorn.access") + ] + original_uvicorn = [ + (logger, list(logger.handlers), logger.level, logger.propagate) + for logger in uvicorn_loggers + ] + stream = io.StringIO() + try: + configure_logging(level="WARNING", json=as_json) + root.handlers[0].setStream(stream) + with pytest.raises(APIError): + await provider.authorize( + _build_request(), + Operation.RUNTIME_TOKEN_EXCHANGE, + context={"target_type": "log_stream", "target_id": sentinel}, + ) + finally: + root.handlers, root.level = original_root + for logger, handlers, level, propagate in original_uvicorn: + logger.handlers, logger.level, logger.propagate = handlers, level, propagate + + output = stream.getvalue() + assert sentinel not in output + if as_json: + rendered = json.loads(output) + assert rendered["target_context"]["target_id"] == f"string:len={len(sentinel)}" + assert rendered["upstream_validation"] == { + "0": {"type": "uuid_parsing", "loc": "body.context.target_id"} + } + assert rendered["upstream_validation_total"] == 1 + assert rendered["upstream_validation_status"] == "parsed" + else: + assert 'target_context={"present":"true"' in output + assert 'upstream_validation={"0":{"type":"uuid_parsing"' in output + assert "upstream_validation_total=1" in output + assert "upstream_validation_status=parsed" in output + + +@pytest.mark.asyncio +async def test_http_upstream_4xx_diagnostics_omit_caller_supplied_values( + caplog: pytest.LogCaptureFixture, +): + """Nothing the caller supplied reaches the log, from either side.""" + sentinel = "SENTINEL-DO-NOT-LOG" + provider = _build_upstream( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + { + "type": "enum", + "loc": ["body", "context", "target_type"], + "msg": f"Input should be 'log_stream', got {sentinel}", + "input": sentinel, + "ctx": {"expected": sentinel}, + } + ] + }, + ) + ) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize( + _build_request(), + Operation.RUNTIME_TOKEN_EXCHANGE, + context={"target_type": sentinel, "target_id": sentinel}, + ) + + record = _rejection_record(caplog) + assert sentinel not in record.getMessage() + assert sentinel not in str(record.__dict__) + + +@pytest.mark.asyncio +async def test_http_upstream_4xx_diagnostics_redact_dynamic_location_and_type( + caplog: pytest.LogCaptureFixture, +): + """Validation paths and kinds can also echo caller-controlled data.""" + sentinel = "SENTINEL-DO-NOT-LOG" * 100 + provider = _build_upstream( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + { + "type": sentinel, + "loc": ["body", "context", sentinel, "target_id", sentinel], + }, + {"type": [sentinel], "loc": ["body", "context", "target_id"]}, + ] + }, + ) + ) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize(_build_request(), Operation.RUNTIME_TOKEN_EXCHANGE) + + record = _rejection_record(caplog) + assert record.__dict__["upstream_validation"] == { + "0": {"type": "other", "loc": "body.context..target_id"}, + "1": {"type": "other", "loc": "body.context.target_id"}, + } + assert sentinel not in record.getMessage() + assert sentinel not in str(record.__dict__) + assert len(record.getMessage()) < 500 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("factory", "expected_total", "expected_status"), + [ + pytest.param( + lambda req: httpx.Response(422, text="not json"), + "unknown", + "unusable", + id="non-json", + ), + pytest.param( + lambda req: httpx.Response(422, json={"detail": "a string, not a list"}), + "unknown", + "unusable", + id="detail-not-a-list", + ), + pytest.param( + lambda req: httpx.Response(422, json=["top-level list"]), + "unknown", + "unusable", + id="body-not-an-object", + ), + pytest.param( + lambda req: httpx.Response(422, json={"detail": ["bare string entry"]}), + 1, + "parsed", + id="entry-not-an-object", + ), + pytest.param( + lambda req: httpx.Response( + 422, + content=b'{"detail":' + b"[" * 10_000 + b"0" + b"]" * 10_000 + b"}", + ), + "unknown", + "unusable", + id="deeply-nested-json", + ), + pytest.param( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + { + "type": "missing", + "loc": ["body", "target_id"], + "input": "x" * (64 * 1024), + } + ] + }, + ), + "unknown", + "oversized", + id="oversized-validation-body", + ), + ], +) +async def test_http_upstream_4xx_diagnostics_tolerate_unexpected_bodies( + caplog: pytest.LogCaptureFixture, + factory, + expected_total: int | str, + expected_status: str, +): + """Unexpected rejection bodies still yield 502 with an accurate count.""" + provider = _build_upstream(factory) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError) as exc_info: + await provider.authorize(_build_request(), Operation.CONTROL_BINDINGS_WRITE) + + assert exc_info.value.status_code == 502 + record = _rejection_record(caplog) + assert record.__dict__["upstream_validation"] == {} + assert record.__dict__["upstream_validation_total"] == expected_total + assert f"upstream_validation_total={expected_total}" in record.getMessage() + assert record.__dict__["upstream_validation_status"] == expected_status + assert f"upstream_validation_status={expected_status}" in record.getMessage() + + +@pytest.mark.asyncio +async def test_http_upstream_4xx_diagnostics_distinguish_empty_validation_list( + caplog: pytest.LogCaptureFixture, +): + """An actual empty validation list has a known count of zero.""" + provider = _build_upstream(lambda req: httpx.Response(422, json={"detail": []})) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize(_build_request(), Operation.CONTROL_BINDINGS_WRITE) + + record = _rejection_record(caplog) + assert record.__dict__["upstream_validation"] == {} + assert record.__dict__["upstream_validation_total"] == 0 + assert record.__dict__["upstream_validation_status"] == "parsed" + + +@pytest.mark.asyncio +async def test_http_upstream_4xx_diagnostics_bound_the_logged_error_list( + caplog: pytest.LogCaptureFixture, +): + """A long rejection body is truncated, and the total says so.""" + provider = _build_upstream( + lambda req: httpx.Response( + 422, + json={ + "detail": [ + {"type": "missing", "loc": ["body", f"field_{index}"]} + for index in range(12) + ] + }, + ) + ) + + with caplog.at_level(logging.WARNING): + with pytest.raises(APIError): + await provider.authorize(_build_request(), Operation.CONTROL_BINDINGS_WRITE) + + record = _rejection_record(caplog) + assert len(record.__dict__["upstream_validation"]) == 5 + assert record.__dict__["upstream_validation_total"] == 12 + assert record.__dict__["upstream_validation_status"] == "parsed" + + @pytest.mark.asyncio async def test_http_upstream_surfaces_rate_limit_distinctly(): """Upstream 429 must surface a rate-limit-specific detail and hint.""" diff --git a/server/tests/test_logging_utils.py b/server/tests/test_logging_utils.py index a816c091..92be0038 100644 --- a/server/tests/test_logging_utils.py +++ b/server/tests/test_logging_utils.py @@ -1,5 +1,7 @@ """Tests for logging utilities.""" +import io +import json import logging from agent_control_server.logging_utils import ( @@ -176,3 +178,34 @@ def test_configure_logging_noops_when_host_owns_logging(monkeypatch) -> None: finally: root.handlers = original_handlers root.setLevel(original_level) + + +def test_configure_json_logging_includes_exception_and_stack(monkeypatch) -> None: + """JSON logs retain traceback and stack information when supplied.""" + monkeypatch.setenv("AGENT_CONTROL_CONFIGURE_LOGGING", "true") + root = logging.getLogger() + original_handlers, original_level = list(root.handlers), root.level + uvicorn_loggers = [ + logging.getLogger(name) for name in ("uvicorn", "uvicorn.error", "uvicorn.access") + ] + original_uvicorn = [ + (logger, list(logger.handlers), logger.level, logger.propagate) + for logger in uvicorn_loggers + ] + stream = io.StringIO() + try: + configure_logging(level="ERROR", json=True) + root.handlers[0].setStream(stream) + try: + raise ValueError("test failure") + except ValueError: + logging.getLogger(__name__).error("operation failed", exc_info=True, stack_info=True) + finally: + root.handlers, root.level = original_handlers, original_level + for logger, handlers, level, propagate in original_uvicorn: + logger.handlers, logger.level, logger.propagate = handlers, level, propagate + + record = json.loads(stream.getvalue()) + assert record["msg"] == "operation failed" + assert "ValueError: test failure" in record["exc_info"] + assert "Stack (most recent call last):" in record["stack_info"]