Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_internal_copy_38013

# Conflicts:
#	tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py
This commit is contained in:
mateo-berri 2026-09-03 17:22:14 -07:00
commit 56cf0cd223
17 changed files with 908 additions and 243 deletions

View file

@ -3,7 +3,7 @@
"limit": 14074
},
"reportArgumentType": {
"limit": 2214
"limit": 2206
},
"reportAssignmentType": {
"limit": 319
@ -57,7 +57,7 @@
"limit": 5601
},
"reportMissingTypeArgument": {
"limit": 15287
"limit": 15285
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44362
"limit": 44360
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38323
"limit": 38311
},
"reportUnknownParameterType": {
"limit": 19624
},
"reportUnknownVariableType": {
"limit": 29861
"limit": 29847
},
"reportUnnecessaryCast": {
"limit": 111

View file

@ -1,8 +1,9 @@
"""Anthropic error format type definitions."""
from collections.abc import Mapping
from typing import Literal
from typing_extensions import Required, TypedDict
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict
# Known Anthropic error types
# Source: https://docs.anthropic.com/en/api/errors
@ -23,6 +24,7 @@ class AnthropicErrorDetail(TypedDict):
type: AnthropicErrorType
message: str
provider_specific_fields: NotRequired[ReadOnly[Mapping[str, object]]]
class AnthropicErrorResponse(TypedDict, total=False):

View file

@ -6,6 +6,7 @@ from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.llms.openai import ChatCompletionSystemMessage
if TYPE_CHECKING:
from litellm.exceptions import ContentPolicyViolationError
@ -36,6 +37,16 @@ def safeguard_refusal_error(model: str, stop_details: Mapping[str, object]) -> "
)
def anthropic_system_to_openai_message(system: object) -> ChatCompletionSystemMessage | None:
"""
Return the Anthropic Messages top-level ``system`` (a string or a list of text
blocks) as an OpenAI-style system message, or None when the request has none.
"""
if not isinstance(system, (str, list)) or not system:
return None
return ChatCompletionSystemMessage(role="system", content=system)
@lru_cache(maxsize=1)
def _anthropic_messages_optional_param_keys() -> frozenset[str]:
"""

View file

@ -9,7 +9,7 @@ from fastapi.responses import JSONResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
from litellm.anthropic_interface.exceptions import AnthropicErrorResponse, AnthropicExceptionMapping
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.anthropic.experimental_pass_through.context_management import (
AnthropicContextManagementError,
@ -30,6 +30,27 @@ from litellm.types.utils import TokenCountResponse
router: Final = APIRouter()
def _anthropic_error_json_response(exc: ProxyException, request: Request) -> JSONResponse:
from litellm.proxy.proxy_server import (
_close_dangling_otel_server_span, # pyright: ignore[reportPrivateUsage] # proxy_server keeps the span-close helper private; error JSONResponses returned by the route must stamp the OTel server span like the global ProxyException handler does
)
status_code: Final = int(exc.code) if exc.code is not None and exc.code.isdigit() else 500
_close_dangling_otel_server_span(request, status_code, exc=exc)
envelope: Final = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=status_code,
raw_message=exc.message,
request_id=request.headers.get("x-request-id"),
)
if not exc.provider_specific_fields:
return JSONResponse(status_code=status_code, content=envelope, headers=exc.headers)
content: Final[AnthropicErrorResponse] = {
**envelope,
"error": {**envelope["error"], "provider_specific_fields": exc.provider_specific_fields},
}
return JSONResponse(status_code=status_code, content=content, headers=exc.headers)
def _strip_total_tokens_from_anthropic_response(response: Any) -> None:
"""Remove the OpenAI-flavored `usage.total_tokens` field that LiteLLM
injects into Anthropic /v1/messages responses.
@ -195,7 +216,7 @@ async def anthropic_response(
verbose_proxy_logger.exception("litellm.proxy.proxy_server.anthropic_response(): Exception occured - %s", e)
if isinstance(e, ProxyException):
raise
return _anthropic_error_json_response(e, request)
# Extract model_id from request metadata (same as success path)
litellm_metadata: Final = data.get("litellm_metadata", {}) or {}
@ -216,15 +237,18 @@ async def anthropic_response(
)
if isinstance(e, HTTPException):
raise proxy_exception_from_http_exception(e, headers)
return _anthropic_error_json_response(proxy_exception_from_http_exception(e, headers), request)
error_msg: Final = f"{e}"
raise ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
headers=headers,
return _anthropic_error_json_response(
ProxyException(
message=getattr(e, "message", error_msg),
type=getattr(e, "type", "None"),
param=getattr(e, "param", "None"),
code=getattr(e, "status_code", 500),
headers=headers,
),
request,
)

View file

@ -1123,13 +1123,15 @@ async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_grou
user_id=user_id,
user_email=user_id, # We don't have email from group membership
user_alias=None,
teams=[], # Teams will be added separately
metadata={"created_via": created_via},
auto_create_key=False,
user_role=default_role,
)
created_user: Final = await new_user(data=new_user_request)
created_user: Final = await new_user(
data=new_user_request,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
verbose_proxy_logger.info("Created user %s via %s", user_id, created_via)
return created_user
@ -1699,7 +1701,7 @@ async def create_user(
user_id=user_id,
user_email=user_data["user_email"],
user_alias=user_data["user_alias"],
teams=user_data["teams"],
teams=user_data["teams"] or None,
metadata=metadata,
auto_create_key=False,
user_role=resolved_role if admin_group is not None else default_role,
@ -1717,6 +1719,7 @@ async def create_user(
created_user: Final = await new_user(
data=new_user_request,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
scim_user: Final = await ScimTransformations.transform_litellm_user_to_scim_user(user=created_user)
@ -1771,22 +1774,25 @@ async def update_user(
roles=user_data["roles"],
)
# SCIM User.groups is readOnly (RFC 7643 4.1.2): IdPs sync membership via /Groups and send
# no groups or `groups: []` on profile PUTs, so empty means unspecified, not "remove from every team"
target_teams: Final = user_data["teams"] or existing_user.teams
await _handle_team_membership_changes(
user_id=user_id,
existing_teams=existing_user.teams or [],
new_teams=user_data["teams"],
existing_teams=existing_user.teams,
new_teams=target_teams,
)
update_data: Final = {
"user_email": user_data["user_email"],
"user_alias": user_data["user_alias"],
"sso_user_id": user_data["sso_user_id"],
"teams": user_data["teams"],
"teams": target_teams,
"metadata": safe_dumps(metadata),
}
admin_group: Final = await _get_scim_admin_group()
if admin_group is not None:
if admin_group is not None and user_data["teams"]:
update_data["user_role"] = _resolve_scim_user_role(
user.groups or [], admin_group, _default_scim_user_role()
)

View file

@ -1719,6 +1719,12 @@ async def azure_proxy_route(
from abc import ABC, abstractmethod
_VERTEX_LOCATION_REQUIRED_DETAIL: Final = (
"No Vertex AI location for this request. Include /projects/<project>/locations/<location>/ in the "
"route, set vertex_location in default_vertex_config (or DEFAULT_VERTEXAI_LOCATION), or add the "
"model to model_list with use_in_pass_through: true."
)
class BaseVertexAIPassThroughHandler(ABC):
@staticmethod
@ -1726,29 +1732,18 @@ class BaseVertexAIPassThroughHandler(ABC):
def get_default_base_target_url(vertex_location: str | None) -> str:
pass
@staticmethod
@abstractmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str:
pass
class VertexAIDiscoveryPassThroughHandler(BaseVertexAIPassThroughHandler):
@staticmethod
def get_default_base_target_url(vertex_location: str | None) -> str:
return "https://discoveryengine.googleapis.com/"
@staticmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str:
return base_target_url
class VertexAIPassThroughHandler(BaseVertexAIPassThroughHandler):
@staticmethod
def get_default_base_target_url(vertex_location: str | None) -> str:
return get_vertex_base_url(vertex_location)
@staticmethod
def update_base_target_url_with_credential_location(base_target_url: str, vertex_location: str | None) -> str:
if vertex_location is None:
raise HTTPException(status_code=400, detail=_VERTEX_LOCATION_REQUIRED_DETAIL)
return get_vertex_base_url(vertex_location)
@ -1971,10 +1966,8 @@ async def _prepare_vertex_auth_headers(
router_credentials: LiteLLM_ManagedVectorStore | None,
vertex_project: str | None,
vertex_location: str | None,
base_target_url: str | None,
get_vertex_pass_through_handler: BaseVertexAIPassThroughHandler,
user_api_key_dict: UserAPIKeyAuth,
) -> tuple[Mapping[str, str], str | None, bool, str | None, str | None]:
) -> tuple[Mapping[str, str], bool, str | None, str | None]:
"""
Prepare authentication headers for Vertex AI pass-through requests.
@ -1984,15 +1977,12 @@ async def _prepare_vertex_auth_headers(
router_credentials: Optional vector store credentials from registry
vertex_project: Vertex project ID
vertex_location: Vertex location
base_target_url: Base URL for the Vertex AI service
get_vertex_pass_through_handler: Handler for the specific Vertex AI service
user_api_key_dict: The caller's resolved authentication, so only the secret that
authenticated them is stripped on the credential-less branch
Returns:
tuple containing:
- headers: dict - Authentication headers to use
- base_target_url: str | None - Updated base target URL
- headers_passed_through: bool - Whether headers were passed through from request
- vertex_project: str | None - Updated vertex project ID
- vertex_location: str | None - Updated vertex location
@ -2045,14 +2035,8 @@ async def _prepare_vertex_auth_headers(
# Add the Authorization header with vendor credentials
headers["Authorization"] = f"Bearer {auth_header}"
if base_target_url is not None:
base_target_url = get_vertex_pass_through_handler.update_base_target_url_with_credential_location(
base_target_url, vertex_location
)
return (
headers,
base_target_url,
headers_passed_through,
vertex_project,
vertex_location,
@ -2145,12 +2129,9 @@ async def _base_vertex_proxy_route(
location=vertex_location,
)
base_target_url = get_vertex_pass_through_handler.get_default_base_target_url(vertex_location)
# Prepare authentication headers
(
headers,
base_target_url,
headers_passed_through,
vertex_project,
vertex_location,
@ -2160,13 +2141,10 @@ async def _base_vertex_proxy_route(
router_credentials=router_credentials,
vertex_project=vertex_project,
vertex_location=vertex_location,
base_target_url=base_target_url,
get_vertex_pass_through_handler=get_vertex_pass_through_handler,
user_api_key_dict=user_api_key_dict,
)
if base_target_url is None:
base_target_url = get_vertex_base_url(vertex_location)
base_target_url: Final = get_vertex_pass_through_handler.get_default_base_target_url(vertex_location)
request_route: Final = encoded_endpoint
verbose_proxy_logger.debug("request_route %s", request_route)

View file

@ -3,7 +3,8 @@ import collections
import json
import os
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from datetime import date, datetime, timedelta, timezone
from itertools import groupby
from types import MappingProxyType
from typing import (
TYPE_CHECKING,
@ -16,7 +17,6 @@ from typing import (
TypeAlias,
TypedDict,
TypeVar,
cast, # noqa: TID251 # prisma group_by returns untyped aggregate mappings
)
import fastapi
@ -201,16 +201,12 @@ class _SessionSpendStats(NamedTuple):
_SessionSpendMap: TypeAlias = Mapping[tuple[str, str], _SessionSpendStats]
class _SpendSumAggregate(TypedDict, total=False):
spend: ReadOnly[float]
class _SpendGroupByRow(TypedDict):
class _SpendDailySummaryRow(TypedDict):
day: ReadOnly[str]
api_key: ReadOnly[str]
user: ReadOnly[str | None]
model: ReadOnly[str]
startTime: ReadOnly[object]
_sum: ReadOnly[_SpendSumAggregate]
spend: ReadOnly[float]
async def _query_raw(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT]:
@ -251,6 +247,66 @@ def _verification_token_table(prisma_client: PrismaClient) -> _VerificationToken
return VerificationTokenRepository(prisma_client).table
def _spend_logs_daily_summary_sql(
*,
start_date_iso: str,
end_date_iso: str,
api_key: str | None,
request_id: str | None,
user_id: str | None,
) -> tuple[str, tuple[object, ...]]:
filter_params: Final[tuple[tuple[str, object], ...]] = tuple(
(column, value)
for column, value in (
("api_key", api_key),
("request_id", request_id),
('"user"', user_id),
)
if value is not None
)
filter_clauses: Final[tuple[str, ...]] = tuple(
f"AND {column} = ${index}" for index, (column, _) in enumerate(filter_params, start=3)
)
filter_sql: Final = "\n".join(filter_clauses)
sql_query: Final = f"""
SELECT
to_char(date_trunc('day', "startTime"), 'YYYY-MM-DD') AS day,
api_key,
"user",
model,
SUM(spend) AS spend
FROM "LiteLLM_SpendLogs"
WHERE "startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') AND "startTime" <= ($2::timestamptz AT TIME ZONE 'UTC')
{filter_sql}
GROUP BY 1, 2, 3, 4
ORDER BY 1
"""
params: Final[tuple[object, ...]] = (
start_date_iso,
end_date_iso,
*(value for _, value in filter_params),
)
return sql_query, params
def _sum_spend_by(
rows: Sequence[_SpendDailySummaryRow], column: Literal["api_key", "user", "model"]
) -> Mapping[str | None, float]:
keys: Final = frozenset(row[column] for row in rows)
return {key: sum(float(row["spend"]) for row in rows if row[column] == key) for key in keys}
def _daily_summary_item(summary_date: date, rows: Sequence[_SpendDailySummaryRow]) -> Mapping[str, object]:
api_key_spend: Final = {key: value for key, value in _sum_spend_by(rows, "api_key").items() if key is not None}
return {
**api_key_spend,
"startTime": summary_date,
"spend": sum(float(row["spend"]) for row in rows),
"users": _sum_spend_by(rows, "user"),
"models": _sum_spend_by(rows, "model"),
}
async def _find_spend_logs(
prisma_client: PrismaClient,
where: Mapping[str, object],
@ -3266,18 +3322,22 @@ async def view_spend_logs(
start_date_iso: Final = start_date_obj.isoformat()
end_date_iso: Final = end_date_obj.isoformat()
filter_query: Final = {
filter_query: Final[
dict[str, object]
] = { # mutable-ok: legacy filters are extended for optional parameters
"startTime": {
"gte": start_date_iso, # Greater than or equal to Start Date
"lte": end_date_iso, # Less than or equal to End Date
}
}
summary_api_key: Final[str | None] = (
prisma_client.hash_token(token=api_key)
if api_key is not None and api_key.startswith("sk-")
else api_key
)
if api_key is not None and isinstance(api_key, str):
if api_key.startswith("sk-"):
filter_query["api_key"] = prisma_client.hash_token(token=api_key)
else:
filter_query["api_key"] = api_key
filter_query["api_key"] = summary_api_key
if request_id is not None and isinstance(request_id, str):
filter_query["request_id"] = request_id
if user_id is not None and isinstance(user_id, str):
@ -3296,58 +3356,34 @@ async def view_spend_logs(
return data
# Legacy behavior: return summarized data (when summarize=true)
# SQL query
response: Final = await SpendLogsRepository(prisma_client).table.group_by(
by=["api_key", "user", "model", "startTime"],
where=filter_query,
sum={
"spend": True,
},
summary_sql_and_params: Final = _spend_logs_daily_summary_sql(
start_date_iso=start_date_iso,
end_date_iso=end_date_iso,
api_key=summary_api_key,
request_id=request_id,
user_id=user_id,
)
sql_query, params = summary_sql_and_params
rows: Final[Sequence[_SpendDailySummaryRow]] = await _query_raw(prisma_client, sql_query, *params)
if len(rows) == 0:
return [] # pyright: ignore[reportUnknownVariableType] # empty summary has no element type
if isinstance(response, list) and len(response) > 0 and isinstance(response[0], dict):
spend_rows: Final = cast(Sequence[_SpendGroupByRow], response) # cast-ok: by/sum fix the shape
result: Final[dict] = {}
for record in spend_rows:
dt_object = datetime.strptime(str(record["startTime"]), "%Y-%m-%dT%H:%M:%S.%fZ")
date = dt_object.date()
if date not in result:
result[date] = {"users": {}, "models": {}}
api_key = record["api_key"]
user_id = record["user"]
model = record["model"]
result[date]["spend"] = result[date].get("spend", 0) + record.get("_sum", {}).get("spend", 0)
result[date][api_key] = result[date].get(api_key, 0) + record.get("_sum", {}).get("spend", 0)
result[date]["users"][user_id] = result[date]["users"].get(user_id, 0) + record.get("_sum", {}).get(
"spend", 0
)
result[date]["models"][model] = result[date]["models"].get(model, 0) + record.get("_sum", {}).get(
"spend", 0
)
return_list: Final = []
final_date = None
for k, v in sorted(result.items()):
return_list.append({**v, "startTime": k})
final_date = k
end_date_date: Final = end_date_obj.date()
if final_date is not None and final_date < end_date_date:
current_date = final_date + timedelta(days=1)
while current_date <= end_date_date:
# Represent current_date as string because original response has it this way
return_list.append(
{
"startTime": current_date,
"spend": 0,
"users": {},
"models": {},
}
) # If no data, will stay as zero
current_date += timedelta(days=1) # Move on to the next day
return return_list
return response
summary_items: Final = tuple(
_daily_summary_item(date.fromisoformat(day), tuple(day_rows))
for day, day_rows in groupby(rows, key=lambda row: row["day"])
)
final_date: Final = date.fromisoformat(rows[-1]["day"])
end_date_date: Final = end_date_obj.date()
padding: Final[tuple[Mapping[str, object], ...]] = tuple(
{
"startTime": final_date + timedelta(days=offset),
"spend": 0,
"users": {},
"models": {},
}
for offset in range(1, (end_date_date - final_date).days + 1)
)
return [*summary_items, *padding]
else:
scoped_filter: Final[dict[str, str]] = {}

View file

@ -197,6 +197,7 @@ from litellm.router_utils.router_callbacks.track_deployment_metrics import (
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolParam,
FileTypes,
OpenAIFileObject,
OpenAIFilesPurpose,
@ -11760,7 +11761,7 @@ class Router:
self,
messages: list[dict[str, str]] | None,
input: str | list | None,
instructions: str | None = None,
request_kwargs: Mapping[str, object] | None = None,
) -> int:
"""
Count input tokens for context-window pre-call checks.
@ -11770,9 +11771,28 @@ class Router:
The Responses payload is normalized to chat messages via the shared
LiteLLMCompletionResponsesConfig transform so the same token_counter path covers
both API surfaces and `instructions` tokens are included in the count.
Prompt content the message list never carries is read from `request_kwargs`:
`tools` (Chat Completions, Responses and Anthropic Messages shapes) and the
Anthropic Messages top-level `system` block.
"""
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
anthropic_system_to_openai_message,
)
extras: Final = request_kwargs if request_kwargs is not None else MappingProxyType({})
raw_instructions: Final = extras.get("instructions")
instructions: Final = raw_instructions if isinstance(raw_instructions, str) else None
raw_tools: Final = extras.get("tools")
tools: Final = (
cast(list[ChatCompletionToolParam], raw_tools) # cast-ok: token_counter formats any tool dict shape
if isinstance(raw_tools, list) and raw_tools
else None
)
system_message: Final = anthropic_system_to_openai_message(extras.get("system"))
if messages is not None:
return litellm.token_counter(messages=messages)
counted_messages: Final = (system_message, *messages) if system_message is not None else messages
return litellm.token_counter(messages=counted_messages, tools=tools)
if input is not None:
from openai.types.responses.response_create_params import ResponseInputParam
@ -11785,7 +11805,10 @@ class Router:
input=typed_input,
responses_api_request={"instructions": instructions} if instructions is not None else {},
)
return litellm.token_counter(messages=cast(list, input_messages)) # cast-ok: transformed chat messages
return litellm.token_counter(
messages=cast(list, input_messages), # cast-ok: transformed chat messages
tools=tools,
)
raise ValueError("Either messages or input must be provided to count tokens")
def _deployment_max_input_tokens(self, model: str, deployment: Mapping[str, object]) -> int | None:
@ -11831,14 +11854,13 @@ class Router:
"""
if messages is None and input is None:
return None
raw_instructions: Final = request_kwargs.get("instructions") if request_kwargs else None
try:
if not self._pre_call_checks_need_token_count(model, healthy_deployments):
return None
return await asyncify(self._count_pre_call_check_tokens)(
messages=cast(list[dict[str, str]] | None, messages), # cast-ok: forwarded to the sync counter
input=cast(str | list | None, input), # cast-ok: forwarded to the sync counter
instructions=raw_instructions if isinstance(raw_instructions, str) else None,
request_kwargs=request_kwargs,
)
except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request
verbose_router_logger.error(
@ -11885,8 +11907,6 @@ class Router:
_rate_limit_error = False
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs)
raw_instructions: Final = request_kwargs.get("instructions") if request_kwargs else None
instructions: Final = raw_instructions if isinstance(raw_instructions, str) else None
has_countable_input: Final = messages is not None or input is not None
## get model group RPM ##
@ -11917,7 +11937,7 @@ class Router:
return _returned_deployments
try:
input_tokens = self._count_pre_call_check_tokens(
messages=messages, input=input, instructions=instructions
messages=messages, input=input, request_kwargs=request_kwargs
)
except Exception as e:
verbose_router_logger.error(

View file

@ -96,7 +96,7 @@
"limit": 10
},
"DTZ007": {
"limit": 17
"limit": 6
},
"DTZ011": {
"limit": 3

View file

@ -763,7 +763,7 @@ class _MigrateDeployHarness:
"_resolve_specific_migration",
staticmethod(self.resolved.append),
)
monkeypatch.setattr(utils_module.subprocess, "run", self._fake_run)
monkeypatch.setattr(utils_module.prisma_toolchain, "run_prisma", self._fake_run)
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
self.baseline_succeeds = True

View file

@ -125,11 +125,12 @@ class TestBlockedResponseUsage:
mock_logging.post_call_failure_hook.assert_awaited_once()
class TestProxyExceptionPassthrough:
class TestProxyExceptionAnthropicEnvelope:
@pytest.mark.asyncio
async def test_anthropic_response_reraises_proxy_exception_unwrapped(self):
"""A 400 ProxyException from request validation must surface as-is,
not be re-wrapped into a code-500 ProxyException."""
async def test_anthropic_response_maps_proxy_exception_to_anthropic_envelope(self):
"""LIT-6468: a 400 ProxyException from request validation must surface as
Anthropic's documented {"type": "error", "error": {...}} envelope with the
original status and message, not the OpenAI {"error": {...}} envelope."""
import litellm.proxy.anthropic_endpoints.endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy._types import ProxyErrorTypes, ProxyException
@ -140,6 +141,8 @@ class TestProxyExceptionPassthrough:
param="metadata",
code=400,
)
request = MagicMock()
request.headers = {"x-request-id": "req_test_6468"}
with (
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})),
@ -151,30 +154,61 @@ class TestProxyExceptionPassthrough:
patch.object(proxy_server, "proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock()
with pytest.raises(ProxyException) as exc_info:
await ep.anthropic_response(
fastapi_response=MagicMock(),
request=MagicMock(),
user_api_key_dict=MagicMock(),
)
response = await ep.anthropic_response(
fastapi_response=MagicMock(),
request=request,
user_api_key_dict=MagicMock(),
)
assert exc_info.value is exc
assert exc_info.value.code == "400"
assert exc_info.value.param == "metadata"
assert response.status_code == 400
body = json.loads(response.body)
assert body == {
"type": "error",
"error": {
"type": "invalid_request_error",
"message": "Invalid type for 'metadata': expected an object, but got a string instead.",
},
"request_id": "req_test_6468",
}
mock_logging.post_call_failure_hook.assert_awaited_once()
@pytest.mark.asyncio
async def test_anthropic_response_maps_429_to_rate_limit_error(self):
"""The Anthropic error type follows the status code (429 -> rate_limit_error),
and a code-less exception falls back to 500 api_error."""
import litellm.proxy.anthropic_endpoints.endpoints as ep
from litellm.proxy._types import ProxyException
request = MagicMock()
request.headers = {}
response = ep._anthropic_error_json_response(
ProxyException(message="Rate limit exceeded", type="rate_limit_error", param=None, code=429),
request,
)
assert response.status_code == 429
assert json.loads(response.body)["error"]["type"] == "rate_limit_error"
fallback = ep._anthropic_error_json_response(
ProxyException(message="boom", type="None", param=None, code=None),
request,
)
assert fallback.status_code == 500
assert json.loads(fallback.body)["error"]["type"] == "api_error"
class TestHttpExceptionDictDetail:
@pytest.mark.asyncio
async def test_anthropic_response_serializes_dict_detail_http_exception(self):
"""LIT-6466: a post_call guardrail's HTTPException(detail=<dict>) must
surface with a clean message plus provider_specific_fields, matching
/v1/chat/completions and /v1/responses, not the str() of the exception."""
"""LIT-6466 + LIT-6468: a post_call guardrail's HTTPException(detail=<dict>)
must surface as Anthropic's {"type": "error", "error": {...}} envelope with
the guardrail's clean message plus provider_specific_fields, not the str()
of the exception and not the OpenAI envelope."""
from fastapi import HTTPException
import litellm.proxy.anthropic_endpoints.endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy._types import UserAPIKeyAuth
detail = {
"error": "Content blocked: keyword 'kumquat' detected",
@ -182,6 +216,8 @@ class TestHttpExceptionDictDetail:
"guardrail": "keyword-block",
}
exc = HTTPException(status_code=400, detail=detail)
request = MagicMock()
request.headers = {}
with (
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: endpoint reads the body via a module function; no injection seam
@ -193,17 +229,19 @@ class TestHttpExceptionDictDetail:
patch.object(proxy_server, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global imported at call time; no injection seam
):
mock_logging.post_call_failure_hook = AsyncMock()
with pytest.raises(ProxyException) as exc_info:
await ep.anthropic_response(
fastapi_response=MagicMock(),
request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(),
)
response = await ep.anthropic_response(
fastapi_response=MagicMock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
assert exc_info.value.message == "Content blocked: keyword 'kumquat' detected"
assert "{'error'" not in exc_info.value.message
assert exc_info.value.provider_specific_fields == detail
assert exc_info.value.code == "400"
assert response.status_code == 400
body = json.loads(response.body)
assert body["type"] == "error"
assert body["error"]["type"] == "invalid_request_error"
assert body["error"]["message"] == "Content blocked: keyword 'kumquat' detected"
assert "{'error'" not in body["error"]["message"]
assert body["error"]["provider_specific_fields"] == detail
mock_logging.post_call_failure_hook.assert_awaited_once()
@ -215,7 +253,7 @@ class TestFailureHookRequestData:
handler must pass that replaced dict, not the raw request body dict."""
import litellm.proxy.anthropic_endpoints.endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy._types import UserAPIKeyAuth
captured = {}
@ -224,18 +262,23 @@ class TestFailureHookRequestData:
captured["processor_data"] = self.data
raise RuntimeError("provider timeout")
request = MagicMock()
request.headers = {}
with (
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={"model": "claude-sonnet"})),
patch.object(ep.ProxyBaseLLMRequestProcessing, "base_process_llm_request", new=fake_process),
patch.object(proxy_server, "proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock()
with pytest.raises(ProxyException):
await ep.anthropic_response(
fastapi_response=MagicMock(),
request=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(),
)
response = await ep.anthropic_response(
fastapi_response=MagicMock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
assert response.status_code == 500
assert json.loads(response.body)["error"]["message"] == "provider timeout"
hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert hook_request_data is captured["processor_data"]

View file

@ -19,11 +19,13 @@ from litellm.proxy._types import (
NewUserResponse,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints.scim.scim_v2 import (
SCIMRosterSyncError,
UserProvisionerHelpers,
_apply_group_patch_updates,
_create_user_if_not_exists,
_extract_group_member_ids,
_extract_ids_from_path_filter,
_handle_group_membership_changes,
@ -37,8 +39,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
delete_group,
delete_user,
get_groups,
get_users,
get_service_provider_config,
get_users,
merge_placeholder,
patch_group,
patch_team_membership,
@ -304,6 +306,85 @@ async def test_create_user_uses_default_internal_user_params_role(mocker, monkey
assert called_args.user_role == LitellmUserRoles.PROXY_ADMIN
def _mock_scim_create_user_deps(mocker: MockerFixture, scim_user: SCIMUser) -> AsyncMock:
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=scim_user),
)
return mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2.new_user",
AsyncMock(return_value=NewUserRequest(user_id=scim_user.userName)),
)
@pytest.mark.asyncio
async def test_create_user_without_groups_defers_to_default_team(mocker: MockerFixture, monkeypatch):
"""IdPs omit groups on POST /Users; teams must stay unset so new_user applies default_internal_user_params.teams"""
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="new-user",
emails=[SCIMUserEmail(value="new@example.com")],
)
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = _mock_scim_create_user_deps(mocker, scim_user)
await create_user(user=scim_user)
assert new_user_mock.call_args.kwargs["data"].teams is None
assert new_user_mock.call_args.kwargs["user_api_key_dict"] == UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
@pytest.mark.asyncio
async def test_create_user_with_groups_keeps_idp_teams(mocker: MockerFixture, monkeypatch):
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="new-user",
emails=[SCIMUserEmail(value="new@example.com")],
groups=[SCIMUserGroup(value="idp-team")],
)
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = _mock_scim_create_user_deps(mocker, scim_user)
await create_user(user=scim_user)
assert new_user_mock.call_args.kwargs["data"].teams == ["idp-team"]
@pytest.mark.asyncio
async def test_create_user_if_not_exists_defers_to_default_team(mocker: MockerFixture, monkeypatch):
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = mocker.patch( # test-quality-ok: new_user is imported inside the helper, not injectable
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
AsyncMock(return_value=NewUserResponse(user_id="group-user", key="k")),
)
created = await _create_user_if_not_exists(user_id="group-user")
assert created is not None
assert new_user_mock.call_args.kwargs["data"].teams is None
assert new_user_mock.call_args.kwargs["user_api_key_dict"] == UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
@pytest.mark.asyncio
async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeypatch):
"""
@ -1176,6 +1257,67 @@ async def test_update_user_success(mocker):
assert call_args[1]["data"]["teams"] == ["new-team"]
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", [None, []], ids=["groups-omitted", "groups-empty"])
async def test_update_user_without_groups_preserves_memberships_and_role(mocker, monkeypatch, groups):
"""Okta profile PUTs carry no `groups` or `groups: []`; neither may drop teams (and their keys) or recompute role"""
from litellm.proxy.proxy_server import proxy_config
async def mock_get_config():
return {"litellm_settings": {"scim_admin_group": "litellm-admins"}}
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False)
existing_user = mocker.MagicMock()
existing_user.teams = ["litellm-admins", "engineering"]
existing_user.metadata = {}
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="okta-user",
name=SCIMUserName(familyName="Renamed", givenName="Okta"),
emails=[SCIMUserEmail(value="okta@example.com")],
**({} if groups is None else {"groups": groups}),
)
response_scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="okta-user",
userName="okta-user",
emails=[SCIMUserEmail(value="okta@example.com")],
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={"user_id": "okta-user"})
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists",
AsyncMock(return_value=existing_user),
)
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=response_scim_user),
)
patch_membership = mocker.patch( # test-quality-ok: roster writes are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
AsyncMock(),
)
result = await update_user(user_id="okta-user", user=scim_user)
assert result == response_scim_user
patch_membership.assert_not_awaited()
update_data = mock_prisma_client.db.litellm_usertable.update.call_args.kwargs["data"]
assert update_data["teams"] == ["litellm-admins", "engineering"]
assert "user_role" not in update_data
@pytest.mark.asyncio
async def test_update_user_not_found(mocker):
"""Should raise 404 when user doesn't exist"""

