diff --git a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py index dd7390a784d..1b74b877824 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py +++ b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py @@ -36,7 +36,9 @@ async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]: if context.user_id is None: return () human: Final = await _delegated_resource_subject(context.user_id) - allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human) + allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers( + human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset() + ) return tuple(sorted(own.intersection(allowed))) except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access raise_identity_failure( @@ -57,7 +59,9 @@ async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] if context.user_id is None: return [] human: Final = await _delegated_resource_subject(context.user_id) - human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human) + human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools( + server_id, human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset() + ) if own is None: return human_tools return own if human_tools is None else sorted(frozenset(own).intersection(human_tools)) 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 aed93111d72..60241d0fc40 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 @@ -1204,7 +1204,7 @@ class MCPRequestHandler: return None @staticmethod - async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth: + async def _reload_admitted_key(key_hash: str, *, check_db_only: bool = False) -> UserAPIKeyAuth: """Reload the live key record an admitted envelope references and re-check live policy. Resolving the current ``UserAPIKeyAuth`` (cache first, then DB) is what stops the @@ -1236,6 +1236,7 @@ class MCPRequestHandler: hashed_token=key_hash, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, + check_db_only=check_db_only, ) except (ProxyException, HTTPException): raise HTTPException(status_code=401, detail="Invalid or expired credential") from None @@ -1841,7 +1842,9 @@ class MCPRequestHandler: return scoped @staticmethod - async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + async def admitted_subject_sources( + auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None + ) -> list[UserAPIKeyAuth]: """The independent sources a keyless admitted subject reaches MCP servers through: their own direct grants, plus every team they are a live roster member of. @@ -1858,6 +1861,8 @@ class MCPRequestHandler: if not auth.user_id or prisma_client is None: return sources for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth): + if allowed_team_ids is not None and team_id not in allowed_team_ids: + continue team_obj = await MCPRequestHandler._roster_team_object(team_id, auth) if team_obj is None: continue @@ -1942,7 +1947,9 @@ class MCPRequestHandler: return team_obj @staticmethod - async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]: + async def admitted_source_grants( + auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None + ) -> list[tuple[UserAPIKeyAuth, set[str]]]: """``(source, the servers that source grants)`` for every source of an admitted subject. THE owner of "which source reaches which server". The reachable union, the per-team throttle @@ -1951,15 +1958,17 @@ class MCPRequestHandler: roster instead of by grant charged unrelated teams' buckets).""" return [ (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) - for source in await MCPRequestHandler.admitted_subject_sources(auth) + for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids) ] @staticmethod - async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]: + async def resolve_admitted_subject_servers( + auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None + ) -> list[str]: """Union of what each of the admitted subject's sources reaches, each answered by the canonical resolver so no rule is reimplemented for this caller shape.""" reachable: Final[set[str]] = set() - for _source, granted in await MCPRequestHandler.admitted_source_grants(auth): + for _source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids): reachable.update(granted) return list(reachable) @@ -2017,7 +2026,9 @@ class MCPRequestHandler: return min((source for source, _ in granting), key=lambda s: s.team_id or "") @staticmethod - async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None: + async def resolve_admitted_subject_tools( + server_id: str, auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None + ) -> list[str] | None: """Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the sources that actually grant that server. @@ -2039,7 +2050,7 @@ class MCPRequestHandler: ) or await MCPRequestHandler.admin_view_unscoped(auth) allowed: Final[set[str]] = set() - for source, granted in await MCPRequestHandler.admitted_source_grants(auth): + for source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids): # The open channel is evaluated against the user's OWN source (team_id is None), so that # source's restrictions apply to it; a team's rules never ride an open-channel server. if server_id not in granted and not (reachable_via_open_channel and source.team_id is None): @@ -3464,7 +3475,8 @@ class MCPRequestHandler: toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions( obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy ) - return list({*expanded_direct_servers, *access_group_servers, *toolset_grants}) + inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions) + return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools}) except Exception as e: if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError): raise diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d421363ee92..be06d2e7321 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -15,6 +15,7 @@ from pydantic import ( Json, JsonValue, PositiveInt, + PrivateAttr, field_validator, model_validator, ) @@ -3334,6 +3335,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob invoked_agent_policy: AgentResponse | None = Field(default=None, exclude=True) agent_invocation_cost: float | None = Field(default=None, exclude=True) billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True) + _managed_delegation_verified: bool = PrivateAttr(default=False) managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True) managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True) agent_caller: AgentCaller | None = Field( diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 4b9238f1868..4a77bb277bd 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -175,7 +175,17 @@ class AgentRequestHandler: or user_api_key_auth is None ): return False - fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True}) + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler + + key_hash: Final = user_api_key_auth.api_key or user_api_key_auth.token + authority: Final = ( + await MCPRequestHandler._reload_admitted_key(key_hash, check_db_only=True) + if key_hash + and user_api_key_auth.managed_agent_policy is None + and not user_api_key_auth.is_session_token + else user_api_key_auth + ) + fresh_auth: Final = authority.model_copy(update={"requires_fresh_policy": True}) explicit: Final = await _granted_agent_ids( fresh_auth, _strict_agent_access, @@ -643,16 +653,18 @@ async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess: return RestrictedAgentAccess(capped) if context.user_id is None: return RestrictedAgentAccess(frozenset()) - human_ids: Final = await verified_human_agent_grants(context.user_id) + human_ids: Final = await verified_human_agent_grants(context.user_id, auth.team_id) return RestrictedAgentAccess(capped.intersection(human_ids)) -async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]: +async def verified_human_agent_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]: from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler if user_id is None: return frozenset() human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True) - sources: Final = await MCPRequestHandler.admitted_subject_sources(human) + sources: Final = await MCPRequestHandler.admitted_subject_sources( + human, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset() + ) human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources)) return frozenset().union(*(_granted_ids(access) for access in human_access)) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py new file mode 100644 index 00000000000..29acd77d78b --- /dev/null +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -0,0 +1,74 @@ +from typing import Final + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore +from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext + + +async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None: + delegation_verified: Final = auth._managed_delegation_verified + auth._managed_delegation_verified = False + if auth.agent_id is None: + return + if store is None: + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + + registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id) + if auth.managed_agent_context is not None or ( + registered is not None and (registered.identity_managed or registered.identity is not None) + ): + raise_identity_failure( + AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database") + ) + return + agent: Final = await store.agent(auth.agent_id) + if isinstance(agent, AgentIdentityFailure): + raise_identity_failure(agent) + if agent is None: + retired: Final = await store.retired_agent(auth.agent_id) + if isinstance(retired, AgentIdentityFailure): + raise_identity_failure(retired) + if auth.managed_agent_context is not None or retired: + raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists")) + return + if not agent.identity_managed: + return + if auth.jwt_claims and auth.managed_agent_context is None: + raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity")) + failure: Final = actor_admission_failure(agent, auth.managed_agent_context) + if failure is not None: + raise_identity_failure(failure) + auth.managed_agent_policy = agent + auth.billing_agent_policy = agent + auth.requires_fresh_policy = True + if ( + auth.managed_agent_context is not None + and auth.managed_agent_context.mode == "delegated" + and not delegation_verified + ): + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants + + grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id, auth.team_id) + if agent.agent_id not in grants: + raise_identity_failure( + AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent") + ) + + +def actor_admission_failure( + agent: AgentResponse, + context: ManagedAgentContext | None, +) -> AgentIdentityFailure | None: + if not agent.enabled or agent.identity is None or not agent.identity.active: + return AgentIdentityFailure(message="Agent execution is disabled") + if context is None: + return AgentIdentityFailure(message="This agent requires its bound identity provider token") + if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision: + return AgentIdentityFailure(message="Agent identity changed during authentication; retry") + if agent.execution_mode not in (context.mode, "both"): + return AgentIdentityFailure(message="Agent is not enabled for this execution mode") + if context.mode == "delegated" and not context.user_id: + return AgentIdentityFailure(message="A verified human subject is required") + return None diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 121f8d0993b..74993110964 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -14,6 +14,7 @@ import math import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence +from functools import partial from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias @@ -3738,6 +3739,8 @@ async def _fetch_key_object_from_db_with_reconnect( parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging | None, deadline_seconds: float | None = None, + *, + check_db_only: bool = False, ) -> BaseModel | None: """ Fetch key object from DB and retry once if a DB connection error can be healed. @@ -3751,6 +3754,7 @@ async def _fetch_key_object_from_db_with_reconnect( prisma_client=prisma_client, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ), name="key", deadline_seconds=deadline_seconds, @@ -3762,10 +3766,13 @@ async def _fetch_key_object_from_db_unbounded( prisma_client: PrismaClient, parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging | None, + *, + check_db_only: bool = False, ) -> BaseModel | None: + fetch: Final = partial(prisma_client.get_data, use_writer=True) if check_db_only else prisma_client.get_data async with db_lookup_gate.current(): try: - return await prisma_client.get_data( + return await fetch( token=hashed_token, table_name="combined_view", parent_otel_span=parent_otel_span, @@ -3787,7 +3794,7 @@ async def _fetch_key_object_from_db_unbounded( lock_timeout_seconds=auth_reconnect_lock_timeout, ) if did_reconnect: - return await prisma_client.get_data( + return await fetch( token=hashed_token, table_name="combined_view", parent_otel_span=parent_otel_span, @@ -3875,6 +3882,8 @@ async def get_key_object( parent_otel_span: Span | None = None, proxy_logging_obj: ProxyLogging | None = None, check_cache_only: bool | None = None, + *, + check_db_only: bool = False, ) -> UserAPIKeyAuth: """ - Check if team id in proxy Team Table @@ -3889,9 +3898,8 @@ async def get_key_object( # Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth # (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB. - user_api_key_auth: Final = await user_api_key_cache.async_get_cache( - key=key, - model_type=UserAPIKeyAuth, + user_api_key_auth: Final = ( + None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth) ) if user_api_key_auth is not None: return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) @@ -3905,6 +3913,7 @@ async def get_key_object( prisma_client=prisma_client, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) if _valid_token is None: @@ -3918,7 +3927,7 @@ async def get_key_object( _response: Final = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True)) # Load object_permission if object_permission_id exists but object_permission is not loaded - if _response.object_permission_id and not _response.object_permission: + if _response.object_permission_id and (check_db_only or not _response.object_permission): try: _response.object_permission = await get_object_permission( object_permission_id=_response.object_permission_id, @@ -3926,14 +3935,20 @@ async def get_key_object( user_api_key_cache=user_api_key_cache, parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, + check_db_only=check_db_only, ) except Exception as e: + if check_db_only: + raise verbose_proxy_logger.debug( "Failed to load object_permission for key with object_permission_id=%s: %s", _response.object_permission_id, e, ) + if check_db_only: + return _response + # save the key object to cache await _cache_key_object( hashed_token=hashed_token, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ed3ec7b4dde..36c6a4c476b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -3204,6 +3204,15 @@ async def _authorize_authenticated_request( # admin-only-route / model-access / budget checks) surface as # ProxyException consistently with pre-refactor behavior. try: + from litellm.proxy.agent_endpoints.auth.managed_authorization import admit_managed_actor + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_auth_obj.agent_id is not None: + await admit_managed_actor( + user_api_key_auth_obj, + AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None, + ) await _run_centralized_common_checks( user_api_key_auth_obj=user_api_key_auth_obj, request=request, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ea294b76e92..b75dbfe14bd 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -4194,6 +4194,8 @@ _PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5) async def _lookup_deprecated_key( db: PrismaWrapper | RoutingPrismaWrapper, hashed_token: str, + *, + check_db_only: bool = False, ) -> str | None: """ Check if a token exists in the deprecated keys table and is still within its grace period. @@ -4205,7 +4207,7 @@ async def _lookup_deprecated_key( now_ts: Final = now.timestamp() # Check cache first - cached: Final = _deprecated_key_cache.get(hashed_token) + cached: Final = None if check_db_only else _deprecated_key_cache.get(hashed_token) if cached is not None: active_token_id, cache_expires_at_ts, revoke_at_ts = cached if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts: @@ -4873,6 +4875,7 @@ class PrismaClient: proxy_logging_obj: ProxyLogging | None = None, budget_id_list: list[str] | None = None, check_deprecated: bool = True, + use_writer: bool = False, ): args_passed_in: Final = locals() start_time: Final = time.time() @@ -5171,12 +5174,20 @@ class PrismaClient: WHERE v.token = $1 """ - response = await self._query_first_with_cached_plan_fallback(sql_query, hashed_token) + response = ( + await self.writer_db.query_first(sql_query, hashed_token) + if use_writer + else await self._query_first_with_cached_plan_fallback(sql_query, hashed_token) + ) # If not found in main table, check deprecated keys (grace period) # check_deprecated=False on the recursive call prevents unbounded chaining if response is None and hashed_token is not None and check_deprecated: - active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token) + active_token_id: Final = await _lookup_deprecated_key( + db=self.writer_db if use_writer else self.db, + hashed_token=hashed_token, + check_db_only=use_writer, + ) if active_token_id: # The recursive call returns a finished # LiteLLM_VerificationTokenView; the dict @@ -5188,6 +5199,7 @@ class PrismaClient: parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, check_deprecated=False, + use_writer=use_writer, ) if deprecated_response is not None: verbose_proxy_logger.debug("Deprecated key used during grace period") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py index 40f6a2487fb..666d0c056c3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py @@ -275,6 +275,7 @@ async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) monkeypatch.setattr(proxy_server, "prisma_client", client) auth: Final = actor(agent_tools, delegated=True) + auth.team_id = "team" assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else []) assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else []) assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == [] @@ -365,3 +366,83 @@ async def test_manager_does_not_replace_managed_policy_failure_with_open_servers with pytest.raises(HTTPException) as failure: await manager.get_allowed_mcp_servers(actor(None, delegated=True)) assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_inline_tool_grant_admits_its_server_without_widening_tools() -> None: + auth: Final = actor(("read",)) + assert auth.managed_agent_policy is not None + auth.managed_agent_policy = auth.managed_agent_policy.model_copy( + update={"object_permission": {"object_permission_id": "tools", "mcp_tool_permissions": {"slack": ["read"]}}} + ) + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"] + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"] + assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selected_team", (None, "selected")) +@pytest.mark.parametrize("selected_grant", (False, True)) +async def test_delegation_never_borrows_another_teams_server_or_tools( + monkeypatch: pytest.MonkeyPatch, selected_team: str | None, selected_grant: bool +) -> None: + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + user: Final = LiteLLM_UserTable(user_id="human", teams=["selected", "other"], organization_memberships=[]) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="selected-grant", + mcp_servers=["slack"] if selected_grant else [], + mcp_tool_permissions={"slack": ["read"]} if selected_grant else {}, + ) + teams: Final = { + name: LiteLLM_TeamTable( + team_id=name, + models=[], + members_with_roles=[{"user_id": "human", "role": "user"}], + object_permission=permission if name == "selected" else LiteLLM_ObjectPermissionTable( + object_permission_id="other-grant", mcp_servers=["slack", "linear"] + ), + ) + for name in ("selected", "other") + } + + async def get_team(team_id: str, **kwargs: object) -> LiteLLM_TeamTable: + return teams[team_id] + + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user)) + monkeypatch.setattr(auth_checks, "get_team_object", get_team) + monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache()) + auth: Final = actor(None, delegated=True) + auth.team_id = selected_team + expected: Final = ["slack"] if selected_team and selected_grant else [] + assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == expected + assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if expected else []) + assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == [] + ordinary: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True) + assert set(await MCPRequestHandler.resolve_admitted_subject_servers(ordinary)) == {"slack", "linear"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("entitlement", ("group", "toolset")) +async def test_managed_mcp_rejects_unavailable_authoritative_entitlements( + monkeypatch: pytest.MonkeyPatch, entitlement: str +) -> None: + client: Final = MagicMock() + client.writer_db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable")) + client.writer_db.litellm_mcptoolsettable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable")) + monkeypatch.setattr(proxy_server, "prisma_client", client) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="entitlements", + mcp_access_groups=["group"] if entitlement == "group" else [], + mcp_toolsets=["toolset"] if entitlement == "toolset" else [], + ) + auth: Final = actor(None) + assert auth.managed_agent_policy is not None + auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"object_permission": permission.model_dump()}) + auth.requires_fresh_policy = True + with pytest.raises(HTTPException) as failure: + await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) + assert failure.value.status_code == 503 + client.db.litellm_mcpservertable.find_many.assert_not_called() + client.db.litellm_mcptoolsettable.find_many.assert_not_called() 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 e5dfeb369e2..d516aba862d 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 @@ -680,7 +680,7 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants( from litellm.proxy._types import LiteLLM_UserTable from litellm.proxy.auth import auth_checks from litellm.types.agents import AgentResponse - from litellm.types.proxy.agent_identity import ManagedAgentContext + from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext database: Final = MagicMock() monkeypatch.setattr(proxy_server, "prisma_client", database) @@ -688,7 +688,7 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants( human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"]) human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants) monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human)) - auth: Final = UserAPIKeyAuth(agent_id="actor") + auth: Final = UserAPIKeyAuth(agent_id="actor", api_key="verified-jwt") auth.managed_agent_policy = AgentResponse( agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump() ) @@ -698,6 +698,15 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants( access: Final = await AgentRequestHandler.resolve_agent_access(auth) assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"})) + target: Final = AgentResponse( + agent_id="shared", agent_name="Shared", agent_card_params={}, identity_managed=True, + identity=AgentIdentityBinding( + agent_id="shared", provider="microsoft_entra", tenant_id="tenant", client_id="client", issuer="issuer", revision="current" + ), + ) + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + assert await AgentRequestHandler.is_agent_allowed("shared", auth) is True + @pytest.mark.asyncio @pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"]) @@ -742,7 +751,7 @@ async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group) monkeypatch.setattr(proxy_server, "prisma_client", client) monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) - assert await verified_human_agent_grants("human") == frozenset({"target"}) + assert await verified_human_agent_grants("human", "team") == frozenset({"target"}) client.writer_db.litellm_usertable.find_unique.return_value = ( human.model_copy(update={"teams": []}) if revoked == "user" else human ) @@ -759,7 +768,7 @@ async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches client.writer_db.litellm_accessgrouptable.find_unique.return_value = ( group.model_copy(update={"access_agent_ids": []}) if grouped else group ) - assert await verified_human_agent_grants("human") == frozenset() + assert await verified_human_agent_grants("human", "team") == frozenset() client.db.litellm_usertable.find_unique.assert_not_called() client.db.litellm_teamtable.find_unique.assert_not_called() client.db.litellm_objectpermissiontable.find_unique.assert_not_called() @@ -885,3 +894,62 @@ async def test_delegation_without_a_verified_human_never_grants_agents(grant: bo auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated") assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset()) assert await verified_human_agent_grants(None) == frozenset() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("change", ("grant", "permission_reference", "groups", "team", "blocked", "expired", "deleted", "outage")) +async def test_managed_target_rechecks_authoritative_key_after_peer_revocation( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + from unittest.mock import MagicMock + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"]) + warm: Final = UserAPIKeyAuth(api_key="a" * 64, token="a" * 64, object_permission_id="grant", object_permission=permission) + target: Final = AgentResponse( + agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True, + identity=AgentIdentityBinding(agent_id="target", provider="microsoft_entra", tenant_id="tenant", client_id="client", issuer="issuer", revision="current"), + ) + client: Final = MagicMock() + client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + client.get_data = AsyncMock(return_value=warm) + client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) + cache: Final = UserApiKeyCache() + cache.set_cache("a" * 64, warm) + monkeypatch.setattr(proxy_server, "prisma_client", client) + monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) + assert await AgentRequestHandler.is_agent_allowed("target", warm) is True + client.get_data.return_value = warm.model_copy(update={ + "object_permission": None, + "object_permission_id": "replacement" if change == "permission_reference" else "grant", + "access_group_ids": [], + "team_id": "new-team" if change == "team" else None, + "blocked": change == "blocked", + "expires": "2000-01-01T00:00:00+00:00" if change == "expired" else None, + }) + client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(update={"agents": []}) + if change == "team": + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.auth import auth_checks + client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission + monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=LiteLLM_TeamTable( + team_id="new-team", object_permission=permission.model_copy(update={"agents": ["other"]}) + ))) + if change == "groups": + warm.object_permission = None + warm.access_group_ids = ["old-group"] + from litellm.proxy.auth import auth_checks + monkeypatch.setattr(auth_checks, "_get_agent_ids_from_access_groups", AsyncMock(return_value=["target"])) + if change == "deleted": + client.get_data.return_value = None + if change == "outage": + client.get_data.side_effect = RuntimeError("writer unavailable") + if change in ("blocked", "expired", "deleted", "outage"): + with pytest.raises((HTTPException, RuntimeError)): + await AgentRequestHandler.is_agent_allowed("target", warm) + else: + assert await AgentRequestHandler.is_agent_allowed("target", warm) is False 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 new file mode 100644 index 00000000000..7747e5eff71 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -0,0 +1,285 @@ +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.auth.managed_authorization import ( + actor_admission_failure, + admit_managed_actor, +) +from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore +from litellm.types.agents import AgentResponse +from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext + +BINDING: Final = AgentIdentityBinding( + agent_id="agent", + provider="microsoft_entra", + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + issuer="issuer", + revision="current", +) + + +def agent(**overrides: object) -> AgentResponse: + return AgentResponse.model_validate( + { + "agent_id": "agent", + "agent_name": "Agent", + "agent_card_params": {}, + "identity": BINDING, + "identity_managed": True, + "execution_mode": "both", + **overrides, + } + ) + + +@pytest.mark.parametrize( + "state", + [ + {"enabled": False}, + {"identity": None}, + {"identity": BINDING.model_copy(update={"active": False})}, + {"execution_mode": "delegated"}, + ], +) +def test_keys_cannot_bypass_lifecycle_or_delegated_only_mode(state: dict[str, object]) -> None: + assert isinstance(actor_admission_failure(agent(**state), None), AgentIdentityFailure) + + +@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"]) +def test_keys_cannot_impersonate_an_entra_bound_agent(mode: str) -> None: + assert isinstance(actor_admission_failure(agent(execution_mode=mode), None), AgentIdentityFailure) + + +@pytest.mark.parametrize( + "context", + [ + ManagedAgentContext(agent_id="agent", binding_revision="previous", mode="autonomous"), + ManagedAgentContext(agent_id="another", binding_revision="current", mode="autonomous"), + ManagedAgentContext(agent_id="agent", binding_revision="current", mode="delegated"), + ], +) +def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context: ManagedAgentContext) -> None: + assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure) + + +@pytest.mark.asyncio +async def test_deleted_agent_key_cannot_fall_back_to_unmanaged_authentication() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_retiredagent.find_unique = AsyncMock(return_value={"original_agent_id": "deleted"}) + with pytest.raises(HTTPException, match="Agent no longer exists"): + await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database)) + database.writer_db.litellm_retiredagent.find_unique.return_value = None + auth: Final = UserAPIKeyAuth(agent_id="legacy-attribution-label") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy is None + database.db.litellm_agentstable.find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None) + database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("unavailable")) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_agent_admission_database_outage_fails_closed() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(side_effect=RuntimeError("DB unavailable")) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_human_authentication_does_not_load_an_agent() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock() + await admit_managed_actor(UserAPIKeyAuth(user_id="human"), AgentIdentityStore.from_client(database)) + database.writer_db.litellm_agentstable.find_unique.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_disabled_agent_key_is_rejected_at_admission() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(enabled=False)) + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("permitted", [True, False]) +async def test_verified_human_still_needs_an_explicit_agent_invocation_grant( + monkeypatch: pytest.MonkeyPatch, + permitted: bool, +) -> None: + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable + from litellm.proxy.auth import auth_checks + + policy: Final = agent() + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + monkeypatch.setattr(proxy_server, "prisma_client", database) + permission: Final = LiteLLM_ObjectPermissionTable( + object_permission_id="human-grants", + agents=["agent"] if permitted else [], + ) + human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission) + monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human)) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.managed_agent_context = ManagedAgentContext( + agent_id="agent", + binding_revision="current", + mode="delegated", + user_id="human", + ) + if permitted: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy == policy + assert auth.billing_agent_policy == policy + else: + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert failure.value.status_code == 403 + + +def test_execution_mode_must_match_verified_token_mode() -> None: + context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context) + assert isinstance(failure, AgentIdentityFailure) + assert "execution mode" in failure.message + + +@pytest.mark.asyncio +async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(execution_mode="autonomous")) + auth: Final = UserAPIKeyAuth(agent_id="agent", jwt_claims={"agent": "agent", "sub": "unrelated-subject"}) + with pytest.raises(HTTPException) as denied: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert denied.value.status_code == 403 + assert auth.managed_agent_policy is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bound", [False, True]) +async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.MonkeyPatch, bound: bool) -> None: + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + + registry: Final = AgentRegistry() + registry.register_agent(agent(identity_managed=bound, identity=BINDING if bound else None)) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + auth: Final = UserAPIKeyAuth(agent_id="agent") + if not bound: + auth.managed_agent_context = ManagedAgentContext( + agent_id="agent", binding_revision="current", mode="autonomous" + ) + with pytest.raises(HTTPException) as denied: + await admit_managed_actor(auth, None) + assert denied.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> None: + policy: Final = agent(execution_mode="autonomous") + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + auth: Final = UserAPIKeyAuth(agent_id="agent", api_key="persisted-key") + with pytest.raises(HTTPException, match="bound identity provider token") as denied: + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert denied.value.status_code == 403 + assert auth.managed_agent_policy is None + assert auth.billing_agent_policy is None + + +@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")]) +def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None: + context: Final = ManagedAgentContext.model_validate( + {"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user} + ) + assert actor_admission_failure(agent(), context) is None + + +@pytest.mark.asyncio +async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None: + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent()) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.managed_agent_policy == agent() + assert auth.billing_agent_policy == agent() + assert auth.user_id is None + + +@pytest.mark.asyncio +async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_next_request() -> None: + """Managed MCP grants (toolsets, access groups) are read through the shared resolvers, which only + bypass the warm cache and the replica when the subject carries requires_fresh_policy""" + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent()) + auth: Final = UserAPIKeyAuth(agent_id="agent") + auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous") + assert auth.requires_fresh_policy is False + await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) + assert auth.requires_fresh_policy is True + + +async def test_jwt_delegation_verification_is_consumed_once_and_cannot_be_supplied_by_a_caller( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy.agent_endpoints.auth import agent_permission_handler + + policy: Final = agent() + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy) + store: Final = AgentIdentityStore.from_client(database) + grants: Final = AsyncMock(return_value=frozenset()) + monkeypatch.setattr(agent_permission_handler, "verified_human_agent_grants", grants) + auth: Final = UserAPIKeyAuth.model_validate({"agent_id": "agent", "_managed_delegation_verified": True}) + assert auth._managed_delegation_verified is False + auth.managed_agent_context = ManagedAgentContext( + agent_id="agent", binding_revision="current", mode="delegated", user_id="human" + ) + auth._managed_delegation_verified = True + assert "_managed_delegation_verified" not in auth.model_dump() + await admit_managed_actor(auth, store) + grants.assert_not_awaited() + assert auth._managed_delegation_verified is False + with pytest.raises(HTTPException) as failure: + await admit_managed_actor(auth, store) + assert failure.value.status_code == 403 + grants.assert_awaited_once_with("human", None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("database_available", (False, True)) +async def test_ordinary_agent_admission_preserves_legacy_authentication( + monkeypatch: pytest.MonkeyPatch, database_available: bool +) -> None: + from litellm.proxy.agent_endpoints import agent_registry + + registry: Final = agent_registry.AgentRegistry() + ordinary: Final = agent(identity_managed=False, identity=None) + registry.register_agent(ordinary) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=ordinary) + auth: Final = UserAPIKeyAuth(agent_id="agent") + await admit_managed_actor(auth, AgentIdentityStore.from_client(database) if database_available else None) + assert auth.agent_id == "agent" + assert auth.managed_agent_policy is None + assert auth.requires_fresh_policy is False diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ac40056c476..db50263a675 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -10160,3 +10160,43 @@ async def test_managed_agent_model_policy_checks_dispatched_model( with pytest.raises((HTTPException, ModelAccessDeniedProxyException)) as failure: await checks assert str(getattr(failure.value, "status_code", getattr(failure.value, "code", None))) == "403" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("reconnect", (False, True)) +async def test_authoritative_key_load_bypasses_warm_key_and_permission_caches(reconnect: bool) -> None: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + + permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="current", agents=["allowed"]) + stale: Final = UserAPIKeyAuth(token="hash", team_id="old-team", object_permission_id="old") + current: Final = UserAPIKeyAuth(token="hash", team_id="new-team", object_permission_id="current") + cache: Final = UserApiKeyCache() + cache.set_cache("hash", stale) + cache.set_cache(object_permission_cache_key("current"), permission.model_copy(update={"agents": ["revoked"]})) + database: Final = MagicMock() + database.get_data = AsyncMock(side_effect=[httpx.ConnectError("reset"), current] if reconnect else [current]) + database.attempt_db_reconnect = AsyncMock(return_value=True) + database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission) + fresh: Final = await get_key_object("hash", database, cache, check_db_only=True) + assert fresh.team_id == "new-team" + assert fresh.object_permission == permission + assert all(call.kwargs["use_writer"] is True for call in database.get_data.await_args_list) + database.db.litellm_objectpermissiontable.find_unique.assert_not_called() + cached: Final = await get_key_object("hash", database, cache) + assert cached.team_id == "old-team" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("missing", (False, True)) +async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailable(missing: bool) -> None: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + database: Final = MagicMock() + database.get_data = AsyncMock(return_value=UserAPIKeyAuth( + object_permission_id="grant", object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"]) + )) + database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=None, side_effect=None if missing else RuntimeError("writer unavailable") + ) + with pytest.raises(Exception, match=r"does not exist|unavailable"): + await get_key_object("hash", database, UserApiKeyCache(), check_db_only=True) 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 d1973e1693b..b2da7f30926 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 @@ -9381,3 +9381,32 @@ def test_identity_prefetch_keys_match_what_auth_reads_for_the_request(): assert _identity_cache_keys("sk-1234", end_user_id=None, key_is_resolved=True) == ( model_access_group_registry_cache_key(), ) + + +@pytest.mark.asyncio +async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monkeypatch: pytest.MonkeyPatch) -> None: + from typing import Final + from litellm.proxy import proxy_server + from litellm.proxy.auth import user_api_key_auth as auth_module + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + target: Final = AgentResponse( + agent_id="bound", agent_name="Bound", agent_card_params={}, identity_managed=True, + identity=AgentIdentityBinding( + agent_id="bound", provider="microsoft_entra", tenant_id="tenant", client_id="client", issuer="issuer", revision="current" + ), + ) + client: Final = MagicMock() + client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + monkeypatch.setattr(proxy_server, "prisma_client", client) + checks: Final = AsyncMock() + monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))) + data: Final = {"model": "allowed", "messages": [{"role": "user", "content": "hello"}]} + request: Final = _alias_request("/v1/chat/completions", data) + with pytest.raises(ProxyException): + await auth_module._authorize_authenticated_request( + UserAPIKeyAuth(agent_id="bound"), request, data, "/v1/chat/completions", "sk-test" + ) + checks.assert_not_awaited() diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py index 672dd1eb674..05c4f9d8a67 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -634,3 +634,28 @@ async def test_query_first_with_cached_plan_fallback_reports_the_reader_generati "reader_served_the_query": 2, "writer_served_the_query": 0, } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("rotated", (False, True)) +async def test_authoritative_combined_key_view_uses_writer_through_rotation( + prisma_client: PrismaClient, rotated: bool +) -> None: + writer: Final = MagicMock() + reader: Final = MagicMock() + active: Final = { + "token": "current-token", "team_id": "current-team", "team_models": None, + "team_blocked": None, "team_members_with_roles": None, "user_id": None, "expires": None, + } + writer.query_first = AsyncMock(side_effect=[None, active] if rotated else [active]) + reader.query_first = AsyncMock(return_value={**active, "team_id": "stale-team"}) + writer.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=SimpleNamespace( + active_token_id="current-token", revoke_at=datetime.now(timezone.utc) + timedelta(hours=1) + )) + prisma_client.db = RoutingPrismaWrapper(writer=writer, reader=reader) + response: Final = await prisma_client.get_data(token="original-token", table_name="combined_view", use_writer=True) + assert isinstance(response, LiteLLM_VerificationTokenView) + assert response.team_id == "current-team" + assert response.token == "current-token" + reader.query_first.assert_not_awaited() + assert writer.query_first.await_count == (2 if rotated else 1)