Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 "<other>"
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.
Expand Down Expand Up @@ -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,
Expand Down
43 changes: 37 additions & 6 deletions server/src/agent_control_server/logging_utils.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import json
import logging

from .config import LoggingSettings
Expand All @@ -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:
Expand Down Expand Up @@ -88,19 +117,21 @@ 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()
root.setLevel(lvl)
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"):
Expand Down
Loading
Loading