test(agents): inject the access group ceiling resolver instead of patching it

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yassin 2026-09-17 20:04:23 +00:00
parent 322262db01
commit d743e08432
7 changed files with 193 additions and 201 deletions

View file

@ -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)))

View file

@ -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

View file

@ -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):

View file

@ -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:

View file

@ -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]."""

View file

@ -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

View file

@ -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 == []