mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
feat(agents): authenticate Entra actors and retain request attribution
This commit is contained in:
parent
97ddb52e0d
commit
b02dfccef4
24 changed files with 862 additions and 156 deletions
|
|
@ -170,6 +170,11 @@ async def identity_from_subject_token(
|
|||
return _refusal_for(denied, denied.message)
|
||||
except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures
|
||||
return _refusal_for(denied, denied)
|
||||
if result.get("agent_id") is not None:
|
||||
return SubjectTokenRefusal(
|
||||
error="invalid_request",
|
||||
description="Agent tokens require direct JWT authentication; this exchange supports users only",
|
||||
)
|
||||
user_id: Final = result["user_id"]
|
||||
if user_id is None:
|
||||
return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from dataclasses import dataclass
|
|||
from datetime import datetime
|
||||
from traceback import walk_tb
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
|
|
@ -14,6 +14,7 @@ import httpx2
|
|||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
from starlette.datastructures import Headers
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
|
||||
|
|
@ -62,7 +63,27 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
|
||||
|
||||
|
||||
class _MCPModelMetadata(TypedDict):
|
||||
model_group: ReadOnly[str]
|
||||
|
||||
|
||||
def _stamp_mcp_tool_metadata(logging_obj: "LiteLLMLoggingObj | None", server_id: str, tool_name: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
if logging_obj is None:
|
||||
return
|
||||
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
server_id
|
||||
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
|
||||
metadata: Final[StandardLoggingMCPToolCall] = {
|
||||
"name": tool_name,
|
||||
"mcp_server_name": server.name if server is not None else server_id,
|
||||
}
|
||||
logging_obj.model_call_details["mcp_tool_call_metadata"] = metadata
|
||||
|
||||
|
||||
MCP_AVAILABLE: bool = True
|
||||
try:
|
||||
|
|
@ -1143,6 +1164,12 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
data["model"] = f"MCP: {tool_name}"
|
||||
model_metadata: Final[_MCPModelMetadata] = {
|
||||
**(data.get("metadata") or MappingProxyType({})),
|
||||
"model_group": f"MCP: {tool_name}",
|
||||
}
|
||||
data["metadata"] = model_metadata
|
||||
proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
_request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below
|
||||
try:
|
||||
|
|
@ -1176,6 +1203,8 @@ if MCP_AVAILABLE:
|
|||
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
|
||||
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
|
||||
|
||||
_stamp_mcp_tool_metadata(logging_obj, server_id, tool_name)
|
||||
|
||||
# Resolve allowed MCP servers with IP filtering
|
||||
(
|
||||
allowed_mcp_servers,
|
||||
|
|
|
|||
|
|
@ -4075,6 +4075,11 @@ class SpendLogsRouterMetadata(TypedDict):
|
|||
|
||||
|
||||
class SpendLogsMetadata(TypedDict):
|
||||
actor_agent_id: ReadOnly[NotRequired[str | None]]
|
||||
target_agent_id: ReadOnly[NotRequired[str | None]]
|
||||
billing_agent_id: ReadOnly[NotRequired[str | None]]
|
||||
agent_execution_mode: ReadOnly[NotRequired[str | None]]
|
||||
verified_human_user_id: ReadOnly[NotRequired[str | None]]
|
||||
autorouter_baseline_observation: ReadOnly[str | None]
|
||||
"""
|
||||
Specific metadata k,v pairs logged to spendlogs for easier cost tracking
|
||||
|
|
@ -4138,6 +4143,7 @@ class SpendLogsPayload(TypedDict):
|
|||
model_id: str | None
|
||||
model_group: str | None
|
||||
mcp_namespaced_tool_name: str | None
|
||||
billing_agent_id: ReadOnly[NotRequired[str | None]]
|
||||
agent_id: str | None
|
||||
api_base: str
|
||||
user: str
|
||||
|
|
@ -5060,6 +5066,7 @@ class JWTAuthBuilderResult(TypedDict):
|
|||
org_id: str | None
|
||||
team_membership: LiteLLM_TeamMembership | None
|
||||
jwt_claims: dict # Decoded JWT token claims (avoids re-decoding)
|
||||
managed_agent_context: ReadOnly[NotRequired[ManagedAgentContext | None]]
|
||||
agent_id: ReadOnly[str | None]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -723,6 +723,8 @@ async def invoke_agent_a2a(
|
|||
detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.",
|
||||
)
|
||||
|
||||
user_api_key_dict.invoked_agent_id = agent.agent_id
|
||||
|
||||
_enforce_inbound_trace_id(agent, request)
|
||||
|
||||
# Get backend URL and agent name
|
||||
|
|
@ -760,6 +762,8 @@ async def invoke_agent_a2a(
|
|||
if "metadata" not in body:
|
||||
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["agent_id"] = agent.agent_id
|
||||
|
||||
body.update(
|
||||
|
|
@ -863,6 +867,7 @@ async def invoke_agent_a2a(
|
|||
# results written by the unified_guardrail hook are captured.
|
||||
logging_obj._defer_async_logging = True
|
||||
response = await asend_message(
|
||||
model=f"a2a_agent/{agent_name}",
|
||||
request=a2a_request,
|
||||
api_base=agent_url,
|
||||
litellm_params=litellm_params,
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ async def route_a2a_agent_request(
|
|||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
)
|
||||
if not is_admin:
|
||||
if not is_admin or agent.identity_managed:
|
||||
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
|
||||
agent_id=agent.agent_id,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
|
|
|
|||
17
litellm/proxy/agent_endpoints/identity.py
Normal file
17
litellm/proxy/agent_endpoints/identity.py
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
LEGACY_IDENTITY_MESSAGE: Final = (
|
||||
"litellm_params.identity is not supported: bind an Entra application through the top-level identity field"
|
||||
)
|
||||
|
||||
|
||||
def has_legacy_identity(params: Mapping[str, object] | None) -> bool:
|
||||
return params is not None and "identity" in params
|
||||
|
||||
|
||||
def reject_legacy_identity(params: Mapping[str, object] | None) -> None:
|
||||
if has_legacy_identity(params):
|
||||
raise HTTPException(400, LEGACY_IDENTITY_MESSAGE)
|
||||
|
|
@ -52,6 +52,10 @@ from litellm.proxy._types import (
|
|||
TeamMemberAddRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
|
||||
from litellm.proxy.agent_endpoints.identity import has_legacy_identity
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||
from litellm.proxy.auth.model_access_denied import (
|
||||
ModelAccessDeniedHTTPException,
|
||||
|
|
@ -67,6 +71,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
|
||||
from .auth_checks import (
|
||||
|
|
@ -157,6 +162,8 @@ class HeaderTeam:
|
|||
class AgentLookup(Protocol):
|
||||
"""The registered-agent lookups a JWT agent claim is matched against."""
|
||||
|
||||
def get_agent_list(self) -> Sequence[AgentResponse]: ...
|
||||
|
||||
def get_agent_by_id(self, agent_id: str) -> AgentResponse | None:
|
||||
"""The agent registered under ``agent_id``, if any."""
|
||||
|
||||
|
|
@ -167,6 +174,9 @@ class AgentLookup(Protocol):
|
|||
class _NoRegisteredAgents:
|
||||
"""The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches."""
|
||||
|
||||
def get_agent_list(self) -> tuple[AgentResponse, ...]:
|
||||
return ()
|
||||
|
||||
def get_agent_by_id(self, agent_id: str) -> None:
|
||||
return None
|
||||
|
||||
|
|
@ -1096,6 +1106,15 @@ class JWTHandler:
|
|||
"options": options or None,
|
||||
}
|
||||
|
||||
def managed_issuer_is_trusted(self, issuer: object) -> bool:
|
||||
if not isinstance(issuer, str):
|
||||
return False
|
||||
configured: Final = self.litellm_jwtauth.issuers or ()
|
||||
for item in configured:
|
||||
if item.issuer == issuer:
|
||||
return bool(item.audience) and not item.disable_audience_validation
|
||||
return issuer == os.getenv("JWT_ISSUER") and bool(os.getenv("JWT_AUDIENCE"))
|
||||
|
||||
def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
|
||||
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
|
||||
if litellm_jwtauth is None:
|
||||
|
|
@ -1488,7 +1507,12 @@ class JWTAuthManager:
|
|||
agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
|
||||
agent_name=agent_claim
|
||||
)
|
||||
if agent is None:
|
||||
if (
|
||||
agent is None
|
||||
or agent.identity_managed
|
||||
or agent.identity is not None
|
||||
or has_legacy_identity(agent.litellm_params)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
|
||||
|
|
@ -2478,12 +2502,39 @@ class JWTAuthManager:
|
|||
"""Resolve and authorize JWT context; only normal admission supplies provisioning."""
|
||||
handler: Final = jwt_handler
|
||||
jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
|
||||
managed: Final = await resolve_managed_agent(jwt_valid_token, prisma_client)
|
||||
if managed is not None:
|
||||
if not handler.managed_issuer_is_trusted(jwt_valid_token.get("iss")):
|
||||
raise HTTPException(403, "Managed agents require trusted JWT issuer and audience validation")
|
||||
if not managed_agent_route_allowed(route, request_method):
|
||||
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
|
||||
evidence: Final = await AgentIdentityStore.from_client(prisma_client).record_authentication(managed)
|
||||
if isinstance(evidence, AgentIdentityFailure):
|
||||
raise_identity_failure(evidence)
|
||||
if managed.mode == "autonomous":
|
||||
return JWTAuthBuilderResult(
|
||||
is_proxy_admin=False,
|
||||
team_id=None,
|
||||
team_object=None,
|
||||
user_id=None,
|
||||
user_email=None,
|
||||
user_object=None,
|
||||
org_id=None,
|
||||
org_object=None,
|
||||
end_user_id=None,
|
||||
end_user_object=None,
|
||||
token=api_key,
|
||||
team_membership=None,
|
||||
jwt_claims=jwt_valid_token,
|
||||
agent_id=managed.agent_id,
|
||||
managed_agent_context=managed,
|
||||
)
|
||||
team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
|
||||
model: Final = request_data.get("model")
|
||||
requested_model: Final = model if isinstance(model, str) else None
|
||||
|
||||
# Check RBAC
|
||||
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
|
||||
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if managed is None else None
|
||||
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
|
||||
|
||||
# Check Scope Based Access
|
||||
|
|
@ -2499,7 +2550,11 @@ class JWTAuthManager:
|
|||
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
|
||||
|
||||
# Get basic user info
|
||||
user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
|
||||
user_id, user_email, valid_user_email = (
|
||||
(managed.user_id, None, None)
|
||||
if managed is not None
|
||||
else await JWTAuthManager.get_user_info(handler, jwt_valid_token)
|
||||
)
|
||||
|
||||
# Get IDs
|
||||
org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
|
||||
|
|
@ -2514,23 +2569,31 @@ class JWTAuthManager:
|
|||
elif rbac_role == LitellmUserRoles.INTERNAL_USER:
|
||||
user_id = object_id
|
||||
|
||||
agent_id: Final = JWTAuthManager.resolve_agent_id(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
agent_registry=handler.agent_lookup,
|
||||
agent_id: Final = (
|
||||
managed.agent_id
|
||||
if managed is not None
|
||||
else JWTAuthManager.resolve_agent_id(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
agent_registry=handler.agent_lookup,
|
||||
)
|
||||
)
|
||||
|
||||
# Check admin access
|
||||
admin_result: Final = await JWTAuthManager.check_admin_access(
|
||||
handler,
|
||||
scopes,
|
||||
route,
|
||||
user_id,
|
||||
org_id,
|
||||
api_key,
|
||||
jwt_valid_token,
|
||||
user_email=user_email,
|
||||
agent_id=agent_id,
|
||||
admin_result: Final = (
|
||||
None
|
||||
if managed is not None
|
||||
else await JWTAuthManager.check_admin_access(
|
||||
handler,
|
||||
scopes,
|
||||
route,
|
||||
user_id,
|
||||
org_id,
|
||||
api_key,
|
||||
jwt_valid_token,
|
||||
user_email=user_email,
|
||||
agent_id=agent_id,
|
||||
)
|
||||
)
|
||||
if admin_result:
|
||||
await JWTAuthManager._attach_team_from_header_for_admin(
|
||||
|
|
@ -2705,13 +2768,13 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
org_alias=org_alias,
|
||||
user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
|
||||
user_id_upsert=provisioning.user_id_upsert if provisioning is not None and managed is None else False,
|
||||
)
|
||||
|
||||
# Derive org_id from org_object if resolved by alias
|
||||
resolved_org_id: Final = org_object.organization_id if org_object else org_id
|
||||
|
||||
if provisioning is not None:
|
||||
if provisioning is not None and managed is None:
|
||||
await JWTAuthManager.sync_user_role_and_teams(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
|
|
@ -2784,7 +2847,7 @@ class JWTAuthManager:
|
|||
)
|
||||
|
||||
## MAP USER TO TEAMS
|
||||
if provisioning is not None:
|
||||
if provisioning is not None and managed is None:
|
||||
await JWTAuthManager.map_user_to_teams(
|
||||
user_object=user_object,
|
||||
team_object=team_object,
|
||||
|
|
@ -2799,7 +2862,9 @@ class JWTAuthManager:
|
|||
)
|
||||
|
||||
# check if user is proxy admin
|
||||
is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN)
|
||||
is_proxy_admin: Final = managed is None and bool(
|
||||
user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
return JWTAuthBuilderResult(
|
||||
is_proxy_admin=is_proxy_admin,
|
||||
|
|
@ -2816,6 +2881,7 @@ class JWTAuthManager:
|
|||
team_membership=team_membership_object,
|
||||
jwt_claims=jwt_valid_token,
|
||||
agent_id=agent_id,
|
||||
managed_agent_context=managed,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2826,11 +2892,13 @@ class JWTAuthManager:
|
|||
"""Keep JWT identity and permission attribution identical across consumers."""
|
||||
user: Final = result["user_object"]
|
||||
admin: Final = result["is_proxy_admin"]
|
||||
return UserAPIKeyAuth(
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=(
|
||||
LitellmUserRoles.PROXY_ADMIN
|
||||
if admin
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
if result.get("managed_agent_context") is not None
|
||||
else LitellmUserRoles(user.user_role)
|
||||
if user is not None and user.user_role is not None
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
|
|
@ -2852,3 +2920,5 @@ class JWTAuthManager:
|
|||
user_id=result["user_id"],
|
||||
),
|
||||
)
|
||||
auth.managed_agent_context = result.get("managed_agent_context")
|
||||
return auth
|
||||
|
|
|
|||
|
|
@ -652,6 +652,8 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
|
|||
# never reaches the fallback.
|
||||
synthetic_scope: Final[dict[str, Any]] = {
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"query_string": ws_scope.get("query_string", b""),
|
||||
"headers": scope_headers,
|
||||
"path": ws_scope.get("path", ""),
|
||||
"state": ws_scope.setdefault("state", {}), # mutable-ok: Starlette's socket state, shared with the request
|
||||
|
|
@ -1653,6 +1655,13 @@ async def _user_api_key_auth_builder(
|
|||
else:
|
||||
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
|
||||
|
||||
from litellm.proxy.agent_endpoints.identity_store import resolve_managed_agent
|
||||
|
||||
if jwt_claims and await resolve_managed_agent(jwt_claims, prisma_client) is not None:
|
||||
raise HTTPException(
|
||||
403, "Managed agents require direct JWT authentication without virtual-key mapping"
|
||||
)
|
||||
|
||||
resolve_result: Final = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
@ -3123,7 +3132,10 @@ async def _reserve_budget_after_common_checks(
|
|||
end_user_id=end_user_id,
|
||||
end_user_object=end_user_object,
|
||||
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
|
||||
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
|
||||
fail_closed_budget_enforcement=(
|
||||
general_settings.get("fail_closed_budget_enforcement") is True
|
||||
or user_api_key_auth_obj.billing_agent_policy is not None
|
||||
),
|
||||
raw_body=await read_raw_json_body(request=request),
|
||||
)
|
||||
if request is not None:
|
||||
|
|
@ -3197,10 +3209,48 @@ 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,
|
||||
invocation_target,
|
||||
managed_agent_route_allowed,
|
||||
managed_inference_request,
|
||||
prepare_agent_invocation,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.proxy_server import general_settings, prisma_client, user_model
|
||||
|
||||
store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None
|
||||
if user_api_key_auth_obj.agent_id is not None:
|
||||
await admit_managed_actor(user_api_key_auth_obj, store)
|
||||
if user_api_key_auth_obj.managed_agent_policy is not None and not managed_agent_route_allowed(
|
||||
route, request.method
|
||||
):
|
||||
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
|
||||
authorized_data: Final = (
|
||||
managed_inference_request(
|
||||
route,
|
||||
request_data,
|
||||
general_settings,
|
||||
user_model,
|
||||
request.path_params.get("model") or request.path_params.get("model_name"),
|
||||
request.query_params.get("model"),
|
||||
)
|
||||
if user_api_key_auth_obj.managed_agent_policy is not None
|
||||
else request_data
|
||||
)
|
||||
target_name: Final = invocation_target(route, authorized_data)
|
||||
if target_name is not None:
|
||||
await prepare_agent_invocation(
|
||||
user_api_key_auth_obj,
|
||||
target_name,
|
||||
store,
|
||||
billable=request_data.get("method")
|
||||
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
|
||||
)
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
request_data=authorized_data,
|
||||
route=route,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -358,6 +358,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
team_id=team_id,
|
||||
end_user_id=end_user_id,
|
||||
call_type=call_type,
|
||||
agent_id=metadata.get("billing_agent_id") or metadata.get("agent_id"),
|
||||
):
|
||||
## UPDATE DATABASE
|
||||
charged: Final = await _update_database_and_spend_counters(
|
||||
|
|
@ -612,6 +613,7 @@ def _should_track_cost_callback(
|
|||
team_id: str | None,
|
||||
end_user_id: str | None,
|
||||
call_type: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the cost callback should be tracked based on the kwargs
|
||||
|
|
@ -628,7 +630,13 @@ def _should_track_cost_callback(
|
|||
if ProxyUpdateSpend.disable_spend_updates() is True:
|
||||
return False
|
||||
|
||||
if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None:
|
||||
if (
|
||||
agent_id is not None
|
||||
or user_api_key is not None
|
||||
or user_id is not None
|
||||
or team_id is not None
|
||||
or end_user_id is not None
|
||||
):
|
||||
return True
|
||||
return call_type in _UNATTRIBUTED_TRACKABLE_CALL_TYPES
|
||||
|
||||
|
|
|
|||
|
|
@ -1659,7 +1659,19 @@ class LiteLLMProxyRequestSetup:
|
|||
_key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None)
|
||||
_existing_agent_id: Final = data[_metadata_variable_name].get("agent_id")
|
||||
_resolved_agent_id: Final = _key_agent_id or _existing_agent_id
|
||||
data[_metadata_variable_name]["agent_id"] = _resolved_agent_id
|
||||
data[_metadata_variable_name]["agent_id"] = user_api_key_dict.invoked_agent_id or _resolved_agent_id
|
||||
managed_context: Final = user_api_key_dict.managed_agent_context
|
||||
data[_metadata_variable_name].update(
|
||||
MappingProxyType(
|
||||
{
|
||||
"actor_agent_id": user_api_key_dict.agent_id,
|
||||
"target_agent_id": user_api_key_dict.invoked_agent_id,
|
||||
"billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id,
|
||||
"agent_execution_mode": managed_context.mode if managed_context else None,
|
||||
"verified_human_user_id": managed_context.user_id if managed_context else None,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
|
||||
user_api_key_dict, "end_user_max_budget", None
|
||||
|
|
|
|||
|
|
@ -77,6 +77,11 @@ _SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)"
|
|||
_SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key"
|
||||
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
|
||||
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
|
||||
_SESSION_REPRESENTATIVE_ORDER_SQL: Final = (
|
||||
f"(call_type = {_AGENT_CALL_TYPE_SQL}) DESC, "
|
||||
f'CASE WHEN call_type = {_AGENT_CALL_TYPE_SQL} THEN "endTime" END DESC NULLS LAST, '
|
||||
f'call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC, request_id'
|
||||
)
|
||||
_BATCH_CALL_TYPES_SQL: Final = "('acreate_batch', 'create_batch', 'aretrieve_batch', 'retrieve_batch')"
|
||||
_SPAN_TYPE_SQL_CONDITIONS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
|
|
@ -2879,7 +2884,7 @@ async def ui_view_spend_logs(
|
|||
p += 1
|
||||
|
||||
# Status filter
|
||||
if status_filter is not None:
|
||||
if status_filter is not None and not (group_by_session is True and not is_search_lookup):
|
||||
if status_filter == "success":
|
||||
sql_conditions.append("(status = 'success' OR status IS NULL)")
|
||||
else:
|
||||
|
|
@ -2925,6 +2930,23 @@ async def ui_view_spend_logs(
|
|||
sql_params.append(f"%{error_message}%")
|
||||
p += 1
|
||||
|
||||
if status_filter is not None and group_by_session is True and not is_search_lookup:
|
||||
session_filter_conditions: Final = " AND ".join(sql_conditions) or "TRUE"
|
||||
sql_conditions.append(
|
||||
f"""({_SESSION_GROUP_KEY_SQL}) IN (
|
||||
SELECT session_key, api_key FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
{_SESSION_KEY_EXPR} AS session_key, api_key, status
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {session_filter_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_outcomes
|
||||
WHERE COALESCE(status, 'success') = ${p}
|
||||
)"""
|
||||
)
|
||||
sql_params.append(status_filter)
|
||||
p += 1
|
||||
|
||||
if (
|
||||
group_by_session is True
|
||||
and not is_v2
|
||||
|
|
@ -2991,7 +3013,7 @@ async def ui_view_spend_logs(
|
|||
{_SPEND_LOG_LIST_COLUMNS}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_representatives
|
||||
ORDER BY {exact_request_id_first}{_order_expr} {_sql_dir}{_nulls_clause}, request_id
|
||||
LIMIT ${p} OFFSET ${p + 1}
|
||||
|
|
@ -3063,7 +3085,7 @@ async def _fetch_session_representatives(
|
|||
next_param_index: int,
|
||||
session_keys: Sequence[tuple[str, str]],
|
||||
) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row
|
||||
"""Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
|
||||
"""Fetch the final agent outcome, or newest non-MCP row, of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
|
||||
rep_query: Final = f"""
|
||||
SELECT * FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
|
|
@ -3073,7 +3095,7 @@ async def _fetch_session_representatives(
|
|||
AND ({_SESSION_GROUP_KEY_SQL}) IN (
|
||||
SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[])
|
||||
)
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_representatives
|
||||
"""
|
||||
rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place
|
||||
|
|
@ -3140,7 +3162,7 @@ async def _ui_session_grouped_spend_logs(
|
|||
page_size``, trimmed to the end of the ``SPEND_LOGS_PAGINATION_COUNT_CAP``
|
||||
window the capped ``total`` promises, so a page never runs past that total
|
||||
and one starting at or past it returns no rows without a query. Each session is represented
|
||||
by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response``
|
||||
by its final agent outcome (or newest non-MCP row), enriched by ``_build_ui_spend_logs_response``
|
||||
exactly like the flat listing, and the response carries
|
||||
``next_session_cursor`` / ``has_more`` while ``total`` counts sessions
|
||||
(capped like the flat total). A page that runs out of sessions while still
|
||||
|
|
|
|||
|
|
@ -795,6 +795,7 @@ def get_logging_payload(
|
|||
model_id=_model_id,
|
||||
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
|
||||
agent_id=agent_id,
|
||||
billing_agent_id=clean_metadata.get("billing_agent_id"),
|
||||
requester_ip_address=clean_metadata.get("requester_ip_address", None),
|
||||
custom_llm_provider=custom_llm_provider or "",
|
||||
messages=_get_messages_for_spend_logs_payload(
|
||||
|
|
|
|||
|
|
@ -7847,7 +7847,7 @@ async def test_load_active_user_by_id_reads_the_row_from_the_database_not_the_ca
|
|||
key="fresh-jwt-user", value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=["team-a"])
|
||||
)
|
||||
proxy_globals.user_api_key_cache = cache
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_r
|
|||
key="stale-cache-user", value=_user(user_id="stale-cache-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="stale-cache-user", teams=["team-a"])
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
|
|
@ -164,7 +164,7 @@ async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_the
|
|||
key="deactivated-user", value=_user(user_id="deactivated-user", teams=["team-a"]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False})
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
|
|
|
|||
|
|
@ -810,6 +810,12 @@ class TestTestConnection:
|
|||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy.management_endpoints import mcp_management_endpoints
|
||||
|
||||
manager = MCPServerManager()
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
saved = MCPServer(
|
||||
server_id="saved-server-id",
|
||||
|
|
@ -1311,8 +1317,9 @@ class TestListToolsRestAPI:
|
|||
session_auth = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="grant-user", user_role="internal_user")
|
||||
admitted_auth = UserAPIKeyAuth(user_id="grant-user", org_id="admitted-org")
|
||||
|
||||
async def fake_reload(user_id):
|
||||
async def fake_reload(user_id, *, requires_fresh_policy=False):
|
||||
assert user_id == "grant-user"
|
||||
assert requires_fresh_policy is False
|
||||
return admitted_auth
|
||||
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -1480,9 +1487,12 @@ class TestListToolsRestAPI:
|
|||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", MCPServerManager())
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
|
|
@ -2414,6 +2424,7 @@ class TestCallToolRestAPI:
|
|||
|
||||
mock_server = MagicMock()
|
||||
mock_server.server_id = "server-1"
|
||||
mock_server.name = "Example server"
|
||||
|
||||
def fake_get_mcp_server_by_id(server_id):
|
||||
return mock_server if server_id == "server-1" else None
|
||||
|
|
@ -2431,6 +2442,11 @@ class TestCallToolRestAPI:
|
|||
raising=False,
|
||||
)
|
||||
|
||||
failure_log = AsyncMock()
|
||||
execute_tool = AsyncMock()
|
||||
monkeypatch.setattr(rest_endpoints, "_safe_fire_mcp_tool_call_failure_logging", failure_log)
|
||||
monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", execute_tool)
|
||||
|
||||
request_payload = {
|
||||
"server_id": "server-1",
|
||||
"name": "demo-tool",
|
||||
|
|
@ -2452,6 +2468,16 @@ class TestCallToolRestAPI:
|
|||
assert exc_info.value.detail["error"] == "access_denied"
|
||||
assert "server server-1" in exc_info.value.detail["message"]
|
||||
|
||||
execute_tool.assert_not_awaited()
|
||||
failure_log.assert_awaited_once()
|
||||
logged_data = failure_log.await_args.args[4]
|
||||
assert logged_data["model"] == "MCP: demo-tool"
|
||||
assert logged_data["metadata"]["model_group"] == "MCP: demo-tool"
|
||||
logging_obj = failure_log.await_args.args[0]
|
||||
assert logging_obj.model_call_details["mcp_tool_call_metadata"] == {
|
||||
"name": "demo-tool", "mcp_server_name": "Example server",
|
||||
}
|
||||
|
||||
async def test_executes_tool_when_allowed(self, monkeypatch):
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
|
||||
# Mock agent
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.agent_id = "test-agent"
|
||||
mock_agent.agent_card_params = {
|
||||
"url": "http://backend-agent:10001",
|
||||
"name": "Test Agent",
|
||||
|
|
@ -72,6 +73,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
"jsonrpc": "2.0",
|
||||
"id": "test-id",
|
||||
"method": "message/send",
|
||||
"metadata": {"model_info": {"id": "caller-supplied-id"}},
|
||||
"params": {
|
||||
"message": {
|
||||
"role": "user",
|
||||
|
|
@ -153,7 +155,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
"litellm.a2a_protocol.asend_message",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
),
|
||||
) as mock_send_message,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{},
|
||||
|
|
@ -190,6 +192,9 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
mock_add_data.assert_called_once()
|
||||
|
||||
# Verify model and custom_llm_provider were set
|
||||
assert mock_send_message.await_args.kwargs["model"] == "a2a_agent/Test Agent"
|
||||
assert captured_data["metadata"]["model_group"] == "a2a_agent/Test Agent"
|
||||
assert captured_data["metadata"]["model_info"] == {"id": mock_agent.agent_id}
|
||||
assert captured_data.get("model") == "a2a_agent/Test Agent"
|
||||
assert captured_data.get("custom_llm_provider") == "a2a_agent"
|
||||
|
||||
|
|
|
|||
27
tests/test_litellm/proxy/agent_endpoints/test_identity.py
Normal file
27
tests/test_litellm/proxy/agent_endpoints/test_identity.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
from collections.abc import Mapping
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.agent_endpoints.identity import has_legacy_identity, reject_legacy_identity
|
||||
|
||||
TENANT = "11111111-1111-4111-8111-111111111111"
|
||||
CLIENT = "22222222-2222-4222-8222-222222222222"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("params", [None, {}, {"model": "gpt-4o", "api_key": "sk-test"}])
|
||||
def test_runtime_params_without_identity_are_accepted(params: Mapping[str, object] | None) -> None:
|
||||
assert has_legacy_identity(params) is False
|
||||
reject_legacy_identity(params)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"identity", [None, {}, {"provider": "microsoft_entra", "tenant_id": TENANT, "client_id": CLIENT}]
|
||||
)
|
||||
def test_legacy_litellm_params_identity_is_rejected(identity: object) -> None:
|
||||
params: Mapping[str, object] = {"model": "gpt-4o", "identity": identity}
|
||||
assert has_legacy_identity(params) is True
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
reject_legacy_identity(params)
|
||||
assert failure.value.status_code == 400
|
||||
assert "top-level identity field" in failure.value.detail
|
||||
|
|
@ -2,15 +2,14 @@ import asyncio
|
|||
import re
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, Optional
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import (
|
||||
DEFAULT_JWKS_STALE_TTL,
|
||||
JWTLiteLLMRoleMap,
|
||||
|
|
@ -26,7 +25,6 @@ from litellm.proxy._types import (
|
|||
RoleBasedPermissions,
|
||||
ScopeMapping,
|
||||
)
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.auth.auth_checks import TeamNotFoundError
|
||||
from litellm.proxy.auth.handle_jwt import (
|
||||
|
|
@ -1637,7 +1635,6 @@ async def test_auth_builder_returns_team_membership_object():
|
|||
@pytest.mark.asyncio
|
||||
async def test_auth_builder_with_oidc_userinfo_enabled():
|
||||
"""Test that auth_builder uses OIDC UserInfo endpoint when enabled"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -1648,9 +1645,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
general_settings = {"enforce_rbac": False}
|
||||
route = "/chat/completions"
|
||||
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
|
||||
# Create JWT handler with OIDC UserInfo enabled
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -1677,18 +1672,12 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
|
||||
# Mock all the dependencies
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
|
||||
) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(
|
||||
JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
|
||||
) as mock_check_rbac,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
|
||||
patch.object(
|
||||
jwt_handler, "get_object_id", return_value=None
|
||||
) as mock_get_object_id,
|
||||
patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"get_user_info",
|
||||
|
|
@ -1696,9 +1685,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
return_value=("test_user_1", "test@example.com", True),
|
||||
) as mock_get_user_info,
|
||||
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
|
||||
patch.object(
|
||||
jwt_handler, "get_end_user_id", return_value=None
|
||||
) as mock_get_end_user_id,
|
||||
patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"check_admin_access",
|
||||
|
|
@ -1711,9 +1698,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(None, None),
|
||||
) as mock_find_team,
|
||||
patch.object(
|
||||
JWTAuthManager, "get_all_team_ids", return_value=set()
|
||||
) as mock_get_all_team_ids,
|
||||
patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"find_team_with_model_access",
|
||||
|
|
@ -1726,15 +1711,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(user_object, None, None, None, user_object.user_id),
|
||||
) as mock_get_objects,
|
||||
patch.object(
|
||||
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
|
||||
) as mock_map_user,
|
||||
patch.object(
|
||||
JWTAuthManager, "validate_object_id", return_value=True
|
||||
) as mock_validate_object,
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
) as mock_sync_user,
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
|
||||
):
|
||||
# Set up mock return values
|
||||
mock_get_userinfo.return_value = userinfo_response
|
||||
|
|
@ -1764,7 +1743,6 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
@pytest.mark.asyncio
|
||||
async def test_auth_builder_with_oidc_userinfo_disabled():
|
||||
"""Test that auth_builder uses JWT validation when OIDC UserInfo is disabled"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -1775,9 +1753,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
general_settings = {"enforce_rbac": False}
|
||||
route = "/chat/completions"
|
||||
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
|
||||
# Create JWT handler with OIDC UserInfo disabled
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -1801,18 +1777,12 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
|
||||
# Mock all the dependencies
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
|
||||
) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(
|
||||
JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
|
||||
) as mock_check_rbac,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
|
||||
patch.object(
|
||||
jwt_handler, "get_object_id", return_value=None
|
||||
) as mock_get_object_id,
|
||||
patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"get_user_info",
|
||||
|
|
@ -1820,9 +1790,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
return_value=("test_user_1", None, None),
|
||||
) as mock_get_user_info,
|
||||
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
|
||||
patch.object(
|
||||
jwt_handler, "get_end_user_id", return_value=None
|
||||
) as mock_get_end_user_id,
|
||||
patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"check_admin_access",
|
||||
|
|
@ -1835,9 +1803,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(None, None),
|
||||
) as mock_find_team,
|
||||
patch.object(
|
||||
JWTAuthManager, "get_all_team_ids", return_value=set()
|
||||
) as mock_get_all_team_ids,
|
||||
patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"find_team_with_model_access",
|
||||
|
|
@ -1850,15 +1816,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(user_object, None, None, None, user_object.user_id),
|
||||
) as mock_get_objects,
|
||||
patch.object(
|
||||
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
|
||||
) as mock_map_user,
|
||||
patch.object(
|
||||
JWTAuthManager, "validate_object_id", return_value=True
|
||||
) as mock_validate_object,
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
) as mock_sync_user,
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
|
||||
):
|
||||
# Set up mock return values
|
||||
mock_auth_jwt.return_value = jwt_response
|
||||
|
|
@ -2631,7 +2591,6 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
|
|||
"""
|
||||
Test that find_and_validate_specific_team_id resolves team by name when team_id is not found
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
|
||||
|
|
@ -2654,9 +2613,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
|
|||
# Mock team object returned by get_team_object_by_alias
|
||||
team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock
|
||||
) as mock_get_by_alias:
|
||||
with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias:
|
||||
mock_get_by_alias.return_value = team_object
|
||||
|
||||
team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
|
|
@ -2685,7 +2642,6 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
"""
|
||||
Test that team_id_jwt_field takes precedence over team_alias_jwt_field
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
|
||||
|
|
@ -2699,9 +2655,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(
|
||||
team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"
|
||||
),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"),
|
||||
)
|
||||
|
||||
# Token with both team_id and team name
|
||||
|
|
@ -2711,9 +2665,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
team_object = LiteLLM_TeamTable(team_id="direct-team-id")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock
|
||||
) as mock_get_by_id,
|
||||
patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id,
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object_by_alias",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -2890,7 +2842,6 @@ async def test_get_objects_resolves_org_by_name():
|
|||
@pytest.mark.asyncio
|
||||
async def test_resolve_jwks_url_passthrough_for_direct_jwks_url():
|
||||
"""Non-discovery URLs are returned unchanged."""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
|
@ -3143,7 +3094,7 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
|
|||
When team_id_jwt_field is a normal field name (no dot-notation) the
|
||||
error message should not contain a spurious bracket-notation hint.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
|
@ -3230,8 +3181,8 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
|
|||
async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
||||
user_id: str,
|
||||
user_teams: list,
|
||||
get_team_object_return: Optional[str],
|
||||
expected_team_id: Optional[str],
|
||||
get_team_object_return: str | None,
|
||||
expected_team_id: str | None,
|
||||
expect_get_team_called: bool,
|
||||
expect_get_membership_called: bool,
|
||||
) -> None:
|
||||
|
|
@ -3244,9 +3195,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
if len(user_teams) == 1 and get_team_object_return == "resolved_row":
|
||||
only = user_teams[0]
|
||||
team_table = LiteLLM_TeamTable(team_id=only)
|
||||
membership = LiteLLM_TeamMembership(
|
||||
user_id=user_id, team_id=only, litellm_budget_table=None
|
||||
)
|
||||
membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None)
|
||||
get_team_return_value = team_table
|
||||
membership_return_value = membership
|
||||
else:
|
||||
|
|
@ -3305,9 +3254,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -3324,9 +3271,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
code = 404 if get_team_object_return == "http_404" else 500
|
||||
mock_get_team.side_effect = HTTPException(
|
||||
status_code=code,
|
||||
detail={
|
||||
"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."
|
||||
},
|
||||
detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."},
|
||||
)
|
||||
else:
|
||||
mock_get_team.return_value = get_team_return_value
|
||||
|
|
@ -4047,7 +3992,7 @@ def _encode_rsa_jwt(
|
|||
issuer: str,
|
||||
audience: str,
|
||||
kid: str,
|
||||
extra_claims: Optional[dict] = None,
|
||||
extra_claims: dict | None = None,
|
||||
) -> str:
|
||||
import time
|
||||
|
||||
|
|
@ -4743,12 +4688,9 @@ async def test_get_objects_team_membership_uses_rebound_user_id():
|
|||
async def fake_get_team_membership(user_id, team_id, *args, **kwargs):
|
||||
captured["user_id"] = user_id
|
||||
captured["team_id"] = team_id
|
||||
return None
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
user_id_jwt_field="email", user_id_upsert=True
|
||||
)
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True)
|
||||
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -5389,7 +5331,7 @@ async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback(
|
|||
assert team_object is None
|
||||
|
||||
|
||||
def _db_fallback_handler(litellm_jwtauth: Optional[LiteLLM_JWTAuth] = None) -> JWTHandler:
|
||||
def _db_fallback_handler(litellm_jwtauth: LiteLLM_JWTAuth | None = None) -> JWTHandler:
|
||||
handler = JWTHandler()
|
||||
handler.litellm_jwtauth = litellm_jwtauth or LiteLLM_JWTAuth()
|
||||
return handler
|
||||
|
|
@ -5447,9 +5389,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
|
|||
"expect_403",
|
||||
),
|
||||
[
|
||||
pytest.param(
|
||||
True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"
|
||||
),
|
||||
pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"),
|
||||
pytest.param(
|
||||
True,
|
||||
["team_a", "team_b"],
|
||||
|
|
@ -5497,8 +5437,8 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
|
|||
async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
||||
fallback_to_db_teams: bool,
|
||||
user_teams: list,
|
||||
header_team_id: Optional[str],
|
||||
expected_team_id: Optional[str],
|
||||
header_team_id: str | None,
|
||||
expected_team_id: str | None,
|
||||
expect_403: bool,
|
||||
) -> None:
|
||||
"""End-to-end auth_builder behavior with no JWT team claims.
|
||||
|
|
@ -5527,9 +5467,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
|||
|
||||
async def call_auth_builder():
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "auth_jwt", new_callable=AsyncMock
|
||||
) as mock_auth_jwt,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock),
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None),
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]),
|
||||
|
|
@ -5569,9 +5507,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -6765,7 +6701,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
|
|||
team_id_upsert=True,
|
||||
)
|
||||
|
||||
upsert_by_team: dict[str, Optional[bool]] = {}
|
||||
upsert_by_team: dict[str, bool | None] = {}
|
||||
|
||||
async def spy_get_team(team_id, **kwargs):
|
||||
upsert_by_team[team_id] = kwargs.get("team_id_upsert")
|
||||
|
|
@ -6800,9 +6736,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -7806,6 +7740,58 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc
|
|||
assert result["team_id"] is None
|
||||
|
||||
|
||||
def _explicit_identity_registry() -> AgentRegistry:
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(AgentResponse(
|
||||
agent_id="explicit-agent-id",
|
||||
agent_name="Readable agent name",
|
||||
agent_card_params={},
|
||||
litellm_params={"identity": {
|
||||
"provider": "microsoft_entra",
|
||||
"tenant_id": "11111111-1111-4111-8111-111111111111",
|
||||
"client_id": "22222222-2222-4222-8222-222222222222",
|
||||
}},
|
||||
))
|
||||
return registry
|
||||
|
||||
|
||||
@pytest.mark.parametrize("claim_field", ["azp", None])
|
||||
def test_runtime_json_cannot_establish_a_managed_identity(claim_field: str | None) -> None:
|
||||
registry: Final = _explicit_identity_registry()
|
||||
handler: Final = _entra_agent_jwt_handler(claim_field)
|
||||
claims: Final = {
|
||||
"iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
|
||||
"tid": "11111111-1111-4111-8111-111111111111",
|
||||
"azp": "22222222-2222-4222-8222-222222222222",
|
||||
}
|
||||
if claim_field is None:
|
||||
assert JWTAuthManager.resolve_agent_id(handler, claims, registry) is None
|
||||
else:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
JWTAuthManager.resolve_agent_id(handler, claims, registry)
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.parametrize("override", [
|
||||
{"iss": "https://attacker.example"},
|
||||
{"tid": "33333333-3333-4333-8333-333333333333"},
|
||||
{"azp": "33333333-3333-4333-8333-333333333333"},
|
||||
{"azp": "explicit-agent-id"},
|
||||
{"azp": "Readable agent name"},
|
||||
])
|
||||
def test_explicit_entra_identity_cannot_be_claimed_via_legacy_lookup(override: Mapping[str, object]) -> None:
|
||||
registry: Final = _explicit_identity_registry()
|
||||
handler: Final = _entra_agent_jwt_handler("azp")
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
JWTAuthManager.resolve_agent_id(handler, {
|
||||
"iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
|
||||
"tid": "11111111-1111-4111-8111-111111111111",
|
||||
"azp": "22222222-2222-4222-8222-222222222222",
|
||||
**override,
|
||||
}, registry)
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("existing_user", [False, True])
|
||||
@pytest.mark.parametrize("warm_cache", [False, True])
|
||||
|
|
@ -7853,3 +7839,187 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning
|
|||
users.create.assert_not_awaited()
|
||||
if existing_user:
|
||||
assert users.find_unique.await_count == (0 if warm_cache else 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
|
||||
@pytest.mark.parametrize("audience_validation", (True, False))
|
||||
@pytest.mark.parametrize(
|
||||
"route,allowed",
|
||||
[
|
||||
("/chat/completions", True), ("/v1/messages", True), ("/v1/responses", True),
|
||||
("/mcp-rest/tools/call", True), ("/a2a/target", True),
|
||||
("/v1/files", False), ("/v1/batches", False), ("/v1/vector_stores", False),
|
||||
("/v1/containers", False), ("/openai/v1/files", False),
|
||||
("/v1/responses/other-response", False), ("/v1/realtime/client_secrets", False),
|
||||
],
|
||||
)
|
||||
async def test_managed_application_uses_persisted_identity_without_provisioning_human(
|
||||
monkeypatch: pytest.MonkeyPatch, mode: str, audience_validation: bool, route: str, allowed: bool
|
||||
) -> None:
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
tenant: Final = "11111111-1111-4111-8111-111111111111"
|
||||
client_id: Final = "22222222-2222-4222-8222-222222222222"
|
||||
principal: Final = "33333333-3333-4333-8333-333333333333"
|
||||
issuer: Final = f"https://login.microsoftonline.com/{tenant}/v2.0"
|
||||
jwks_url: Final = "https://login.microsoftonline.test/managed-keys"
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
|
||||
private_key, jwk = _get_rsa_key_and_jwk(kid="managed-key")
|
||||
cache: Final = DualCache()
|
||||
cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(None, cache, LiteLLM_JWTAuth(user_id_upsert=True))
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="stable-id",
|
||||
provider="microsoft_entra",
|
||||
issuer=issuer,
|
||||
tenant_id=tenant,
|
||||
client_id=client_id,
|
||||
service_principal_id=principal,
|
||||
revision="revision-one",
|
||||
required_roles=("Agent.Invoke",),
|
||||
)
|
||||
agent: Final = AgentResponse.model_validate(
|
||||
{
|
||||
"agent_id": "stable-id",
|
||||
"agent_name": "A readable name",
|
||||
"agent_card_params": {},
|
||||
"identity": binding,
|
||||
"identity_managed": True,
|
||||
"execution_mode": mode,
|
||||
}
|
||||
)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
|
||||
database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1)
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
|
||||
database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None)
|
||||
database.db.litellm_usertable.upsert = AsyncMock()
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key,
|
||||
issuer=issuer,
|
||||
audience="api://gateway",
|
||||
kid="managed-key",
|
||||
extra_claims={
|
||||
"tid": tenant,
|
||||
"azp": client_id,
|
||||
"oid": principal,
|
||||
"roles": ["Agent.Invoke"],
|
||||
"idtyp": "app",
|
||||
},
|
||||
)
|
||||
arguments: Final = dict(
|
||||
api_key=token,
|
||||
jwt_handler=handler,
|
||||
request_data={},
|
||||
general_settings={},
|
||||
route=route,
|
||||
prisma_client=database,
|
||||
user_api_key_cache=cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
if not audience_validation:
|
||||
monkeypatch.delenv("JWT_AUDIENCE")
|
||||
if mode == "delegated" or not audience_validation or not allowed:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await JWTAuthManager.auth_builder(**arguments)
|
||||
assert failure.value.status_code == 403
|
||||
else:
|
||||
result: Final = await JWTAuthManager.auth_builder(**arguments)
|
||||
auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
|
||||
assert auth.agent_id == "stable-id"
|
||||
assert auth.user_id is None
|
||||
assert auth.team_id is None
|
||||
assert auth.managed_agent_context is not None
|
||||
assert auth.managed_agent_context.mode == "autonomous"
|
||||
assert result["is_proxy_admin"] is False
|
||||
database.db.litellm_usertable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("claim_value", ["managed", "Readable managed agent"])
|
||||
def test_legacy_claim_cannot_select_a_top_level_entra_binding(claim_value: str) -> None:
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(
|
||||
AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Readable managed agent",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
)
|
||||
)
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
JWTAuthManager.resolve_agent_id(_entra_agent_jwt_handler("agent"), {"agent": claim_value}, registry)
|
||||
assert denied.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["human", "config-agent", "managed-agent"])
|
||||
async def test_database_free_jwt_admission_with_entra_shaped_claims(monkeypatch: pytest.MonkeyPatch, kind: str) -> None:
|
||||
issuer: Final = "https://login.microsoftonline.com/test-tenant/v2.0"
|
||||
jwks_url: Final = "https://login.microsoftonline.test/config-only-keys"
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
|
||||
private_key, jwk = _get_rsa_key_and_jwk(kid="config-key")
|
||||
cache: Final = DualCache()
|
||||
cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(
|
||||
AgentResponse(
|
||||
agent_id="configured",
|
||||
agent_name="Configured",
|
||||
agent_card_params={},
|
||||
identity_managed=kind == "managed-agent",
|
||||
)
|
||||
)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"]))
|
||||
handler.bind_agent_lookup(registry)
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key,
|
||||
issuer=issuer,
|
||||
audience="api://gateway",
|
||||
kid="config-key",
|
||||
extra_claims={
|
||||
"tid": "test-tenant",
|
||||
"azp": "application",
|
||||
"scope": "litellm_proxy_admin",
|
||||
**({"agent": "configured"} if kind != "human" else {}),
|
||||
},
|
||||
)
|
||||
arguments: Final = dict(
|
||||
api_key=token,
|
||||
jwt_handler=handler,
|
||||
request_data={},
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
prisma_client=None,
|
||||
user_api_key_cache=cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
if kind == "managed-agent":
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await JWTAuthManager.auth_builder(**arguments)
|
||||
assert denied.value.status_code == 403
|
||||
else:
|
||||
result: Final = await JWTAuthManager.auth_builder(**arguments)
|
||||
auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
|
||||
assert auth.agent_id == ("configured" if kind == "config-agent" else None)
|
||||
assert auth.managed_agent_context is None
|
||||
assert result["is_proxy_admin"] is True
|
||||
|
|
|
|||
|
|
@ -9278,6 +9278,8 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
|
|||
)
|
||||
|
||||
async def auth_that_reserves(request, api_key):
|
||||
assert request.method == "GET"
|
||||
assert request.query_params.get("model") == "gpt-realtime"
|
||||
request.state.budget_reservation = reservation
|
||||
return UserAPIKeyAuth(token="hashed", budget_reservation=reservation)
|
||||
|
||||
|
|
@ -9290,3 +9292,149 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
|
|||
assert result.budget_reservation == reservation
|
||||
assert websocket.state.budget_reservation is reservation
|
||||
assert websocket.scope["state"]["budget_reservation"] is reservation
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("invoke", [False, True])
|
||||
async def test_centralized_authorization_preserves_database_free_config_agents(monkeypatch, invoke: bool):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": None,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
registry = AgentRegistry()
|
||||
registry.load_agents_from_config(
|
||||
[{"agent_name": "config-agent", "agent_card_params": {"name": "Config", "url": "http://localhost:9999"}}]
|
||||
)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
registered = registry.get_agent_by_name("config-agent")
|
||||
model = "a2a/config-agent" if invoke else "test-model"
|
||||
auth = UserAPIKeyAuth(agent_id=registered.agent_id, jwt_claims={"agent": "config-agent"}, models=[model])
|
||||
data = {"model": model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
assert (
|
||||
await _authorize_authenticated_request(
|
||||
auth, _alias_request("/v1/chat/completions", data), data, "/v1/chat/completions", "jwt-token"
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert auth.managed_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_virtual_key_cannot_access_provider_resource_routes(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
policy = AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Managed",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="application",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
object_permission={"models": ["test-model"]},
|
||||
)
|
||||
database = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
request = _alias_request("/v1/files", {})
|
||||
request.scope["method"] = "GET"
|
||||
auth = UserAPIKeyAuth(agent_id="managed", api_key="persisted-key", models=["test-model"])
|
||||
with pytest.raises(ProxyException) as denied:
|
||||
await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key")
|
||||
assert denied.value.code == "403"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("requested", [None, "test-model"])
|
||||
@pytest.mark.parametrize("grant_default", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"route,settings,cli_model",
|
||||
[
|
||||
("/v1/chat/completions", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/responses", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/messages", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/moderations", {"moderation_model": "forbidden-model"}, None),
|
||||
("/v1/audio/transcriptions", {"moderation_model": "forbidden-model"}, None),
|
||||
("/v1/audio/speech", {}, "forbidden-model"),
|
||||
("/v1/chat/completions", {}, "forbidden-model"),
|
||||
("/v1/images/generations", {"image_generation_model": "forbidden-model"}, None),
|
||||
("/v1/images/edits", {"image_generation_model": "forbidden-model"}, None),
|
||||
],
|
||||
)
|
||||
async def test_managed_agent_cannot_bypass_grants_with_server_default(
|
||||
monkeypatch, requested, route, settings, cli_model, grant_default
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext
|
||||
|
||||
policy = AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Managed",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="application",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
object_permission={"models": ["test-model", "forbidden-model"] if grant_default else ["test-model"]},
|
||||
)
|
||||
database = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"general_settings": settings,
|
||||
"user_model": cli_model,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
data = {"messages": [{"role": "user", "content": "hi"}], **({"model": requested} if requested else {})}
|
||||
auth = UserAPIKeyAuth(agent_id="managed")
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="managed", binding_revision="revision", mode="autonomous"
|
||||
)
|
||||
if not grant_default:
|
||||
with pytest.raises(ProxyException) as denied:
|
||||
await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
|
||||
assert denied.value.code == "403"
|
||||
assert "forbidden-model" in denied.value.message
|
||||
return
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
|
||||
new_callable=AsyncMock,
|
||||
) as reserve:
|
||||
reserve.return_value = None
|
||||
assert (
|
||||
await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
|
||||
is None
|
||||
)
|
||||
reserve.assert_awaited_once()
|
||||
assert reserve.call_args.kwargs["request_body"]["model"] == "forbidden-model"
|
||||
|
|
|
|||
|
|
@ -2712,3 +2712,34 @@ async def test_track_cost_callback_failure_alert_never_carries_request_metadata_
|
|||
assert "headers" in failure_debug_lines[0]
|
||||
else:
|
||||
assert failure_debug_lines == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("identity_field", ["agent_id", "billing_agent_id"])
|
||||
async def test_autonomous_llm_callback_persists_without_human_or_key(identity_field: str) -> None:
|
||||
kwargs: Final = {
|
||||
"call_type": "acompletion",
|
||||
"model": "test-model",
|
||||
"response_cost": 0.01,
|
||||
"litellm_params": {"metadata": {identity_field: "autonomous-agent"}},
|
||||
}
|
||||
with patch( # test-quality-ok: callback invokes this module-level persistence boundary without an injection seam
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback._update_database_and_spend_counters",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
) as persist:
|
||||
await _ProxyDBLogger()._PROXY_track_cost_callback(
|
||||
kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now()
|
||||
)
|
||||
persist.assert_awaited_once()
|
||||
assert persist.call_args.kwargs["response_cost"] == 0.01
|
||||
assert persist.call_args.kwargs["user_id"] is None
|
||||
assert persist.call_args.kwargs["user_api_key"] is None
|
||||
assert persist.call_args.kwargs["kwargs"]["litellm_params"]["metadata"][identity_field] == "autonomous-agent"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("agent_id,expected", [(None, False), ("autonomous-agent", True)])
|
||||
def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id: str | None, expected: bool) -> None:
|
||||
assert _should_track_cost_callback(
|
||||
user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id
|
||||
) is expected
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import datetime
|
|||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import sqlite3
|
||||
from datetime import timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -3762,7 +3763,7 @@ class TestSpendLogsPayload:
|
|||
"model": "gpt-4o",
|
||||
"user": "",
|
||||
"team_id": "",
|
||||
"metadata": '{"applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
|
||||
"metadata": '{"actor_agent_id": null, "target_agent_id": null, "billing_agent_id": null, "agent_execution_mode": null, "verified_human_user_id": null, "applied_guardrails": [], "attempted_fallbacks": null, "original_model_group": null, "batch_models": null, "batch_successful_requests": null, "batch_failed_requests": null, "mcp_tool_call_metadata": null, "vector_store_request_metadata": null, "routing_decision": null, "internal_call_origin": null, "guardrail_information": null, "compression_savings": null, "litellm_gateway_injected_cache": null, "router_metadata": null, "autorouter_savings_estimate": null, "autorouter_baseline_observation": null, "azure_spillover": null, "usage_object": {"completion_tokens": 20, "prompt_tokens": 10, "total_tokens": 30, "completion_tokens_details": null, "prompt_tokens_details": null}, "model_map_information": {"model_map_key": "gpt-4o", "model_map_value": {"key": "gpt-4o", "max_tokens": 16384, "max_input_tokens": 128000, "max_output_tokens": 16384, "input_cost_per_token": 2.5e-06, "cache_creation_input_token_cost": null, "cache_read_input_token_cost": 1.25e-06, "input_cost_per_character": null, "input_cost_per_token_above_128k_tokens": null, "input_cost_per_token_above_200k_tokens": null, "input_cost_per_query": null, "input_cost_per_second": null, "input_cost_per_audio_token": null, "input_cost_per_token_batches": 1.25e-06, "output_cost_per_token_batches": 5e-06, "output_cost_per_token": 1e-05, "output_cost_per_audio_token": null, "output_cost_per_character": null, "output_cost_per_token_above_128k_tokens": null, "output_cost_per_character_above_128k_tokens": null, "output_cost_per_token_above_200k_tokens": null, "output_cost_per_second": null, "output_cost_per_reasoning_token": null, "output_cost_per_image": null, "output_vector_size": null, "litellm_provider": "openai", "mode": "chat", "supports_system_messages": true, "supports_response_schema": true, "supports_vision": true, "supports_function_calling": true, "supports_tool_choice": true, "supports_assistant_prefill": false, "supports_prompt_caching": true, "supports_audio_input": false, "supports_audio_output": false, "supports_pdf_input": false, "supports_embedding_image_input": false, "supports_native_streaming": null, "supports_web_search": true, "supports_reasoning": false, "search_context_cost_per_query": {"search_context_size_low": 0.03, "search_context_size_medium": 0.035, "search_context_size_high": 0.05}, "tpm": null, "rpm": null, "supported_openai_params": ["frequency_penalty", "logit_bias", "logprobs", "top_logprobs", "max_tokens", "max_completion_tokens", "modalities", "prediction", "n", "presence_penalty", "seed", "stop", "stream", "stream_options", "temperature", "top_p", "tools", "tool_choice", "function_call", "functions", "max_retries", "extra_headers", "parallel_tool_calls", "audio", "response_format", "user"]}}, "additional_usage_values": {"completion_tokens_details": null, "prompt_tokens_details": null}}',
|
||||
"cache_key": "Cache OFF",
|
||||
"spend": 0.00022500000000000002,
|
||||
"total_tokens": 30,
|
||||
|
|
@ -3781,6 +3782,7 @@ class TestSpendLogsPayload:
|
|||
"status": "success",
|
||||
"mcp_namespaced_tool_name": None,
|
||||
"agent_id": None,
|
||||
"billing_agent_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -6586,9 +6588,7 @@ def test_key_spend_report_scopes_to_caller_key(client, monkeypatch):
|
|||
|
||||
|
||||
def test_key_spend_report_scopes_a_cli_session_to_the_per_user_alias_not_the_login_token(client, monkeypatch):
|
||||
mock_prisma = _spend_report_mock_prisma(
|
||||
query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}]
|
||||
)
|
||||
mock_prisma = _spend_report_mock_prisma(query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
|
|
@ -7138,9 +7138,8 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc
|
|||
rep_call = emitted[2]
|
||||
assert f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in rep_call[0]
|
||||
assert (
|
||||
f"ORDER BY {SESSION_GROUP_KEY_SQL}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
|
||||
in rep_call[0]
|
||||
), "the session representative must prefer the newest non-MCP call"
|
||||
f"ORDER BY {SESSION_GROUP_KEY_SQL}, " + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
) in rep_call[0]
|
||||
assert rep_call[-2] == ["sess-1", "req-solo"]
|
||||
assert rep_call[-1] == ["hashed-key", "hashed-key"]
|
||||
finally:
|
||||
|
|
@ -7630,6 +7629,39 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo
|
|||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("parent_status", "child_status", "expected"),
|
||||
[("failure", "success", "failure"), ("success", "failure", "success")],
|
||||
)
|
||||
def test_session_representative_uses_completed_agent_outcome(parent_status, child_status, expected):
|
||||
with sqlite3.connect(":memory:") as connection:
|
||||
connection.execute(
|
||||
'CREATE TABLE logs (request_id TEXT, call_type TEXT, status TEXT, "startTime" TEXT, "endTime" TEXT)'
|
||||
)
|
||||
connection.executemany(
|
||||
"INSERT INTO logs VALUES (?, ?, ?, ?, ?)",
|
||||
(
|
||||
("parent", "asend_message", parent_status, "10:00:00", "10:00:05"),
|
||||
("nested-agent", "asend_message", child_status, "10:00:01", "10:00:03"),
|
||||
("llm", "acompletion", "success", "10:00:02", "10:00:04"),
|
||||
("tool", "call_mcp_tool", child_status, "10:00:04", "10:00:04"),
|
||||
),
|
||||
)
|
||||
result = connection.execute(
|
||||
"SELECT request_id, status FROM logs ORDER BY "
|
||||
+ spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
+ " LIMIT 1"
|
||||
).fetchone()
|
||||
assert result == ("parent", expected)
|
||||
connection.execute("DELETE FROM logs WHERE call_type = 'asend_message'")
|
||||
fallback = connection.execute(
|
||||
"SELECT request_id, status FROM logs ORDER BY "
|
||||
+ spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
+ " LIMIT 1"
|
||||
).fetchone()
|
||||
assert fallback == ("llm", "success")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calculate_spend_unpriced_model_returns_400():
|
||||
model = "openrouter/unit-test-unpriced-model"
|
||||
|
|
|
|||
|
|
@ -539,7 +539,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
|
|||
async def mock_query_raw(sql_query, *params):
|
||||
if "COUNT(*) AS total_count" in sql_query:
|
||||
return [{"total_count": 60}]
|
||||
if "DISTINCT ON" in sql_query:
|
||||
if "AS session_representatives" in sql_query:
|
||||
return representative_rows
|
||||
return session_rows
|
||||
|
||||
|
|
@ -584,9 +584,11 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
|
|||
|
||||
rep_sql = emitted[2][0]
|
||||
assert f"DISTINCT ON ({group_key})" in rep_sql, f"page must return one row per session. SQL was:\n{rep_sql}"
|
||||
assert f"ORDER BY {group_key}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in rep_sql, (
|
||||
"the session representative must prefer the newest non-MCP call"
|
||||
)
|
||||
assert (
|
||||
f"ORDER BY {group_key}, (call_type = 'asend_message') DESC, "
|
||||
"CASE WHEN call_type = 'asend_message' THEN \"endTime\" END DESC NULLS LAST, "
|
||||
"call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
|
||||
) in rep_sql, "the session representative must prefer the final agent outcome, then the newest non-MCP call"
|
||||
assert "COUNT(*) OVER ()" not in rep_sql
|
||||
|
||||
assert [row["request_id"] for row in response["data"]] == ["req-1", "req-2"]
|
||||
|
|
|
|||
|
|
@ -5156,6 +5156,24 @@ def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_recei
|
|||
)
|
||||
|
||||
|
||||
def test_failed_agent_request_keeps_registered_display_name():
|
||||
agent_model: Final = "a2a_agent/Research Agent"
|
||||
payload: Final = get_logging_payload(
|
||||
kwargs={
|
||||
"model": agent_model,
|
||||
"call_type": "asend_message",
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": agent_model, "model_info": {"id": "registered-agent"}, "status": "failure"}
|
||||
},
|
||||
},
|
||||
response_obj=ValueError("Agent action denied"),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
assert payload["model"] == agent_model
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["model_id"] == "registered-agent"
|
||||
|
||||
_CLI_SESSION_ALIAS: Final = "cli-session-alice"
|
||||
_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
|
||||
|
||||
|
|
@ -5281,3 +5299,21 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None:
|
|||
assert result["autorouter_savings_estimate"] == recorded
|
||||
absent: Final = _get_spend_logs_metadata({"autorouter_savings_estimate": supplied}) # mutable-ok: legacy metadata helper accepts dicts
|
||||
assert absent["autorouter_savings_estimate"] is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("billing_agent", [None, "authenticated-agent"])
|
||||
def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None:
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"litellm_params": {"metadata": {
|
||||
"user_api_key": "test-key",
|
||||
"agent_id": "header-selected-agent",
|
||||
"billing_agent_id": billing_agent,
|
||||
}},
|
||||
}
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs, response_obj={"id": "request"},
|
||||
start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
assert payload["agent_id"] == "header-selected-agent"
|
||||
assert payload["billing_agent_id"] == billing_agent
|
||||
|
|
|
|||
|
|
@ -149,6 +149,9 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re
|
|||
prisma_client.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
side_effect=[None, _DbAgentRow("a2a-sibling-replica-agent-id", agent_name)]
|
||||
)
|
||||
prisma_client.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=_DbAgentRow("a2a-sibling-replica-agent-id", agent_name)
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue