Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions sdks/python/src/agent_control/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -448,6 +448,7 @@ def init(
server_url: str | None = None,
api_key: str | None = None,
api_key_header: str | None = None,
runtime_token_header: str | None = None,
controls_file: str | None = None,
steps: list[StepSchemaDict] | None = None,
conflict_mode: Literal["strict", "overwrite"] = "overwrite",
Expand Down Expand Up @@ -488,6 +489,12 @@ def init(
api_key: Optional API key for authentication (defaults to AGENT_CONTROL_API_KEY env var)
api_key_header: Optional HTTP header name for API key authentication
(defaults to AGENT_CONTROL_API_KEY_HEADER env var or X-API-Key)
runtime_token_header: Optional HTTP header the runtime token is sent on
for evaluation requests (defaults to AGENT_CONTROL_RUNTIME_TOKEN_HEADER
env var or Authorization). Point this at a dedicated header (e.g.
X-Agent-Control-Runtime-Token) when the server runs behind a gateway
that reserves Authorization for its own identity JWT. The server must
be configured to read the same header.
controls_file: Optional explicit path to controls.yaml (auto-discovered if not provided)
steps: Optional list of step schemas for registration:
[{"type": "tool", "name": "search", "input_schema": {...}, "output_schema": {...}}]
Expand Down Expand Up @@ -588,6 +595,9 @@ async def handle(message: str):
state.server_url = server_url or os.getenv('AGENT_CONTROL_URL') or 'http://localhost:8000'
state.api_key = api_key
state.api_key_header = resolved_api_key_header
# Stored raw (may be None); AgentControlClient resolves the env-var
# fallback and blank-handling so the rule stays in one place.
state.runtime_token_header = runtime_token_header
state.runtime_token_cache.clear()
state.target_type = target_type
state.target_id = target_id
Expand Down Expand Up @@ -746,6 +756,7 @@ def _reset_state() -> None:
state.server_url = None
state.api_key = None
state.api_key_header = None
state.runtime_token_header = None
state.runtime_token_cache.clear()
state.target_type = None
state.target_id = None
Expand Down
1 change: 1 addition & 0 deletions sdks/python/src/agent_control/_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ def __init__(self) -> None:
self.server_url: str | None = None
self.api_key: str | None = None
self.api_key_header: str | None = None
self.runtime_token_header: str | None = None
self.runtime_token_cache = RuntimeTokenCache()
# Optional target context fixed at init() time; both fields are set
# together or both remain None.
Expand Down
72 changes: 61 additions & 11 deletions sdks/python/src/agent_control/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@
_logger = logging.getLogger(__name__)

_RUNTIME_AUTH_MODE_ENV_VAR = "AGENT_CONTROL_RUNTIME_AUTH_MODE"
_RUNTIME_TOKEN_HEADER_ENV_VAR = "AGENT_CONTROL_RUNTIME_TOKEN_HEADER"
_DEFAULT_RUNTIME_TOKEN_HEADER = "Authorization"
_DEFAULT_RUNTIME_TOKEN_REFRESH_MARGIN_SECONDS = 30
_AUTO_RUNTIME_TOKEN_FALLBACK_STATUSES = {404, 500, 502, 503, 504}
_GLOBAL_RUNTIME_TOKEN_FALLBACK_STATUSES = {404}
Expand All @@ -35,19 +37,34 @@ def _runtime_cache_identity(api_key: str | None, api_key_header: str) -> str:


class _AgentControlAuth(httpx.Auth):
"""Attach local API-key credentials unless a request already has Bearer auth."""
"""Attach local API-key credentials unless the request already carries them.

The API key is suppressed only when the request already presents a bearer
credential on ``Authorization`` or already carries the API key on its own
header. When the runtime token rides a dedicated header, the API key is
left in place: it may be the outer credential the request needs to
authenticate at a gateway before Agent Control verifies the runtime token,
and it rides a different header so there is no collision.
"""

def __init__(self, api_key: str | None, header_name: str = "X-API-Key") -> None:
def __init__(
self,
api_key: str | None,
header_name: str = "X-API-Key",
) -> None:
self._api_key = api_key
self._header_name = header_name

def auth_flow(
self,
request: httpx.Request,
) -> Generator[httpx.Request, httpx.Response, None]:
if self._api_key and "Authorization" not in request.headers:
if self._header_name not in request.headers:
request.headers[self._header_name] = self._api_key
if (
self._api_key
and "Authorization" not in request.headers
and self._header_name not in request.headers
):
request.headers[self._header_name] = self._api_key
yield request


Expand Down Expand Up @@ -102,6 +119,7 @@ def __init__(
runtime_token_cache: RuntimeTokenCache | None = None,
runtime_token_refresh_margin_seconds: int = (_DEFAULT_RUNTIME_TOKEN_REFRESH_MARGIN_SECONDS),
transport: httpx.AsyncBaseTransport | None = None,
runtime_token_header: str | None = None,
):
"""
Initialize the client.
Expand All @@ -125,6 +143,13 @@ def __init__(
runtime_token_refresh_margin_seconds: Refresh cached runtime tokens
before this many seconds of validity remain.
transport: Optional httpx transport, primarily for tests.
runtime_token_header: HTTP header the runtime token is sent on.
Comment thread
josjeon marked this conversation as resolved.
Defaults to ``Authorization``; the
AGENT_CONTROL_RUNTIME_TOKEN_HEADER environment variable
overrides the default. Point this at a dedicated header (e.g.
``X-Agent-Control-Runtime-Token``) when the server runs behind
a gateway that reserves ``Authorization`` for its own identity
JWT. The server must be configured to read the same header.
"""
resolved_base_url = base_url or os.environ.get(
self.BASE_URL_ENV_VAR, "http://localhost:8000"
Expand All @@ -140,6 +165,24 @@ def __init__(
self._runtime_cache_identity = _runtime_cache_identity(self._api_key, self._api_key_header)
configured_runtime_mode = runtime_auth_mode or os.environ.get(_RUNTIME_AUTH_MODE_ENV_VAR)
self._runtime_auth_mode = normalize_runtime_auth_mode(configured_runtime_mode)
# Explicit blank param is a hard error; a blank env var falls back to
# the default (mirrors the server's _resolve_runtime_token_header).
if runtime_token_header is not None and not runtime_token_header.strip():
raise ValueError("runtime_token_header must not be blank.")
env_runtime_token_header = os.environ.get(_RUNTIME_TOKEN_HEADER_ENV_VAR)
if env_runtime_token_header is not None and not env_runtime_token_header.strip():
env_runtime_token_header = None
self._runtime_token_header = (
runtime_token_header
or env_runtime_token_header
or _DEFAULT_RUNTIME_TOKEN_HEADER
).strip()
# Bearer only on Authorization; a dedicated header carries the raw token
# so it can't collide with the gateway's Authorization JWT. Must stay in
# sync with LocalJwtVerifyProvider._require_bearer on the server.
self._runtime_token_use_bearer = (
self._runtime_token_header.lower() == _DEFAULT_RUNTIME_TOKEN_HEADER.lower()
)
if runtime_token_refresh_margin_seconds < 0:
raise ValueError("runtime_token_refresh_margin_seconds must be >= 0.")
self._runtime_token_refresh_margin_seconds = runtime_token_refresh_margin_seconds
Expand Down Expand Up @@ -198,7 +241,10 @@ async def __aenter__(self) -> "AgentControlClient":
base_url=self.base_url,
timeout=self.timeout,
headers=self._get_headers(),
auth=_AgentControlAuth(self._api_key, self._api_key_header),
auth=_AgentControlAuth(
self._api_key,
self._api_key_header,
),
transport=self._transport,
event_hooks={"response": [self._check_server_version]},
)
Expand Down Expand Up @@ -295,18 +341,22 @@ async def post_runtime_evaluation(

return response

def _format_runtime_token(self, token: str) -> str:
"""Sole place the Bearer prefix is applied (see __init__ for the rule)."""
return f"Bearer {token}" if self._runtime_token_use_bearer else token

def _merge_runtime_headers(
self,
headers: dict[str, str] | None,
runtime_authorization: str | None,
) -> dict[str, str] | None:
"""Merge caller headers with an optional Bearer token."""
"""Merge caller headers with an optional runtime token header."""
if headers is None and runtime_authorization is None:
return None

merged = dict(headers or {})
if runtime_authorization is not None:
merged["Authorization"] = runtime_authorization
merged[self._runtime_token_header] = runtime_authorization
return merged

async def _runtime_authorization(
Expand Down Expand Up @@ -350,7 +400,7 @@ async def _runtime_authorization(
refresh_margin_seconds=self._runtime_token_refresh_margin_seconds,
)
if cached is not None:
return f"Bearer {cached.token}"
return self._format_runtime_token(cached.token)

exchange_lock = self._runtime_token_cache.exchange_lock(
self.base_url,
Expand Down Expand Up @@ -379,7 +429,7 @@ async def _runtime_authorization(
refresh_margin_seconds=self._runtime_token_refresh_margin_seconds,
)
if cached is not None:
return f"Bearer {cached.token}"
return self._format_runtime_token(cached.token)

token = await self._exchange_runtime_token(
target_type=target_type,
Expand All @@ -388,7 +438,7 @@ async def _runtime_authorization(
)
if token is None:
return None
return f"Bearer {token}"
return self._format_runtime_token(token)

async def _exchange_runtime_token(
self,
Expand Down
1 change: 1 addition & 0 deletions sdks/python/src/agent_control/control_decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,6 +280,7 @@ async def _evaluate(
base_url=server_url,
api_key=state.api_key,
api_key_header=state.api_key_header,
runtime_token_header=state.runtime_token_header,
runtime_token_cache=state.runtime_token_cache,
) as client:
# If we have controls, use local evaluation which handles both SDK and server controls
Expand Down
1 change: 1 addition & 0 deletions sdks/python/src/agent_control/evaluation.py
Original file line number Diff line number Diff line change
Expand Up @@ -560,6 +560,7 @@ async def evaluate_controls(
base_url=state.server_url,
api_key=state.api_key,
api_key_header=state.api_key_header,
runtime_token_header=state.runtime_token_header,
runtime_token_cache=state.runtime_token_cache,
) as client:
return await check_evaluation_with_local(
Expand Down
171 changes: 171 additions & 0 deletions sdks/python/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -1081,3 +1081,174 @@ async def test_check_server_version_ignores_missing_header() -> None:

# Then: no warning is emitted
mock_warning.assert_not_called()


# ---------------------------------------------------------------------------
# HYBIM-741: configurable runtime-token header (gateway Authorization collision)
# ---------------------------------------------------------------------------


def test_runtime_token_header_defaults_to_authorization() -> None:
client = AgentControlClient(base_url="https://agent-control.test")
assert client._runtime_token_header == "Authorization"
assert client._runtime_token_use_bearer is True


def test_runtime_token_header_param_selects_dedicated_header() -> None:
client = AgentControlClient(
base_url="https://agent-control.test",
runtime_token_header="X-Agent-Control-Runtime-Token",
)
assert client._runtime_token_header == "X-Agent-Control-Runtime-Token"
assert client._runtime_token_use_bearer is False


def test_runtime_token_header_env_var_overrides_default(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv(
"AGENT_CONTROL_RUNTIME_TOKEN_HEADER", "X-Agent-Control-Runtime-Token"
)
client = AgentControlClient(base_url="https://agent-control.test")
assert client._runtime_token_header == "X-Agent-Control-Runtime-Token"
assert client._runtime_token_use_bearer is False


def test_runtime_token_header_rejects_blank_param() -> None:
with pytest.raises(ValueError, match="runtime_token_header"):
AgentControlClient(
base_url="https://agent-control.test", runtime_token_header=" "
)


def test_runtime_token_header_whitespace_env_falls_back(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("AGENT_CONTROL_RUNTIME_TOKEN_HEADER", " ")
client = AgentControlClient(base_url="https://agent-control.test")
assert client._runtime_token_header == "Authorization"
assert client._runtime_token_use_bearer is True


@pytest.mark.asyncio
async def test_runtime_evaluation_sends_raw_token_on_custom_header() -> None:
"""Custom header carries the raw token; Authorization stays free for the
gateway JWT; the API key is preserved on its own header as the outer
credential (it rides X-API-Key, so there is no collision)."""
seen: dict[str, str | None] = {}
expires_at = (datetime.now(UTC) + timedelta(minutes=5)).isoformat()

def handler(request: httpx.Request) -> httpx.Response:
if request.url.path.endswith("/runtime-token-exchange"):
assert request.headers.get("X-API-Key") == "test-key"
return httpx.Response(
200,
json={
"token": "runtime-token",
"expires_at": expires_at,
"target_type": "log_stream",
"target_id": "ls-1",
"scopes": ["runtime.use"],
},
)
seen["runtime"] = request.headers.get("X-Agent-Control-Runtime-Token")
seen["authorization"] = request.headers.get("Authorization")
seen["api_key"] = request.headers.get("X-API-Key")
return httpx.Response(200, json={"is_safe": True, "confidence": 1.0})

transport = httpx.MockTransport(handler)
async with AgentControlClient(
base_url="https://agent-control.test",
api_key="test-key",
runtime_auth_mode="jwt",
runtime_token_header="X-Agent-Control-Runtime-Token",
transport=transport,
) as client:
response = await client.post_runtime_evaluation(
json={"target_type": "log_stream", "target_id": "ls-1"},
target_type="log_stream",
target_id="ls-1",
)

assert response.status_code == 200
assert seen["runtime"] == "runtime-token"
assert seen["authorization"] is None
# The API key is retained on its own header: with the runtime token on a
# dedicated header, X-API-Key can still serve as the outer gateway
# credential without colliding with either the runtime token or the
# gateway's Authorization JWT.
assert seen["api_key"] == "test-key"


@pytest.mark.asyncio
async def test_runtime_evaluation_default_sends_bearer_on_authorization() -> None:
"""Default (unset) behavior is unchanged: Bearer token on Authorization."""
seen: dict[str, str | None] = {}
expires_at = (datetime.now(UTC) + timedelta(minutes=5)).isoformat()

def handler(request: httpx.Request) -> httpx.Response:
if request.url.path.endswith("/runtime-token-exchange"):
return httpx.Response(
200,
json={
"token": "runtime-token",
"expires_at": expires_at,
"target_type": "log_stream",
"target_id": "ls-1",
"scopes": ["runtime.use"],
},
)
seen["authorization"] = request.headers.get("Authorization")
return httpx.Response(200, json={"is_safe": True, "confidence": 1.0})

transport = httpx.MockTransport(handler)
async with AgentControlClient(
base_url="https://agent-control.test",
api_key="test-key",
runtime_auth_mode="jwt",
transport=transport,
) as client:
await client.post_runtime_evaluation(
json={"target_type": "log_stream", "target_id": "ls-1"},
target_type="log_stream",
target_id="ls-1",
)

assert seen["authorization"] == "Bearer runtime-token"


@pytest.mark.asyncio
async def test_runtime_evaluation_custom_header_fallback_keeps_api_key() -> None:
"""Custom-header mode must still fall back to the API key when the exchange
is unavailable: with no runtime token minted, the dedicated header is absent
and X-API-Key must authenticate the evaluation request."""
exchange_calls = 0
seen: dict[str, str | None] = {}

def handler(request: httpx.Request) -> httpx.Response:
nonlocal exchange_calls
if request.url.path.endswith("/runtime-token-exchange"):
exchange_calls += 1
return httpx.Response(503, json={"detail": "runtime auth disabled"})
seen["runtime"] = request.headers.get("X-Agent-Control-Runtime-Token")
seen["api_key"] = request.headers.get("X-API-Key")
return httpx.Response(200, json={"is_safe": True, "confidence": 1.0})

transport = httpx.MockTransport(handler)
async with AgentControlClient(
base_url="https://agent-control.test",
api_key="test-key",
runtime_auth_mode="auto",
runtime_token_header="X-Agent-Control-Runtime-Token",
transport=transport,
) as client:
response = await client.post_runtime_evaluation(
json={"target_type": "log_stream", "target_id": "ls-1"},
target_type="log_stream",
target_id="ls-1",
)
assert response.status_code == 200

assert exchange_calls == 1
assert seen["runtime"] is None
assert seen["api_key"] == "test-key"
Loading