feat(agents): authenticate Entra actors and retain request attribution

This commit is contained in:
Joshua Valluru 2026-09-26 11:55:03 -07:00
parent 97ddb52e0d
commit b02dfccef4
24 changed files with 862 additions and 156 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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