mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
56cf0cd223
17 changed files with 908 additions and 243 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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]] = {}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -96,7 +96,7 @@
|
|||
"limit": 10
|
||||
},
|
||||
"DTZ007": {
|
||||
"limit": 17
|
||||
"limit": 6
|
||||
},
|
||||
"DTZ011": {
|
||||
"limit": 3
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue