mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): enforce target policy and preserve trusted authentication
This commit is contained in:
parent
30c70adb64
commit
0fd4e93c48
12 changed files with 267 additions and 34 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue