From 562664f117abaa7a98b6b6f85877cf5acf8d9ffd Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 31 Aug 2026 16:23:36 -0700 Subject: [PATCH 01/10] fix(anthropic_endpoints): return Anthropic type:error envelope for /v1/messages errors --- .../exceptions/exceptions.py | 4 +- .../proxy/anthropic_endpoints/endpoints.py | 42 +++++-- .../anthropic_endpoints/test_endpoints.py | 111 ++++++++++++------ 3 files changed, 113 insertions(+), 44 deletions(-) diff --git a/litellm/anthropic_interface/exceptions/exceptions.py b/litellm/anthropic_interface/exceptions/exceptions.py index ae333d1f4ad..91bcf82f455 100644 --- a/litellm/anthropic_interface/exceptions/exceptions.py +++ b/litellm/anthropic_interface/exceptions/exceptions.py @@ -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): diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 7f0045c1d93..b243b737b0a 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -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, ) diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index c83ba142011..f809fadc879 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -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=) 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=) + 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"] From 57da95a77cbbc27a01372786798acae2a7987489 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 15:27:47 -0700 Subject: [PATCH 02/10] fix(proxy): apply default_vertex_config location before building the Vertex passthrough base URL Routes without /projects//locations// built the upstream host from the URL's still-empty location and 500ed even with default_vertex_config set. Build the base URL once after the configured project and location are applied, drop the hook that re-derived it afterwards, and answer 400 with a fix-it message when no location is available at all. Resolves LIT-6905 --- .../llm_passthrough_endpoints.py | 42 ++--- .../test_llm_pass_through_endpoints.py | 147 ++++++++++++++++-- .../test_vertex_passthrough_load_balancing.py | 25 +-- 3 files changed, 147 insertions(+), 67 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 29f216fd450..688123c9d41 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -1659,6 +1659,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//locations// 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 @@ -1666,29 +1672,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) @@ -1911,10 +1906,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. @@ -1924,15 +1917,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 @@ -1985,14 +1975,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, @@ -2085,12 +2069,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, @@ -2100,13 +2081,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) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 1d4b0264879..e37060493f7 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -319,9 +319,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 @@ -427,9 +424,6 @@ class TestVertexAIPassThroughHandler: 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_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -530,9 +524,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 @@ -1308,9 +1299,6 @@ class TestVertexAIDiscoveryPassThroughHandler: 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_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -3650,7 +3638,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 ( @@ -4234,6 +4221,140 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: assert "sk-master-1234" not in " ".join(f"{name}:{value}" for name, value in forwarded.items()) +class TestVertexPassthroughDefaultLocationOnShortRoutes: + """Regression coverage for LIT-6905. + + ``default_vertex_config`` carries the project and location, yet a route that + omits ``/projects//locations//`` built the upstream base URL + from the still-unresolved URL location and 500ed with ``vertex_location is + required``. The base URL must be built after the configured location is + applied, and a request with no location anywhere must fail with a clean 400 + that says where a location can come from, never a 500. + """ + + PROJECT = "test-project" + SHORT_ROUTE = "publishers/google/models/gemini-2.5-flash:generateContent" + + async def _forward( + self, + monkeypatch, + endpoint: str, + default_config: dict | None, + headers: list[tuple[bytes, bytes]], + ) -> tuple[HTTPException | None, dict]: + from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import ( + PassthroughEndpointRouter, + ) + + async def receive(): + return {"type": "http.request", "body": b"{}", "more_body": False} + + request = Request( + { + "type": "http", + "method": "POST", + "path": f"/vertex_ai/{endpoint}", + "headers": headers, + "query_string": b"", + }, + receive=receive, + ) + + captured: dict = {} + + def fake_create_pass_through_route(**kwargs): + captured.update(kwargs) + return AsyncMock(return_value={"status": "success"}) + + router = PassthroughEndpointRouter() + if default_config is not None: + router.set_default_vertex_config(dict(default_config)) + module = "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 = Mock() + mock_credentials.token = "test-token" + caller: Final = UserAPIKeyAuth(api_key="test-key") + raised: HTTPException | None = None + 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", side_effect=fake_create_pass_through_route), + mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=caller)), + ): + try: + await vertex_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=Response(), + user_api_key_dict=caller, + ) + except HTTPException as exc: + raised = exc + return raised, captured + + @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 + ): + raised, captured = 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")], + ) + assert raised is None + assert str(captured["target"]) == expected_target + assert captured["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): + raised, captured = await self._forward(monkeypatch, self.SHORT_ROUTE, default_config, headers) + assert not captured, "a request with no location must never reach the upstream forwarder" + assert raised is not None + assert raised.status_code == 400 + assert "/projects//locations//" in str(raised.detail) + assert "default_vertex_config" in str(raised.detail) + + class TestGetAzureAISearchIndexFromEndpoint: """The operable index is only the segment right after ``indexes``. diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index 6735f2a3780..e8fd5579631 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -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", From 7bdd148f38f35e5baed4bced6fd980dd77a83bdd Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 3 Sep 2026 16:10:44 -0700 Subject: [PATCH 03/10] test(proxy-extras): fake run_prisma instead of subprocess.run in the migrate deploy harness --- tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py index 9ffb57924b6..3fab20a28ad 100644 --- a/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py +++ b/tests/litellm-proxy-extras/test_litellm_proxy_extras_utils.py @@ -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 From dc98901dc1645391986e3434a72cd256617837cf Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 3 Sep 2026 19:53:32 +0000 Subject: [PATCH 04/10] fix(scim): apply default_internal_user_params.teams to SCIM-created users Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/scim/scim_v2.py | 3 +- .../scim/test_scim_v2_endpoints.py | 80 ++++++++++++++++++- 2 files changed, 80 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 069f86c852c..0f0124ee9d9 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -1123,7 +1123,6 @@ 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, @@ -1699,7 +1698,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, diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 1697b77b99a..f4627f82506 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -24,6 +24,7 @@ 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 +38,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 +305,83 @@ 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 + + +@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 + + @pytest.mark.asyncio async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeypatch): """ From 0429339204ac41f2c0420d1693f78155244b6205 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 3 Sep 2026 20:16:09 +0000 Subject: [PATCH 05/10] fix(scim): pass proxy admin auth to new_user so default team add succeeds Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_endpoints/scim/scim_v2.py | 6 +++++- .../management_endpoints/scim/test_scim_v2_endpoints.py | 3 +++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 0f0124ee9d9..98770cf9c2b 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -1128,7 +1128,10 @@ async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_grou 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 @@ -1716,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) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index f4627f82506..88f67cfe0e4 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -19,6 +19,7 @@ from litellm.proxy._types import ( NewUserResponse, ProxyErrorTypes, ProxyException, + UserAPIKeyAuth, ) from litellm.proxy.management_endpoints.scim.scim_v2 import ( SCIMRosterSyncError, @@ -342,6 +343,7 @@ async def test_create_user_without_groups_defers_to_default_team(mocker: MockerF 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 @@ -380,6 +382,7 @@ async def test_create_user_if_not_exists_defers_to_default_team(mocker: MockerFi 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 From 07dd8a7e47957a020841439da35da1d287998e08 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 3 Sep 2026 15:46:44 -0700 Subject: [PATCH 06/10] fix(scim): keep team memberships when PUT /Users carries no groups Okta sends profile updates as full PUTs with no groups or groups: [], since SCIM User.groups is readOnly and membership is synced through /Groups. The PUT handler diffed that empty list against the stored teams, removed the user from every team (which also deletes their team keys) and recomputed the role from an empty group list. Treat an empty groups list on PUT as unspecified: keep the stored teams and leave the role alone. Explicit non-empty groups still replace memberships as before Claude-Session: https://claude.ai/code/session_01CqwUV4Ywnu5aUjXx1UhJrM --- .../management_endpoints/scim/scim_v2.py | 11 ++-- .../scim/test_scim_v2_endpoints.py | 61 +++++++++++++++++++ type-discipline-budget.json | 2 +- 3 files changed, 69 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 98770cf9c2b..ceb67e3eee8 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -1774,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() ) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 88f67cfe0e4..60f9a1a55e2 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -1257,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""" diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 2f85128b4b6..78090779109 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22328 }, "LIT002": { - "limit": 26760 + "limit": 26758 }, "LIT003": { "limit": 261 From 0b7773dd44aaf5c7da2e591e995ff3a41688ba0b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 3 Sep 2026 16:38:48 -0700 Subject: [PATCH 07/10] fix(router): count tools and Anthropic system prompt in context-window pre-call check (#39663) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/utils.py | 11 ++ litellm/router.py | 36 +++- tests/test_litellm/test_router.py | 160 +++++++++++++++++- 3 files changed, 198 insertions(+), 9 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py index 9deff950724..242300c7b6d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/utils.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/utils.py @@ -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]: """ diff --git a/litellm/router.py b/litellm/router.py index f33dfbba7bf..dea9aa62729 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, @@ -11762,7 +11763,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. @@ -11772,9 +11773,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 @@ -11787,7 +11807,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: @@ -11833,14 +11856,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( @@ -11887,8 +11909,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 ## @@ -11919,7 +11939,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( diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 228588d974f..f7f0d79b4fd 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -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 From a330bc98a68725d7e5afa078ac5eb05d0cfcf8ff Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 16:39:03 -0700 Subject: [PATCH 08/10] test(vertex-passthrough): inject the forwarder into the short-route regression helper --- .../test_llm_pass_through_endpoints.py | 74 ++++++++----------- 1 file changed, 30 insertions(+), 44 deletions(-) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index e37060493f7..5154f738e9a 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -4222,26 +4222,21 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: class TestVertexPassthroughDefaultLocationOnShortRoutes: - """Regression coverage for LIT-6905. - - ``default_vertex_config`` carries the project and location, yet a route that - omits ``/projects//locations//`` built the upstream base URL - from the still-unresolved URL location and 500ed with ``vertex_location is - required``. The base URL must be built after the configured location is - applied, and a request with no location anywhere must fail with a clean 400 - that says where a location can come from, never a 500. - """ - 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]], - ) -> tuple[HTTPException | None, dict]: + forwarder: Mock, + ) -> None: from litellm.proxy.pass_through_endpoints.passthrough_endpoint_router import ( PassthroughEndpointRouter, ) @@ -4249,7 +4244,7 @@ class TestVertexPassthroughDefaultLocationOnShortRoutes: async def receive(): return {"type": "http.request", "body": b"{}", "more_body": False} - request = Request( + request: Final = Request( { "type": "http", "method": "POST", @@ -4259,41 +4254,29 @@ class TestVertexPassthroughDefaultLocationOnShortRoutes: }, receive=receive, ) - - captured: dict = {} - - def fake_create_pass_through_route(**kwargs): - captured.update(kwargs) - return AsyncMock(return_value={"status": "success"}) - - router = PassthroughEndpointRouter() + router: Final = PassthroughEndpointRouter() if default_config is not None: router.set_default_vertex_config(dict(default_config)) - module = "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints" + 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 = Mock() + mock_credentials: Final = Mock() mock_credentials.token = "test-token" caller: Final = UserAPIKeyAuth(api_key="test-key") - raised: HTTPException | None = None 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", side_effect=fake_create_pass_through_route), + mock.patch(f"{module}.create_pass_through_route", new=forwarder), mock.patch(f"{module}.user_api_key_auth", new=AsyncMock(return_value=caller)), ): - try: - await vertex_proxy_route( - endpoint=endpoint, - request=request, - fastapi_response=Response(), - user_api_key_dict=caller, - ) - except HTTPException as exc: - raised = exc - return raised, captured + await vertex_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=Response(), + user_api_key_dict=caller, + ) @pytest.mark.asyncio @pytest.mark.parametrize( @@ -4325,15 +4308,17 @@ class TestVertexPassthroughDefaultLocationOnShortRoutes: async def test_default_vertex_config_location_fills_routes_without_project_and_location( self, monkeypatch, endpoint, location, expected_target ): - raised, captured = await self._forward( + 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, ) - assert raised is None - assert str(captured["target"]) == expected_target - assert captured["custom_headers"]["Authorization"] == "Bearer test-token" + 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( @@ -4347,12 +4332,13 @@ class TestVertexPassthroughDefaultLocationOnShortRoutes: ], ) async def test_no_location_anywhere_is_a_400_not_a_500(self, monkeypatch, default_config, headers): - raised, captured = await self._forward(monkeypatch, self.SHORT_ROUTE, default_config, headers) - assert not captured, "a request with no location must never reach the upstream forwarder" - assert raised is not None - assert raised.status_code == 400 - assert "/projects//locations//" in str(raised.detail) - assert "default_vertex_config" in str(raised.detail) + 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//locations//" in str(raised.value.detail) + assert "default_vertex_config" in str(raised.value.detail) class TestGetAzureAISearchIndexFromEndpoint: From 39a17898ffdb55e4b54c0a7fe750b4d5afe49ea6 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Thu, 3 Sep 2026 19:45:09 -0400 Subject: [PATCH 09/10] test(proxy-extras): repoint the migrate-deploy harness at the run_prisma seam (#39673) From cf3af0f486590633ee875f6d120c416509fad5d2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 3 Sep 2026 23:53:59 +0000 Subject: [PATCH 10/10] perf(spend): group /spend/logs summary by day in Postgres instead of per-row Prisma group_by (#39351) * perf(spend): group /spend/logs summary by day in Postgres instead of per-row Prisma group_by Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(spend): simplify /spend/logs daily summary aggregation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): preserve spend logs response schema Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): compare spend log range bounds as naive UTC timestamps Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: cover spend logs summary edge cases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: cover spend summary request filters Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 10 +- .../spend_management_endpoints.py | 164 +++++++----- ruff-strict-budget.json | 2 +- .../test_spend_management_endpoints.py | 235 +++++++++++++++--- type-discipline-budget.json | 6 +- 5 files changed, 303 insertions(+), 114 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 9a9b1138a7d..cb3756575f7 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -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 diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index b86a877e8f9..dff100bdea7 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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]] = {} diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 24c0ff6b181..8763318b4eb 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -96,7 +96,7 @@ "limit": 10 }, "DTZ007": { - "limit": 17 + "limit": 6 }, "DTZ011": { "limit": 3 diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 73a29afd9b9..329a33eb440 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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 [] diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 78090779109..ab1a793e09d 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22328 }, "LIT002": { - "limit": 26758 + "limit": 26750 }, "LIT003": { "limit": 261 @@ -27,10 +27,10 @@ "limit": 0 }, "LIT010": { - "limit": 16470 + "limit": 16468 }, "LIT011": { - "limit": 5516 + "limit": 5514 }, "LIT012": { "limit": 4489