From bc5ec233549f37e4ae6f55a862baa28d8e99b615 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 18:19:37 -0700 Subject: [PATCH] fix(agents): bind dispatch and pricing to resolved admission --- .../proxy/agent_endpoints/a2a_endpoints.py | 13 ++- litellm/proxy/agent_endpoints/a2a_routing.py | 11 +- .../auth/managed_authorization.py | 64 ++++++++++- litellm/proxy/auth/user_api_key_auth.py | 44 ++++++-- litellm/proxy/route_llm_request.py | 14 +++ .../auth/test_managed_authorization.py | 38 ++++++- .../agent_endpoints/test_a2a_endpoints.py | 78 ++++++++++++-- .../proxy/auth/test_user_api_key_auth.py | 100 +++++++++++++++++- .../proxy/test_route_a2a_models.py | 100 ++++++++++++++++++ .../unit/a2a_protocol/test_cost_calculator.py | 50 ++++++++- 10 files changed, 480 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 6ddcd20d919..d66d9835081 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 26fdf83e0b0..9d91fbfff7d 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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"]) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 9fd1dfa7499..d945c6e1ad2 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -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 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a684b5ebcf3..3d13309a682 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 42ac74cae33..a2dfc95750c 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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) diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index 88160562c9b..bc84e9ddd03 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -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 diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index a5d0d0a3ecc..8f905b2fc96 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -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" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index d69b3e72183..b514e78984f 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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" diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 0429dd97a1c..8783ab1b61e 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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/" diff --git a/tests/unit/a2a_protocol/test_cost_calculator.py b/tests/unit/a2a_protocol/test_cost_calculator.py index 1db6b2bfdce..62dc077054e 100644 --- a/tests/unit/a2a_protocol/test_cost_calculator.py +++ b/tests/unit/a2a_protocol/test_cost_calculator.py @@ -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()