fix(agents): enforce target policy and preserve trusted authentication

This commit is contained in:
Joshua Valluru 2026-09-30 11:46:27 -07:00
parent 30c70adb64
commit 0fd4e93c48
12 changed files with 267 additions and 34 deletions

View file

@ -3320,6 +3320,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
# single-owner so its meaning stays trustworthy.
mcp_session_resource_server_id: str | None = Field(default=None, exclude=True)
mcp_toolset_id: str | None = Field(default=None, exclude=True)
authenticated_by_custom_auth: bool = Field(default=False, exclude=True)
via_virtual_key: bool = Field(
default=False,
exclude=True,
@ -3381,6 +3382,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
values.pop("mcp_session_resource_server_id", None)
values.pop("mcp_toolset_id", None)
values.pop("via_virtual_key", None)
values.pop("authenticated_by_custom_auth", None)
values.pop("agent_caller", None)
values.pop("managed_agent_context", None)
values.pop("managed_agent_policy", None)

View file

@ -762,7 +762,9 @@ async def invoke_agent_a2a(
body["metadata"] = {}
body["metadata"]["agent_id"] = agent.agent_id
body["metadata"]["model_group"] = f"a2a_agent/{agent_name}"
body["metadata"]["model_info"] = {"id": agent.agent_id}
body["metadata"]["model_info"] = { # mutable-ok: request hooks mutate metadata before JSON logging
"id": agent.agent_id
}
body["agent_id"] = agent.agent_id
body.update(

View file

@ -200,6 +200,7 @@ class AgentRequestHandler:
if key_hash
and managed_agent_policy(user_api_key_auth) is None
and not user_api_key_auth.is_session_token
and not user_api_key_auth.authenticated_by_custom_auth
else user_api_key_auth
)
fresh_auth: Final = authority.model_copy(

View file

@ -1,7 +1,7 @@
from collections.abc import Mapping
from itertools import product
from types import MappingProxyType
from typing import Annotated, Final
from typing import Annotated, Final, Literal
from pydantic import Field, TypeAdapter, ValidationError
@ -62,6 +62,22 @@ _MANAGED_MCP_ROUTES: Final = tuple(
)
_MODEL_ROUTE_KINDS: Final[
Mapping[str, Literal["image_generation", "image_edit", "moderation", "speech", "body", "path"]]
] = MappingProxyType(
{
"/images/generations": "image_generation",
"/images/edits": "image_edit",
"/moderations": "moderation",
"/audio/transcriptions": "moderation",
"/audio/speech": "speech",
"/rerank": "body",
"/messages/count_tokens": "body",
":countTokens": "path",
}
)
def managed_agent_route_allowed(route: str, method: str | None) -> bool:
from litellm.proxy.auth.route_checks import RouteChecks
@ -92,26 +108,12 @@ def managed_inference_request(
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
return {**body, "model": model}
return {**body, "model": model} # mutable-ok: centralized auth hooks add request tags and budget metadata
if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
return dict(body)
return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
kind: Final = (
"image_generation"
if route.endswith("/images/generations")
else "image_edit"
if route.endswith("/images/edits")
else "moderation"
if route.endswith(("/moderations", "/audio/transcriptions"))
else "speech"
if route.endswith("/audio/speech")
else "body"
if route.endswith(("/rerank", "/messages/count_tokens"))
else "path"
if route.endswith(":countTokens")
else "completion"
)
kind: Final = next((kind for suffix, kind in _MODEL_ROUTE_KINDS.items() if route.endswith(suffix)), "completion")
endpoint_model: Final = path_model or (
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
)
@ -120,7 +122,7 @@ def managed_inference_request(
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
return {**body, "model": effective}
return {**body, "model": effective} # mutable-ok: centralized auth hooks add request tags and budget metadata
def managed_agent_policy(auth: "UserAPIKeyAuth | None") -> AgentResponse | None:
@ -204,12 +206,12 @@ _INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
model: Final = body.get("model")
if isinstance(model, str) and model.startswith("a2a/"):
return model.removeprefix("a2a/") or None
components: Final = tuple(route.strip("/").split("/"))
path: Final = components[1:] if components and components[0] == "v1" else components
return path[1] if len(path) >= 2 and path[0] == "a2a" else None
if len(path) >= 2 and path[0] == "a2a":
return path[1] or None
model: Final = body.get("model")
return model.removeprefix("a2a/") or None if isinstance(model, str) and model.startswith("a2a/") else None
async def prepare_agent_invocation(

View file

@ -2965,7 +2965,7 @@ class JWTAuthManager:
),
)
auth.managed_agent_context = result.get("managed_agent_context")
auth._managed_delegation_verified = (
auth._managed_delegation_verified = ( # pyright: ignore[reportPrivateUsage] # JWT admission produces the one-shot proof consumed by managed authorization
auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated"
)
return auth

View file

@ -1561,6 +1561,7 @@ async def _user_api_key_auth_builder(
route=route,
parent_otel_span=parent_otel_span,
)
validated.authenticated_by_custom_auth = True
return validated
elif response is not None and isinstance(response, str):
api_key = response
@ -1576,6 +1577,7 @@ async def _user_api_key_auth_builder(
route=route,
parent_otel_span=parent_otel_span,
)
validated.authenticated_by_custom_auth = True
return validated
### LITELLM-DEFINED AUTH FUNCTION ###

View file

@ -2313,6 +2313,11 @@ async def _complete_cli_sso_callback_session(
status_code=500,
detail="Could not resolve team model grants for this login. Please try again",
)
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject
await enroll_microsoft_subject(
request.scope.get("litellm_microsoft_interactive_subject"), user_info.user_id, prisma_client
)
resolved_teams: Final = _cli_sso_session_teams(team_details)
attribution_metadata: Final = build_cli_sso_attribution_metadata(result=result)
if attribution_metadata:

View file

@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
actor_admission_failure,
admit_managed_actor,
@ -78,6 +78,7 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
{
"managed_agent_context": context,
"requires_fresh_policy": True,
"authenticated_by_custom_auth": True,
"mcp_explicit_grants_only": True,
"managed_agent_policy": agent(),
"billing_agent_policy": agent(),
@ -86,6 +87,8 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
}
)
assert auth.requires_fresh_policy is False
assert auth.authenticated_by_custom_auth is False
assert "authenticated_by_custom_auth" not in auth.model_dump()
assert auth.mcp_explicit_grants_only is False
assert "mcp_explicit_grants_only" not in auth.model_dump()
assert auth.managed_agent_context is None
@ -161,6 +164,9 @@ async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
"route,body,expected",
[
("/a2a/agent", {}, "agent"),
("/a2a/expensive", {"model": "a2a/cheap"}, "expensive"),
("/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
("/v1/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
("/v1/a2a/agent/", {}, "agent"),
("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"),
("/v1/chat/completions", {"model": "a2a/"}, None),
@ -496,6 +502,8 @@ async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_
auth: Final = UserAPIKeyAuth(agent_id="agent")
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
assert auth.requires_fresh_policy is False
assert auth.authenticated_by_custom_auth is False
assert "authenticated_by_custom_auth" not in auth.model_dump()
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert auth.requires_fresh_policy is True
@ -545,3 +553,27 @@ async def test_ordinary_agent_admission_preserves_legacy_authentication(
assert auth.agent_id == "agent"
assert auth.managed_agent_policy is None
assert auth.requires_fresh_policy is False
assert auth.authenticated_by_custom_auth is False
assert "authenticated_by_custom_auth" not in auth.model_dump()
@pytest.mark.parametrize(
"route",
tuple(dict.fromkeys(
LiteLLMRoutes.openai_routes.value
+ LiteLLMRoutes.anthropic_routes.value
+ LiteLLMRoutes.google_routes.value
)),
)
def test_registered_inference_routes_have_an_explicit_managed_access_decision(route: str) -> None:
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
normalized: Final = route.removeprefix("/openai").removeprefix("/v1beta").removeprefix("/v1")
unsupported: Final = normalized.startswith((
"/videos", "/batches", "/files", "/fine_tuning", "/assistants", "/threads", "/utils/",
"/vector_stores", "/vector_store/", "/search", "/containers", "/skills", "/claude-code/",
"/interactions", "/agents", "/responses/{", "/responses/input_tokens",
"/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions",
)) or normalized in ("/models", "/cursor/models", "/cursor/v1/models")
concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model")
assert managed_agent_route_allowed(concrete, None) is not unsupported, route

View file

@ -8055,7 +8055,10 @@ def test_managed_issuer_requires_configured_audience_validation(
@pytest.mark.asyncio
async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(monkeypatch: pytest.MonkeyPatch) -> None:
@pytest.mark.parametrize("authentication_write", ["success", "revoked", "unavailable"])
async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(
monkeypatch: pytest.MonkeyPatch, authentication_write: str
) -> None:
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.types.proxy.agent_identity import AgentIdentityBinding
@ -8093,6 +8096,16 @@ async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(mo
database.writer_db.litellm_agentidentity.find_unique.assert_awaited_once()
assert database.writer_db.litellm_agentstable.find_unique.await_count == 2
assert database.writer_db.litellm_agentidentity.update_many.await_count == 2
if authentication_write != "success":
database.writer_db.litellm_agentidentity.update_many.return_value = 0
database.writer_db.litellm_agentidentity.update_many.side_effect = (
RuntimeError("storage unavailable") if authentication_write == "unavailable" else None
)
with pytest.raises(HTTPException) as failed_write:
await JWTAuthManager.authorize_jwt(**arguments)
assert failed_write.value.status_code == (503 if authentication_write == "unavailable" else 403)
assert database.writer_db.litellm_agentidentity.update_many.await_count == 3
return
database.writer_db.litellm_agentstable.find_unique.return_value = agent.model_copy(update={"enabled": False})
with pytest.raises(HTTPException) as denied:
await JWTAuthManager.authorize_jwt(**arguments)

View file

@ -9453,16 +9453,18 @@ async def test_managed_actor_cannot_access_provider_resource_routes(monkeypatch,
request.scope["method"] = "GET"
from litellm.types.proxy.agent_identity import ManagedAgentContext
auth = UserAPIKeyAuth(
agent_id="managed", api_key="persisted-key", models=["test-model"],
managed_agent_context=(
ManagedAgentContext(agent_id="managed", binding_revision="revision", mode="autonomous")
if verified_identity else None
),
)
auth = UserAPIKeyAuth(agent_id="managed", api_key="persisted-key", models=["test-model"])
if verified_identity:
auth.managed_agent_context = ManagedAgentContext(
agent_id="managed", binding_revision="revision", mode="autonomous"
)
with pytest.raises(ProxyException) as denied:
await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key")
assert denied.value.code == "403"
if verified_identity:
assert denied.value.message == "Agent identities can only access inference and agent discovery routes"
else:
assert denied.value.message == "This agent requires its bound identity provider token"
@pytest.mark.asyncio
@ -9613,3 +9615,76 @@ async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monk
UserAPIKeyAuth(agent_id="bound"), request, data, "/v1/chat/completions", "sk-test"
)
checks.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("enterprise", [False, True])
@pytest.mark.parametrize("credential", ["custom-credential", "sk-custom-credential"])
@pytest.mark.parametrize("granted", [False, True])
async def test_custom_auth_grants_reach_managed_targets_without_a_virtual_key_row(
monkeypatch: pytest.MonkeyPatch, enterprise: bool, credential: str, granted: bool
) -> None:
import importlib
from typing import Final
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityBinding
target: Final = AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True,
identity=AgentIdentityBinding(
agent_id="target", provider="microsoft_entra", tenant_id="tenant", client_id="client",
issuer="issuer", revision="current",
),
)
registry: Final = AgentRegistry()
registry.register_agent(target)
trusted: Final = UserAPIKeyAuth(
api_key=credential, object_permission={"object_permission_id": "custom", "agents": ["target"] if granted else ["other"]}
)
custom: Final = AsyncMock(return_value=trusted)
database: Final = MagicMock()
database.get_data = AsyncMock(return_value=None)
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
for name, value in {
**_proxy_server_attrs_for_custom_auth(user_custom_auth=None if enterprise else custom),
"prisma_client": database,
}.items():
monkeypatch.setattr(proxy_server, name, value)
module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(module, "enterprise_custom_auth", custom if enterprise else None)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False)
admitted: Final = await _user_api_key_auth_builder(
request=_alias_request("/a2a/target/message/send", {}), api_key=f"Bearer {credential}",
azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
azure_apim_header=None, request_data={},
)
assert await AgentRequestHandler.is_agent_allowed("target", admitted) is granted
custom.assert_awaited_once()
database.get_data.assert_not_awaited()
@pytest.mark.asyncio
async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(monkeypatch: pytest.MonkeyPatch) -> None:
import importlib
from typing import Final
from litellm.proxy import proxy_server
custom: Final = AsyncMock(return_value="sk-master-key")
for name, value in _proxy_server_attrs_for_custom_auth(user_custom_auth=custom).items():
monkeypatch.setattr(proxy_server, name, value)
module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
monkeypatch.setattr(module, "enterprise_custom_auth", custom)
admitted: Final = await _user_api_key_auth_builder(
request=_alias_request("/v1/chat/completions", {}), api_key="Bearer external-credential",
azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
azure_apim_header=None, request_data={},
)
assert admitted.authenticated_by_custom_auth is False
assert admitted.via_virtual_key is True

View file

@ -7656,3 +7656,53 @@ async def test_managed_invocations_enforce_actor_and_target_rate_policies(
for scope in stash.reserved_scopes:
if scope[0] in ("agent", "agent_session"):
assert increments[handler.create_rate_limit_keys(*scope, "tokens")] == 3 - stash.reserved_tokens
@pytest.mark.parametrize("route", ["/a2a/expensive", "/a2a/expensive/message/send", "/v1/a2a/expensive/message/send"])
async def test_a2a_url_target_owns_invocation_fee_and_request_limit(
monkeypatch: pytest.MonkeyPatch, route: str
) -> None:
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.auth.managed_authorization import invocation_target, prepare_agent_invocation
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.types.agents import AgentResponse
expensive: Final = AgentResponse(
agent_id="expensive", agent_name="Expensive", agent_card_params={}, rpm_limit=1,
litellm_params={"cost_per_query": 0.25},
)
cheap: Final = AgentResponse(
agent_id="cheap", agent_name="Cheap", agent_card_params={}, rpm_limit=100,
litellm_params={"cost_per_query": 0.01},
)
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(expensive)
registry.register_agent(cheap)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
side_effect=lambda where, include: {"expensive": expensive, "cheap": cheap}[where["agent_id"]]
)
monkeypatch.setattr(proxy_server, "prisma_client", database)
auth: Final = UserAPIKeyAuth(agent_id="caller")
auth.managed_agent_policy = AgentResponse(
agent_id="caller", agent_name="Caller", agent_card_params={},
object_permission={"object_permission_id": "both-targets", "agents": ["expensive", "cheap"]},
)
body: Final = {"model": "a2a/cheap"}
target: Final = invocation_target(route, body)
assert target is not None
await prepare_agent_invocation(auth, target, AgentIdentityStore.from_client(database))
assert auth.invoked_agent_id == "expensive"
assert auth.invoked_agent_policy == expensive
assert auth.agent_invocation_cost == pytest.approx(0.25)
cache: Final = DualCache()
limiter: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
await _rpm_request(limiter, cache, auth, "a2a/cheap")
with pytest.raises(HTTPException) as denied:
await _rpm_request(limiter, cache, auth, "a2a/cheap")
assert denied.value.status_code == 429
assert "expensive" in str(denied.value.detail)

View file

@ -2996,6 +2996,7 @@ class TestCLIKeyRegenerationFlow:
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "https://proxy.example.com/"
mock_user_info = LiteLLM_UserTable(
@ -3159,6 +3160,7 @@ class TestCLIKeyRegenerationFlow:
# Mock request
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
# Test data
@ -7107,6 +7109,7 @@ class TestCliSsoAttributionMetadata:
from litellm.proxy.management_endpoints.types import CustomOpenID
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
session_key = "cli-session-new-user"
mock_user_info = LiteLLM_UserTable(
@ -7221,6 +7224,7 @@ class TestCliSsoAttributionMetadata:
)
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://internal-proxy.local/"
session_key = "cli-session-4567890"
mock_user_info = LiteLLM_UserTable(
@ -8824,6 +8828,7 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
assertion = assertion_from_sso_login(_ema_id_token(), None)
assert assertion is not None
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
user_info = MagicMock()
@ -9062,6 +9067,7 @@ async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "cid")
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
user_info = MagicMock()
@ -9137,6 +9143,7 @@ def _cli_callback_kwargs(flow):
def _cli_callback_request():
mock_request = MagicMock(spec=Request)
mock_request.scope = {}
mock_request.base_url = "http://localhost:4000/"
return mock_request
@ -9441,3 +9448,45 @@ class TestSessionTokenCookie:
resp = Response()
set_session_token_cookie(resp, _make_http_request(), "jwt-token-value")
assert "Secure" in self._cookie(resp)
@pytest.mark.asyncio
@pytest.mark.parametrize("trusted", [False, True])
@pytest.mark.parametrize("storage_available", [False, True])
async def test_cli_sign_in_enrolls_only_verified_subjects_before_completing(
monkeypatch: pytest.MonkeyPatch, trusted: bool, storage_available: bool
) -> None:
from typing import Final
from litellm.proxy.management_endpoints import ui_sso
from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
flow: Final[dict[str, object]] = {}
kwargs: Final = _cli_callback_kwargs(flow)
subject: Final = MicrosoftInteractiveSubject(issuer="issuer", tenant_id="tenant", oid="subject")
kwargs["request"].scope = {"litellm_microsoft_interactive_subject": subject if trusted else subject.model_dump()}
table: Final = kwargs["prisma_client"].writer_db.litellm_verifiedsubject
table.upsert = AsyncMock(
return_value=SimpleNamespace(kind="human", user_id="cli-user-id", verified_via="sso_interactive"),
side_effect=None if storage_available else RuntimeError("storage unavailable"),
)
monkeypatch.setattr(ui_sso, "get_user_info_from_db", AsyncMock(return_value=_cli_callback_user_info([])))
monkeypatch.setattr(ui_sso, "fetch_cli_sso_team_details", AsyncMock(return_value=()))
monkeypatch.setattr(ui_sso, "retain_sso_identity_assertion_for_ema", AsyncMock())
if trusted and not storage_available:
with pytest.raises(HTTPException) as error:
await ui_sso._complete_cli_sso_callback_session(**kwargs)
assert error.value.status_code == 503
assert "sso_complete" not in flow
return
response: Final = await ui_sso._complete_cli_sso_callback_session(**kwargs)
assert response.status_code == 200
assert flow["session_data"]["user_id"] == "cli-user-id"
if trusted:
table.upsert.assert_awaited_once_with(
where={"issuer_tenant_id_oid": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject"}},
data={"create": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject",
"user_id": "cli-user-id", "verified_via": "sso_interactive"}, "update": {}},
)
else:
table.upsert.assert_not_awaited()