diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index bbb3d30864f..05661584a6b 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -44,6 +44,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, user_api_key_has_admin_view, ) +from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( + CeilingResolver, + resolve_agent_access_group_ceiling, +) from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.auth.user_api_key_auth import ( _get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth @@ -184,6 +188,24 @@ def _has_client_supplied_mcp_auth( return bool(mcp_auth_header) or bool(mcp_server_auth_headers) +def _agent_capped_servers( + allowed_mcp_servers: Sequence[str], + agent_servers: Sequence[str], + agent_access_group_servers: frozenset[str] | None, +) -> tuple[str, ...] | None: + """Servers left once the agent's object_permission and attached access groups both cap the + key/team result, or None when the agent restricts nothing. An attached group set naming no + server is an empty ceiling, not an absent one, so it denies every server.""" + if not agent_servers and agent_access_group_servers is None: + return None + return tuple( + s + for s in allowed_mcp_servers + if (not agent_servers or s in agent_servers) + and (agent_access_group_servers is None or s in agent_access_group_servers) + ) + + def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool: """True when this auth is a keyless subject admitted by the gateway session / bridge user path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``. @@ -1546,21 +1568,14 @@ class MCPRequestHandler: # Check agent permissions if agent_id is set on the key ######################################################### if user_api_key_auth and user_api_key_auth.agent_id: - allowed_mcp_servers_for_agent: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_agent( - user_api_key_auth + agent_capped: Final = _agent_capped_servers( + allowed_mcp_servers, + await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth), + await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth), ) - agent_access_group_servers: Final = await MCPRequestHandler._get_agent_access_group_server_ceiling( - user_api_key_auth - ) - if len(allowed_mcp_servers_for_agent) > 0 or agent_access_group_servers is not None: + if agent_capped is not None: has_lower_level_mcp_restrictions = True - # Intersect: agent can only use servers allowed by key/team AND agent config AND agent access groups - allowed_mcp_servers = [ - s - for s in allowed_mcp_servers - if (len(allowed_mcp_servers_for_agent) == 0 or s in allowed_mcp_servers_for_agent) - and (agent_access_group_servers is None or s in agent_access_group_servers) - ] + allowed_mcp_servers = list(agent_capped) verbose_logger.debug( "Applied agent intersection filter. Final allowed servers: %s", allowed_mcp_servers ) @@ -3148,6 +3163,7 @@ class MCPRequestHandler: @staticmethod async def _get_agent_access_group_server_ceiling( user_api_key_auth: UserAPIKeyAuth, + resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling, ) -> frozenset[str] | None: """ Server IDs the agent's attached unified access groups (``LiteLLM_AgentsTable.access_group_ids``) @@ -3157,13 +3173,10 @@ class MCPRequestHandler: from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, ) - from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( - resolve_agent_access_group_ceiling, - ) if not user_api_key_auth.agent_id: return None - ceiling: Final = await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id) + ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id) if ceiling is None: return None return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids))) diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py index f5adc897e11..67bb43638e0 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py +++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py @@ -33,6 +33,9 @@ class AgentAccessGroupCeiling: agent_ids: frozenset[str] +CeilingResolver: TypeAlias = Callable[[str], Awaitable[AgentAccessGroupCeiling | None]] # mutable-ok: Callable params + + async def _load_agent(agent_id: str) -> AgentResponse | None: from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 11d2a68072c..1759090a29a 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -19,6 +19,10 @@ from litellm.proxy._types import ( LitellmUserRoles, UserAPIKeyAuth, ) +from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( + CeilingResolver, + resolve_agent_access_group_ceiling, +) from litellm.repositories.table_repositories import AgentsRepository from litellm.types.agents import AgentResponse @@ -61,6 +65,7 @@ class AgentRequestHandler: @staticmethod async def resolve_agent_access( user_api_key_auth: UserAPIKeyAuth | None = None, + resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling, ) -> AgentAccess: """ Resolve the agents the given user/key may reach. @@ -71,7 +76,7 @@ class AgentRequestHandler: never widen what it reaches. """ key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth) - agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth) + agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling) if agent_ceiling is None: return key_team_access match key_team_access: @@ -104,13 +109,12 @@ class AgentRequestHandler: @staticmethod async def _agent_access_group_ceiling( user_api_key_auth: UserAPIKeyAuth | None, + resolve_ceiling: CeilingResolver, ) -> frozenset[str] | None: """Stable IDs of the agents the calling agent's attached access groups allow; None when none attached.""" - from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_agent_access_group_ceiling - if user_api_key_auth is None or not user_api_key_auth.agent_id: return None - ceiling: Final = await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id) + ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id) if ceiling is None: return None return _to_stable_ids(ceiling.agent_ids) @@ -119,6 +123,7 @@ class AgentRequestHandler: async def is_agent_allowed( agent_id: str, user_api_key_auth: UserAPIKeyAuth | None = None, + resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling, ) -> bool: """ Check if a specific agent is allowed for the given user/key. @@ -132,7 +137,7 @@ class AgentRequestHandler: """ from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry - match await AgentRequestHandler.resolve_agent_access(user_api_key_auth): + match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling): case UnrestrictedAgentAccess(): return True case RestrictedAgentAccess(allowed_agent_ids): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 5c53ee49717..137cf849389 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -68,6 +68,10 @@ from litellm.proxy._types import ( SpecialModelNames, UserAPIKeyAuth, ) +from litellm.proxy.agent_endpoints.auth.agent_access_groups import ( + CeilingResolver, + resolve_agent_access_group_ceiling, +) from litellm.proxy.auth.budget_throttle import ( budget_throttle_percentage, should_throttle_budget_exceeded, @@ -4199,15 +4203,14 @@ async def _check_agent_access_group_model_access( model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str] valid_token: UserAPIKeyAuth | None, llm_router: Router | None, + resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling, ) -> Literal[True]: """Raises when the key's agent has access groups attached and none of them names the model. Attached groups that name no model deny every model; ``_can_object_call_model`` would read an empty allowlist as unrestricted.""" - from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_agent_access_group_ceiling - if not model or valid_token is None or not valid_token.agent_id: return True - ceiling: Final = await resolve_agent_access_group_ceiling(valid_token.agent_id) + ceiling: Final = await resolve_ceiling(valid_token.agent_id) if ceiling is None: return True if not ceiling.models: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 2f8e3d1cb82..f8e648fd383 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -12,6 +12,7 @@ from starlette.datastructures import Headers from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, UnloadableEntitlementError, + _agent_capped_servers, _is_mcp_admitted_user_subject, ) from litellm.proxy._types import ( @@ -4169,6 +4170,27 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission(): global_mcp_server_manager.registry.pop("direct-server", None) +@pytest.mark.parametrize( + ("agent_servers", "group_ceiling", "expected"), + [ + ([], frozenset({"server_1"}), ("server_1",)), + ([], frozenset({"server_1", "server_2", "server_3"}), ("server_1", "server_2")), + ([], frozenset(), ()), + (["server_2"], frozenset({"server_1", "server_2"}), ("server_2",)), + (["server_1"], frozenset({"server_2"}), ()), + (["server_1"], None, ("server_1",)), + ], +) +def test_agent_capped_servers_intersects_agent_config_and_access_groups(agent_servers, group_ceiling, expected): + """The agent's attached access groups cap the key/team servers alongside its own + object_permission; groups naming no server deny all.""" + assert _agent_capped_servers(["server_1", "server_2"], agent_servers, group_ceiling) == expected + + +def test_agent_capped_servers_without_agent_restrictions_is_uncapped(): + assert _agent_capped_servers(["server_1", "server_2"], [], None) is None + + @pytest.mark.asyncio class TestAgentMCPPermissions: """Test agent-level MCP server and tool permission intersection.""" @@ -4208,64 +4230,45 @@ class TestAgentMCPPermissions: assert sorted(result) == ["server_1", "server_2"] mock_agent.assert_called_once_with(user_api_key_auth) - @pytest.mark.parametrize( - ("group_ceiling", "expected"), - [ - (frozenset({"server_1"}), ["server_1"]), - (frozenset({"server_1", "server_2", "server_3"}), ["server_1", "server_2"]), - (frozenset(), []), - ], - ) - async def test_get_allowed_mcp_servers_agent_access_group_ceiling(self, group_ceiling, expected): - """The agent's attached access groups cap the key/team servers; groups naming no server deny all.""" - user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-ag") - with ( - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key", return_value=["server_1", "server_2"]), - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", return_value=[]), - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", return_value=[]), - patch.object(MCPRequestHandler, "_get_agent_access_group_server_ceiling", return_value=group_ceiling), - ): - access = await MCPRequestHandler.get_mcp_server_access(user_api_key_auth=user_api_key_auth) - assert sorted(access.server_ids) == expected - assert access.scope == "scoped" - - async def test_get_allowed_mcp_servers_agent_without_access_groups_is_uncapped(self): - user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-ag") - with ( - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key", return_value=["server_1", "server_2"]), - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", return_value=[]), - patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", return_value=[]), - patch.object(MCPRequestHandler, "_get_agent_access_group_server_ceiling", return_value=None), - ): - result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=user_api_key_auth) - assert sorted(result) == ["server_1", "server_2"] - async def test_agent_access_group_server_ceiling_expands_group_servers(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer - ceiling = AgentAccessGroupCeiling( - access_group_ids=("ag-1",), - models=frozenset(), - mcp_server_ids=frozenset({"server_1"}), - agent_ids=frozenset(), - ) - with ( - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=ceiling), - ), - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" - ) as mock_manager, - ): - mock_manager.expand_permission_list.return_value = ["server_1"] - result = await MCPRequestHandler._get_agent_access_group_server_ceiling( - UserAPIKeyAuth(api_key="test-key", agent_id="agent-ag") + asked: list[str] = [] + + async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None: + asked.append(agent_id) + return AgentAccessGroupCeiling( + access_group_ids=("ag-1",), + models=frozenset(), + mcp_server_ids=frozenset({"aliased-server"}), + agent_ids=frozenset(), ) - assert result == frozenset({"server_1"}) - mock_manager.expand_permission_list.assert_called_once_with(["server_1"]) - assert await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k")) is None + global_mcp_server_manager.registry["ag-server-id"] = MCPServer( + server_id="ag-server-id", + name="ag-server", + server_name="ag-server", + alias="aliased-server", + url="https://ag-server.example.com", + transport=MCPTransport.http, + ) + try: + result = await MCPRequestHandler._get_agent_access_group_server_ceiling( + UserAPIKeyAuth(api_key="test-key", agent_id="agent-ag"), resolve + ) + finally: + global_mcp_server_manager.registry.pop("ag-server-id", None) + + assert result == frozenset({"ag-server-id"}) + assert asked == ["agent-ag"] + assert ( + await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve) + is None + ) + assert asked == ["agent-ag"] async def test_get_allowed_mcp_servers_key_team_agent_intersection(self): """Key allows [1, 2], agent allows [2, 3]. Result = [2].""" diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a8a55d332b6..2a98e6e4feb 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -11,9 +11,9 @@ import pytest from litellm.constants import UI_SESSION_TOKEN_TEAM_ID -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry -from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling +from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling, CeilingResolver from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( AgentAccess, AgentRequestHandler, @@ -159,86 +159,74 @@ class TestAgentRequestHandler: ), agent_id @staticmethod - def _ceiling(agent_ids: frozenset[str]) -> AgentAccessGroupCeiling: - return AgentAccessGroupCeiling( - access_group_ids=("ag-1",), - models=frozenset(), - mcp_server_ids=frozenset(), - agent_ids=agent_ids, + def _ceiling_resolver(agent_ids: frozenset[str] | None) -> tuple[CeilingResolver, list[str]]: + """A resolver that records the agent ids it was asked about and answers with a fixed + ceiling, or None when the agent has no access groups attached.""" + asked: Final[list[str]] = [] + + async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None: + asked.append(agent_id) + if agent_ids is None: + return None + return AgentAccessGroupCeiling( + access_group_ids=("ag-1",), models=frozenset(), mcp_server_ids=frozenset(), agent_ids=agent_ids + ) + + return resolve, asked + + @staticmethod + def _key_granting(agent_ids: list[str], agent_id: str | None) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id=agent_id, + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="obj-1", agents=agent_ids), ) async def test_agent_access_groups_cap_an_otherwise_unrestricted_key(self): """A key with no agent grant of its own may still only reach the agents its agent's attached access groups name.""" agent_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="caller-agent") + resolve, asked = self._ceiling_resolver(frozenset({"agent-beta"})) - with ( - patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()), - patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()), - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=self._ceiling(frozenset({"agent-beta"}))), - ) as mock_ceiling, - ): - assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess( - frozenset({"agent-beta"}) - ) - assert await AgentRequestHandler.is_agent_allowed("agent-beta", agent_key) is True - assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key) is False - mock_ceiling.assert_called_with("caller-agent") - - async def test_agent_access_groups_intersect_with_key_and_team_grants(self): - agent_key: Final = UserAPIKeyAuth( - api_key="test-key", user_id="test-user", team_id="test-team", agent_id="caller-agent" + assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess( + frozenset({"agent-beta"}) ) + assert await AgentRequestHandler.is_agent_allowed("agent-beta", agent_key, resolve) is True + assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key, resolve) is False + assert asked == ["caller-agent"] * 3 - with ( - patch.object( - AgentRequestHandler, - "_get_allowed_agents_for_key", - return_value=RestrictedAgentAccess(frozenset({"agent-alpha", "agent-beta"})), - ), - patch.object( - AgentRequestHandler, - "_get_allowed_agents_for_team", - return_value=RestrictedAgentAccess(frozenset({"agent-alpha", "agent-beta", "agent-gamma"})), - ), - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=self._ceiling(frozenset({"agent-beta", "agent-gamma"}))), - ), - ): - assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess( - frozenset({"agent-beta"}) - ) + async def test_agent_access_groups_intersect_with_key_grants(self): + agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent") + resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"})) + + assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess( + frozenset({"agent-beta"}) + ) + assert await AgentRequestHandler.is_agent_allowed("agent-gamma", agent_key, resolve) is False async def test_agent_access_groups_naming_no_agent_deny_every_agent(self): agent_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="caller-agent") + resolve, _ = self._ceiling_resolver(frozenset()) - with ( - patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()), - patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()), - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=self._ceiling(frozenset())), - ), - ): - assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess(frozenset()) - assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key) is False + assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess(frozenset()) + assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key, resolve) is False + + async def test_agent_without_access_groups_keeps_key_grants(self): + agent_key: Final = self._key_granting(["agent-alpha"], agent_id="caller-agent") + resolve, asked = self._ceiling_resolver(None) + + assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess( + frozenset({"agent-alpha"}) + ) + assert asked == ["caller-agent"] async def test_key_without_agent_never_consults_agent_access_groups(self): plain_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + resolve, asked = self._ceiling_resolver(frozenset()) - with ( - patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()), - patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()), - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=self._ceiling(frozenset())), - ) as mock_ceiling, - ): - assert await AgentRequestHandler.resolve_agent_access(plain_key) == UnrestrictedAgentAccess() - mock_ceiling.assert_not_called() + assert await AgentRequestHandler.resolve_agent_access(plain_key, resolve) == UnrestrictedAgentAccess() + assert asked == [] async def test_empty_access_group_denies_every_agent(self): """LIT-5143: a key restricted to an access group that resolves to no agents is diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ad1742db4b7..b9b9a786d93 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -32,11 +32,13 @@ from litellm.proxy._types import ( UserAPIKeyAuth, WebhookEvent, ) +from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling, CeilingResolver from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _cache_management_object, _can_object_call_model, _can_object_call_vector_stores, + _check_agent_access_group_model_access, _check_end_user_budget, _check_team_member_budget, _fetch_key_object_from_db_with_reconnect, @@ -8466,89 +8468,64 @@ def test_request_skips_budget_checks_extends_route_rule_with_zero_cost_models() # Agent access group model ceiling -def _agent_model_ceiling(models: frozenset[str]): - from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling +def _agent_model_ceiling_resolver( + models: frozenset[str] | None, +) -> tuple[CeilingResolver, list[str]]: + """Resolver that records the agent ids it was asked about and answers with a fixed model + ceiling, or None when the agent has no access groups attached.""" + asked: Final[list[str]] = [] - return AgentAccessGroupCeiling( - access_group_ids=("ag-1",), models=models, mcp_server_ids=frozenset(), agent_ids=frozenset() - ) + async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None: + asked.append(agent_id) + if models is None: + return None + return AgentAccessGroupCeiling( + access_group_ids=("ag-1",), models=models, mcp_server_ids=frozenset(), agent_ids=frozenset() + ) - -async def _run_common_checks_for_agent_key(model: str, valid_token: UserAPIKeyAuth): - from fastapi import Request - - from litellm.proxy.auth.auth_checks import common_checks - - return await common_checks( - request_body={"model": model, "messages": [{"role": "user", "content": "hi"}]}, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route="/chat/completions", - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=valid_token, - request=MagicMock(spec=Request), - ) + return resolve, asked @pytest.mark.asyncio -async def test_common_checks_agent_access_groups_cap_models_even_when_key_allows_them(): +async def test_agent_access_groups_cap_models_even_when_key_allows_them(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"]) + resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"})) - with patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=_agent_model_ceiling(frozenset({"gpt-5"}))), - ): - assert await _run_common_checks_for_agent_key("gpt-5", agent_key) is True + assert await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve) is True - with pytest.raises(ProxyException) as exc_info: - await _run_common_checks_for_agent_key("claude-sonnet", agent_key) + with pytest.raises(ProxyException) as exc_info: + await _check_agent_access_group_model_access("claude-sonnet", agent_key, None, resolve) assert exc_info.value.type == ProxyErrorTypes.agent_model_access_denied assert exc_info.value.code == str(status.HTTP_403_FORBIDDEN) + assert asked == ["agent-1", "agent-1"] @pytest.mark.asyncio -async def test_common_checks_agent_access_groups_naming_no_model_deny_every_model(): +async def test_agent_access_groups_naming_no_model_deny_every_model(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[]) + resolve, _ = _agent_model_ceiling_resolver(frozenset()) - with ( - patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=_agent_model_ceiling(frozenset())), - ), - pytest.raises(ProxyException) as exc_info, - ): - await _run_common_checks_for_agent_key("gpt-5", agent_key) + with pytest.raises(ProxyException) as exc_info: + await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve) assert exc_info.value.type == ProxyErrorTypes.agent_model_access_denied @pytest.mark.asyncio -async def test_common_checks_agent_without_access_groups_adds_no_model_ceiling(): +async def test_agent_without_access_groups_adds_no_model_ceiling(): agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"]) + resolve, asked = _agent_model_ceiling_resolver(None) - with patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=None), - ) as mock_ceiling: - assert await _run_common_checks_for_agent_key("gpt-5", agent_key) is True - assert await _run_common_checks_for_agent_key("claude-sonnet", agent_key) is True - - mock_ceiling.assert_called_with("agent-1") + assert await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve) is True + assert await _check_agent_access_group_model_access("claude-sonnet", agent_key, None, resolve) is True + assert asked == ["agent-1", "agent-1"] @pytest.mark.asyncio -async def test_common_checks_key_without_agent_never_consults_agent_access_groups(): +async def test_key_without_agent_never_consults_agent_access_groups(): plain_key: Final = UserAPIKeyAuth(token="plain-token", models=["gpt-5"]) + resolve, asked = _agent_model_ceiling_resolver(frozenset()) - with patch( - "litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling", - new=AsyncMock(return_value=_agent_model_ceiling(frozenset())), - ) as mock_ceiling: - assert await _run_common_checks_for_agent_key("gpt-5", plain_key) is True - - mock_ceiling.assert_not_called() + assert await _check_agent_access_group_model_access("gpt-5", plain_key, None, resolve) is True + assert asked == []