mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): bind dispatch and pricing to resolved admission
This commit is contained in:
parent
e3f672064c
commit
bc5ec23354
10 changed files with 480 additions and 32 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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/"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue