fix(agents): bind dispatch and pricing to resolved admission

This commit is contained in:
Joshua Valluru 2026-09-30 18:19:37 -07:00
parent e3f672064c
commit bc5ec23354
10 changed files with 480 additions and 32 deletions

View file

@ -33,6 +33,7 @@ from litellm.proxy.a2a.version_convert import (
normalize_request_params,
normalize_stream_event,
)
from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
from litellm.proxy.agent_endpoints.databricks_oauth import (
DATABRICKS_OAUTH_PARAM,
resolve_databricks_app_auth_header,
@ -706,9 +707,10 @@ async def invoke_agent_a2a(
params.pop(key)
# Find the agent
agent: Final = await _get_agent(agent_id)
if agent is None:
registered: Final = await _get_agent(agent_id)
if registered is None:
return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404)
agent: Final = agent_invocation_policy(user_api_key_dict, registered)
served_version: Final = _served_version(agent, request, original_method)
@ -732,7 +734,12 @@ async def invoke_agent_a2a(
agent_name: Final = agent_card_params.get("name", agent_id)
# Get litellm_params (may include custom_llm_provider for completion bridge)
litellm_params: dict[str, object] = agent.litellm_params or {}
litellm_params: dict[
str, object
] = { # mutable-ok: A2A SDK and completion bridge accept provider parameters as a dict
**(agent.litellm_params or MappingProxyType({})),
"cost_per_query": user_api_key_dict.agent_invocation_cost,
}
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
# Hand the authenticated key hash to the completion bridge so provider

View file

@ -13,12 +13,16 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
from litellm.types.agents import AgentResponse
async def route_a2a_agent_request(
data: dict,
route_type: str,
user_api_key_dict: UserAPIKeyAuth | None = None,
*,
registered_agent: AgentResponse | None = None,
) -> Any | None:
"""
Route A2A agent requests directly to litellm with injected API base.
@ -47,12 +51,14 @@ async def route_a2a_agent_request(
agent_name: Final = model_name[4:]
# Look up agent in registry
agent: Final = await get_agent_with_read_through(agent_name)
if agent is None:
registered: Final = registered_agent or await get_agent_with_read_through(agent_name)
if registered is None:
verbose_proxy_logger.error("[A2A] Agent '%s' not found in registry", agent_name)
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
agent: Final = agent_invocation_policy(user_api_key_dict, registered)
# Verify the caller is permitted to use this agent (admins bypass the check)
is_admin: Final = user_api_key_dict is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
@ -76,6 +82,7 @@ async def route_a2a_agent_request(
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
# Inject API base and route to litellm
data.pop("litellm_params", None)
data["api_base"] = agent.agent_card_params["url"]
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, data["api_base"])

View file

@ -99,12 +99,18 @@ def managed_inference_request(
cli_model: str | None,
path_model: object = None,
query_model: object = None,
*,
auth: UserAPIKeyAuth | None = None,
require_model: bool = True,
model_group_alias: object = None,
) -> dict[str, object]:
from litellm.proxy.auth.route_checks import RouteChecks
if route in _MANAGED_REALTIME_ROUTES:
model: Final = query_model or body.get("model")
if not isinstance(model, str) or not model:
if not require_model:
return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
@ -117,8 +123,32 @@ def managed_inference_request(
endpoint_model: Final = path_model or (
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
)
effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
from litellm.proxy.common_utils.model_listing_utils import CallerAliases, alias_target
from litellm.proxy.litellm_pre_call_utils import (
_update_model_if_key_alias_exists,
_update_model_if_team_alias_exists,
)
aliased_body: Final = dict(body) # mutable-ok: existing alias helpers rewrite their request copy
if auth is not None:
_update_model_if_team_alias_exists(aliased_body, auth)
_update_model_if_key_alias_exists(aliased_body, auth)
selected: Final = resolve_inference_model(aliased_body.get("model"), settings, cli_model, endpoint_model, kind=kind)
import litellm
aliased: Final = (
alias_target(selected, CallerAliases((), (litellm.model_alias_map, auth.aliases))) or selected
if isinstance(selected, str) and auth is not None
else selected
)
from litellm.router_utils.common_utils import resolve_model_group_alias
effective: Final = (
(resolve_model_group_alias(model_group_alias, aliased) or aliased) if isinstance(aliased, str) else aliased
)
if not isinstance(effective, str) or not effective:
if not require_model:
return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
raise_identity_failure(
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
)
@ -256,6 +286,10 @@ async def prepare_agent_invocation(
effective: Final = target if target is not None else registered
pricing: Final = effective.litellm_params or MappingProxyType({})
fixed_fee: Final = pricing.get("cost_per_query")
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
auth.invoked_agent_id = effective.agent_id
auth.invoked_agent_policy = effective
if (
not effective.identity_managed
and effective.litellm_budget_table is None
@ -264,10 +298,6 @@ async def prepare_agent_invocation(
and fixed_fee is None
):
return
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
auth.invoked_agent_id = effective.agent_id
auth.invoked_agent_policy = effective
if (
billable
and auth.agent_id is None
@ -304,3 +334,27 @@ async def prepare_agent_invocation(
)
)
auth.agent_invocation_cost = fee
def agent_invocation_policy(auth: UserAPIKeyAuth | None, registered: AgentResponse) -> AgentResponse:
admitted: Final = auth.invoked_agent_policy if auth is not None else None
if admitted is not None and auth is not None:
matching: Final = auth.invoked_agent_id == registered.agent_id == admitted.agent_id
captured_price: Final = (admitted.litellm_params or MappingProxyType({})).get(
"cost_per_query"
) is None or auth.agent_invocation_cost is not None
if matching and captured_price:
return admitted
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch does not match its admission")
)
if (
registered.identity_managed
or registered.identity is not None
or registered.litellm_budget_table is not None
or (registered.litellm_params or MappingProxyType({})).get("cost_per_query") is not None
):
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Agent dispatch requires a matching admission")
)
return registered

View file

@ -3233,7 +3233,14 @@ async def _authorize_authenticated_request(
prepare_agent_invocation,
)
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.proxy_server import general_settings, prisma_client, user_model
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
prisma_client,
proxy_config,
proxy_logging_obj,
user_model,
)
store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None
if user_api_key_auth_obj.agent_id is not None:
@ -3242,19 +3249,34 @@ async def _authorize_authenticated_request(
route, request.method
):
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
authorized_data: Final = (
managed_inference_request(
route,
request_data,
general_settings,
user_model,
request.path_params.get("model") or request.path_params.get("model_name"),
request.query_params.get("model"),
router_settings: Final = (
await proxy_config.get_hierarchical_router_settings(
user_api_key_dict=user_api_key_auth_obj,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
)
if user_api_key_auth_obj.managed_agent_policy is not None
if llm_router is not None and RouteChecks.is_llm_api_route(route=route)
else None
)
inference_data: Final = managed_inference_request(
route,
request_data,
general_settings,
user_model,
request.path_params.get("model") or request.path_params.get("model_name"),
request.query_params.get("model"),
model_group_alias=router_settings.get("model_group_alias")
if isinstance(router_settings, Mapping)
else None,
auth=user_api_key_auth_obj,
require_model=user_api_key_auth_obj.managed_agent_policy is not None,
)
target_name: Final = invocation_target(route, inference_data)
authorized_data: Final = (
inference_data
if target_name is not None or user_api_key_auth_obj.managed_agent_policy is not None
else request_data
)
target_name: Final = invocation_target(route, authorized_data)
if target_name is not None:
await prepare_agent_invocation(
user_api_key_auth_obj,

View file

@ -477,6 +477,20 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
data.pop("enable_tag_filtering", None)
if _is_a2a_agent_model(data.get("model")):
from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
registered_agent: Final = await get_agent_with_read_through(data["model"][4:])
if registered_agent is not None:
agent_response: Final = await route_a2a_agent_request(
data, route_type, user_api_key_dict=user_api_key_dict, registered_agent=registered_agent
)
if agent_response is not None:
return agent_response
if user_api_key_dict is not None and user_api_key_dict.invoked_agent_policy is not None:
raise HTTPException(503, "Agent dispatch does not match its admission")
team_id: Final = get_team_id_from_data(data)
router_model_names: Final = llm_router.model_names if llm_router is not None else []
is_proxy_admin_without_team: Final = team_id is None and _is_proxy_admin_request(data)

View file

@ -534,7 +534,8 @@ async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: p
await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
assert auth.managed_agent_policy is None
assert auth.billing_agent_policy is None
assert auth.invoked_agent_id is None
assert auth.invoked_agent_id == "agent"
assert auth.invoked_agent_policy is not None
@pytest.mark.asyncio
@ -744,3 +745,38 @@ async def test_unmanaged_invocation_rejects_invalid_configured_fees(
assert exc.value.status_code == 503
assert "Agent invocation price is invalid" in str(exc.value.detail)
assert auth.agent_invocation_cost is None
@pytest.mark.parametrize("route", ("/realtime", "/v1/chat/completions", "/v1/files"))
def test_ordinary_requests_without_models_keep_existing_validation(route: str) -> None:
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
assert managed_inference_request(route, {}, {}, None, require_model=False) == {}
@pytest.mark.parametrize("state", (
{"identity_managed": True},
{"identity": BINDING},
{"litellm_budget_table": {"budget_id": "budget", "max_budget": 0.5}},
{"litellm_params": {"cost_per_query": 0.25}},
))
@pytest.mark.parametrize("has_auth", (False, True))
def test_protected_agent_dispatch_requires_admission(state: dict[str, object], has_auth: bool) -> None:
from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
registered: Final = agent(**{"identity": None, "identity_managed": False, **state})
with pytest.raises(HTTPException, match="admission") as exc:
agent_invocation_policy(UserAPIKeyAuth() if has_auth else None, registered)
assert exc.value.status_code == 503
def test_paid_agent_dispatch_requires_the_captured_fee() -> None:
from litellm.proxy.agent_endpoints.auth.managed_authorization import agent_invocation_policy
policy: Final = agent(identity=None, identity_managed=False, litellm_params={"cost_per_query": 0.25})
auth: Final = UserAPIKeyAuth()
auth.invoked_agent_id = policy.agent_id
auth.invoked_agent_policy = policy
with pytest.raises(HTTPException, match="admission") as exc:
agent_invocation_policy(auth, policy)
assert exc.value.status_code == 503

View file

@ -16,7 +16,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.agents import AgentCaller
from litellm.types.agents import AgentCaller, AgentResponse
AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]]
@ -25,6 +25,8 @@ AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]]
class CapturedAgentCall:
request_id: object
agent_extra_headers: dict[str, str] | None
cost_per_query: object
api_base: object
@pytest.mark.asyncio
@ -58,7 +60,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
}
# Mock agent
mock_agent = MagicMock()
mock_agent = _make_agent_mock()
mock_agent.agent_id = "test-agent"
mock_agent.agent_card_params = {
"url": "http://backend-agent:10001",
@ -211,7 +213,7 @@ async def test_invoke_agent_a2a_handles_none_agent_card_params():
"""
from litellm.proxy._types import UserAPIKeyAuth
mock_agent = MagicMock()
mock_agent = _make_agent_mock()
mock_agent.agent_card_params = None
mock_agent.litellm_params = None
@ -295,7 +297,7 @@ async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge():
resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}}
return resp
mock_agent = MagicMock()
mock_agent = _make_agent_mock()
mock_agent.agent_id = "lf-agent"
mock_agent.agent_name = "lf-agent"
# No URL: the bridge derives the endpoint from the LangFlow agent config.
@ -376,6 +378,9 @@ def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
agent.litellm_params = {}
agent.static_headers = None
agent.extra_headers = None
agent.identity_managed = False
agent.identity = None
agent.litellm_budget_table = None
return agent
@ -394,7 +399,7 @@ def _make_request_mock(method: str, params: Mapping[str, object], request_id: ob
def _base_patches(
agent: MagicMock, add_litellm_data: AddLiteLLMData | None = None
agent: MagicMock | AgentResponse, add_litellm_data: AddLiteLLMData | None = None
) -> list[AbstractContextManager[object]]:
return [
patch(
@ -437,7 +442,7 @@ async def _invoke_message_method(
mock_request: MagicMock,
user_api_key_dict: UserAPIKeyAuth,
add_litellm_data: AddLiteLLMData | None = None,
agent: MagicMock | None = None,
agent: MagicMock | AgentResponse | None = None,
) -> CapturedAgentCall:
from fastapi.responses import JSONResponse
@ -490,7 +495,12 @@ async def _invoke_message_method(
kwargs: Final = downstream.call_args.kwargs
request_id: Final = kwargs["request"].__dict__["id"] if is_send else kwargs["request_id"]
return CapturedAgentCall(request_id=request_id, agent_extra_headers=kwargs.get("agent_extra_headers"))
return CapturedAgentCall(
request_id=request_id,
agent_extra_headers=kwargs.get("agent_extra_headers"),
cost_per_query=kwargs["litellm_params"].get("cost_per_query"),
api_base=kwargs["api_base"],
)
@pytest.mark.asyncio
@ -2689,3 +2699,57 @@ def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_
)
assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"}
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ("message/send", "message/stream"))
@pytest.mark.parametrize("fee", (None, 0.0, 0.25))
async def test_native_dispatch_keeps_the_admitted_price_and_destination(
monkeypatch: pytest.MonkeyPatch,
method: str,
fee: float | None,
) -> None:
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
admitted: Final = AgentResponse(
agent_id="test-agent",
agent_name="test-agent",
agent_card_params={"url": "https://admitted.test/"},
litellm_params={"cost_per_query": fee} if fee is not None else {},
)
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(admitted)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
await prepare_agent_invocation(auth, admitted.agent_id, None)
changed: Final = admitted.model_copy(
update={
"litellm_params": {"cost_per_query": 0.75},
"agent_card_params": {"url": "https://changed.test/"},
}
)
captured: Final = await _invoke_message_method(
method,
_make_request_mock(method, _HELLO_MESSAGE_PARAMS),
auth,
agent=changed,
)
assert captured.cost_per_query == fee
assert captured.api_base == "https://admitted.test/"
@pytest.mark.asyncio
async def test_native_dispatch_returns_not_found_for_a_removed_agent(monkeypatch: pytest.MonkeyPatch) -> None:
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
monkeypatch.setattr(proxy_server, "prisma_client", None)
response: Final = await invoke_agent_a2a(
agent_id="removed", request=_make_request_mock("message/send", _HELLO_MESSAGE_PARAMS),
fastapi_response=MagicMock(), user_api_key_dict=UserAPIKeyAuth(),
)
assert response.status_code == 404
assert json.loads(response.body)["error"]["message"] == "Agent 'removed' not found"

View file

@ -6267,6 +6267,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it():
"type": "http",
"headers": [(b"content-type", b"application/json")],
"method": "POST",
"query_string": b"",
}
)
request._url = URL(url="/chat/completions")
@ -6321,6 +6322,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder()
"type": "http",
"headers": [(b"content-type", b"application/json")],
"method": "POST",
"query_string": b"",
}
)
request._url = URL(url="/chat/completions")
@ -8403,7 +8405,7 @@ _DDTRACE_AUTH_PROBE = dedent(
async def auth(api_key):
request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions"})
request = Request(scope={"type": "http", "headers": [], "method": "POST", "path": "/chat/completions", "query_string": b""})
request._url = URL(url="/chat/completions")
try:
await user_api_key_auth(
@ -9811,3 +9813,99 @@ def test_free_model_only_waives_budgets_without_a_paid_agent_invocation(
)
is skipped
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"selection",
(
"body",
"query",
"path",
"cli",
"default",
"key-alias",
"team-alias",
"global-alias",
"alias-chain",
"query-alias",
"query-over-alias",
"different-agent",
"router-alias",
"query-router-alias",
),
)
async def test_agent_admission_prices_the_model_selected_for_dispatch(
monkeypatch: pytest.MonkeyPatch,
selection: str,
) -> None:
from typing import Final
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
from litellm.types.agents import AgentResponse
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(
AgentResponse(
agent_id="paid",
agent_name="Paid",
agent_card_params={},
litellm_params={"cost_per_query": 0.25},
)
)
registry.register_agent(
AgentResponse(
agent_id="other",
agent_name="Other",
agent_card_params={},
litellm_params={"cost_per_query": 0.75},
)
)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
monkeypatch.setattr(litellm, "model_alias_map", {"global": "a2a/paid"})
for name, value in {
**_proxy_attrs_for_centralized_checks(),
"user_model": "a2a/paid" if selection == "cli" else None,
"llm_router": litellm.Router(model_list=[]) if "router-alias" in selection else None,
"general_settings": {"completion_model": "a2a/paid"} if selection == "default" else {},
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
}.items():
monkeypatch.setattr(proxy_server, name, value)
body: Final = {
"model": {
"body": "a2a/paid",
"key-alias": "alias",
"team-alias": "team-alias",
"global-alias": "global",
"alias-chain": "team-alias",
"query-over-alias": "alias",
"different-agent": "a2a/other",
"router-alias": "router-alias",
}.get(selection, "gpt-4o"),
"messages": [{"role": "user", "content": "Hello"}],
}
route: Final = "/openai/deployments/a2a/paid/chat/completions" if selection == "path" else "/v1/chat/completions"
request: Final = _alias_request(route, body, path_params={"model": "a2a/paid"} if selection == "path" else {})
if selection in ("query", "query-over-alias", "different-agent", "query-alias"):
request.scope["query_string"] = b"model=alias" if selection == "query-alias" else b"model=a2a%2Fpaid"
if selection == "query-router-alias":
request.scope["query_string"] = b"model=router-alias"
auth: Final = UserAPIKeyAuth(
router_settings={"model_group_alias": {"router-alias": "a2a/paid"}},
user_role="proxy_admin",
aliases={"alias": "global" if selection == "alias-chain" else "a2a/paid"},
team_model_aliases={"team-alias": "alias" if selection == "alias-chain" else "a2a/paid"},
)
with patch(
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
new_callable=AsyncMock,
return_value=None,
) as reserve:
assert await _authorize_authenticated_request(auth, request, body, route, "test-key") is None
assert auth.invoked_agent_id == "paid"
assert auth.agent_invocation_cost == pytest.approx(0.25)
reserve.assert_awaited_once()
assert reserve.call_args.kwargs["valid_token"].agent_invocation_cost == pytest.approx(0.25)
assert reserve.call_args.kwargs["request_body"]["model"] == "a2a/paid"

View file

@ -183,3 +183,103 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re
assert call_kwargs["model"] == f"a2a/{agent_name}"
assert call_kwargs["api_base"] == "http://sibling-db-agent.example.com"
prisma_client.db.litellm_agentstable.find_unique.assert_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"override", ("api_key", "api_base", "user_config", "router_settings_override", "deployment", "no-router")
)
async def test_registered_agent_dispatch_owns_the_admitted_destination_and_fee(monkeypatch: pytest.MonkeyPatch, override: str) -> None:
from typing import Final
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
from litellm.types.agents import AgentResponse
agent: Final = AgentResponse(
agent_id="paid",
agent_name="paid",
agent_card_params={"url": "https://registered.test/"},
litellm_params={"cost_per_query": 0.25},
)
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(agent)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
router: Final = _router_without_models()
if override == "deployment":
router.is_recognized_model.return_value = True
provider: Final = AsyncMock(return_value={"id": "reply"})
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
with patch("litellm.acompletion", provider):
await prepare_agent_invocation(auth, "paid", None)
pending: Final = await route_request(
data={
"model": "a2a/paid",
"messages": [{"role": "user", "content": "Hi"}],
"cost_per_query": 99.0,
**(
{override: {} if override in ("user_config", "router_settings_override") else "override"}
if override not in ("deployment", "no-router")
else {}
),
},
llm_router=None if override == "no-router" else router,
user_model=None,
route_type="acompletion",
user_api_key_dict=auth,
)
assert await pending == {"id": "reply"}
provider.assert_awaited_once()
assert provider.call_args.kwargs["api_base"] == "https://registered.test/"
assert provider.call_args.kwargs["cost_per_query"] == 0.25
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ("a2a/paid", "gpt-4o"))
async def test_routing_overrides_cannot_dispatch_without_matching_agent_admission(
monkeypatch: pytest.MonkeyPatch, model: str,
) -> None:
from typing import Final
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
from litellm.types.agents import AgentResponse
registry: Final = agent_registry.AgentRegistry()
registry.register_agent(AgentResponse(
agent_id="paid", agent_name="paid", agent_card_params={"url": "https://agent.test/"},
litellm_params={"cost_per_query": 0.25},
))
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
monkeypatch.setattr(proxy_server, "prisma_client", None)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
if model == "gpt-4o":
await prepare_agent_invocation(auth, "paid", None)
provider: Final = Mock(return_value=None)
with patch("litellm.acompletion", provider), pytest.raises(HTTPException, match="admission") as exc:
await route_request(
data={"model": model, "api_base": "https://override.test/", "messages": [{"role": "user", "content": "Hi"}]},
llm_router=None, user_model=None, route_type="acompletion", user_api_key_dict=auth,
)
assert exc.value.status_code == 503
provider.assert_not_called()
@pytest.mark.asyncio
async def test_unregistered_direct_agent_keeps_explicit_endpoint_routing(monkeypatch: pytest.MonkeyPatch) -> None:
from typing import Final
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
monkeypatch.setattr(proxy_server, "prisma_client", None)
provider: Final = AsyncMock(return_value={"id": "direct-reply"})
with patch("litellm.acompletion", provider):
pending: Final = await route_request(
data={"model": "a2a/direct", "api_base": "https://direct.test/", "messages": [{"role": "user", "content": "Hi"}]},
llm_router=None, user_model=None, route_type="acompletion",
)
assert await pending == {"id": "direct-reply"}
assert provider.call_args.kwargs["api_base"] == "https://direct.test/"

