diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index be06d2e7321..6c9e1511b00 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3320,6 +3320,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # single-owner so its meaning stays trustworthy. mcp_session_resource_server_id: str | None = Field(default=None, exclude=True) mcp_toolset_id: str | None = Field(default=None, exclude=True) + authenticated_by_custom_auth: bool = Field(default=False, exclude=True) via_virtual_key: bool = Field( default=False, exclude=True, @@ -3381,6 +3382,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob values.pop("mcp_session_resource_server_id", None) values.pop("mcp_toolset_id", None) values.pop("via_virtual_key", None) + values.pop("authenticated_by_custom_auth", None) values.pop("agent_caller", None) values.pop("managed_agent_context", None) values.pop("managed_agent_policy", None) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index c88c6f2570a..6ddcd20d919 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -762,7 +762,9 @@ async def invoke_agent_a2a( body["metadata"] = {} body["metadata"]["agent_id"] = agent.agent_id body["metadata"]["model_group"] = f"a2a_agent/{agent_name}" - body["metadata"]["model_info"] = {"id": agent.agent_id} + body["metadata"]["model_info"] = { # mutable-ok: request hooks mutate metadata before JSON logging + "id": agent.agent_id + } body["agent_id"] = agent.agent_id body.update( diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index d8f6f0f4454..54cc26885ab 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -200,6 +200,7 @@ class AgentRequestHandler: if key_hash and managed_agent_policy(user_api_key_auth) is None and not user_api_key_auth.is_session_token + and not user_api_key_auth.authenticated_by_custom_auth else user_api_key_auth ) fresh_auth: Final = authority.model_copy( diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 69ea634f933..17d988127ec 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -1,7 +1,7 @@ from collections.abc import Mapping from itertools import product from types import MappingProxyType -from typing import Annotated, Final +from typing import Annotated, Final, Literal from pydantic import Field, TypeAdapter, ValidationError @@ -62,6 +62,22 @@ _MANAGED_MCP_ROUTES: Final = tuple( ) +_MODEL_ROUTE_KINDS: Final[ + Mapping[str, Literal["image_generation", "image_edit", "moderation", "speech", "body", "path"]] +] = MappingProxyType( + { + "/images/generations": "image_generation", + "/images/edits": "image_edit", + "/moderations": "moderation", + "/audio/transcriptions": "moderation", + "/audio/speech": "speech", + "/rerank": "body", + "/messages/count_tokens": "body", + ":countTokens": "path", + } +) + + def managed_agent_route_allowed(route: str, method: str | None) -> bool: from litellm.proxy.auth.route_checks import RouteChecks @@ -92,26 +108,12 @@ def managed_inference_request( raise_identity_failure( AgentIdentityFailure(message="Managed inference requires an explicit or configured model") ) - return {**body, "model": model} + return {**body, "model": model} # mutable-ok: centralized auth hooks add request tags and budget metadata if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS): - return dict(body) + return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model - kind: Final = ( - "image_generation" - if route.endswith("/images/generations") - else "image_edit" - if route.endswith("/images/edits") - else "moderation" - if route.endswith(("/moderations", "/audio/transcriptions")) - else "speech" - if route.endswith("/audio/speech") - else "body" - if route.endswith(("/rerank", "/messages/count_tokens")) - else "path" - if route.endswith(":countTokens") - else "completion" - ) + kind: Final = next((kind for suffix, kind in _MODEL_ROUTE_KINDS.items() if route.endswith(suffix)), "completion") endpoint_model: Final = path_model or ( query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None ) @@ -120,7 +122,7 @@ def managed_inference_request( raise_identity_failure( AgentIdentityFailure(message="Managed inference requires an explicit or configured model") ) - return {**body, "model": effective} + return {**body, "model": effective} # mutable-ok: centralized auth hooks add request tags and budget metadata def managed_agent_policy(auth: "UserAPIKeyAuth | None") -> AgentResponse | None: @@ -204,12 +206,12 @@ _INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan def invocation_target(route: str, body: Mapping[str, object]) -> str | None: - model: Final = body.get("model") - if isinstance(model, str) and model.startswith("a2a/"): - return model.removeprefix("a2a/") or None components: Final = tuple(route.strip("/").split("/")) path: Final = components[1:] if components and components[0] == "v1" else components - return path[1] if len(path) >= 2 and path[0] == "a2a" else None + if len(path) >= 2 and path[0] == "a2a": + return path[1] or None + model: Final = body.get("model") + return model.removeprefix("a2a/") or None if isinstance(model, str) and model.startswith("a2a/") else None async def prepare_agent_invocation( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 03f8e763041..4448d860217 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -2965,7 +2965,7 @@ class JWTAuthManager: ), ) auth.managed_agent_context = result.get("managed_agent_context") - auth._managed_delegation_verified = ( + auth._managed_delegation_verified = ( # pyright: ignore[reportPrivateUsage] # JWT admission produces the one-shot proof consumed by managed authorization auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated" ) return auth diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 8595f1b3120..5194f62cf78 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1561,6 +1561,7 @@ async def _user_api_key_auth_builder( route=route, parent_otel_span=parent_otel_span, ) + validated.authenticated_by_custom_auth = True return validated elif response is not None and isinstance(response, str): api_key = response @@ -1576,6 +1577,7 @@ async def _user_api_key_auth_builder( route=route, parent_otel_span=parent_otel_span, ) + validated.authenticated_by_custom_auth = True return validated ### LITELLM-DEFINED AUTH FUNCTION ### diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index c995b5aedd4..2a22077eb99 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2313,6 +2313,11 @@ async def _complete_cli_sso_callback_session( status_code=500, detail="Could not resolve team model grants for this login. Please try again", ) + from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject + + await enroll_microsoft_subject( + request.scope.get("litellm_microsoft_interactive_subject"), user_info.user_id, prisma_client + ) resolved_teams: Final = _cli_sso_session_teams(team_details) attribution_metadata: Final = build_cli_sso_attribution_metadata(result=result) if attribution_metadata: diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py index f137f2713f9..eee985f0aca 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth from litellm.proxy.agent_endpoints.auth.managed_authorization import ( actor_admission_failure, admit_managed_actor, @@ -78,6 +78,7 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None: { "managed_agent_context": context, "requires_fresh_policy": True, + "authenticated_by_custom_auth": True, "mcp_explicit_grants_only": True, "managed_agent_policy": agent(), "billing_agent_policy": agent(), @@ -86,6 +87,8 @@ def test_caller_cannot_construct_trusted_subject_or_policy() -> None: } ) assert auth.requires_fresh_policy is False + assert auth.authenticated_by_custom_auth is False + assert "authenticated_by_custom_auth" not in auth.model_dump() assert auth.mcp_explicit_grants_only is False assert "mcp_explicit_grants_only" not in auth.model_dump() assert auth.managed_agent_context is None @@ -161,6 +164,9 @@ async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None: "route,body,expected", [ ("/a2a/agent", {}, "agent"), + ("/a2a/expensive", {"model": "a2a/cheap"}, "expensive"), + ("/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"), + ("/v1/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"), ("/v1/a2a/agent/", {}, "agent"), ("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"), ("/v1/chat/completions", {"model": "a2a/"}, None), @@ -496,6 +502,8 @@ async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_ 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 + assert auth.authenticated_by_custom_auth is False + assert "authenticated_by_custom_auth" not in auth.model_dump() await admit_managed_actor(auth, AgentIdentityStore.from_client(database)) assert auth.requires_fresh_policy is True @@ -545,3 +553,27 @@ async def test_ordinary_agent_admission_preserves_legacy_authentication( assert auth.agent_id == "agent" assert auth.managed_agent_policy is None assert auth.requires_fresh_policy is False + assert auth.authenticated_by_custom_auth is False + assert "authenticated_by_custom_auth" not in auth.model_dump() + + +@pytest.mark.parametrize( + "route", + tuple(dict.fromkeys( + LiteLLMRoutes.openai_routes.value + + LiteLLMRoutes.anthropic_routes.value + + LiteLLMRoutes.google_routes.value + )), +) +def test_registered_inference_routes_have_an_explicit_managed_access_decision(route: str) -> None: + from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed + + normalized: Final = route.removeprefix("/openai").removeprefix("/v1beta").removeprefix("/v1") + unsupported: Final = normalized.startswith(( + "/videos", "/batches", "/files", "/fine_tuning", "/assistants", "/threads", "/utils/", + "/vector_stores", "/vector_store/", "/search", "/containers", "/skills", "/claude-code/", + "/interactions", "/agents", "/responses/{", "/responses/input_tokens", + "/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions", + )) or normalized in ("/models", "/cursor/models", "/cursor/v1/models") + concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model") + assert managed_agent_route_allowed(concrete, None) is not unsupported, route diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 9be982483ac..df576fc27d0 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -8055,7 +8055,10 @@ def test_managed_issuer_requires_configured_audience_validation( @pytest.mark.asyncio -async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(monkeypatch: pytest.MonkeyPatch) -> None: +@pytest.mark.parametrize("authentication_write", ["success", "revoked", "unavailable"]) +async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy( + monkeypatch: pytest.MonkeyPatch, authentication_write: str +) -> None: from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.types.proxy.agent_identity import AgentIdentityBinding @@ -8093,6 +8096,16 @@ async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(mo database.writer_db.litellm_agentidentity.find_unique.assert_awaited_once() assert database.writer_db.litellm_agentstable.find_unique.await_count == 2 assert database.writer_db.litellm_agentidentity.update_many.await_count == 2 + if authentication_write != "success": + database.writer_db.litellm_agentidentity.update_many.return_value = 0 + database.writer_db.litellm_agentidentity.update_many.side_effect = ( + RuntimeError("storage unavailable") if authentication_write == "unavailable" else None + ) + with pytest.raises(HTTPException) as failed_write: + await JWTAuthManager.authorize_jwt(**arguments) + assert failed_write.value.status_code == (503 if authentication_write == "unavailable" else 403) + assert database.writer_db.litellm_agentidentity.update_many.await_count == 3 + return database.writer_db.litellm_agentstable.find_unique.return_value = agent.model_copy(update={"enabled": False}) with pytest.raises(HTTPException) as denied: await JWTAuthManager.authorize_jwt(**arguments) 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 c106b64f587..ef6832ef77b 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 @@ -9453,16 +9453,18 @@ async def test_managed_actor_cannot_access_provider_resource_routes(monkeypatch, request.scope["method"] = "GET" from litellm.types.proxy.agent_identity import ManagedAgentContext - auth = UserAPIKeyAuth( - agent_id="managed", api_key="persisted-key", models=["test-model"], - managed_agent_context=( - ManagedAgentContext(agent_id="managed", binding_revision="revision", mode="autonomous") - if verified_identity else None - ), - ) + auth = UserAPIKeyAuth(agent_id="managed", api_key="persisted-key", models=["test-model"]) + if verified_identity: + auth.managed_agent_context = ManagedAgentContext( + agent_id="managed", binding_revision="revision", mode="autonomous" + ) with pytest.raises(ProxyException) as denied: await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key") assert denied.value.code == "403" + if verified_identity: + assert denied.value.message == "Agent identities can only access inference and agent discovery routes" + else: + assert denied.value.message == "This agent requires its bound identity provider token" @pytest.mark.asyncio @@ -9613,3 +9615,76 @@ async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monk UserAPIKeyAuth(agent_id="bound"), request, data, "/v1/chat/completions", "sk-test" ) checks.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("enterprise", [False, True]) +@pytest.mark.parametrize("credential", ["custom-credential", "sk-custom-credential"]) +@pytest.mark.parametrize("granted", [False, True]) +async def test_custom_auth_grants_reach_managed_targets_without_a_virtual_key_row( + monkeypatch: pytest.MonkeyPatch, enterprise: bool, credential: str, granted: bool +) -> None: + import importlib + from typing import Final + + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler + from litellm.types.agents import AgentResponse + from litellm.types.proxy.agent_identity import AgentIdentityBinding + + 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", + ), + ) + registry: Final = AgentRegistry() + registry.register_agent(target) + trusted: Final = UserAPIKeyAuth( + api_key=credential, object_permission={"object_permission_id": "custom", "agents": ["target"] if granted else ["other"]} + ) + custom: Final = AsyncMock(return_value=trusted) + database: Final = MagicMock() + database.get_data = AsyncMock(return_value=None) + database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target) + for name, value in { + **_proxy_server_attrs_for_custom_auth(user_custom_auth=None if enterprise else custom), + "prisma_client": database, + }.items(): + monkeypatch.setattr(proxy_server, name, value) + module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setattr(module, "enterprise_custom_auth", custom if enterprise else None) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False) + admitted: Final = await _user_api_key_auth_builder( + request=_alias_request("/a2a/target/message/send", {}), api_key=f"Bearer {credential}", + azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, + azure_apim_header=None, request_data={}, + ) + assert await AgentRequestHandler.is_agent_allowed("target", admitted) is granted + custom.assert_awaited_once() + database.get_data.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(monkeypatch: pytest.MonkeyPatch) -> None: + import importlib + from typing import Final + + from litellm.proxy import proxy_server + + custom: Final = AsyncMock(return_value="sk-master-key") + for name, value in _proxy_server_attrs_for_custom_auth(user_custom_auth=custom).items(): + monkeypatch.setattr(proxy_server, name, value) + module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth") + monkeypatch.setattr(module, "enterprise_custom_auth", custom) + admitted: Final = await _user_api_key_auth_builder( + request=_alias_request("/v1/chat/completions", {}), api_key="Bearer external-credential", + azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None, + azure_apim_header=None, request_data={}, + ) + assert admitted.authenticated_by_custom_auth is False + assert admitted.via_virtual_key is True diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 81c48868b9a..b546b9eb965 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -7656,3 +7656,53 @@ async def test_managed_invocations_enforce_actor_and_target_rate_policies( for scope in stash.reserved_scopes: if scope[0] in ("agent", "agent_session"): assert increments[handler.create_rate_limit_keys(*scope, "tokens")] == 3 - stash.reserved_tokens + + +@pytest.mark.parametrize("route", ["/a2a/expensive", "/a2a/expensive/message/send", "/v1/a2a/expensive/message/send"]) +async def test_a2a_url_target_owns_invocation_fee_and_request_limit( + monkeypatch: pytest.MonkeyPatch, route: str +) -> None: + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy import proxy_server + from litellm.proxy.agent_endpoints import agent_registry + from litellm.proxy.agent_endpoints.auth.managed_authorization import invocation_target, prepare_agent_invocation + from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore + from litellm.types.agents import AgentResponse + + expensive: Final = AgentResponse( + agent_id="expensive", agent_name="Expensive", agent_card_params={}, rpm_limit=1, + litellm_params={"cost_per_query": 0.25}, + ) + cheap: Final = AgentResponse( + agent_id="cheap", agent_name="Cheap", agent_card_params={}, rpm_limit=100, + litellm_params={"cost_per_query": 0.01}, + ) + registry: Final = agent_registry.AgentRegistry() + registry.register_agent(expensive) + registry.register_agent(cheap) + monkeypatch.setattr(agent_registry, "global_agent_registry", registry) + database: Final = MagicMock() + database.writer_db.litellm_agentstable.find_unique = AsyncMock( + side_effect=lambda where, include: {"expensive": expensive, "cheap": cheap}[where["agent_id"]] + ) + monkeypatch.setattr(proxy_server, "prisma_client", database) + auth: Final = UserAPIKeyAuth(agent_id="caller") + auth.managed_agent_policy = AgentResponse( + agent_id="caller", agent_name="Caller", agent_card_params={}, + object_permission={"object_permission_id": "both-targets", "agents": ["expensive", "cheap"]}, + ) + body: Final = {"model": "a2a/cheap"} + target: Final = invocation_target(route, body) + assert target is not None + await prepare_agent_invocation(auth, target, AgentIdentityStore.from_client(database)) + assert auth.invoked_agent_id == "expensive" + assert auth.invoked_agent_policy == expensive + assert auth.agent_invocation_cost == pytest.approx(0.25) + cache: Final = DualCache() + limiter: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache)) + await _rpm_request(limiter, cache, auth, "a2a/cheap") + with pytest.raises(HTTPException) as denied: + await _rpm_request(limiter, cache, auth, "a2a/cheap") + assert denied.value.status_code == 429 + assert "expensive" in str(denied.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 971019160f9..9cdf5e9d6ff 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -2996,6 +2996,7 @@ class TestCLIKeyRegenerationFlow: from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "https://proxy.example.com/" mock_user_info = LiteLLM_UserTable( @@ -3159,6 +3160,7 @@ class TestCLIKeyRegenerationFlow: # Mock request mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://internal-proxy.local/" # Test data @@ -7107,6 +7109,7 @@ class TestCliSsoAttributionMetadata: from litellm.proxy.management_endpoints.types import CustomOpenID mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://internal-proxy.local/" session_key = "cli-session-new-user" mock_user_info = LiteLLM_UserTable( @@ -7221,6 +7224,7 @@ class TestCliSsoAttributionMetadata: ) mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://internal-proxy.local/" session_key = "cli-session-4567890" mock_user_info = LiteLLM_UserTable( @@ -8824,6 +8828,7 @@ async def test_cli_completion_persists_assertion_under_db_user_id(): assertion = assertion_from_sso_login(_ema_id_token(), None) assert assertion is not None mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://localhost:4000/" user_info = MagicMock() @@ -9062,6 +9067,7 @@ async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog): monkeypatch.setenv("MICROSOFT_CLIENT_ID", "cid") mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://localhost:4000/" user_info = MagicMock() @@ -9137,6 +9143,7 @@ def _cli_callback_kwargs(flow): def _cli_callback_request(): mock_request = MagicMock(spec=Request) + mock_request.scope = {} mock_request.base_url = "http://localhost:4000/" return mock_request @@ -9441,3 +9448,45 @@ class TestSessionTokenCookie: resp = Response() set_session_token_cookie(resp, _make_http_request(), "jwt-token-value") assert "Secure" in self._cookie(resp) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("trusted", [False, True]) +@pytest.mark.parametrize("storage_available", [False, True]) +async def test_cli_sign_in_enrolls_only_verified_subjects_before_completing( + monkeypatch: pytest.MonkeyPatch, trusted: bool, storage_available: bool +) -> None: + from typing import Final + + from litellm.proxy.management_endpoints import ui_sso + from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject + + flow: Final[dict[str, object]] = {} + kwargs: Final = _cli_callback_kwargs(flow) + subject: Final = MicrosoftInteractiveSubject(issuer="issuer", tenant_id="tenant", oid="subject") + kwargs["request"].scope = {"litellm_microsoft_interactive_subject": subject if trusted else subject.model_dump()} + table: Final = kwargs["prisma_client"].writer_db.litellm_verifiedsubject + table.upsert = AsyncMock( + return_value=SimpleNamespace(kind="human", user_id="cli-user-id", verified_via="sso_interactive"), + side_effect=None if storage_available else RuntimeError("storage unavailable"), + ) + monkeypatch.setattr(ui_sso, "get_user_info_from_db", AsyncMock(return_value=_cli_callback_user_info([]))) + monkeypatch.setattr(ui_sso, "fetch_cli_sso_team_details", AsyncMock(return_value=())) + monkeypatch.setattr(ui_sso, "retain_sso_identity_assertion_for_ema", AsyncMock()) + if trusted and not storage_available: + with pytest.raises(HTTPException) as error: + await ui_sso._complete_cli_sso_callback_session(**kwargs) + assert error.value.status_code == 503 + assert "sso_complete" not in flow + return + response: Final = await ui_sso._complete_cli_sso_callback_session(**kwargs) + assert response.status_code == 200 + assert flow["session_data"]["user_id"] == "cli-user-id" + if trusted: + table.upsert.assert_awaited_once_with( + where={"issuer_tenant_id_oid": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject"}}, + data={"create": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject", + "user_id": "cli-user-id", "verified_via": "sso_interactive"}, "update": {}}, + ) + else: + table.upsert.assert_not_awaited()