View file

@ -297,9 +297,6 @@ class TestVertexAIPassThroughHandler:
mock_handler.get_default_base_target_url.return_value = (
f"https://{test_location}-aiplatform.googleapis.com/"
)
mock_handler.update_base_target_url_with_credential_location = Mock(
return_value=f"https://{test_location}-aiplatform.googleapis.com/"
)
mock_get_handler.return_value = mock_handler
# Mock create_pass_through_route to return a function that returns a mock response
@ -398,9 +395,8 @@ class TestVertexAIPassThroughHandler:
# Mock the vertex handler for global location
mock_handler = Mock()
mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com/"
mock_handler.update_base_target_url_with_credential_location = Mock(
return_value="https://aiplatform.googleapis.com/"
mock_handler.get_default_base_target_url.return_value = (
"https://aiplatform.googleapis.com/"
)
mock_get_handler.return_value = mock_handler
@ -498,9 +494,6 @@ class TestVertexAIPassThroughHandler:
mock_handler.get_default_base_target_url.return_value = (
f"https://{default_location}-aiplatform.googleapis.com/"
)
mock_handler.update_base_target_url_with_credential_location = Mock(
return_value=f"https://{default_location}-aiplatform.googleapis.com/"
)
mock_get_handler.return_value = mock_handler
# Mock create_pass_through_route to return a function that returns a mock response
@ -1225,9 +1218,8 @@ class TestVertexAIDiscoveryPassThroughHandler:
# Mock the discovery handler
mock_handler = Mock()
mock_handler.get_default_base_target_url.return_value = "https://discoveryengine.googleapis.com"
mock_handler.update_base_target_url_with_credential_location = Mock(
return_value="https://discoveryengine.googleapis.com"
mock_handler.get_default_base_target_url.return_value = (
"https://discoveryengine.googleapis.com"
)
mock_get_handler.return_value = mock_handler
@ -3442,7 +3434,6 @@ class TestVertexRawPredictStreamingClassification:
base_url = "https://us-east5-aiplatform.googleapis.com/"
mock_handler = Mock()
mock_handler.get_default_base_target_url.return_value = base_url
mock_handler.update_base_target_url_with_credential_location = Mock(return_value=base_url)
module = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints"
with (
@ -4030,6 +4021,126 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak:
assert "sk-master-1234" not in " ".join(f"{name}:{value}" for name, value in forwarded.items())
class TestVertexPassthroughDefaultLocationOnShortRoutes:
PROJECT = "test-project"
SHORT_ROUTE = "publishers/google/models/gemini-2.5-flash:generateContent"
@staticmethod
def _forwarder() -> Mock:
return Mock(return_value=AsyncMock(return_value={"status": "success"}))
async def _forward(
self,
monkeypatch,
endpoint: str,
default_config: dict | None,
headers: list[tuple[bytes, bytes]],
forwarder: Mock,
) -> None:
from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import (
PassthroughEndpointRouter,
)
async def receive():
return {"type": "http.request", "body": b"{}", "more_body": False}
request: Final = Request(
{
"type": "http",
"method": "POST",
"path": f"/vertex_ai/{endpoint}",
"headers": headers,
"query_string": b"",
},
receive=receive,
)
router: Final = PassthroughEndpointRouter()
if default_config is not None:
router.set_default_vertex_config(dict(default_config))
module: Final = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints"
monkeypatch.setattr(f"{module}.passthrough_endpoint_router", router)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
mock_credentials: Final = Mock()
mock_credentials.token = "test-token"
caller: Final = UserAPIKeyAuth(api_key="test-key")
with (
mock.patch( # test-quality-ok: the route mints its Google token through its own VertexBase, nothing injects the credential loader
"litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth",
return_value=(mock_credentials, self.PROJECT),
),
mock.patch(f"{module}.create_pass_through_route", new=forwarder),
mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=caller)),
):
await vertex_proxy_route(
endpoint=endpoint,
request=request,
fastapi_response=Response(),
user_api_key_dict=caller,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("endpoint", "location", "expected_target"),
[
(
SHORT_ROUTE,
"global",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/" + SHORT_ROUTE,
),
(
f"v1/{SHORT_ROUTE}",
"global",
"https://aiplatform.googleapis.com/v1/projects/test-project/locations/global/" + SHORT_ROUTE,
),
(
f"v1beta1/{SHORT_ROUTE}",
"global",
"https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/" + SHORT_ROUTE,
),
(
SHORT_ROUTE,
"us-central1",
"https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/"
+ SHORT_ROUTE,
),
],
)
async def test_default_vertex_config_location_fills_routes_without_project_and_location(
self, monkeypatch, endpoint, location, expected_target
):
forwarder: Final = self._forwarder()
await self._forward(
monkeypatch,
endpoint,
{"vertex_project": self.PROJECT, "vertex_location": location, "vertex_credentials": "test-creds"},
[(b"content-type", b"application/json"), (b"authorization", b"Bearer test-key")],
forwarder,
)
forwarded: Final = forwarder.call_args.kwargs
assert str(forwarded["target"]) == expected_target
assert forwarded["custom_headers"]["Authorization"] == "Bearer test-token"
@pytest.mark.asyncio
@pytest.mark.parametrize(
("default_config", "headers"),
[
(None, [(b"content-type", b"application/json"), (b"authorization", b"Bearer ya29.byo-google-oauth")]),
(
{"vertex_project": PROJECT, "vertex_credentials": "test-creds"},
[(b"content-type", b"application/json"), (b"authorization", b"Bearer test-key")],
),
],
)
async def test_no_location_anywhere_is_a_400_not_a_500(self, monkeypatch, default_config, headers):
forwarder: Final = self._forwarder()
with pytest.raises(HTTPException) as raised:
await self._forward(monkeypatch, self.SHORT_ROUTE, default_config, headers, forwarder)
forwarder.assert_not_called()
assert raised.value.status_code == 400
assert "/projects/<project>/locations/<location>/" in str(raised.value.detail)
assert "default_vertex_config" in str(raised.value.detail)
class TestGetAzureAISearchIndexFromEndpoint:
"""The operable index is only the segment right after ``indexes``.

View file

@ -4,6 +4,7 @@ import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
VertexAIPassThroughHandler,
_base_vertex_proxy_route,
_upstream_headers_for_vertex_route,
)
@ -20,6 +21,7 @@ async def test_vertex_passthrough_load_balancing():
mock_request = MagicMock()
mock_response = MagicMock()
mock_handler = MagicMock()
mock_handler.get_default_base_target_url.return_value = "https://test.url"
# Mock the router
mock_router = MagicMock()
@ -68,7 +70,6 @@ async def test_vertex_passthrough_load_balancing():
mock_pt_router.get_vertex_credentials.return_value = MagicMock()
mock_prep_headers.return_value = (
{},
"https://test.url",
False,
"test-project-lb",
"us-central1-lb",
@ -290,12 +291,6 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header():
mock_vertex_credentials.vertex_location = "us-central1"
mock_vertex_credentials.vertex_credentials = "test-credentials"
# Create mock handler
mock_handler = MagicMock()
mock_handler.update_base_target_url_with_credential_location.return_value = (
"https://us-central1-aiplatform.googleapis.com"
)
with (
patch.object(
VertexBase,
@ -313,7 +308,6 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header():
# Call the function
(
headers,
base_target_url,
headers_passed_through,
vertex_project,
vertex_location,
@ -323,8 +317,6 @@ async def test_vertex_passthrough_forwards_anthropic_beta_header():
router_credentials=None,
vertex_project="test-project",
vertex_location="us-central1",
base_target_url="https://us-central1-aiplatform.googleapis.com",
get_vertex_pass_through_handler=mock_handler,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-litellm-secret-key"),
)
@ -394,7 +386,6 @@ async def test_vertex_passthrough_drops_anthropic_beta_only_on_count_tokens(
"content-type": "application/json",
"Authorization": "Bearer vertex-access-token",
},
"https://aiplatform.googleapis.com",
False,
"test-project",
"global",
@ -406,7 +397,7 @@ async def test_vertex_passthrough_drops_anthropic_beta_only_on_count_tokens(
endpoint=f"{VERTEX_ANTHROPIC_MODELS_PREFIX}{model_segment}",
request=MagicMock(),
fastapi_response=MagicMock(),
get_vertex_pass_through_handler=MagicMock(),
get_vertex_pass_through_handler=VertexAIPassThroughHandler(),
)
upstream_headers = mock_create_route.call_args.kwargs["custom_headers"]
@ -473,12 +464,6 @@ async def test_vertex_passthrough_does_not_forward_litellm_auth_token():
mock_vertex_credentials.vertex_location = "us-central1"
mock_vertex_credentials.vertex_credentials = "test-credentials"
# Create mock handler
mock_handler = MagicMock()
mock_handler.update_base_target_url_with_credential_location.return_value = (
"https://us-central1-aiplatform.googleapis.com"
)
with (
patch.object(
VertexBase,
@ -495,7 +480,6 @@ async def test_vertex_passthrough_does_not_forward_litellm_auth_token():
(
headers,
_base_target_url,
_headers_passed_through,
_vertex_project,
_vertex_location,
@ -505,8 +489,6 @@ async def test_vertex_passthrough_does_not_forward_litellm_auth_token():
router_credentials=None,
vertex_project="test-project",
vertex_location="us-central1",
base_target_url="https://us-central1-aiplatform.googleapis.com",
get_vertex_pass_through_handler=mock_handler,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-litellm-secret-key"),
)
@ -742,7 +724,6 @@ async def test_vertex_passthrough_custom_model_name_replaced_in_url():
mock_pt_router.get_vertex_credentials.return_value = MagicMock()
mock_prep_headers.return_value = (
{},
"https://global-aiplatform.googleapis.com",
False,
"nv-gcpllmgwit-20250411173346",
"global",

View file

@ -5,14 +5,12 @@ import hashlib
import json
import re
from datetime import timezone
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from unittest.mock import AsyncMock, MagicMock, patch
import litellm
import litellm.proxy.proxy_server as ps
@ -3325,7 +3323,7 @@ def _compare_nested_dicts(
return differences
# Check for keys in actual but not in expected
for key in actual.keys():
for key in actual:
current_path = f"{path}.{key}" if path else key
if current_path not in ignore_keys and key not in expected:
differences.append(f"Extra key in actual: {current_path}")
@ -3495,24 +3493,22 @@ async def test_view_spend_logs_summarize_parameter(client, monkeypatch):
# Return individual log entries when summarize=false
return mock_spend_logs
async def group_by(self, *args, **kwargs):
# Return grouped data when summarize=true
# Simplified mock response for grouped data
async def query_raw(self, sql_query, *params):
yesterday = datetime.datetime.now(timezone.utc) - timedelta(days=1)
return [
{
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-3.5-turbo",
"startTime": yesterday.strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
"_sum": {"spend": 0.05},
"day": yesterday.date().isoformat(),
"spend": 0.05,
},
{
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-4",
"startTime": yesterday.strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
"_sum": {"spend": 0.10},
"day": yesterday.date().isoformat(),
"spend": 0.10,
},
]
@ -3850,47 +3846,30 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
"""
from datetime import datetime, timedelta, timezone
# This simulates the summarized data that Prisma's `group_by` would return.
mock_summarized_response = [
{
"api_key": "sk-test-key",
"user": "test_user_1",
"model": "gpt-4",
"startTime": (datetime.now(timezone.utc) - timedelta(days=1)).strftime(
"%Y-%m-%dT%H:%M:%S.%fZ"
),
"_sum": {"spend": 0.15},
"day": (datetime.now(timezone.utc) - timedelta(days=1)).date().isoformat(),
"spend": 0.15,
}
]
# This mock class will replace the real Prisma client.
class MockDB:
def __init__(self):
self.litellm_spendlogs = self
async def group_by(self, *args, **kwargs):
# We assert that the `gte` and `lte` values are strings in ISO format.
# If they were datetime objects, this test would fail.
where_clause = kwargs.get("where", {})
start_time_filter = where_clause.get("startTime", {})
assert "gte" in start_time_filter
assert "lte" in start_time_filter
assert isinstance(start_time_filter["gte"], str)
assert isinstance(start_time_filter["lte"], str)
assert "T" in start_time_filter["gte"] # Check for ISO format 'T' separator
# If the assertions pass, return the mock response.
async def query_raw(self, sql_query, *params):
assert isinstance(params[0], str)
assert isinstance(params[1], str)
assert "T" in params[0]
assert "T" in params[1]
return mock_summarized_response
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
# Apply the monkeypatch to replace the real prisma_client with our mock.
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
# Define a date range for the test.
start_date = (datetime.now(timezone.utc) - timedelta(days=2)).strftime("%Y-%m-%d")
end_date = datetime.now(timezone.utc).strftime("%Y-%m-%d")
@ -3898,8 +3877,6 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
# Call the endpoint with both start and end dates.
# We don't need `summarize=true` as it's the default.
response = client.get(
"/spend/logs",
params={
@ -3909,11 +3886,9 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
headers={"Authorization": "Bearer sk-test"},
)
# ASSERTIONS
assert response.status_code == 200
data = response.json()
# Check that the response is not empty and has the summarized structure.
assert isinstance(data, list)
assert len(data) > 0
assert "startTime" in data[0]
@ -3924,6 +3899,183 @@ async def test_view_spend_logs_with_date_range_summarized(client, monkeypatch):
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_view_spend_logs_summarize_groups_by_day_in_sql(client, monkeypatch):
mock_rows = [
{
"day": "2024-01-01",
"api_key": "hashed::sk-abc",
"user": "u1",
"model": "gpt-4",
"spend": 0.1,
},
{
"day": "2024-01-01",
"api_key": "hashed::sk-abc",
"user": "u1",
"model": "gpt-4o",
"spend": 0.2,
},
]
class MockDB:
def __init__(self):
self.captured_sql = None
self.captured_params = None
async def query_raw(self, sql_query, *params):
self.captured_sql = sql_query
self.captured_params = params
return mock_rows
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
def hash_token(self, token):
return "hashed::" + token
mock_prisma_client = MockPrismaClient()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs",
params={
"start_date": "2024-01-01",
"end_date": "2024-01-03",
"api_key": "sk-abc",
"request_id": "req-123",
"user_id": "u1",
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
sql = mock_prisma_client.db.captured_sql
assert "date_trunc('day'" in sql
assert "GROUP BY" in sql
assert "find_many" not in sql
assert not hasattr(mock_prisma_client.db, "group_by")
assert mock_prisma_client.db.captured_params == (
"2024-01-01T00:00:00+00:00",
"2024-01-03T00:00:00+00:00",
"hashed::sk-abc",
"req-123",
"u1",
)
assert len(data) == 3
assert data[0]["startTime"] == "2024-01-01"
assert data[0]["spend"] == pytest.approx(0.3)
assert data[0]["models"] == {"gpt-4": 0.1, "gpt-4o": 0.2}
assert data[0]["users"] == {"u1": pytest.approx(0.3)}
assert data[0]["hashed::sk-abc"] == pytest.approx(0.3)
assert data[1] == {
"startTime": "2024-01-02",
"spend": 0,
"users": {},
"models": {},
}
assert data[2] == {
"startTime": "2024-01-03",
"spend": 0,
"users": {},
"models": {},
}
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_view_spend_logs_summarize_empty_rows(client, monkeypatch):
class MockDB:
async def query_raw(self, sql_query, *params):
return []
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs",
params={"start_date": "2024-01-01", "end_date": "2024-01-01"},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
assert response.json() == []
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_view_spend_logs_summarize_unhashed_api_key_without_padding(client, monkeypatch):
mock_rows = [
{
"day": "2024-01-01",
"api_key": "plain-key",
"user": "u1",
"model": "gpt-4",
"spend": 0.4,
}
]
class MockDB:
def __init__(self):
self.captured_params = None
async def query_raw(self, sql_query, *params):
self.captured_params = params
return mock_rows
class MockPrismaClient:
def __init__(self):
self.db = MockDB()
mock_prisma_client = MockPrismaClient()
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
try:
response = client.get(
"/spend/logs",
params={
"start_date": "2024-01-01",
"end_date": "2024-01-01",
"api_key": "plain-key",
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200
data = response.json()
assert mock_prisma_client.db.captured_params == (
"2024-01-01T00:00:00+00:00",
"2024-01-01T00:00:00+00:00",
"plain-key",
)
assert data == [
{
"startTime": "2024-01-01",
"spend": pytest.approx(0.4),
"plain-key": pytest.approx(0.4),
"users": {"u1": pytest.approx(0.4)},
"models": {"gpt-4": pytest.approx(0.4)},
}
]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_with_error_code(client):
"""Test filtering spend logs by error code"""
@ -4832,13 +4984,14 @@ class _CaptureFilterDB:
def __init__(self):
self.litellm_spendlogs = self
self.captured_where = None
self.captured_params = None
async def find_many(self, *args, **kwargs):
self.captured_where = kwargs.get("where")
return []
async def group_by(self, *args, **kwargs):
self.captured_where = kwargs.get("where")
async def query_raw(self, sql_query, *params):
self.captured_params = params
return []

View file

@ -3855,7 +3855,7 @@ def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch):
input_only_tokens = router._count_pre_call_check_tokens(messages=None, input=short_input)
with_instructions_tokens = router._count_pre_call_check_tokens(
messages=None, input=short_input, instructions=long_instructions
messages=None, input=short_input, request_kwargs={"instructions": long_instructions}
)
assert with_instructions_tokens > input_only_tokens
@ -3871,6 +3871,164 @@ def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch):
)
_OVERSIZED_TOOL_DESCRIPTION = "look up the answer in the knowledge base. " * 40
@pytest.mark.parametrize(
"prompt_kwargs, tool",
[
pytest.param(
{"messages": [{"role": "user", "content": "hi"}]},
{
"type": "function",
"function": {
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"parameters": {"type": "object", "properties": {"q": {"type": "string"}}},
},
},
id="chat_completions_tool",
),
pytest.param(
{"input": "hi"},
{
"type": "function",
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"parameters": {"type": "object", "properties": {"q": {"type": "string"}}},
},
id="responses_tool",
),
pytest.param(
{"messages": [{"role": "user", "content": "hi"}]},
{
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"input_schema": {"type": "object", "properties": {"q": {"type": "string"}}},
},
id="anthropic_messages_tool",
),
],
)
def test_pre_call_checks_counts_tool_definition_tokens(monkeypatch, prompt_kwargs, tool):
"""
Tool definitions are sent to the model as prompt tokens but never appear in
`messages` or `input`. A request whose prompt alone fits the context window but
whose prompt plus `tools` exceeds it must be rejected before dispatch, for the
Chat Completions, Responses and Anthropic Messages tool shapes alike.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
]
prompt_only_tokens = router._count_pre_call_check_tokens(
messages=prompt_kwargs.get("messages"), input=prompt_kwargs.get("input")
)
monkeypatch.setattr(
router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": prompt_only_tokens}
)
assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, **prompt_kwargs)) == 1
with pytest.raises(litellm.ContextWindowExceededError):
router._pre_call_checks(
model="m",
healthy_deployments=deployments,
request_kwargs={"tools": [tool]},
**prompt_kwargs,
)
@pytest.mark.parametrize(
"system",
[
pytest.param("You are a meticulous assistant. " * 40, id="system_string"),
pytest.param(
[{"type": "text", "text": "You are a meticulous assistant. " * 40}],
id="system_blocks",
),
],
)
def test_pre_call_checks_counts_anthropic_system_tokens(monkeypatch, system):
"""
The Anthropic Messages API carries the system prompt as a top-level `system` field,
not as a message. Its tokens reach the model, so a request whose `messages` fit but
whose `messages` plus `system` exceed the context window must be rejected.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
]
messages = [{"role": "user", "content": "hi"}]
messages_only_tokens = router._count_pre_call_check_tokens(messages=messages, input=None)
monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": messages_only_tokens})
assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, messages=messages)) == 1
with pytest.raises(litellm.ContextWindowExceededError):
router._pre_call_checks(
model="m",
healthy_deployments=deployments,
messages=messages,
request_kwargs={"system": system},
)
@pytest.mark.asyncio
async def test_aanthropic_messages_enforces_context_window_with_system_and_tools():
"""
End-to-end router regression for /v1/messages: a request whose only oversized
content lives in the top-level `system` field or in `tools` must trip the pre-call
context-window check instead of being dispatched (the deployment uses mock_response,
so reaching the provider handler would return a response rather than raise).
"""
router = litellm.Router(
model_list=[
{
"model_name": "small-ctx",
"litellm_params": {"model": "anthropic/claude-3-5-haiku-20241022", "mock_response": "hi"},
"model_info": {"max_input_tokens": 20},
}
],
enable_pre_call_checks=True,
)
messages = [{"role": "user", "content": "hi"}]
response = await router.aanthropic_messages(model="small-ctx", messages=messages, max_tokens=5)
assert response is not None
with pytest.raises(litellm.ContextWindowExceededError):
await router.aanthropic_messages(
model="small-ctx",
messages=messages,
max_tokens=5,
system="You are a meticulous assistant. " * 40,
)
with pytest.raises(litellm.ContextWindowExceededError):
await router.aanthropic_messages(
model="small-ctx",
messages=messages,
max_tokens=5,
tools=[
{
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"input_schema": {"type": "object", "properties": {"q": {"type": "string"}}},
}
],
)
def test_count_pre_call_check_tokens_across_api_surfaces():
"""
_count_pre_call_check_tokens must count tokens from chat `messages`, a Responses

View file

@ -3,7 +3,7 @@
"limit": 22328
},
"LIT002": {
"limit": 26760
"limit": 26750
},
"LIT003": {
"limit": 261
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16470
"limit": 16468
},
"LIT011": {
"limit": 5516
"limit": 5514
},
"LIT012": {
"limit": 4489