View file

@ -456,11 +456,12 @@ async def test_asend_message_streaming_triggers_callbacks():
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", (False, True))
@pytest.mark.parametrize("claimed_fee", (None, -1000.0, 0.0, 99.0))
@pytest.mark.parametrize("fee_field", ("cost_per_query", "litellm_params"))
@pytest.mark.parametrize("changed_after_admission", (False, True))
@pytest.mark.parametrize("configured_fee", (None, 0.0, 0.25))
async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing(
monkeypatch: pytest.MonkeyPatch, stream: bool, claimed_fee: float | None,
changed_after_admission: bool, configured_fee: float | None
changed_after_admission: bool, configured_fee: float | None, fee_field: str
) -> None:
import json
from typing import Final
@ -509,7 +510,8 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
pending: Final = await route_a2a_agent_request(
data={"model": "a2a/fee-target", "messages": [{"role": "user", "content": "Hello"}],
"stream": stream, "client": client,
**({"cost_per_query": claimed_fee} if claimed_fee is not None else {})},
**({fee_field: claimed_fee if fee_field == "cost_per_query" else {"cost_per_query": claimed_fee}}
if claimed_fee is not None else {})},
route_type="acompletion", user_api_key_dict=auth,
)
response: Final = await pending
@ -526,3 +528,47 @@ async def test_chat_adapter_settles_the_admitted_agent_fee_without_model_pricing
assert logger.response_cost == pytest.approx(configured_fee)
finally:
await client.close()
@pytest.mark.asyncio
@pytest.mark.parametrize("admitted_target", (None, "other"))
async def test_chat_agent_dispatch_rejects_missing_or_different_admission(
monkeypatch: pytest.MonkeyPatch,
admitted_target: str | None,
) -> None:
from typing import Final
from unittest.mock import Mock
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
from litellm.types.agents import AgentResponse
registry: Final = agent_registry.AgentRegistry()
for name in ("paid", "other"):
registry.register_agent(
AgentResponse(
agent_id=name,
agent_name=name,
agent_card_params={"url": "https://agent.test/"},
litellm_params={"cost_per_query": 0.25},
)
)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
auth: Final = UserAPIKeyAuth(user_role="proxy_admin")
if admitted_target is not None:
await prepare_agent_invocation(auth, admitted_target, None)
provider: Final = Mock(return_value=None)
monkeypatch.setattr(litellm, "acompletion", provider)
with pytest.raises(HTTPException) as exc:
await route_a2a_agent_request(
data={"model": "a2a/paid", "messages": [{"role": "user", "content": "Hello"}]},
route_type="acompletion",
user_api_key_dict=auth,
)
assert exc.value.status_code == 503
assert "admission" in str(exc.value.detail).lower()
provider.assert_not_called()