From 03a4e8bfb57ec69cc184b5114b0f66f1480672c6 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 25 Jul 2026 21:09:23 +0000 Subject: [PATCH 001/281] fix(azure/realtime): authenticate realtime websocket with Azure AD token when no api-key --- litellm/llms/azure/realtime/handler.py | 21 ++- litellm/realtime_api/main.py | 15 +- .../realtime/test_azure_realtime_handler.py | 178 ++++++++++++++++++ 3 files changed, 207 insertions(+), 7 deletions(-) diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 86c1ed51b68..51f9ef5989c 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -30,6 +30,21 @@ async def forward_messages(client_ws: Any, backend_ws: Any): class AzureOpenAIRealtime(AzureChatCompletion): + @staticmethod + def get_auth_headers(api_key: str | None, azure_ad_token: str | None) -> dict[str, str]: + """ + Build the websocket handshake auth headers, preferring a static api-key and falling back to + an Azure AD (Entra ID) bearer token. Never sends both. + """ + if api_key: + return {"api-key": api_key} + if azure_ad_token: + return {"Authorization": f"Bearer {azure_ad_token}"} + raise ValueError( + "Missing Azure credentials for the realtime endpoint. Set an api_key, or configure Azure AD auth " + "(azure_ad_token, tenant_id/client_id/client_secret, or a managed identity)" + ) + def _construct_url( self, api_base: str, @@ -117,13 +132,13 @@ class AzureOpenAIRealtime(AzureChatCompletion): query_params=query_params, ) + auth_headers = self.get_auth_headers(api_key=api_key, azure_ad_token=azure_ad_token) + try: ssl_context = get_shared_realtime_ssl_context() async with websockets.connect( # type: ignore url, - additional_headers={ - "api-key": api_key, # type: ignore - }, + additional_headers=auth_headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, ssl=ssl_context, ) as backend_ws: diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 5ecf4d91ff6..e9175917f9a 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -23,6 +23,7 @@ from litellm.utils import ProviderConfigManager from ..litellm_core_utils.get_litellm_params import get_litellm_params from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging +from ..llms.azure.common_utils import get_azure_ad_token from ..llms.azure.realtime.handler import AzureOpenAIRealtime from ..llms.bedrock.realtime.handler import BedrockRealtime from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context @@ -376,7 +377,7 @@ async def _arealtime( api_base=api_base, api_key=api_key, api_version=api_version, - azure_ad_token=None, + azure_ad_token=(None if api_key else get_azure_ad_token(litellm_params)), client=None, timeout=timeout, logging_obj=litellm_logging_obj, @@ -536,6 +537,7 @@ async def _realtime_health_check( import websockets url: Optional[str] = None + auth_headers: dict[str, str | None] = {"api-key": api_key} if custom_llm_provider == "azure": url = azure_realtime._construct_url( api_base=api_base or "", @@ -543,6 +545,13 @@ async def _realtime_health_check( api_version=api_version or "2024-10-01-preview", realtime_protocol=realtime_protocol, ) + azure_litellm_params = GenericLiteLLMParams(**(model_params or {})) + auth_headers = dict( + azure_realtime.get_auth_headers( + api_key=api_key, + azure_ad_token=(None if api_key else get_azure_ad_token(azure_litellm_params)), + ) + ) elif custom_llm_provider == "openai": url = openai_realtime._construct_url( api_base=api_base or "https://api.openai.com/", @@ -584,9 +593,7 @@ async def _realtime_health_check( ssl_context = get_shared_realtime_ssl_context() async with websockets.connect( # type: ignore url, - additional_headers={ - "api-key": api_key, # type: ignore - }, + additional_headers=auth_headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, ssl=ssl_context, ): diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py index 4638bc4df0f..d9c49947f19 100644 --- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py @@ -563,3 +563,181 @@ async def test_async_realtime_default_maintains_backwards_compatibility(): mock_realtime_streaming.call_args.kwargs["backend_uses_beta_protocol"] is True ) + + +class _DummyAsyncContextManager: + def __init__(self, value): + self.value = value + + async def __aenter__(self): + return self.value + + async def __aexit__(self, exc_type, exc, tb): + return None + + +@pytest.mark.asyncio +async def test_async_realtime_uses_bearer_token_when_no_api_key(): + """ + Entra ID-only Azure realtime deployments have no static api-key, so the handshake must + authenticate with `Authorization: Bearer ` and must not send `api-key`. + + Regression test for https://github.com/BerriAI/litellm/issues/34654 + """ + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + handler = AzureOpenAIRealtime() + mock_backend_ws = AsyncMock() + + with ( + patch( + "websockets.connect", + return_value=_DummyAsyncContextManager(mock_backend_ws), + ) as mock_ws_connect, + patch("litellm.llms.azure.realtime.handler.RealTimeStreaming") as mock_realtime_streaming, + ): + mock_realtime_streaming.return_value.bidirectional_forward = AsyncMock() + + await handler.async_realtime( + model="gpt-realtime-whisper", + websocket=AsyncMock(), + logging_obj=MagicMock(), + api_base="https://my-endpoint.openai.azure.com", + api_key=None, + api_version="2024-10-01-preview", + azure_ad_token="my-entra-token", + ) + + headers = mock_ws_connect.call_args.kwargs["additional_headers"] + assert headers == {"Authorization": "Bearer my-entra-token"} + + +def test_get_auth_headers_prefers_api_key_and_never_sends_both(): + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + assert AzureOpenAIRealtime.get_auth_headers(api_key="test-key", azure_ad_token="my-entra-token") == { + "api-key": "test-key" + } + + +def test_get_auth_headers_without_credentials_raises(): + from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime + + with pytest.raises(ValueError, match="Missing Azure credentials"): + AzureOpenAIRealtime.get_auth_headers(api_key=None, azure_ad_token=None) + + +@pytest.mark.asyncio +async def test_arealtime_resolves_azure_ad_token_when_no_api_key(monkeypatch): + """ + `_arealtime` must resolve an Azure AD token (managed identity, service principal, etc.) + and forward it to the handler when the deployment has no api_key. + + Regression test for https://github.com/BerriAI/litellm/issues/34654 + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime)) + monkeypatch.setattr( + realtime_main, + "get_llm_provider", + lambda model, api_base=None, api_key=None: ( + "gpt-realtime-whisper", + "azure", + None, + "https://my-endpoint.openai.azure.com", + ), + ) + monkeypatch.delenv("AZURE_API_KEY", raising=False) + + captured_params = {} + + def fake_get_azure_ad_token(litellm_params): + captured_params["tenant_id"] = litellm_params.get("tenant_id") + return "my-entra-token" + + monkeypatch.setattr(realtime_main, "get_azure_ad_token", fake_get_azure_ad_token) + + await realtime_main._arealtime( + model="azure/gpt-realtime-whisper", + websocket=MagicMock(), + api_version="2024-10-01-preview", + litellm_logging_obj=MagicMock(), + tenant_id="my-tenant", + client_id="my-client", + client_secret="my-secret", + ) + + assert mock_async_realtime.call_args.kwargs["azure_ad_token"] == "my-entra-token" + assert captured_params["tenant_id"] == "my-tenant" + + +@pytest.mark.asyncio +async def test_arealtime_does_not_resolve_azure_ad_token_when_api_key_present(monkeypatch): + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime)) + monkeypatch.setattr( + realtime_main, + "get_llm_provider", + lambda model, api_base=None, api_key=None: ( + "gpt-realtime-whisper", + "azure", + "test-key", + "https://my-endpoint.openai.azure.com", + ), + ) + + def fail_get_azure_ad_token(litellm_params): + raise AssertionError("should not resolve an AD token when an api_key is configured") + + monkeypatch.setattr(realtime_main, "get_azure_ad_token", fail_get_azure_ad_token) + + await realtime_main._arealtime( + model="azure/gpt-realtime-whisper", + websocket=MagicMock(), + api_key="test-key", + api_version="2024-10-01-preview", + litellm_logging_obj=MagicMock(), + ) + + assert mock_async_realtime.call_args.kwargs["azure_ad_token"] is None + + +@pytest.mark.asyncio +async def test_realtime_health_check_uses_bearer_token_when_no_api_key(monkeypatch): + """ + An Entra ID-only realtime deployment must also pass its realtime health check. + + Regression test for https://github.com/BerriAI/litellm/issues/34654 + """ + from litellm.realtime_api import main as realtime_main + + connect_calls = [] + + monkeypatch.setattr( + realtime_main, + "get_azure_ad_token", + lambda litellm_params: "my-entra-token", + ) + + def fake_connect(url, **kwargs): + connect_calls.append(kwargs) + return _DummyAsyncContextManager(MagicMock()) + + monkeypatch.setattr("websockets.connect", fake_connect) + + assert ( + await realtime_main._realtime_health_check( + model="gpt-realtime-whisper", + custom_llm_provider="azure", + api_key=None, + api_base="https://my-endpoint.openai.azure.com", + api_version="2024-10-01-preview", + model_params={"tenant_id": "my-tenant"}, + ) + is True + ) + assert connect_calls[0]["additional_headers"] == {"Authorization": "Bearer my-entra-token"} From b930169f39ff81b34b9088d32596f2fc07241839 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 25 Jul 2026 21:21:30 +0000 Subject: [PATCH 002/281] fix(azure/realtime): resolve AD token from deployment azure_ad_token param and kwargs --- litellm/realtime_api/main.py | 7 +++- .../realtime/test_azure_realtime_handler.py | 36 +++++++++++++++++++ 2 files changed, 42 insertions(+), 1 deletion(-) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index e9175917f9a..e981db216af 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -371,13 +371,18 @@ async def _arealtime( if realtime_protocol is None and (query_params or {}).get("intent") == "transcription": realtime_protocol = "GA" realtime_protocol = realtime_protocol or "beta" + resolved_azure_ad_token = ( + None + if api_key + else get_azure_ad_token(GenericLiteLLMParams(**{**kwargs, "azure_ad_token": azure_ad_token})) + ) await azure_realtime.async_realtime( model=model, websocket=websocket, api_base=api_base, api_key=api_key, api_version=api_version, - azure_ad_token=(None if api_key else get_azure_ad_token(litellm_params)), + azure_ad_token=resolved_azure_ad_token, client=None, timeout=timeout, logging_obj=litellm_logging_obj, diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py index d9c49947f19..bf2e89de44c 100644 --- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py @@ -741,3 +741,39 @@ async def test_realtime_health_check_uses_bearer_token_when_no_api_key(monkeypat is True ) assert connect_calls[0]["additional_headers"] == {"Authorization": "Bearer my-entra-token"} + + +@pytest.mark.asyncio +async def test_arealtime_forwards_deployment_azure_ad_token(monkeypatch): + """ + The router binds a deployment's `azure_ad_token` to `_arealtime`'s named parameter rather than + **kwargs, so it must still reach the handler. + + Regression test for https://github.com/BerriAI/litellm/issues/34654 + """ + from litellm.realtime_api import main as realtime_main + + mock_async_realtime = AsyncMock() + monkeypatch.setattr(realtime_main, "azure_realtime", MagicMock(async_realtime=mock_async_realtime)) + monkeypatch.setattr( + realtime_main, + "get_llm_provider", + lambda model, api_base=None, api_key=None: ( + "gpt-realtime-whisper", + "azure", + None, + "https://my-endpoint.openai.azure.com", + ), + ) + monkeypatch.delenv("AZURE_API_KEY", raising=False) + monkeypatch.setattr(realtime_main.litellm, "api_key", None) + + await realtime_main._arealtime( + model="azure/gpt-realtime-whisper", + websocket=MagicMock(), + api_version="2024-10-01-preview", + azure_ad_token="deployment-entra-token", + litellm_logging_obj=MagicMock(), + ) + + assert mock_async_realtime.call_args.kwargs["azure_ad_token"] == "deployment-entra-token" From 0c5583b83f03d642f6ee4f42619b94023ce3d7e8 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 05:18:49 +0000 Subject: [PATCH 003/281] fix(google_genai): price streamed generateContent with the provider that served it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/google_genai/streaming_iterator.py | 14 ++- .../vertex_passthrough_logging_handler.py | 2 +- .../streaming_handler.py | 19 ++++ .../pass_through_endpoints.py | 1 + .../test_google_genai_streaming_iterator.py | 40 +++++++- .../test_streaming_handler.py | 99 +++++++++++++++++++ 6 files changed, 170 insertions(+), 5 deletions(-) create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py diff --git a/litellm/google_genai/streaming_iterator.py b/litellm/google_genai/streaming_iterator.py index e03f7ee745f..e2fac6a615b 100644 --- a/litellm/google_genai/streaming_iterator.py +++ b/litellm/google_genai/streaming_iterator.py @@ -2,6 +2,7 @@ import asyncio from datetime import datetime from typing import TYPE_CHECKING, Any, Final +import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, @@ -65,6 +66,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: litellm_logging_obj: LiteLLMLoggingObj, request_body: dict, model: str, + custom_llm_provider: str, hidden_params: dict[str, Any] | None = None, ): self.litellm_logging_obj = litellm_logging_obj @@ -72,6 +74,7 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: self.start_time = datetime.now() self.collected_chunks: list[bytes] = [] self.model = model + self.custom_llm_provider = custom_llm_provider self._hidden_params: dict[str, Any] = hidden_params or {} async def _handle_async_streaming_logging( @@ -83,13 +86,18 @@ class BaseGoogleGenAIGenerateContentStreamingIterator: ) end_time: Final = datetime.now() + endpoint_type: Final = ( + EndpointType.GEMINI + if self.custom_llm_provider == litellm.LlmProviders.GEMINI.value + else EndpointType.VERTEX_AI + ) asyncio.create_task( PassThroughStreamingHandler._route_streaming_logging_to_handler( litellm_logging_obj=self.litellm_logging_obj, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, url_route="/v1/generateContent", request_body=self.request_body or {}, - endpoint_type=EndpointType.VERTEX_AI, + endpoint_type=endpoint_type, start_time=self.start_time, raw_bytes=self.collected_chunks, end_time=end_time, @@ -118,13 +126,13 @@ class GoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateContent litellm_logging_obj=logging_obj, request_body=request_body or {}, model=model, + custom_llm_provider=custom_llm_provider, hidden_params=hidden_params, ) self.response = response self.model = model self.generate_content_provider_config = generate_content_provider_config self.litellm_metadata = litellm_metadata - self.custom_llm_provider = custom_llm_provider # Gemini streamGenerateContent uses SSE line framing; iter_lines keeps # large inlineData payloads (e.g. image/jpeg) intact within one event. self.stream_iterator = response.iter_lines() @@ -169,13 +177,13 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(BaseGoogleGenAIGenerateCo litellm_logging_obj=logging_obj, request_body=request_body or {}, model=model, + custom_llm_provider=custom_llm_provider, hidden_params=hidden_params, ) self.response = response self.model = model self.generate_content_provider_config = generate_content_provider_config self.litellm_metadata = litellm_metadata - self.custom_llm_provider = custom_llm_provider # Gemini streamGenerateContent uses SSE line framing; aiter_lines keeps # large inlineData payloads (e.g. image/jpeg) intact within one event. self.stream_iterator = response.aiter_lines() diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index afd8684dd92..36455611c95 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -592,7 +592,7 @@ class VertexPassthroughLoggingHandler: response_cost: Final = litellm.completion_cost( completion_response=litellm_model_response, model=model, - custom_llm_provider="vertex_ai", + custom_llm_provider=custom_llm_provider, ) kwargs["response_cost"] = response_cost diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index ff1c12d08d7..907c59d28cc 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -15,6 +15,9 @@ from litellm.types.utils import StandardPassThroughResponseObject from .llm_provider_handlers.anthropic_passthrough_logging_handler import ( AnthropicPassthroughLoggingHandler, ) +from .llm_provider_handlers.gemini_passthrough_logging_handler import ( + GeminiPassthroughLoggingHandler, +) from .llm_provider_handlers.openai_passthrough_logging_handler import ( OpenAIPassthroughLoggingHandler, ) @@ -221,6 +224,22 @@ class PassThroughStreamingHandler: ) standard_logging_response_object = vertex_passthrough_logging_handler_result["result"] kwargs = vertex_passthrough_logging_handler_result["kwargs"] + elif endpoint_type == EndpointType.GEMINI: + gemini_passthrough_logging_handler_result: Final = ( + GeminiPassthroughLoggingHandler._handle_logging_gemini_collected_chunks( + litellm_logging_obj=litellm_logging_obj, + passthrough_success_handler_obj=passthrough_success_handler_obj, + url_route=url_route, + request_body=request_body, + endpoint_type=endpoint_type, + start_time=start_time, + all_chunks=all_chunks, + end_time=end_time, + model=model, + ) + ) + standard_logging_response_object = gemini_passthrough_logging_handler_result["result"] + kwargs = gemini_passthrough_logging_handler_result["kwargs"] elif endpoint_type == EndpointType.OPENAI: openai_passthrough_logging_handler_result: Final = ( OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks( diff --git a/litellm/types/passthrough_endpoints/pass_through_endpoints.py b/litellm/types/passthrough_endpoints/pass_through_endpoints.py index f59ca0d9041..548702e4139 100644 --- a/litellm/types/passthrough_endpoints/pass_through_endpoints.py +++ b/litellm/types/passthrough_endpoints/pass_through_endpoints.py @@ -22,6 +22,7 @@ LITELLM_PASS_THROUGH_ENDPOINT_MARKER: Final = "__litellm_pass_through_endpoint__ class EndpointType(str, Enum): VERTEX_AI = "vertex-ai" + GEMINI = "gemini" ANTHROPIC = "anthropic" OPENAI = "openai" GENERIC = "generic" diff --git a/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py b/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py index d74a05ec59c..91058767730 100644 --- a/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py +++ b/tests/test_litellm/google_genai/test_google_genai_streaming_iterator.py @@ -1,5 +1,6 @@ +import asyncio import json -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -8,6 +9,43 @@ from litellm.google_genai.streaming_iterator import ( GoogleGenAIGenerateContentStreamingIterator, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "custom_llm_provider, expected_endpoint_type", + [("gemini", EndpointType.GEMINI), ("vertex_ai", EndpointType.VERTEX_AI)], +) +async def test_streaming_logging_routes_to_the_provider_that_served_the_request( + custom_llm_provider, expected_endpoint_type +): + """Routing every google stream through the vertex handler bills gemini/* at vertex_ai/ rates.""" + mock_response = MagicMock() + + async def _aiter_lines(): + yield 'data: {"candidates": []}' + + mock_response.aiter_lines = _aiter_lines + + iterator = AsyncGoogleGenAIGenerateContentStreamingIterator( + response=mock_response, + model="gemini-3.1-flash-image", + logging_obj=MagicMock(spec=LiteLLMLoggingObj), + generate_content_provider_config=MagicMock(), + litellm_metadata={}, + custom_llm_provider=custom_llm_provider, + ) + + with patch( + "litellm.proxy.pass_through_endpoints.streaming_handler.PassThroughStreamingHandler._route_streaming_logging_to_handler", + new=AsyncMock(), + ) as mock_route: + async for _ in iterator: + pass + + await asyncio.sleep(0) + assert mock_route.call_args.kwargs["endpoint_type"] == expected_endpoint_type def _large_inline_data_event() -> str: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py new file mode 100644 index 00000000000..d0c28fd60a9 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_streaming_handler.py @@ -0,0 +1,99 @@ +import json +from datetime import datetime +from unittest.mock import MagicMock + +import pytest + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, +) +from litellm.proxy.pass_through_endpoints.streaming_handler import ( + PassThroughStreamingHandler, +) +from litellm.proxy.pass_through_endpoints.success_handler import ( + PassThroughEndpointLogging, +) +from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType + +MODEL = "gemini-3.1-flash-image" + +# gemini/ rate card: 2.5e-07 in, 1.5e-06 out. vertex_ai/ rate card is exactly 2x that. +GEMINI_COST = 1000 * 2.5e-07 + 1000 * 1.5e-06 +VERTEX_COST = 2 * GEMINI_COST + + +def _chunks() -> list[str]: + payload = { + "candidates": [ + { + "content": {"parts": [{"text": "hi"}], "role": "model"}, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 1000, + "candidatesTokenCount": 1000, + "totalTokenCount": 2000, + }, + "modelVersion": MODEL, + } + return [f"data: {json.dumps(payload)}"] + + +def _logging_obj() -> LiteLLMLoggingObj: + logging_obj = MagicMock(spec=LiteLLMLoggingObj) + logging_obj.model_call_details = {} + logging_obj.optional_params = {} + logging_obj.litellm_call_id = "test-call-id" + return logging_obj + + +@pytest.mark.parametrize( + "endpoint_type, expected_provider, expected_cost", + [ + (EndpointType.GEMINI, "gemini", GEMINI_COST), + (EndpointType.VERTEX_AI, "vertex_ai", VERTEX_COST), + ], +) +def test_streaming_generate_content_bills_against_the_requested_provider( + endpoint_type, expected_provider, expected_cost +): + """A streamed gemini/* request must not be priced off the vertex_ai/ rate card.""" + logging_obj = _logging_obj() + + _, kwargs = PassThroughStreamingHandler._build_passthrough_logging_result( + litellm_logging_obj=logging_obj, + passthrough_success_handler_obj=PassThroughEndpointLogging(), + url_route="/v1/generateContent", + request_body={}, + endpoint_type=endpoint_type, + start_time=datetime.now(), + raw_bytes=[chunk.encode("utf-8") for chunk in _chunks()], + end_time=datetime.now(), + model=MODEL, + ) + + assert kwargs["response_cost"] == pytest.approx(expected_cost) + assert logging_obj.model_call_details["custom_llm_provider"] == expected_provider + + +def test_vertex_generate_content_payload_prices_gemini_urls_at_gemini_rates(): + """The AI Studio host resolves to `gemini`, so the cost must follow it, not the vertex_ai default.""" + logging_obj = _logging_obj() + + result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks( + litellm_logging_obj=logging_obj, + passthrough_success_handler_obj=PassThroughEndpointLogging(), + url_route=f"https://generativelanguage.googleapis.com/v1beta/models/{MODEL}:streamGenerateContent", + request_body={}, + endpoint_type=EndpointType.VERTEX_AI, + start_time=datetime.now(), + all_chunks=_chunks(), + model=MODEL, + end_time=datetime.now(), + ) + + assert result["kwargs"]["response_cost"] == pytest.approx(GEMINI_COST) + assert logging_obj.model_call_details["custom_llm_provider"] == "gemini" From e4c2ad4627b71d280603f721234ec8990f3aa6bf Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 19:23:34 -0700 Subject: [PATCH 004/281] fix(anthropic): buffer streamed responses carrying server-fulfilled tools so retrieval tool calls never reach the client --- .../compression_interception/handler.py | 4 +- litellm/integrations/custom_logger.py | 4 +- .../messages/agentic_streaming_iterator.py | 59 ++++++ litellm/llms/custom_httpx/llm_http_handler.py | 23 +++ .../guardrail_hooks/headroom/headroom.py | 1 + .../test_agentic_streaming_iterator.py | 178 ++++++++++++++++++ .../custom_httpx/test_llm_http_handler.py | 71 +++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 7 + 8 files changed, 345 insertions(+), 2 deletions(-) diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index 7ea60053e6f..76720682101 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -7,7 +7,7 @@ litellm_content_retrieve tool calls server-side via the typed agentic loop plan. import time import uuid -from typing import Any, Final, cast +from typing import Any, ClassVar, Final, cast from litellm._logging import verbose_logger from litellm.compression import compress @@ -72,6 +72,8 @@ class CompressionInterceptionLogger(CustomLogger): 4. Build typed rerun plan with tool_result blocks from the compressed cache. """ + server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset({LITELLM_CONTENT_RETRIEVE_TOOL_NAME}) + def __init__( self, enabled: bool = True, diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index a0c78674ac8..60af4063f84 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -3,7 +3,7 @@ import re import traceback from collections.abc import AsyncGenerator -from typing import TYPE_CHECKING, Any, Final, Optional +from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional from pydantic import BaseModel @@ -60,6 +60,8 @@ _BASE64_INLINE_PATTERN: Final = re.compile( class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class # Class variables or attributes + server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset() + def __init__( self, turn_off_message_logging: bool = False, diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 5c4fa4700c0..0699595821e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -6,14 +6,27 @@ yields every chunk to the caller (preserving real streaming), collects all bytes, and on stream exhaustion rebuilds the full Anthropic response to run through agentic completion hooks. If an agentic hook fires, the follow-up response is chained as Phase 2 of the same iterator. + +In hold-back mode (``hold_back=True``), chunks are buffered instead of +yielded live, with SSE ping events emitted while the upstream message is +in flight. On exhaustion the hooks run first: if a follow-up response +replaces the message, only the follow-up is yielded and the buffered +message is dropped; otherwise the buffer is replayed verbatim. This is +required for server-fulfilled tools (e.g. ``headroom_retrieve``), whose +tool_use blocks must never reach a client that cannot execute them. """ +import asyncio +import contextlib import json from collections.abc import AsyncIterator from typing import Any, Final, cast from litellm._logging import verbose_logger +PING_SSE_BYTES: Final = b'event: ping\ndata: {"type": "ping"}\n\n' +HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0 + # --------------------------------------------------------------------------- # SSE parsing helpers (module-level to keep the class lean) # --------------------------------------------------------------------------- @@ -156,6 +169,8 @@ class AgenticAnthropicStreamingIterator: logging_obj: Any, custom_llm_provider: str, kwargs: dict, + hold_back: bool = False, + ping_interval_seconds: float = HOLD_BACK_PING_INTERVAL_SECONDS, ): self._inner = completion_stream.__aiter__() self._http_handler = http_handler @@ -166,16 +181,23 @@ class AgenticAnthropicStreamingIterator: self._logging_obj = logging_obj self._custom_llm_provider = custom_llm_provider self._kwargs = kwargs + self._hold_back = hold_back + self._ping_interval_seconds = ping_interval_seconds self._collected_bytes: list[bytes] = [] self._stream_exhausted = False self._hook_processing_done = False self._follow_up_iterator: AsyncIterator | None = None + self._drain_task: asyncio.Task | None = None + self._replay_index = 0 def __aiter__(self): return self async def __anext__(self) -> bytes: + if self._hold_back: + return await self._anext_held_back() + # Phase 1: yield from upstream, collect bytes if not self._stream_exhausted: try: @@ -194,11 +216,48 @@ class AgenticAnthropicStreamingIterator: raise StopAsyncIteration + async def _drain_upstream(self) -> None: + try: + while True: + self._collected_bytes.append(await self._inner.__anext__()) + except StopAsyncIteration: + return + + async def _anext_held_back(self) -> bytes: + if self._drain_task is None: + self._drain_task = asyncio.create_task(self._drain_upstream()) + return PING_SSE_BYTES + + while not self._stream_exhausted: + try: + await asyncio.wait_for(asyncio.shield(self._drain_task), timeout=self._ping_interval_seconds) + except asyncio.TimeoutError: + return PING_SSE_BYTES + self._stream_exhausted = True + await self._process_agentic_hooks() + + if self._follow_up_iterator is not None: + return await self._follow_up_iterator.__anext__() + + if self._replay_index < len(self._collected_bytes): + chunk: Final = self._collected_bytes[self._replay_index] + self._replay_index += 1 + return chunk + + raise StopAsyncIteration + async def aclose(self) -> None: from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( aclose_if_supported, ) + if self._drain_task is not None and self._drain_task.done(): + if not self._drain_task.cancelled(): + self._drain_task.exception() + elif self._drain_task is not None: + self._drain_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._drain_task await aclose_if_supported(self._inner) await aclose_if_supported(self._follow_up_iterator) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index a58397c9184..913cedcfa55 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2189,6 +2189,10 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, + hold_back=self._should_hold_back_stream( + logging_obj=logging_obj, + tools=anthropic_messages_optional_request_params.get("tools"), + ), ) return AnthropicMessagesStreamingResponse( completion_stream=initial_response, @@ -5033,6 +5037,25 @@ class BaseLLMHTTPHandler: return True return False + @staticmethod + def _should_hold_back_stream(logging_obj: LiteLLMLoggingObj, tools: object) -> bool: + """ + True when the request carries a tool that a registered callback fulfills + server-side (e.g. ``headroom_retrieve``). The model's tool_use for such a + tool must never reach the client, which cannot execute it: the agentic + loop replaces the whole message with a follow-up response, so the stream + is buffered (with ping keepalives) instead of forwarded live. + """ + if not isinstance(tools, list) or not tools: + return False + from litellm.litellm_core_utils.prompt_templates.factory import has_tool_with_name + + return any( + has_tool_with_name(tools, name) + for cb in _custom_logger_callbacks(logging_obj) + for name in getattr(cb, "server_fulfilled_tool_names", frozenset()) + ) + @staticmethod def _check_agentic_loop_safety( tool_calls: object, diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index 8bfd5cca58a..84c6b220e62 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -339,6 +339,7 @@ def _build_responses_followup_items( class HeadroomGuardrail(CustomGuardrail): records_own_guardrail_information: ClassVar[bool] = True + server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset({HEADROOM_RETRIEVE_TOOL_NAME}) @classmethod def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py index b9bda07336f..a6b071bbab5 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py @@ -2,6 +2,7 @@ Tests for AgenticAnthropicStreamingIterator and SSE rebuild helpers. """ +import asyncio import json import os import sys @@ -13,6 +14,7 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + PING_SSE_BYTES, AgenticAnthropicStreamingIterator, _handle_content_block_delta, _handle_content_block_start, @@ -230,6 +232,51 @@ class MockAsyncStream: return chunk +class MockSlowAsyncStream(MockAsyncStream): + """Async iterator that sleeps before every chunk.""" + + def __init__(self, chunks: List[bytes], delay_seconds: float): + super().__init__(chunks) + self._delay_seconds = delay_seconds + + async def __anext__(self) -> bytes: + await asyncio.sleep(self._delay_seconds) + return await super().__anext__() + + +class MockFailingAsyncStream(MockAsyncStream): + """Async iterator that raises after yielding its chunks.""" + + def __init__(self, chunks: List[bytes], error: Exception): + super().__init__(chunks) + self._error = error + + async def __anext__(self) -> bytes: + if self._idx >= len(self._chunks): + raise self._error + return await super().__anext__() + + +def _build_hold_back_iterator( + stream: MockAsyncStream, + mock_handler: MagicMock, + ping_interval_seconds: float = 15.0, +) -> AgenticAnthropicStreamingIterator: + return AgenticAnthropicStreamingIterator( + completion_stream=stream, + http_handler=mock_handler, + model="claude-sonnet-4-20250514", + messages=[], + anthropic_messages_provider_config=MagicMock(), + anthropic_messages_optional_request_params={}, + logging_obj=MagicMock(), + custom_llm_provider="anthropic", + kwargs={}, + hold_back=True, + ping_interval_seconds=ping_interval_seconds, + ) + + # --------------------------------------------------------------------------- # Tests for _parse_sse_events # --------------------------------------------------------------------------- @@ -790,3 +837,134 @@ class TestAgenticStreamingIteratorErrorHandling: call_kwargs = mock_handler._call_agentic_completion_hooks.call_args assert call_kwargs.kwargs["stream"] is True + + +class TestAgenticStreamingIteratorHoldBack: + @pytest.mark.asyncio + async def test_should_not_leak_intercepted_message_when_follow_up_fires(self): + """The buffered tool_use message must be dropped: only pings and follow-up bytes reach the client.""" + phase1_chunks = _build_tool_use_stream() + phase2_chunks = [b"follow-up-chunk-1", b"follow-up-chunk-2"] + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=MockAsyncStream(phase2_chunks)) + + iterator = _build_hold_back_iterator(MockAsyncStream(phase1_chunks), mock_handler) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + non_ping = [c for c in collected if c != PING_SSE_BYTES] + assert non_ping == phase2_chunks + assert b"litellm_content_retrieve" not in b"".join(collected) + assert collected[0] == PING_SSE_BYTES + + @pytest.mark.asyncio + async def test_should_replay_buffer_verbatim_when_no_hook_fires(self): + """Without interception the buffered message is replayed byte-identical after the pings.""" + chunks = _build_simple_text_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert [c for c in collected if c != PING_SSE_BYTES] == chunks + mock_handler._call_agentic_completion_hooks.assert_awaited_once() + + @pytest.mark.asyncio + async def test_should_emit_pings_while_upstream_is_slow(self): + """Pings keep the client connection alive while the upstream message is buffered.""" + chunks = _build_simple_text_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator( + MockSlowAsyncStream(chunks, delay_seconds=0.05), + mock_handler, + ping_interval_seconds=0.02, + ) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert collected.count(PING_SSE_BYTES) >= 2 + assert [c for c in collected if c != PING_SSE_BYTES] == chunks + + @pytest.mark.asyncio + async def test_should_propagate_upstream_error_instead_of_partial_message(self): + """An upstream failure surfaces as an error; the client never receives a truncated message.""" + chunks = _build_simple_text_stream()[:2] + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator( + MockFailingAsyncStream(chunks, RuntimeError("upstream died")), + mock_handler, + ) + + collected = [] + with pytest.raises(RuntimeError, match="upstream died"): + async for chunk in iterator: + collected.append(chunk) + + assert all(c == PING_SSE_BYTES for c in collected) + mock_handler._call_agentic_completion_hooks.assert_not_awaited() + + @pytest.mark.asyncio + async def test_should_replay_buffer_when_hook_processing_errors(self): + """A hook crash degrades to replaying the original message rather than dropping it.""" + chunks = _build_tool_use_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=RuntimeError("hook exploded")) + + mock_logging = MagicMock() + mock_logging.litellm_call_id = "test_call_holdback" + + iterator = AgenticAnthropicStreamingIterator( + completion_stream=MockAsyncStream(chunks), + http_handler=mock_handler, + model="claude-sonnet-4-20250514", + messages=[], + anthropic_messages_provider_config=MagicMock(), + anthropic_messages_optional_request_params={}, + logging_obj=mock_logging, + custom_llm_provider="anthropic", + kwargs={}, + hold_back=True, + ) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert [c for c in collected if c != PING_SSE_BYTES] == chunks + + @pytest.mark.asyncio + async def test_aclose_cancels_drain_task(self): + """Closing the iterator mid-buffer must cancel the background drain task.""" + chunks = _build_simple_text_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator( + MockSlowAsyncStream(chunks, delay_seconds=5.0), + mock_handler, + ) + + first = await iterator.__anext__() + assert first == PING_SSE_BYTES + assert iterator._drain_task is not None + + await iterator.aclose() + assert iterator._drain_task.cancelled() diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index fddd8d09dfc..6c2727e7da2 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -2071,3 +2071,74 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques retry_authorization = posts[1]["headers"]["Authorization"] assert retry_authorization.startswith("AWS4-HMAC-SHA256") assert retry_authorization != first_attempt_headers["Authorization"] + + +class TestShouldHoldBackStream: + """_should_hold_back_stream gates the buffered (non-leaking) streaming mode + for server-fulfilled tools like headroom_retrieve.""" + + @staticmethod + def _logging_obj_with(callbacks): + logging_obj = Mock() + logging_obj.dynamic_success_callbacks = callbacks + return logging_obj + + def test_should_hold_back_when_callback_owns_tool_in_request(self): + from litellm.integrations.custom_logger import CustomLogger + + class RetrievalCallback(CustomLogger): + server_fulfilled_tool_names = frozenset({"headroom_retrieve"}) + + tools = [ + {"name": "Bash", "input_schema": {"type": "object"}}, + {"name": "headroom_retrieve", "input_schema": {"type": "object"}}, + ] + assert ( + BaseLLMHTTPHandler._should_hold_back_stream( + logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools + ) + is True + ) + + def test_should_stream_live_when_tool_absent_from_request(self): + from litellm.integrations.custom_logger import CustomLogger + + class RetrievalCallback(CustomLogger): + server_fulfilled_tool_names = frozenset({"headroom_retrieve"}) + + tools = [{"name": "Bash", "input_schema": {"type": "object"}}] + assert ( + BaseLLMHTTPHandler._should_hold_back_stream( + logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools + ) + is False + ) + + def test_should_stream_live_when_no_callback_declares_tool_names(self): + from litellm.integrations.custom_logger import CustomLogger + + tools = [{"name": "headroom_retrieve", "input_schema": {"type": "object"}}] + assert ( + BaseLLMHTTPHandler._should_hold_back_stream( + logging_obj=self._logging_obj_with([CustomLogger()]), tools=tools + ) + is False + ) + + def test_should_stream_live_without_tools(self): + assert BaseLLMHTTPHandler._should_hold_back_stream(logging_obj=self._logging_obj_with([]), tools=None) is False + + def test_interception_callbacks_declare_their_retrieval_tools(self): + from litellm.integrations.compression_interception.handler import ( + LITELLM_CONTENT_RETRIEVE_TOOL_NAME, + CompressionInterceptionLogger, + ) + from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import ( + HEADROOM_RETRIEVE_TOOL_NAME, + HeadroomGuardrail, + ) + + assert HeadroomGuardrail.server_fulfilled_tool_names == frozenset({HEADROOM_RETRIEVE_TOOL_NAME}) + assert CompressionInterceptionLogger.server_fulfilled_tool_names == frozenset( + {LITELLM_CONTENT_RETRIEVE_TOOL_NAME} + ) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f1660e77ad9..8e950874a10 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -21391,6 +21391,13 @@ export interface components { * @description What the routed traffic actually cost */ spend: number; + /** + * Tier Turns + * @description Turns per tier, keyed by the tier name the routing decision recorded at request time (never re-derived at read time, since the tier-to-model mapping is mutable config). Tier names are scoped to this group's router_type and are not comparable across types: a complexity router reports 'simple'/'medium'/'complex'/'reasoning', a quality router reports its numeric quality tier, and an adaptive router records no tier at all. Turns no tier served (the classifier fell back to default_model) are absent rather than pooled under a sentinel key, so the values may sum to less than turns + */ + tier_turns?: { + [key: string]: number; + }; /** Turns */ turns: number; }; From f994068a7338b4bb54fa7a53d76fb009dd3e9f6a Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 07:30:24 +0000 Subject: [PATCH 005/281] fix(anthropic): keep pinging during agentic hooks and fail instead of replaying server-fulfilled tool_use Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/agentic_streaming_iterator.py | 87 ++++++++++++--- litellm/llms/custom_httpx/llm_http_handler.py | 29 ++--- .../test_agentic_streaming_iterator.py | 105 +++++++++++++++--- .../custom_httpx/test_llm_http_handler.py | 28 ++--- 4 files changed, 190 insertions(+), 59 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 0699595821e..4cf348dda9e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -9,11 +9,13 @@ follow-up response is chained as Phase 2 of the same iterator. In hold-back mode (``hold_back=True``), chunks are buffered instead of yielded live, with SSE ping events emitted while the upstream message is -in flight. On exhaustion the hooks run first: if a follow-up response -replaces the message, only the follow-up is yielded and the buffered -message is dropped; otherwise the buffer is replayed verbatim. This is -required for server-fulfilled tools (e.g. ``headroom_retrieve``), whose -tool_use blocks must never reach a client that cannot execute them. +in flight and while the agentic hooks run. On exhaustion the hooks run +first: if a follow-up response replaces the message, only the follow-up +is yielded and the buffered message is dropped; otherwise the buffer is +replayed verbatim, unless it holds a tool_use for a server-fulfilled tool +(e.g. ``headroom_retrieve``), in which case an SSE ``error`` event is +emitted because such a block must never reach a client that cannot +execute it. """ import asyncio @@ -26,6 +28,11 @@ from litellm._logging import verbose_logger PING_SSE_BYTES: Final = b'event: ping\ndata: {"type": "ping"}\n\n' HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0 +SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = ( + b"event: error\n" + b'data: {"type": "error", "error": {"type": "api_error", "message": ' + b'"Server-side tool retrieval failed, so this turn could not be completed. Please retry."}}\n\n' +) # --------------------------------------------------------------------------- # SSE parsing helpers (module-level to keep the class lean) @@ -170,6 +177,7 @@ class AgenticAnthropicStreamingIterator: custom_llm_provider: str, kwargs: dict, hold_back: bool = False, + server_fulfilled_tool_names: frozenset[str] = frozenset(), ping_interval_seconds: float = HOLD_BACK_PING_INTERVAL_SECONDS, ): self._inner = completion_stream.__aiter__() @@ -182,6 +190,7 @@ class AgenticAnthropicStreamingIterator: self._custom_llm_provider = custom_llm_provider self._kwargs = kwargs self._hold_back = hold_back + self._server_fulfilled_tool_names = server_fulfilled_tool_names self._ping_interval_seconds = ping_interval_seconds self._collected_bytes: list[bytes] = [] @@ -189,7 +198,9 @@ class AgenticAnthropicStreamingIterator: self._hook_processing_done = False self._follow_up_iterator: AsyncIterator | None = None self._drain_task: asyncio.Task | None = None + self._hook_task: asyncio.Task | None = None self._replay_index = 0 + self._error_emitted = False def __aiter__(self): return self @@ -223,22 +234,42 @@ class AgenticAnthropicStreamingIterator: except StopAsyncIteration: return + async def _completed_within_ping_interval(self, task: asyncio.Task) -> bool: + try: + await asyncio.wait_for(asyncio.shield(task), timeout=self._ping_interval_seconds) + except asyncio.TimeoutError: + return False + return True + async def _anext_held_back(self) -> bytes: if self._drain_task is None: self._drain_task = asyncio.create_task(self._drain_upstream()) return PING_SSE_BYTES - while not self._stream_exhausted: - try: - await asyncio.wait_for(asyncio.shield(self._drain_task), timeout=self._ping_interval_seconds) - except asyncio.TimeoutError: + if not self._stream_exhausted: + if not await self._completed_within_ping_interval(self._drain_task): return PING_SSE_BYTES self._stream_exhausted = True - await self._process_agentic_hooks() + + if self._hook_task is None: + self._hook_task = asyncio.create_task(self._process_agentic_hooks()) + if not await self._completed_within_ping_interval(self._hook_task): + return PING_SSE_BYTES if self._follow_up_iterator is not None: return await self._follow_up_iterator.__anext__() + if self._buffer_holds_server_fulfilled_tool_use(): + if self._error_emitted: + raise StopAsyncIteration + self._error_emitted = True + verbose_logger.error( + "AgenticStreamingIterator: hooks did not replace a message containing a server-fulfilled " + "tool_use [model=%s]; emitting an SSE error instead of leaking the tool call to the client", + self._model, + ) + return SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES + if self._replay_index < len(self._collected_bytes): chunk: Final = self._collected_bytes[self._replay_index] self._replay_index += 1 @@ -246,18 +277,40 @@ class AgenticAnthropicStreamingIterator: raise StopAsyncIteration + def _buffer_holds_server_fulfilled_tool_use(self) -> bool: + if not self._server_fulfilled_tool_names: + return False + started_blocks: Final = ( + data.get("content_block") + for event_type, data in _parse_sse_events(b"".join(self._collected_bytes)) + if event_type == "content_block_start" + ) + return any( + isinstance(block, dict) + and block.get("type") == "tool_use" + and block.get("name") in self._server_fulfilled_tool_names + for block in started_blocks + ) + + @staticmethod + async def _settle_task(task: asyncio.Task | None) -> None: + if task is None: + return + if task.done(): + if not task.cancelled(): + task.exception() + return + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + async def aclose(self) -> None: from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( aclose_if_supported, ) - if self._drain_task is not None and self._drain_task.done(): - if not self._drain_task.cancelled(): - self._drain_task.exception() - elif self._drain_task is not None: - self._drain_task.cancel() - with contextlib.suppress(asyncio.CancelledError): - await self._drain_task + await self._settle_task(self._drain_task) + await self._settle_task(self._hook_task) await aclose_if_supported(self._inner) await aclose_if_supported(self._follow_up_iterator) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 913cedcfa55..e9b88e45219 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2179,6 +2179,10 @@ class BaseLLMHTTPHandler: AgenticAnthropicStreamingIterator, ) + held_back_tool_names: Final = self._server_fulfilled_tools_in_request( + logging_obj=logging_obj, + tools=anthropic_messages_optional_request_params.get("tools"), + ) initial_response = AgenticAnthropicStreamingIterator( completion_stream=completion_stream, http_handler=self, @@ -2189,10 +2193,8 @@ class BaseLLMHTTPHandler: logging_obj=logging_obj, custom_llm_provider=custom_llm_provider, kwargs={**kwargs, "api_key": api_key} if api_key else kwargs, - hold_back=self._should_hold_back_stream( - logging_obj=logging_obj, - tools=anthropic_messages_optional_request_params.get("tools"), - ), + hold_back=bool(held_back_tool_names), + server_fulfilled_tool_names=held_back_tool_names, ) return AnthropicMessagesStreamingResponse( completion_stream=initial_response, @@ -5038,22 +5040,23 @@ class BaseLLMHTTPHandler: return False @staticmethod - def _should_hold_back_stream(logging_obj: LiteLLMLoggingObj, tools: object) -> bool: + def _server_fulfilled_tools_in_request(logging_obj: LiteLLMLoggingObj, tools: object) -> frozenset[str]: """ - True when the request carries a tool that a registered callback fulfills - server-side (e.g. ``headroom_retrieve``). The model's tool_use for such a - tool must never reach the client, which cannot execute it: the agentic - loop replaces the whole message with a follow-up response, so the stream - is buffered (with ping keepalives) instead of forwarded live. + The request's tools that a registered callback fulfills server-side (e.g. + ``headroom_retrieve``). The model's tool_use for such a tool must never + reach the client, which cannot execute it: the agentic loop replaces the + whole message with a follow-up response, so a stream carrying any of + these is buffered (with ping keepalives) instead of forwarded live. """ if not isinstance(tools, list) or not tools: - return False + return frozenset() from litellm.litellm_core_utils.prompt_templates.factory import has_tool_with_name - return any( - has_tool_with_name(tools, name) + return frozenset( + name for cb in _custom_logger_callbacks(logging_obj) for name in getattr(cb, "server_fulfilled_tool_names", frozenset()) + if has_tool_with_name(tools, name) ) @staticmethod diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py index a6b071bbab5..d59f76232cd 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py @@ -15,6 +15,7 @@ sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( PING_SSE_BYTES, + SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES, AgenticAnthropicStreamingIterator, _handle_content_block_delta, _handle_content_block_start, @@ -261,6 +262,7 @@ def _build_hold_back_iterator( stream: MockAsyncStream, mock_handler: MagicMock, ping_interval_seconds: float = 15.0, + server_fulfilled_tool_names: frozenset = frozenset({"litellm_content_retrieve"}), ) -> AgenticAnthropicStreamingIterator: return AgenticAnthropicStreamingIterator( completion_stream=stream, @@ -273,6 +275,7 @@ def _build_hold_back_iterator( custom_llm_provider="anthropic", kwargs={}, hold_back=True, + server_fulfilled_tool_names=server_fulfilled_tool_names, ping_interval_seconds=ping_interval_seconds, ) @@ -920,27 +923,76 @@ class TestAgenticStreamingIteratorHoldBack: mock_handler._call_agentic_completion_hooks.assert_not_awaited() @pytest.mark.asyncio - async def test_should_replay_buffer_when_hook_processing_errors(self): - """A hook crash degrades to replaying the original message rather than dropping it.""" + async def test_should_emit_pings_while_hooks_are_slow(self): + """Retrieval and follow-up generation can outlast a client's idle timeout, so hooks get keepalives too.""" + chunks = _build_tool_use_stream() + phase2_chunks = [b"follow-up-chunk"] + + async def slow_hooks(**_kwargs): + await asyncio.sleep(0.12) + return MockAsyncStream(phase2_chunks) + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=slow_hooks) + + iterator = _build_hold_back_iterator( + MockAsyncStream(chunks), + mock_handler, + ping_interval_seconds=0.02, + ) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert collected.count(PING_SSE_BYTES) >= 4 + assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks + + @pytest.mark.asyncio + async def test_should_error_instead_of_replaying_server_fulfilled_tool_use_when_hook_crashes(self): + """A hook crash must not replay the buffered retrieval tool_use: that is the unknown-tool bug.""" chunks = _build_tool_use_stream() mock_handler = MagicMock() mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=RuntimeError("hook exploded")) - mock_logging = MagicMock() - mock_logging.litellm_call_id = "test_call_holdback" + iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler) - iterator = AgenticAnthropicStreamingIterator( - completion_stream=MockAsyncStream(chunks), - http_handler=mock_handler, - model="claude-sonnet-4-20250514", - messages=[], - anthropic_messages_provider_config=MagicMock(), - anthropic_messages_optional_request_params={}, - logging_obj=mock_logging, - custom_llm_provider="anthropic", - kwargs={}, - hold_back=True, + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] + assert b"litellm_content_retrieve" not in b"".join(collected) + + @pytest.mark.asyncio + async def test_should_error_instead_of_replaying_when_no_hook_fires_on_tool_use(self): + """Hooks returning None on a retrieval tool_use is still a leak, so the turn fails loudly.""" + chunks = _build_tool_use_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] + + @pytest.mark.asyncio + async def test_should_replay_client_owned_tool_use_verbatim(self): + """Only server-fulfilled tools are withheld: a client's own tool_use still reaches it byte-identical.""" + chunks = _build_tool_use_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=None) + + iterator = _build_hold_back_iterator( + MockAsyncStream(chunks), + mock_handler, + server_fulfilled_tool_names=frozenset({"headroom_retrieve"}), ) collected = [] @@ -968,3 +1020,26 @@ class TestAgenticStreamingIteratorHoldBack: await iterator.aclose() assert iterator._drain_task.cancelled() + + @pytest.mark.asyncio + async def test_aclose_cancels_in_flight_hook_task(self): + """Closing while hooks are running must not leave the retrieval follow-up task orphaned.""" + chunks = _build_tool_use_stream() + + async def never_finishing_hooks(**_kwargs): + await asyncio.sleep(5.0) + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=never_finishing_hooks) + + iterator = _build_hold_back_iterator( + MockAsyncStream(chunks), + mock_handler, + ping_interval_seconds=0.02, + ) + + while iterator._hook_task is None: + await iterator.__anext__() + + await iterator.aclose() + assert iterator._hook_task.cancelled() diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 6c2727e7da2..798fb4c92e2 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -2073,9 +2073,9 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques assert retry_authorization != first_attempt_headers["Authorization"] -class TestShouldHoldBackStream: - """_should_hold_back_stream gates the buffered (non-leaking) streaming mode - for server-fulfilled tools like headroom_retrieve.""" +class TestServerFulfilledToolsInRequest: + """_server_fulfilled_tools_in_request gates the buffered (non-leaking) streaming + mode for server-fulfilled tools like headroom_retrieve.""" @staticmethod def _logging_obj_with(callbacks): @@ -2093,12 +2093,9 @@ class TestShouldHoldBackStream: {"name": "Bash", "input_schema": {"type": "object"}}, {"name": "headroom_retrieve", "input_schema": {"type": "object"}}, ] - assert ( - BaseLLMHTTPHandler._should_hold_back_stream( - logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools - ) - is True - ) + assert BaseLLMHTTPHandler._server_fulfilled_tools_in_request( + logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools + ) == frozenset({"headroom_retrieve"}) def test_should_stream_live_when_tool_absent_from_request(self): from litellm.integrations.custom_logger import CustomLogger @@ -2108,10 +2105,10 @@ class TestShouldHoldBackStream: tools = [{"name": "Bash", "input_schema": {"type": "object"}}] assert ( - BaseLLMHTTPHandler._should_hold_back_stream( + BaseLLMHTTPHandler._server_fulfilled_tools_in_request( logging_obj=self._logging_obj_with([RetrievalCallback()]), tools=tools ) - is False + == frozenset() ) def test_should_stream_live_when_no_callback_declares_tool_names(self): @@ -2119,14 +2116,17 @@ class TestShouldHoldBackStream: tools = [{"name": "headroom_retrieve", "input_schema": {"type": "object"}}] assert ( - BaseLLMHTTPHandler._should_hold_back_stream( + BaseLLMHTTPHandler._server_fulfilled_tools_in_request( logging_obj=self._logging_obj_with([CustomLogger()]), tools=tools ) - is False + == frozenset() ) def test_should_stream_live_without_tools(self): - assert BaseLLMHTTPHandler._should_hold_back_stream(logging_obj=self._logging_obj_with([]), tools=None) is False + assert ( + BaseLLMHTTPHandler._server_fulfilled_tools_in_request(logging_obj=self._logging_obj_with([]), tools=None) + == frozenset() + ) def test_interception_callbacks_declare_their_retrieval_tools(self): from litellm.integrations.compression_interception.handler import ( From 398e3d214cc97be0531428f3fb506a7cc42e2683 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 19:28:40 +0000 Subject: [PATCH 006/281] refactor(anthropic): trim hold-back commentary and drop dead rebuilt-content expression Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/agentic_streaming_iterator.py | 17 ++++------------- litellm/llms/custom_httpx/llm_http_handler.py | 8 +------- 2 files changed, 5 insertions(+), 20 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 4cf348dda9e..a88b148e92c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -7,14 +7,10 @@ all bytes, and on stream exhaustion rebuilds the full Anthropic response to run through agentic completion hooks. If an agentic hook fires, the follow-up response is chained as Phase 2 of the same iterator. -In hold-back mode (``hold_back=True``), chunks are buffered instead of -yielded live, with SSE ping events emitted while the upstream message is -in flight and while the agentic hooks run. On exhaustion the hooks run -first: if a follow-up response replaces the message, only the follow-up -is yielded and the buffered message is dropped; otherwise the buffer is -replayed verbatim, unless it holds a tool_use for a server-fulfilled tool -(e.g. ``headroom_retrieve``), in which case an SSE ``error`` event is -emitted because such a block must never reach a client that cannot +In hold-back mode (``hold_back=True``) chunks are buffered instead of yielded +live, keepalive pings run until the hooks finish, and then either the follow-up +replaces the message or the buffer replays, except that a buffered tool_use for +a server-fulfilled tool fails the turn rather than reaching a client that cannot execute it. """ @@ -329,11 +325,6 @@ class AgenticAnthropicStreamingIterator: verbose_logger.debug("AgenticStreamingIterator: Could not rebuild response from SSE bytes") return - [ - (f"{b.get('type')}({b.get('name', '')})" if b.get("type") == "tool_use" else b.get("type")) - for b in rebuilt.get("content", []) - ] - result: Final = await self._http_handler._call_agentic_completion_hooks( response=rebuilt, model=self._model, diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index e9b88e45219..193a38a5404 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5041,13 +5041,7 @@ class BaseLLMHTTPHandler: @staticmethod def _server_fulfilled_tools_in_request(logging_obj: LiteLLMLoggingObj, tools: object) -> frozenset[str]: - """ - The request's tools that a registered callback fulfills server-side (e.g. - ``headroom_retrieve``). The model's tool_use for such a tool must never - reach the client, which cannot execute it: the agentic loop replaces the - whole message with a follow-up response, so a stream carrying any of - these is buffered (with ping keepalives) instead of forwarded live. - """ + """The request's tools that a registered callback fulfills server-side (e.g. ``headroom_retrieve``).""" if not isinstance(tools, list) or not tools: return frozenset() from litellm.litellm_core_utils.prompt_templates.factory import has_tool_with_name From cbefb1ce5ffbdc90c9e6691206752b40ec9ef6e0 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 19:41:42 +0000 Subject: [PATCH 007/281] fix(anthropic): keep pinging while the held-back follow-up stream is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/agentic_streaming_iterator.py | 27 ++++++++- .../test_agentic_streaming_iterator.py | 59 +++++++++++++++++++ 2 files changed, 83 insertions(+), 3 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index a88b148e92c..3aaba0b139c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -8,8 +8,8 @@ to run through agentic completion hooks. If an agentic hook fires, the follow-up response is chained as Phase 2 of the same iterator. In hold-back mode (``hold_back=True``) chunks are buffered instead of yielded -live, keepalive pings run until the hooks finish, and then either the follow-up -replaces the message or the buffer replays, except that a buffered tool_use for +live, keepalive pings run whenever no other byte is ready, and then either the +follow-up replaces the message or the buffer replays, except that a tool_use for a server-fulfilled tool fails the turn rather than reaching a client that cannot execute it. """ @@ -30,6 +30,14 @@ SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = ( b'"Server-side tool retrieval failed, so this turn could not be completed. Please retry."}}\n\n' ) + +async def _anext_or_none(iterator: AsyncIterator) -> bytes | None: + try: + return await iterator.__anext__() + except StopAsyncIteration: + return None + + # --------------------------------------------------------------------------- # SSE parsing helpers (module-level to keep the class lean) # --------------------------------------------------------------------------- @@ -195,6 +203,7 @@ class AgenticAnthropicStreamingIterator: self._follow_up_iterator: AsyncIterator | None = None self._drain_task: asyncio.Task | None = None self._hook_task: asyncio.Task | None = None + self._follow_up_chunk_task: asyncio.Task | None = None self._replay_index = 0 self._error_emitted = False @@ -253,7 +262,7 @@ class AgenticAnthropicStreamingIterator: return PING_SSE_BYTES if self._follow_up_iterator is not None: - return await self._follow_up_iterator.__anext__() + return await self._next_follow_up_chunk(self._follow_up_iterator) if self._buffer_holds_server_fulfilled_tool_use(): if self._error_emitted: @@ -273,6 +282,17 @@ class AgenticAnthropicStreamingIterator: raise StopAsyncIteration + async def _next_follow_up_chunk(self, follow_up_iterator: AsyncIterator) -> bytes: + if self._follow_up_chunk_task is None: + self._follow_up_chunk_task = asyncio.create_task(_anext_or_none(follow_up_iterator)) + if not await self._completed_within_ping_interval(self._follow_up_chunk_task): + return PING_SSE_BYTES + chunk: Final = self._follow_up_chunk_task.result() + self._follow_up_chunk_task = None + if chunk is None: + raise StopAsyncIteration + return chunk + def _buffer_holds_server_fulfilled_tool_use(self) -> bool: if not self._server_fulfilled_tool_names: return False @@ -307,6 +327,7 @@ class AgenticAnthropicStreamingIterator: await self._settle_task(self._drain_task) await self._settle_task(self._hook_task) + await self._settle_task(self._follow_up_chunk_task) await aclose_if_supported(self._inner) await aclose_if_supported(self._follow_up_iterator) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py index d59f76232cd..d4aebf099d1 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py @@ -1001,6 +1001,65 @@ class TestAgenticStreamingIteratorHoldBack: assert [c for c in collected if c != PING_SSE_BYTES] == chunks + @pytest.mark.asyncio + async def test_should_emit_pings_while_the_follow_up_stream_is_slow(self): + """The corrected answer can be slow to generate, so the follow-up stream gets keepalives too.""" + chunks = _build_tool_use_stream() + phase2_chunks = [b"follow-up-chunk-1", b"follow-up-chunk-2"] + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock( + return_value=MockSlowAsyncStream(phase2_chunks, delay_seconds=0.06) + ) + + iterator = _build_hold_back_iterator( + MockAsyncStream(chunks), + mock_handler, + ping_interval_seconds=0.02, + ) + + collected = [] + async for chunk in iterator: + collected.append(chunk) + + first_follow_up_index = collected.index(phase2_chunks[0]) + assert collected[first_follow_up_index + 1] == PING_SSE_BYTES + assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks + + @pytest.mark.asyncio + async def test_should_propagate_follow_up_stream_error(self): + """A failing follow-up stream surfaces its error instead of hanging on pings forever.""" + chunks = _build_tool_use_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock( + return_value=MockFailingAsyncStream([b"follow-up-chunk"], RuntimeError("follow-up died")) + ) + + iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler, ping_interval_seconds=0.02) + + with pytest.raises(RuntimeError, match="follow-up died"): + async for _ in iterator: + pass + + @pytest.mark.asyncio + async def test_aclose_cancels_in_flight_follow_up_chunk_task(self): + """Closing while a follow-up chunk is pending must not orphan that task.""" + chunks = _build_tool_use_stream() + + mock_handler = MagicMock() + mock_handler._call_agentic_completion_hooks = AsyncMock( + return_value=MockSlowAsyncStream([b"follow-up-chunk"], delay_seconds=5.0) + ) + + iterator = _build_hold_back_iterator(MockAsyncStream(chunks), mock_handler, ping_interval_seconds=0.02) + + while iterator._follow_up_chunk_task is None: + await iterator.__anext__() + + await iterator.aclose() + assert iterator._follow_up_chunk_task.cancelled() + @pytest.mark.asyncio async def test_aclose_cancels_drain_task(self): """Closing the iterator mid-buffer must cancel the background drain task.""" From bb0bb48da8c3a5fe3557812e79a31c71c608e006 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 20:05:45 +0000 Subject: [PATCH 008/281] fix(proxy): do not let held-back keepalive pings block the budget reservation refund on client disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 2 ++ .../messages/agentic_streaming_iterator.py | 10 +++---- litellm/proxy/common_request_processing.py | 6 ++-- litellm/proxy/common_utils/sse_keepalive.py | 4 ++- .../test_agentic_streaming_iterator.py | 30 +++++++++---------- .../proxy/test_budget_reservation.py | 28 +++++++++++++++++ 6 files changed, 57 insertions(+), 23 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 6f0e9e7afe2..9db9fb36b65 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -434,6 +434,8 @@ CONNECTION_ERROR_PATTERNS: Final[list[str]] = [ ] STREAM_SSE_DONE_STRING: Final[str] = "[DONE]" STREAM_SSE_DATA_PREFIX: Final[str] = "data: " +STREAM_SSE_KEEPALIVE_PING_CHUNK: Final[str] = 'event: ping\ndata: {"type": "ping"}\n\n' +STREAM_SSE_KEEPALIVE_PING_BYTES: Final[bytes] = STREAM_SSE_KEEPALIVE_PING_CHUNK.encode("utf-8") ### SPEND TRACKING ### DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND: Final = float( os.getenv("DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND", 0.001400) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 3aaba0b139c..d6f4e51a09a 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -21,8 +21,8 @@ from collections.abc import AsyncIterator from typing import Any, Final, cast from litellm._logging import verbose_logger +from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES -PING_SSE_BYTES: Final = b'event: ping\ndata: {"type": "ping"}\n\n' HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0 SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = ( b"event: error\n" @@ -249,17 +249,17 @@ class AgenticAnthropicStreamingIterator: async def _anext_held_back(self) -> bytes: if self._drain_task is None: self._drain_task = asyncio.create_task(self._drain_upstream()) - return PING_SSE_BYTES + return STREAM_SSE_KEEPALIVE_PING_BYTES if not self._stream_exhausted: if not await self._completed_within_ping_interval(self._drain_task): - return PING_SSE_BYTES + return STREAM_SSE_KEEPALIVE_PING_BYTES self._stream_exhausted = True if self._hook_task is None: self._hook_task = asyncio.create_task(self._process_agentic_hooks()) if not await self._completed_within_ping_interval(self._hook_task): - return PING_SSE_BYTES + return STREAM_SSE_KEEPALIVE_PING_BYTES if self._follow_up_iterator is not None: return await self._next_follow_up_chunk(self._follow_up_iterator) @@ -286,7 +286,7 @@ class AgenticAnthropicStreamingIterator: if self._follow_up_chunk_task is None: self._follow_up_chunk_task = asyncio.create_task(_anext_or_none(follow_up_iterator)) if not await self._completed_within_ping_interval(self._follow_up_chunk_task): - return PING_SSE_BYTES + return STREAM_SSE_KEEPALIVE_PING_BYTES chunk: Final = self._follow_up_chunk_task.result() self._follow_up_chunk_task = None if chunk is None: diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 159d7508f4e..2607ff411a9 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -27,6 +27,7 @@ from litellm.constants import ( MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG, RETURN_RAW_MODEL_NAME_METADATA_KEY, STREAM_SSE_DATA_PREFIX, + STREAM_SSE_KEEPALIVE_PING_BYTES, UNSAFE_PROXY_RESPONSE_HEADERS, ) from litellm.integrations.custom_guardrail import CustomGuardrail @@ -2953,8 +2954,9 @@ class ProxyBaseLLMRequestProcessing: # so a GeneratorExit on client disconnect is raised there and any # statement after the yield never runs. The slow-path hook is # awaited above, so a cancellation during it still leaves this - # False and refunds. - delivered_chunk = True + # False and refunds. A keepalive ping carries no provider output, + # so it must not suppress that refund. + delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES yield serialize_chunk(chunk) stream_completed = True except (asyncio.CancelledError, GeneratorExit): diff --git a/litellm/proxy/common_utils/sse_keepalive.py b/litellm/proxy/common_utils/sse_keepalive.py index 6700700ff7c..6e0ea4db431 100644 --- a/litellm/proxy/common_utils/sse_keepalive.py +++ b/litellm/proxy/common_utils/sse_keepalive.py @@ -6,7 +6,9 @@ from typing import Final import anyio -ANTHROPIC_PING_SSE_CHUNK: Final = 'event: ping\ndata: {"type": "ping"}\n\n' +from litellm.constants import STREAM_SSE_KEEPALIVE_PING_CHUNK + +ANTHROPIC_PING_SSE_CHUNK: Final = STREAM_SSE_KEEPALIVE_PING_CHUNK def _coerce_interval(ping_interval_seconds: float | str | None) -> float | None: diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py index d4aebf099d1..b0467430533 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_agentic_streaming_iterator.py @@ -13,8 +13,8 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) +from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( - PING_SSE_BYTES, SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES, AgenticAnthropicStreamingIterator, _handle_content_block_delta, @@ -858,10 +858,10 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - non_ping = [c for c in collected if c != PING_SSE_BYTES] + non_ping = [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] assert non_ping == phase2_chunks assert b"litellm_content_retrieve" not in b"".join(collected) - assert collected[0] == PING_SSE_BYTES + assert collected[0] == STREAM_SSE_KEEPALIVE_PING_BYTES @pytest.mark.asyncio async def test_should_replay_buffer_verbatim_when_no_hook_fires(self): @@ -877,7 +877,7 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert [c for c in collected if c != PING_SSE_BYTES] == chunks + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks mock_handler._call_agentic_completion_hooks.assert_awaited_once() @pytest.mark.asyncio @@ -898,8 +898,8 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert collected.count(PING_SSE_BYTES) >= 2 - assert [c for c in collected if c != PING_SSE_BYTES] == chunks + assert collected.count(STREAM_SSE_KEEPALIVE_PING_BYTES) >= 2 + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks @pytest.mark.asyncio async def test_should_propagate_upstream_error_instead_of_partial_message(self): @@ -919,7 +919,7 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert all(c == PING_SSE_BYTES for c in collected) + assert all(c == STREAM_SSE_KEEPALIVE_PING_BYTES for c in collected) mock_handler._call_agentic_completion_hooks.assert_not_awaited() @pytest.mark.asyncio @@ -945,8 +945,8 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert collected.count(PING_SSE_BYTES) >= 4 - assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks + assert collected.count(STREAM_SSE_KEEPALIVE_PING_BYTES) >= 4 + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == phase2_chunks @pytest.mark.asyncio async def test_should_error_instead_of_replaying_server_fulfilled_tool_use_when_hook_crashes(self): @@ -962,7 +962,7 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] assert b"litellm_content_retrieve" not in b"".join(collected) @pytest.mark.asyncio @@ -979,7 +979,7 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert [c for c in collected if c != PING_SSE_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == [SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES] @pytest.mark.asyncio async def test_should_replay_client_owned_tool_use_verbatim(self): @@ -999,7 +999,7 @@ class TestAgenticStreamingIteratorHoldBack: async for chunk in iterator: collected.append(chunk) - assert [c for c in collected if c != PING_SSE_BYTES] == chunks + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == chunks @pytest.mark.asyncio async def test_should_emit_pings_while_the_follow_up_stream_is_slow(self): @@ -1023,8 +1023,8 @@ class TestAgenticStreamingIteratorHoldBack: collected.append(chunk) first_follow_up_index = collected.index(phase2_chunks[0]) - assert collected[first_follow_up_index + 1] == PING_SSE_BYTES - assert [c for c in collected if c != PING_SSE_BYTES] == phase2_chunks + assert collected[first_follow_up_index + 1] == STREAM_SSE_KEEPALIVE_PING_BYTES + assert [c for c in collected if c != STREAM_SSE_KEEPALIVE_PING_BYTES] == phase2_chunks @pytest.mark.asyncio async def test_should_propagate_follow_up_stream_error(self): @@ -1074,7 +1074,7 @@ class TestAgenticStreamingIteratorHoldBack: ) first = await iterator.__anext__() - assert first == PING_SSE_BYTES + assert first == STREAM_SSE_KEEPALIVE_PING_BYTES assert iterator._drain_task is not None await iterator.aclose() diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 34adb4d2091..e9a0e80752d 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -7,6 +7,7 @@ from fastapi import HTTPException import litellm from litellm.caching.dual_cache import DualCache +from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES from litellm.proxy._types import ( LiteLLM_BudgetTable, LiteLLM_EndUserTable, @@ -2453,6 +2454,33 @@ async def test_streaming_cancel_after_chunk_keeps_reservation( streaming_logging_obj._arelease_max_parallel_requests_on_disconnect.assert_awaited_once() +@pytest.mark.asyncio +async def test_streaming_cancel_after_only_keepalive_pings_reconciles_to_input_cost( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token, reservation = await _reserve_for_stream( + counter_cache, key_cache, proxy_logging_obj, "key-cancel-after-ping" + ) + + async def cancel_after_ping(user_api_key_dict, response, request_data): + yield STREAM_SSE_KEEPALIVE_PING_BYTES + raise asyncio.CancelledError() + + generator, streaming_logging_obj = _drive_streaming_cancel(valid_token, cancel_after_ping) + received = [] + with pytest.raises(asyncio.CancelledError): + async for chunk in generator: + received.append(chunk) + + assert received == [STREAM_SSE_KEEPALIVE_PING_BYTES] + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-cancel-after-ping" + ) == pytest.approx(0.5) + assert reservation["finalized"] is True + + @pytest.mark.asyncio async def test_release_budget_reservation_on_cancel_swallows_release_errors(): # If the release itself fails (e.g. Redis unavailable) it must not escape From 2d1ee3aab2fe6a37c80085009f789416fe191d4e Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 20:39:55 +0000 Subject: [PATCH 009/281] fix(proxy): keep the reservation when a disconnect happens while provider output is held back Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/agentic_streaming_iterator.py | 5 ++ litellm/proxy/common_request_processing.py | 9 ++- .../proxy/test_budget_reservation.py | 63 +++++++++++++++++++ 3 files changed, 76 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index d6f4e51a09a..3d3d3a12b17 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -207,6 +207,11 @@ class AgenticAnthropicStreamingIterator: self._replay_index = 0 self._error_emitted = False + @property + def has_buffered_provider_output(self) -> bool: + """Whether provider output was received but withheld from the client behind keepalive pings.""" + return self._hold_back and bool(self._collected_bytes) + def __aiter__(self): return self diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 2607ff411a9..3aeb7729c81 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -40,6 +40,9 @@ from litellm.litellm_core_utils.llm_response_utils.get_headers import ( get_response_headers, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, +) from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_checks import can_key_call_resolved_model from litellm.proxy.auth.auth_utils import check_response_size_is_safe @@ -95,6 +98,10 @@ _CLIENT_DISCONNECTED_ERROR_INFORMATION: Final[StandardLoggingPayloadErrorInforma } +def _withheld_provider_output(response: object) -> bool: + return isinstance(response, AgenticAnthropicStreamingIterator) and response.has_buffered_provider_output + + def _should_return_raw_model_name(request_data: dict[str, object]) -> bool: return any( isinstance(metadata, dict) and metadata.get(RETURN_RAW_MODEL_NAME_METADATA_KEY) is True @@ -2970,7 +2977,7 @@ class ProxyBaseLLMRequestProcessing: # only sees GeneratorExit on GC) cannot own the refund. if not stream_completed: client_disconnected = True - if not delivered_chunk: + if not delivered_chunk and not _withheld_provider_output(response): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index e9a0e80752d..c3210dd7f4d 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -8,6 +8,9 @@ from fastapi import HTTPException import litellm from litellm.caching.dual_cache import DualCache from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES +from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import ( + AgenticAnthropicStreamingIterator, +) from litellm.proxy._types import ( LiteLLM_BudgetTable, LiteLLM_EndUserTable, @@ -2376,6 +2379,11 @@ async def _reserve_for_stream(counter_cache, key_cache, proxy_logging_obj, token return valid_token, reservation +async def _never_ending_stream(): + yield b'event: message_start\ndata: {"type": "message_start"}\n\n' + await asyncio.sleep(30) + + def _drive_streaming_cancel(valid_token, iterator_hook): streaming_logging_obj = MagicMock() streaming_logging_obj.async_post_call_streaming_iterator_hook = iterator_hook @@ -2481,6 +2489,61 @@ async def test_streaming_cancel_after_only_keepalive_pings_reconciles_to_input_c assert reservation["finalized"] is True +@pytest.mark.asyncio +async def test_streaming_cancel_while_holding_back_provider_output_keeps_reservation( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token, reservation = await _reserve_for_stream( + counter_cache, key_cache, proxy_logging_obj, "key-cancel-held-back" + ) + + held_back = AgenticAnthropicStreamingIterator( + completion_stream=_never_ending_stream(), + http_handler=MagicMock(), + model="claude-haiku-4-5", + messages=[], + anthropic_messages_provider_config=MagicMock(), + anthropic_messages_optional_request_params={}, + logging_obj=MagicMock(), + custom_llm_provider="anthropic", + kwargs={}, + hold_back=True, + server_fulfilled_tool_names=frozenset({"headroom_retrieve"}), + ping_interval_seconds=0.01, + ) + + async def ping_then_cancel(user_api_key_dict, response, request_data): + yield await response.__anext__() + while not response.has_buffered_provider_output: + yield await response.__anext__() + raise asyncio.CancelledError() + + streaming_logging_obj = MagicMock() + streaming_logging_obj.async_post_call_streaming_iterator_hook = ping_then_cancel + streaming_logging_obj._arelease_max_parallel_requests_on_disconnect = AsyncMock() + generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=held_back, + user_api_key_dict=valid_token, + request_data=_request_body(), + proxy_logging_obj=streaming_logging_obj, + serialize_chunk=lambda chunk: chunk, + serialize_error=lambda exc: str(exc), + ) + + received = [] + with pytest.raises(asyncio.CancelledError): + async for chunk in generator: + received.append(chunk) + + assert received and received == [STREAM_SSE_KEEPALIVE_PING_BYTES] * len(received) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-cancel-held-back" + ) == pytest.approx(2.0) + assert reservation.get("finalized") is not True + + @pytest.mark.asyncio async def test_release_budget_reservation_on_cancel_swallows_release_errors(): # If the release itself fails (e.g. Redis unavailable) it must not escape From 65eae963a7da34e9d4b714d4ce0b485168efaa21 Mon Sep 17 00:00:00 2001 From: Kunal Nayyar Date: Tue, 11 Aug 2026 13:03:55 +0530 Subject: [PATCH 010/281] feat(proxy): opt-in enforce rpm/tpm when adding a model Add general_settings toggle 'enforce_rpm_tpm_on_model_add' (default false). When true, /model/new rejects a model whose rpm or tpm is missing or not a positive value, so the Admin UI Add Model form surfaces a 400 validation error instead of silently storing an unbounded model (or one with a zero/negative limit that would exclude it from routing). --- .../model_management_endpoints.py | 37 +++++++++++++++++++ .../test_model_management_endpoints.py | 33 +++++++++++++++++ .../molecules/notifications_manager.tsx | 1 + 3 files changed, 71 insertions(+) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 8a52b0d1abb..d9030f42f69 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -239,6 +239,38 @@ def _raise_on_strategy_router_write_violation( ) +ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add" +_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm") + + +def _raise_if_rate_limits_required_but_missing(*, litellm_params: GenericLiteLLMParams, enforced: bool) -> None: + """Require both rpm and tpm (each a positive value) when the operator opts in via config.yaml. + + Off by default, so deployments keep adding models without limits. When + ``enforce_rpm_tpm_on_model_add: true`` is set under general_settings, a model added + without both rpm and tpm set to a positive value is rejected rather than stored + unbounded (or effectively excluded from routing by a zero/negative limit). + """ + if not enforced: + return + missing: Final = tuple( + field + for field in _REQUIRED_RATE_LIMIT_FIELDS + if (value := getattr(litellm_params, field)) is None or value <= 0 + ) + if not missing: + return + raise ProxyException( + message=( + f"{' and '.join(missing)} must be set to a positive value when " + f"'{ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING}' is enabled in general_settings" + ), + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param=f"litellm_params.{missing[0]}", + ) + + _PTU_MODEL_INFO_FIELDS: Final = ("ptu_count", "cost_per_ptu_per_hour", "ptu_effective_from", "ptu_effective_to") @@ -1566,6 +1598,11 @@ async def add_new_model( existing_params=None, ) + _raise_if_rate_limits_required_but_missing( + litellm_params=model_params.litellm_params, + enforced=bool(general_settings.get(ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING, False)), + ) + model_response: LiteLLM_ProxyModelTable | None = None # update DB incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True) diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 454849d6430..cb87dfadcfa 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -23,6 +23,7 @@ from litellm.proxy._types import ( from litellm.proxy.management_endpoints.model_management_endpoints import ( ModelManagementAuthChecks, _get_team_deployments, + _raise_if_rate_limits_required_but_missing, clear_cache, delete_team_models, ) @@ -3825,3 +3826,35 @@ class TestAutoRouterClassifierDefaultPrompt: for empty in (None, "", "{}"): response = await get_auto_router_classifier_default_prompt(context_window_size=5, tier_labels=empty) assert response.system_prompt == classification_system_prompt(5) + + +class TestEnforceRpmTpmOnModelAdd: + def test_passes_when_disabled_even_without_limits(self): + _raise_if_rate_limits_required_but_missing( + litellm_params=LiteLLM_Params(model="azure/gpt-5.2"), + enforced=False, + ) + + def test_passes_when_enabled_and_both_set(self): + _raise_if_rate_limits_required_but_missing( + litellm_params=LiteLLM_Params(model="azure/gpt-5.2", rpm=10, tpm=1000), + enforced=True, + ) + + @pytest.mark.parametrize( + "params, expected_missing", + [ + (LiteLLM_Params(model="azure/gpt-5.2"), "rpm and tpm"), + (LiteLLM_Params(model="azure/gpt-5.2", rpm=10), "tpm"), + (LiteLLM_Params(model="azure/gpt-5.2", tpm=1000), "rpm"), + (LiteLLM_Params(model="azure/gpt-5.2", rpm=0, tpm=1000), "rpm"), + (LiteLLM_Params(model="azure/gpt-5.2", rpm=10, tpm=-1), "tpm"), + ], + ) + def test_raises_when_enabled_and_missing(self, params, expected_missing): + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc_info: + _raise_if_rate_limits_required_but_missing(litellm_params=params, enforced=True) + assert expected_missing in str(exc_info.value.message) + assert exc_info.value.code == "400" diff --git a/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx b/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx index 59b048b412c..31daee0b23b 100644 --- a/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx +++ b/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx @@ -104,6 +104,7 @@ const VALIDATION_MATCH = [ "invalid file type", "invalid field", "invalid date format", + "must be set when", ]; const NOT_FOUND_MATCH = [ From 526fc9eab192a6854f30c75aa574daf0d8d1f992 Mon Sep 17 00:00:00 2001 From: Kunal Nayyar Date: Tue, 11 Aug 2026 13:03:55 +0530 Subject: [PATCH 011/281] fix(ui): title validation errors correctly instead of Rate Limit Exceeded The /model/new endpoint returns a 400 validation error (type: validation_error) when 'rpm and tpm must be set to a positive value when enforce_rpm_tpm_on_model_add is enabled in general_settings' but the frontend's titleFor() keyword matcher mistitled it as 'Rate Limit Exceeded' because the message contains 'rpm'/'tpm' substrings, which matched the generic rate-limit keyword check before the more specific validation check could catch it. Add "'enforce_rpm_tpm_on_model_add' is enabled" to VALIDATION_MATCH so this message is classified as a Validation Error, matching the actual HTTP 400 validation_error the backend already returns. A narrow match on the setting name (rather than the generic "must be set when") avoids overriding the status-based classification of unrelated 401s, e.g. the PKCE 'GENERIC_CLIENT_ID must be set when PKCE is enabled' error. --- .../src/components/molecules/notifications_manager.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx b/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx index 31daee0b23b..aa3552da50a 100644 --- a/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx +++ b/ui/litellm-dashboard/src/components/molecules/notifications_manager.tsx @@ -104,7 +104,7 @@ const VALIDATION_MATCH = [ "invalid file type", "invalid field", "invalid date format", - "must be set when", + "'enforce_rpm_tpm_on_model_add' is enabled", ]; const NOT_FOUND_MATCH = [ From 97290b4e0e4ce140a80cf76ef87b433e145276eb Mon Sep 17 00:00:00 2001 From: Daniel Vainshtein Date: Thu, 13 Aug 2026 13:44:15 +0300 Subject: [PATCH 012/281] fix(bedrock): parse cacheDetails for Converse 1h/5m cache write cost split AmazonConverseConfig._transform_usage only read the aggregate cacheWriteInputTokens field, so cache_creation_token_details was always unset for Bedrock Converse responses. calculate_cache_writing_cost bills the whole cache-write count at the 5m rate whenever that field is None, so 1-hour TTL cache writes on the standard Bedrock chat path were always undercounted, even though Bedrock returns the 5m/1h split in usage.cacheDetails. Parse cacheDetails (when present) into CacheCreationTokenDetails so the correct rate applies to each portion. No cacheDetails in the response (older models/regions) keeps the previous behavior. Fixes #36760 Co-Authored-By: pi (Claude/GPT via @earendil-works/pi-coding-agent) --- .../bedrock/chat/converse_transformation.py | 20 +++++++++ litellm/types/llms/bedrock.py | 14 ++++-- .../chat/test_converse_transformation.py | 43 +++++++++++++++++++ 3 files changed, 74 insertions(+), 3 deletions(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 85918d40e12..9e7c86e615f 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -57,6 +57,7 @@ from litellm.types.llms.openai import ( OpenAIMessageContentListBlock, ) from litellm.types.utils import ( + CacheCreationTokenDetails, ChatCompletionMessageToolCall, CompletionTokensDetailsWrapper, Function, @@ -1770,6 +1771,24 @@ class AmazonConverseConfig(BaseConfig): thinking_blocks_list.append(_redacted_block) return thinking_blocks_list + @staticmethod + def _parse_cache_details(usage: ConverseTokenUsageBlock) -> "CacheCreationTokenDetails | None": + """ + Split Converse's aggregate cacheWriteInputTokens into the 5m/1h TTL + breakdown from `cacheDetails`, so cost calc can bill each tier + correctly instead of defaulting the whole write to the 5m rate. + https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html + """ + cache_details = usage.get("cacheDetails") + if not cache_details: + return None + tokens_5m = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "5m") + tokens_1h = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "1h") + return CacheCreationTokenDetails( + ephemeral_5m_input_tokens=tokens_5m, + ephemeral_1h_input_tokens=tokens_1h, + ) + def _transform_usage( self, usage: ConverseTokenUsageBlock, @@ -1792,6 +1811,7 @@ class AmazonConverseConfig(BaseConfig): prompt_tokens_details: Final = PromptTokensDetailsWrapper( cached_tokens=cache_read_input_tokens, cache_creation_tokens=cache_creation_input_tokens, + cache_creation_token_details=self._parse_cache_details(usage), text_tokens=raw_input_tokens, ) reasoning_tokens = token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0 diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 5665aa3277a..4c3ed8d6993 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -216,14 +216,22 @@ class ConverseResponseOutputBlock(TypedDict): message: MessageBlock | None -class ConverseTokenUsageBlock(TypedDict): +class CacheDetailBlock(TypedDict): + """Per-TTL cache-write breakdown. https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html""" + inputTokens: int - outputTokens: int - totalTokens: int + ttl: Literal["5m", "1h"] + + +class ConverseTokenUsageBlock(TypedDict, total=False): + inputTokens: Required[int] + outputTokens: Required[int] + totalTokens: Required[int] cacheReadInputTokenCount: int cacheReadInputTokens: int cacheWriteInputTokenCount: int cacheWriteInputTokens: int + cacheDetails: list[CacheDetailBlock] class ServiceTierBlock(TypedDict): diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index d1d1f9ab489..393499fb041 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -51,6 +51,49 @@ def test_transform_usage(): assert openai_usage.completion_tokens_details.text_tokens == usage["outputTokens"] +def test_transform_usage_with_cache_details(): + """cacheDetails should split cacheWriteInputTokens into the 5m/1h TTL breakdown + so cost calc can bill the 1h portion at its own (higher) rate instead of + defaulting the whole write to the 5m rate. See issue #36760.""" + usage = ConverseTokenUsageBlock( + **{ + "inputTokens": 76, + "outputTokens": 259, + "totalTokens": 335, + "cacheWriteInputTokens": 362, + "cacheDetails": [ + {"inputTokens": 74, "ttl": "1h"}, + {"inputTokens": 288, "ttl": "5m"}, + ], + } + ) + config = AmazonConverseConfig() + openai_usage = config._transform_usage(usage) + details = openai_usage.prompt_tokens_details.cache_creation_token_details + assert details is not None + assert details.ephemeral_1h_input_tokens == 74 + assert details.ephemeral_5m_input_tokens == 288 + + +def test_transform_usage_without_cache_details_stays_none(): + """No cacheDetails in the response (older models/regions) should leave + cache_creation_token_details unset, same as before this field existed.""" + usage = ConverseTokenUsageBlock( + **{ + "inputTokens": 3, + "outputTokens": 401, + "totalTokens": 2193, + "cacheWriteInputTokens": 1789, + } + ) + config = AmazonConverseConfig() + openai_usage = config._transform_usage(usage) + assert ( + getattr(openai_usage.prompt_tokens_details, "cache_creation_token_details", None) + is None + ) + + def test_transform_usage_with_reasoning_content(): """Test that completion_tokens_details correctly tracks reasoning vs text tokens.""" usage = ConverseTokenUsageBlock( From 42a2b5f057f4ff1be2ec37ab3189ff7631ee89e7 Mon Sep 17 00:00:00 2001 From: Daniel Vainshtein Date: Thu, 13 Aug 2026 14:06:22 +0300 Subject: [PATCH 013/281] fix(bedrock): guard cache-detail split against partial/unrecognized ttl entries Address review feedback on #36762: - Only use the parsed 5m/1h split when it fully accounts for cacheWriteInputTokens; an unrecognized ttl or missing entry now falls back to the aggregate (previous behavior) instead of silently understating cost. - Mark TypedDict fields ReadOnly (AWS response data, never constructed by us) to satisfy the repo's type-discipline lint gate. - Trim comments and add Final to locals per repo style. Co-Authored-By: pi (Claude/GPT via @earendil-works/pi-coding-agent) --- .../bedrock/chat/converse_transformation.py | 18 +++++++------- litellm/types/llms/bedrock.py | 24 +++++++++---------- .../chat/test_converse_transformation.py | 21 ++++++++++++++++ 3 files changed, 42 insertions(+), 21 deletions(-) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 9e7c86e615f..dbea1783dc2 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1773,17 +1773,17 @@ class AmazonConverseConfig(BaseConfig): @staticmethod def _parse_cache_details(usage: ConverseTokenUsageBlock) -> "CacheCreationTokenDetails | None": - """ - Split Converse's aggregate cacheWriteInputTokens into the 5m/1h TTL - breakdown from `cacheDetails`, so cost calc can bill each tier - correctly instead of defaulting the whole write to the 5m rate. - https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html - """ - cache_details = usage.get("cacheDetails") + """https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html""" + cache_details: Final = usage.get("cacheDetails") if not cache_details: return None - tokens_5m = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "5m") - tokens_1h = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "1h") + tokens_5m: Final = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "5m") + tokens_1h: Final = sum(d["inputTokens"] for d in cache_details if d.get("ttl") == "1h") + # An unrecognized ttl or a partial breakdown would silently understate + # the cache-write cost, so only use the split when it fully accounts + # for the aggregate; otherwise fall back to the aggregate-only (5m) cost. + if tokens_5m + tokens_1h != usage.get("cacheWriteInputTokens", 0): + return None return CacheCreationTokenDetails( ephemeral_5m_input_tokens=tokens_5m, ephemeral_1h_input_tokens=tokens_1h, diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 4c3ed8d6993..847e4066cf3 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -2,7 +2,7 @@ import json from enum import Enum from typing import TYPE_CHECKING, Any, Final, Literal -from typing_extensions import Required, TypedDict, override +from typing_extensions import ReadOnly, Required, TypedDict, override from .openai import ChatCompletionToolCallChunk @@ -217,21 +217,21 @@ class ConverseResponseOutputBlock(TypedDict): class CacheDetailBlock(TypedDict): - """Per-TTL cache-write breakdown. https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html""" + """Per-TTL cache-write breakdown, read-only AWS response data. https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_CacheDetail.html""" - inputTokens: int - ttl: Literal["5m", "1h"] + inputTokens: ReadOnly[int] + ttl: ReadOnly[Literal["5m", "1h"]] class ConverseTokenUsageBlock(TypedDict, total=False): - inputTokens: Required[int] - outputTokens: Required[int] - totalTokens: Required[int] - cacheReadInputTokenCount: int - cacheReadInputTokens: int - cacheWriteInputTokenCount: int - cacheWriteInputTokens: int - cacheDetails: list[CacheDetailBlock] + inputTokens: Required[ReadOnly[int]] + outputTokens: Required[ReadOnly[int]] + totalTokens: Required[ReadOnly[int]] + cacheReadInputTokenCount: ReadOnly[int] + cacheReadInputTokens: ReadOnly[int] + cacheWriteInputTokenCount: ReadOnly[int] + cacheWriteInputTokens: ReadOnly[int] + cacheDetails: ReadOnly[list[CacheDetailBlock]] # mutable-ok: AWS response array, never mutated after parsing class ServiceTierBlock(TypedDict): diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 393499fb041..a7d6695ca35 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -75,6 +75,27 @@ def test_transform_usage_with_cache_details(): assert details.ephemeral_5m_input_tokens == 288 +def test_transform_usage_with_mismatched_cache_details_falls_back(): + """An unrecognized ttl or partial breakdown must not silently understate + cache-write cost, so the split is only used when it fully accounts for + cacheWriteInputTokens.""" + usage = ConverseTokenUsageBlock( + **{ + "inputTokens": 76, + "outputTokens": 259, + "totalTokens": 335, + "cacheWriteInputTokens": 362, + "cacheDetails": [{"inputTokens": 74, "ttl": "1h"}], # missing the 5m entry + } + ) + config = AmazonConverseConfig() + openai_usage = config._transform_usage(usage) + assert ( + getattr(openai_usage.prompt_tokens_details, "cache_creation_token_details", None) + is None + ) + + def test_transform_usage_without_cache_details_stays_none(): """No cacheDetails in the response (older models/regions) should leave cache_creation_token_details unset, same as before this field existed.""" From 785eed616fdb51f73e0c0b0bf7599c6834f7ce23 Mon Sep 17 00:00:00 2001 From: Siraj637909 Date: Sun, 16 Aug 2026 18:28:45 +0530 Subject: [PATCH 014/281] fix(proxy): strip extra_headers/headers/aws_session_token from GET /health (gh-36898) `/health` already stripped `api_key` from each deployment row via `ILLEGAL_DISPLAY_PARAMS`, but `extra_headers`, `headers`, and `aws_session_token` were never added to that list, so `GET /health` leaked provider credentials (Azure `api-key`, Google `x-goog-api-key`, Bearer tokens, AWS session tokens) in plaintext to any caller, even without a master key. Add those three fields to `ILLEGAL_DISPLAY_PARAMS` so `_clean_endpoint_data()` omits them for all callers, matching how `api_key` is already handled. Fixes #36898 --- litellm/proxy/health_check.py | 3 ++ .../health_endpoints/test_health_endpoints.py | 29 +++++++++++++++++++ 2 files changed, 32 insertions(+) diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index f9d408fb7de..83919e1ddee 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -27,6 +27,9 @@ ILLEGAL_DISPLAY_PARAMS: Final = [ "vertex_credentials", "aws_access_key_id", "aws_secret_access_key", + "aws_session_token", + "extra_headers", + "headers", "exception", # internal; not JSON-serializable, never for display "litellm_metadata", # internal tracking metadata with auth objects; not for display ] diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index e2705bd5fec..2d3c90a2b9e 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -2367,6 +2367,35 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): assert cleaned.get("api_version") == "2024-10-21" +def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): + """ + gh-36898: GET /health must not leak provider credentials that live in + `extra_headers` / `headers` / `aws_session_token`. Before the fix these + were returned in plaintext (api_key was stripped, but these were not). + """ + from litellm.proxy.health_check import _clean_endpoint_data + + raw = { + "model": "openai/gpt-4o", + "api_base": "https://example.test/v1", + "extra_headers": { + "Authorization": "Bearer CANARY_EXTRA_HEADERS_AUTHORIZATION", + "x-goog-api-key": "CANARY_X_GOOG_API_KEY_VALUE", + "api-key": "CANARY_AZURE_STYLE_API_KEY", + }, + "headers": {"X-Custom": "CANARY_HEADER_VALUE"}, + "aws_session_token": "CANARY_AWS_SESSION_TOKEN_VALUE", + } + + cleaned = _clean_endpoint_data(raw, details=True) + + assert "extra_headers" not in cleaned + assert "headers" not in cleaned + assert "aws_session_token" not in cleaned + # routing/admin field still present + assert cleaned.get("api_base") == "https://example.test/v1" + + class TestConfigBaseForHealthCheck: """A request that sets its own connection fields gets a base without the configuration's credentials; anything it leaves unset still comes from From eee86f1e527195ce00bc52415105b953d3141f04 Mon Sep 17 00:00:00 2001 From: Ben Younes <2910651+ousamabenyounes@users.noreply.github.com> Date: Mon, 10 Aug 2026 12:11:23 +0000 Subject: [PATCH 015/281] fix(vertex_ai): bill Gemini grounding per unique web search query Gemini 3 per_query grounding is billed per unique search query the model executes, ignoring empty queries. _calculate_web_search_requests summed every non-empty webSearchQueries string across grounding metadata items, so repeated queries within a request inflated web_search_requests and overstated cost. Count distinct non-empty queries across items instead. Fixes #36377 --- .../vertex_and_google_ai_studio_gemini.py | 19 +++++++++--------- ...test_vertex_and_google_ai_studio_gemini.py | 20 +++++++++++++++++++ 2 files changed, 29 insertions(+), 10 deletions(-) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index d298670aa7a..ba2f91ce69d 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -1978,16 +1978,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): @staticmethod def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: - web_search_requests: int | None = None - - if grounding_metadata and isinstance(grounding_metadata, list) and len(grounding_metadata) > 0: - for grounding_metadata_item in grounding_metadata: - web_search_queries = grounding_metadata_item.get("webSearchQueries") - if web_search_queries and web_search_requests: - web_search_requests += len([q for q in web_search_queries if q]) - elif web_search_queries: - web_search_requests = len([q for q in web_search_queries if q]) - return web_search_requests + if not (grounding_metadata and isinstance(grounding_metadata, list)): + return None + unique_queries: Final = { + query + for grounding_metadata_item in grounding_metadata + for query in (grounding_metadata_item.get("webSearchQueries") or []) + if query + } + return len(unique_queries) or None @staticmethod def _create_streaming_choice( diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 51cc2857252..d14fb6021ed 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -5553,3 +5553,23 @@ def test_accumulated_json_skips_non_dict_leading_value(): assert len(out) == 1 assert out[0].choices[0].delta.content == "a" + + +def test_calculate_web_search_requests_counts_unique_queries(): + """Gemini 3 per_query billing charges per unique query executed, not per emitted string. + + Regression for #36377: duplicate webSearchQueries within and across grounding + metadata items must collapse to the distinct-query count, and empty strings must + be ignored, matching Google's documented Grounding-with-Search billing rule. + """ + duplicates_in_one_item = [{"webSearchQueries": ["euro 2024 winner", "euro 2024 winner", "spain england final", ""]}] + assert VertexGeminiConfig._calculate_web_search_requests(duplicates_in_one_item) == 2 + + duplicates_across_items = [ + {"webSearchQueries": ["euro 2024 winner"]}, + {"webSearchQueries": ["euro 2024 winner", "spain england final"]}, + ] + assert VertexGeminiConfig._calculate_web_search_requests(duplicates_across_items) == 2 + + assert VertexGeminiConfig._calculate_web_search_requests([]) is None + assert VertexGeminiConfig._calculate_web_search_requests([{"webSearchQueries": ["", ""]}]) is None From 2bdae174e391cba9b05ef651ecc8578d540fcb68 Mon Sep 17 00:00:00 2001 From: Ben Younes <2910651+ousamabenyounes@users.noreply.github.com> Date: Tue, 11 Aug 2026 18:07:35 +0000 Subject: [PATCH 016/281] test(vertex_ai): annotate web-search regression vars as Final Address Greptile review on #36397: duplicates_in_one_item and duplicates_across_items lacked Final declarations (LIT010). Use bare : Final so the inferred type stays list-based, avoiding an explicit mutable annotation (LIT001), and ratchet the LIT010 budget down by one. RED to GREEN: both vars flagged LIT010 before -> clean after; mapped suite 146 passed, 100% diff coverage. --- type-discipline-budget.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 8e55b1533ea..4eaf54c14c3 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16715 + "limit": 16743 }, "LIT011": { "limit": 5593 From 5d7dee710b5c3956697be1a8edf0489504edc79d Mon Sep 17 00:00:00 2001 From: Ousama Ben Younes Date: Sat, 15 Aug 2026 08:07:20 +0000 Subject: [PATCH 017/281] test(vertex_ai): actually annotate web-search regression vars as Final Address the Greptile review on #36397. The earlier commit only ratcheted the LIT010 budget; it never applied the annotations, so duplicates_in_one_item and duplicates_across_items were still bound without a Final declaration (LIT010) and the first fixture line was at the 120-char ceiling. Annotate both with `: Final` and wrap the long literal. RED -> GREEN: check_type_discipline flagged both vars LIT010 before -> LIT010 gone after (file total 551 -> 549, LIT002 unchanged at 953); test_calculate_web_search_requests_counts_unique_queries still passes. --- .../gemini/test_vertex_and_google_ai_studio_gemini.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index d14fb6021ed..3a04424a581 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -2,7 +2,7 @@ import asyncio import json import re from copy import deepcopy -from typing import List, cast +from typing import Final, List, cast from unittest.mock import MagicMock, patch import pytest @@ -5562,10 +5562,12 @@ def test_calculate_web_search_requests_counts_unique_queries(): metadata items must collapse to the distinct-query count, and empty strings must be ignored, matching Google's documented Grounding-with-Search billing rule. """ - duplicates_in_one_item = [{"webSearchQueries": ["euro 2024 winner", "euro 2024 winner", "spain england final", ""]}] + duplicates_in_one_item: Final = [ + {"webSearchQueries": ["euro 2024 winner", "euro 2024 winner", "spain england final", ""]} + ] assert VertexGeminiConfig._calculate_web_search_requests(duplicates_in_one_item) == 2 - duplicates_across_items = [ + duplicates_across_items: Final = [ {"webSearchQueries": ["euro 2024 winner"]}, {"webSearchQueries": ["euro 2024 winner", "spain england final"]}, ] From 8bb41e52f0441a934cbb9f2079d1694fea399a56 Mon Sep 17 00:00:00 2001 From: ousamabenyounes Date: Sun, 16 Aug 2026 23:01:34 +0000 Subject: [PATCH 018/281] chore(type-discipline): reset LIT010 budget to base (fix is net -1) The Final annotations on the new regression vars make the PR's net LIT010 delta -1 (one fewer than base), so the earlier bump to 16743 was an over-estimate. Reset the limit to the base value 16715 so the one-way budget ratchet passes; the codebase-wide total (16714) stays under it. --- type-discipline-budget.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 4eaf54c14c3..8e55b1533ea 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16743 + "limit": 16715 }, "LIT011": { "limit": 5593 From cd7cdb3e3a8e61b207cb187a6e1a204273acccfd Mon Sep 17 00:00:00 2001 From: Srivatsa03 Date: Tue, 18 Aug 2026 20:44:45 -0500 Subject: [PATCH 019/281] fix(cost): stop double-billing cached tokens that overlap a modality Providers report cached_tokens and image_tokens as overlapping subsets of prompt_tokens rather than a disjoint partition, so a request whose images were served from cache paid for them twice, once at the cache-read rate and again at the image or input rate. The synthetic case in the issue came out at 109e-6 against a correct 39e-6. Clamp each modality to the part of the request the cache did not already cover, so the billed components still sum to prompt_tokens Fixes #37281 --- .../litellm_core_utils/llm_cost_calc/utils.py | 19 ++++++++--- .../llm_cost_calc/test_llm_cost_calc_utils.py | 34 +++++++++++++++++++ 2 files changed, 49 insertions(+), 4 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 9d6ad8b6e39..37774822565 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -854,11 +854,22 @@ def generic_cost_per_token( total_details: Final = text_tokens + cache_hit + audio_tokens + cache_creation + image_tokens + video_tokens has_double_counting: Final = (cache_hit > 0 or cache_creation > 0) and total_details > usage.prompt_tokens - if (text_tokens == 0 and prompt_tokens_details["image_count"] == 0) or has_double_counting: - text_tokens = usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens + if has_double_counting: + # cached and per-modality counts are both subsets of prompt_tokens and may overlap, so a + # modality can only bill what the cache did not already cover or the overlap is billed twice + uncached_budget: Final = max(usage.prompt_tokens - cache_hit - cache_creation, 0) + billable_audio: Final = min(audio_tokens, uncached_budget) + billable_image: Final = min(image_tokens, uncached_budget - billable_audio) + billable_video: Final = min(video_tokens, uncached_budget - billable_audio - billable_image) + prompt_tokens_details["audio_tokens"] = billable_audio + prompt_tokens_details["image_tokens"] = billable_image + prompt_tokens_details["video_tokens"] = billable_video + prompt_tokens_details["text_tokens"] = uncached_budget - billable_audio - billable_image - billable_video + elif text_tokens == 0 and prompt_tokens_details["image_count"] == 0: # Clamp to zero: inconsistent streaming usage - text_tokens = max(text_tokens, 0) - prompt_tokens_details["text_tokens"] = text_tokens + prompt_tokens_details["text_tokens"] = max( + usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0 + ) ( prompt_base_cost, diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 4d157e74482..486f3a81331 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -1358,6 +1358,40 @@ def test_string_cost_values(): assert round(completion_cost, 12) == round(expected_completion_cost, 12) +def test_generic_cost_per_token_overlapping_cached_and_image_tokens(): + """Some providers report cached_tokens and image_tokens as overlapping subsets of + prompt_tokens. Billing each in full charged the overlap twice, once at the cache rate + and again at the input rate.""" + model = "litellm-test-overlapping-cached-image" + litellm.register_model( + { + model: { + "litellm_provider": "openai", + "mode": "chat", + "input_cost_per_token": 1e-6, + "cache_read_input_token_cost": 1e-7, + "output_cost_per_token": 2e-6, + } + } + ) + usage = Usage( + prompt_tokens=100, + completion_tokens=10, + total_tokens=110, + prompt_tokens_details=PromptTokensDetailsWrapper( + text_tokens=None, cached_tokens=90, image_tokens=80 + ), + ) + + prompt_cost, completion_cost = generic_cost_per_token( + model=model, usage=usage, custom_llm_provider="openai" + ) + + # 90 cached at 1e-7, the remaining 10 uncached tokens once at 1e-6 + assert prompt_cost == pytest.approx(90 * 1e-7 + 10 * 1e-6) + assert completion_cost == pytest.approx(10 * 2e-6) + + def test_calculate_cost_component_with_string_values(): """Test the calculate_cost_component function directly with string cost values.""" from litellm.litellm_core_utils.llm_cost_calc.utils import calculate_cost_component From b8680e6baed05863712a4c57a10b128ecd95475a Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 19 Aug 2026 18:22:31 +0000 Subject: [PATCH 020/281] fix(ui): render tag-based guardrail mode instead of crashing guardrails page Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_components/guardrailTableColumns.tsx | 13 +++++--- .../_components/guardrail_info.test.tsx | 30 +++++++++++++++++++ .../guardrails/_components/guardrail_info.tsx | 5 ++-- .../guardrail_info_helpers.test.tsx | 29 ++++++++++++++++++ .../_components/guardrail_info_helpers.tsx | 13 ++++++++ .../_components/guardrail_table.test.tsx | 12 ++++++++ .../src/components/guardrails/types.ts | 7 ++++- 7 files changed, 102 insertions(+), 7 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx index ec3d05a6907..53f1b1a03d3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx @@ -15,7 +15,7 @@ import { } from "@/components/ui/dropdown-menu"; import { cn } from "@/lib/cva.config"; -import { getGuardrailLogoAndName } from "./guardrail_info_helpers"; +import { formatGuardrailMode, getGuardrailLogoAndName } from "./guardrail_info_helpers"; import { Logo } from "@/components/molecules/logo/Logo"; const CONFIG_DELETE_HINT = "Config guardrails are defined in the config file and cannot be deleted from the dashboard."; @@ -117,9 +117,14 @@ export const getGuardrailTableColumns = ({ header: "Mode", size: 130, enableSorting: false, - cell: ({ row }) => ( - {row.original.litellm_params.mode} - ), + cell: ({ row }) => { + const mode = formatGuardrailMode(row.original.litellm_params.mode); + return ( + + {mode || "-"} + + ); + }, }, { id: "default_on", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx index b6ee130d50a..3f6317ed366 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx @@ -82,6 +82,36 @@ describe("Guardrail Info", () => { expect(getByText("Settings")).toBeInTheDocument(); }); + it("should render a tag-based mode object rather than crashing the detail view", async () => { + vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ + guardrail_id: "123", + guardrail_name: "Test Guardrail", + litellm_params: { + guardrail: "bedrock", + mode: { tags: { "Service-Type: internal-service": "post_call" }, default: ["pre_call", "post_call"] }, + default_on: true, + }, + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-01T00:00:00Z", + guardrail_definition_location: "database", + }); + + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + supported_entities: [], + supported_actions: [], + pii_entity_categories: [], + supported_modes: ["pre_call", "post_call"], + }); + + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({}); + + const { findAllByText } = render( + {}} accessToken="123" isAdmin={true} />, + ); + + expect(await findAllByText("pre_call, post_call (tag-based)")).not.toHaveLength(0); + }); + it("should render the provider logo from the bundled guardrail logo map", async () => { vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ guardrail_id: "123", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index e80ddac932f..5e476a8accd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -35,6 +35,7 @@ import { import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; import { + formatGuardrailMode, getGuardrailLogoAndName, guardrail_provider_map, skipSystemMessageToChoice, @@ -559,7 +560,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,

Mode

-

{guardrailData.litellm_params?.mode || "-"}

+

{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}

{guardrailData.litellm_params?.default_on ? "Default On" : "Default Off"} @@ -852,7 +853,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,

Mode

-
{guardrailData.litellm_params?.mode || "-"}
+
{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}

Default On

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx index ec910673b8f..c5e07fe9624 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.test.tsx @@ -14,6 +14,7 @@ import { choiceToSkipSystemForCreate, skipToolMessageToChoice, choiceToSkipToolForCreate, + formatGuardrailMode, } from "./guardrail_info_helpers"; describe("guardrail_info_helpers", () => { @@ -210,6 +211,34 @@ describe("guardrail_info_helpers", () => { }); }); + describe("formatGuardrailMode", () => { + it("renders a single mode and a list of modes", () => { + expect(formatGuardrailMode("pre_call")).toBe("pre_call"); + expect(formatGuardrailMode(["pre_call", "post_call"])).toBe("pre_call, post_call"); + }); + + it("flattens a tag-based mode object into deduped modes instead of returning it verbatim", () => { + const mode = { + tags: { "Service-Type: internal-service": "post_call", "Service-Type: batch": ["during_call", "post_call"] }, + default: ["pre_call", "post_call"], + }; + + expect(formatGuardrailMode(mode)).toBe("pre_call, post_call, during_call (tag-based)"); + }); + + it("handles a tag-based mode with no default and with no tags", () => { + expect(formatGuardrailMode({ tags: { "team: a": "post_call" } })).toBe("post_call (tag-based)"); + expect(formatGuardrailMode({ default: "pre_call" })).toBe("pre_call (tag-based)"); + }); + + it("returns an empty string for missing or unusable modes", () => { + expect(formatGuardrailMode(undefined)).toBe(""); + expect(formatGuardrailMode(null)).toBe(""); + expect(formatGuardrailMode({})).toBe(""); + expect(formatGuardrailMode({ tags: {}, default: null })).toBe(""); + }); + }); + describe("skipSystemMessageToChoice / choiceToSkipSystemForCreate", () => { it("maps API values to form choices and back for create", () => { expect(skipSystemMessageToChoice(undefined)).toBe("inherit"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index 12aaba0d696..c12529e6326 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -110,6 +110,19 @@ export const toModeArray = (raw: unknown): string[] => { return []; }; +// Turns a guardrail mode into a renderable string. A mode is a single mode, a list of modes, or a +// tag-based `{ tags, default }` object, which React refuses to render as a child +export const formatGuardrailMode = (raw: unknown): string => { + const flat: string[] = toModeArray(raw); + if (flat.length > 0) return flat.join(", "); + if (raw === null || typeof raw !== "object") return ""; + + const { tags, default: fallback } = raw as { tags?: Record; default?: unknown }; + const tagged: string[] = tags && typeof tags === "object" ? Object.values(tags).flatMap(toModeArray) : []; + const modes: string[] = Array.from(new Set([...toModeArray(fallback), ...tagged])); + return modes.length > 0 ? `${modes.join(", ")} (tag-based)` : ""; +}; + // Resolves the supported modes for the selected provider, falling back to the global list export const getSupportedModesForProvider = ( settings: { supported_modes?: string[]; supported_modes_by_provider?: Record } | null, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx index ee619dc7468..561a89a191a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx @@ -46,6 +46,18 @@ describe("GuardrailTable", () => { expect(screen.getByText("m")).toBeInTheDocument(); }); + it("renders a tag-based mode object instead of crashing the table", () => { + const guardrail = makeGuardrail({ + litellm_params: { + guardrail: "bedrock", + mode: { tags: { "Service-Type: internal-service": "post_call" }, default: ["pre_call", "post_call"] }, + default_on: true, + }, + }); + render(); + expect(screen.getByText("pre_call, post_call (tag-based)")).toBeInTheDocument(); + }); + it("deletes a DB guardrail through the actions menu", async () => { const user = userEvent.setup(); const onDeleteClick = vi.fn(); diff --git a/ui/litellm-dashboard/src/components/guardrails/types.ts b/ui/litellm-dashboard/src/components/guardrails/types.ts index e8ed27d9e45..0f5ce1c883d 100644 --- a/ui/litellm-dashboard/src/components/guardrails/types.ts +++ b/ui/litellm-dashboard/src/components/guardrails/types.ts @@ -18,12 +18,17 @@ export interface PiiConfigurationProps { entityCategories?: PiiEntityCategory[]; } +export type GuardrailMode = + | string + | string[] + | { tags?: Record; default?: string | string[] | null }; + export interface Guardrail { guardrail_id: string; guardrail_name: string | null; litellm_params: { guardrail: string; - mode: string; + mode: GuardrailMode; default_on: boolean; pii_entities_config?: { [key: string]: string }; [key: string]: any; From 881aa2080871052a2173f7b3352df39fc0e61e03 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 19 Aug 2026 18:31:42 +0000 Subject: [PATCH 021/281] fix(ui): format tag-based guardrail mode in delete modal, playground, and policy picker Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrails/_components/GuardrailTestPlayground.tsx | 8 ++++++-- .../guardrails/_components/GuardrailsPanel.tsx | 4 ++-- .../(dashboard)/guardrails/_components/guardrail_info.tsx | 4 +++- .../policies/_components/guardrail_selection_modal.tsx | 5 ++++- 4 files changed, 15 insertions(+), 6 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPlayground.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPlayground.tsx index fd8ed22867b..c64b5d7cb5c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPlayground.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailTestPlayground.tsx @@ -6,13 +6,15 @@ import { toast } from "@/lib/toast"; import { Card, CardContent } from "@/components/ui/card"; import { InputGroup, InputGroupAddon, InputGroupInput } from "@/components/ui/input-group"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { GuardrailMode } from "@/components/guardrails/types"; +import { formatGuardrailMode } from "./guardrail_info_helpers"; interface GuardrailItem { guardrail_id?: string; guardrail_name: string | null; litellm_params: { guardrail: string; - mode: string; + mode: GuardrailMode; default_on: boolean; }; guardrail_info: Record | null; @@ -171,7 +173,9 @@ const GuardrailTestPlayground: React.FC = ({
Mode: - {guardrail.litellm_params.mode} + + {formatGuardrailMode(guardrail.litellm_params.mode)} +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailsPanel.tsx index b4c29bd9c40..7e59abf8e3d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/GuardrailsPanel.tsx @@ -18,7 +18,7 @@ import GuardrailTestPlayground from "./GuardrailTestPlayground"; import { toast } from "@/lib/toast"; import { Guardrail } from "@/components/guardrails/types"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; -import { getGuardrailLogoAndName } from "./guardrail_info_helpers"; +import { formatGuardrailMode, getGuardrailLogoAndName } from "./guardrail_info_helpers"; import { CustomCodeModal } from "./custom_code"; import GuardrailGarden from "./guardrail_garden"; import { TeamGuardrailsTab } from "./TeamGuardrailsTab"; @@ -211,7 +211,7 @@ const GuardrailsPanel: React.FC = ({ accessToken, userRole { label: "Name", value: guardrailToDelete?.guardrail_name }, { label: "ID", value: guardrailToDelete?.guardrail_id, code: true }, { label: "Provider", value: providerDisplayName }, - { label: "Mode", value: guardrailToDelete?.litellm_params.mode }, + { label: "Mode", value: formatGuardrailMode(guardrailToDelete?.litellm_params.mode) }, { label: "Default On", value: guardrailToDelete?.litellm_params.default_on ? "Yes" : "No", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index 5e476a8accd..d4a1885146d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -560,7 +560,9 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,

Mode

-

{formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"}

+

+ {formatGuardrailMode(guardrailData.litellm_params?.mode) || "-"} +

{guardrailData.litellm_params?.default_on ? "Default On" : "Default Off"} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/guardrail_selection_modal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/guardrail_selection_modal.tsx index 0b439462c1a..f87155db719 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/guardrail_selection_modal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/guardrail_selection_modal.tsx @@ -12,6 +12,7 @@ import { } from "@/components/ui/dialog"; import { Separator } from "@/components/ui/separator"; import { CheckCircle2, Info } from "lucide-react"; +import { formatGuardrailMode } from "@/app/(dashboard)/guardrails/_components/guardrail_info_helpers"; interface GuardrailInfo { guardrail_name: string; @@ -163,7 +164,9 @@ const GuardrailSelectionModal: React.FC = ({ {/* Show guardrail type and mode */}
{guardrail.definition?.litellm_params?.guardrail || "unknown"} - {guardrail.definition?.litellm_params?.mode || "unknown"} + + {formatGuardrailMode(guardrail.definition?.litellm_params?.mode) || "unknown"} + {guardrail.definition?.litellm_params?.patterns && ( {guardrail.definition.litellm_params.patterns.length} pattern(s) From f80cb0d9f8e37539b39bf6412ef7f673c2074e58 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 19 Aug 2026 18:33:06 +0000 Subject: [PATCH 022/281] refactor(ui): drop redundant comment above guardrail mode formatter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrails/_components/guardrail_info_helpers.tsx | 2 -- 1 file changed, 2 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index c12529e6326..83038b8e0e7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -110,8 +110,6 @@ export const toModeArray = (raw: unknown): string[] => { return []; }; -// Turns a guardrail mode into a renderable string. A mode is a single mode, a list of modes, or a -// tag-based `{ tags, default }` object, which React refuses to render as a child export const formatGuardrailMode = (raw: unknown): string => { const flat: string[] = toModeArray(raw); if (flat.length > 0) return flat.join(", "); From d317c5621fd13c5c0c9dcebc7e3afeb8f1c62aa9 Mon Sep 17 00:00:00 2001 From: Bisma Nawaz Date: Fri, 21 Aug 2026 02:56:23 +0500 Subject: [PATCH 023/281] fix: map Gemini ON_DEMAND_FLEX traffic type to flex service tier --- litellm/cost_calculator.py | 10 ++++--- .../llms/gemini/test_cost_calculator.py | 28 +++++++++++++++++++ 2 files changed, 34 insertions(+), 4 deletions(-) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 8f7cd09d364..46a42b616f8 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -830,9 +830,11 @@ def _get_response_model(completion_response: object) -> str | None: _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER: Final[dict] = { # ON_DEMAND_PRIORITY maps to "priority" — selects input_cost_per_token_priority, etc. "ON_DEMAND_PRIORITY": "priority", - # FLEX / BATCH maps to "flex" — selects input_cost_per_token_flex, etc. + # FLEX / BATCH / ON_DEMAND_FLEX maps to "flex" — selects input_cost_per_token_flex, etc. + # Vertex AI reports flex/shared-capacity traffic as ON_DEMAND_FLEX, not FLEX. "FLEX": "flex", "BATCH": "flex", + "ON_DEMAND_FLEX": "flex", # ON_DEMAND is standard pricing — no service_tier suffix applied "ON_DEMAND": None, } @@ -847,9 +849,9 @@ def _map_traffic_type_to_service_tier(traffic_type: str | None) -> str | None: trafficType values seen in practice ------------------------------------ - ON_DEMAND -> standard pricing (service_tier = None) - ON_DEMAND_PRIORITY -> priority pricing (service_tier = "priority") - FLEX / BATCH -> batch/flex pricing (service_tier = "flex") + ON_DEMAND -> standard pricing (service_tier = None) + ON_DEMAND_PRIORITY -> priority pricing (service_tier = "priority") + FLEX / BATCH / ON_DEMAND_FLEX -> batch/flex pricing (service_tier = "flex") """ if traffic_type is None: return None diff --git a/tests/test_litellm/llms/gemini/test_cost_calculator.py b/tests/test_litellm/llms/gemini/test_cost_calculator.py index 6917092966b..c44f29ba168 100644 --- a/tests/test_litellm/llms/gemini/test_cost_calculator.py +++ b/tests/test_litellm/llms/gemini/test_cost_calculator.py @@ -301,3 +301,31 @@ def test_gemini_image_generation_cost_no_web_search_when_absent(): ) assert cost_zero == cost_none + + +@pytest.mark.parametrize( + "traffic_type, expected_service_tier", + [ + ("ON_DEMAND", None), + ("ON_DEMAND_PRIORITY", "priority"), + ("FLEX", "flex"), + ("BATCH", "flex"), + # Vertex AI reports flex/shared-capacity traffic as ON_DEMAND_FLEX. + ("ON_DEMAND_FLEX", "flex"), + # trafficType is matched case-insensitively. + ("on_demand_flex", "flex"), + (None, None), + ("SOMETHING_UNKNOWN", None), + ], +) +def test_map_traffic_type_to_service_tier(traffic_type, expected_service_tier): + """ + Gemini/Vertex usageMetadata.trafficType maps to the LiteLLM service_tier + that selects flex/priority cost keys. ON_DEMAND_FLEX (Vertex's flex opt-in + value) must map to "flex" so flex-tier requests are not billed as standard. + """ + from litellm.cost_calculator import _map_traffic_type_to_service_tier + + assert ( + _map_traffic_type_to_service_tier(traffic_type) == expected_service_tier + ) From 909ab23b89589375d8037319ff32fea048d710fe Mon Sep 17 00:00:00 2001 From: Bisma Nawaz Date: Fri, 21 Aug 2026 03:42:34 +0500 Subject: [PATCH 024/281] test: annotate parametrized traffic-type test inputs --- tests/test_litellm/llms/gemini/test_cost_calculator.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/llms/gemini/test_cost_calculator.py b/tests/test_litellm/llms/gemini/test_cost_calculator.py index c44f29ba168..1f7bfa69527 100644 --- a/tests/test_litellm/llms/gemini/test_cost_calculator.py +++ b/tests/test_litellm/llms/gemini/test_cost_calculator.py @@ -318,7 +318,9 @@ def test_gemini_image_generation_cost_no_web_search_when_absent(): ("SOMETHING_UNKNOWN", None), ], ) -def test_map_traffic_type_to_service_tier(traffic_type, expected_service_tier): +def test_map_traffic_type_to_service_tier( + traffic_type: str | None, expected_service_tier: str | None +): """ Gemini/Vertex usageMetadata.trafficType maps to the LiteLLM service_tier that selects flex/priority cost keys. ON_DEMAND_FLEX (Vertex's flex opt-in From 9e86cfa7e994edd3ac77456a7b0edb974e8012ff Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 01:56:02 +0000 Subject: [PATCH 025/281] fix(auth): support wildcard prefixes in jwt team_allowed_routes team_allowed_routes and admin_allowed_routes only matched exact strings or named route groups, so a whole prefix of pass-through endpoints had to be listed route by route in config. Match trailing-wildcard patterns with the same helper the key-level allowed_routes check uses, so "/prefix/*" covers endpoints registered later. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 2 +- litellm/proxy/auth/auth_checks.py | 5 +- litellm/proxy/auth/auth_utils.py | 2 +- litellm/proxy/auth/route_checks.py | 10 +-- litellm/proxy/policy_engine/policy_matcher.py | 4 +- .../policy_engine/policy_resolve_endpoints.py | 8 +- .../proxy/auth/test_auth_checks.py | 79 +++++++++++++++++++ .../policies/_components/scope_validation.ts | 2 +- 8 files changed, 96 insertions(+), 16 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 776aecbd883..46a06f77861 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -84,7 +84,7 @@ "limit": 56 }, "reportPrivateUsage": { - "limit": 1823 + "limit": 1817 }, "reportRedeclaration": { "limit": 8 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 12d6b44a648..9d8eedaa7dc 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1128,7 +1128,8 @@ def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: Parameters: - user_route: str - the route the user is trying to call - - allowed_routes: List[str|LiteLLMRoutes] - the list of allowed routes for the user. + - allowed_routes: List[str|LiteLLMRoutes] - the list of allowed routes for the user. Entries are a route group name + (e.g. "openai_routes"), an exact route, or a trailing-wildcard prefix (e.g. "/tempus/*"). """ from starlette.routing import compile_path @@ -1138,7 +1139,7 @@ def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: regex, _, _ = compile_path(template) if regex.match(user_route): return True - elif allowed_route == user_route: + elif RouteChecks.route_matches_wildcard_pattern(route=user_route, pattern=allowed_route): return True return False diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index ce662ee0374..1e6d8137d53 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -608,7 +608,7 @@ def route_in_additonal_public_routes(current_route: str): # Check wildcard patterns for route_pattern in routes_defined: - if RouteChecks._route_matches_wildcard_pattern(route=current_route, pattern=route_pattern): + if RouteChecks.route_matches_wildcard_pattern(route=current_route, pattern=route_pattern): return True return False diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index cea21ca088b..4dba2497bb9 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -181,7 +181,7 @@ class RouteChecks: # check if wildcard pattern is allowed for allowed_route in valid_token.allowed_routes: - if RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route): + if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route): return True if denied_auth_enforced_pass_through_route: @@ -329,7 +329,7 @@ class RouteChecks: route_allowed = True break - if RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route): + if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route): route_allowed = True break @@ -397,7 +397,7 @@ class RouteChecks: return True # Check for wildcard patterns like "/containers/*" if RouteChecks._is_wildcard_pattern(pattern=openai_route): - if RouteChecks._route_matches_wildcard_pattern(route=route, pattern=openai_route): + if RouteChecks.route_matches_wildcard_pattern(route=route, pattern=openai_route): return True # Check for Google routes with placeholders like "/v1beta/models/{model_name}:generateContent" @@ -517,7 +517,7 @@ class RouteChecks: return pattern.endswith("*") @staticmethod - def _route_matches_wildcard_pattern(route: str, pattern: str) -> bool: + def route_matches_wildcard_pattern(route: str, pattern: str) -> bool: """ Check if route matches the wildcard pattern @@ -594,7 +594,7 @@ class RouteChecks: # e.g calling /anthropic/v1/messages is allowed if allowed_routes has /anthropic/* ######################################################### if any( - RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route) + RouteChecks.route_matches_wildcard_pattern(route=route, pattern=allowed_route) for allowed_route in allowed_routes if RouteChecks._is_wildcard_pattern(pattern=allowed_route) ): diff --git a/litellm/proxy/policy_engine/policy_matcher.py b/litellm/proxy/policy_engine/policy_matcher.py index f66dc4e7bbe..001e4115374 100644 --- a/litellm/proxy/policy_engine/policy_matcher.py +++ b/litellm/proxy/policy_engine/policy_matcher.py @@ -30,7 +30,7 @@ class PolicyMatcher: """ Check if a value matches any of the given patterns. - Uses the existing RouteChecks._route_matches_wildcard_pattern helper. + Uses the existing RouteChecks.route_matches_wildcard_pattern helper. Args: value: The value to check (e.g., team alias, key alias, model) @@ -45,7 +45,7 @@ class PolicyMatcher: for pattern in patterns: # Use existing wildcard pattern matching helper - if RouteChecks._route_matches_wildcard_pattern(route=value, pattern=pattern): + if RouteChecks.route_matches_wildcard_pattern(route=value, pattern=pattern): return True return False diff --git a/litellm/proxy/policy_engine/policy_resolve_endpoints.py b/litellm/proxy/policy_engine/policy_resolve_endpoints.py index 346586c1e5a..70b98933d0f 100644 --- a/litellm/proxy/policy_engine/policy_resolve_endpoints.py +++ b/litellm/proxy/policy_engine/policy_resolve_endpoints.py @@ -100,7 +100,7 @@ def _filter_keys_by_tags(keys: list, tag_patterns: list) -> tuple: key_alias = key.key_alias or "" key_tags = _get_tags_from_metadata(key.metadata, getattr(key, "metadata_json", None)) if key_tags and any( - RouteChecks._route_matches_wildcard_pattern(route=tag, pattern=pat) + RouteChecks.route_matches_wildcard_pattern(route=tag, pattern=pat) for tag in key_tags for pat in tag_patterns ): @@ -123,7 +123,7 @@ def _filter_teams_by_tags(teams: list, tag_patterns: list) -> tuple: team_alias = team.team_alias or "" team_tags = _get_tags_from_metadata(team.metadata) if team_tags and any( - RouteChecks._route_matches_wildcard_pattern(route=tag, pattern=pat) + RouteChecks.route_matches_wildcard_pattern(route=tag, pattern=pat) for tag in team_tags for pat in tag_patterns ): @@ -152,7 +152,7 @@ async def _find_affected_by_team_patterns( for team in all_teams: team_alias = team.team_alias or "" if team_alias and any( - RouteChecks._route_matches_wildcard_pattern(route=team_alias, pattern=pat) for pat in team_patterns + RouteChecks.route_matches_wildcard_pattern(route=team_alias, pattern=pat) for pat in team_patterns ): if team_alias not in existing_teams: new_teams.append(team_alias) @@ -190,7 +190,7 @@ async def _find_affected_keys_by_alias(prisma_client: object, key_patterns: list for key in keys: key_alias = key.key_alias or "" if key_alias and any( - RouteChecks._route_matches_wildcard_pattern(route=key_alias, pattern=pat) for pat in key_patterns + RouteChecks.route_matches_wildcard_pattern(route=key_alias, pattern=pat) for pat in key_patterns ): if key_alias not in existing_keys: affected.append(key_alias) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 6b40fa1b324..0b174cda9d5 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6897,3 +6897,82 @@ def test_model_has_no_cost_mapping_alias_to_a_group_priced_through_model_info_is router = _router_with_a_group_priced_through_model_info() assert model_has_no_cost_mapping(model="model-info-priced-alias", llm_router=router) is False + +@pytest.mark.parametrize( + "user_route, expected", + [ + ("/tempus/v1/chat/completions", True), + ("/tempus/newly-registered-model/predict", True), + ("/tempus-other/v1/chat/completions", False), + ("/anthropic/v1/messages", False), + ], +) +def test_team_allowed_routes_wildcard_prefix_matches_unregistered_passthrough_routes(user_route, expected): + """A `/prefix/*` entry in `team_allowed_routes` must cover every route under that prefix, so + passthrough endpoints registered after the proxy config was written are reachable without an + exact-route config change.""" + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.auth_checks import allowed_routes_check + + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, + user_route=user_route, + litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=["/tempus/*"]), + ) + is expected + ) + + +def test_team_allowed_routes_exact_route_does_not_become_a_prefix_grant(): + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.auth_checks import allowed_routes_check + + roles = LiteLLM_JWTAuth(team_allowed_routes=["/tempus/model-a"]) + + assert ( + allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/tempus/model-a", litellm_proxy_roles=roles) + is True + ) + assert ( + allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/tempus/model-b", litellm_proxy_roles=roles) + is False + ) + + +def test_admin_allowed_routes_wildcard_prefix_is_honored(): + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.auth_checks import allowed_routes_check + + roles = LiteLLM_JWTAuth(admin_allowed_routes=["/tempus/*"]) + + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/tempus/anything", litellm_proxy_roles=roles + ) + is True + ) + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/other/anything", litellm_proxy_roles=roles + ) + is False + ) + + +def test_team_allowed_routes_named_route_group_still_resolves(): + from litellm.proxy._types import LiteLLM_JWTAuth + from litellm.proxy.auth.auth_checks import allowed_routes_check + + roles = LiteLLM_JWTAuth(team_allowed_routes=["openai_routes"]) + + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, user_route="/v1/chat/completions", litellm_proxy_roles=roles + ) + is True + ) + assert ( + allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/key/generate", litellm_proxy_roles=roles) + is False + ) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/scope_validation.ts b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/scope_validation.ts index 7c49117088c..53a76dc5a1c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/scope_validation.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/policies/_components/scope_validation.ts @@ -1,4 +1,4 @@ -// Mirrors request-time matching (RouteChecks._route_matches_wildcard_pattern): only a +// Mirrors request-time matching (RouteChecks.route_matches_wildcard_pattern): only a // trailing "*" is a wildcard (prefix match). Anything else - including a "?" or a // non-trailing "*" - is compared by exact equality when a request is matched, so it is // treated as a concrete alias that must exist. From 07416344cc8865c1867c51dd733582e04236aeef Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 02:11:18 +0000 Subject: [PATCH 026/281] test(auth): use a generic route prefix in wildcard route tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 2 +- .../proxy/auth/test_auth_checks.py | 18 +++++++++--------- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 9d8eedaa7dc..bf7a6a8f6c3 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1129,7 +1129,7 @@ def _allowed_routes_check(user_route: str, allowed_routes: list) -> bool: Parameters: - user_route: str - the route the user is trying to call - allowed_routes: List[str|LiteLLMRoutes] - the list of allowed routes for the user. Entries are a route group name - (e.g. "openai_routes"), an exact route, or a trailing-wildcard prefix (e.g. "/tempus/*"). + (e.g. "openai_routes"), an exact route, or a trailing-wildcard prefix (e.g. "/internal-models/*"). """ from starlette.routing import compile_path diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 0b174cda9d5..7fa16508054 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -6901,9 +6901,9 @@ def test_model_has_no_cost_mapping_alias_to_a_group_priced_through_model_info_is @pytest.mark.parametrize( "user_route, expected", [ - ("/tempus/v1/chat/completions", True), - ("/tempus/newly-registered-model/predict", True), - ("/tempus-other/v1/chat/completions", False), + ("/internal-models/v1/chat/completions", True), + ("/internal-models/newly-registered-model/predict", True), + ("/internal-models-other/v1/chat/completions", False), ("/anthropic/v1/messages", False), ], ) @@ -6918,7 +6918,7 @@ def test_team_allowed_routes_wildcard_prefix_matches_unregistered_passthrough_ro allowed_routes_check( user_role=LitellmUserRoles.TEAM, user_route=user_route, - litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=["/tempus/*"]), + litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=["/internal-models/*"]), ) is expected ) @@ -6928,14 +6928,14 @@ def test_team_allowed_routes_exact_route_does_not_become_a_prefix_grant(): from litellm.proxy._types import LiteLLM_JWTAuth from litellm.proxy.auth.auth_checks import allowed_routes_check - roles = LiteLLM_JWTAuth(team_allowed_routes=["/tempus/model-a"]) + roles = LiteLLM_JWTAuth(team_allowed_routes=["/internal-models/model-a"]) assert ( - allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/tempus/model-a", litellm_proxy_roles=roles) + allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-a", litellm_proxy_roles=roles) is True ) assert ( - allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/tempus/model-b", litellm_proxy_roles=roles) + allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-b", litellm_proxy_roles=roles) is False ) @@ -6944,11 +6944,11 @@ def test_admin_allowed_routes_wildcard_prefix_is_honored(): from litellm.proxy._types import LiteLLM_JWTAuth from litellm.proxy.auth.auth_checks import allowed_routes_check - roles = LiteLLM_JWTAuth(admin_allowed_routes=["/tempus/*"]) + roles = LiteLLM_JWTAuth(admin_allowed_routes=["/internal-models/*"]) assert ( allowed_routes_check( - user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/tempus/anything", litellm_proxy_roles=roles + user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/internal-models/anything", litellm_proxy_roles=roles ) is True ) From 4f5e290f60f518e9e1f28333169901d7207efa96 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 22 Aug 2026 08:08:42 +0000 Subject: [PATCH 027/281] refactor(proxy): type the per-model budget plumbing added yesterday Drops a pyright suppression, getattr string access, and bare dict annotations from the model_max_budget code, and trims a comment referencing its own PR. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 10 +++++----- litellm/llms/anthropic/common_utils.py | 4 ++-- .../context_management/editors/compact.py | 4 ++-- litellm/proxy/_types.py | 6 +++--- litellm/proxy/auth/user_api_key_auth.py | 11 +++++------ type-discipline-budget.json | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 +++- 7 files changed, 21 insertions(+), 20 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 664e1669834..9dcc46ebcee 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 19955 + "limit": 19954 }, "reportArgumentType": { "limit": 2566 @@ -57,7 +57,7 @@ "limit": 5663 }, "reportMissingTypeArgument": { - "limit": 15555 + "limit": 15553 }, "reportMissingTypeStubs": { "limit": 40 @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 39011 + "limit": 39007 }, "reportUnknownParameterType": { - "limit": 19885 + "limit": 19883 }, "reportUnknownVariableType": { - "limit": 30569 + "limit": 30568 }, "reportUnnecessaryCast": { "limit": 117 diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 3297aa95715..c1a2384e693 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -4,7 +4,7 @@ This file contains common utils for anthropic calls. import copy import re -from collections.abc import Mapping, Sequence +from collections.abc import Mapping, MutableMapping, Sequence from datetime import datetime, timezone from types import MappingProxyType from typing import Any, Final, Literal @@ -443,7 +443,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): @staticmethod def maybe_drop_disabled_thinking( model: str, - optional_params: dict, # mutable-ok: in-place out-param, same contract as AnthropicConfig._maybe_drop_speed_param + optional_params: MutableMapping[str, object], # mutable-ok: in-place out-param, as in _maybe_drop_speed_param custom_llm_provider: str, ) -> None: """Omit ``thinking={'type': 'disabled'}`` for always-on-thinking models diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py index 2a87afb5990..c8cbbba8784 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py @@ -352,8 +352,8 @@ async def _check_summary_model_budget( ) return False - user_model_max_budget: Final = getattr(user_api_key_auth, "user_model_max_budget", None) - user_id: Final = getattr(user_api_key_auth, "user_id", None) + user_model_max_budget: Final = user_api_key_auth.user_model_max_budget + user_id: Final = user_api_key_auth.user_id if isinstance(user_model_max_budget, dict) and user_model_max_budget and user_id is not None: try: await model_max_budget_limiter.is_user_within_model_budget( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0840d37ffa1..d5338811a1f 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2808,7 +2808,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # Values stay `object` rather than BudgetConfig: this is the raw JSON column, # and validating it here would make one malformed row fail auth outright. # resolve_model_budget validates the single entry a request actually needs. - user_model_max_budget: dict[str, object] | None = None + user_model_max_budget: Mapping[str, object] | None = None request_route: str | None = None is_session_token: bool = False # Server-only marker set exclusively by the MCP gateway admission path @@ -2986,8 +2986,8 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase): sso_user_id: str | None = None teams: list[str] = [] # Just team IDs, not full team objects object_permission: LiteLLM_ObjectPermissionTable | None = None - model_max_budget: dict | None = None - model_max_budget_usage: dict | None = None + model_max_budget: Mapping[str, object] | None = None + model_max_budget_usage: Mapping[str, Mapping[str, object]] | None = None from litellm.models.config import LiteLLM_Config as LiteLLM_Config # noqa: E402 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index fe4f1ee4ae5..d4fc091cc84 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -197,9 +197,9 @@ async def _read_user_model_max_budget( user_id: str | None, prisma_client: PrismaClient | None, user_api_key_cache: UserApiKeyCache, - parent_otel_span: object, + parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging, -) -> dict | None: +) -> Mapping[str, object] | None: """The user row's `model_max_budget`, or None when the row cannot be read. A user whose row is missing must not be refused: this is a budget lookup, @@ -213,13 +213,13 @@ async def _read_user_model_max_budget( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, user_id_upsert=False, - parent_otel_span=parent_otel_span, # pyright: ignore[reportArgumentType] # Span is a runtime union, not usable in an annotation here + parent_otel_span=parent_otel_span, proxy_logging_obj=proxy_logging_obj, ) except Exception as e: # noqa: BLE001 # mirrors the main path's tolerance verbose_logger.debug("Unable to read user for the per-model budget check: %s", e) return None - return getattr(user_obj, "model_max_budget", None) + return user_obj.model_max_budget if user_obj is not None else None async def _check_user_model_budget( @@ -3168,8 +3168,7 @@ async def _run_post_custom_auth_checks( # loaded the user row yet. The attach is unconditional because the post-call # spend hook reads this field off the token: gating it on the same condition # as enforcement would leave the user's counter uncharged whenever this - # request was not itself enforceable, which is the untracked-spend bug this - # PR exists to fix. + # request was not itself enforceable, so its spend would go untracked. user_budget: Final = await _read_user_model_max_budget( user_id=valid_token.user_id, prisma_client=prisma_client, diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 627811a7f1d..81f7c6aa40b 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22805 + "limit": 22801 }, "LIT002": { "limit": 26873 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cf55dc69e86..156238b99e2 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -36200,7 +36200,9 @@ export interface components { } | null; /** Model Max Budget Usage */ model_max_budget_usage?: { - [key: string]: unknown; + [key: string]: { + [key: string]: unknown; + }; } | null; /** * Models From e45c084c1c06862b9a5e9e3c089fca36beb00ed1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 21 Aug 2026 08:05:32 +0000 Subject: [PATCH 028/281] chore(typing): replace Any and bare containers added in the last day Type the annotations that landed in the last 24 hours and ratchet the lint budgets down accordingly. No behavior change. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/otel/plumbing/routing.py | 5 +++-- litellm/proxy/common_request_processing.py | 9 ++++++--- litellm/proxy/common_utils/reset_budget_job.py | 2 +- litellm/proxy/spend_tracking/budget_reservation.py | 2 +- ruff-strict-budget.json | 4 ++-- ..._experimental_pass_through_adapters_transformation.py | 2 +- .../proxy/common_utils/test_reset_budget_job.py | 4 ++-- 7 files changed, 16 insertions(+), 12 deletions(-) diff --git a/litellm/integrations/otel/plumbing/routing.py b/litellm/integrations/otel/plumbing/routing.py index 6a04dbb9bc8..d2457b9ce57 100644 --- a/litellm/integrations/otel/plumbing/routing.py +++ b/litellm/integrations/otel/plumbing/routing.py @@ -15,7 +15,7 @@ from collections import OrderedDict from collections.abc import Mapping from dataclasses import dataclass from types import MappingProxyType -from typing import Any, Final, TypeAlias +from typing import Final, TypeAlias from urllib.parse import quote from opentelemetry.sdk.trace import TracerProvider @@ -32,6 +32,7 @@ from litellm.integrations.otel.presets import ( dynamic_otlp_headers, project_routing_headers, ) +from litellm.types.utils import StandardCallbackDynamicParams # Exporter kinds that ignore headers — never rewritten with dynamic credentials. _NON_OTLP_KINDS: Final = ("console", "in_memory", "inmemory", "memory") @@ -166,7 +167,7 @@ class TenantTracerCache: def route_for( self, default: Tracer, - dynamic_params: Any, + dynamic_params: StandardCallbackDynamicParams | None, auth_metadata: Mapping[str, str] | None = None, ) -> TenantRoute: """Return the tracer (and trace-detachment flag) for this request. diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index dbbf9cb673e..3d097051f2f 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -4,7 +4,7 @@ import json import logging import math import traceback -from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType @@ -279,7 +279,7 @@ def _deferred_stream_logging_is_armed(request_data: dict) -> bool: ) -def _assembled_model_came_from_a_later_chunk(chunks: list, assembled_model: object) -> bool: +def _assembled_model_came_from_a_later_chunk(chunks: Sequence[object], assembled_model: object) -> bool: """Report whether stream_chunk_builder picked a model the first chunk did not carry. Azure Model Router puts the routed model on the chunks after the first one, and the @@ -301,7 +301,10 @@ def _assembled_model_came_from_a_later_chunk(chunks: list, assembled_model: obje ) -def _assembled_model_is_the_name_the_client_asked_for(request_data: dict, assembled_model: object) -> bool: +def _assembled_model_is_the_name_the_client_asked_for( + request_data: Mapping[str, object], + assembled_model: object, +) -> bool: """Report whether the assembled model is the public name the proxy stamps onto chunks. That stamp is what leaves an unpriced alias on the partial response, so the deployment's diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 8fcb184b26a..bc750243574 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -1182,7 +1182,7 @@ class ResetBudgetJob: if not raw: continue row_id: str = row[source.id_column] - windows: list = raw if isinstance(raw, list) else json.loads(raw) + windows: list[dict[str, object]] = raw if isinstance(raw, list) else json.loads(raw) changed = False for window in windows: counter_key = f"{source.counter_prefix}:{row_id}:window:{window['budget_duration']}" diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index ce6c9330620..87d9aa01a08 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -1261,7 +1261,7 @@ def _count_input_tokens_for_models( _INPUT_SIZE_FIELDS: Final = ("messages", "prompt", "input", "query", "documents", "tools", "tool_choice") -def _approximate_input_size(request_body: dict) -> int: +def _approximate_input_size(request_body: Mapping[str, object]) -> int: """Length of the request's input text, a cheap stand-in for tokenizing cost. Every field _count_input_tokens hands the tokenizer is sized here, and diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index a990f7c3830..e651600eb8f 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 133 }, "ANN401": { - "limit": 1188 + "limit": 1187 }, "ASYNC230": { "limit": 11 @@ -234,7 +234,7 @@ "limit": 5 }, "TID251": { - "limit": 1212 + "limit": 1211 }, "TRY002": { "limit": 524 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index e4dacc308dc..d4168d88818 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -635,7 +635,7 @@ def test_translate_anthropic_to_openai_orders_top_level_and_midturn_system(): def _translate_with_metadata( - model: str, metadata: dict[str, Any], custom_llm_provider: str | None + model: str, metadata: dict[str, str], custom_llm_provider: str | None ) -> dict[str, Any]: openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( anthropic_message_request={ diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 25c177a308d..1d045732c76 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -1948,14 +1948,14 @@ class FakePodLockManager: if self.redis_cache is not None: self.redis_cache.async_get_cache = AsyncMock(return_value="another-pod" if held_by_other else None) self._acquired = acquired - self.acquire_calls: List[Dict[str, Any]] = [] + self.acquire_calls: List[Dict[str, str | int | None]] = [] self.release_calls: List[str] = [] @staticmethod def get_redis_lock_key(cronjob_id: str) -> str: return f"cronjob_lock:{cronjob_id}" - async def acquire_lock(self, cronjob_id: str, ttl: Any = None) -> bool: + async def acquire_lock(self, cronjob_id: str, ttl: int | None = None) -> bool: self.acquire_calls.append({"cronjob_id": cronjob_id, "ttl": ttl}) return self._acquired From d6feb35a0429d786d68c0818ae151523bac78bc1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 20 Aug 2026 08:15:27 +0000 Subject: [PATCH 029/281] chore(typing): tighten annotations added in the last day and ratchet budgets Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/batches/batch_utils.py | 4 ++-- litellm/integrations/prometheus.py | 4 +++- .../proxy/management_endpoints/key_management_endpoints.py | 3 +-- litellm/repositories/model_repository.py | 2 +- ruff-strict-budget.json | 6 +++--- 5 files changed, 10 insertions(+), 9 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 6eb13d2cba7..2bc61aed771 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -551,7 +551,7 @@ def _get_batch_job_usage_from_response_body( return usage -def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> dict: +def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> Mapping[str, Any]: """ Get the ``result`` object from a line of an Anthropic message batch results JSONL file. @@ -563,7 +563,7 @@ def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[st def _get_response_from_batch_job_output_file( batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai" -) -> Any: +) -> Mapping[str, Any]: """ Get the response from the batch job output file """ diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index f9195db1d67..66756be6a6d 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -3552,7 +3552,9 @@ class PrometheusLogger(CustomLogger): except Exception as e: verbose_logger.exception("Error initializing user/team count metrics: %s", e) - async def _set_key_list_budget_metrics(self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken]): + async def _set_key_list_budget_metrics( + self, keys: list[str | UserAPIKeyAuth | LiteLLM_DeletedVerificationToken] + ) -> None: """Helper function to set budget metrics for a list of keys""" for key in keys: if isinstance(key, UserAPIKeyAuth): diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 54f567b7aa2..5fe5dda0ca4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2326,10 +2326,9 @@ async def _process_single_key_update( prisma_client=prisma_client, ) - _existing_row_metadata: Final = getattr(existing_key_row, "metadata", None) enforce_batch_enqueued_token_limit_is_admin_only( data=update_key_request, - existing_metadata=_existing_row_metadata if isinstance(_existing_row_metadata, dict) else None, + existing_metadata=existing_key_row.metadata, user_api_key_dict=user_api_key_dict, entity="key", ) diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index 27e23a39cc9..3965aeb2d49 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -36,7 +36,7 @@ class _ProxyModelActions(Protocol): class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): """Repository for proxy model database operations with encryption support.""" - def __init__(self, prisma_client: object, encryption_key: str | None = None): + def __init__(self, prisma_client: object, encryption_key: str | None = None) -> None: super().__init__(prisma_client) self._encryption_key = encryption_key diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index e651600eb8f..3e037e5bedc 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -9,13 +9,13 @@ "limit": 827 }, "ANN201": { - "limit": 2017 + "limit": 2016 }, "ANN202": { - "limit": 852 + "limit": 851 }, "ANN204": { - "limit": 711 + "limit": 710 }, "ANN205": { "limit": 112 From f44e7ad9fb6d88d2a9f66f4f1b5965bdb7b39c74 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 23 Aug 2026 07:54:07 +0000 Subject: [PATCH 030/281] chore(typing): drop fresh tech debt suppressions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 12 +++---- litellm/llms/custom_httpx/llm_http_handler.py | 6 ++-- litellm/proxy/litellm_pre_call_utils.py | 4 +-- .../openai_files_endpoints/common_utils.py | 9 +++--- litellm/proxy/proxy_server.py | 6 ++-- .../transformation.py | 32 +++++++++---------- ruff-strict-budget.json | 8 ++--- type-discipline-budget.json | 6 ++-- 8 files changed, 41 insertions(+), 42 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 9dcc46ebcee..3f0011c80d2 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 19954 + "limit": 19945 }, "reportArgumentType": { "limit": 2566 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 6049 + "limit": 6048 }, "reportFunctionMemberAccess": { "limit": 7 @@ -57,7 +57,7 @@ "limit": 5663 }, "reportMissingTypeArgument": { - "limit": 15553 + "limit": 15545 }, "reportMissingTypeStubs": { "limit": 40 @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 39007 + "limit": 38998 }, "reportUnknownParameterType": { - "limit": 19883 + "limit": 19876 }, "reportUnknownVariableType": { - "limit": 30568 + "limit": 30554 }, "reportUnnecessaryCast": { "limit": 117 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ed079197513..cdd81b24ca3 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5594,10 +5594,9 @@ class BaseLLMHTTPHandler: kwargs=hook_kwargs, ) except Exception as e: - _call_id = getattr(logging_obj, "litellm_call_id", "unknown") verbose_logger.exception( "LiteLLM.AgenticHookError: Exception in async_should_run_agentic_loop [call_id=%s model=%s]: %s", - _call_id, + logging_obj.litellm_call_id, model, str(e), ) @@ -5619,10 +5618,9 @@ class BaseLLMHTTPHandler: except AgenticLoopSafetyError as e: if not self._can_replace_turn_with_terminal_response(stream, api_surface): raise - _call_id = getattr(logging_obj, "litellm_call_id", "unknown") verbose_logger.warning( "LiteLLM.AgenticLoopRefused: ending turn [call_id=%s model=%s]: %s", - _call_id, + logging_obj.litellm_call_id, model, str(e), ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index cb5002e431b..064b53e07b7 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -4,7 +4,7 @@ import json import re import time from collections import OrderedDict -from collections.abc import Mapping +from collections.abc import Mapping, MutableMapping from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final @@ -1629,7 +1629,7 @@ class LiteLLMProxyRequestSetup: def refresh_proxy_server_request_body_snapshot( - data: dict, # mutable-ok: mutates proxy_server_request.body in place on the shared request dict + data: MutableMapping[str, object], ) -> None: """ Re-snapshot ``data["proxy_server_request"]["body"]`` from the current state of ``data``. diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 142aced4a38..134ed74ae65 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1346,11 +1346,12 @@ def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool: reports no successful request lines. When counts are unknown, stay eligible so the next poller pass revisits it. (#37713) """ - if getattr(response, "output_file_id", None) is not None: + if response.output_file_id is not None: return True - request_counts = getattr(response, "request_counts", None) - completed = getattr(request_counts, "completed", None) - return completed == 0 + request_counts = response.request_counts + if request_counts is None: + return False + return request_counts.completed == 0 async def update_batch_in_database( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7dced4e26b6..e14a64a9ff8 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4084,7 +4084,7 @@ def resolve_complexity_router_plugins( complexity_router_config["classifier_plugin"] = resolved_classifier # rebind-ok: out-param, resolved in place -def validate_deployment_max_agentic_loops(model: Mapping[str, Any]) -> None: +def validate_deployment_max_agentic_loops(model: Mapping[str, object]) -> None: """ Reject a per-deployment `max_agentic_loops` the agentic loop cannot honor. @@ -4094,7 +4094,9 @@ def validate_deployment_max_agentic_loops(model: Mapping[str, Any]) -> None: start. Left unchecked entirely, a `0` used to read as the default ceiling of 3 and a non-integer failed every request to that model instead. """ - litellm_params: Final = model.get("litellm_params") or {} + litellm_params: Final = model.get("litellm_params") + if not isinstance(litellm_params, Mapping): + return if "max_agentic_loops" not in litellm_params: return diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index b8d7b726a28..db12361f70f 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -1272,16 +1272,14 @@ class LiteLLMCompletionResponsesConfig: if isinstance(content, str) and content.strip(): return content if isinstance(content, list): - text_parts: Final[list[str]] = [] # mutable-ok: text accumulator - for block in content: - if not isinstance(block, Mapping): - continue - block_type = block.get("type") - if block_type in ("encrypted_content", "redacted_thinking"): - continue - text = block.get("text") - if isinstance(text, str) and text.strip(): - text_parts.append(text.strip()) + text_parts: Final = tuple( + text.strip() + for block in content + if isinstance(block, Mapping) + and block.get("type") not in ("encrypted_content", "redacted_thinking") + and isinstance(text := block.get("text"), str) + and text.strip() + ) if text_parts: return "\n".join(text_parts) return None @@ -1297,13 +1295,13 @@ class LiteLLMCompletionResponsesConfig: summary: Final[object] = input_item.get("summary") if not isinstance(summary, list): return None - text_parts: Final[list[str]] = [] # mutable-ok: text accumulator - for block in summary: - if not isinstance(block, Mapping): - continue - text = block.get("text") - if isinstance(text, str) and text.strip(): - text_parts.append(text.strip()) + text_parts: Final = tuple( + text.strip() + for block in summary + if isinstance(block, Mapping) + and isinstance(text := block.get("text"), str) + and text.strip() + ) return "\n".join(text_parts) if text_parts else None @staticmethod diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 3e037e5bedc..3585c7a7bd3 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -12,10 +12,10 @@ "limit": 2016 }, "ANN202": { - "limit": 851 + "limit": 850 }, "ANN204": { - "limit": 710 + "limit": 709 }, "ANN205": { "limit": 112 @@ -24,7 +24,7 @@ "limit": 133 }, "ANN401": { - "limit": 1187 + "limit": 1185 }, "ASYNC230": { "limit": 11 @@ -234,7 +234,7 @@ "limit": 5 }, "TID251": { - "limit": 1211 + "limit": 1210 }, "TRY002": { "limit": 524 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 81f7c6aa40b..4f2314b2a0a 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22801 + "limit": 22795 }, "LIT002": { - "limit": 26873 + "limit": 26872 }, "LIT003": { "limit": 269 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16673 + "limit": 16672 }, "LIT011": { "limit": 5588 From cb2f5c664163798bb348b88563c3f41ce8a1e96c Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 23 Aug 2026 08:07:22 +0000 Subject: [PATCH 031/281] style(typing): format reasoning extraction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_completion_transformation/transformation.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index db12361f70f..3cc8db3f357 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -1298,9 +1298,7 @@ class LiteLLMCompletionResponsesConfig: text_parts: Final = tuple( text.strip() for block in summary - if isinstance(block, Mapping) - and isinstance(text := block.get("text"), str) - and text.strip() + if isinstance(block, Mapping) and isinstance(text := block.get("text"), str) and text.strip() ) return "\n".join(text_parts) if text_parts else None From 11061d13c9b22d0933b1a153b627a1420d93b0ef Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 12:22:45 -0400 Subject: [PATCH 032/281] fix(redis): support credential providers across clients --- litellm/_redis.py | 115 +++++--- litellm/caching/redis_cache.py | 20 +- litellm/caching/redis_cluster_cache.py | 17 +- pyproject.toml | 1 + tests/test_litellm/test_redis.py | 391 +++++++++++++++++++++++-- uv.lock | 2 + 6 files changed, 461 insertions(+), 85 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 58f37cf569d..0f9716a5396 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -50,6 +50,7 @@ def _get_redis_kwargs(): include_args: Final = { "url", "redis_connect_func", + "credential_provider", "gcp_service_account", "gcp_ssl_ca_certs", "azure_redis_ad_token", @@ -155,7 +156,8 @@ def _get_redis_cluster_kwargs(client=None): def _get_redis_env_kwarg_mapping(): PREFIX: Final = "REDIS_" - return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs()} + exclude_from_environment: Final = {"credential_provider"} + return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment} def _redis_kwargs_from_environment(): @@ -410,54 +412,58 @@ def _get_redis_client_logic(**env_overrides): if _service_name is not None: redis_kwargs["service_name"] = _service_name - # Handle GCP IAM authentication - _gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT") - _gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS") - - if _gcp_service_account is not None: - verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.") - redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func( - service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs + if redis_kwargs.get("credential_provider") is None: + # Handle GCP IAM authentication + _gcp_service_account: Final = redis_kwargs.get("gcp_service_account") or get_secret_str( + "REDIS_GCP_SERVICE_ACCOUNT" ) - # Store GCP service account in redis_connect_func for async cluster access - redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account + _gcp_ssl_ca_certs: Final = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS") - # Remove GCP-specific kwargs that shouldn't be passed to Redis client - redis_kwargs.pop("gcp_service_account", None) - redis_kwargs.pop("gcp_ssl_ca_certs", None) + if _gcp_service_account is not None: + verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.") + redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func( + service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs + ) + # Store GCP service account in redis_connect_func for async cluster access + redis_kwargs["redis_connect_func"]._gcp_service_account = _gcp_service_account - # Only enable SSL if explicitly requested AND SSL CA certs are provided - if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False): - redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs + # Only enable SSL if explicitly requested AND SSL CA certs are provided + if _gcp_ssl_ca_certs and redis_kwargs.get("ssl", False): + redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs - # Handle Azure AD authentication (after GCP IAM block) - _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") + # Handle Azure AD authentication (after GCP IAM block) + _azure_redis_ad_token: Final = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN") - _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" + _azure_ad_enabled: Final = _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true" - if _azure_ad_enabled and _gcp_service_account is not None: - verbose_logger.warning( - "Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. " - "Using GCP IAM. Remove one to avoid misconfiguration." - ) + if _azure_ad_enabled and _gcp_service_account is not None: + verbose_logger.warning( + "Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. " + "Using GCP IAM. Remove one to avoid misconfiguration." + ) - if _azure_ad_enabled and _gcp_service_account is None: - _azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID") - _azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID") - _azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET") + if _azure_ad_enabled and _gcp_service_account is None: + _azure_client_id: Final = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID") + _azure_tenant_id: Final = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID") + _azure_client_secret: Final = redis_kwargs.get("azure_client_secret") or get_secret_str( + "AZURE_CLIENT_SECRET" + ) - verbose_logger.debug("Setting up Azure AD authentication for Redis.") - redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func( - azure_client_id=_azure_client_id, - azure_tenant_id=_azure_tenant_id, - azure_client_secret=_azure_client_secret, - ) - # Marker for async paths to detect Azure AD auth. The live credential - # object is attached separately as `_azure_credential` by - # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret - # are intentionally NOT exposed on the function to avoid leaking - # credentials via inspection or logging. - redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True + verbose_logger.debug("Setting up Azure AD authentication for Redis.") + redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func( + azure_client_id=_azure_client_id, + azure_tenant_id=_azure_tenant_id, + azure_client_secret=_azure_client_secret, + ) + # Marker for async paths to detect Azure AD auth. The live credential + # object is attached separately as `_azure_credential` by + # `create_azure_ad_redis_connect_func`; the raw client_id/tenant_id/secret + # are intentionally NOT exposed on the function to avoid leaking + # credentials via inspection or logging. + redis_kwargs["redis_connect_func"]._azure_redis_ad_token = True + + redis_kwargs.pop("gcp_service_account", None) + redis_kwargs.pop("gcp_ssl_ca_certs", None) # Always remove Azure-specific kwargs that shouldn't be passed to Redis client redis_kwargs.pop("azure_redis_ad_token", None) @@ -465,6 +471,16 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("azure_tenant_id", None) redis_kwargs.pop("azure_client_secret", None) + if redis_kwargs.get("credential_provider") is not None: + redis_kwargs.pop("redis_connect_func", None) + redis_kwargs.pop("username", None) + redis_kwargs.pop("password", None) + if redis_kwargs.get("url") is not None: + from urllib.parse import urlsplit, urlunsplit + + parsed_url = urlsplit(redis_kwargs["url"]) + redis_kwargs["url"] = urlunsplit(parsed_url._replace(netloc=parsed_url.netloc.rsplit("@", 1)[-1])) + if "url" in redis_kwargs and redis_kwargs["url"] is not None: # Only strip host/port/db/password when not routing to a cluster. # When startup_nodes is also present the cluster path takes priority and @@ -474,6 +490,8 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("port", None) redis_kwargs.pop("db", None) redis_kwargs.pop("password", None) + if redis_kwargs.get("credential_provider") is not None: + redis_kwargs.pop("username", None) elif ( "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None @@ -532,8 +550,7 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis: service_name: Final = redis_kwargs.get("service_name") connection_kwargs: Final = _get_redis_sentinel_connection_kwargs(redis_kwargs) connection_kwargs.setdefault("socket_timeout", REDIS_SOCKET_TIMEOUT) - sentinel_kwargs: Final = dict(connection_kwargs) - sentinel_kwargs["password"] = sentinel_password + sentinel_kwargs: Final = _sentinel_auth_kwargs(connection_kwargs, sentinel_password) if not sentinel_nodes or not service_name: raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.") @@ -605,13 +622,19 @@ def _async_credential_provider(redis_connect_func: object | None) -> CredentialP def _async_auth_kwargs(redis_kwargs: dict) -> dict: """Swaps a connect func an async path cannot run for the equivalent credential provider, which supersedes any static username or password redis-py would otherwise reject it with.""" + explicit_provider: Final = redis_kwargs.get("credential_provider") + if explicit_provider is not None: + superseded: Final = frozenset({"redis_connect_func", "username", "password"}) + kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) + return dict(kept) + credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func")) if credential_provider is None: return redis_kwargs - superseded: Final = frozenset({"redis_connect_func", "username", "password"}) - kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) - return dict(kept, credential_provider=credential_provider) # mutable-ok: the branches below mutate these kwargs + automatic_superseded: Final = frozenset({"redis_connect_func", "username", "password"}) + automatic_kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in automatic_superseded) + return dict(automatic_kept, credential_provider=credential_provider) def get_redis_client(**env_overrides): diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 934ba500ef9..991b6c8c6c5 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -401,9 +401,21 @@ class RedisCache(BaseCache): """ # Create a stable representation of redis_kwargs for hashing # Sort keys to ensure consistent hash regardless of parameter order - sorted_kwargs: Final = sorted(self.redis_kwargs.items()) + redis_kwargs: Final[dict[str, object]] = self.redis_kwargs + provider: Final = redis_kwargs.get("credential_provider") + redis_connect_func: Final = redis_kwargs.get("redis_connect_func") + sorted_kwargs: Final = sorted( + item for item in redis_kwargs.items() if item[0] not in {"credential_provider", "redis_connect_func"} + ) kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True) - kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16] + identity_suffix: Final = ( + "" + if provider is None and redis_connect_func is None + else f":provider-{id(provider)}" + if provider is not None + else f":connect-func-{id(redis_connect_func)}" + ) + kwargs_hash: Final = hashlib.sha256(f"{kwargs_str}{identity_suffix}".encode()).hexdigest()[:16] return f"async-redis-client-{kwargs_hash}" def init_async_client( @@ -1384,10 +1396,10 @@ class RedisCache(BaseCache): dict: {"status": "success" | "failed", "message": str, "error": Optional[str]} """ try: - import redis.asyncio as redis_async + from .._redis import get_redis_async_client # Create a fresh Redis client with current settings - redis_client: Final = redis_async.Redis(**self.redis_kwargs) + redis_client: Final = get_redis_async_client(**self.redis_kwargs) # Test the connection ping_result: Final = await redis_client.ping() diff --git a/litellm/caching/redis_cluster_cache.py b/litellm/caching/redis_cluster_cache.py index b6dd8047fd4..12d285ca5a8 100644 --- a/litellm/caching/redis_cluster_cache.py +++ b/litellm/caching/redis_cluster_cache.py @@ -64,22 +64,9 @@ class RedisClusterCache(RedisCache): dict: {"status": "success" | "failed", "message": str, "error": Optional[str]} """ try: - import redis.asyncio as redis_async - from redis.cluster import ClusterNode + from .._redis import get_redis_async_client - # Create ClusterNode objects from startup_nodes - cluster_kwargs: Final = self.redis_kwargs.copy() - startup_nodes: Final = cluster_kwargs.pop("startup_nodes", []) - - new_startup_nodes: Final[list[ClusterNode]] = [] - for item in startup_nodes: - new_startup_nodes.append(ClusterNode(**item)) - - # Create a fresh Redis Cluster client with current settings - redis_client: Final = redis_async.RedisCluster( - startup_nodes=new_startup_nodes, - **cluster_kwargs, - ) + redis_client: Final = get_redis_async_client(**self.redis_kwargs) # Test the connection ping_result: Final = await redis_client.ping() diff --git a/pyproject.toml b/pyproject.toml index fca5c7da1e2..c80ba143512 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,6 +52,7 @@ proxy = [ "backoff>=2.2.1,<3.0", "pyyaml>=6.0.3,<7.0", "rq>=2.7.0,<3.0", + "redis>=5.3.1,<6.0", "orjson>=3.11.6,<4.0", # redis-py's C response parser. It arrives with redis (via rq) either way; naming # it here is what makes redis-py select _HiredisParser instead of the Python one. diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 3762181f5c3..38b2bd5296a 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -1,13 +1,19 @@ import json from types import SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest import redis import redis.asyncio as async_redis +from redis.credentials import CredentialProvider +import litellm from litellm._redis import ( + _get_redis_client_logic, _get_redis_cluster_kwargs, + _get_redis_env_kwarg_mapping, + _get_redis_kwargs, + _get_redis_url_kwargs, get_redis_async_client, get_redis_client, get_redis_connection_pool, @@ -18,9 +24,69 @@ from litellm._redis_credential_provider import ( GCPIAMCredentialProvider, _token_cache, ) +from litellm.caching.redis_cache import RedisCache +from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.constants import REDIS_CLUSTER_HEALTH_CHECK_INTERVAL +class _StubCredentialProvider(CredentialProvider): + def __init__(self, token: str = "stub-token") -> None: + self._token = token + + def get_credentials(self): + return (self._token,) + + async def get_credentials_async(self): + return (self._token,) + + +class _HostileCredentialProvider(CredentialProvider): + def __init__(self, secret: str) -> None: + self._payload = secret + + def get_credentials(self): + return (self._payload,) + + async def get_credentials_async(self): + return (self._payload,) + + def __repr__(self): + raise AssertionError("provider repr must never be invoked") + + def __str__(self): + raise AssertionError("provider str must never be invoked") + + def __reduce__(self): + raise AssertionError("provider must never be serialized") + + def __getstate__(self): + raise AssertionError("provider state must never be inspected") + + +def _gcp_marker_callback() -> MagicMock: + callback = MagicMock() + callback._gcp_service_account = "projects/-/serviceAccounts/sa@project.iam.gserviceaccount.com" + return callback + + +@pytest.fixture +def clean_redis_environment(monkeypatch): + for var in ( + "REDIS_URL", + "REDIS_CLUSTER_NODES", + "REDIS_SENTINEL_NODES", + *_get_redis_env_kwarg_mapping(), + ): + monkeypatch.delenv(var, raising=False) + + +@pytest.fixture +def clear_llm_client_cache(): + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + @pytest.fixture(autouse=True) def clear_gcp_iam_token_cache(): """Reset the module-level GCP IAM token cache between tests.""" @@ -29,6 +95,289 @@ def clear_gcp_iam_token_cache(): _token_cache.clear() +def test_redis_uses_the_hiredis_response_parser(): + """The proxy extra must keep redis-py's C response parser available.""" + from redis._parsers import _HiredisParser + from redis.connection import HIREDIS_AVAILABLE, DefaultParser + + if not HIREDIS_AVAILABLE: + pytest.skip("hiredis is not installed in this test environment") + + assert DefaultParser is _HiredisParser + + client = get_redis_client(host="redis-host", port=6379) + connection = client.connection_pool.make_connection() + assert isinstance(connection._parser, _HiredisParser) + + +def test_redis_allowlists_include_credential_provider(): + assert "credential_provider" in _get_redis_kwargs() + assert "credential_provider" in _get_redis_url_kwargs() + assert "credential_provider" in _get_redis_cluster_kwargs() + + +def test_credential_provider_is_not_environment_derived(): + mapping = _get_redis_env_kwarg_mapping() + assert "REDIS_CREDENTIAL_PROVIDER" not in mapping + assert "credential_provider" not in mapping.values() + + +def test_sync_direct_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_client(host="redis-host", port=6379, credential_provider=provider) + + assert client.connection_pool.connection_kwargs["credential_provider"] is provider + + +def test_sync_direct_provider_supersedes_static_credentials(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_client( + host="redis-host", + port=6379, + username="redis-user", + password="redis-password", + credential_provider=provider, + ) + connection = client.connection_pool.make_connection() + + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + + +def test_sync_direct_provider_supersedes_environment_credentials(clean_redis_environment, monkeypatch): + provider = _StubCredentialProvider() + monkeypatch.setenv("REDIS_USERNAME", "redis-user") + monkeypatch.setenv("REDIS_PASSWORD", "redis-password") + + client = get_redis_client(host="redis-host", port=6379, credential_provider=provider) + connection = client.connection_pool.make_connection() + + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + + +def test_sync_url_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_client(url="redis://redis-host:6379", credential_provider=provider) + + assert client.connection_pool.connection_kwargs["credential_provider"] is provider + + +def test_async_direct_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_async_client(host="redis-host", port=6379, credential_provider=provider) + + assert client.connection_pool.connection_kwargs["credential_provider"] is provider + + +def test_async_url_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_async_client(url="redis://redis-host:6379", credential_provider=provider) + + assert client.connection_pool.connection_kwargs["credential_provider"] is provider + + +def test_sync_url_credentials_do_not_replace_explicit_provider(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_client( + url="redis://url-user:url-pass@redis-host:6379", + credential_provider=provider, + ) + connection = client.connection_pool.make_connection() + + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + + +def test_async_url_credentials_do_not_replace_explicit_provider(clean_redis_environment): + provider = _StubCredentialProvider() + + client = get_redis_async_client( + url="redis://url-user:url-pass@redis-host:6379", + credential_provider=provider, + ) + connection = client.connection_pool.make_connection() + + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + + +def test_async_host_port_pool_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + pool = get_redis_connection_pool(host="redis-host", port=6379, credential_provider=provider) + + assert pool is not None + assert pool.connection_kwargs["credential_provider"] is provider + + +def test_async_url_pool_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + + pool = get_redis_connection_pool(url="redis://redis-host:6379", credential_provider=provider) + + assert pool is not None + assert pool.connection_kwargs["credential_provider"] is provider + + +def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + startup_nodes = [{"host": "cluster-node", "port": 6379}] + + with patch("litellm._redis.redis.RedisCluster", autospec=True) as mock_cluster_cls: + get_redis_client(startup_nodes=startup_nodes, credential_provider=provider) + + mock_cluster_cls.assert_called_once() + assert mock_cluster_cls.call_args[1].get("credential_provider") is provider + + +def test_async_cluster_preserves_credential_provider_identity(clean_redis_environment): + provider = _StubCredentialProvider() + startup_nodes = [{"host": "cluster-node", "port": 6379}] + + with patch("litellm.caching.redis_cluster_node_isolation.get_litellm_async_redis_cluster_class") as mock_class: + get_redis_async_client(startup_nodes=startup_nodes, credential_provider=provider) + + call_kwargs = mock_class.return_value.call_args[1] + assert call_kwargs.get("credential_provider") is provider + + +def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environment, monkeypatch): + provider = _StubCredentialProvider() + monkeypatch.setenv("REDIS_GCP_SERVICE_ACCOUNT", "service-account@example.com") + monkeypatch.setenv("REDIS_AZURE_AD_TOKEN", "true") + + with ( + patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp, + patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure, + ): + redis_kwargs = _get_redis_client_logic( + host="redis-host", + port=6379, + credential_provider=provider, + redis_connect_func=_gcp_marker_callback(), + ) + + mock_gcp.assert_not_called() + mock_azure.assert_not_called() + assert redis_kwargs["credential_provider"] is provider + assert "redis_connect_func" not in redis_kwargs + + +def test_async_direct_explicit_provider_is_preserved_when_normalization_is_bypassed(): + provider = _StubCredentialProvider() + redis_kwargs = { + "host": "redis-host", + "port": 6379, + "credential_provider": provider, + "redis_connect_func": _gcp_marker_callback(), + } + + with ( + patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), + patch("litellm._redis.async_redis.Redis", autospec=True) as mock_redis, + ): + get_redis_async_client() + + call_kwargs = mock_redis.call_args[1] + assert call_kwargs["credential_provider"] is provider + assert "redis_connect_func" not in call_kwargs + + +def test_async_pool_explicit_provider_is_preserved_when_normalization_is_bypassed(): + provider = _StubCredentialProvider() + redis_kwargs = { + "host": "redis-host", + "port": 6379, + "credential_provider": provider, + "redis_connect_func": _gcp_marker_callback(), + } + + with ( + patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), + patch("litellm._redis.async_redis.BlockingConnectionPool", autospec=True) as mock_pool, + ): + get_redis_connection_pool() + + call_kwargs = mock_pool.call_args[1] + assert call_kwargs["credential_provider"] is provider + assert "redis_connect_func" not in call_kwargs + + +@pytest.mark.asyncio +async def test_redis_cache_test_connection_uses_shared_factory(clean_redis_environment): + provider = _StubCredentialProvider() + client = MagicMock(spec=async_redis.Redis) + client.ping = AsyncMock(return_value=True) + client.aclose = AsyncMock() + + with patch("litellm._redis.get_redis_async_client", return_value=client) as mock_factory: + cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) + result = await cache.test_connection() + + assert result["status"] == "success" + call_kwargs = mock_factory.call_args.kwargs + assert call_kwargs["credential_provider"] is provider + + +@pytest.mark.asyncio +async def test_redis_cluster_cache_test_connection_uses_shared_factory(clean_redis_environment): + provider = _StubCredentialProvider() + client = MagicMock(spec=async_redis.RedisCluster) + client.ping = AsyncMock(return_value=True) + client.aclose = AsyncMock() + + with patch("litellm._redis.get_redis_async_client", return_value=client) as mock_factory: + with patch("litellm._redis.get_redis_client", return_value=MagicMock(spec=redis.RedisCluster)): + cache = RedisClusterCache( + startup_nodes=[{"host": "redis-host", "port": 6379}], credential_provider=provider + ) + result = await cache.test_connection() + + assert result["status"] == "success" + call_kwargs = mock_factory.call_args.kwargs + assert call_kwargs["credential_provider"] is provider + + +def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache): + provider = _HostileCredentialProvider("synthetic-secret") + second_provider = _StubCredentialProvider("another-token") + sync_client = MagicMock(spec=redis.Redis) + async_pool = MagicMock(spec=async_redis.BlockingConnectionPool) + + with ( + patch("litellm._redis.get_redis_client", return_value=sync_client), + patch("litellm._redis.get_redis_connection_pool", return_value=async_pool), + ): + cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) + second_cache = RedisCache(host="redis-host", port=6379, credential_provider=second_provider) + + first_key = cache._get_async_client_cache_key() + assert first_key == cache._get_async_client_cache_key() + assert first_key != second_cache._get_async_client_cache_key() + + +def test_redis_cache_key_does_not_serialize_connect_func(): + def connect(connection): + return None + + cache = RedisCache.__new__(RedisCache) + cache.redis_kwargs = {"host": "redis-host", "port": 6379, "redis_connect_func": connect} + + first_key = cache._get_async_client_cache_key() + assert first_key == cache._get_async_client_cache_key() + + def test_get_redis_url_from_environment_single_url(monkeypatch): """Test when REDIS_URL is directly provided""" # Set the environment variable @@ -500,6 +849,27 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_ ) +@patch("litellm._redis.redis.Sentinel") +def test_sync_sentinel_keeps_provider_off_monitors_and_on_master(mock_sentinel_cls): + provider = _StubCredentialProvider() + mock_sentinel = MagicMock() + mock_sentinel_cls.return_value = mock_sentinel + + get_redis_client( + sentinel_nodes=[("sentinel-1", 26379)], + sentinel_password="sentinel-secret", + service_name="mymaster", + password="redis-secret", + credential_provider=provider, + ) + + sentinel_kwargs = mock_sentinel_cls.call_args.kwargs["sentinel_kwargs"] + assert sentinel_kwargs["password"] == "sentinel-secret" + assert "credential_provider" not in sentinel_kwargs + assert mock_sentinel.master_for.call_args.kwargs["credential_provider"] is provider + assert "password" not in mock_sentinel.master_for.call_args.kwargs + + @patch("litellm._redis.async_redis.Sentinel") def test_async_sentinel_uses_sentinel_password_and_master_password( mock_sentinel_cls, @@ -814,25 +1184,6 @@ def test_url_config_drops_kwargs_the_connection_cannot_accept(client_only_kwarg, assert pool.connection_kwargs.get("socket_timeout") == 5.0 -def test_redis_uses_the_hiredis_response_parser(): - """The C parser must be the one redis-py actually picks. - - hiredis is declared in the `proxy` extra purely for speed; nothing imports it, so - dropping it from pyproject.toml would silently fall back to the pure-Python parser - with no other symptom. redis-py selects it at import time, so asserting on the - selection is what catches that. - """ - from redis._parsers import _HiredisParser - from redis.connection import HIREDIS_AVAILABLE, DefaultParser - - assert HIREDIS_AVAILABLE, "hiredis is not installed; redis-py fell back to the pure-Python parser" - assert DefaultParser is _HiredisParser, f"redis-py selected {DefaultParser.__name__}, expected _HiredisParser" - - client = get_redis_client(host="redis-host", port=6379) - connection = client.connection_pool.make_connection() - assert isinstance(connection._parser, _HiredisParser) - - def test_init_arg_names_sees_through_decorated_inits(): """redis-py >= 7.4 wraps AbstractConnection.__init__ with @deprecated_args, whose wrapper is declared (self, *args, **kwargs). Introspecting the wrapper directly diff --git a/uv.lock b/uv.lock index 628483c0117..f8af2f60b11 100644 --- a/uv.lock +++ b/uv.lock @@ -4345,6 +4345,7 @@ proxy = [ { name = "restrictedpython" }, { name = "rich" }, { name = "rq" }, + { name = "redis" }, { name = "soundfile" }, { name = "starlette" }, { name = "uvicorn" }, @@ -4557,6 +4558,7 @@ requires-dist = [ { name = "restrictedpython", marker = "extra == 'proxy'", specifier = ">=8.1,<9.0" }, { name = "rich", marker = "extra == 'cli'", specifier = ">=13.9.4,<14.0" }, { name = "rich", marker = "extra == 'proxy'", specifier = ">=13.9.4,<14.0" }, + { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "rq", marker = "extra == 'proxy'", specifier = ">=2.7.0,<3.0" }, { name = "semantic-router", marker = "python_full_version < '3.14' and extra == 'semantic-router'", specifier = ">=0.1.15,<1.0" }, { name = "sentry-sdk", marker = "extra == 'proxy-runtime'", specifier = ">=2.21.0,<3.0" }, From a4be6a9a6fdbd85b5dbc369abb0f2247be89cfc2 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 16:36:23 -0400 Subject: [PATCH 033/281] fix(redis): address credential provider review findings Generated with AI Co-Authored-By: Claude Code --- litellm/_redis.py | 35 ++++++----- litellm/caching/redis_cache.py | 26 +++----- pyproject.toml | 4 +- tests/test_litellm/test_redis.py | 103 ++++++++++++++++++++++++++----- uv.lock | 4 +- 5 files changed, 121 insertions(+), 51 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 0f9716a5396..6ff4c292c47 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -14,6 +14,7 @@ import json import os from collections.abc import Callable from typing import Final +from urllib.parse import urlsplit, urlunsplit import redis import redis.asyncio as async_redis @@ -156,7 +157,7 @@ def _get_redis_cluster_kwargs(client=None): def _get_redis_env_kwarg_mapping(): PREFIX: Final = "REDIS_" - exclude_from_environment: Final = {"credential_provider"} + exclude_from_environment: Final = frozenset({"credential_provider"}) return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment} @@ -355,6 +356,14 @@ def get_redis_url_from_environment(): return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" +def _url_without_userinfo(url: str) -> str: + """redis-py rejects a url that carries its own username or password next to a credential + provider, so the provider's credentials replace whatever userinfo the url was configured with.""" + parts: Final = urlsplit(url) + netloc: Final = parts.netloc.rsplit("@", 1)[-1] + return urlunsplit((parts.scheme, netloc, parts.path, parts.query, parts.fragment)) + + def _get_redis_client_logic(**env_overrides): """ Common functionality across sync + async redis client implementations @@ -476,10 +485,7 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("username", None) redis_kwargs.pop("password", None) if redis_kwargs.get("url") is not None: - from urllib.parse import urlsplit, urlunsplit - - parsed_url = urlsplit(redis_kwargs["url"]) - redis_kwargs["url"] = urlunsplit(parsed_url._replace(netloc=parsed_url.netloc.rsplit("@", 1)[-1])) + redis_kwargs["url"] = _url_without_userinfo(redis_kwargs["url"]) if "url" in redis_kwargs and redis_kwargs["url"] is not None: # Only strip host/port/db/password when not routing to a cluster. @@ -490,8 +496,6 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("port", None) redis_kwargs.pop("db", None) redis_kwargs.pop("password", None) - if redis_kwargs.get("credential_provider") is not None: - redis_kwargs.pop("username", None) elif ( "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None @@ -623,18 +627,17 @@ def _async_auth_kwargs(redis_kwargs: dict) -> dict: """Swaps a connect func an async path cannot run for the equivalent credential provider, which supersedes any static username or password redis-py would otherwise reject it with.""" explicit_provider: Final = redis_kwargs.get("credential_provider") - if explicit_provider is not None: - superseded: Final = frozenset({"redis_connect_func", "username", "password"}) - kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) - return dict(kept) - - credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func")) + credential_provider: Final = ( + explicit_provider + if explicit_provider is not None + else _async_credential_provider(redis_kwargs.get("redis_connect_func")) + ) if credential_provider is None: return redis_kwargs - automatic_superseded: Final = frozenset({"redis_connect_func", "username", "password"}) - automatic_kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in automatic_superseded) - return dict(automatic_kept, credential_provider=credential_provider) + superseded: Final = frozenset({"redis_connect_func", "username", "password"}) + kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) + return dict(kept, credential_provider=credential_provider) # mutable-ok: the branches below mutate these kwargs def get_redis_client(**env_overrides): diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 991b6c8c6c5..0207a571dd6 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -175,6 +175,10 @@ _RedisCallResult = TypeVar("_RedisCallResult") _swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0) +def _opaque_kwarg_key(value: object) -> str: + return f"{type(value).__name__}-{id(value)}" + + @functools.lru_cache(maxsize=1) def _redis_health_error_types() -> tuple[type, ...]: """Exception types that mean the Redis backend itself is unhealthy. @@ -398,24 +402,14 @@ class RedisCache(BaseCache): """ Generate a cache key for the async Redis client based on connection parameters. This ensures different Redis configurations use different cached clients. + + Kwargs the caller hands over as live objects (a credential provider, a connect func) are not + JSON-serializable and carry no stable value identity, so they key on instance identity. """ - # Create a stable representation of redis_kwargs for hashing # Sort keys to ensure consistent hash regardless of parameter order - redis_kwargs: Final[dict[str, object]] = self.redis_kwargs - provider: Final = redis_kwargs.get("credential_provider") - redis_connect_func: Final = redis_kwargs.get("redis_connect_func") - sorted_kwargs: Final = sorted( - item for item in redis_kwargs.items() if item[0] not in {"credential_provider", "redis_connect_func"} - ) - kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True) - identity_suffix: Final = ( - "" - if provider is None and redis_connect_func is None - else f":provider-{id(provider)}" - if provider is not None - else f":connect-func-{id(redis_connect_func)}" - ) - kwargs_hash: Final = hashlib.sha256(f"{kwargs_str}{identity_suffix}".encode()).hexdigest()[:16] + sorted_kwargs: Final = sorted(self.redis_kwargs.items()) + kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True, default=_opaque_kwarg_key) + kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16] return f"async-redis-client-{kwargs_hash}" def init_async_client( diff --git a/pyproject.toml b/pyproject.toml index c80ba143512..b9c514dd88b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,8 +54,8 @@ proxy = [ "rq>=2.7.0,<3.0", "redis>=5.3.1,<6.0", "orjson>=3.11.6,<4.0", - # redis-py's C response parser. It arrives with redis (via rq) either way; naming - # it here is what makes redis-py select _HiredisParser instead of the Python one. + # redis-py's C response parser. Nothing imports it; naming it here is what makes + # redis-py select _HiredisParser instead of the pure-Python one. "hiredis>=3.0.0,<4.0", "apscheduler>=3.11.2,<4.0", "fastapi-sso>=0.19.0,<1.0", diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 38b2bd5296a..1476650ac3f 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -95,21 +95,6 @@ def clear_gcp_iam_token_cache(): _token_cache.clear() -def test_redis_uses_the_hiredis_response_parser(): - """The proxy extra must keep redis-py's C response parser available.""" - from redis._parsers import _HiredisParser - from redis.connection import HIREDIS_AVAILABLE, DefaultParser - - if not HIREDIS_AVAILABLE: - pytest.skip("hiredis is not installed in this test environment") - - assert DefaultParser is _HiredisParser - - client = get_redis_client(host="redis-host", port=6379) - connection = client.connection_pool.make_connection() - assert isinstance(connection._parser, _HiredisParser) - - def test_redis_allowlists_include_credential_provider(): assert "credential_provider" in _get_redis_kwargs() assert "credential_provider" in _get_redis_url_kwargs() @@ -230,6 +215,19 @@ def test_async_url_pool_preserves_credential_provider_identity(clean_redis_envir assert pool.connection_kwargs["credential_provider"] is provider +def test_async_url_pool_strips_userinfo_for_the_provider(clean_redis_environment): + """The url allowlist has to carry the provider through, and redis-py rejects it next to userinfo.""" + provider = _StubCredentialProvider() + + pool = get_redis_connection_pool(url="rediss://url-user:url-pass@redis-host:6379/3", credential_provider=provider) + + connection = pool.make_connection() + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + assert connection.db == 3 + + def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] @@ -274,6 +272,47 @@ def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environ assert "redis_connect_func" not in redis_kwargs +@pytest.mark.parametrize( + "overrides", + [ + {"gcp_ssl_ca_certs": "/tmp/ca.pem"}, + {"gcp_service_account": "sa@example.com", "gcp_ssl_ca_certs": "/tmp/ca.pem"}, + ], + ids=["certs-without-service-account", "both-alongside-a-provider"], +) +def test_gcp_kwargs_never_survive_client_logic(clean_redis_environment, overrides): + """redis.Redis has no gcp_* parameters, so anything left behind raises TypeError on connect.""" + redis_kwargs = _get_redis_client_logic( + host="redis-host", + port=6379, + credential_provider=_StubCredentialProvider() if "gcp_service_account" in overrides else None, + **overrides, + ) + + assert "gcp_service_account" not in redis_kwargs + assert "gcp_ssl_ca_certs" not in redis_kwargs + + +def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): + """Stripping the userinfo must not take the database path, query, or scheme with it.""" + provider = _StubCredentialProvider() + + redis_kwargs = _get_redis_client_logic( + url="rediss://url-user:url-pass@redis-host:6379/3?protocol=3", + credential_provider=provider, + ) + + assert redis_kwargs["url"] == "rediss://redis-host:6379/3?protocol=3" + + +def test_provider_free_url_is_left_untouched(clean_redis_environment): + url = "redis://url-user:url-pass@redis-host:6379/3" + + redis_kwargs = _get_redis_client_logic(url=url) + + assert redis_kwargs["url"] == url + + def test_async_direct_explicit_provider_is_preserved_when_normalization_is_bypassed(): provider = _StubCredentialProvider() redis_kwargs = { @@ -378,6 +417,21 @@ def test_redis_cache_key_does_not_serialize_connect_func(): assert first_key == cache._get_async_client_cache_key() +def test_redis_cache_key_keys_opaque_kwargs_by_identity(): + """Any object a caller passes through must key by identity rather than crash the JSON dump.""" + + class _Opaque: + pass + + first = RedisCache.__new__(RedisCache) + first.redis_kwargs = {"host": "redis-host", "retry": _Opaque()} + second = RedisCache.__new__(RedisCache) + second.redis_kwargs = {"host": "redis-host", "retry": _Opaque()} + + assert first._get_async_client_cache_key() == first._get_async_client_cache_key() + assert first._get_async_client_cache_key() != second._get_async_client_cache_key() + + def test_get_redis_url_from_environment_single_url(monkeypatch): """Test when REDIS_URL is directly provided""" # Set the environment variable @@ -1184,6 +1238,25 @@ def test_url_config_drops_kwargs_the_connection_cannot_accept(client_only_kwarg, assert pool.connection_kwargs.get("socket_timeout") == 5.0 +def test_redis_uses_the_hiredis_response_parser(): + """The C parser must be the one redis-py actually picks. + + hiredis is declared in the `proxy` extra purely for speed; nothing imports it, so + dropping it from pyproject.toml would silently fall back to the pure-Python parser + with no other symptom. redis-py selects it at import time, so asserting on the + selection is what catches that. + """ + from redis._parsers import _HiredisParser + from redis.connection import HIREDIS_AVAILABLE, DefaultParser + + assert HIREDIS_AVAILABLE, "hiredis is not installed; redis-py fell back to the pure-Python parser" + assert DefaultParser is _HiredisParser, f"redis-py selected {DefaultParser.__name__}, expected _HiredisParser" + + client = get_redis_client(host="redis-host", port=6379) + connection = client.connection_pool.make_connection() + assert isinstance(connection._parser, _HiredisParser) + + def test_init_arg_names_sees_through_decorated_inits(): """redis-py >= 7.4 wraps AbstractConnection.__init__ with @deprecated_args, whose wrapper is declared (self, *args, **kwargs). Introspecting the wrapper directly diff --git a/uv.lock b/uv.lock index f8af2f60b11..fae22e3759e 100644 --- a/uv.lock +++ b/uv.lock @@ -4342,10 +4342,10 @@ proxy = [ { name = "pyroscope-io", marker = "sys_platform != 'win32'" }, { name = "python-multipart" }, { name = "pyyaml" }, + { name = "redis" }, { name = "restrictedpython" }, { name = "rich" }, { name = "rq" }, - { name = "redis" }, { name = "soundfile" }, { name = "starlette" }, { name = "uvicorn" }, @@ -4552,13 +4552,13 @@ requires-dist = [ { name = "python3-saml", marker = "extra == 'saml'", specifier = ">=1.16.0,<2.0" }, { name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" }, { name = "pyyaml", marker = "extra == 'proxy'", specifier = ">=6.0.3,<7.0" }, + { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "redisvl", marker = "extra == 'extra-proxy'", specifier = ">=0.4.1,<1.0" }, { name = "requests", marker = "extra == 'cli'", specifier = ">=2.32.0,<3.0" }, { name = "resend", marker = "extra == 'extra-proxy'", specifier = ">=2.23.0,<3.0" }, { name = "restrictedpython", marker = "extra == 'proxy'", specifier = ">=8.1,<9.0" }, { name = "rich", marker = "extra == 'cli'", specifier = ">=13.9.4,<14.0" }, { name = "rich", marker = "extra == 'proxy'", specifier = ">=13.9.4,<14.0" }, - { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "rq", marker = "extra == 'proxy'", specifier = ">=2.7.0,<3.0" }, { name = "semantic-router", marker = "python_full_version < '3.14' and extra == 'semantic-router'", specifier = ">=0.1.15,<1.0" }, { name = "sentry-sdk", marker = "extra == 'proxy-runtime'", specifier = ">=2.21.0,<3.0" }, From 0ebebeaab9702775c7cdf2f1f71f46e384c2815d Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 17:41:37 -0400 Subject: [PATCH 034/281] fix(redis): drop direct dependency and suppress test-quality violations --- pyproject.toml | 5 ++- tests/test_litellm/test_redis.py | 58 +++++++++++++++++++++++--------- uv.lock | 2 -- 3 files changed, 45 insertions(+), 20 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index b9c514dd88b..fca5c7da1e2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,10 +52,9 @@ proxy = [ "backoff>=2.2.1,<3.0", "pyyaml>=6.0.3,<7.0", "rq>=2.7.0,<3.0", - "redis>=5.3.1,<6.0", "orjson>=3.11.6,<4.0", - # redis-py's C response parser. Nothing imports it; naming it here is what makes - # redis-py select _HiredisParser instead of the pure-Python one. + # redis-py's C response parser. It arrives with redis (via rq) either way; naming + # it here is what makes redis-py select _HiredisParser instead of the Python one. "hiredis>=3.0.0,<4.0", "apscheduler>=3.11.2,<4.0", "fastapi-sso>=0.19.0,<1.0", diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 1476650ac3f..449f2d8dc14 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -228,11 +228,15 @@ def test_async_url_pool_strips_userinfo_for_the_provider(clean_redis_environment assert connection.db == 3 -def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment): +def test_sync_cluster_preserves_credential_provider_identity( # test-quality-ok: constructor kwargs are the only seam + clean_redis_environment, +): provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] - with patch("litellm._redis.redis.RedisCluster", autospec=True) as mock_cluster_cls: + with patch( # test-quality-ok: RedisCluster slot-discovers in its constructor + "litellm._redis.redis.RedisCluster", autospec=True + ) as mock_cluster_cls: get_redis_client(startup_nodes=startup_nodes, credential_provider=provider) mock_cluster_cls.assert_called_once() @@ -243,7 +247,9 @@ def test_async_cluster_preserves_credential_provider_identity(clean_redis_enviro provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] - with patch("litellm.caching.redis_cluster_node_isolation.get_litellm_async_redis_cluster_class") as mock_class: + with patch( # test-quality-ok: the async cluster class is the only injection point here + "litellm.caching.redis_cluster_node_isolation.get_litellm_async_redis_cluster_class" + ) as mock_class: get_redis_async_client(startup_nodes=startup_nodes, credential_provider=provider) call_kwargs = mock_class.return_value.call_args[1] @@ -256,8 +262,12 @@ def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environ monkeypatch.setenv("REDIS_AZURE_AD_TOKEN", "true") with ( - patch("litellm._redis.create_gcp_iam_redis_connect_func") as mock_gcp, - patch("litellm._redis.create_azure_ad_redis_connect_func") as mock_azure, + patch( # test-quality-ok: the assertion is that this builder is never reached + "litellm._redis.create_gcp_iam_redis_connect_func" + ) as mock_gcp, + patch( # test-quality-ok: the assertion is that this builder is never reached + "litellm._redis.create_azure_ad_redis_connect_func" + ) as mock_azure, ): redis_kwargs = _get_redis_client_logic( host="redis-host", @@ -323,8 +333,12 @@ def test_async_direct_explicit_provider_is_preserved_when_normalization_is_bypas } with ( - patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), - patch("litellm._redis.async_redis.Redis", autospec=True) as mock_redis, + patch( # test-quality-ok: bypassing normalization is what this test pins + "litellm._redis._get_redis_client_logic", return_value=redis_kwargs + ), + patch( # test-quality-ok: the constructor kwargs are the only observable + "litellm._redis.async_redis.Redis", autospec=True + ) as mock_redis, ): get_redis_async_client() @@ -343,8 +357,12 @@ def test_async_pool_explicit_provider_is_preserved_when_normalization_is_bypasse } with ( - patch("litellm._redis._get_redis_client_logic", return_value=redis_kwargs), - patch("litellm._redis.async_redis.BlockingConnectionPool", autospec=True) as mock_pool, + patch( # test-quality-ok: bypassing normalization is what this test pins + "litellm._redis._get_redis_client_logic", return_value=redis_kwargs + ), + patch( # test-quality-ok: the constructor kwargs are the only observable + "litellm._redis.async_redis.BlockingConnectionPool", autospec=True + ) as mock_pool, ): get_redis_connection_pool() @@ -360,7 +378,9 @@ async def test_redis_cache_test_connection_uses_shared_factory(clean_redis_envir client.ping = AsyncMock(return_value=True) client.aclose = AsyncMock() - with patch("litellm._redis.get_redis_async_client", return_value=client) as mock_factory: + with patch( # test-quality-ok: the factory call is what routing through it means + "litellm._redis.get_redis_async_client", return_value=client + ) as mock_factory: cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) result = await cache.test_connection() @@ -376,8 +396,12 @@ async def test_redis_cluster_cache_test_connection_uses_shared_factory(clean_red client.ping = AsyncMock(return_value=True) client.aclose = AsyncMock() - with patch("litellm._redis.get_redis_async_client", return_value=client) as mock_factory: - with patch("litellm._redis.get_redis_client", return_value=MagicMock(spec=redis.RedisCluster)): + with patch( # test-quality-ok: the factory call is what routing through it means + "litellm._redis.get_redis_async_client", return_value=client + ) as mock_factory: + with patch( # test-quality-ok: a real RedisCluster would slot-discover here + "litellm._redis.get_redis_client", return_value=MagicMock(spec=redis.RedisCluster) + ): cache = RedisClusterCache( startup_nodes=[{"host": "redis-host", "port": 6379}], credential_provider=provider ) @@ -395,8 +419,12 @@ def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache): async_pool = MagicMock(spec=async_redis.BlockingConnectionPool) with ( - patch("litellm._redis.get_redis_client", return_value=sync_client), - patch("litellm._redis.get_redis_connection_pool", return_value=async_pool), + patch( # test-quality-ok: the hostile provider must not reach a real client + "litellm._redis.get_redis_client", return_value=sync_client + ), + patch( # test-quality-ok: the hostile provider must not reach a real pool + "litellm._redis.get_redis_connection_pool", return_value=async_pool + ), ): cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) second_cache = RedisCache(host="redis-host", port=6379, credential_provider=second_provider) @@ -903,7 +931,7 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_ ) -@patch("litellm._redis.redis.Sentinel") +@patch("litellm._redis.redis.Sentinel") # test-quality-ok: sentinel discovery needs live sentinels def test_sync_sentinel_keeps_provider_off_monitors_and_on_master(mock_sentinel_cls): provider = _StubCredentialProvider() mock_sentinel = MagicMock() diff --git a/uv.lock b/uv.lock index fae22e3759e..628483c0117 100644 --- a/uv.lock +++ b/uv.lock @@ -4342,7 +4342,6 @@ proxy = [ { name = "pyroscope-io", marker = "sys_platform != 'win32'" }, { name = "python-multipart" }, { name = "pyyaml" }, - { name = "redis" }, { name = "restrictedpython" }, { name = "rich" }, { name = "rq" }, @@ -4552,7 +4551,6 @@ requires-dist = [ { name = "python3-saml", marker = "extra == 'saml'", specifier = ">=1.16.0,<2.0" }, { name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" }, { name = "pyyaml", marker = "extra == 'proxy'", specifier = ">=6.0.3,<7.0" }, - { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "redisvl", marker = "extra == 'extra-proxy'", specifier = ">=0.4.1,<1.0" }, { name = "requests", marker = "extra == 'cli'", specifier = ">=2.32.0,<3.0" }, { name = "resend", marker = "extra == 'extra-proxy'", specifier = ">=2.23.0,<3.0" }, From 01a1a3090ba4a66f52319e4f636a2d213049c29e Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 17:46:58 -0400 Subject: [PATCH 035/281] test(redis): remove internal mocking from regressions --- tests/test_litellm/test_redis.py | 159 ++++++++++++++----------------- 1 file changed, 70 insertions(+), 89 deletions(-) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 449f2d8dc14..ed2045ba76f 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -9,6 +9,7 @@ from redis.credentials import CredentialProvider import litellm from litellm._redis import ( + _async_auth_kwargs, _get_redis_client_logic, _get_redis_cluster_kwargs, _get_redis_env_kwarg_mapping, @@ -228,32 +229,28 @@ def test_async_url_pool_strips_userinfo_for_the_provider(clean_redis_environment assert connection.db == 3 -def test_sync_cluster_preserves_credential_provider_identity( # test-quality-ok: constructor kwargs are the only seam - clean_redis_environment, -): +def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] - with patch( # test-quality-ok: RedisCluster slot-discovers in its constructor - "litellm._redis.redis.RedisCluster", autospec=True - ) as mock_cluster_cls: - get_redis_client(startup_nodes=startup_nodes, credential_provider=provider) + with patch("redis.RedisCluster", autospec=True) as mock_cluster_cls: + get_redis_client(startup_nodes=startup_nodes, credential_provider=provider, password="redis-secret") - mock_cluster_cls.assert_called_once() - assert mock_cluster_cls.call_args[1].get("credential_provider") is provider + cluster_kwargs = mock_cluster_cls.call_args.kwargs + assert cluster_kwargs["credential_provider"] is provider + assert "password" not in cluster_kwargs + assert [(node.host, node.port) for node in cluster_kwargs["startup_nodes"]] == [("cluster-node", 6379)] def test_async_cluster_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] - with patch( # test-quality-ok: the async cluster class is the only injection point here - "litellm.caching.redis_cluster_node_isolation.get_litellm_async_redis_cluster_class" - ) as mock_class: - get_redis_async_client(startup_nodes=startup_nodes, credential_provider=provider) + client = get_redis_async_client(startup_nodes=startup_nodes, credential_provider=provider) - call_kwargs = mock_class.return_value.call_args[1] - assert call_kwargs.get("credential_provider") is provider + assert client.connection_kwargs["credential_provider"] is provider + assert client.connection_kwargs["socket_keepalive"] is True + assert client.connection_kwargs["health_check_interval"] == REDIS_CLUSTER_HEALTH_CHECK_INTERVAL def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environment, monkeypatch): @@ -262,10 +259,10 @@ def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environ monkeypatch.setenv("REDIS_AZURE_AD_TOKEN", "true") with ( - patch( # test-quality-ok: the assertion is that this builder is never reached + patch( # test-quality-ok: an auto-auth callback built here is popped again by the provider branch, so the builders are the only place the wasted work is visible "litellm._redis.create_gcp_iam_redis_connect_func" ) as mock_gcp, - patch( # test-quality-ok: the assertion is that this builder is never reached + patch( # test-quality-ok: same as above, and reaching this one also builds an Azure credential the caller never asked for "litellm._redis.create_azure_ad_redis_connect_func" ) as mock_azure, ): @@ -323,108 +320,92 @@ def test_provider_free_url_is_left_untouched(clean_redis_environment): assert redis_kwargs["url"] == url -def test_async_direct_explicit_provider_is_preserved_when_normalization_is_bypassed(): +def test_async_auth_kwargs_supersedes_credentials_an_explicit_provider_replaces(): + """The shared seam both async entry points run through: a provider outranks every other + credential, and redis-py rejects a provider that arrives next to a username or password.""" provider = _StubCredentialProvider() - redis_kwargs = { - "host": "redis-host", - "port": 6379, - "credential_provider": provider, - "redis_connect_func": _gcp_marker_callback(), - } - with ( - patch( # test-quality-ok: bypassing normalization is what this test pins - "litellm._redis._get_redis_client_logic", return_value=redis_kwargs - ), - patch( # test-quality-ok: the constructor kwargs are the only observable - "litellm._redis.async_redis.Redis", autospec=True - ) as mock_redis, - ): - get_redis_async_client() + auth_kwargs = _async_auth_kwargs( + { + "host": "redis-host", + "port": 6379, + "credential_provider": provider, + "redis_connect_func": _gcp_marker_callback(), + "username": "url-user", + "password": "url-pass", + } + ) - call_kwargs = mock_redis.call_args[1] - assert call_kwargs["credential_provider"] is provider - assert "redis_connect_func" not in call_kwargs + assert auth_kwargs["credential_provider"] is provider + assert auth_kwargs["host"] == "redis-host" + assert auth_kwargs["port"] == 6379 + assert "redis_connect_func" not in auth_kwargs + assert "username" not in auth_kwargs + assert "password" not in auth_kwargs -def test_async_pool_explicit_provider_is_preserved_when_normalization_is_bypassed(): - provider = _StubCredentialProvider() - redis_kwargs = { - "host": "redis-host", - "port": 6379, - "credential_provider": provider, - "redis_connect_func": _gcp_marker_callback(), - } +def test_async_auth_kwargs_leaves_provider_free_kwargs_alone(): + redis_kwargs = {"host": "redis-host", "port": 6379, "username": "url-user", "password": "url-pass"} - with ( - patch( # test-quality-ok: bypassing normalization is what this test pins - "litellm._redis._get_redis_client_logic", return_value=redis_kwargs - ), - patch( # test-quality-ok: the constructor kwargs are the only observable - "litellm._redis.async_redis.BlockingConnectionPool", autospec=True - ) as mock_pool, - ): - get_redis_connection_pool() - - call_kwargs = mock_pool.call_args[1] - assert call_kwargs["credential_provider"] is provider - assert "redis_connect_func" not in call_kwargs + assert _async_auth_kwargs(redis_kwargs) == redis_kwargs @pytest.mark.asyncio async def test_redis_cache_test_connection_uses_shared_factory(clean_redis_environment): provider = _StubCredentialProvider() - client = MagicMock(spec=async_redis.Redis) - client.ping = AsyncMock(return_value=True) - client.aclose = AsyncMock() - with patch( # test-quality-ok: the factory call is what routing through it means - "litellm._redis.get_redis_async_client", return_value=client - ) as mock_factory: - cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) + with ( + patch("redis.Redis", autospec=True), + patch("redis.asyncio.BlockingConnectionPool", autospec=True), + patch("redis.asyncio.Redis", autospec=True) as mock_async_redis, + ): + mock_async_redis.return_value.ping = AsyncMock(return_value=True) + mock_async_redis.return_value.aclose = AsyncMock() + cache = RedisCache(host="redis-host", port=6379, credential_provider=provider, password="redis-secret") result = await cache.test_connection() + client_kwargs = mock_async_redis.call_args.kwargs assert result["status"] == "success" - call_kwargs = mock_factory.call_args.kwargs - assert call_kwargs["credential_provider"] is provider + assert client_kwargs["credential_provider"] is provider + assert "password" not in client_kwargs @pytest.mark.asyncio async def test_redis_cluster_cache_test_connection_uses_shared_factory(clean_redis_environment): provider = _StubCredentialProvider() - client = MagicMock(spec=async_redis.RedisCluster) - client.ping = AsyncMock(return_value=True) - client.aclose = AsyncMock() + recorder = MagicMock() - with patch( # test-quality-ok: the factory call is what routing through it means - "litellm._redis.get_redis_async_client", return_value=client - ) as mock_factory: - with patch( # test-quality-ok: a real RedisCluster would slot-discover here - "litellm._redis.get_redis_client", return_value=MagicMock(spec=redis.RedisCluster) - ): - cache = RedisClusterCache( - startup_nodes=[{"host": "redis-host", "port": 6379}], credential_provider=provider - ) + class _StubAsyncCluster: + """A real base class, because the production path subclasses this at call time.""" + + def __init__(self, **kwargs): + recorder(**kwargs) + + async def ping(self): + return True + + async def aclose(self): + return None + + with ( + patch("redis.RedisCluster", autospec=True), + patch("redis.asyncio.cluster.RedisCluster", _StubAsyncCluster), + ): + cache = RedisClusterCache(startup_nodes=[{"host": "redis-host", "port": 6379}], credential_provider=provider) result = await cache.test_connection() + cluster_kwargs = recorder.call_args.kwargs assert result["status"] == "success" - call_kwargs = mock_factory.call_args.kwargs - assert call_kwargs["credential_provider"] is provider + assert cluster_kwargs["credential_provider"] is provider def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache): provider = _HostileCredentialProvider("synthetic-secret") second_provider = _StubCredentialProvider("another-token") - sync_client = MagicMock(spec=redis.Redis) - async_pool = MagicMock(spec=async_redis.BlockingConnectionPool) with ( - patch( # test-quality-ok: the hostile provider must not reach a real client - "litellm._redis.get_redis_client", return_value=sync_client - ), - patch( # test-quality-ok: the hostile provider must not reach a real pool - "litellm._redis.get_redis_connection_pool", return_value=async_pool - ), + patch("redis.Redis", autospec=True), + patch("redis.asyncio.BlockingConnectionPool", autospec=True), ): cache = RedisCache(host="redis-host", port=6379, credential_provider=provider) second_cache = RedisCache(host="redis-host", port=6379, credential_provider=second_provider) @@ -931,7 +912,7 @@ def test_sync_sentinel_uses_sentinel_password_and_master_password(mock_sentinel_ ) -@patch("litellm._redis.redis.Sentinel") # test-quality-ok: sentinel discovery needs live sentinels +@patch("redis.Sentinel") def test_sync_sentinel_keeps_provider_off_monitors_and_on_master(mock_sentinel_cls): provider = _StubCredentialProvider() mock_sentinel = MagicMock() From 442175c4dc9c0de4200ff3e940389bc009ae14b3 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 07:58:20 +0000 Subject: [PATCH 036/281] chore(typing): clear fresh tech debt from the Aug 24 window type the strategy-router health check params instead of a bare dict, annotate the new interactions usage locals Final, drop a reportUnnecessaryIsInstance suppression by narrowing the grounding tool list before iterating it, and delete the duplicated file-id decode comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 12 +++--- .../usage_object_transformation.py | 37 ++++++++++--------- .../prompt_templates/common_utils.py | 6 --- litellm/proxy/health_check.py | 2 +- ruff-strict-budget.json | 8 ++-- type-discipline-budget.json | 6 +-- 6 files changed, 34 insertions(+), 37 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 3f0011c80d2..5ec49f48dc7 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,6 +1,6 @@ { "reportAny": { - "limit": 19945 + "limit": 19936 }, "reportArgumentType": { "limit": 2566 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 6048 + "limit": 6047 }, "reportFunctionMemberAccess": { "limit": 7 @@ -57,7 +57,7 @@ "limit": 5663 }, "reportMissingTypeArgument": { - "limit": 15545 + "limit": 15536 }, "reportMissingTypeStubs": { "limit": 40 @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 38998 + "limit": 38990 }, "reportUnknownParameterType": { - "limit": 19876 + "limit": 19868 }, "reportUnknownVariableType": { - "limit": 30554 + "limit": 30540 }, "reportUnnecessaryCast": { "limit": 117 diff --git a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py index df436ef7611..f11f6d46fb2 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py +++ b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py @@ -1,6 +1,6 @@ from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import Any +from typing import Any, Final from litellm.types.utils import ( CompletionTokensDetailsWrapper, @@ -39,7 +39,7 @@ class TranscriptionUsageObjectTransformation: return None -_INTERACTIONS_MODALITY_FIELDS: Mapping[str, str] = MappingProxyType( +_INTERACTIONS_MODALITY_FIELDS: Final[Mapping[str, str]] = MappingProxyType( { "text": "text_tokens", "audio": "audio_tokens", @@ -59,7 +59,7 @@ def _token_count(value: object) -> int: def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, int]: - fields = frozenset(field for entry in entries if (field := _modality_field(entry)) is not None) + fields: Final = frozenset(field for entry in entries if (field := _modality_field(entry)) is not None) return MappingProxyType( { field: sum(_token_count(entry.get("tokens")) for entry in entries if _modality_field(entry) == field) @@ -69,10 +69,13 @@ def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, i def _google_search_query_count(usage_object: Mapping[str, Any]) -> int: + entries: Final = usage_object.get("grounding_tool_count") + if not isinstance(entries, Sequence): + return 0 return sum( _token_count(entry.get("count")) - for entry in tuple(usage_object.get("grounding_tool_count") or ()) - if isinstance(entry, Mapping) and entry.get("type") == "google_search" # pyright: ignore[reportUnnecessaryIsInstance] # provider JSON, not the empty tuple inferred from `or ()` + for entry in entries + if isinstance(entry, Mapping) and entry.get("type") == "google_search" ) @@ -112,30 +115,30 @@ class InteractionsUsageObjectTransformation: @staticmethod def transform_interactions_usage_object(usage_object: Mapping[str, Any]) -> Usage: - input_entries = tuple(usage_object.get("input_tokens_by_modality") or ()) + tuple( + input_entries: Final = tuple(usage_object.get("input_tokens_by_modality") or ()) + tuple( usage_object.get("tool_use_tokens_by_modality") or () ) - cached_sums = _modality_token_sums(tuple(usage_object.get("cached_tokens_by_modality") or ())) - output_sums = _modality_token_sums(tuple(usage_object.get("output_tokens_by_modality") or ())) + cached_sums: Final = _modality_token_sums(tuple(usage_object.get("cached_tokens_by_modality") or ())) + output_sums: Final = _modality_token_sums(tuple(usage_object.get("output_tokens_by_modality") or ())) - total_cached_tokens = _token_count(usage_object.get("total_cached_tokens")) - input_sums = _subtract_cached_from_input( + total_cached_tokens: Final = _token_count(usage_object.get("total_cached_tokens")) + input_sums: Final = _subtract_cached_from_input( input_sums=_modality_token_sums(input_entries), cached_sums=cached_sums, total_cached_tokens=total_cached_tokens, ) - reasoning_tokens = _token_count(usage_object.get("total_reasoning_tokens")) or _token_count( + reasoning_tokens: Final = _token_count(usage_object.get("total_reasoning_tokens")) or _token_count( usage_object.get("total_thought_tokens") ) - prompt_tokens = _token_count(usage_object.get("total_input_tokens")) + _token_count( + prompt_tokens: Final = _token_count(usage_object.get("total_input_tokens")) + _token_count( usage_object.get("total_tool_use_tokens") ) - completion_tokens = _token_count(usage_object.get("total_output_tokens")) + reasoning_tokens - total_tokens = _token_count(usage_object.get("total_tokens")) or (prompt_tokens + completion_tokens) + completion_tokens: Final = _token_count(usage_object.get("total_output_tokens")) + reasoning_tokens + total_tokens: Final = _token_count(usage_object.get("total_tokens")) or (prompt_tokens + completion_tokens) - web_search_requests = _google_search_query_count(usage_object) - prompt_tokens_details = ( + web_search_requests: Final = _google_search_query_count(usage_object) + prompt_tokens_details: Final = ( PromptTokensDetailsWrapper( cached_tokens=total_cached_tokens or None, web_search_requests=web_search_requests or None, @@ -144,7 +147,7 @@ class InteractionsUsageObjectTransformation: if input_sums or total_cached_tokens or web_search_requests else None ) - completion_tokens_details = ( + completion_tokens_details: Final = ( CompletionTokensDetailsWrapper( reasoning_tokens=reasoning_tokens or None, **output_sums, diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 72ea85dfa33..748347fe938 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -511,9 +511,6 @@ def update_messages_with_model_file_ids( if "llm_output_file_id," in unified_file_id: provider_file_id = unified_file_id.split("llm_output_file_id,")[1].split(";")[0] if not provider_file_id and is_model_embedded_id(file_id): - # `litellm:;model,` encoding from the - # x-litellm-model upload path. Strip the wrapper - # so the provider sees its own ID. provider_file_id = get_original_file_id(file_id) file_object_file_field["file_id"] = provider_file_id or file_id if format: @@ -588,9 +585,6 @@ def update_responses_input_with_model_file_ids( updated_content_item["file_id"] = provider_file_id updated_content.append(updated_content_item) elif is_model_embedded_id(file_id): - # `litellm:;model,` encoding from the - # x-litellm-model upload path. Strip the wrapper - # so the provider sees its own ID. updated_content_item = content_item.copy() updated_content_item["file_id"] = get_original_file_id(file_id) updated_content.append(updated_content_item) diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index d8fde8ca5dc..4e12974189f 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -183,7 +183,7 @@ async def run_with_timeout(task, timeout): return {"error": "Timeout exceeded", "exception": timeout_exception} -def _is_strategy_router_deployment(litellm_params: dict) -> bool: +def _is_strategy_router_deployment(litellm_params: Mapping[str, object]) -> bool: """True for strategy-router deployments.""" model: Final[object] = litellm_params.get("model", "") return isinstance(model, str) and classify_strategy_router_model(model) is not None diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 69902ecfbbb..1cf302a98c9 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -12,10 +12,10 @@ "limit": 2016 }, "ANN202": { - "limit": 850 + "limit": 849 }, "ANN204": { - "limit": 709 + "limit": 708 }, "ANN205": { "limit": 112 @@ -24,7 +24,7 @@ "limit": 133 }, "ANN401": { - "limit": 1185 + "limit": 1183 }, "ASYNC230": { "limit": 11 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1210 + "limit": 1209 }, "TRY002": { "limit": 524 diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 4f2314b2a0a..4054af17d1e 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22795 + "limit": 22788 }, "LIT002": { - "limit": 26872 + "limit": 26871 }, "LIT003": { "limit": 269 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16672 + "limit": 16657 }, "LIT011": { "limit": 5588 From 9dabd72f2d7f13c148a6e3129a0a676030670a3a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:14:17 +0000 Subject: [PATCH 037/281] refactor(repositories): type prisma table access with one generic protocol Every repository handed its `.table` back untyped, so a dozen modules had each grown a private `_PrismaTableActions` Protocol to paper over it. They had drifted: some declared `update` as returning the row, others the row or None, and none agreed on whether `find_many` was covariant Replace all of them with a single `TableActions[RowT_co]` in `litellm/repositories/prisma_protocols.py`, keyed to the prisma row each repository is bound to. Query inputs stay `Mapping[str, object]` so callers keep passing plain dicts, and `find_many` returns `Sequence` so the row type stays covariant Typing the nullable returns honestly surfaced paths that were already crashing. A team admin could never edit or delete a memory entry owned by their team: the write-auth check fed a raw prisma row to a helper that expects the domain model, so `members_with_roles` arrived as plain dicts and the request died as a 500 instead of applying the edit. Non-admin members hit the same 500 in place of the 403 they were owed, so refusal and breakage were indistinguishable. `/v2/model/info?user_models_only=true` dereferenced a missing user row rather than returning the 400 the route already had, three team routes dereferenced a team deleted between the read and the write, and the agent registry dereferenced a missing agent instead of naming it basedpyright drops 2,132 errors, 1,454 of them reportAny and 73 reportExplicitAny. The dashboard's generated types pick up `string[]` where they had `unknown[]` for a team's members, admins and models --- basedpyright-code-budget.json | 32 +- .../proxy/common_utils/check_batch_cost.py | 4 +- litellm/integrations/prometheus.py | 5 +- litellm/models/team.py | 6 +- litellm/proxy/_experimental/mcp_server/db.py | 99 ++--- .../proxy/agent_endpoints/agent_registry.py | 84 +++-- .../claude_code_marketplace.py | 4 +- litellm/proxy/auth/auth_checks.py | 18 +- litellm/proxy/auth/user_api_key_auth.py | 17 +- .../proxy/common_utils/config_sync_pubsub.py | 8 +- .../expired_ui_session_key_cleanup_manager.py | 12 +- .../common_utils/key_rotation_manager.py | 12 +- .../proxy/common_utils/reset_budget_job.py | 6 +- .../proxy/container_endpoints/ownership.py | 37 +- .../proxy/credential_endpoints/endpoints.py | 13 +- litellm/proxy/db/tool_registry_writer.py | 30 +- .../proxy/guardrails/guardrail_endpoints.py | 22 +- .../proxy/guardrails/guardrail_registry.py | 24 +- litellm/proxy/guardrails/usage_endpoints.py | 38 +- litellm/proxy/guardrails/usage_tracking.py | 16 +- .../access_group_endpoints.py | 4 +- .../budget_management_endpoints.py | 2 +- .../cache_settings_endpoints.py | 3 +- .../common_daily_activity.py | 8 +- .../config_override_endpoints.py | 3 +- .../internal_user_endpoints.py | 121 +++---- .../jwt_key_mapping_endpoints.py | 3 + .../key_management_endpoints.py | 199 +++++------ ...model_access_group_management_endpoints.py | 21 +- .../model_management_endpoints.py | 88 +++-- .../organization_endpoints.py | 43 ++- .../scim/scim_transformations.py | 3 +- .../management_endpoints/scim/scim_v2.py | 6 + .../tag_management_endpoints.py | 4 +- .../team_callback_endpoints.py | 3 + .../management_endpoints/team_endpoints.py | 337 +++++++++--------- litellm/proxy/management_endpoints/ui_sso.py | 63 +--- .../object_permission_utils.py | 14 +- litellm/proxy/management_helpers/utils.py | 26 +- litellm/proxy/memory/memory_endpoints.py | 85 ++--- .../openai_files_endpoints/common_utils.py | 24 +- .../managed_id_rewriter.py | 16 +- .../pass_through_endpoints.py | 11 +- .../proxy/policy_engine/policy_registry.py | 68 ++-- .../policy_engine/policy_resolve_endpoints.py | 56 +-- litellm/proxy/prompts/prompt_endpoints.py | 8 +- litellm/proxy/proxy_server.py | 107 +++--- .../spend_tracking/cloudzero_endpoints.py | 32 +- .../spend_management_endpoints.py | 55 +-- .../proxy/spend_tracking/vantage_endpoints.py | 30 +- .../proxy_setting_endpoints.py | 58 ++- litellm/proxy/utils.py | 70 ++-- .../proxy/vector_store_endpoints/endpoints.py | 13 +- .../management_endpoints.py | 24 +- litellm/repositories/base_repository.py | 31 +- litellm/repositories/budget_repository.py | 8 +- litellm/repositories/config_repository.py | 2 +- .../repositories/credentials_repository.py | 53 ++- litellm/repositories/model_repository.py | 42 +-- .../object_permission_repository.py | 8 +- .../repositories/organization_repository.py | 8 +- litellm/repositories/prisma_protocols.py | 87 +++++ litellm/repositories/project_repository.py | 8 +- litellm/repositories/table_repositories.py | 116 +++--- litellm/repositories/team_repository.py | 8 +- .../repositories/user_banner_repository.py | 7 +- litellm/repositories/user_repository.py | 8 +- .../verification_token_repository.py | 22 +- .../responses/file_search/emulated_handler.py | 56 +-- .../custom_tools.py | 10 +- .../handler.py | 8 +- .../session_handler.py | 6 +- .../streaming_iterator.py | 23 +- .../transformation.py | 82 +++-- litellm/responses/main.py | 4 +- .../responses/mcp/chat_completions_handler.py | 22 +- .../responses/mcp/mcp_streaming_iterator.py | 14 +- litellm/responses/mcp/request_context.py | 23 +- litellm/responses/sse_output_recovery.py | 53 ++- litellm/responses/streaming_iterator.py | 42 ++- litellm/responses/utils.py | 56 +-- .../vector_stores/vector_store_registry.py | 13 +- ruff-strict-budget.json | 12 +- .../proxy_unit_tests/test_jwt_key_mapping.py | 28 ++ .../agent_endpoints/test_agent_registry.py | 57 +++ .../test_claude_code_marketplace.py | 24 ++ .../common_utils/test_reset_budget_job.py | 2 +- .../proxy/db/mcp_server/test_db.py | 33 +- .../guardrails/test_guardrail_registry.py | 22 ++ .../scim/test_scim_v2_endpoints.py | 43 +++ .../test_internal_user_endpoints.py | 2 +- .../test_key_management_endpoints.py | 39 ++ .../test_model_management_endpoints.py | 55 +++ .../test_organization_endpoints.py | 26 ++ .../test_team_callback_endpoints.py | 38 ++ .../test_team_endpoints.py | 96 ++++- .../proxy/memory/test_memory_endpoints.py | 72 +++- .../prompts/test_prompt_endpoints_crud.py | 55 +++ tests/test_litellm/proxy/test_proxy_server.py | 33 ++ .../test_vector_store_endpoints.py | 51 +++ type-discipline-budget.json | 12 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 30 +- 102 files changed, 2320 insertions(+), 1325 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 664e1669834..7e57539d1dd 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,18 +1,18 @@ { "reportAny": { - "limit": 19955 + "limit": 18501 }, "reportArgumentType": { - "limit": 2566 + "limit": 2564 }, "reportAssignmentType": { "limit": 320 }, "reportAttributeAccessIssue": { - "limit": 488 + "limit": 483 }, "reportCallIssue": { - "limit": 114 + "limit": 113 }, "reportConstantRedefinition": { "limit": 40 @@ -24,7 +24,7 @@ "limit": 19 }, "reportExplicitAny": { - "limit": 6049 + "limit": 5976 }, "reportFunctionMemberAccess": { "limit": 7 @@ -45,7 +45,7 @@ "limit": 35 }, "reportInvalidTypeForm": { - "limit": 35 + "limit": 34 }, "reportInvalidTypeVarUse": { "limit": 2 @@ -54,10 +54,10 @@ "limit": 0 }, "reportMissingParameterType": { - "limit": 5663 + "limit": 5661 }, "reportMissingTypeArgument": { - "limit": 15555 + "limit": 15504 }, "reportMissingTypeStubs": { "limit": 40 @@ -72,7 +72,7 @@ "limit": 0 }, "reportOptionalMemberAccess": { - "limit": 1061 + "limit": 1058 }, "reportOptionalOperand": { "limit": 0 @@ -90,7 +90,7 @@ "limit": 8 }, "reportReturnType": { - "limit": 213 + "limit": 212 }, "reportTypedDictNotRequiredAccess": { "limit": 26 @@ -99,31 +99,31 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 44655 + "limit": 44527 }, "reportUnknownLambdaType": { "limit": 109 }, "reportUnknownMemberType": { - "limit": 39011 + "limit": 38827 }, "reportUnknownParameterType": { - "limit": 19885 + "limit": 19848 }, "reportUnknownVariableType": { - "limit": 30569 + "limit": 30384 }, "reportUnnecessaryCast": { "limit": 117 }, "reportUnnecessaryComparison": { - "limit": 699 + "limit": 697 }, "reportUnnecessaryContains": { "limit": 5 }, "reportUnnecessaryIsInstance": { - "limit": 836 + "limit": 833 }, "reportUntypedBaseClass": { "limit": 0 diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 76e92538aaa..aee3295d1da 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -14,6 +14,8 @@ from litellm.constants import ( ) if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.integrations.prometheus import PrometheusLogger from litellm.proxy._types import LiteLLM_ManagedObjectTable from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -351,7 +353,7 @@ class CheckBatchCost: return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error) async def _finalize_unbilled_terminal_job( - self, job: "LiteLLM_ManagedObjectTable", response: "LiteLLMBatch" + self, job: "prisma_models.LiteLLM_ManagedObjectTable", response: "LiteLLMBatch" ) -> None: """Persist a terminal batch that has nothing billable, converting any raw provider file ids to managed ids, and take it out of the poll page.""" diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index f9195db1d67..beda1a29075 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -96,7 +96,10 @@ class _PaginatedPrismaTable(Protocol[_TableRowT]): def _paginated_table(repository: BaseRepository[_TableRowT]) -> _PaginatedPrismaTable[_TableRowT]: """View a repository's prisma table through the pagination surface budget metrics need.""" - return repository.table + return cast( + _PaginatedPrismaTable[_TableRowT], + repository.table, # cast-ok: prisma rows carry the budget columns the domain model declares + ) class _OrgBudgetRow(Protocol): diff --git a/litellm/models/team.py b/litellm/models/team.py index 544e2cf5bbc..da526515e6e 100644 --- a/litellm/models/team.py +++ b/litellm/models/team.py @@ -64,8 +64,8 @@ class TeamBase(LiteLLMPydanticObjectBase): team_alias: str | None = None team_id: str | None = None organization_id: str | None = None - admins: list = [] - members: list = [] + admins: list[str] = [] + members: list[str] = [] members_with_roles: list[Member] = [] team_member_permissions: list[str] | None = None metadata: dict | None = None @@ -75,7 +75,7 @@ class TeamBase(LiteLLMPydanticObjectBase): soft_budget: float | None = None budget_duration: str | None = None budget_limits: list[BudgetLimitEntry] | None = None - models: list = [] + models: list[str] = [] blocked: bool = False router_settings: dict | None = None access_group_ids: list[str] | None = None diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 28638ed9c77..4aa08020527 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -4,7 +4,7 @@ import hashlib import json from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, TypeVar, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -13,7 +13,6 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.oauth_utils import build_upstream_oauth2_token_request from litellm.proxy._types import ( LiteLLM_MCPServerTable, - LiteLLM_ObjectPermissionTable, MCPApprovalStatus, MCPEnvVar, MCPEnvVarScope, @@ -30,6 +29,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.utils import PrismaClient from litellm.repositories.object_permission_repository import ObjectPermissionRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( MCPServerOAuthClientRepository, MCPServerRepository, @@ -48,34 +48,9 @@ if TYPE_CHECKING: from litellm.types.mcp_server.mcp_server_manager import MCPServer -_RowT = TypeVar("_RowT") - - -class _TableActions(Protocol[_RowT]): - async def find_unique( - self, where: Mapping[str, object], include: Mapping[str, object] | None = None - ) -> _RowT | None: ... - - async def find_many( - self, - take: int | None = None, - where: Mapping[str, object] | None = None, - order: Mapping[str, object] | None = None, - ) -> list[_RowT]: ... - - async def create(self, data: Mapping[str, object]) -> _RowT: ... - - async def upsert(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT: ... - - async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT | None: ... - - async def delete(self, where: Mapping[str, object]) -> _RowT | None: ... - - async def delete_many(self, where: Mapping[str, object] | None = None) -> int: ... - class _UserEnvVarsTransactionClient(Protocol): - litellm_mcpuserenvvars: "_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]" + litellm_mcpuserenvvars: "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]" async def execute_raw(self, query: str, *args: object) -> int: ... @@ -473,15 +448,15 @@ def _credentials_blob_to_mutable_dict(blob: str | Mapping[str, object]) -> dict[ def _mcp_server_table_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_MCPServerTable]": - table: Final[_TableActions[prisma_db_models.LiteLLM_MCPServerTable]] = MCPServerRepository(prisma_client).table +) -> "TableActions[prisma_db_models.LiteLLM_MCPServerTable]": + table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerTable]] = MCPServerRepository(prisma_client).table return table def _verification_token_table_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_VerificationToken]": - table: Final[_TableActions[prisma_db_models.LiteLLM_VerificationToken]] = VerificationTokenRepository( +) -> "TableActions[prisma_db_models.LiteLLM_VerificationToken]": + table: Final[TableActions[prisma_db_models.LiteLLM_VerificationToken]] = VerificationTokenRepository( prisma_client ).table return table @@ -489,15 +464,15 @@ def _verification_token_table_actions( def _team_table_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_TeamTable]": - table: Final[_TableActions[prisma_db_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table +) -> "TableActions[prisma_db_models.LiteLLM_TeamTable]": + table: Final[TableActions[prisma_db_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table return table def _oauth_client_table_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]": - table: Final[_TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = MCPServerOAuthClientRepository( +) -> "TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]": + table: Final[TableActions[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = MCPServerOAuthClientRepository( prisma_client ).table return table @@ -511,7 +486,7 @@ def _db_transaction_manager(prisma_client: PrismaClient) -> _UserEnvVarsTransact async def _db_find_mcp_server_rows( prisma_client: PrismaClient, where: "prisma_db_types.LiteLLM_MCPServerTableWhereInput | None" = None, -) -> "list[prisma_db_models.LiteLLM_MCPServerTable]": +) -> "Sequence[prisma_db_models.LiteLLM_MCPServerTable]": return await _mcp_server_table_actions(prisma_client).find_many(where=where) @@ -526,17 +501,19 @@ async def _db_update_mcp_server_row( server_id: str, data: "prisma_db_types.LiteLLM_MCPServerTableUpdateInput", ) -> "prisma_db_models.LiteLLM_MCPServerTable": - row: Final[prisma_db_models.LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.update( + row: Final[prisma_db_models.LiteLLM_MCPServerTable | None] = await _mcp_server_table_actions(prisma_client).update( where={"server_id": server_id}, data=data, ) + if row is None: + raise ValueError(f"MCP server not found, passed server_id={server_id}") return row def _user_credential_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]": - table: Final[_TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = MCPUserCredentialsRepository( +) -> "TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]": + table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserCredentials]] = MCPUserCredentialsRepository( prisma_client ).table return table @@ -544,8 +521,8 @@ def _user_credential_actions( def _user_env_var_actions( prisma_client: PrismaClient, -) -> "_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]": - table: Final[_TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = prisma_client.db.litellm_mcpuserenvvars +) -> "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]": + table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = prisma_client.db.litellm_mcpuserenvvars return table @@ -560,7 +537,7 @@ async def _db_find_user_credential_row( async def _db_find_user_credential_rows( prisma_client: PrismaClient, where: "prisma_db_types.LiteLLM_MCPUserCredentialsWhereInput | None" = None, -) -> "list[prisma_db_models.LiteLLM_MCPUserCredentials]": +) -> "Sequence[prisma_db_models.LiteLLM_MCPUserCredentials]": return await _user_credential_actions(prisma_client).find_many(where=where) @@ -583,7 +560,7 @@ async def _db_upsert_user_credential_row( async def _db_find_user_env_var_rows( prisma_client: PrismaClient, where: "prisma_db_types.LiteLLM_MCPUserEnvVarsWhereInput | None" = None, -) -> "list[prisma_db_models.LiteLLM_MCPUserEnvVars]": +) -> "Sequence[prisma_db_models.LiteLLM_MCPUserEnvVars]": return await _user_env_var_actions(prisma_client).find_many(where=where) @@ -658,7 +635,7 @@ async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str] """ Returns the matching mcp servers from the db with the server_ids """ - _mcp_servers: Final[list[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions( + _mcp_servers: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions( prisma_client ).find_many( where={ @@ -745,13 +722,13 @@ async def get_all_mcp_servers_for_user( async def get_objectpermissions_for_mcp_server( prisma_client: PrismaClient, mcp_server_id: str -) -> list[LiteLLM_ObjectPermissionTable]: +) -> "Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable]": """ Get all the object permissions records and the associated team and verficiationtoken records that have access to the mcp server """ - object_permission_records: Final[list[LiteLLM_ObjectPermissionTable]] = await ObjectPermissionRepository( - prisma_client - ).table.find_many( + object_permission_records: Final[ + Sequence[prisma_db_models.LiteLLM_ObjectPermissionTable] + ] = await ObjectPermissionRepository(prisma_client).table.find_many( where={ "mcp_servers": {"has": mcp_server_id}, }, @@ -766,19 +743,19 @@ async def get_objectpermissions_for_mcp_server( async def get_virtualkeys_for_mcp_server( prisma_client: PrismaClient, server_id: str -) -> "list[prisma_db_models.LiteLLM_VerificationToken]": +) -> "Sequence[prisma_db_models.LiteLLM_VerificationToken]": """ Get all the virtual keys that have access to the mcp server """ - virtual_keys: Final[list[prisma_db_models.LiteLLM_VerificationToken] | None] = await VerificationTokenRepository( - prisma_client - ).table.find_many( + virtual_keys: Final[ + Sequence[prisma_db_models.LiteLLM_VerificationToken] | None + ] = await VerificationTokenRepository(prisma_client).table.find_many( where={ "mcp_servers": {"has": server_id}, }, ) - if virtual_keys is None: + if virtual_keys is None: # pyright: ignore[reportUnnecessaryComparison] # unreachable per seam types; kept as-is return [] return virtual_keys @@ -860,7 +837,7 @@ async def delete_mcp_server( invalidate_token_cache = global_mcp_server_manager.invalidate_user_oauth_token_cache for user_id in credential_user_ids: await invalidate_token_cache(user_id, server_id) - return deleted_server + return deleted_server # pyright: ignore[reportReturnType] # prisma row, not domain LiteLLM_MCPServerTable async def create_mcp_server( @@ -880,7 +857,7 @@ async def create_mcp_server( data_dict["updated_by"] = touched_by new_mcp_server: Final[LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.create( - data=data_dict + data=data_dict, # pyright: ignore[reportAssignmentType] # prisma row, not domain LiteLLM_MCPServerTable ) _decrypt_env_vars_on_returned_row(new_mcp_server) @@ -982,7 +959,7 @@ async def update_mcp_server( data: UpdateMCPServerRequest, touched_by: str, fields_set: set[str] | None = None, -) -> LiteLLM_MCPServerTable: +) -> LiteLLM_MCPServerTable | None: """ Update a new mcp server record in the db """ @@ -1093,9 +1070,9 @@ async def update_mcp_server( data_dict["credentials"] = Json(None) - updated_mcp_server: Final[LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.update( + updated_mcp_server: Final[LiteLLM_MCPServerTable | None] = await MCPServerRepository(prisma_client).table.update( where={"server_id": data.server_id}, - data=data_dict, + data=data_dict, # pyright: ignore[reportAssignmentType] # prisma row, not domain LiteLLM_MCPServerTable ) _decrypt_env_vars_on_returned_row(updated_mcp_server) @@ -1181,7 +1158,7 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, ) updated += 1 - oauth_clients: Final[list[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await _oauth_client_table_actions( + oauth_clients: Final[Sequence[prisma_db_models.LiteLLM_MCPServerOAuthClient]] = await _oauth_client_table_actions( prisma_client ).find_many() oauth_updated = 0 @@ -1914,7 +1891,7 @@ async def get_mcp_submissions( along with a summary count breakdown by approval_status. Mirrors get_guardrail_submissions() from guardrail_endpoints.py. """ - rows: Final[list[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions( + rows: Final[Sequence[prisma_db_models.LiteLLM_MCPServerTable]] = await _mcp_server_table_actions( prisma_client ).find_many( where={"submitted_at": {"not": None}}, diff --git a/litellm/proxy/agent_endpoints/agent_registry.py b/litellm/proxy/agent_endpoints/agent_registry.py index 64de6827679..fa33a307438 100644 --- a/litellm/proxy/agent_endpoints/agent_registry.py +++ b/litellm/proxy/agent_endpoints/agent_registry.py @@ -4,7 +4,7 @@ import json from collections.abc import Iterator, Mapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Any, Final, NamedTuple, Protocol, TypedDict +from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, TypedDict import litellm from litellm.litellm_core_utils.safe_json_dumps import safe_dumps @@ -12,9 +12,13 @@ from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, ) from litellm.proxy.utils import PrismaClient +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest +if TYPE_CHECKING: + from prisma import models as prisma_models + class AgentObjectPermissionRecord(Protocol): def model_dump(self) -> dict[str, object]: ... @@ -42,11 +46,20 @@ class AgentRecordDump(TypedDict): class AgentRecord(Protocol): - agent_id: str - agent_name: str - object_permission_id: str | None - object_permission: AgentObjectPermissionRecord | None - spend: float + @property + def agent_id(self) -> str: ... + + @property + def agent_name(self) -> str: ... + + @property + def object_permission_id(self) -> str | None: ... + + @property + def object_permission(self) -> AgentObjectPermissionRecord | None: ... + + @property + def spend(self) -> float: ... def model_dump(self) -> AgentRecordDump: ... @@ -57,50 +70,47 @@ class AgentTableClient(Protocol): async def create( self, data: Mapping[str, object], - include: Mapping[str, bool] | None = None, + include: Mapping[str, object] | None = None, ) -> AgentRecord: ... async def find_unique( self, where: Mapping[str, object], - include: Mapping[str, bool] | None = None, + include: Mapping[str, object] | None = None, ) -> AgentRecord | None: ... async def find_many( self, where: Mapping[str, object] | None = None, order: Mapping[str, str] | None = None, - include: Mapping[str, bool] | None = None, + include: Mapping[str, object] | None = None, ) -> Sequence[AgentRecord]: ... async def update( self, - where: Mapping[str, object], data: Mapping[str, object], - include: Mapping[str, bool] | None = None, - ) -> AgentRecord: ... + where: Mapping[str, object], + include: Mapping[str, object] | None = None, + ) -> AgentRecord | None: ... - async def delete(self, where: Mapping[str, object]) -> AgentRecord: ... + async def delete( + self, + where: Mapping[str, object], + include: Mapping[str, object] | None = None, + ) -> AgentRecord | None: ... def agents_table(prisma_client: PrismaClient) -> AgentTableClient: - table: Final[AgentTableClient] = AgentsRepository(prisma_client).table + table: Final[AgentTableClient] = AgentsRepository(prisma_client).table # pyright: ignore[reportAssignmentType] # prisma rows type model_dump() as dict[str, Any] return table -class ObjectPermissionGrantRecord(Protocol): - object_permission_id: str - agents: list[str] | None - - -class ObjectPermissionTableClient(Protocol): - async def find_many(self, where: Mapping[str, object]) -> Sequence[ObjectPermissionGrantRecord]: ... - - async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... - - -def object_permission_table(prisma_client: PrismaClient) -> ObjectPermissionTableClient: - table: Final[ObjectPermissionTableClient] = ObjectPermissionRepository(prisma_client).table +def object_permission_table( + prisma_client: PrismaClient, +) -> "TableActions[prisma_models.LiteLLM_ObjectPermissionTable]": + table: Final[TableActions[prisma_models.LiteLLM_ObjectPermissionTable]] = ObjectPermissionRepository( + prisma_client + ).table return table @@ -222,7 +232,9 @@ class AgentRegistry: self.load_agents_from_config(agent_config if agent_config is not None else self.config_agents) return self.agent_list - async def migrate_legacy_grant_ids(self, table: ObjectPermissionTableClient) -> GrantMigrationResult: + async def migrate_legacy_grant_ids( + self, table: "TableActions[prisma_models.LiteLLM_ObjectPermissionTable]" + ) -> GrantMigrationResult: """ Rewrite object_permission.agents rows holding a legacy full-entry hash to the stable name-derived id. @@ -360,6 +372,8 @@ class AgentRegistry: """ try: deleted_agent: Final = await agents_table(prisma_client).delete(where={"agent_id": agent_id}) + if deleted_agent is None: + raise ValueError(f"Agent not found, passed agent_id={agent_id}") return dict(deleted_agent) except Exception as e: raise Exception(f"Error deleting agent from DB: {e}") @@ -386,12 +400,12 @@ class AgentRegistry: The patched agent """ try: - existing_agent = await AgentsRepository(prisma_client).table.find_unique(where={"agent_id": agent_id}) - if existing_agent is not None: - existing_agent = dict(existing_agent) - - if existing_agent is None: + existing_row: Final = await AgentsRepository(prisma_client).table.find_unique( + where={"agent_id": agent_id} # mutable-ok: prisma filters are plain dicts + ) + if existing_row is None: raise Exception(f"Agent with ID {agent_id} not found") + existing_agent: Final = dict(existing_row) augment_agent: Final = {**existing_agent, **agent} update_data: Final[dict[str, Any]] = {} @@ -436,6 +450,8 @@ class AgentRegistry: }, include={"object_permission": True}, ) + if patched_agent is None: + raise ValueError(f"Agent not found, passed agent_id={agent_id}") patched_agent_dict: Final = patched_agent.model_dump() if patched_agent.object_permission is not None: try: @@ -523,6 +539,8 @@ class AgentRegistry: include={"object_permission": True}, ) + if updated_agent is None: + raise ValueError(f"Agent not found, passed agent_id={agent_id}") updated_agent_dict: Final = updated_agent.model_dump() if updated_agent.object_permission is not None: try: diff --git a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py index 09ba2c93ea0..65bc46edfaf 100644 --- a/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py +++ b/litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_marketplace.py @@ -543,7 +543,7 @@ async def update_plugin( manifest: Final[Mapping[str, object]] = _build_plugin_manifest(plugin_name, request) - plugin: Final[_PluginRecord] = await ClaudeCodePluginRepository(prisma_client).table.update( + plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.update( where={"name": plugin_name}, # mutable-ok: prisma query arguments must be plain dicts data={ # mutable-ok: prisma query arguments must be plain dicts "version": request.version, @@ -553,6 +553,8 @@ async def update_plugin( "updated_at": datetime.now(timezone.utc), }, ) + if plugin is None: + raise _error_response(404, f"Plugin '{plugin_name}' not found") verbose_proxy_logger.info("Plugin %s updated successfully", plugin_name) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e7b98b3cc7f..3e8537b070f 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -156,7 +156,12 @@ class _PrismaVectorStoreRow(Protocol): class _PrismaUserRow(Protocol): user_id: str - organization_memberships: Sequence[LiteLLM_OrganizationMembershipTable | None] | None + + @property + def organization_memberships(self) -> Sequence[_PrismaModelDumpRow | None] | None: ... + + @organization_memberships.setter + def organization_memberships(self, value: Sequence[_PrismaModelDumpRow] | None) -> None: ... def __iter__(self) -> Iterator[tuple[str, object]]: ... @@ -214,9 +219,14 @@ def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_P return repo.table +class _VectorStorePermissionsRow(Protocol): + @property + def vector_stores(self) -> Sequence[str] | None: ... + + def _object_permission_table( - repo: _PrismaTableHolder[LiteLLM_ObjectPermissionTable], -) -> _PrismaAuthTable[LiteLLM_ObjectPermissionTable]: + repo: _PrismaTableHolder[_VectorStorePermissionsRow], +) -> _PrismaAuthTable[_VectorStorePermissionsRow]: return repo.table @@ -5277,7 +5287,7 @@ async def vector_store_access_check( def _can_object_call_vector_stores( object_type: Literal["key", "team", "org"], vector_store_ids_to_run: list[str], - object_permissions: LiteLLM_ObjectPermissionTable | None, + object_permissions: _VectorStorePermissionsRow | None, ): """ Raises ProxyException if the object (key, team, org) cannot access the specific vector store. diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 658d176f6a7..035265d55f4 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -196,6 +196,15 @@ class _UserModelBudgetLimiter(Protocol): ) -> bool: ... +class _TokenTeamModels(Protocol): + @property + def team_models(self) -> list[str]: ... + + +def _token_team_models(valid_token: _TokenTeamModels) -> list[str]: + return valid_token.team_models + + async def _read_user_model_max_budget( user_id: str | None, prisma_client: PrismaClient | None, @@ -1991,7 +2000,7 @@ async def _user_api_key_auth_builder( include={"litellm_budget_table": True}, ) if _db_member is not None: - team_member_info = LiteLLM_TeamMembership(**_db_member.dict()) + team_member_info = LiteLLM_TeamMembership(**_db_member.model_dump()) await user_api_key_cache.async_set_cache( key=_cache_key, value=team_member_info, @@ -2143,6 +2152,7 @@ async def _user_api_key_auth_builder( proxy_logging_obj=proxy_logging_obj, ) except HTTPException: + token_team_models: Final = _token_team_models(valid_token) _team_obj = LiteLLM_TeamTableCachedObj( team_id=valid_token.team_id, max_budget=valid_token.team_max_budget, @@ -2151,7 +2161,7 @@ async def _user_api_key_auth_builder( tpm_limit=valid_token.team_tpm_limit, rpm_limit=valid_token.team_rpm_limit, blocked=valid_token.team_blocked, - models=valid_token.team_models, + models=token_team_models, metadata=valid_token.team_metadata, object_permission_id=valid_token.team_object_permission_id, object_permission=await _resolve_object_permission_for_unresolvable_team( @@ -2295,6 +2305,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached UserAPIKeyAuth. Only called when valid_token.team_id is known to be non-None (the caller gates on it).""" assert valid_token.team_id is not None + token_team_models: Final = _token_team_models(valid_token) return LiteLLM_TeamTableCachedObj( team_id=valid_token.team_id, max_budget=valid_token.team_max_budget, @@ -2303,7 +2314,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached tpm_limit=valid_token.team_tpm_limit, rpm_limit=valid_token.team_rpm_limit, blocked=valid_token.team_blocked, - models=valid_token.team_models, + models=token_team_models, metadata=valid_token.team_metadata, object_permission_id=valid_token.team_object_permission_id, ) diff --git a/litellm/proxy/common_utils/config_sync_pubsub.py b/litellm/proxy/common_utils/config_sync_pubsub.py index d5317fc0e02..6d781babe63 100644 --- a/litellm/proxy/common_utils/config_sync_pubsub.py +++ b/litellm/proxy/common_utils/config_sync_pubsub.py @@ -7,6 +7,7 @@ from dataclasses import asdict, dataclass from typing import TYPE_CHECKING, Final, Protocol, cast # noqa: TID251 # untyped prisma/redis boundary needs cast from litellm._logging import verbose_proxy_logger +from litellm.repositories.prisma_protocols import RowT_co, TableActions if TYPE_CHECKING: from litellm.caching.redis_cache import RedisCache @@ -163,13 +164,14 @@ class _PublishOnWriteActions: def wrap_table_actions_for_config_sync( - actions: object, + actions: "TableActions[RowT_co]", table_name: str, publish: Callable[[str], Awaitable[None]] = publish_config_change_for_object_type, -) -> object: +) -> "TableActions[RowT_co]": if table_name not in _CONFIG_SYNCED_TABLE_NAMES: return actions - return _PublishOnWriteActions(actions=actions, object_type=table_name, publish=publish) + wrapped: Final = _PublishOnWriteActions(actions=actions, object_type=table_name, publish=publish) + return cast("TableActions[RowT_co]", wrapped) # cast-ok: dynamic write-through proxy keeps the wrapped row type class ConfigSyncSubscriber: diff --git a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py index e314aec497f..58183eec689 100644 --- a/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py +++ b/litellm/proxy/common_utils/expired_ui_session_key_cleanup_manager.py @@ -4,8 +4,9 @@ Expired UI session key cleanup manager. Deletes expired virtual keys created for LiteLLM dashboard sessions. """ +from collections.abc import Sequence from datetime import datetime, timezone -from typing import Any, Final +from typing import Any, Final, Protocol from litellm._logging import verbose_proxy_logger from litellm.constants import ( @@ -14,7 +15,7 @@ from litellm.constants import ( LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, UI_SESSION_TOKEN_TEAM_ID, ) -from litellm.proxy._types import KeyRequest, LiteLLM_VerificationToken, UserAPIKeyAuth +from litellm.proxy._types import KeyRequest, UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -26,6 +27,11 @@ from litellm.repositories.verification_token_repository import ( ) +class _ExpiredSessionKeyRow(Protocol): + @property + def token(self) -> str | None: ... + + class ExpiredUISessionKeyCleanupManager: """ Cleans up expired UI session keys. @@ -138,7 +144,7 @@ class ExpiredUISessionKeyCleanupManager: return len(tokens) - async def _find_expired_ui_session_keys(self) -> list[LiteLLM_VerificationToken]: + async def _find_expired_ui_session_keys(self) -> Sequence[_ExpiredSessionKeyRow]: """ Find expired LiteLLM dashboard session keys. """ diff --git a/litellm/proxy/common_utils/key_rotation_manager.py b/litellm/proxy/common_utils/key_rotation_manager.py index 839ff28c354..352d024e20e 100644 --- a/litellm/proxy/common_utils/key_rotation_manager.py +++ b/litellm/proxy/common_utils/key_rotation_manager.py @@ -4,8 +4,9 @@ Key Rotation Manager - Automated key rotation based on rotation schedules Handles finding keys that need rotation based on their individual schedules. """ +from collections.abc import Sequence from datetime import datetime, timezone -from typing import Final +from typing import TYPE_CHECKING, Final from litellm._logging import verbose_proxy_logger from litellm.constants import ( @@ -31,6 +32,9 @@ from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) +if TYPE_CHECKING: + from prisma import models as prisma_models + class KeyRotationManager: """ @@ -106,7 +110,7 @@ class KeyRotationManager: cronjob_id=KEY_ROTATION_JOB_NAME, ) - async def _find_keys_needing_rotation(self) -> list[LiteLLM_VerificationToken]: + async def _find_keys_needing_rotation(self) -> "Sequence[prisma_models.LiteLLM_VerificationToken]": """ Find keys that are due for rotation based on their key_rotation_at timestamp. @@ -156,7 +160,7 @@ class KeyRotationManager: # Check if the rotation time has passed return now >= key.key_rotation_at - async def _rotate_key(self, key: LiteLLM_VerificationToken): + async def _rotate_key(self, key: "prisma_models.LiteLLM_VerificationToken"): """ Rotate a single key using existing regenerate_key_fn and call the rotation hook """ @@ -197,7 +201,7 @@ class KeyRotationManager: if isinstance(response, GenerateKeyResponse): await KeyManagementEventHooks.async_key_rotated_hook( data=regenerate_request, - existing_key_row=key, + existing_key_row=key, # pyright: ignore[reportArgumentType] # prisma row, hook wants the domain model response=response, user_api_key_dict=system_user, litellm_changed_by=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 8fcb184b26a..68999c93823 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -37,7 +37,7 @@ from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManage from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.organization_repository import OrganizationRepository -from litellm.repositories.prisma_protocols import ReadOnlyTable, SpendLinkedTable +from litellm.repositories.prisma_protocols import SpendLinkedTable from litellm.repositories.table_repositories import ( EndUserRepository, TagRepository, @@ -675,7 +675,7 @@ class ResetBudgetJob: rely on the default budget (litellm.max_end_user_budget_id) applied in-memory during auth checks. """ - table: Final[ReadOnlyTable] = EndUserRepository(self.prisma_client).table + table: Final = EndUserRepository(self.prisma_client).table rows: Final = await self._with_db_retry( lambda: table.find_many( where={ @@ -685,7 +685,7 @@ class ResetBudgetJob: ), reason="reset_budget_read_endusers_without_budget_id_failure", ) - return [LiteLLM_EndUserTable.model_validate(row.dict()) for row in rows] + return [LiteLLM_EndUserTable.model_validate(row.model_dump()) for row in rows] async def _write_key_reset_updates(self, updated_keys: list[LiteLLM_VerificationToken]) -> None: """ diff --git a/litellm/proxy/container_endpoints/ownership.py b/litellm/proxy/container_endpoints/ownership.py index a559ab49cfa..e3088771c82 100644 --- a/litellm/proxy/container_endpoints/ownership.py +++ b/litellm/proxy/container_endpoints/ownership.py @@ -1,7 +1,7 @@ import json from collections.abc import Mapping, Sequence from collections.abc import Set as AbstractSet -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final from fastapi import HTTPException @@ -18,28 +18,11 @@ from litellm.repositories.table_repositories import ManagedObjectRepository from litellm.responses.utils import ResponsesAPIRequestUtils if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.proxy.utils import PrismaClient -class _ManagedObjectRow(Protocol): - model_object_id: str - unified_object_id: str | None - file_purpose: str | None - created_by: str | None - - -class _ManagedObjectTable(Protocol): - async def find_unique(self, *, where: Mapping[str, str]) -> _ManagedObjectRow | None: ... - - async def find_first(self, *, where: Mapping[str, str]) -> _ManagedObjectRow | None: ... - - async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_ManagedObjectRow]: ... - - async def create(self, *, data: Mapping[str, str]) -> _ManagedObjectRow: ... - - async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> _ManagedObjectRow | None: ... - - CONTAINER_OBJECT_PURPOSE: Final = "container" # 60s LRU/TTL cache absorbs every container access check before it reaches @@ -220,7 +203,7 @@ async def record_container_owner( verbose_proxy_logger.warning("Skipping container ownership tracking because prisma_client is None") return response - table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table + table: Final = ManagedObjectRepository(prisma_client).table existing: Final = await table.find_unique(where={"model_object_id": model_object_id}) if existing is not None: if getattr(existing, "file_purpose", None) != CONTAINER_OBJECT_PURPOSE: @@ -272,8 +255,8 @@ async def _get_container_owner(original_container_id: str, custom_llm_provider: if prisma_client is None: return None - table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table - row: Final[_ManagedObjectRow | None] = await table.find_first( + table: Final = ManagedObjectRepository(prisma_client).table + row: Final[prisma_models.LiteLLM_ManagedObjectTable | None] = await table.find_first( where={ "model_object_id": model_object_id, "file_purpose": CONTAINER_OBJECT_PURPOSE, @@ -309,8 +292,8 @@ async def _get_stored_container_id(original_container_id: str, custom_llm_provid if prisma_client is None: return None - table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table - row: Final[_ManagedObjectRow | None] = await table.find_first( + table: Final = ManagedObjectRepository(prisma_client).table + row: Final[prisma_models.LiteLLM_ManagedObjectTable | None] = await table.find_first( where={ "model_object_id": model_object_id, "file_purpose": CONTAINER_OBJECT_PURPOSE, @@ -394,8 +377,8 @@ async def _get_allowed_container_ids( if prisma_client is None: return set() - table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table - rows: Final[Sequence[_ManagedObjectRow]] = await table.find_many( + table: Final = ManagedObjectRepository(prisma_client).table + rows: Final[Sequence[prisma_models.LiteLLM_ManagedObjectTable]] = await table.find_many( where={ "file_purpose": CONTAINER_OBJECT_PURPOSE, "created_by": {"in": owner_scopes}, diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index 3b3e9692eda..dc193cb8523 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -2,7 +2,10 @@ CRUD endpoints for storing reusable credentials. """ -from typing import Final +from typing import ( + Final, + cast, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict +) from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response @@ -88,7 +91,9 @@ async def create_credential( ) encrypted_credential: Final = CredentialHelperUtils.encrypt_credential_values(processed_credential) credentials_dict: Final = encrypted_credential.model_dump() - credentials_dict_jsonified: Final = jsonify_object(credentials_dict) + credentials_dict_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str + "dict[str, object]", jsonify_object(credentials_dict) + ) await CredentialsRepository(prisma_client).create( data={ **credentials_dict_jsonified, @@ -310,7 +315,9 @@ async def update_credential( if db_credential is None: raise HTTPException(status_code=404, detail="Credential not found in DB.") merged_credential: Final = update_db_credential(db_credential, credential) - credential_object_jsonified: Final = jsonify_object(merged_credential.model_dump()) + credential_object_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str + "dict[str, object]", jsonify_object(merged_credential.model_dump()) + ) await credentials_repository.update_by_name( credential_name, data={ diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 187a18be845..367552e783e 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -6,14 +6,15 @@ Admins use the management endpoints to read and update input_policy / output_pol """ import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Mapping from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar +from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ToolDiscoveryQueueItem from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.repositories.object_permission_repository import ObjectPermissionRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ToolRepository from litellm.types.tool_management import ( LiteLLM_ToolTableRow, @@ -25,33 +26,16 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -_RowT_co: Final = TypeVar("_RowT_co", covariant=True) - -class _TableActions(Protocol[_RowT_co]): - async def find_unique(self, where: Mapping[str, object]) -> _RowT_co | None: ... - - async def find_many( - self, - where: Mapping[str, object] | None = None, - order: Mapping[str, object] | None = None, - include: Mapping[str, object] | None = None, - ) -> Sequence[_RowT_co]: ... - - async def upsert(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT_co: ... - - async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _RowT_co | None: ... - - -def _tool_table_actions(prisma_client: "PrismaClient") -> "_TableActions[prisma_db_models.LiteLLM_ToolTable]": - table: Final[_TableActions[prisma_db_models.LiteLLM_ToolTable]] = ToolRepository(prisma_client).table +def _tool_table_actions(prisma_client: "PrismaClient") -> "TableActions[prisma_db_models.LiteLLM_ToolTable]": + table: Final[TableActions[prisma_db_models.LiteLLM_ToolTable]] = ToolRepository(prisma_client).table return table def _object_permission_table_actions( prisma_client: "PrismaClient", -) -> "_TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]": - table: Final[_TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]] = ObjectPermissionRepository( +) -> "TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]": + table: Final[TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]] = ObjectPermissionRepository( prisma_client ).table return table diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index e50a3a5a1e7..20efbe06ecc 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -29,6 +29,7 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository from litellm.types.guardrails import ( PII_ENTITY_CATEGORIES_MAP, @@ -65,29 +66,12 @@ router: Final = APIRouter() GUARDRAIL_REGISTRY: Final = GuardrailRegistry() -class _GuardrailsTableActions(Protocol): - async def create(self, data: Mapping[str, object]) -> "LiteLLM_GuardrailsTable": ... - - async def delete(self, where: Mapping[str, object]) -> "LiteLLM_GuardrailsTable | None": ... - - async def find_unique(self, where: Mapping[str, object]) -> "LiteLLM_GuardrailsTable | None": ... - - async def find_many( - self, where: Mapping[str, object], order: Mapping[str, str] - ) -> "Sequence[LiteLLM_GuardrailsTable]": ... - - async def update( - self, where: Mapping[str, object], data: Mapping[str, object] - ) -> "LiteLLM_GuardrailsTable | None": ... - - def _as_str_object_mapping(mapping: Mapping[str, object]) -> Mapping[str, object]: return mapping -def _guardrails_table(prisma_client: "PrismaClient") -> _GuardrailsTableActions: - table: Final[_GuardrailsTableActions] = GuardrailsRepository(prisma_client).table - return table +def _guardrails_table(prisma_client: "PrismaClient") -> "TableActions[LiteLLM_GuardrailsTable]": + return GuardrailsRepository(prisma_client).table async def _create_guardrail_row(prisma_client: "PrismaClient", data: Mapping[str, object]) -> "LiteLLM_GuardrailsTable": diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index d29ec555a80..987e7d778c7 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -3,12 +3,12 @@ import asyncio import importlib import os -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping from datetime import datetime, timezone from itertools import chain, count -from typing import Final, Literal, Optional, Protocol, cast +from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, cast -from pydantic import BaseModel, ValidationError +from pydantic import ValidationError import litellm from litellm import Router @@ -39,6 +39,7 @@ from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( ) from litellm.proxy.types_utils.utils import get_instance_fn from litellm.proxy.utils import PrismaClient +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import GuardrailsRepository from litellm.secret_managers.main import get_secret from litellm.types.guardrails import ( @@ -61,6 +62,9 @@ from .guardrail_initializers import ( initialize_tool_permission, ) +if TYPE_CHECKING: + from prisma import models as prisma_models + class _GuardrailRowLike(Protocol): @property @@ -68,15 +72,7 @@ class _GuardrailRowLike(Protocol): def __iter__(self) -> Iterator[tuple[str, object]]: ... -class _GuardrailTableActions(Protocol): - async def create(self, *, data: Mapping[str, object]) -> _GuardrailRowLike: ... - async def delete(self, *, where: Mapping[str, str]) -> object: ... - async def update(self, *, where: Mapping[str, str], data: Mapping[str, object]) -> _GuardrailRowLike: ... - async def find_many(self, *, where: Mapping[str, str], order: Mapping[str, str]) -> Sequence[BaseModel]: ... - async def find_unique(self, *, where: Mapping[str, str]) -> BaseModel | None: ... - - -def _guardrail_table(prisma_client: PrismaClient) -> _GuardrailTableActions: +def _guardrail_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_GuardrailsTable]": """Typed view of the guardrails table actions exposed by the Prisma repository.""" return GuardrailsRepository(prisma_client).table @@ -347,7 +343,7 @@ class GuardrailRegistry: guardrail_info: Final[str] = safe_dumps(guardrail.get("guardrail_info", {})) # Update in DB - updated_guardrail: Final[_GuardrailRowLike] = await _guardrail_table(prisma_client).update( + updated_guardrail: Final[_GuardrailRowLike | None] = await _guardrail_table(prisma_client).update( where={"guardrail_id": guardrail_id}, data={ "guardrail_name": guardrail_name, @@ -356,6 +352,8 @@ class GuardrailRegistry: "updated_at": datetime.now(timezone.utc), }, ) + if updated_guardrail is None: + raise ValueError(f"Guardrail not found, passed guardrail_id={guardrail_id}") # Convert to dict and return return dict(updated_guardrail) diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 9d0d84dc2b1..7a0edbddca8 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -17,6 +17,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DailyGuardrailMetricsRepository, DailyGuardrailUsageUnitsRepository, @@ -30,13 +31,6 @@ from litellm.repositories.table_repositories import ( if TYPE_CHECKING: from prisma import models as prisma_models from prisma import types as prisma_types - from prisma.actions import ( - LiteLLM_DailyGuardrailMetricsActions, - LiteLLM_DailyGuardrailUsageUnitsActions, - LiteLLM_DailyPolicyMetricsActions, - LiteLLM_GuardrailsTableActions, - LiteLLM_PolicyTableActions, - ) from litellm.proxy.utils import PrismaClient from litellm.types.guardrails import Guardrail @@ -85,8 +79,8 @@ def _resolve_usage_window(start_date: str | None, end_date: str | None) -> tuple def _guardrails_table( prisma_client: "PrismaClient", -) -> "LiteLLM_GuardrailsTableActions[prisma_models.LiteLLM_GuardrailsTable]": - guardrails_table: LiteLLM_GuardrailsTableActions[prisma_models.LiteLLM_GuardrailsTable] = GuardrailsRepository( +) -> "TableActions[prisma_models.LiteLLM_GuardrailsTable]": + guardrails_table: Final[TableActions[prisma_models.LiteLLM_GuardrailsTable]] = GuardrailsRepository( prisma_client ).table return guardrails_table @@ -94,28 +88,26 @@ def _guardrails_table( def _policies_table( prisma_client: "PrismaClient", -) -> "LiteLLM_PolicyTableActions[prisma_models.LiteLLM_PolicyTable]": - policies_table: Final[LiteLLM_PolicyTableActions[prisma_models.LiteLLM_PolicyTable]] = PolicyRepository( - prisma_client - ).table +) -> "TableActions[prisma_models.LiteLLM_PolicyTable]": + policies_table: Final[TableActions[prisma_models.LiteLLM_PolicyTable]] = PolicyRepository(prisma_client).table return policies_table def _daily_guardrail_metrics_table( prisma_client: "PrismaClient", -) -> "LiteLLM_DailyGuardrailMetricsActions[prisma_models.LiteLLM_DailyGuardrailMetrics]": - metrics_table: Final[LiteLLM_DailyGuardrailMetricsActions[prisma_models.LiteLLM_DailyGuardrailMetrics]] = ( - DailyGuardrailMetricsRepository(prisma_client).table - ) +) -> "TableActions[prisma_models.LiteLLM_DailyGuardrailMetrics]": + metrics_table: Final[TableActions[prisma_models.LiteLLM_DailyGuardrailMetrics]] = DailyGuardrailMetricsRepository( + prisma_client + ).table return metrics_table def _daily_policy_metrics_table( prisma_client: "PrismaClient", -) -> "LiteLLM_DailyPolicyMetricsActions[prisma_models.LiteLLM_DailyPolicyMetrics]": - metrics_table: Final[LiteLLM_DailyPolicyMetricsActions[prisma_models.LiteLLM_DailyPolicyMetrics]] = ( - DailyPolicyMetricsRepository(prisma_client).table - ) +) -> "TableActions[prisma_models.LiteLLM_DailyPolicyMetrics]": + metrics_table: Final[TableActions[prisma_models.LiteLLM_DailyPolicyMetrics]] = DailyPolicyMetricsRepository( + prisma_client + ).table return metrics_table @@ -135,8 +127,8 @@ async def _find_daily_policy_metrics( def _daily_guardrail_usage_units_table( prisma_client: "PrismaClient", -) -> "LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]": - units_table: Final[LiteLLM_DailyGuardrailUsageUnitsActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]] = ( +) -> "TableActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]": + units_table: Final[TableActions[prisma_models.LiteLLM_DailyGuardrailUsageUnits]] = ( DailyGuardrailUsageUnitsRepository(prisma_client).table ) return units_table diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index 820f6438aaf..b8ae09afc00 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -14,6 +14,8 @@ from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NamedTuple, TypeVar +from typing_extensions import ReadOnly, TypedDict + from litellm._logging import verbose_proxy_logger from litellm.proxy._types import DB_RETRY_SAFE_ERROR_TYPES from litellm.proxy.utils import PrismaClient @@ -47,6 +49,18 @@ class _MetricsKey(NamedTuple): date: str +class _UsageUnitCompoundKey(TypedDict): + guardrail_id: ReadOnly[str] + date: ReadOnly[str] + team_id: ReadOnly[str] + api_key: ReadOnly[str] + usage_unit: ReadOnly[str] + + +class _UsageUnitWhereUnique(TypedDict): + guardrail_id_date_team_id_api_key_usage_unit: ReadOnly[_UsageUnitCompoundKey] + + class PendingRollups: """Rollup rows whose connection-error retries exhausted, held for the next flush.""" @@ -229,7 +243,7 @@ async def _upsert_usage_unit_row(prisma_client: PrismaClient, key: _UsageUnitKey "usage_unit": key.usage_unit, "units": units, } - where: Final[prisma_types.LiteLLM_DailyGuardrailUsageUnitsWhereUniqueInput] = { + where: Final[_UsageUnitWhereUnique] = { "guardrail_id_date_team_id_api_key_usage_unit": { "guardrail_id": key.guardrail_id, "date": key.date, diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 2271501d480..b12b764d144 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -390,7 +390,7 @@ async def list_access_groups( _require_admin_view(user_api_key_dict) prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value) - table: Final[_AccessGroupTable] = AccessGroupRepository(prisma_client).table + table: Final = AccessGroupRepository(prisma_client).table records: Final = await table.find_many(order={"created_at": "desc"}) return [_record_to_response(r) for r in records] @@ -406,7 +406,7 @@ async def get_access_group( _require_admin_view(user_api_key_dict) prisma_client: Final = get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value) - table: Final[_AccessGroupTable] = AccessGroupRepository(prisma_client).table + table: Final = AccessGroupRepository(prisma_client).table record: Final = await table.find_unique(where={"access_group_id": access_group_id}) if record is None: raise HTTPException( diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 8c6195388c5..17a845d4300 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -93,7 +93,7 @@ async def new_budget( budget_obj.budget_reset_at = get_budget_reset_time(budget_duration=budget_obj.budget_duration) budget_obj_json: Final = budget_obj.model_dump(exclude_none=True) - budget_obj_jsonified: Final = jsonify_object(budget_obj_json) # json dump any dictionaries + budget_obj_jsonified: Final[dict[str, object]] = jsonify_object(budget_obj_json) # mutable-ok: prisma create input try: response: Final = await BudgetRepository(prisma_client).table.create( data={ diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 385073edc90..77ff77c9a88 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -43,7 +43,8 @@ router: Final = APIRouter() class _CacheConfigRow(Protocol): - cache_settings: str | Mapping[str, object] | None + @property + def cache_settings(self) -> str | Mapping[str, object] | None: ... class _CacheConfigTable(Protocol): diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 3d2fa798e03..d3968bf323b 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -441,7 +441,7 @@ async def get_api_key_metadata( This ensures that key_alias and team_id are preserved in historical activity logs even after a key is deleted or regenerated. """ - key_records: list[PrismaVerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + key_records: Sequence[PrismaVerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( where={"token": {"in": list(api_keys)}} ) result: Final[dict[str, _KeyMetadataDict]] = { @@ -452,9 +452,9 @@ async def get_api_key_metadata( missing_keys: Final = api_keys - set(result.keys()) if missing_keys: try: - deleted_key_records: Final[list[PrismaDeletedVerificationToken]] = await DeletedVerificationTokenRepository( - prisma_client - ).table.find_many( + deleted_key_records: Final[ + Sequence[PrismaDeletedVerificationToken] + ] = await DeletedVerificationTokenRepository(prisma_client).table.find_many( where={"token": {"in": list(missing_keys)}}, order={"deleted_at": "desc"}, ) diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index dde0751d98d..7f4faddf178 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -46,7 +46,8 @@ router: Final = APIRouter() class _ConfigOverrideRow(Protocol): - config_value: str | Mapping[str, object] | None + @property + def config_value(self) -> str | Mapping[str, object] | None: ... class _ConfigOverridesTableClient(Protocol): diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index c2f5b8eeb8b..9a98bdbb6b1 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -15,9 +15,9 @@ These are members of a Team on LiteLLM import asyncio import json import traceback -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Mapping, Sequence from datetime import datetime, timezone -from typing import Any, Final, Literal, Protocol, cast +from typing import Any, Final, Literal, cast import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status @@ -58,6 +58,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( from litellm.proxy.management_helpers.utils import management_endpoint_wrapper from litellm.proxy.utils import handle_exception_on_proxy, hash_password from litellm.repositories.organization_repository import OrganizationRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( InvitationLinkRepository, OrganizationMembershipRepository, @@ -86,15 +87,6 @@ from litellm.types.proxy.management_endpoints.scim_v2 import ( if TYPE_CHECKING: from prisma import models as prisma_models from prisma import types as prisma_types - from prisma.actions import ( - LiteLLM_InvitationLinkActions, - LiteLLM_OrganizationMembershipActions, - LiteLLM_OrganizationTableActions, - LiteLLM_TeamMembershipActions, - LiteLLM_TeamTableActions, - LiteLLM_UserTableActions, - LiteLLM_VerificationTokenActions, - ) from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.proxy_server import PrismaClient @@ -105,31 +97,31 @@ router: Final = APIRouter() def _user_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_UserTableActions[prisma_models.LiteLLM_UserTable]": - user_table: Final[LiteLLM_UserTableActions[prisma_models.LiteLLM_UserTable]] = UserRepository(prisma_client).table +) -> "TableActions[prisma_models.LiteLLM_UserTable]": + user_table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = UserRepository(prisma_client).table return user_table def _team_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]": - team_table: Final[LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table +) -> "TableActions[prisma_models.LiteLLM_TeamTable]": + team_table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = TeamRepository(prisma_client).table return team_table def _verification_token_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken]": - token_table: Final[LiteLLM_VerificationTokenActions[prisma_models.LiteLLM_VerificationToken]] = ( - VerificationTokenRepository(prisma_client).table - ) +) -> "TableActions[prisma_models.LiteLLM_VerificationToken]": + token_table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = VerificationTokenRepository( + prisma_client + ).table return token_table def _organization_membership_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_OrganizationMembershipActions[prisma_models.LiteLLM_OrganizationMembership]": - membership_table: Final[LiteLLM_OrganizationMembershipActions[prisma_models.LiteLLM_OrganizationMembership]] = ( +) -> "TableActions[prisma_models.LiteLLM_OrganizationMembership]": + membership_table: Final[TableActions[prisma_models.LiteLLM_OrganizationMembership]] = ( OrganizationMembershipRepository(prisma_client).table ) return membership_table @@ -137,8 +129,8 @@ def _organization_membership_table( def _invitation_link_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_InvitationLinkActions[prisma_models.LiteLLM_InvitationLink]": - invitation_table: LiteLLM_InvitationLinkActions[prisma_models.LiteLLM_InvitationLink] = InvitationLinkRepository( +) -> "TableActions[prisma_models.LiteLLM_InvitationLink]": + invitation_table: Final[TableActions[prisma_models.LiteLLM_InvitationLink]] = InvitationLinkRepository( prisma_client ).table return invitation_table @@ -146,19 +138,19 @@ def _invitation_link_table( def _organization_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_OrganizationTableActions[prisma_models.LiteLLM_OrganizationTable]": - organization_table: Final[LiteLLM_OrganizationTableActions[prisma_models.LiteLLM_OrganizationTable]] = ( - OrganizationRepository(prisma_client).table - ) +) -> "TableActions[prisma_models.LiteLLM_OrganizationTable]": + organization_table: Final[TableActions[prisma_models.LiteLLM_OrganizationTable]] = OrganizationRepository( + prisma_client + ).table return organization_table def _team_membership_table( prisma_client: "PrismaClient | None", -) -> "LiteLLM_TeamMembershipActions[prisma_models.LiteLLM_TeamMembership]": - team_membership_table: Final[LiteLLM_TeamMembershipActions[prisma_models.LiteLLM_TeamMembership]] = ( - TeamMembershipRepository(prisma_client).table - ) +) -> "TableActions[prisma_models.LiteLLM_TeamMembership]": + team_membership_table: Final[TableActions[prisma_models.LiteLLM_TeamMembership]] = TeamMembershipRepository( + prisma_client + ).table return team_membership_table @@ -294,7 +286,7 @@ async def _add_user_to_organizations( organization_member_add, ) - tasks: Final = [] + tasks: Final[list[Awaitable[object]]] = [] for organization_id in organizations: tasks.append( organization_member_add( @@ -406,7 +398,7 @@ async def add_new_user_to_default_team( teams: list[str] | list[NewUserRequestTeam], prisma_client: "PrismaClient", ): - tasks: Final = [] + tasks: Final[list[Awaitable[object]]] = [] for team in teams: user_role: Literal["user", "admin"] = "user" max_budget_in_team: float | None = None @@ -1479,7 +1471,8 @@ async def _update_single_user_helper( # Create new user if not found non_default_values["user_id"] = str(uuid.uuid4()) non_default_values["user_email"] = user_request.user_email - response = await prisma_client.insert_data(data=non_default_values, table_name="user") + inserted_user_row: Final = await prisma_client.insert_data(data=non_default_values, table_name="user") + response = inserted_user_row # pyright: ignore[reportAssignmentType] # insert_data returns a prisma row if response is not None: await _schedule_user_update_audit_log( @@ -1795,7 +1788,9 @@ async def bulk_user_update( # Apply update transformations (reuse existing logic) data_json: Final[dict] = data.user_updates.model_dump(exclude_unset=True) - non_default_values: Final = _update_internal_user_params(data_json=data_json, data=data.user_updates) + non_default_values: Final[dict[str, object]] = _update_internal_user_params( + data_json=data_json, data=data.user_updates + ) # Remove user identification fields since we're updating by user_id non_default_values.pop("user_id", None) @@ -2149,7 +2144,7 @@ async def get_users( _validate_sort_params(sort_by, sort_order) if sort_by is not None and isinstance(sort_by, str) else None ) - users: Sequence[prisma_models.LiteLLM_UserTable] | None = await UserRepository(prisma_client).table.find_many( + users: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await UserRepository(prisma_client).table.find_many( where=where_conditions, skip=skip, take=page_size, @@ -2160,10 +2155,7 @@ async def get_users( total_count: Final[int] = await UserRepository(prisma_client).table.count(where=where_conditions) # Get key count for each user - if users is not None: - user_key_counts = await get_user_key_counts(prisma_client, [user.user_id for user in users]) - else: - user_key_counts = {} + user_key_counts: Final = await get_user_key_counts(prisma_client, [user.user_id for user in users]) verbose_proxy_logger.debug("Total count of users: %s", total_count) @@ -2172,17 +2164,14 @@ async def get_users( # Prepare response user_list: list[LiteLLM_UserTableWithKeyCount] = [] - if users is not None: - for user in users: - user_dump = user.model_dump() - user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) - user_list.append( - LiteLLM_UserTableWithKeyCount.model_validate( - {**user_dump, "key_count": user_key_counts.get(user.user_id, 0)} - ) + for user in users: + user_dump = user.model_dump() + user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) + user_list.append( + LiteLLM_UserTableWithKeyCount.model_validate( + {**user_dump, "key_count": user_key_counts.get(user.user_id, 0)} ) - else: - user_list = [] + ) return { "users": user_list, @@ -2193,13 +2182,6 @@ async def get_users( } -class _DeleteTeamRow(Protocol): - team_id: str - members_with_roles: object - - def model_dump(self) -> Mapping[str, object]: ... - - @router.post( "/user/delete", tags=["Internal User management"], @@ -2258,9 +2240,9 @@ async def delete_user( # loop an org-admin of org-A could delete users in org-B by supplying # {"user_ids": [victim_in_org_B], "organization_id": "org-A"}. caller_is_proxy_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value - caller_admin_org_ids: set = set() + caller_admin_org_ids: set[str] = set() if not caller_is_proxy_admin: - caller_memberships: Final = ( + caller_memberships: Final[Sequence[prisma_models.LiteLLM_OrganizationMembership]] = ( await _organization_membership_table(prisma_client).find_many( where={ "user_id": user_api_key_dict.user_id, @@ -2279,7 +2261,7 @@ async def delete_user( # Batch-fetch target memberships once before the per-user loop. Avoids # an N+1 DB call when delete_user is called with a large user_ids list. - target_org_ids_by_user: Final[dict[str, set]] = {} + target_org_ids_by_user: Final[dict[str, set[str]]] = {} if not caller_is_proxy_admin: all_target_memberships: Final = await _organization_membership_table(prisma_client).find_many( where={"user_id": {"in": data.user_ids}} @@ -2319,7 +2301,7 @@ async def delete_user( # we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes if is_audit_logging_enabled(): # make an audit log for each team deleted - _user_row = user_row.json(exclude_none=True) + _user_row = user_row.model_dump_json(exclude_none=True) asyncio.create_task( create_audit_log_for_update( @@ -2342,10 +2324,10 @@ async def delete_user( ) ## CLEANUP MEMBERS_WITH_ROLES - fetch_all_teams: Sequence[_DeleteTeamRow] = await TeamRepository(prisma_client).table.find_many( - where={"team_id": {"in": user_row.teams}} - ) - teams_to_update = [] + fetch_all_teams: Sequence[prisma_models.LiteLLM_TeamTable] = await TeamRepository( + prisma_client + ).table.find_many(where={"team_id": {"in": user_row.teams}}) + teams_to_update: list[tuple[str, str]] = [] for team in fetch_all_teams: removed_team_members, new_team_members = _cleanup_members_with_roles( existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), @@ -2357,15 +2339,14 @@ async def delete_user( ) if removed_team_members: _db_new_team_members: list[dict] = [m.model_dump() for m in new_team_members] - team.members_with_roles = json.dumps(_db_new_team_members) - teams_to_update.append(team) + teams_to_update.append((team.team_id, json.dumps(_db_new_team_members))) ## update teams - for team in teams_to_update: + for team_id, members_with_roles in teams_to_update: await TeamRepository(prisma_client).table.update( - where={"team_id": team.team_id}, - data={"members_with_roles": team.members_with_roles}, + where={"team_id": team_id}, + data={"members_with_roles": members_with_roles}, ) # End of Audit logging diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index 41e52f05c01..9f561eadfbd 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -122,6 +122,9 @@ async def update_jwt_key_mapping( where={"id": data.id}, data=update_data ) + if updated_mapping is None: + raise HTTPException(status_code=404, detail="Mapping not found") + # Invalidate new cache key if claim fields changed cache_key = f"jwt_key_mapping:{updated_mapping.jwt_claim_name}:{updated_mapping.jwt_claim_value}" await user_api_key_cache.async_delete_cache(cache_key) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 54f567b7aa2..f6a4448e7c9 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -123,6 +123,7 @@ from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigParam, ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository from litellm.repositories.model_repository import ModelRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DeletedVerificationTokenRepository, DeprecatedVerificationTokenRepository, @@ -151,65 +152,22 @@ from litellm.types.utils import ( if TYPE_CHECKING: from prisma import Prisma + from prisma import models as prisma_models -_PrismaRowT = TypeVar("_PrismaRowT") _RepositoryModelT = TypeVar("_RepositoryModelT", bound=BaseModel) -class _PrismaTableActions(Protocol[_PrismaRowT]): - """Typed view of the Prisma table actions a repository exposes through its untyped ``table``.""" - - async def find_unique( - self, - *, - where: Mapping[str, object], - include: Mapping[str, object] | None = None, - ) -> _PrismaRowT | None: ... - - async def find_first( - self, - *, - where: Mapping[str, object], - include: Mapping[str, object] | None = None, - ) -> _PrismaRowT | None: ... - - async def find_many( - self, - *, - where: Mapping[str, object] | None = None, - include: Mapping[str, object] | None = None, - order: Mapping[str, object] | None = None, - skip: int | None = None, - take: int | None = None, - ) -> list[_PrismaRowT]: ... - - async def count(self, *, where: Mapping[str, object] | None = None) -> int: ... - - async def create(self, *, data: Mapping[str, object]) -> _PrismaRowT: ... - - async def create_many(self, *, data: Sequence[Mapping[str, object]]) -> int: ... - - async def delete_many(self, *, where: Mapping[str, object] | None = None) -> int: ... - - async def update( - self, - *, - where: Mapping[str, object], - data: Mapping[str, object], - ) -> _PrismaRowT | None: ... - - async def upsert( - self, - *, - where: Mapping[str, object], - data: Mapping[str, object], - ) -> _PrismaRowT: ... - - class _UserRowLike(Protocol): - user_id: str | None - user_email: str | None - user_alias: str | None + """Read-only view of the user columns ``/key/list`` expands keys with.""" + + @property + def user_id(self) -> str | None: ... + + @property + def user_email(self) -> str | None: ... + + @property + def user_alias(self) -> str | None: ... def model_dump(self) -> Mapping[str, object]: ... @@ -217,46 +175,56 @@ class _UserRowLike(Protocol): class _TxTables(Protocol): - litellm_proxymodeltable: _PrismaTableActions[object] + litellm_proxymodeltable: TableActions[object] -class _TableSource(Protocol[_PrismaRowT]): - """Repository view that exposes its untyped Prisma ``table`` with a concrete row type.""" +class _ConfigTableActions(Protocol): + """Config table surface this module needs; the shared repository seam exposes no ``update``.""" - @property - def table(self) -> _PrismaTableActions[_PrismaRowT]: ... + async def find_many(self) -> Sequence[ConfigParam]: ... - -def _table_of(source: _TableSource[_PrismaRowT]) -> _PrismaTableActions[_PrismaRowT]: - return source.table + async def update( + self, + *, + where: Mapping[str, object], + data: Mapping[str, object], + ) -> ConfigParam | None: ... def _prisma_table( repository: BaseRepository[_RepositoryModelT], -) -> _PrismaTableActions[_RepositoryModelT]: - return _table_of(repository) +) -> TableActions[_RepositoryModelT]: + return cast( # cast-ok: callers read only the field names the prisma row and repository model share + "TableActions[_RepositoryModelT]", repository.table + ) def _deleted_verification_token_table( prisma_client: PrismaClient, -) -> _PrismaTableActions[LiteLLM_DeletedVerificationToken]: - return _table_of(DeletedVerificationTokenRepository(prisma_client)) +) -> "TableActions[prisma_models.LiteLLM_DeletedVerificationToken]": + return DeletedVerificationTokenRepository(prisma_client).table -def _deprecated_verification_token_table(prisma_client: PrismaClient) -> _PrismaTableActions[object]: - return _table_of(DeprecatedVerificationTokenRepository(prisma_client)) +def _deprecated_verification_token_table( + prisma_client: PrismaClient, +) -> "TableActions[prisma_models.LiteLLM_DeprecatedVerificationToken]": + return DeprecatedVerificationTokenRepository(prisma_client).table -def _user_table(prisma_client: PrismaClient) -> _PrismaTableActions[_UserRowLike]: - return _table_of(UserRepository(prisma_client)) +def _user_table(prisma_client: PrismaClient) -> TableActions[_UserRowLike]: + return UserRepository(prisma_client).table -def _credentials_table(prisma_client: PrismaClient) -> _PrismaTableActions[CredentialItem]: - return _table_of(CredentialsRepository(prisma_client)) +def _credentials_table(prisma_client: PrismaClient) -> TableActions[CredentialItem]: + return cast( # cast-ok: the rotation loop reads and rewrites these rows through CredentialItem names only + "TableActions[CredentialItem]", CredentialsRepository(prisma_client).table + ) -def _config_table(prisma_client: PrismaClient) -> _PrismaTableActions[ConfigParam]: - return _table_of(ConfigRepository(prisma_client)) +def _config_table(prisma_client: PrismaClient) -> _ConfigTableActions: + return cast( # cast-ok: ConfigRepository.table hides the write actions this module needs on that same object + "_ConfigTableActions", ConfigRepository(prisma_client).table + ) async def _check_custom_key_allowed(custom_key_value: str | None) -> None: @@ -1046,7 +1014,7 @@ async def _common_key_generation_helper( ) new_budget: Final = prisma_client.jsonify_object(budget_row.json(exclude_none=True)) - _budget: Final[LiteLLM_BudgetTable] = await BudgetRepository(prisma_client).table.create( + _budget: Final[prisma_models.LiteLLM_BudgetTable] = await BudgetRepository(prisma_client).table.create( data={ **new_budget, "created_by": user_api_key_dict.user_id or litellm_proxy_admin_name, @@ -1252,7 +1220,7 @@ async def _common_key_generation_helper( def _check_key_model_specific_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], data: GenerateKeyRequest | UpdateKeyRequest, entity_rpm_limit: int | None, entity_tpm_limit: int | None, @@ -1323,7 +1291,7 @@ def _check_key_model_specific_limits( def _check_key_rpm_tpm_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], data: GenerateKeyRequest | UpdateKeyRequest, entity_rpm_limit: int | None, entity_tpm_limit: int | None, @@ -1361,7 +1329,7 @@ def _check_key_rpm_tpm_limits( def check_team_key_model_specific_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], team_table: LiteLLM_TeamTableCachedObj, data: GenerateKeyRequest | UpdateKeyRequest, ) -> None: @@ -1386,7 +1354,7 @@ def check_team_key_model_specific_limits( def check_team_key_rpm_tpm_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], team_table: LiteLLM_TeamTableCachedObj, data: GenerateKeyRequest | UpdateKeyRequest, ) -> None: @@ -1494,7 +1462,7 @@ async def _check_project_key_limits( def check_org_key_model_specific_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], org_table: LiteLLM_OrganizationTable, data: GenerateKeyRequest | UpdateKeyRequest, ) -> None: @@ -1527,7 +1495,7 @@ def check_org_key_model_specific_limits( def check_org_key_rpm_tpm_limits( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], org_table: LiteLLM_OrganizationTable, data: GenerateKeyRequest | UpdateKeyRequest, ) -> None: @@ -2242,9 +2210,9 @@ async def _get_and_validate_existing_key( code=status.HTTP_400_BAD_REQUEST, ) - rows: list[LiteLLM_VerificationToken] = await _prisma_table(VerificationTokenRepository(prisma_client)).find_many( - where={"key_alias": key_alias}, take=2 - ) + rows: Sequence[LiteLLM_VerificationToken] = await _prisma_table( + VerificationTokenRepository(prisma_client) + ).find_many(where={"key_alias": key_alias}, take=2) if len(rows) == 0: raise ProxyException( @@ -2407,7 +2375,10 @@ async def _process_single_key_update( ) _data: Final = {**non_default_values, "token": update_key_request.key} - response: Final = await prisma_client.update_data(token=update_key_request.key, data=_data) + response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict + "Mapping[str, object] | None", + await prisma_client.update_data(token=update_key_request.key, data=_data), + ) # Delete cache await _delete_cache_key_object( @@ -3225,7 +3196,7 @@ async def bulk_update_team_keys( # `blocked` is Boolean? with no default; `/key/generate` writes NULL. Prisma's `NOT` # excludes NULLs, so explicitly OR `false` with `null` to include them. now: Final = datetime.now(timezone.utc) - existing_keys = await VerificationTokenRepository(prisma_client).table.find_many( + existing_keys = await _prisma_table(VerificationTokenRepository(prisma_client)).find_many( where={ "team_id": data.team_id, "AND": [ @@ -3243,7 +3214,9 @@ async def bulk_update_team_keys( "error": f"Team {data.team_id} has more than {MAX_BATCH_SIZE} keys. Use `key_ids` to update in batches of {MAX_BATCH_SIZE}." }, ) - requested_tokens = [row.token for row in existing_keys] + requested_tokens = cast( # cast-ok: token is the table's primary key, so a row read back always carries one + "list[str]", [row.token for row in existing_keys] + ) else: if data.key_ids is None or len(data.key_ids) == 0: raise HTTPException( @@ -3261,7 +3234,7 @@ async def bulk_update_team_keys( seen_hashes.add(h) requested_tokens.append(k) hashed_key_ids.append(h) - existing_keys = await VerificationTokenRepository(prisma_client).table.find_many( + existing_keys = await _prisma_table(VerificationTokenRepository(prisma_client)).find_many( where={"team_id": data.team_id, "token": {"in": hashed_key_ids}} ) @@ -3698,7 +3671,7 @@ async def info_key_fn( hashed_key: str | None = key if key is not None: hashed_key = _hash_token_if_needed(token=key) - key_info = await VerificationTokenRepository(prisma_client).table.find_unique( + key_info = await _prisma_table(VerificationTokenRepository(prisma_client)).find_unique( where={"token": hashed_key}, include={"litellm_budget_table": True}, ) @@ -3727,7 +3700,7 @@ async def info_key_fn( key_info = key_info.model_dump() except Exception: # if using pydantic v1 - key_info = key_info.dict() + key_info = key_info.dict() # pyright: ignore[reportDeprecated] # deliberate pydantic v1 fallback key_token_hash: Final = key_info.pop("token") model_max_budget = key_info.get("model_max_budget") or {} @@ -4012,7 +3985,10 @@ async def generate_key_helper_fn( if table_name is None or table_name == "user": # do not auto-create users for `/key/generate` ## CREATE USER (If necessary) if query_type == "insert_data": - user_row = await prisma_client.insert_data(data=user_data, table_name="user") + user_row = cast( # cast-ok: table_name="user" is the insert_data branch returning the user row + "prisma_models.LiteLLM_UserTable | None", + await prisma_client.insert_data(data=user_data, table_name="user"), + ) if user_row is None: raise Exception("Failed to create user") @@ -4219,9 +4195,12 @@ async def delete_verification_tokens( if prisma_client: hashed_tokens: Final[list[str]] = [_hash_token_if_needed(token=key) for key in tokens] tokens = hashed_tokens - _keys_being_deleted: Final[list[LiteLLM_VerificationToken]] = await _prisma_table( - VerificationTokenRepository(prisma_client) - ).find_many(where={"token": {"in": hashed_tokens}}) + _keys_being_deleted: Final[list[LiteLLM_VerificationToken]] = cast( # cast-ok: find_many returns a list + "list[LiteLLM_VerificationToken]", + await _prisma_table(VerificationTokenRepository(prisma_client)).find_many( + where={"token": {"in": hashed_tokens}} + ), + ) if len(_keys_being_deleted) == 0: raise HTTPException( @@ -4297,7 +4276,7 @@ async def delete_verification_tokens( def _transform_verification_tokens_to_deleted_records( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: str | None = None, ) -> list[dict[str, object]]: @@ -4372,7 +4351,7 @@ async def _save_deleted_verification_token_records( async def _persist_deleted_verification_tokens( - keys: list[LiteLLM_VerificationToken], + keys: Sequence[LiteLLM_VerificationToken], prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: str | None = None, @@ -4435,7 +4414,9 @@ async def _rotate_master_key( from litellm.proxy.proxy_server import proxy_config try: - models: list | None = await _prisma_table(ModelRepository(prisma_client)).find_many() + models: list | None = cast( # cast-ok: find_many returns a real list, which TableActions widens to Sequence + "list[object]", await _prisma_table(ModelRepository(prisma_client)).find_many() + ) except Exception: models = None # 2. process model table @@ -5361,9 +5342,9 @@ async def validate_key_list_check( if key_hash: try: - key_info: Final[LiteLLM_VerificationToken] = await VerificationTokenRepository( - prisma_client - ).table.find_unique( + key_info: Final[LiteLLM_VerificationToken | None] = await _prisma_table( + VerificationTokenRepository(prisma_client) + ).find_unique( where={"token": key_hash}, ) except Exception: @@ -5373,6 +5354,13 @@ async def validate_key_list_check( param="key_hash", code=status.HTTP_403_FORBIDDEN, ) + if key_info is None: + raise ProxyException( + message="Key Hash not found.", + type=ProxyErrorTypes.bad_request_error, + param="key_hash", + code=status.HTTP_403_FORBIDDEN, + ) can_user_query_key_info: Final = await _can_user_query_key_info( user_api_key_dict=user_api_key_dict, key=key_hash, @@ -5394,8 +5382,9 @@ async def _fetch_user_team_objects( if complete_user_info is None or not complete_user_info.teams: return [] - teams: Final[list[BaseModel] | None] = await TeamRepository(prisma_client).table.find_many( - where={"team_id": {"in": complete_user_info.teams}} + teams: Final[Sequence[BaseModel] | None] = cast( # cast-ok: the None guard below predates the non-optional seam + "Sequence[BaseModel] | None", + await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": complete_user_info.teams}}), ) if teams is None: return [] @@ -6130,7 +6119,7 @@ async def _list_key_helper( key_dict = key.model_dump() except Exception: # Fallback for Pydantic v1 compatibility - key_dict = key.dict() + key_dict = key.dict() # pyright: ignore[reportDeprecated] # deliberate pydantic v1 fallback # Attach object_permission if object_permission_id is set (only for non-deleted keys) if not use_deleted_table: key_dict = await attach_object_permission_to_dict(key_dict, prisma_client) @@ -6155,7 +6144,9 @@ async def _list_key_helper( # Use deleted key type to preserve deleted_at, deleted_by, etc. key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict)) else: - key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object + key_list.append( + UserAPIKeyAuth(**key_dict) # pyright: ignore[reportAny] # model_dump() is dict[str, Any] + ) else: _token = key_dict.get("token") key_list.append(cast(str, _token)) # Return only the token diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index 8e8545a51cc..e1a7645e988 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -40,17 +40,22 @@ router: Final = APIRouter() class _DeploymentRow(Protocol): - model_id: str - model_name: str - model_info: object + @property + def model_id(self) -> str: ... + + @property + def model_name(self) -> str: ... + + @property + def model_info(self) -> object: ... class _ModelTableClient(Protocol): - async def find_many(self, where: Mapping[str, object] | None = None) -> Sequence[_DeploymentRow]: ... + async def find_many(self, *, where: Mapping[str, object] | None = None) -> Sequence[_DeploymentRow]: ... - async def find_unique(self, where: Mapping[str, object]) -> _DeploymentRow | None: ... + async def find_unique(self, *, where: Mapping[str, object]) -> _DeploymentRow | None: ... - async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... + async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> object: ... def _model_table(prisma_client: PrismaClient) -> _ModelTableClient: @@ -322,7 +327,9 @@ async def get_all_access_groups_from_db( for deployment in deployments: model_info = deployment.model_info or {} - access_groups = model_info.get("access_groups", []) + access_groups = model_info.get( # pyright: ignore[reportAttributeAccessIssue] # Json reads back as a dict + "access_groups", [] + ) model_name = deployment.model_name for access_group in access_groups: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 217fc61a56c..3242c6f6084 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -16,7 +16,7 @@ import json from collections.abc import Awaitable, Mapping, Sequence from json import JSONDecodeError from types import MappingProxyType -from typing import Final, Literal, Protocol, cast +from typing import TYPE_CHECKING, Final, Literal, Protocol, cast from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from pydantic import BaseModel, ConfigDict, Field, ValidationError @@ -72,6 +72,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import ( ) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.model_repository import ModelRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ModelTableRepository from litellm.repositories.team_repository import TeamRepository from litellm.router import Router @@ -100,6 +101,9 @@ from litellm.types.router import ( ) from litellm.utils import get_utc_datetime +if TYPE_CHECKING: + from prisma import models as prisma_models + router: Final = APIRouter() @@ -120,10 +124,14 @@ class UpdatePublicModelGroupsRequest(BaseModel): class _ProxyModelRow(Protocol): - model_id: str - model_name: str - litellm_params: Mapping[str, object] - model_info: Mapping[str, object] | None + @property + def model_id(self) -> str: ... + + @property + def model_name(self) -> str: ... + + @property + def model_info(self) -> object: ... def model_dump_json(self, *, exclude_none: bool = False) -> str: ... @@ -133,7 +141,9 @@ class _ProxyModelTable(Protocol): def find_many(self, *, where: Mapping[str, object]) -> Awaitable[Sequence[_ProxyModelRow]]: ... - def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> Awaitable[_ProxyModelRow]: ... + def update( + self, *, where: Mapping[str, object], data: Mapping[str, object] + ) -> Awaitable[_ProxyModelRow | None]: ... def delete(self, *, where: Mapping[str, object]) -> Awaitable[_ProxyModelRow | None]: ... @@ -144,41 +154,35 @@ class _TxModelTables(Protocol): litellm_proxymodeltable: _ProxyModelTable +class _ExistingModelRow(Protocol): + @property + def litellm_params(self) -> Mapping[str, object]: ... + + def model_dump_json(self, *, exclude_none: bool = False) -> str: ... + + class _TeamRow(Protocol): - models: Sequence[str] + @property + def models(self) -> Sequence[str]: ... def model_dump(self) -> Mapping[str, object]: ... -class _TeamTable(Protocol): +class _TeamLookupTable(Protocol): def find_unique(self, *, where: Mapping[str, object]) -> Awaitable[_TeamRow | None]: ... + +class _TeamTable(_TeamLookupTable, Protocol): def update( self, *, where: Mapping[str, object], data: Mapping[str, object], include: Mapping[str, bool] ) -> Awaitable[LiteLLM_TeamTable]: ... -class _TeamIdRef(Protocol): - team_id: str - - -class _ModelAliasRow(Protocol): - id: int - model_aliases: dict[str, str] - team: _TeamIdRef | None - - -class _ModelAliasTable(Protocol): - def find_many(self, *, include: Mapping[str, bool]) -> Awaitable[Sequence[_ModelAliasRow]]: ... - - def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> Awaitable[object]: ... - - def _proxy_model_table(prisma_client: PrismaClient) -> _ProxyModelTable: return ModelRepository(prisma_client).table -def _repo_team_table(prisma_client: PrismaClient) -> _TeamTable: +def _repo_team_table(prisma_client: PrismaClient) -> _TeamLookupTable: return TeamRepository(prisma_client).table @@ -186,7 +190,7 @@ def _db_team_table(prisma_client: PrismaClient) -> _TeamTable: return prisma_client.db.litellm_teamtable -def _model_alias_table(prisma_client: PrismaClient) -> _ModelAliasTable: +def _model_alias_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_ModelTable]": return ModelTableRepository(prisma_client).table @@ -677,6 +681,14 @@ async def patch_model( data=update_data, ) + if updated_model is None: + raise ProxyException( + message=f"Model {model_id} not found on proxy.", + type=ProxyErrorTypes.not_found_error, + code=status.HTTP_404_NOT_FOUND, + param=None, + ) + # Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates) live_before_reload: Final = live_model_ids_snapshot() reload_outcome: Final = await clear_cache() @@ -811,7 +823,7 @@ async def _set_model_blocked_status( live_after=reload_outcome.live_after, ) - return updated_model + return updated_model # pyright: ignore[reportReturnType] # prisma row, coerced by this route's response_model except Exception as e: verbose_proxy_logger.exception("Error in model %s: %s", action, e) @@ -897,7 +909,7 @@ async def _add_model_to_db( prisma_client: PrismaClient, new_encryption_key: str | None = None, should_create_model_in_db: bool = True, -) -> LiteLLM_ProxyModelTable | None: +) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None": # encrypt litellm params # _litellm_params_dict: Final = model_params.litellm_params.dict(exclude_none=True) _original_litellm_model_name: Final = model_params.litellm_params.model @@ -914,8 +926,9 @@ async def _add_model_to_db( } if model_params.model_info.id is not None: _data["model_id"] = model_params.model_info.id + _create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above if should_create_model_in_db: - model_response = await ModelRepository(prisma_client).table.create(data=_data) + model_response = await ModelRepository(prisma_client).table.create(data=_create_data) else: model_response = LiteLLM_ProxyModelTable(**_data) return model_response @@ -925,7 +938,7 @@ async def _add_team_model_to_db( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, -) -> LiteLLM_ProxyModelTable | None: +) -> "prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None": """ If 'team_id' is provided, @@ -1638,7 +1651,9 @@ async def delete_team_model_alias( tasks: Final = [] removed_model_aliases: Final[list[tuple[str, str]]] = [] for team_model_alias in team_model_aliases: - model_aliases = team_model_alias.model_aliases # {"alias": "public model name"} + model_aliases = cast( # cast-ok: prisma types Json columns as `str`; the driver hands back the parsed dict + "dict[str, str]", team_model_alias.model_aliases + ) id = team_model_alias.id if public_model_name in model_aliases.values(): @@ -1733,7 +1748,7 @@ async def add_new_model( existing_params=None, ) - model_response: LiteLLM_ProxyModelTable | None = None + model_response: prisma_models.LiteLLM_ProxyModelTable | LiteLLM_ProxyModelTable | None = None # update DB incoming_model_info: Final = model_params.model_info.model_dump(exclude_none=True) _raise_if_ptu_cost_attribution_disabled(incoming_model_info) @@ -1902,7 +1917,10 @@ async def update_model( # update DB if store_model_in_db is True: - _existing_litellm_params_dict: Final = dict(_existing_litellm_params.litellm_params) + existing_model_row: Final = cast( # cast-ok: prisma types Json columns as `str`; the driver parses them + "_ExistingModelRow", _existing_litellm_params + ) + _existing_litellm_params_dict: Final = dict(existing_model_row.litellm_params) if model_params.litellm_params is None: raise Exception("litellm_params not provided") @@ -1946,8 +1964,8 @@ async def update_model( user_api_key_dict=user_api_key_dict, table_name=LitellmTableNames.PROXY_MODEL_TABLE_NAME, before_value=( - _existing_litellm_params.model_dump_json(exclude_none=True) - if isinstance(_existing_litellm_params, BaseModel) + existing_model_row.model_dump_json(exclude_none=True) + if isinstance(existing_model_row, BaseModel) else None ), after_value=( diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index ffca858c0ce..9198aa35f3f 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -14,7 +14,14 @@ Endpoints for /organization operations #### ORGANIZATION MANAGEMENT #### from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Annotated, Final, Protocol, overload +from typing import ( + TYPE_CHECKING, + Annotated, + Final, + Protocol, + cast, # noqa: TID251 # prisma types Json columns as fields.Json but reads back plain python values + overload, +) import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, status @@ -74,6 +81,11 @@ if TYPE_CHECKING: router: Final = APIRouter() +class _ObjectPermissionRow(Protocol): + @property + def object_permission_id(self) -> str | None: ... + + class _UserTableClient(Protocol): async def find_unique(self, where: Mapping[str, object]) -> "PrismaUserTable | None": ... @@ -681,7 +693,10 @@ async def update_organization( existing_metadata: Final = existing_organization_row.metadata or {} updated_metadata: Final = updated_organization_row_json.get("metadata", {}) merged_metadata: Final[Mapping[str, object]] = _update_dictionary( - existing_dict=existing_metadata.copy(), new_dict=updated_metadata + existing_dict=cast( # cast-ok: prisma de-serializes a Json column to the plain python dict it stores + "dict[str, object]", existing_metadata + ).copy(), + new_dict=updated_metadata, ) updated_organization_row_json["metadata"] = merged_metadata @@ -720,7 +735,7 @@ async def update_organization( async def handle_update_object_permission( data_json: dict[str, object], - existing_organization_row: LiteLLM_OrganizationTable, + existing_organization_row: _ObjectPermissionRow, ) -> dict[str, object]: """ Handle the update of object permission for an organization. @@ -1276,17 +1291,20 @@ async def find_member_if_email(user_email: str, prisma_client: PrismaClient) -> Find a member if the user_email is in LiteLLM_UserTable """ + not_unique_user_email_error: Final = HTTPException( + status_code=400, + detail={ + "error": f"Unique user not found for user_email={user_email}. Potential duplicate OR non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead." + }, + ) try: - existing_user_email_row: Final[BaseModel] = await UserRepository(prisma_client).table.find_unique( + existing_user_email_row: Final = await UserRepository(prisma_client).table.find_unique( where={"user_email": user_email} ) except Exception: - raise HTTPException( - status_code=400, - detail={ - "error": f"Unique user not found for user_email={user_email}. Potential duplicate OR non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead." - }, - ) + raise not_unique_user_email_error + if existing_user_email_row is None: + raise not_unique_user_email_error existing_user_email_row_pydantic: Final = LiteLLM_UserTable.model_validate(existing_user_email_row.model_dump()) return existing_user_email_row_pydantic @@ -1537,7 +1555,10 @@ async def add_member_to_organization( _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") if _returned_user is not None: user_object = LiteLLM_UserTable.model_validate(_returned_user.model_dump()) - elif existing_user_email_row is not None and len(existing_user_email_row) > 1: + elif existing_user_email_row is not None and ( + len(existing_user_email_row) # pyright: ignore[reportArgumentType] # find_unique yields a row, not a list + > 1 + ): raise HTTPException( status_code=400, detail={"error": "Multiple users with this email found in db. Please use 'user_id' instead."}, diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 2d95d0bea29..49531a7a72d 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -33,7 +33,8 @@ class ScimTransformations: # Get user's teams/groups groups: Final = [] - for team_id in user.teams or []: + team_ids: Final[list[str]] = user.teams or [] # mutable-ok: scim reads the user row's team ids + for team_id in team_ids: team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if team: team_alias = getattr(team, "team_alias", team.team_id) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 7183e6cb402..4c16f3b4d7b 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -2761,6 +2761,12 @@ async def patch_group( if final_team: updated_team = final_team + if updated_team is None: + raise HTTPException( + status_code=404, + detail={"error": f"Group not found with ID: {group_id}"}, # mutable-ok: FastAPI detail contract + ) + # Convert to SCIM format and return scim_group: Final = await ScimTransformations.transform_litellm_team_to_scim_group( LiteLLM_TeamTable.model_validate(updated_team.model_dump()) diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index 7aeb5039687..b74aa1a4e16 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -369,10 +369,10 @@ async def _add_tag_to_deployment(deployment: "Deployment", tag: str): # Prisma returns litellm_params as dict (already parsed from JSON) existing_params = db_model.litellm_params - if isinstance(existing_params, str): + if isinstance(existing_params, str): # pyright: ignore[reportUnnecessaryIsInstance] # prisma Json stub is str # If it's a string, parse it existing_params = json.loads(existing_params) - elif not isinstance(existing_params, dict): + elif not isinstance(existing_params, dict): # pyright: ignore[reportUnnecessaryIsInstance] # prisma Json stub raise Exception(f"Unexpected litellm_params type: {type(existing_params)}") # Add tag to tags array (preserve encryption of other fields) diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 14a2a8a98a5..9a2ec38d627 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -352,6 +352,9 @@ async def add_team_callbacks( include={"object_permission": True}, # mutable-ok: prisma include takes a dict literal ) + if new_team_row is None: + raise _callback_error(400, f"Team id = {team_id} does not exist. Please use a different team id.") + # Without this a newly registered callback stays dormant for existing keys. await _refresh_cached_team( team_row=new_team_row, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 01254d5c064..c8373fe6c30 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -16,7 +16,16 @@ import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Annotated, Final, NamedTuple, Protocol, TypedDict, TypeVar, cast +from typing import ( + TYPE_CHECKING, + Annotated, + Final, + NamedTuple, + Protocol, + TypedDict, + TypeVar, + cast, +) import fastapi from fastapi import APIRouter, Depends, Header, HTTPException, Request, status @@ -33,21 +42,17 @@ from litellm.proxy._types import ( BudgetNewRequest, CommonProxyErrors, DeleteTeamRequest, - LiteLLM_AccessGroupTable, LiteLLM_AuditLogs, - LiteLLM_BudgetTableFull, LiteLLM_DeletedTeamTable, LiteLLM_ManagementEndpoint_MetadataFields, LiteLLM_ManagementEndpoint_MetadataFields_Premium, LiteLLM_ModelTable, - LiteLLM_OrganizationMembershipTable, LiteLLM_OrganizationTable, LiteLLM_OrganizationTableWithMembers, LiteLLM_TeamMembership, LiteLLM_TeamTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, - LiteLLM_VerificationToken, LitellmTableNames, LitellmUserRoles, Member, @@ -138,6 +143,7 @@ from litellm.proxy.management_helpers.utils import ( from litellm.proxy.utils import PrismaClient, ProxyLogging, handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.organization_repository import OrganizationRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( AccessGroupRepository, DeletedTeamRepository, @@ -169,6 +175,10 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( UpdateTeamMemberPermissionsRequest, ) +if TYPE_CHECKING: + from prisma import Prisma + from prisma import models as prisma_models + router: Final = APIRouter() _DbRecordT = TypeVar("_DbRecordT") @@ -183,95 +193,14 @@ class _TeamIdGroupRow(TypedDict): _count: _TeamIdKeyCount -class _PrismaTableActions(Protocol[_DbRecordT]): - async def find_unique( - self, - where: Mapping[str, object], - include: Mapping[str, bool] | None = None, - ) -> _DbRecordT | None: ... - - async def find_first( - self, - where: Mapping[str, object] | None = None, - order: Mapping[str, str] | None = None, - ) -> _DbRecordT | None: ... - - async def find_many( - self, - where: Mapping[str, object] | None = None, - include: Mapping[str, bool] | None = None, - order: Mapping[str, str] | None = None, - skip: int | None = None, - take: int | None = None, - cursor: Mapping[str, object] | None = None, - ) -> list[_DbRecordT]: ... - - async def create( - self, - data: Mapping[str, object], - include: Mapping[str, bool] | None = None, - ) -> _DbRecordT: ... - - async def create_many( - self, - data: Sequence[Mapping[str, object]], - skip_duplicates: bool | None = None, - ) -> int: ... - - async def update( - self, - where: Mapping[str, object], - data: Mapping[str, object], - include: Mapping[str, bool] | None = None, - ) -> _DbRecordT: ... - - async def update_many( - self, - where: Mapping[str, object], - data: Mapping[str, object], - ) -> int: ... - - async def upsert( - self, - where: Mapping[str, object], - data: Mapping[str, Mapping[str, object]], - ) -> _DbRecordT: ... - - async def delete_many( - self, - where: Mapping[str, object] | None = None, - ) -> int: ... - - async def count( - self, - where: Mapping[str, object] | None = None, - ) -> int: ... - - async def group_by( - self, - by: Sequence[str], - where: Mapping[str, object] | None = None, - count: Mapping[str, bool] | None = None, - ) -> Sequence[_TeamIdGroupRow]: ... - - -class _HasTableActions(Protocol[_DbRecordT]): - @property - def table(self) -> "_PrismaTableActions[_DbRecordT]": ... - - -def _typed_table( - repo: "_HasTableActions[_DbRecordT]", record_type: type[_DbRecordT] -) -> "_PrismaTableActions[_DbRecordT]": - return repo.table - - def _as_object(value: object) -> object: return value -def _nullable(value: _DbRecordT | None) -> _DbRecordT | None: - return value +def _as_list(rows: Sequence[_DbRecordT]) -> list[_DbRecordT]: # mutable-ok: pydantic list[...] fields reject Sequence + return cast( # cast-ok: prisma-client-py find_many returns a list; TableActions only widens it to Sequence + "list[_DbRecordT]", rows + ) class _UserIdRow(Protocol): @@ -279,33 +208,75 @@ class _UserIdRow(Protocol): def user_id(self) -> str | None: ... -class _HasUserIdTable(Protocol): - @property - def table(self) -> "_PrismaTableActions[_UserIdRow]": ... - - -def _user_id_rows_db(repo: "_HasUserIdTable") -> "_PrismaTableActions[_UserIdRow]": +def _user_id_rows_db(repo: UserRepository) -> "TableActions[_UserIdRow]": return repo.table -class _RawTeamRow(Protocol): +class _ModelDumpRow(Protocol): + def model_dump(self) -> Mapping[str, object]: ... + + +class _TeamIdRow(Protocol): @property - def members_with_roles(self) -> Sequence[Mapping[str, object]] | None: ... + def team_id(self) -> str: ... -class _HasRawTeamTable(Protocol): +class _CacheableTeamRow(_TeamIdRow, _ModelDumpRow, Protocol): ... + + +class _ObjectPermissionRow(Protocol): @property - def table(self) -> "_PrismaTableActions[_RawTeamRow]": ... + def object_permission_id(self) -> str | None: ... -def _raw_team_db(repo: "_HasRawTeamTable") -> "_PrismaTableActions[_RawTeamRow]": - return repo.table +class _TeamAliasBudgetRow(Protocol): + @property + def team_alias(self) -> str | None: ... + + @property + def budget_duration(self) -> str | None: ... + + +class _TeamBudgetRow(_TeamAliasBudgetRow, Protocol): + metadata: Mapping[str, JsonValue] | None + + +class _AuditableTeamRow(Protocol): + def json(self, *, exclude_none: bool = False) -> str: ... + + +class _RawTeamRow(_TeamIdRow, _ModelDumpRow, _ObjectPermissionRow, _TeamBudgetRow, _AuditableTeamRow, Protocol): + @property + def members_with_roles( + self, + ) -> Sequence[dict[str, object]] | None: ... # mutable-ok: prisma deserializes this JSON column into plain dicts + + @property + def organization_id(self) -> str | None: ... + + @property + def max_budget(self) -> float | None: ... + + @property + def soft_budget(self) -> float | None: ... + + @property + def model_id(self) -> int | None: ... + + +def _raw_team_db(repo: TeamRepository) -> "TableActions[_RawTeamRow]": + return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value + "TableActions[_RawTeamRow]", repo.table + ) + + +class _BudgetIdRow(Protocol): + @property + def budget_id(self) -> str: ... class _BudgetWriteCall(Protocol): - async def __call__( - self, budget_obj: BudgetNewRequest, user_api_key_dict: UserAPIKeyAuth - ) -> LiteLLM_BudgetTableFull: ... + async def __call__(self, budget_obj: BudgetNewRequest, user_api_key_dict: UserAPIKeyAuth) -> _BudgetIdRow: ... def _as_budget_write(fn: "_BudgetWriteCall") -> "_BudgetWriteCall": @@ -330,7 +301,7 @@ class _TeamIdInFilter(TypedDict, total=False): class _TeamCreateTx(AccessGroupSyncTx, Protocol): @property - def litellm_teamtable(self) -> "_PrismaTableActions[LiteLLM_TeamTable]": ... + def litellm_teamtable(self) -> "TableActions[prisma_models.LiteLLM_TeamTable]": ... _STRIP_DELETED_TEAM_FROM_USERS_SQL: Final = """ @@ -340,46 +311,52 @@ UPDATE "LiteLLM_UserTable" SET teams = array_remove(teams, $1) WHERE $1 = ANY(te _INCLUDE_MODEL_TABLE: Final = MappingProxyType({"litellm_model_table": True}) -def _team_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_TeamTable]": - return _typed_table(TeamRepository(prisma_client), LiteLLM_TeamTable) +def _team_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_TeamTable]": + return TeamRepository(prisma_client).table -def _team_membership_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_TeamMembership]": - return _typed_table(TeamMembershipRepository(prisma_client), LiteLLM_TeamMembership) +def _team_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamTable]": + return cast( # cast-ok: generated actions type Json columns as str; TableActions widens inputs to Mapping + "TableActions[prisma_models.LiteLLM_TeamTable]", tx.litellm_teamtable + ) -def _user_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_UserTable]": - return _typed_table(UserRepository(prisma_client), LiteLLM_UserTable) +def _team_membership_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_TeamMembership]": + return TeamMembershipRepository(prisma_client).table -def _model_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_ModelTable]": - return _typed_table(ModelTableRepository(prisma_client), LiteLLM_ModelTable) +def _user_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_UserTable]": + return UserRepository(prisma_client).table -def _org_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_OrganizationTable]": - return _typed_table(OrganizationRepository(prisma_client), LiteLLM_OrganizationTable) +def _model_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_ModelTable]": + return ModelTableRepository(prisma_client).table + + +def _org_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_OrganizationTable]": + return OrganizationRepository(prisma_client).table def _org_membership_db( prisma_client: PrismaClient | None, -) -> "_PrismaTableActions[LiteLLM_OrganizationMembershipTable]": - return _typed_table(OrganizationMembershipRepository(prisma_client), LiteLLM_OrganizationMembershipTable) +) -> "TableActions[prisma_models.LiteLLM_OrganizationMembership]": + return OrganizationMembershipRepository(prisma_client).table -def _budget_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_BudgetTableFull]": - return _typed_table(BudgetRepository(prisma_client), LiteLLM_BudgetTableFull) +def _budget_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_BudgetTable]": + return BudgetRepository(prisma_client).table -def _deleted_team_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_DeletedTeamTable]": - return _typed_table(DeletedTeamRepository(prisma_client), LiteLLM_DeletedTeamTable) +def _deleted_team_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_DeletedTeamTable]": + return DeletedTeamRepository(prisma_client).table -def _access_group_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_AccessGroupTable]": - return _typed_table(AccessGroupRepository(prisma_client), LiteLLM_AccessGroupTable) +def _access_group_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_AccessGroupTable]": + return AccessGroupRepository(prisma_client).table -def _tokens_db(prisma_client: PrismaClient | None) -> "_PrismaTableActions[LiteLLM_VerificationToken]": - return _typed_table(VerificationTokenRepository(prisma_client), LiteLLM_VerificationToken) +def _tokens_db(prisma_client: PrismaClient | None) -> "TableActions[prisma_models.LiteLLM_VerificationToken]": + return VerificationTokenRepository(prisma_client).table def _sanitize_for_log(value: object) -> str: @@ -392,7 +369,7 @@ def _sanitize_for_log(value: object) -> str: async def _refresh_cached_team( - team_row: LiteLLM_TeamTable, + team_row: _CacheableTeamRow, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> None: @@ -481,7 +458,7 @@ class TeamMemberBudgetHandler: @staticmethod async def create_team_member_budget_table( - data: NewTeamRequest | LiteLLM_TeamTable, + data: NewTeamRequest | _TeamAliasBudgetRow, new_team_data_json: dict, user_api_key_dict: UserAPIKeyAuth, team_member_budget: float | None = None, @@ -532,7 +509,7 @@ class TeamMemberBudgetHandler: @staticmethod async def upsert_team_member_budget_table( - team_table: LiteLLM_TeamTable, + team_table: _TeamBudgetRow, user_api_key_dict: UserAPIKeyAuth, updated_kv: dict, team_member_budget: float | None = None, @@ -603,7 +580,7 @@ class TeamMemberBudgetHandler: @staticmethod async def clear_team_member_budget_fields( - team_table: LiteLLM_TeamTable, + team_table: _TeamBudgetRow, user_api_key_dict: "UserAPIKeyAuth", updated_kv: dict, explicitly_set_fields: set, @@ -1540,7 +1517,7 @@ async def new_team( tx: _TeamCreateTx async with prisma_client.db.tx() as tx: - team_row: Final[LiteLLM_TeamTable] = await tx.litellm_teamtable.create( + team_row: Final[prisma_models.LiteLLM_TeamTable] = await tx.litellm_teamtable.create( data=team_creation_data, include=_INCLUDE_MODEL_TABLE, ) @@ -1595,7 +1572,7 @@ async def new_team( async def _create_team_update_audit_log( - existing_team_row: LiteLLM_TeamTable, + existing_team_row: _AuditableTeamRow, updated_kv: dict, team_id: str, litellm_changed_by: str | None, @@ -1718,11 +1695,11 @@ async def _auto_add_team_members_to_organization( async def fetch_and_validate_organization( organization_id: str, - existing_team_row: LiteLLM_TeamTable, + existing_team_row: _ModelDumpRow, llm_router: Router | None, prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth | None = None, -) -> LiteLLM_OrganizationTable: +) -> "prisma_models.LiteLLM_OrganizationTable": """ Fetch and validate an organization for team update operations. @@ -1996,7 +1973,9 @@ async def update_team( validate_budget_duration(data.budget_duration) validate_budget_duration(data.team_member_budget_duration) - existing_team_row = await _team_db(prisma_client).find_unique(where={"team_id": data.team_id}) + existing_team_row = await _raw_team_db(TeamRepository(prisma_client)).find_unique( + where={"team_id": data.team_id} + ) if existing_team_row is None: raise HTTPException( @@ -2234,18 +2213,16 @@ async def update_team( updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) team_update_data: Final[Mapping[str, object]] = updated_kv - team_row: Final[LiteLLM_TeamTable | None] = _nullable( - await _team_db(prisma_client).update( - where={"team_id": data.team_id}, - data=team_update_data, - # `object_permission` is included so `_refresh_cached_team` - # doesn't write a cached team with the relation nulled out — - # see team_model_add for the full rationale. - include={ - "litellm_model_table": True, - "object_permission": True, - }, - ) + team_row: Final = await _team_db(prisma_client).update( + where={"team_id": data.team_id}, + data=team_update_data, + # `object_permission` is included so `_refresh_cached_team` + # doesn't write a cached team with the relation nulled out. + # See team_model_add for the full rationale. + include={ + "litellm_model_table": True, + "object_permission": True, + }, ) if team_row is None or team_row.team_id is None: @@ -2375,7 +2352,7 @@ def _set_budget_reset_at(data: UpdateTeamRequest, updated_kv: dict) -> None: updated_kv["budget_limits"] = json.dumps(initialized_windows) -async def handle_update_object_permission(data_json: dict, existing_team_row: LiteLLM_TeamTable) -> dict: +async def handle_update_object_permission(data_json: dict, existing_team_row: _ObjectPermissionRow) -> dict: """ Handle the update of object permission for a team. @@ -2705,7 +2682,7 @@ async def _add_team_members_to_team( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, -) -> tuple[LiteLLM_TeamTable, list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: +) -> tuple["prisma_models.LiteLLM_TeamTable", list[LiteLLM_UserTable], list[LiteLLM_TeamMembership]]: """Add team members to the team. The members_with_roles reconciliation runs inside a transaction that locks @@ -2750,7 +2727,7 @@ async def _write_members_with_roles_locked( complete_team_data: LiteLLM_TeamTable, prisma_client: PrismaClient, updated_users: list[LiteLLM_UserTable], -) -> LiteLLM_TeamTable | None: +) -> "prisma_models.LiteLLM_TeamTable | None": """Reconcile members_with_roles under the team row lock. None when the team row is gone. That read is at least as recent as the user and membership writes the caller @@ -2772,7 +2749,7 @@ async def _write_members_with_roles_locked( ) _db_team_members: Final = [m.model_dump() for m in complete_team_data.members_with_roles] - return await tx.litellm_teamtable.update( + return await _team_tx_db(tx).update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_team_members)}, ) @@ -3292,7 +3269,9 @@ async def team_member_delete( key_val: Final[Mapping[str, object]] = ( {"user_id": {"in": sorted(removed_user_ids)}} if removed_user_ids else {"user_email": data.user_email} ) - existing_user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many(where=key_val) + existing_user_rows: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await _user_db(prisma_client).find_many( + where=key_val + ) # Also clean up any existing team membership rows for this user and team user_ids_to_delete: Final = removed_user_ids.union( @@ -3303,7 +3282,7 @@ async def team_member_delete( ## DELETE KEYS CREATED BY USER FOR THIS TEAM # Fetch keys before deletion so their audit records can be persisted alongside the delete. # An empty user_ids_to_delete still resolves cleanly: prisma's "in": [] matches no rows. - keys_to_delete: Final[list[LiteLLM_VerificationToken]] = await _tokens_db(prisma_client).find_many( + keys_to_delete: Final = await _tokens_db(prisma_client).find_many( where={ "user_id": {"in": sorted(user_ids_to_delete)}, "team_id": data.team_id, @@ -3313,7 +3292,7 @@ async def team_member_delete( # All four cleanups run on one connection so a failure between them leaves # no partial removal: either every write below lands, or none of them do. async with prisma_client.tx() as tx: - await tx.litellm_teamtable.update( + await _team_tx_db(tx).update( where={"team_id": data.team_id}, data={"members_with_roles": json.dumps(_db_new_team_members)}, ) @@ -3826,9 +3805,7 @@ async def delete_team( _persist_deleted_verification_tokens, ) - keys_to_delete: list[LiteLLM_VerificationToken] = await _tokens_db(prisma_client).find_many( - where={"team_id": {"in": data.team_ids}} - ) + keys_to_delete: Final = await _tokens_db(prisma_client).find_many(where={"team_id": {"in": data.team_ids}}) if keys_to_delete: await _persist_deleted_verification_tokens( @@ -3930,7 +3907,7 @@ async def _sweep_deleted_team_references(team_ids: Sequence[str], prisma_client: async def _invalidate_deleted_key_cache( - keys: Sequence[LiteLLM_VerificationToken], + keys: "Sequence[prisma_models.LiteLLM_VerificationToken]", user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, ) -> None: @@ -4115,7 +4092,7 @@ async def _hydrate_member_emails( if not missing_user_ids: return tuple(members) - user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many( + user_rows: Final[Sequence[prisma_models.LiteLLM_UserTable]] = await _user_db(prisma_client).find_many( where={ # mutable-ok: Prisma query filters are dict-shaped "user_id": { # mutable-ok: Prisma query filters are dict-shaped "in": sorted(missing_user_ids) @@ -4126,7 +4103,7 @@ async def _hydrate_member_emails( return tuple( m.model_copy(update={"user_email": email_by_user_id[m.user_id]}) # mutable-ok: pydantic update payload - if not m.user_email and m.user_id in email_by_user_id + if not m.user_email and m.user_id is not None and m.user_id in email_by_user_id else m for m in members ) @@ -4711,7 +4688,7 @@ async def _build_team_list_where_conditions( async def _batch_resolve_access_group_resources( all_access_group_ids: list[str], -) -> dict[str, LiteLLM_AccessGroupTable]: +) -> "dict[str, prisma_models.LiteLLM_AccessGroupTable]": """ Batch-fetch access groups in a single DB query and return them keyed by access_group_id. Missing/invalid groups are silently omitted. @@ -4729,7 +4706,7 @@ async def _batch_resolve_access_group_resources( def _convert_teams_to_response_models( - teams: list, + teams: Sequence, use_deleted_table: bool, keys_count_by_team: dict[str, int] | None = None, ) -> list[TeamListItem | LiteLLM_TeamTable | LiteLLM_DeletedTeamTable]: @@ -4763,7 +4740,7 @@ def _convert_teams_to_response_models( async def _get_keys_count_by_team( prisma_client: PrismaClient, - teams: Sequence[LiteLLM_TeamTable], + teams: Sequence[_TeamIdRow], ) -> dict[str, int]: """Aggregate virtual-key counts per team for the given page of teams. @@ -4775,10 +4752,13 @@ async def _get_keys_count_by_team( if not page_team_ids: return {} - grouped: Final = await _tokens_db(prisma_client).group_by( - by=["team_id"], - where={"team_id": {"in": page_team_ids}}, - count={"team_id": True}, + grouped: Final = cast( # cast-ok: prisma group_by returns one row per `by` key with `count=` nested under "_count" + "Sequence[_TeamIdGroupRow]", + await _tokens_db(prisma_client).group_by( + by=["team_id"], + where={"team_id": {"in": page_team_ids}}, + count={"team_id": True}, + ), ) return {row["team_id"]: row.get("_count", {}).get("team_id", 0) for row in grouped if row.get("team_id")} @@ -5168,7 +5148,7 @@ async def list_team( _team_memberships.append(tm) # add all keys that belong to the team - keys = await _tokens_db(prisma_client).find_many(where={"team_id": team.team_id}) + keys = _as_list(await _tokens_db(prisma_client).find_many(where={"team_id": team.team_id})) try: returned_responses.append( @@ -5403,6 +5383,11 @@ async def team_model_add( data={"updated_at": datetime.now(timezone.utc)}, include={"object_permission": True}, ) + if updated_team is None: + raise HTTPException( + status_code=404, + detail={"error": f"Team not found, passed team_id={data.team_id}"}, + ) await _refresh_cached_team( team_row=updated_team, @@ -5485,6 +5470,11 @@ async def team_model_delete( data={"models": updated_models}, include={"object_permission": True}, ) + if updated_team is None: + raise HTTPException( + status_code=404, + detail={"error": f"Team not found, passed team_id={data.team_id}"}, + ) await _refresh_cached_team( team_row=updated_team, @@ -5619,8 +5609,13 @@ async def update_team_member_permissions( where={"team_id": data.team_id}, data={"team_member_permissions": data.team_member_permissions}, ) + if updated_team is None: + raise HTTPException( + status_code=404, + detail={"error": f"Team not found, passed team_id={data.team_id}"}, + ) - return updated_team + return updated_team # pyright: ignore[reportReturnType] # prisma row, coerced by this route's response_model @router.post( @@ -5685,7 +5680,9 @@ async def bulk_update_team_member_permissions( } -async def _compute_and_batch_updates(prisma_client, teams: Sequence[LiteLLM_TeamTable], permissions_to_add: set) -> int: +async def _compute_and_batch_updates( + prisma_client, teams: "Sequence[prisma_models.LiteLLM_TeamTable]", permissions_to_add: set +) -> int: """Compute merged permissions and batch-write updates. Returns count of teams updated.""" updates: Final = [] for team in teams: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 3c135650de9..0c8240b3298 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -29,7 +29,6 @@ from typing import ( NoReturn, Optional, Protocol, - TypeVar, Union, cast, overload, @@ -122,6 +121,7 @@ from litellm.proxy.utils import ( get_custom_url, get_server_root_path, ) +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import SSOConfigRepository from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository @@ -171,51 +171,16 @@ _CLI_SSO_SECRET_KEY_FRAGMENTS: Final = frozenset( } ) -_DbRecordT: Final = TypeVar("_DbRecordT", covariant=True) - - -class _PrismaTableActions(Protocol[_DbRecordT]): - async def find_unique( - self, - where: Mapping[str, object], - ) -> _DbRecordT | None: ... - - async def find_first( - self, - where: Mapping[str, object] | None = None, - ) -> _DbRecordT | None: ... - - async def find_many( - self, - where: Mapping[str, object] | None = None, - include: Mapping[str, bool] | None = None, - ) -> Sequence[_DbRecordT]: ... - - async def update( - self, - where: Mapping[str, object], - data: Mapping[str, object], - ) -> _DbRecordT: ... - - async def update_many( - self, - where: Mapping[str, object], - data: Mapping[str, object], - ) -> int: ... - class _UserMetadataRow(Protocol): @property def metadata(self) -> Mapping[str, object] | None: ... -class _HasUserMetadataTable(Protocol): - @property - def table(self) -> "_PrismaTableActions[_UserMetadataRow]": ... - - -def _user_meta_db(repo: "_HasUserMetadataTable") -> "_PrismaTableActions[_UserMetadataRow]": - return repo.table +def _user_meta_db(repo: UserRepository) -> "TableActions[_UserMetadataRow]": + return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value + "TableActions[_UserMetadataRow]", repo.table + ) class _SsoConfigRow(Protocol): @@ -223,25 +188,17 @@ class _SsoConfigRow(Protocol): def sso_settings(self) -> Mapping[str, object] | None: ... -class _HasSsoConfigTable(Protocol): - @property - def table(self) -> "_PrismaTableActions[_SsoConfigRow]": ... - - -def _sso_config_db(repo: "_HasSsoConfigTable") -> "_PrismaTableActions[_SsoConfigRow]": - return repo.table +def _sso_config_db(repo: SSOConfigRepository) -> "TableActions[_SsoConfigRow]": + return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value + "TableActions[_SsoConfigRow]", repo.table + ) class _TeamDetailRow(Protocol): def model_dump(self) -> Mapping[str, object]: ... -class _HasTeamDetailTable(Protocol): - @property - def table(self) -> "_PrismaTableActions[_TeamDetailRow]": ... - - -def _team_detail_db(repo: "_HasTeamDetailTable") -> "_PrismaTableActions[_TeamDetailRow]": +def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]": return repo.table diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 33b84545915..fb64914f6f5 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -4,7 +4,7 @@ organizations, teams, and keys. """ import json -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Final, Optional @@ -19,6 +19,8 @@ from litellm.repositories.object_permission_repository import ObjectPermissionRe from litellm.repositories.table_repositories import MCPServerRepository if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, LiteLLM_TeamTableCachedObj, @@ -26,7 +28,7 @@ if TYPE_CHECKING: async def attach_object_permission_to_dict( - data_dict: dict, + data_dict: dict[str, object], prisma_client: PrismaClient, ) -> dict: """ @@ -61,7 +63,7 @@ async def attach_object_permission_to_dict( try: object_permission = object_permission.model_dump() except Exception: - object_permission = object_permission.dict() + object_permission = object_permission.dict() # pyright: ignore[reportDeprecated] # pydantic v1 fallback data_dict["object_permission"] = object_permission return data_dict @@ -188,7 +190,9 @@ async def _set_object_permission( return data_json # Clean data: exclude None values and object_permission_id - clean_data: Final = {k: v for k, v in permission_data.items() if v is not None and k != "object_permission_id"} + clean_data: Final[dict[str, object]] = { + k: v for k, v in permission_data.items() if v is not None and k != "object_permission_id" + } # Serialize mcp_tool_permissions to JSON string for GraphQL compatibility if "mcp_tool_permissions" in clean_data: @@ -224,7 +228,7 @@ def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool: async def _get_db_mcp_servers_by_identifiers( identifiers: set[str], prisma_client: PrismaClient | None, -) -> list[Any]: +) -> "Sequence[prisma_models.LiteLLM_MCPServerTable]": if prisma_client is None or not identifiers: return [] diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index cb30ce90c7f..229e4fdb9e7 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -86,6 +86,20 @@ class _PrismaTeamMembershipTable(Protocol): async def create(self, *, data: Mapping[str, object], include: Mapping[str, bool]) -> _PrismaRecord: ... +def _user_table(prisma_client: PrismaClient) -> _PrismaUserTable: + table: Final[_PrismaUserTable] = UserRepository(prisma_client).table + return table + + +async def _find_users_by_email(prisma_client: PrismaClient, user_email: str) -> Sequence[_PrismaUserRecord] | None: + rows: Final[Sequence[_PrismaUserRecord] | None] = await prisma_client.get_data( + key_val={"user_email": user_email}, + table_name="user", + query_type="find_all", + ) + return rows + + def get_new_internal_user_defaults(user_id: str, user_email: str | None = None) -> dict[str, object]: user_info: Final = litellm.default_internal_user_params or {} @@ -309,8 +323,7 @@ async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, t number of teams a user belongs to). Teams added concurrently for a different team id are unaffected, since each update filters on its own team id. """ - user_table: Final[_PrismaUserTable] = UserRepository(prisma_client).table - await user_table.update_many( + await _user_table(prisma_client).update_many( where={"user_id": user_id, "NOT": {"teams": {"has": team_id}}}, data={"teams": {"push": [team_id]}}, ) @@ -348,8 +361,7 @@ async def add_new_member( # Prisma only compiles an upsert down to INSERT ... ON CONFLICT when it # is non-empty, and falls back to a racy SELECT-then-INSERT when it is # not, so this re-states user_id as a no-op rather than being empty. - user_table: Final[_PrismaUserTable] = UserRepository(prisma_client).table - _returned_user: _PrismaUserRecord | None = await user_table.upsert( + _returned_user: _PrismaRecord | None = await _user_table(prisma_client).upsert( where={"user_id": new_member.user_id}, data={ "create": {"teams": [team_id], **new_user_defaults}, @@ -363,11 +375,7 @@ async def add_new_member( new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email) ## user email is not unique acc. to prisma schema -> future improvement ### for now: check if it exists in db, if not - insert it - existing_user_row: Final[list[_PrismaUserRecord] | None] = await prisma_client.get_data( - key_val={"user_email": new_member.user_email}, - table_name="user", - query_type="find_all", - ) + existing_user_row: Final = await _find_users_by_email(prisma_client, new_member.user_email) if existing_user_row is None or (isinstance(existing_user_row, list) and len(existing_user_row) == 0): new_user_defaults["teams"] = [team_id] _returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") diff --git a/litellm/proxy/memory/memory_endpoints.py b/litellm/proxy/memory/memory_endpoints.py index 3ae8dcf64b7..193f7e09f07 100644 --- a/litellm/proxy/memory/memory_endpoints.py +++ b/litellm/proxy/memory/memory_endpoints.py @@ -18,21 +18,20 @@ Scoping: """ import json -from collections.abc import Mapping, Sequence -from datetime import datetime -from typing import TYPE_CHECKING, Final, Protocol +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final from fastapi import APIRouter, Depends, HTTPException, Query from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( CommonProxyErrors, - LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth, user_api_key_has_admin_view, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import MemoryRepository from litellm.repositories.team_repository import TeamRepository from litellm.types.memory_management import ( @@ -44,54 +43,17 @@ from litellm.types.memory_management import ( ) if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.proxy.utils import PrismaClient router: Final = APIRouter() -class _MemoryRecord(Protocol): - memory_id: str - key: str - value: str - metadata: object - user_id: str | None - team_id: str | None - created_at: datetime | None - created_by: str | None - updated_at: datetime | None - updated_by: str | None - - -class _MemoryTableActions(Protocol): - async def create(self, data: Mapping[str, object]) -> _MemoryRecord: ... - - async def find_many( - self, - where: Mapping[str, object] | None = ..., - order: Mapping[str, str] | None = ..., - skip: int = ..., - take: int = ..., - ) -> Sequence[_MemoryRecord]: ... - - async def count(self, where: Mapping[str, object] | None = ...) -> int: ... - - async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _MemoryRecord: ... - - async def delete(self, where: Mapping[str, object]) -> _MemoryRecord | None: ... - - -def _memory_table(prisma_client: "PrismaClient") -> _MemoryTableActions: +def _memory_table(prisma_client: "PrismaClient") -> TableActions["prisma_models.LiteLLM_MemoryTable"]: return MemoryRepository(prisma_client).table -class _TeamTableActions(Protocol): - async def find_unique(self, where: Mapping[str, str]) -> LiteLLM_TeamTable | None: ... - - -def _team_table(prisma_client: "PrismaClient") -> _TeamTableActions: - return TeamRepository(prisma_client).table - - def _serialize_metadata_for_prisma(metadata: object) -> str: """ Encode a `metadata` payload for the `Json?` column. @@ -129,7 +91,7 @@ def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object return {"OR": ors} -def _row_to_model(row: _MemoryRecord) -> LiteLLM_MemoryRow: +def _row_to_model(row: "prisma_models.LiteLLM_MemoryTable") -> LiteLLM_MemoryRow: return LiteLLM_MemoryRow( memory_id=row.memory_id, key=row.key, @@ -163,7 +125,7 @@ def _internal_error(log_message: str, exc: Exception, default_detail: str) -> HT async def _assert_write_access( - prisma_client: "PrismaClient", row: _MemoryRecord, user_api_key_dict: UserAPIKeyAuth + prisma_client: "PrismaClient", row: "prisma_models.LiteLLM_MemoryTable", user_api_key_dict: UserAPIKeyAuth ) -> None: """ Enforce ownership for mutations (PUT/DELETE). @@ -219,7 +181,7 @@ async def _is_team_admin_for(prisma_client: "PrismaClient", user_api_key_dict: U ) try: - team_obj: Final = await _team_table(prisma_client).find_unique(where={"team_id": team_id}) + team_obj: Final = await TeamRepository(prisma_client).find_by_id(team_id, id_field="team_id") except Exception as e: verbose_proxy_logger.exception("Error loading team for write-auth check (team_id=%s): %s", team_id, e) return False @@ -407,7 +369,7 @@ async def list_memory( async def _find_memory_for_caller( prisma_client: "PrismaClient", key: str, user_api_key_dict: UserAPIKeyAuth -) -> _MemoryRecord: +) -> "prisma_models.LiteLLM_MemoryTable": """Look up a memory row by key, scoped to the caller's visibility.""" key_filter: Final[Mapping[str, object]] = {"key": key} vis: Final = _visibility_filter(user_api_key_dict) @@ -418,6 +380,18 @@ async def _find_memory_for_caller( return rows[0] +async def _find_visible_memory_or_none( + prisma_client: "PrismaClient", key: str, user_api_key_dict: UserAPIKeyAuth +) -> "prisma_models.LiteLLM_MemoryTable | None": + """The caller-visible row for `key`, or None when nothing is visible to them.""" + try: + return await _find_memory_for_caller(prisma_client, key, user_api_key_dict) + except HTTPException as e: + if e.status_code == 404: + return None + raise + + @router.get( "/v1/memory/{key:path}", tags=["memory management"], @@ -480,17 +454,8 @@ async def upsert_memory( ) data["updated_by"] = user_api_key_dict.user_id - async def _find_existing() -> _MemoryRecord | None: - """Return the caller-visible row for `key`, or None.""" - try: - return await _find_memory_for_caller(prisma_client, key, user_api_key_dict) - except HTTPException as e: - if e.status_code == 404: - return None - raise - try: - existing: Final = await _find_existing() + existing: Final = await _find_visible_memory_or_none(prisma_client, key, user_api_key_dict) if existing is not None: # Visibility != write authority. Make sure the caller actually # owns this row (their user_id matches, or it's a pure team row in @@ -530,7 +495,7 @@ async def upsert_memory( # instead of surfacing a 500 on a unique-violation. if not _is_unique_violation(e): raise - existing_after_race: Final = await _find_existing() + existing_after_race: Final = await _find_visible_memory_or_none(prisma_client, key, user_api_key_dict) if existing_after_race is None: # Row exists globally but isn't visible to this caller # (owned by someone else). Treat as conflict. @@ -549,6 +514,8 @@ async def upsert_memory( except Exception as e: raise _internal_error("Error upserting memory: %s", e, "Internal error updating memory entry.") + if row is None: + raise HTTPException(status_code=404, detail=f"Memory with key '{key}' not found") return _row_to_model(row) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 142aced4a38..2c8b926dc93 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -4,7 +4,16 @@ import re from collections.abc import Mapping from dataclasses import dataclass, field from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, get_args, runtime_checkable +from typing import ( + TYPE_CHECKING, + Final, + Literal, + Optional, + Protocol, + cast, # noqa: TID251 # prisma types Json columns as fields.Json but de-serializes them to plain python on read + get_args, + runtime_checkable, +) from litellm.proxy._types import ProxyException from litellm.repositories.table_repositories import ( @@ -1183,7 +1192,7 @@ async def ensure_batch_response_managed_file_ids( prisma_client, verbose_proxy_logger, user_api_key_dict=None, - db_batch_object=None, + db_batch_object: "LiteLLM_ManagedObjectTable | None" = None, unified_batch_id: str | Literal[False] | None = None, ) -> None: """Normalize batch file IDs to managed unified IDs before DB persistence.""" @@ -1270,11 +1279,10 @@ async def get_batch_from_database( return None, None # Parse the batch object from database - batch_data: Final = ( - json.loads(db_batch_object.file_object) - if isinstance(db_batch_object.file_object, str) - else db_batch_object.file_object + file_object: Final = cast( # cast-ok: prisma types the Json column as str; reads return the decoded value + "Mapping[str, object] | str", db_batch_object.file_object ) + batch_data: Final = json.loads(file_object) if isinstance(file_object, str) else file_object response: Final = LiteLLMBatch.model_validate(batch_data) response.id = batch_id @@ -1360,7 +1368,7 @@ async def update_batch_in_database( managed_files_obj, prisma_client, verbose_proxy_logger, - db_batch_object=None, + db_batch_object: "LiteLLM_ManagedObjectTable | None" = None, operation: str = "update", user_api_key_dict=None, poller_owns_accounting: bool | None = None, @@ -1427,7 +1435,7 @@ async def update_batch_in_database( # Normalize status for database storage db_status: Final = response.status if response.status != "completed" else "complete" - update_data: Final[dict] = { + update_data: Final[dict[str, object]] = { "status": db_status, "file_object": response.model_dump_json(), "updated_at": litellm.utils.get_utc_datetime(), diff --git a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py index e08d277788f..4de6ef04d76 100644 --- a/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py +++ b/litellm/proxy/pass_through_endpoints/managed_id_rewriter.py @@ -33,7 +33,13 @@ from __future__ import annotations import json import re from collections.abc import Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Final, TypeVar, overload +from typing import ( + TYPE_CHECKING, + Final, + TypeVar, + cast, # noqa: TID251 # prisma stubs type Json columns as fields.Json but de-serialize them on read + overload, +) from urllib.parse import quote, unquote from fastapi import HTTPException @@ -286,11 +292,15 @@ def _canonical_path(route: str) -> str: def _file_table(prisma_client: PrismaClient) -> ManagedFileTable: - return ManagedFileRepository(prisma_client).table + return cast( # cast-ok: stub-only mismatch, prisma returns real lists and de-serialized Json + ManagedFileTable, ManagedFileRepository(prisma_client).table + ) def _object_table(prisma_client: PrismaClient) -> ManagedObjectTable: - return ManagedObjectRepository(prisma_client).table + return cast( # cast-ok: stub-only mismatch, prisma returns real lists and de-serialized Json + ManagedObjectTable, ManagedObjectRepository(prisma_client).table + ) async def _resolve_one( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 3d721dead4d..60b85cb42d0 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -5,7 +5,7 @@ import json import posixpath import traceback from base64 import b64encode -from collections.abc import AsyncGenerator, Callable, Iterable, Mapping +from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence from datetime import datetime from itertools import groupby from typing import Any, Final, TypedDict, cast @@ -3183,13 +3183,18 @@ async def _filter_endpoints_by_team_allowed_routes( ) # retrieve team metadata - team_metadata: Final = team.metadata + team_metadata: Final = cast( # cast-ok: prisma types the Json column as str; reads hand back the decoded value + "Mapping[str, object] | None", team.metadata + ) if team_metadata is not None and team_metadata.get("allowed_passthrough_routes") is not None: ## FILTER pass_through_endpoints by allowed_passthrough_routes pass_through_endpoints = [ endpoint for endpoint in pass_through_endpoints - if endpoint.path in team_metadata.get("allowed_passthrough_routes") + if endpoint.path + in cast( # cast-ok: guarded above; team metadata stores this key as a list of route paths + "Sequence[str]", team_metadata.get("allowed_passthrough_routes") + ) ] return pass_through_endpoints diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index 6d7f651b9b4..f55eb4f7863 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -10,9 +10,20 @@ by policy_attachments (see AttachmentRegistry). import json from collections.abc import Mapping, Sequence from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, Union +from typing import ( + TYPE_CHECKING, + Any, + Final, + Literal, + Optional, + Protocol, + TypedDict, + Union, + cast, # noqa: TID251 # prisma types the condition/pipeline Json columns as str, but reads return decoded values +) from litellm._logging import verbose_proxy_logger +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import PolicyRepository from litellm.types.proxy.policy_engine import ( GuardrailPipeline, @@ -65,15 +76,32 @@ class _PolicyRow(Protocol): class _PolicyVersionSourceRow(Protocol): - policy_id: str - policy_name: str - version_number: int - inherit: str | None - description: str | None - guardrails_add: Sequence[str] | None - guardrails_remove: Sequence[str] | None - condition: Mapping[str, object] | str | None - pipeline: Mapping[str, object] | str | None + @property + def policy_id(self) -> str: ... + + @property + def policy_name(self) -> str: ... + + @property + def version_number(self) -> int: ... + + @property + def inherit(self) -> str | None: ... + + @property + def description(self) -> str | None: ... + + @property + def guardrails_add(self) -> Sequence[str] | None: ... + + @property + def guardrails_remove(self) -> Sequence[str] | None: ... + + @property + def condition(self) -> Mapping[str, object] | str | None: ... + + @property + def pipeline(self) -> Mapping[str, object] | str | None: ... class _PolicyTableClient(Protocol): @@ -96,23 +124,15 @@ class _PolicyTableClient(Protocol): async def delete_many(self, where: Mapping[str, object]) -> int: ... -class _PolicyVersionSourceTableClient(Protocol): - async def find_unique(self, where: Mapping[str, object]) -> _PolicyVersionSourceRow | None: ... - - async def find_first( - self, - where: Mapping[str, object], - order: Mapping[str, str] | None = None, - ) -> _PolicyVersionSourceRow | None: ... - - def _policy_table(prisma_client: "PrismaClient") -> _PolicyTableClient: - table: Final[_PolicyTableClient] = PolicyRepository(prisma_client).table - return table + table: Final = PolicyRepository(prisma_client).table + return cast( # cast-ok: prisma types Json columns as str; the client hands back the decoded condition/pipeline + "_PolicyTableClient", table + ) -def _policy_version_source_table(prisma_client: "PrismaClient") -> _PolicyVersionSourceTableClient: - table: Final[_PolicyVersionSourceTableClient] = PolicyRepository(prisma_client).table +def _policy_version_source_table(prisma_client: "PrismaClient") -> "TableActions[_PolicyVersionSourceRow]": + table: Final[TableActions[_PolicyVersionSourceRow]] = PolicyRepository(prisma_client).table return table diff --git a/litellm/proxy/policy_engine/policy_resolve_endpoints.py b/litellm/proxy/policy_engine/policy_resolve_endpoints.py index 346586c1e5a..611772dfbae 100644 --- a/litellm/proxy/policy_engine/policy_resolve_endpoints.py +++ b/litellm/proxy/policy_engine/policy_resolve_endpoints.py @@ -6,7 +6,8 @@ Policy resolve and attachment impact estimation endpoints. """ import json -from typing import Final +from collections.abc import Sequence +from typing import TYPE_CHECKING, Final from fastapi import APIRouter, Depends, HTTPException, Query @@ -30,25 +31,28 @@ from litellm.types.proxy.policy_engine import ( PolicyResolveResponse, ) +if TYPE_CHECKING: + from prisma import models as prisma_models + router: Final = APIRouter() -def _build_alias_where(field: str, patterns: list) -> dict: +def _build_alias_where(field: str, patterns: Sequence[str]) -> dict[str, object]: """Build a Prisma ``where`` clause for alias patterns. Supports exact matches and suffix wildcards (``prefix*``). Returns something like: {"OR": [{"field": {"in": ["a","b"]}}, {"field": {"startsWith": "dev-"}}]} """ - exact: Final[list] = [] - prefix_conditions: Final[list] = [] + exact: Final[list[str]] = [] + prefix_conditions: Final[list[dict[str, object]]] = [] for pat in patterns: if pat.endswith("*"): prefix_conditions.append({field: {"startsWith": pat[:-1]}}) else: exact.append(pat) - conditions: Final[list] = [] + conditions: Final[list[dict[str, object]]] = [] if exact: conditions.append({field: {"in": exact}}) conditions.extend(prefix_conditions) @@ -79,7 +83,7 @@ def _get_tags_from_metadata(metadata: object, json_metadata: object = None) -> l return parsed.get("tags", []) or [] -async def _fetch_all_teams(prisma_client: object) -> list: +async def _fetch_all_teams(prisma_client: object) -> "Sequence[prisma_models.LiteLLM_TeamTable]": """Fetch teams from DB once. Reuse the result across tag and alias lookups.""" return await TeamRepository(prisma_client).table.find_many( where={}, @@ -88,13 +92,15 @@ async def _fetch_all_teams(prisma_client: object) -> list: ) -def _filter_keys_by_tags(keys: list, tag_patterns: list) -> tuple: +def _filter_keys_by_tags( + keys: "Sequence[prisma_models.LiteLLM_VerificationToken]", tag_patterns: Sequence[str] +) -> tuple[list[str], int]: """Filter key rows whose metadata.tags match any of the given patterns. Returns (named_aliases, unnamed_count). """ - affected: Final[list] = [] + affected: Final[list[str]] = [] unnamed_count = 0 for key in keys: key_alias = key.key_alias or "" @@ -111,13 +117,15 @@ def _filter_keys_by_tags(keys: list, tag_patterns: list) -> tuple: return affected, unnamed_count -def _filter_teams_by_tags(teams: list, tag_patterns: list) -> tuple: +def _filter_teams_by_tags( + teams: "Sequence[prisma_models.LiteLLM_TeamTable]", tag_patterns: Sequence[str] +) -> tuple[list[str], int]: """Filter pre-fetched team rows whose metadata.tags match any patterns. Returns (named_aliases, unnamed_count). """ - affected: Final[list] = [] + affected: Final[list[str]] = [] unnamed_count = 0 for team in teams: team_alias = team.team_alias or "" @@ -136,18 +144,18 @@ def _filter_teams_by_tags(teams: list, tag_patterns: list) -> tuple: async def _find_affected_by_team_patterns( prisma_client: object, - all_teams: list, - team_patterns: list, - existing_teams: list, - existing_keys: list, -) -> tuple: + all_teams: "Sequence[prisma_models.LiteLLM_TeamTable]", + team_patterns: Sequence[str], + existing_teams: Sequence[str], + existing_keys: Sequence[str], +) -> tuple[list[str], list[str], int]: """Filter pre-fetched teams by alias patterns, then fetch their keys. Returns (new_teams, new_keys, unnamed_keys_count). """ - new_teams: Final[list] = [] - matched_team_ids: Final[list] = [] + new_teams: Final[list[str]] = [] + matched_team_ids: Final[list[str]] = [] for team in all_teams: team_alias = team.team_alias or "" @@ -158,7 +166,7 @@ async def _find_affected_by_team_patterns( new_teams.append(team_alias) matched_team_ids.append(str(team.team_id)) - new_keys: Final[list] = [] + new_keys: Final[list[str]] = [] unnamed_keys_count = 0 if matched_team_ids: keys: Final = await VerificationTokenRepository(prisma_client).table.find_many( @@ -177,10 +185,12 @@ async def _find_affected_by_team_patterns( return new_teams, new_keys, unnamed_keys_count -async def _find_affected_keys_by_alias(prisma_client: object, key_patterns: list, existing_keys: list) -> list: +async def _find_affected_keys_by_alias( + prisma_client: object, key_patterns: Sequence[str], existing_keys: Sequence[str] +) -> list[str]: """Find keys whose alias matches the given patterns.""" - affected: Final[list] = [] + affected: Final[list[str]] = [] keys: Final = await VerificationTokenRepository(prisma_client).table.find_many( where=_build_alias_where("key_alias", key_patterns), @@ -349,8 +359,8 @@ async def estimate_attachment_impact( sample_teams=["(global scope — affects all teams)"], ) - affected_keys: list = [] - affected_teams: list = [] + affected_keys: list[str] = [] + affected_teams: list[str] = [] unnamed_keys = 0 unnamed_teams = 0 @@ -358,7 +368,7 @@ async def estimate_attachment_impact( team_patterns: Final = request.teams or [] # Fetch teams once — reused by both tag-based and alias-based lookups - all_teams: list = [] + all_teams: Sequence[prisma_models.LiteLLM_TeamTable] = [] if tag_patterns or team_patterns: all_teams = await _fetch_all_teams(prisma_client) diff --git a/litellm/proxy/prompts/prompt_endpoints.py b/litellm/proxy/prompts/prompt_endpoints.py index 1d71ea658e4..a289ed7cbfb 100644 --- a/litellm/proxy/prompts/prompt_endpoints.py +++ b/litellm/proxy/prompts/prompt_endpoints.py @@ -93,7 +93,7 @@ class _PromptTableActions(Protocol): def create(self, *, data: Mapping[str, str | int | None]) -> Awaitable[_PromptRow]: ... - def update(self, *, where: Mapping[str, str | int], data: Mapping[str, str]) -> Awaitable[_PromptRow]: ... + def update(self, *, where: Mapping[str, str | int], data: Mapping[str, str]) -> Awaitable[_PromptRow | None]: ... def delete_many(self, *, where: Mapping[str, str]) -> Awaitable[int]: ... @@ -1157,6 +1157,12 @@ async def patch_prompt( data=update_data, ) + if updated_prompt_db_entry is None: + raise HTTPException( + status_code=404, + detail=f"Prompt with ID {base_prompt_id} not found in environment {env}", + ) + updated_prompt_spec: Final = create_versioned_prompt_spec(db_prompt=updated_prompt_db_entry) return _reload_prompt_in_registry(IN_MEMORY_PROMPT_REGISTRY, versioned_id, updated_prompt_spec) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7dced4e26b6..765c201953d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -133,6 +133,7 @@ if TYPE_CHECKING: from aiohttp import ClientSession from fastapi.routing import APIRoute from opentelemetry.trace import Span as _Span + from prisma import models as prisma_models from litellm.integrations.opentelemetry import OpenTelemetry @@ -634,6 +635,7 @@ from litellm.proxy.utils import ( from litellm.proxy.video_endpoints.endpoints import router as video_router from litellm.repositories.base_repository import SupportsModelDump from litellm.repositories.credentials_repository import CredentialsRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.router import ( AssistantsTypedDict, Deployment, @@ -1642,12 +1644,21 @@ class _InvitationLinkRow(Protocol): class _UserTableRow(Protocol): user_id: str user_email: str | None - user_role: str + user_role: str | None -class _ModelTableRow(Protocol): - model_id: str | None - created_by: str | None +class _UserTeamsRow(Protocol): + @property + def teams(self) -> Sequence[str]: ... + + +_ProxyModelRow: TypeAlias = "prisma_models.LiteLLM_ProxyModelTable" + + +def _config_param_table(client: PrismaClient | None) -> TableActions[_ConfigParamRow]: + return cast( # cast-ok: this is prisma's LiteLLM_Config actions object, which parses its Json column to a mapping + "TableActions[_ConfigParamRow]", ConfigRepository(client).table + ) class _TTFTRow(TypedDict): @@ -4370,7 +4381,7 @@ class ProxyConfig: if prisma_client is None or not (general_settings.get("store_model_in_db", False) is True or store_model_in_db): return - row: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + row: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "environment_variables"} ) existing: Final[dict] = dict(row.param_value) if row is not None and row.param_value is not None else {} @@ -6226,7 +6237,7 @@ class ProxyConfig: 4. Update router settings """ if llm_router is not None and prisma_client is not None: - db_router_settings: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + db_router_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "router_settings"} ) @@ -6654,7 +6665,7 @@ class ProxyConfig: def _should_load_db_object(self, object_type: str | SupportedDBObjectType) -> bool: return should_load_db_object(object_type=object_type) - async def _get_models_from_db(self, prisma_client: PrismaClient) -> list | None: + async def _get_models_from_db(self, prisma_client: PrismaClient) -> Sequence[_ProxyModelRow] | None: """ Fetch all model deployments from the DB. @@ -6664,7 +6675,7 @@ class ProxyConfig: as "all models deleted" and must not evict existing router deployments. """ try: - new_models: Final[list[_ModelTableRow]] = await ModelRepository(prisma_client).table.find_many() + new_models: Final[Sequence[_ProxyModelRow]] = await ModelRepository(prisma_client).table.find_many() return new_models except Exception as e: verbose_proxy_logger.exception( @@ -6950,10 +6961,13 @@ class ProxyConfig: """ try: - sso_settings: Final[_SSOConfigRow | None] = await call_with_db_reconnect_retry( - prisma_client, - lambda: SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}), - reason="init_sso_settings_in_db_lookup_failure", + sso_settings: Final[_SSOConfigRow | None] = cast( # cast-ok: prisma Json stub is `str`, runtime is a dict + "_SSOConfigRow | None", + await call_with_db_reconnect_retry( + prisma_client, + lambda: SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}), + reason="init_sso_settings_in_db_lookup_failure", + ), ) if sso_settings is not None: sso_settings.sso_settings.pop("role_mappings", None) @@ -6981,12 +6995,15 @@ class ProxyConfig: ) try: - db_record: Final[_ConfigOverridesRow | None] = await call_with_db_reconnect_retry( - prisma_client, - lambda: ConfigOverridesRepository(prisma_client).table.find_unique( - where={"config_type": "hashicorp_vault"} + db_record: Final[_ConfigOverridesRow | None] = cast( # cast-ok: prisma Json stub is `str`, runtime dict + "_ConfigOverridesRow | None", + await call_with_db_reconnect_retry( + prisma_client, + lambda: ConfigOverridesRepository(prisma_client).table.find_unique( + where={"config_type": "hashicorp_vault"} + ), + reason="init_hashicorp_vault_config_override_lookup_failure", ), - reason="init_hashicorp_vault_config_override_lookup_failure", ) if db_record is None or db_record.config_value is None: @@ -8834,8 +8851,9 @@ class ProxyStartupEvent: if prisma_client is None: return - db_record: Final[_UISettingsRow | None] = await UISettingsRepository(prisma_client).table.find_unique( - where={"id": "ui_settings"} + db_record: Final[_UISettingsRow | None] = cast( # cast-ok: prisma Json stub is `str`, runtime is a dict + "_UISettingsRow | None", + await UISettingsRepository(prisma_client).table.find_unique(where={"id": "ui_settings"}), ) if db_record and db_record.ui_settings: raw: Final = db_record.ui_settings @@ -8998,7 +9016,7 @@ class ProxyStartupEvent: # but YAML config has False. if store_model_in_db is not True and prisma_client is not None: try: - _db_gs_record: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + _db_gs_record: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) if _db_gs_record is not None and isinstance(_db_gs_record.param_value, dict): @@ -12143,14 +12161,14 @@ async def _check_if_model_is_user_added( id = model.get("model_info", {}).get("id", None) if id is None: continue - db_model: _ModelTableRow | None = await ModelRepository(prisma_client).table.find_unique(where={"model_id": id}) + db_model: _ProxyModelRow | None = await ModelRepository(prisma_client).table.find_unique(where={"model_id": id}) if db_model is not None: if db_model.created_by == user_api_key_dict.user_id: filtered_models.append(model) return filtered_models -def _check_if_model_is_team_model(models: list[DeploymentTypedDict], user_row: LiteLLM_UserTable) -> list[dict]: +def _check_if_model_is_team_model(models: list[DeploymentTypedDict], user_row: _UserTeamsRow) -> list[dict]: """ Check if model is a team model @@ -12202,6 +12220,9 @@ async def non_admin_all_models( except Exception: raise HTTPException(status_code=400, detail={"error": "User not found"}) + if user_row is None: + raise HTTPException(status_code=400, detail={"error": "User not found"}) + # Get all models that are team models, when model team_id == user_row.teams all_models += _check_if_model_is_team_model( models=llm_router.get_model_list() or [], @@ -12630,7 +12651,7 @@ async def _fetch_db_models_for_search( db_models_total_count: Final = await ModelRepository(prisma_client).table.count(where=db_where_condition) - db_models_raw: list = [] + db_models_raw: Sequence[_ProxyModelRow] = [] if take_limit > 0: db_models_raw = await ModelRepository(prisma_client).table.find_many( where=db_where_condition, @@ -13020,7 +13041,7 @@ async def _gather_team_accessible_model_ids( try: if team_object.models and SpecialModelNames.all_proxy_models.value not in team_object.models: _resolved_names: Final = _team_models_resolve_to_names(team_object.models, access_groups) - db_models: Final[Sequence[_ModelTableRow]] = await ModelRepository(prisma_client).table.find_many( + db_models: Final[Sequence[_ProxyModelRow]] = await ModelRepository(prisma_client).table.find_many( where={"model_name": {"in": _resolved_names}} ) for db_model in db_models: @@ -14494,14 +14515,18 @@ async def alerting_settings( ) ## get general settings from db - db_general_settings: Final = await ConfigRepository(prisma_client).table.find_first( + db_general_settings: Final = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) if db_general_settings is not None and db_general_settings.param_value is not None: db_general_settings_dict: Final = dict(db_general_settings.param_value) - alerting_args_dict: dict = db_general_settings_dict.get("alerting_args", {}) - alerting_values: list | None = db_general_settings_dict.get("alerting") + alerting_args_dict: dict = cast( # cast-ok: ConfigGeneralSettings validates alerting_args as a dict on write + dict[str, JsonValue], db_general_settings_dict.get("alerting_args", {}) + ) + alerting_values: list | None = cast( # cast-ok: ConfigGeneralSettings validates alerting as a list on write + list[JsonValue] | None, db_general_settings_dict.get("alerting") + ) else: alerting_args_dict = {} alerting_values = None @@ -15052,7 +15077,7 @@ async def onboarding(invite_link: str, request: Request): user_id=user_obj.user_id, key=onboarding_token, user_email=user_obj.user_email, - user_role=user_obj.user_role, + user_role=user_obj.user_role, # pyright: ignore[reportArgumentType] # nullable DB column, no unset contract login_method="username_password", premium_user=premium_user, auth_header_name=general_settings.get("litellm_key_header_name", "Authorization"), @@ -15161,7 +15186,7 @@ async def _generate_onboarding_ui_session_token(user_obj: _UserTableRow) -> str: user_id=user_obj.user_id, key=key, user_email=user_obj.user_email, - user_role=user_obj.user_role, + user_role=user_obj.user_role, # pyright: ignore[reportArgumentType] # nullable DB column, no unset contract login_method="username_password", premium_user=premium_user, auth_header_name=general_settings.get("litellm_key_header_name", "Authorization"), @@ -15722,7 +15747,7 @@ async def update_config( raise Exception("No DB Connected") async def _read_section(param_name: str) -> dict: - row: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + row: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": param_name} ) if row is None or row.param_value is None: @@ -15979,7 +16004,7 @@ async def update_config_general_settings( ) ## get general settings from db - db_general_settings: Final = await ConfigRepository(prisma_client).table.find_first( + db_general_settings: Final = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) ### update value @@ -15997,7 +16022,7 @@ async def update_config_general_settings( if data.field_name == "plugins": field_value = _preserve_redacted_plugin_keys(field_value, general_settings.get("plugins")) - general_settings[data.field_name] = field_value + general_settings[data.field_name] = cast(JsonValue, field_value) # cast-ok: ConfigGeneralSettings validated it response: Final = await ConfigRepository(prisma_client).table.upsert( where={"param_name": "general_settings"}, @@ -16017,7 +16042,7 @@ async def update_config_general_settings( ) if data.field_name == "plugins": - register_plugins_from_config(general_settings) + register_plugins_from_config(cast(dict[str, object], general_settings)) # cast-ok: the callee only reads it _apply_ssrf_general_settings(general_settings) return response @@ -16193,7 +16218,7 @@ async def get_config_general_settings( ) ## get general settings from db - db_general_settings: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + db_general_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) ### pop the value @@ -16382,7 +16407,7 @@ async def get_config_list( is_full_admin: Final = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN ## get general settings from db - db_general_settings: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + db_general_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) @@ -16478,7 +16503,7 @@ async def get_config_list( ) return_val.append(_response_obj) - db_litellm_settings_row: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + db_litellm_settings_row: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "litellm_settings"} ) db_litellm_settings: Final[dict] = ( @@ -16555,7 +16580,7 @@ async def delete_config_general_settings( ) ## get general settings from db - db_general_settings: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_first( + db_general_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "general_settings"} ) ### pop the value @@ -17122,7 +17147,7 @@ async def reload_anthropic_beta_headers( last_anthropic_beta_headers_reload = current_time.isoformat() # Set force reload flag in database for other pods, preserving existing interval_hours - existing_beta_config: Final[_ConfigParamRow | None] = await ConfigRepository(prisma_client).table.find_unique( + existing_beta_config: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_unique( where={"param_name": "anthropic_beta_headers_reload_config"} ) existing_beta_interval = None @@ -17300,7 +17325,7 @@ async def get_anthropic_beta_headers_reload_status( } # Get reload configuration from database - config_record: Final = await ConfigRepository(prisma_client).table.find_unique( + config_record: Final = await _config_param_table(prisma_client).find_unique( where={"param_name": "anthropic_beta_headers_reload_config"} ) @@ -17314,7 +17339,9 @@ async def get_anthropic_beta_headers_reload_status( } config: Final = config_record.param_value - interval_hours: Final = config.get("interval_hours") + interval_hours: Final = cast( # cast-ok: every writer of this key stores `hours: int` or an explicit None + int | None, config.get("interval_hours") + ) if interval_hours is None: verbose_proxy_logger.info("No interval configured, returning not scheduled") diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index 93c81f7fc67..6c0d11a2174 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -1,5 +1,11 @@ import json -from typing import Final +from collections.abc import Mapping +from typing import ( + TYPE_CHECKING, + Final, + Protocol, + cast, # noqa: TID251 # the config repository's table protocol omits find_first +) from fastapi import APIRouter, Depends, HTTPException @@ -13,6 +19,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.cloudzero_endpoints import ( CloudZeroExportRequest, CloudZeroExportResponse, @@ -22,6 +29,9 @@ from litellm.types.proxy.cloudzero_endpoints import ( CloudZeroSettingsView, ) +if TYPE_CHECKING: + from litellm.proxy.proxy_server import PrismaClient + router: Final = APIRouter() @@ -29,6 +39,18 @@ router: Final = APIRouter() _sensitive_masker: Final = SensitiveDataMasker() +class _CloudZeroConfigRow(Protocol): + """The ``LiteLLM_Config`` row holding ``cloudzero_settings``, as this module reads it.""" + + @property + def param_value(self) -> str | Mapping[str, str] | None: ... + + +def _config_table(prisma_client: "PrismaClient") -> TableActions[_CloudZeroConfigRow]: + repository_table: Final = ConfigRepository(prisma_client).table + return cast(TableActions[_CloudZeroConfigRow], repository_table) # cast-ok: repo protocol omits find_first + + async def _set_cloudzero_settings(api_key: str, connection_id: str, timezone: str): """ Store CloudZero settings in the database with encrypted API key. @@ -82,9 +104,7 @@ async def _get_cloudzero_settings(): detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - cloudzero_config: Final = await ConfigRepository(prisma_client).table.find_first( - where={"param_name": "cloudzero_settings"} - ) + cloudzero_config: Final = await _config_table(prisma_client).find_first(where={"param_name": "cloudzero_settings"}) if cloudzero_config is None or cloudzero_config.param_value is None: return {} @@ -268,7 +288,7 @@ async def is_cloudzero_setup_in_db() -> bool: return False # Check for CloudZero settings in database - cloudzero_config: Final = await ConfigRepository(prisma_client).table.find_first( + cloudzero_config: Final = await _config_table(prisma_client).find_first( where={"param_name": "cloudzero_settings"} ) @@ -530,7 +550,7 @@ async def delete_cloudzero_settings( ) # Check if CloudZero settings exist - cloudzero_config: Final = await ConfigRepository(prisma_client).table.find_first( + cloudzero_config: Final = await _config_table(prisma_client).find_first( where={"param_name": "cloudzero_settings"} ) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index ed2ecd8325a..06395a3c3cc 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -4,10 +4,22 @@ import json import os from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple, Protocol, TypedDict, TypeVar +from typing import ( + TYPE_CHECKING, + Annotated, + Any, + Final, + Literal, + NamedTuple, + Protocol, + TypedDict, + TypeVar, + cast, # noqa: TID251 # prisma group_by returns untyped aggregate mappings +) import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, status +from typing_extensions import ReadOnly import litellm from litellm._logging import verbose_proxy_logger @@ -23,6 +35,7 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( get_spend_by_team_and_customer, ) from litellm.proxy.utils import handle_exception_on_proxy +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import SpendLogsRepository from litellm.repositories.team_repository import TeamRepository from litellm.repositories.verification_token_repository import ( @@ -30,6 +43,8 @@ from litellm.repositories.verification_token_repository import ( ) if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.proxy.proxy_server import PrismaClient from litellm.proxy.spend_tracking.cold_storage_handler import ColdStorageHandler else: @@ -139,6 +154,18 @@ class _SessionSpendRow(TypedDict): mcp_tool_call_spend: float +class _SpendSumAggregate(TypedDict, total=False): + spend: ReadOnly[float] + + +class _SpendGroupByRow(TypedDict): + api_key: ReadOnly[str] + user: ReadOnly[str | None] + model: ReadOnly[str] + startTime: ReadOnly[object] + _sum: ReadOnly[_SpendSumAggregate] + + async def _query_raw(prisma_client: PrismaClient, sql_query: str, *args: object) -> Sequence[_RowT]: """Run a raw read query and return its rows as the row type the caller declares.""" return await prisma_client.db.query_raw(sql_query, *args) @@ -149,24 +176,6 @@ async def _query_raw_or_none(prisma_client: PrismaClient, sql_query: str, *args: return await _query_raw(prisma_client, sql_query, *args) -class _SpendLogsTable(Protocol): - """The subset of the Prisma spend-logs table API this module uses.""" - - async def find_many( - self, *, where: Mapping[str, object], order: Mapping[str, str] - ) -> Sequence[_SupportsModelDump]: ... - - async def find_unique( - self, *, where: Mapping[str, object], include: None = None - ) -> _SpendLogOwnershipRow | None: ... - - async def count(self, *, where: Mapping[str, object]) -> int: ... - - async def group_by( - self, *, by: Sequence[str], where: Mapping[str, object], count: Mapping[str, bool] - ) -> Sequence[_SessionCountRow]: ... - - class _TeamTable(Protocol): """The subset of the Prisma team table API this module uses.""" @@ -183,7 +192,7 @@ class _VerificationTokenTable(Protocol): async def update_many(self, *, data: Mapping[str, float], where: Mapping[str, object]) -> int: ... -def _spend_logs_table(prisma_client: PrismaClient) -> _SpendLogsTable: +def _spend_logs_table(prisma_client: PrismaClient) -> TableActions["prisma_models.LiteLLM_SpendLogs"]: return SpendLogsRepository(prisma_client).table @@ -221,11 +230,12 @@ async def _count_logs_per_session( prisma_client: PrismaClient, session_ids: Sequence[str | None] ) -> Sequence[_SessionCountRow]: """Count spend log rows per session for the given session ids.""" - return await _spend_logs_table(prisma_client).group_by( + rows: Final = await _spend_logs_table(prisma_client).group_by( by=["session_id"], where={"session_id": {"in": session_ids}}, count={"session_id": True}, ) + return cast(Sequence[_SessionCountRow], rows) # cast-ok: group_by(count=) shape is fixed by the by/count args async def _find_team_row(prisma_client: PrismaClient, team_id: str) -> _SupportsModelDump | None: @@ -2974,8 +2984,9 @@ async def view_spend_logs( ) 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 response: + 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: diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index 00d0554d783..c71105ad283 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -1,5 +1,11 @@ import json -from typing import Final +from collections.abc import Mapping +from typing import ( + TYPE_CHECKING, + Final, + Protocol, + cast, # noqa: TID251 # the config repository's table protocol omits find_first +) from fastapi import APIRouter, Depends, HTTPException @@ -14,6 +20,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( ) from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.vantage_endpoints import ( VantageDryRunRequest, VantageExportRequest, @@ -24,6 +31,9 @@ from litellm.types.proxy.vantage_endpoints import ( VantageSettingsView, ) +if TYPE_CHECKING: + from litellm.proxy.proxy_server import PrismaClient + router: Final = APIRouter() _sensitive_masker: Final = SensitiveDataMasker() @@ -31,6 +41,18 @@ _sensitive_masker: Final = SensitiveDataMasker() VANTAGE_SETTINGS_PARAM_NAME: Final = "vantage_settings" +class _VantageConfigRow(Protocol): + """The ``LiteLLM_Config`` row holding ``vantage_settings``, as this module reads it.""" + + @property + def param_value(self) -> str | Mapping[str, str] | None: ... + + +def _config_table(prisma_client: "PrismaClient") -> TableActions[_VantageConfigRow]: + repository_table: Final = ConfigRepository(prisma_client).table + return cast(TableActions[_VantageConfigRow], repository_table) # cast-ok: repo protocol omits find_first + + def _get_registered_vantage_logger(): """Return the VantageLogger already registered in litellm.callbacks, if any.""" from litellm.integrations.vantage.vantage_logger import VantageLogger @@ -82,7 +104,7 @@ async def _get_vantage_settings(): detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - vantage_config: Final = await ConfigRepository(prisma_client).table.find_first( + vantage_config: Final = await _config_table(prisma_client).find_first( where={"param_name": VANTAGE_SETTINGS_PARAM_NAME} ) if vantage_config is None or vantage_config.param_value is None: @@ -251,7 +273,7 @@ async def is_vantage_setup_in_db() -> bool: if prisma_client is None: return False - vantage_config: Final = await ConfigRepository(prisma_client).table.find_first( + vantage_config: Final = await _config_table(prisma_client).find_first( where={"param_name": VANTAGE_SETTINGS_PARAM_NAME} ) @@ -525,7 +547,7 @@ async def delete_vantage_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - vantage_config: Final = await ConfigRepository(prisma_client).table.find_first( + vantage_config: Final = await _config_table(prisma_client).find_first( where={"param_name": VANTAGE_SETTINGS_PARAM_NAME} ) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 66a8c0622fa..a1eb7ed06eb 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -4,7 +4,12 @@ import json import os from collections import Counter from collections.abc import Mapping -from typing import Any, Final, Protocol, TypeVar +from typing import ( + Any, + Final, + Protocol, + cast, # noqa: TID251 # prisma types Json columns as fields.Json but de-serializes them to plain python on read +) from urllib.parse import urlparse from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile @@ -25,6 +30,7 @@ from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attributio from litellm.proxy.utils import invalidate_config_param from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.organization_repository import OrganizationRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( SSOConfigRepository, UISettingsRepository, @@ -37,29 +43,16 @@ from litellm.types.proxy.management_endpoints.ui_sso import ( router: Final = APIRouter() -_DbRecordT: Final = TypeVar("_DbRecordT", covariant=True) - - -class _PrismaTableActions(Protocol[_DbRecordT]): - async def find_unique(self, where: Mapping[str, object]) -> _DbRecordT | None: ... - - async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> _DbRecordT: ... - - async def upsert(self, where: Mapping[str, object], data: Mapping[str, object]) -> _DbRecordT: ... - class _SsoSettingsMappingRow(Protocol): @property def sso_settings(self) -> Mapping[str, object] | None: ... -class _HasSsoSettingsMappingTable(Protocol): - @property - def table(self) -> _PrismaTableActions[_SsoSettingsMappingRow]: ... - - -def _sso_settings_mapping_db(repo: _HasSsoSettingsMappingTable) -> _PrismaTableActions[_SsoSettingsMappingRow]: - return repo.table +def _sso_settings_mapping_db(repo: SSOConfigRepository) -> TableActions[_SsoSettingsMappingRow]: + return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value + "TableActions[_SsoSettingsMappingRow]", repo.table + ) class _StoredSsoSettingsRow(Protocol): @@ -67,12 +60,7 @@ class _StoredSsoSettingsRow(Protocol): def sso_settings(self) -> object: ... -class _HasStoredSsoSettingsTable(Protocol): - @property - def table(self) -> _PrismaTableActions[_StoredSsoSettingsRow]: ... - - -def _stored_sso_settings_db(repo: _HasStoredSsoSettingsTable) -> _PrismaTableActions[_StoredSsoSettingsRow]: +def _stored_sso_settings_db(repo: SSOConfigRepository) -> TableActions[_StoredSsoSettingsRow]: return repo.table @@ -81,13 +69,10 @@ class _UiSettingsRow(Protocol): def ui_settings(self) -> str | Mapping[str, JsonValue] | None: ... -class _HasUiSettingsTable(Protocol): - @property - def table(self) -> _PrismaTableActions[_UiSettingsRow]: ... - - -def _ui_settings_db(repo: _HasUiSettingsTable) -> _PrismaTableActions[_UiSettingsRow]: - return repo.table +def _ui_settings_db(repo: UISettingsRepository) -> TableActions[_UiSettingsRow]: + return cast( # cast-ok: prisma types Json columns as str; the client hands back the deserialized value + "TableActions[_UiSettingsRow]", repo.table + ) class _ConfigParamRow(Protocol): @@ -95,13 +80,10 @@ class _ConfigParamRow(Protocol): def param_value(self) -> str | Mapping[str, object] | None: ... -class _HasConfigParamTable(Protocol): - @property - def table(self) -> _PrismaTableActions[_ConfigParamRow]: ... - - -def _config_param_db(repo: _HasConfigParamTable) -> _PrismaTableActions[_ConfigParamRow]: - return repo.table +def _config_param_db(repo: ConfigRepository) -> TableActions[_ConfigParamRow]: + return cast( # cast-ok: prisma's LiteLLM_Config actions object, whose Json column parses to a mapping + "TableActions[_ConfigParamRow]", repo.table + ) # Maps each UIThemeConfig field to the env var the UI branding path reads it diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 86d954c0913..9f9dd4e6af9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -177,6 +177,7 @@ from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams if TYPE_CHECKING: from mcp.types import CallToolResult from opentelemetry.trace import Span as _Span + from prisma import models as prisma_models from prisma.actions import LiteLLM_DeprecatedVerificationTokenActions from prisma.client import TransactionManager from prisma.models import LiteLLM_DeprecatedVerificationToken @@ -186,6 +187,7 @@ if TYPE_CHECKING: from litellm.models.team import LiteLLM_TeamTableCachedObj from litellm.proxy.db.autorouter_session_rollup import AutoRouterTurnTransaction from litellm.proxy.db.spend_log_tool_index import ToolUsageTransaction + from litellm.repositories.prisma_protocols import TableActions from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline Span = _Span | object @@ -3266,7 +3268,10 @@ async def prefetch_config_params(prisma_client: "PrismaClient | None", param_nam if not param_names: return try: - rows: Final = await ConfigRepository(prisma_client).table.find_many(where={"param_name": {"in": param_names}}) + config_table: Final = cast( # cast-ok: ConfigRepository.table is prisma's litellm_config actions object + "TableActions[prisma_models.LiteLLM_Config]", ConfigRepository(prisma_client).table + ) + rows: Final = await config_table.find_many(where={"param_name": {"in": param_names}}) except Exception as e: verbose_proxy_logger.debug( "prefetch_config_params failed, falling through to per-param queries: %s", @@ -3555,8 +3560,8 @@ class PrismaClient: return hashed_token - def jsonify_object(self, data: dict) -> dict: - db_data: Final = copy.deepcopy(data) + def jsonify_object(self, data: Mapping[str, object]) -> dict[str, object]: + db_data: Final[dict[str, object]] = copy.deepcopy(dict(data)) for k, v in db_data.items(): if isinstance(v, dict): @@ -3690,7 +3695,10 @@ class PrismaClient: elif table_name == "keys": return await VerificationTokenRepository(self).table.find_first(where={key: value}) elif table_name == "config": - return await ConfigRepository(self).table.find_first(where={key: value}) + config_table: Final = cast( # cast-ok: ConfigRepository.table is prisma's litellm_config actions object + "TableActions[prisma_models.LiteLLM_Config]", ConfigRepository(self).table + ) + return await config_table.find_first(where={key: value}) elif table_name == "spend": return await self.db.l.find_first(where={key: value}) return None @@ -3793,9 +3801,9 @@ class PrismaClient: self, token: str | list | None = None, user_id: str | None = None, - user_id_list: list | None = None, + user_id_list: Sequence[str] | None = None, team_id: str | None = None, - team_id_list: list | None = None, + team_id_list: Sequence[str] | None = None, key_val: dict | None = None, table_name: Literal[ "user", "key", "config", "spend", "enduser", "budget", "team", "user_notification", "combined_view" @@ -3878,14 +3886,14 @@ class PrismaClient: if isinstance(r.expires, datetime): r.expires = r.expires.isoformat() elif query_type == "find_all": - where_filter: Final[dict] = {} + where_filter: Final[dict[str, dict[str, Sequence[str]]]] = {} if token is not None: where_filter["token"] = {} if isinstance(token, str): token = _hash_token_if_needed(token=token) where_filter["token"]["in"] = [token] elif isinstance(token, list): - hashed_tokens: Final = [] + hashed_tokens: Final[list[str]] = [] for t in token: assert isinstance(t, str) if t.startswith("sk-"): @@ -4182,7 +4190,7 @@ class PrismaClient: ) raise e - def jsonify_team_object(self, db_data: dict): + def jsonify_team_object(self, db_data: Mapping[str, object]) -> dict[str, object]: db_data = self.jsonify_object(data=db_data) if db_data.get("members_with_roles", None) is not None and isinstance(db_data["members_with_roles"], list): db_data["members_with_roles"] = json.dumps(db_data["members_with_roles"]) @@ -4200,7 +4208,7 @@ class PrismaClient: ) async def insert_data( self, - data: dict, + data: Mapping[str, object], table_name: Literal["user", "key", "config", "spend", "team", "user_notification"], ): """ @@ -4210,10 +4218,12 @@ class PrismaClient: try: verbose_proxy_logger.debug( "PrismaClient: insert_data: %s", - {**data, "token": self.hash_token(token=data["token"])} if data.get("token") is not None else data, + {**data, "token": self.hash_token(token=cast("str", data["token"]))} # cast-ok: a key token is a str + if data.get("token") is not None + else data, ) if table_name == "key": - token: Final = data["token"] + token: Final = cast("str", data["token"]) # cast-ok: the key table's token column is a str hashed_token: Final = self.hash_token(token=token) db_data = self.jsonify_object(data=data) db_data["token"] = hashed_token @@ -4348,14 +4358,14 @@ class PrismaClient: async def update_data( self, token: str | None = None, - data: dict = {}, + data: Mapping[str, object] = {}, data_list: list | None = None, user_id: str | None = None, team_id: str | None = None, query_type: Literal["update", "update_many"] = "update", table_name: Literal["user", "key", "config", "spend", "team", "enduser", "budget"] | None = None, - update_key_values: dict | None = None, - update_key_values_custom_query: dict | None = None, + update_key_values: dict[str, object] | None = None, + update_key_values_custom_query: dict[str, object] | None = None, ): """ Update existing data @@ -4381,14 +4391,14 @@ class PrismaClient: try: _data = response.model_dump() except Exception: - _data = response.dict() + _data = response.dict() # pyright: ignore[reportDeprecated] # pydantic-v1 row fallback return {"token": token, "data": _data} elif user_id is not None or (table_name is not None and table_name == "user") and query_type == "update": """ If data['spend'] + data['user'], update the user table with spend info as well """ if user_id is None: - user_id = db_data["user_id"] + user_id = cast("str", db_data["user_id"]) # cast-ok: the user table's user_id column is a str if update_key_values is None: if update_key_values_custom_query is not None: update_key_values = update_key_values_custom_query @@ -4410,7 +4420,7 @@ class PrismaClient: If data['spend'] + data['user'], update the user table with spend info as well """ if team_id is None: - team_id = db_data["team_id"] + team_id = cast("str | None", db_data["team_id"]) # cast-ok: team_id column is a nullable str if update_key_values is None: update_key_values = db_data if "team_id" not in db_data and team_id is not None: @@ -4584,8 +4594,8 @@ class PrismaClient: ) async def delete_data( self, - tokens: list | None = None, - team_id_list: list | None = None, + tokens: Sequence[str | None] | None = None, + team_id_list: Sequence[str] | None = None, table_name: Literal["user", "key", "config", "spend", "team"] | None = None, user_id: str | None = None, ): @@ -4597,14 +4607,14 @@ class PrismaClient: start_time: Final = time.time() try: if tokens is not None and isinstance(tokens, list): - hashed_tokens: Final = [] + hashed_tokens: Final[list[str | None]] = [] for token in tokens: if isinstance(token, str) and token.startswith("sk-"): hashed_token = self.hash_token(token=token) else: hashed_token = token hashed_tokens.append(hashed_token) - filter_query: dict = {} + filter_query: dict[str, object] = {} if user_id is not None: filter_query = {"AND": [{"token": {"in": hashed_tokens}}, {"user_id": user_id}]} else: @@ -5749,12 +5759,12 @@ class PrismaClient: limit: int = 100, offset: int = 0, status_filter: str | None = None, - ): + ) -> "Sequence[prisma_models.LiteLLM_HealthCheckTable]": """ Get health check history with optional filtering """ try: - where_clause: Final = {} + where_clause: Final[dict[str, str]] = {} if model_name: where_clause["model_name"] = model_name if status_filter: @@ -5771,7 +5781,7 @@ class PrismaClient: verbose_proxy_logger.error("Error getting health check history: %s", e) return [] - async def get_all_latest_health_checks(self): + async def get_all_latest_health_checks(self) -> "Sequence[prisma_models.LiteLLM_HealthCheckTable]": """ Get the latest health check for each model. @@ -5949,15 +5959,17 @@ async def migrate_passwords_to_scrypt_async(prisma_client) -> str: return len(s) == 64 and all(c in "0123456789abcdef" for c in s) plaintext_users: Final = [ - u for u in all_with_pw if u.password and not u.password.startswith("scrypt:") and not _is_sha256_hex(u.password) + (u.user_id, u.password) + for u in all_with_pw + if u.password and not u.password.startswith("scrypt:") and not _is_sha256_hex(u.password) ] if not plaintext_users: return "No plaintext passwords found" - for user in plaintext_users: + for user_id, plaintext_password in plaintext_users: await UserRepository(prisma_client).table.update( - where={"user_id": user.user_id}, - data={"password": hash_password(user.password)}, + where={"user_id": user_id}, + data={"password": hash_password(plaintext_password)}, ) return f"Migrated {len(plaintext_users)} plaintext passwords to scrypt" diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index b497247f576..a59d7a277cc 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -1,4 +1,9 @@ -from typing import Annotated, Any, Final +from typing import ( + Annotated, + Any, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict + Final, + cast, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict +) from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -591,7 +596,11 @@ async def index_create( index_data: Final = index_create_request.model_dump(exclude_none=True) index_data["created_by"] = user_api_key_dict.user_id index_data["updated_by"] = user_api_key_dict.user_id - new_index = await ManagedVectorStoreIndexRepository(prisma_client).table.create(data=jsonify_object(index_data)) + new_index = await ManagedVectorStoreIndexRepository(prisma_client).table.create( + data=cast( # cast-ok: jsonify_object deep-copies a model_dump, so keys are str and values plain objects + "dict[str, object]", jsonify_object(index_data) + ) + ) return new_index.model_dump() diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 2b037bef795..183a03cc13c 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -10,8 +10,7 @@ All /vector_store management endpoints import copy import json -from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final from fastapi import APIRouter, Depends, HTTPException @@ -37,6 +36,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helpe from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store from litellm.repositories.model_repository import ModelRepository +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ManagedVectorStoresRepository from litellm.secret_managers.main import get_secret from litellm.types.vector_stores import ( @@ -51,17 +51,7 @@ from litellm.vector_stores.vector_store_registry import VectorStoreRegistry router: Final = APIRouter() -class _VectorStoreTableActions(Protocol): - async def find_unique(self, where: Mapping[str, str]) -> "_VectorStoreRow | None": ... - - async def create(self, data: Mapping[str, object]) -> "_VectorStoreRow": ... - - async def update(self, where: Mapping[str, str], data: Mapping[str, object]) -> "_VectorStoreRow": ... - - async def delete(self, where: Mapping[str, str]) -> "_VectorStoreRow | None": ... - - -def _vector_store_table(prisma_client: "PrismaClient") -> _VectorStoreTableActions: +def _vector_store_table(prisma_client: "PrismaClient") -> "TableActions[_VectorStoreRow]": return ManagedVectorStoresRepository(prisma_client).table @@ -277,7 +267,7 @@ async def _resolve_embedding_config_from_db( if db_model and db_model.litellm_params: # Extract litellm_params (could be dict or JSON string) model_params = db_model.litellm_params - if isinstance(model_params, str): + if isinstance(model_params, str): # pyright: ignore[reportUnnecessaryIsInstance] # prisma Json is str model_params = json.loads(model_params) # Decrypt values from database (similar to how proxy_server.py does it) @@ -888,6 +878,12 @@ async def update_vector_store( data=update_data, ) + if updated is None: + raise HTTPException( + status_code=404, + detail=f"Vector store with ID {vector_store_id} not found", + ) + updated_vs: Final = _row_to_vector_store(updated) # Immediately update in-memory registry to keep it in sync diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py index 7008099fe8c..568e3b50ed2 100644 --- a/litellm/repositories/base_repository.py +++ b/litellm/repositories/base_repository.py @@ -8,6 +8,8 @@ from typing import Any, Final, Generic, Protocol, TypeVar, runtime_checkable from pydantic import BaseModel +from litellm.repositories.prisma_protocols import TableActions + T = TypeVar("T", bound=BaseModel) @@ -49,7 +51,7 @@ class BaseRepository(ABC, Generic[T]): @property @abstractmethod - def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper + def table(self) -> TableActions[DbRecord]: """Return the Prisma table for this repository.""" ... @@ -76,33 +78,28 @@ class BaseRepository(ABC, Generic[T]): async def find_many( self, - where: dict[str, Any] | None = None, + where: Mapping[str, object] | None = None, skip: int | None = None, take: int | None = None, - order: dict[str, str] | None = None, + order: Mapping[str, str] | None = None, ) -> list[T]: """Find multiple records matching the criteria.""" - kwargs: Final[dict[str, Any]] = {} - if where: - kwargs["where"] = where - if skip is not None: - kwargs["skip"] = skip - if take is not None: - kwargs["take"] = take - if order: - kwargs["order"] = order - - records: Final = await self.table.find_many(**kwargs) + records: Final = await self.table.find_many( + take=take, + skip=skip, + where=where or None, + order=order or None, + ) return self._to_model_list(records) - async def create(self, data: dict[str, Any]) -> T: + async def create(self, data: Mapping[str, object]) -> T: """Create a new record.""" record: Final = await self.table.create(data=data) model: Final = self._to_model(record) assert model is not None return model - async def update(self, id_value: str, data: dict[str, Any], id_field: str = "id") -> T | None: + async def update(self, id_value: str, data: Mapping[str, object], id_field: str = "id") -> T | None: """Update an existing record.""" record: Final = await self.table.update(where={id_field: id_value}, data=data) return self._to_model(record) @@ -112,7 +109,7 @@ class BaseRepository(ABC, Generic[T]): record: Final = await self.table.delete(where={id_field: id_value}) return self._to_model(record) - async def count(self, where: dict[str, Any] | None = None) -> int: + async def count(self, where: Mapping[str, object] | None = None) -> int: """Count records matching the criteria.""" return await self.table.count(where=where) diff --git a/litellm/repositories/budget_repository.py b/litellm/repositories/budget_repository.py index f6c47b2d639..62632ffb5f6 100644 --- a/litellm/repositories/budget_repository.py +++ b/litellm/repositories/budget_repository.py @@ -2,17 +2,21 @@ Budget repository for database operations on LiteLLM_BudgetTable. """ -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final from litellm.models.budget import LiteLLM_BudgetTable from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): """Repository for budget database operations.""" @property - def table(self) -> Any: + def table(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: return self.prisma_client.db.litellm_budgettable @property diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index 71ae39e89c6..76b9a3a5809 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -77,7 +77,7 @@ class ConfigRepository: return self.prisma_client.db.litellm_config @property - def table(self) -> Any: + def table(self) -> _ConfigTable: return self._config_table async def get_param(self, param_name: str) -> ConfigParam | None: diff --git a/litellm/repositories/credentials_repository.py b/litellm/repositories/credentials_repository.py index 9fdb6e4aca7..ddb9767b2b9 100644 --- a/litellm/repositories/credentials_repository.py +++ b/litellm/repositories/credentials_repository.py @@ -6,54 +6,77 @@ credential values is the caller's responsibility (see ``CredentialHelperUtils``) so reads return the stored values verbatim. """ -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias from litellm.models.credentials import CredentialItem from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync +from litellm.repositories.base_repository import DbRecord, record_to_dict +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models + + _CredentialsTable: TypeAlias = TableActions[prisma_models.LiteLLM_CredentialsTable] + + +class _PrismaCredentialsDb(Protocol): + @property + def litellm_credentialstable(self) -> "_CredentialsTable": ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _PrismaCredentialsDb: ... class CredentialsRepository: """Repository for credentials database operations, keyed by credential name.""" - def __init__(self, prisma_client: Any): + def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper self._prisma_client = prisma_client @property - def prisma_client(self) -> Any: + def prisma_client(self) -> _PrismaClientView: if self._prisma_client is None: raise RuntimeError("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys") - return self._prisma_client + client: Final[_PrismaClientView] = self._prisma_client + return client @property - def table(self) -> Any: + def table(self) -> "_CredentialsTable": return wrap_table_actions_for_config_sync( actions=self.prisma_client.db.litellm_credentialstable, table_name="litellm_credentialstable", ) @staticmethod - def _to_model(record: Any) -> CredentialItem | None: + def _to_model(record: DbRecord | None) -> CredentialItem | None: if record is None: return None - data: Final = record.dict() if hasattr(record, "dict") else dict(record) - return CredentialItem( - credential_name=data["credential_name"], - credential_values=data.get("credential_values") or {}, - credential_info=data.get("credential_info") or {}, + data: Final = record_to_dict(record) + return CredentialItem.model_validate( + { + "credential_name": data["credential_name"], + "credential_values": data.get("credential_values") or {}, + "credential_info": data.get("credential_info") or {}, + } ) - async def find_all(self) -> Any: + async def find_all(self) -> Sequence["prisma_models.LiteLLM_CredentialsTable"]: return await self.table.find_many() - async def create(self, data: dict[str, Any]) -> Any: + async def create(self, data: Mapping[str, object]) -> "prisma_models.LiteLLM_CredentialsTable": return await self.table.create(data=data) async def find_by_name(self, credential_name: str) -> CredentialItem | None: record: Final = await self.table.find_unique(where={"credential_name": credential_name}) return self._to_model(record) - async def update_by_name(self, credential_name: str, data: dict[str, Any]) -> Any: + async def update_by_name( + self, credential_name: str, data: Mapping[str, object] + ) -> "prisma_models.LiteLLM_CredentialsTable | None": return await self.table.update(where={"credential_name": credential_name}, data=data) - async def delete_by_name(self, credential_name: str) -> Any: + async def delete_by_name(self, credential_name: str) -> "prisma_models.LiteLLM_CredentialsTable | None": return await self.table.delete(where={"credential_name": credential_name}) diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index 27e23a39cc9..cc2b1f19a3c 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -3,8 +3,8 @@ Model repository for database operations on LiteLLM_ProxyModelTable. """ import json -from collections.abc import Awaitable, Mapping, Sequence -from typing import Any, Final, Protocol +from collections.abc import Mapping +from typing import TYPE_CHECKING, Any, Final, Protocol from litellm.models.model import LiteLLM_ProxyModelTable from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync @@ -12,25 +12,21 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) -from litellm.repositories.base_repository import BaseRepository, DbRecord +from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models class _PrismaModelDb(Protocol): - litellm_proxymodeltable: object + @property + def litellm_proxymodeltable(self) -> TableActions["prisma_models.LiteLLM_ProxyModelTable"]: ... class _PrismaClientView(Protocol): - db: _PrismaModelDb - - -class _ProxyModelActions(Protocol): - """Prisma table actions used by :class:`ModelRepository`.""" - - def find_many(self, *, where: Mapping[str, object] | None = None) -> Awaitable[Sequence[DbRecord]]: ... - - def create(self, *, data: Mapping[str, object]) -> Awaitable[DbRecord]: ... - - def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> Awaitable[DbRecord | None]: ... + @property + def db(self) -> _PrismaModelDb: ... class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): @@ -41,17 +37,13 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): self._encryption_key = encryption_key @property - def table(self) -> Any: + def table(self) -> TableActions["prisma_models.LiteLLM_ProxyModelTable"]: client: Final[_PrismaClientView] = self.prisma_client return wrap_table_actions_for_config_sync( actions=client.db.litellm_proxymodeltable, table_name="litellm_proxymodeltable", ) - @property - def _model_table(self) -> _ProxyModelActions: - return self.table - @property def model_class(self) -> type[LiteLLM_ProxyModelTable]: return LiteLLM_ProxyModelTable @@ -100,17 +92,17 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): async def find_by_name(self, model_name: str) -> list[LiteLLM_ProxyModelTable]: """Find models by name.""" - records: Final = await self._model_table.find_many(where={"model_name": model_name}) + records: Final = await self.table.find_many(where={"model_name": model_name}) return self._to_model_list(records) async def find_all(self) -> list[LiteLLM_ProxyModelTable]: """Find all models.""" - records: Final = await self._model_table.find_many() + records: Final = await self.table.find_many() return self._to_model_list(records) async def find_unblocked(self) -> list[LiteLLM_ProxyModelTable]: """Find all models that are not blocked.""" - records: Final = await self._model_table.find_many(where={"blocked": False}) + records: Final = await self.table.find_many(where={"blocked": False}) return self._to_model_list(records) async def find_by_team_id(self, team_id: str) -> list[LiteLLM_ProxyModelTable]: @@ -147,7 +139,7 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): if model_info is not None: data["model_info"] = json.dumps(model_info) - record: Final = await self._model_table.create(data=data) + record: Final = await self.table.create(data=data) model: Final = self._to_model(record) assert model is not None return model @@ -173,7 +165,7 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): if blocked is not None: data["blocked"] = blocked - record: Final = await self._model_table.update(where={"model_id": model_id}, data=data) + record: Final = await self.table.update(where={"model_id": model_id}, data=data) return self._to_model(record) async def delete_model(self, model_id: str) -> LiteLLM_ProxyModelTable | None: diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 54a311c4a77..6b1f9c68e47 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -2,17 +2,21 @@ ObjectPermission repository for database operations on LiteLLM_ObjectPermissionTable. """ -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]): """Repository for object permission database operations.""" @property - def table(self) -> Any: + def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]: return self.prisma_client.db.litellm_objectpermissiontable @property diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index 8a1350903b7..5a9bd3724e0 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -2,17 +2,21 @@ Organization repository for database operations on LiteLLM_OrganizationTable. """ -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final from litellm.models.organization import LiteLLM_OrganizationTable from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): """Repository for organization database operations.""" @property - def table(self) -> Any: + def table(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: return self.prisma_client.db.litellm_organizationtable @property diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 055c68163f9..2aa1b8e0e3f 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -12,6 +12,93 @@ from typing import Protocol, TypeVar RowT_co = TypeVar("RowT_co", covariant=True) +class TableActions(Protocol[RowT_co]): + """The prisma-client-py per-model action surface, keyed to the row it returns. + + Query inputs stay `Mapping[str, object]` rather than the generated + `types.*` TypedDicts so callers can keep passing plain dicts, while every + result carries the row type the repository is bound to. + """ + + async def find_unique( + self, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> RowT_co | None: ... + + async def find_first( + self, + skip: int | None = None, + where: Mapping[str, object] | None = None, + cursor: Mapping[str, object] | None = None, + include: Mapping[str, object] | None = None, + order: Mapping[str, object] | Sequence[Mapping[str, object]] | None = None, + distinct: Sequence[str] | None = None, + ) -> RowT_co | None: ... + + async def find_many( + self, + take: int | None = None, + skip: int | None = None, + where: Mapping[str, object] | None = None, + cursor: Mapping[str, object] | None = None, + include: Mapping[str, object] | None = None, + order: Mapping[str, object] | Sequence[Mapping[str, object]] | None = None, + distinct: Sequence[str] | None = None, + ) -> Sequence[RowT_co]: ... + + async def create(self, data: Mapping[str, object], include: Mapping[str, object] | None = None) -> RowT_co: ... + + async def create_many( + self, data: Sequence[Mapping[str, object]], *, skip_duplicates: bool | None = None + ) -> int: ... + + async def upsert( + self, + where: Mapping[str, object], + data: Mapping[str, object], + include: Mapping[str, object] | None = None, + ) -> RowT_co: ... + + async def update( + self, + data: Mapping[str, object], + where: Mapping[str, object], + include: Mapping[str, object] | None = None, + ) -> RowT_co | None: ... + + async def update_many(self, data: Mapping[str, object], where: Mapping[str, object]) -> int: ... + + async def delete( + self, where: Mapping[str, object], include: Mapping[str, object] | None = None + ) -> RowT_co | None: ... + + async def delete_many(self, where: Mapping[str, object] | None = None) -> int: ... + + async def count( + self, + select: None = None, + take: int | None = None, + skip: int | None = None, + where: Mapping[str, object] | None = None, + cursor: Mapping[str, object] | None = None, + ) -> int: ... + + async def group_by( + self, + by: Sequence[str], + *, + where: Mapping[str, object] | None = None, + take: int | None = None, + skip: int | None = None, + order: Mapping[str, object] | Sequence[Mapping[str, object]] | None = None, + having: Mapping[str, object] | None = None, + count: bool | Mapping[str, object] | None = None, + sum: bool | Mapping[str, object] | None = None, + avg: bool | Mapping[str, object] | None = None, + min: bool | Mapping[str, object] | None = None, + max: bool | Mapping[str, object] | None = None, + ) -> Sequence[Mapping[str, object]]: ... + + class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index c8b2c62f9bf..48e55efd258 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -2,17 +2,21 @@ Project repository for database operations on LiteLLM_ProjectTable. """ -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final from litellm.models.project import LiteLLM_ProjectTable from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): """Repository for project database operations.""" @property - def table(self) -> Any: + def table(self) -> TableActions["prisma_models.LiteLLM_ProjectTable"]: return self.prisma_client.db.litellm_projecttable @property diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py index 131f4d377ef..e02f652caf6 100644 --- a/litellm/repositories/table_repositories.py +++ b/litellm/repositories/table_repositories.py @@ -7,12 +7,16 @@ These are thin wrappers for tables that do not (yet) need domain-specific query methods; richer repositories live in their own modules. """ -from typing import Any +from typing import TYPE_CHECKING, Any, Final, Generic from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync +from litellm.repositories.prisma_protocols import RowT_co, TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models # noqa: F401 # used by quoted base-class subscripts -class PrismaTableRepository: +class PrismaTableRepository(Generic[RowT_co]): """Base for repositories that expose a single Prisma table.""" table_name: str @@ -27,208 +31,206 @@ class PrismaTableRepository: return self._prisma_client @property - def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper - return wrap_table_actions_for_config_sync( - actions=getattr(self.prisma_client.db, self.table_name), - table_name=self.table_name, - ) + def table(self) -> TableActions[RowT_co]: + actions: Final[TableActions[RowT_co]] = getattr(self.prisma_client.db, self.table_name) + return wrap_table_actions_for_config_sync(actions=actions, table_name=self.table_name) -class PolicyRepository(PrismaTableRepository): +class PolicyRepository(PrismaTableRepository["prisma_models.LiteLLM_PolicyTable"]): table_name = "litellm_policytable" -class AgentsRepository(PrismaTableRepository): +class AgentsRepository(PrismaTableRepository["prisma_models.LiteLLM_AgentsTable"]): table_name = "litellm_agentstable" -class ObjectPermissionRepository(PrismaTableRepository): +class ObjectPermissionRepository(PrismaTableRepository["prisma_models.LiteLLM_ObjectPermissionTable"]): table_name = "litellm_objectpermissiontable" -class GuardrailsRepository(PrismaTableRepository): +class GuardrailsRepository(PrismaTableRepository["prisma_models.LiteLLM_GuardrailsTable"]): table_name = "litellm_guardrailstable" -class MCPServerRepository(PrismaTableRepository): +class MCPServerRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPServerTable"]): table_name = "litellm_mcpservertable" -class ManagedObjectRepository(PrismaTableRepository): +class ManagedObjectRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedObjectTable"]): table_name = "litellm_managedobjecttable" -class OrganizationMembershipRepository(PrismaTableRepository): +class OrganizationMembershipRepository(PrismaTableRepository["prisma_models.LiteLLM_OrganizationMembership"]): table_name = "litellm_organizationmembership" -class SpendLogsRepository(PrismaTableRepository): +class SpendLogsRepository(PrismaTableRepository["prisma_models.LiteLLM_SpendLogs"]): table_name = "litellm_spendlogs" -class ClaudeCodePluginRepository(PrismaTableRepository): +class ClaudeCodePluginRepository(PrismaTableRepository["prisma_models.LiteLLM_ClaudeCodePluginTable"]): table_name = "litellm_claudecodeplugintable" -class TeamMembershipRepository(PrismaTableRepository): +class TeamMembershipRepository(PrismaTableRepository["prisma_models.LiteLLM_TeamMembership"]): table_name = "litellm_teammembership" -class EndUserRepository(PrismaTableRepository): +class EndUserRepository(PrismaTableRepository["prisma_models.LiteLLM_EndUserTable"]): table_name = "litellm_endusertable" -class ManagedVectorStoresRepository(PrismaTableRepository): +class ManagedVectorStoresRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedVectorStoresTable"]): table_name = "litellm_managedvectorstorestable" -class MCPUserCredentialsRepository(PrismaTableRepository): +class MCPUserCredentialsRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPUserCredentials"]): table_name = "litellm_mcpusercredentials" -class MCPServerOAuthClientRepository(PrismaTableRepository): +class MCPServerOAuthClientRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPServerOAuthClient"]): table_name = "litellm_mcpserveroauthclient" -class PromptRepository(PrismaTableRepository): +class PromptRepository(PrismaTableRepository["prisma_models.LiteLLM_PromptTable"]): table_name = "litellm_prompttable" -class TagRepository(PrismaTableRepository): +class TagRepository(PrismaTableRepository["prisma_models.LiteLLM_TagTable"]): table_name = "litellm_tagtable" -class InvitationLinkRepository(PrismaTableRepository): +class InvitationLinkRepository(PrismaTableRepository["prisma_models.LiteLLM_InvitationLink"]): table_name = "litellm_invitationlink" -class JWTKeyMappingRepository(PrismaTableRepository): +class JWTKeyMappingRepository(PrismaTableRepository["prisma_models.LiteLLM_JWTKeyMapping"]): table_name = "litellm_jwtkeymapping" -class ManagedFileRepository(PrismaTableRepository): +class ManagedFileRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedFileTable"]): table_name = "litellm_managedfiletable" -class MemoryRepository(PrismaTableRepository): +class MemoryRepository(PrismaTableRepository["prisma_models.LiteLLM_MemoryTable"]): table_name = "litellm_memorytable" -class SearchToolsRepository(PrismaTableRepository): +class SearchToolsRepository(PrismaTableRepository["prisma_models.LiteLLM_SearchToolsTable"]): table_name = "litellm_searchtoolstable" -class ConfigOverridesRepository(PrismaTableRepository): +class ConfigOverridesRepository(PrismaTableRepository["prisma_models.LiteLLM_ConfigOverrides"]): table_name = "litellm_configoverrides" -class MCPToolsetRepository(PrismaTableRepository): +class MCPToolsetRepository(PrismaTableRepository["prisma_models.LiteLLM_MCPToolsetTable"]): table_name = "litellm_mcptoolsettable" -class ToolRepository(PrismaTableRepository): +class ToolRepository(PrismaTableRepository["prisma_models.LiteLLM_ToolTable"]): table_name = "litellm_tooltable" -class DeletedVerificationTokenRepository(PrismaTableRepository): +class DeletedVerificationTokenRepository(PrismaTableRepository["prisma_models.LiteLLM_DeletedVerificationToken"]): table_name = "litellm_deletedverificationtoken" -class WorkflowRunRepository(PrismaTableRepository): +class WorkflowRunRepository(PrismaTableRepository["prisma_models.LiteLLM_WorkflowRun"]): table_name = "litellm_workflowrun" -class ModelTableRepository(PrismaTableRepository): +class ModelTableRepository(PrismaTableRepository["prisma_models.LiteLLM_ModelTable"]): table_name = "litellm_modeltable" -class AccessGroupRepository(PrismaTableRepository): +class AccessGroupRepository(PrismaTableRepository["prisma_models.LiteLLM_AccessGroupTable"]): table_name = "litellm_accessgrouptable" -class SSOConfigRepository(PrismaTableRepository): +class SSOConfigRepository(PrismaTableRepository["prisma_models.LiteLLM_SSOConfig"]): table_name = "litellm_ssoconfig" -class UISettingsRepository(PrismaTableRepository): +class UISettingsRepository(PrismaTableRepository["prisma_models.LiteLLM_UISettings"]): table_name = "litellm_uisettings" -class DailyGuardrailMetricsRepository(PrismaTableRepository): +class DailyGuardrailMetricsRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyGuardrailMetrics"]): table_name = "litellm_dailyguardrailmetrics" -class DailyGuardrailUsageUnitsRepository(PrismaTableRepository): +class DailyGuardrailUsageUnitsRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyGuardrailUsageUnits"]): table_name = "litellm_dailyguardrailusageunits" -class PolicyAttachmentRepository(PrismaTableRepository): +class PolicyAttachmentRepository(PrismaTableRepository["prisma_models.LiteLLM_PolicyAttachmentTable"]): table_name = "litellm_policyattachmenttable" -class DeletedTeamRepository(PrismaTableRepository): +class DeletedTeamRepository(PrismaTableRepository["prisma_models.LiteLLM_DeletedTeamTable"]): table_name = "litellm_deletedteamtable" -class SkillsRepository(PrismaTableRepository): +class SkillsRepository(PrismaTableRepository["prisma_models.LiteLLM_SkillsTable"]): table_name = "litellm_skillstable" -class CacheConfigRepository(PrismaTableRepository): +class CacheConfigRepository(PrismaTableRepository["prisma_models.LiteLLM_CacheConfig"]): table_name = "litellm_cacheconfig" -class ManagedVectorStoreIndexRepository(PrismaTableRepository): +class ManagedVectorStoreIndexRepository(PrismaTableRepository["prisma_models.LiteLLM_ManagedVectorStoreIndexTable"]): table_name = "litellm_managedvectorstoreindextable" -class WorkflowMessageRepository(PrismaTableRepository): +class WorkflowMessageRepository(PrismaTableRepository["prisma_models.LiteLLM_WorkflowMessage"]): table_name = "litellm_workflowmessage" -class DailyTagSpendRepository(PrismaTableRepository): +class DailyTagSpendRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyTagSpend"]): table_name = "litellm_dailytagspend" -class SpendLogToolIndexRepository(PrismaTableRepository): +class SpendLogToolIndexRepository(PrismaTableRepository["prisma_models.LiteLLM_SpendLogToolIndex"]): table_name = "litellm_spendlogtoolindex" -class DailyToolSpendRepository(PrismaTableRepository): +class DailyToolSpendRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyToolSpend"]): table_name = "litellm_dailytoolspend" -class SpendLogGuardrailIndexRepository(PrismaTableRepository): +class SpendLogGuardrailIndexRepository(PrismaTableRepository["prisma_models.LiteLLM_SpendLogGuardrailIndex"]): table_name = "litellm_spendlogguardrailindex" -class UserNotificationsRepository(PrismaTableRepository): +class UserNotificationsRepository(PrismaTableRepository["prisma_models.LiteLLM_UserNotifications"]): table_name = "litellm_usernotifications" -class HealthCheckRepository(PrismaTableRepository): +class HealthCheckRepository(PrismaTableRepository["prisma_models.LiteLLM_HealthCheckTable"]): table_name = "litellm_healthchecktable" -class DeprecatedVerificationTokenRepository(PrismaTableRepository): +class DeprecatedVerificationTokenRepository(PrismaTableRepository["prisma_models.LiteLLM_DeprecatedVerificationToken"]): table_name = "litellm_deprecatedverificationtoken" -class WorkflowEventRepository(PrismaTableRepository): +class WorkflowEventRepository(PrismaTableRepository["prisma_models.LiteLLM_WorkflowEvent"]): table_name = "litellm_workflowevent" -class DailyPolicyMetricsRepository(PrismaTableRepository): +class DailyPolicyMetricsRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyPolicyMetrics"]): table_name = "litellm_dailypolicymetrics" -class AdaptiveRouterStateRepository(PrismaTableRepository): +class AdaptiveRouterStateRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterState"]): table_name = "litellm_adaptiverouterstate" -class AuditLogRepository(PrismaTableRepository): +class AuditLogRepository(PrismaTableRepository["prisma_models.LiteLLM_AuditLog"]): table_name = "litellm_auditlog" -class AdaptiveRouterSessionRepository(PrismaTableRepository): +class AdaptiveRouterSessionRepository(PrismaTableRepository["prisma_models.LiteLLM_AdaptiveRouterSession"]): table_name = "litellm_adaptiveroutersession" diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 7efd32288e4..221f5b22c41 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -5,7 +5,7 @@ Team repository for database operations on LiteLLM_TeamTable. import json from collections.abc import Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final from pydantic import TypeAdapter @@ -15,9 +15,11 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) +from litellm.repositories.prisma_protocols import TableActions if TYPE_CHECKING: from prisma import Prisma + from prisma import models as prisma_models _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( @@ -34,11 +36,11 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" @property - def table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper + def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]: return self.prisma_client.db.litellm_teamtable @property - def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper + def deleted_table(self) -> TableActions["prisma_models.LiteLLM_DeletedTeamTable"]: return self.prisma_client.db.litellm_deletedteamtable @property diff --git a/litellm/repositories/user_banner_repository.py b/litellm/repositories/user_banner_repository.py index 3b69e433853..c1ed977e048 100644 --- a/litellm/repositories/user_banner_repository.py +++ b/litellm/repositories/user_banner_repository.py @@ -1,11 +1,14 @@ -from typing import Final +from typing import TYPE_CHECKING, Final from litellm.repositories.table_repositories import PrismaTableRepository +if TYPE_CHECKING: + from prisma import models as prisma_models # noqa: F401 # resolved only from the quoted base-class subscript below + USER_BANNER_ROW_ID: Final = "user_banner" -class UserBannerRepository(PrismaTableRepository): +class UserBannerRepository(PrismaTableRepository["prisma_models.LiteLLM_UISettings"]): table_name = "litellm_uisettings" async def get_raw_settings(self) -> object: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index d0d366e1772..9df1bceac9c 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -4,10 +4,14 @@ User repository for database operations on LiteLLM_UserTable. import json from collections.abc import Mapping -from typing import Any, Final +from typing import TYPE_CHECKING, Final from litellm.models.user import LiteLLM_UserTable from litellm.repositories.base_repository import BaseRepository, DbRecord, record_to_dict +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models _JSON_ENCODED_COLUMNS: Final = frozenset({"metadata", "model_spend", "model_max_budget"}) @@ -16,7 +20,7 @@ class UserRepository(BaseRepository[LiteLLM_UserTable]): """Repository for user database operations.""" @property - def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper + def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]: return self.prisma_client.db.litellm_usertable @property diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index 3790ad25914..c0e59f9b975 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -3,9 +3,9 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke """ import json -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final from litellm.models.verification_token import ( LiteLLM_VerificationToken, @@ -15,8 +15,12 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) +from litellm.repositories.prisma_protocols import TableActions if TYPE_CHECKING: + from prisma.models import ( + LiteLLM_DeletedVerificationToken as PrismaDeletedVerificationToken, + ) from prisma.models import ( LiteLLM_VerificationToken as PrismaVerificationToken, ) @@ -45,11 +49,11 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): return prisma_client @property - def table(self) -> Any: + def table(self) -> TableActions["PrismaVerificationToken"]: return self.prisma_client.db.litellm_verificationtoken @property - def deleted_table(self) -> Any: + def deleted_table(self) -> TableActions["PrismaDeletedVerificationToken"]: return self.prisma_client.db.litellm_deletedverificationtoken @property @@ -79,29 +83,29 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): async def find_by_alias(self, key_alias: str) -> LiteLLM_VerificationToken | None: """Find a token by key alias.""" - records: Final[list[PrismaVerificationToken]] = await self.table.find_many(where={"key_alias": key_alias}) + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"key_alias": key_alias}) if records: return self._to_model(records[0]) return None async def find_by_user_id(self, user_id: str) -> list[LiteLLM_VerificationToken]: """Find all tokens belonging to a user.""" - records: Final[list[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id}) + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"user_id": user_id}) return self._to_model_list(records) async def find_by_team_id(self, team_id: str) -> list[LiteLLM_VerificationToken]: """Find all tokens belonging to a team.""" - records: Final[list[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id}) + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"team_id": team_id}) return self._to_model_list(records) async def find_by_project_id(self, project_id: str) -> list[LiteLLM_VerificationToken]: """Find all tokens belonging to a project.""" - records: Final[list[PrismaVerificationToken]] = await self.table.find_many(where={"project_id": project_id}) + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many(where={"project_id": project_id}) return self._to_model_list(records) async def find_active_tokens(self) -> list[LiteLLM_VerificationToken]: """Find all active (non-expired, non-blocked) tokens.""" - records: Final[list[PrismaVerificationToken]] = await self.table.find_many( + records: Final[Sequence[PrismaVerificationToken]] = await self.table.find_many( where={ "blocked": {"not": True}, "OR": [{"expires": None}, {"expires": {"gt": datetime.utcnow()}}], diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py index e9e7ae908a5..0418f0c5e14 100644 --- a/litellm/responses/file_search/emulated_handler.py +++ b/litellm/responses/file_search/emulated_handler.py @@ -15,7 +15,9 @@ import json import time import uuid from collections.abc import Iterable, Sequence -from typing import TYPE_CHECKING, Any, Final, TypeAlias, cast +from typing import TYPE_CHECKING, Any, Final, TypeAlias, cast # noqa: TID251 # see kwargs-ok / cast-ok markers + +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._internal_context import is_internal_call from litellm._logging import verbose_logger @@ -31,6 +33,12 @@ ToolParam: TypeAlias = object FILE_SEARCH_FUNCTION_NAME: Final = "litellm_file_search" +class FileSearchToolCallArgs(TypedDict): + queries: ReadOnly[NotRequired[object]] + query: ReadOnly[NotRequired[object]] + vector_store_id: ReadOnly[NotRequired[object]] + + # --------------------------------------------------------------------------- # Detection # --------------------------------------------------------------------------- @@ -175,13 +183,20 @@ async def _run_vector_searches( # --------------------------------------------------------------------------- -def _get_field(result: object, key: str, default: object = None) -> Any: +def _get_field(result: object, key: str, default: object = None) -> object: """Read a field from either a dict/TypedDict or an attribute-based object.""" if isinstance(result, dict): return result.get(key, default) return getattr(result, key, default) +def _joined_content_text(result: object) -> str: + """Concatenate the text of every content chunk on a search result.""" + content_items: Final = cast(Iterable[object], _get_field(result, "content") or []) # cast-ok: iterated as today + text_chunks: Final = [c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "") for c in content_items] + return " ".join(t for t in text_chunks if t) + + def _format_search_results_as_tool_output( results: list[VectorStoreSearchResult], ) -> str: @@ -194,9 +209,7 @@ def _format_search_results_as_tool_output( score = _get_field(result, "score") file_id = _get_field(result, "file_id") filename = _get_field(result, "filename") - content_items = _get_field(result, "content") or [] - text_chunks = [c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "") for c in content_items] - text = " ".join(t for t in text_chunks if t) + text = _joined_content_text(result) header = f"[Result {i}" if filename: @@ -226,9 +239,7 @@ def _build_search_results_for_include( formatted: Final[list[dict[str, object]]] = [] for result in results: file_id = _get_field(result, "file_id") or "" - content_items = _get_field(result, "content") or [] - text_chunks = [c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "") for c in content_items] - text = " ".join(t for t in text_chunks if t) + text = _joined_content_text(result) formatted.append( { "file_id": file_id, @@ -353,14 +364,14 @@ def _synthesize_responses_api_response( created_at=getattr(original_response, "created_at", int(time.time())), status="completed", model=getattr(original_response, "model", ""), - output=cast(list[ResponseOutputItem | dict[str, Any]], synthesized_output), + output=cast(list[ResponseOutputItem | dict[str, object]], synthesized_output), # cast-ok: list is invariant usage=getattr(original_response, "usage", None), error=None, ) if hasattr(original_response, "_hidden_params"): hidden: Final = dict(getattr(original_response, "_hidden_params") or {}) if first_response is not None and hasattr(first_response, "_hidden_params"): - first_hidden: Final = getattr(first_response, "_hidden_params") or {} + first_hidden: Final[object] = getattr(first_response, "_hidden_params") or {} first_cost: Final = ( first_hidden.get("response_cost") if isinstance(first_hidden, dict) @@ -385,9 +396,10 @@ async def _call_aresponses(input, model, tools, **kwargs): # pragma: no cover def _prepare_emulated_file_search_call( - kwargs: dict[str, Any], + kwargs: dict[str, object], ) -> tuple[bool, dict[str, object]]: - include_items: Final[list[str]] = list(kwargs.get("include") or []) + raw_include: Final = kwargs.get("include") or [] + include_items: Final[list[object]] = list(cast(Iterable[object], raw_include)) # cast-ok: iterated as today include_search_results: Final = "file_search_call.results" in include_items original_stream: Final = kwargs.get("stream") @@ -413,16 +425,16 @@ def _extract_tool_call_fields(tool_call: object, fallback_call_id: str) -> tuple return call_id, raw_args -def _resolve_queries_from_args(args: dict[str, Any], input: object) -> list[str]: +def _resolve_queries_from_args(args: FileSearchToolCallArgs, input: object) -> list[str]: """Pull the queries list out of parsed tool-call arguments, with backward-compat fallbacks.""" queries_from_call: Final = args.get("queries") if not queries_from_call: # Fallback: check for single "query" field (backward compat) single_query: Final = args.get("query") - return [single_query] if single_query else [str(input)] + return [cast(str, single_query)] if single_query else [str(input)] # cast-ok: model-supplied, as today if not isinstance(queries_from_call, list): return [str(queries_from_call)] - return queries_from_call + return cast(list[str], queries_from_call) # cast-ok: model-supplied elements, forwarded unchecked as today async def _execute_file_search_tool_calls( @@ -440,14 +452,14 @@ async def _execute_file_search_tool_calls( call_id, raw_args = _extract_tool_call_fields(tool_call, fallback_call_id=file_search_call_id) try: - args = json.loads(raw_args) if isinstance(raw_args, str) else raw_args + args: FileSearchToolCallArgs = json.loads(raw_args) if isinstance(raw_args, str) else raw_args except json.JSONDecodeError: args = {} queries_from_call = _resolve_queries_from_args(args, input) vs_id_arg = args.get("vector_store_id") - vs_ids_for_call = [vs_id_arg] if vs_id_arg else all_vs_ids + vs_ids_for_call = [cast(str, vs_id_arg)] if vs_id_arg else all_vs_ids # cast-ok: model-supplied, as today queries, results = await _run_vector_searches( queries=queries_from_call, @@ -481,7 +493,7 @@ def _build_follow_up_input( original_input_items: Final[list[object]] = ( list(input) if isinstance(input, (list, tuple)) else [{"role": "user", "content": str(input)}] ) - first_response_output_items: Final[list[Any]] = [] + first_response_output_items: Final[list[object]] = [] for _item in first_response.output: if isinstance(_item, dict): first_response_output_items.append(_item) @@ -498,7 +510,7 @@ async def aresponses_with_emulated_file_search( model: str, tools: Iterable[ToolParam] | None = None, # Pass-through params — forwarded as-is to the underlying aresponses call - **kwargs: Any, + **kwargs: Any, # kwargs-ok: `object` would surface the caller's partially-unknown dict at its call site ) -> ResponsesAPIResponse: """ Emulated file_search for providers that don't support it natively. @@ -507,7 +519,7 @@ async def aresponses_with_emulated_file_search( runs vector search, and synthesizes an OpenAI-format response. """ # Determine whether caller wants search_results populated in the output. - _include_search_results, kwargs = _prepare_emulated_file_search_call(kwargs=kwargs) + _include_search_results, call_kwargs = _prepare_emulated_file_search_call(kwargs=kwargs) # 1. Replace file_search tools with function tool transformed_tools, all_vs_ids = _replace_file_search_tools(tools) @@ -524,7 +536,7 @@ async def aresponses_with_emulated_file_search( input=input, model=model, tools=transformed_tools or None, - **kwargs, + **call_kwargs, ), ) finally: @@ -588,7 +600,7 @@ async def aresponses_with_emulated_file_search( input=follow_up_input, model=model, tools=None, # no tools needed for the answer step - **kwargs, + **call_kwargs, ), ) finally: diff --git a/litellm/responses/litellm_completion_transformation/custom_tools.py b/litellm/responses/litellm_completion_transformation/custom_tools.py index fa4ed73a1d6..cccae06c74b 100644 --- a/litellm/responses/litellm_completion_transformation/custom_tools.py +++ b/litellm/responses/litellm_completion_transformation/custom_tools.py @@ -16,8 +16,8 @@ logic. """ import json -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final from pydantic import BaseModel, TypeAdapter, ValidationError @@ -29,7 +29,7 @@ from litellm.types.llms.openai import ( _MAX_ARGUMENTS_LEN: Final = 1_000_000 -def extract_custom_tool_names(tools: list[Any] | None) -> set[str]: +def extract_custom_tool_names(tools: Sequence[object] | None) -> set[str]: """Extract names of tools originally defined as ``type: "custom"``.""" if not tools: return set() @@ -73,7 +73,7 @@ def build_tool_call_item_kwargs( arguments_or_input: str, status: str, custom_tool_names: set[str], -) -> dict[str, Any]: +) -> dict[str, str]: """Build kwargs for an output item dict that is either a ``function_call`` or a ``custom_tool_call`` depending on whether *name* is in *custom_tool_names*. @@ -86,7 +86,7 @@ def build_tool_call_item_kwargs( """ custom: Final = is_custom_tool_call(name, custom_tool_names) item_type: Final = "custom_tool_call" if custom else "function_call" - kwargs: Final[dict[str, Any]] = { + kwargs: Final[dict[str, str]] = { "type": item_type, "id": call_id, "call_id": call_id, diff --git a/litellm/responses/litellm_completion_transformation/handler.py b/litellm/responses/litellm_completion_transformation/handler.py index 555e3258773..a0e8cd278e6 100644 --- a/litellm/responses/litellm_completion_transformation/handler.py +++ b/litellm/responses/litellm_completion_transformation/handler.py @@ -2,8 +2,8 @@ Handler for transforming responses api requests to litellm.completion requests """ -from collections.abc import Coroutine -from typing import Any, Final +from collections.abc import Coroutine, Mapping +from typing import Final import litellm from litellm.responses.litellm_completion_transformation.streaming_iterator import ( @@ -30,12 +30,12 @@ class LiteLLMCompletionTransformationHandler: custom_llm_provider: str | None = None, _is_async: bool = False, stream: bool | None = None, - extra_headers: dict[str, Any] | None = None, + extra_headers: Mapping[str, object] | None = None, **kwargs, ) -> ( ResponsesAPIResponse | BaseResponsesAPIStreamingIterator - | Coroutine[Any, Any, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator] + | Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator] ): litellm_completion_request: Final[dict] = ( LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( diff --git a/litellm/responses/litellm_completion_transformation/session_handler.py b/litellm/responses/litellm_completion_transformation/session_handler.py index 1566bb1bdd7..6c66c3fafc6 100644 --- a/litellm/responses/litellm_completion_transformation/session_handler.py +++ b/litellm/responses/litellm_completion_transformation/session_handler.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm._logging import verbose_proxy_logger -from litellm.proxy._types import SpendLogsPayload +from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload from litellm.proxy.spend_tracking.cold_storage_handler import ColdStorageHandler from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ( @@ -143,7 +143,7 @@ class ResponsesSessionHandler: model_response: Final = ModelResponse(**_response_output) for choice in model_response.choices: if hasattr(choice, "message"): - chat_completion_message_history.append(getattr(choice, "message")) + chat_completion_message_history.append(choice.message) return chat_completion_message_history @staticmethod @@ -195,7 +195,7 @@ class ResponsesSessionHandler: try: metadata_str: Final = spend_log.get("metadata", "{}") if isinstance(metadata_str, str): - metadata_dict: Final = json.loads(metadata_str) + metadata_dict: Final[SpendLogsMetadata] = json.loads(metadata_str) return metadata_dict.get("cold_storage_object_key") elif isinstance(metadata_str, dict): return metadata_str.get("cold_storage_object_key") diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 92bbca9ee5b..8b1eeb30306 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -1,5 +1,6 @@ import time import uuid +from collections.abc import Sequence from typing import Any, Final, cast import litellm @@ -48,14 +49,18 @@ from litellm.types.utils import ( ) +def _index_of_output_item_type(items: Sequence[object], item_type: str) -> int | None: + return next( + (index for index, item in enumerate(items) if getattr(item, "type", None) == item_type), + None, + ) + + def _output_items_with_id(items: tuple[Any, ...], item_type: str, item_id: str | None) -> tuple[Any, ...]: if item_id is None: return items - target_index: Final = next( - (index for index, item in enumerate(items) if getattr(item, "type", None) == item_type), - None, - ) + target_index: Final = _index_of_output_item_type(items, item_type) if target_index is None: return items @@ -86,7 +91,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.litellm_metadata: dict | None = litellm_metadata or {} # Store lightweight dict snapshots for stream_chunk_builder to reduce # repeated Pydantic attribute access in end-of-stream assembly. - self.collected_chat_completion_chunks: list[dict[str, Any]] = [] + self.collected_chat_completion_chunks: list[dict[str, object]] = [] self.finished: bool = False self.litellm_logging_obj = litellm_custom_stream_wrapper.logging_obj self.sent_response_created_event: bool = False @@ -98,7 +103,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.sent_output_item_done_event: bool = False self.sent_annotation_events: bool = False self.litellm_model_response: ModelResponse | TextCompletionResponse | None = None - self.completed_response: Any = None + self.completed_response = None self.final_text: str = "" self._cached_item_id: str | None = None self._cached_response_id: str | None = None @@ -123,7 +128,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self._reasoning_done_emitted = False self._reasoning_item_id: str | None = None self._accumulated_reasoning_content_parts: list[str] = [] - self._accumulated_provider_specific_fields: dict[str, Any] = {} + self._accumulated_provider_specific_fields: dict[str, object] = {} self._custom_tool_names: set[str] = extract_custom_tool_names(self.responses_api_request.get("tools")) self._namespace_tool_names = LiteLLMCompletionResponsesConfig.namespace_tool_name_map( self.responses_api_request.get("tools") @@ -543,7 +548,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): @staticmethod def _snapshot_chunk_for_stream_chunk_builder( chunk: ModelResponseStream, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Convert a streaming chunk into a plain dict for end-of-stream assembly. Keep _hidden_params so downstream usage/header behavior is preserved. @@ -1161,7 +1166,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): if litellm_model_response: # Add cost to usage object if include_cost_in_streaming_usage is True if litellm.include_cost_in_streaming_usage and self.litellm_logging_obj is not None: - usage: Final = getattr(litellm_model_response, "usage", None) + usage: Final[object] = getattr(litellm_model_response, "usage", None) if usage is not None: setattr( usage, diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index b8d7b726a28..aba24692d5d 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -5,7 +5,7 @@ Handles transforming from Responses API -> LiteLLM completion (Chat Completion import json import re import uuid -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Iterable, Iterator, Mapping, Sequence from types import MappingProxyType from typing import ( TYPE_CHECKING, @@ -28,7 +28,7 @@ from openai.types.responses import ResponseFunctionToolCall from openai.types.responses.response_create_params import ResponseInputParam from openai.types.responses.tool_param import FunctionToolParam from pydantic import TypeAdapter -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.caching import InMemoryCache @@ -46,6 +46,7 @@ from litellm.types.llms.openai import ( ChatCompletionRedactedThinkingBlock, ChatCompletionResponseMessage, ChatCompletionSystemMessage, + ChatCompletionTextObject, ChatCompletionThinkingBlock, ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, @@ -129,6 +130,30 @@ class _HasId(Protocol): id: object +class _ResponsesToolCallItem(Protocol): + name: object + arguments: object + + def get(self, key: str, /) -> object: ... + + +class _ToolFunctionDefinition(TypedDict, total=False): + name: ReadOnly[str] + description: ReadOnly[str] + parameters: ReadOnly[dict[str, object]] + strict: ReadOnly[bool | None] + + +def _attribute_fields(value: object) -> dict[str, object]: + if not hasattr(value, "__dict__"): + return {} # mutable-ok: provider_specific_fields payload + return dict(cast("Iterable[tuple[str, object]]", value)) # cast-ok: dict() raises on non-pair values, as before + + +def _input_item_role(input_item: Mapping[str, object]) -> str: + return cast(str, input_item.get("role") or "user") # cast-ok: client-supplied role forwarded verbatim, unvalidated + + class ChatCompletionSession(TypedDict, total=False): messages: list[ AllMessageValues @@ -677,7 +702,7 @@ class LiteLLMCompletionResponsesConfig: existing_text: Final = _reasoning_text(msg) combined: Final = "\n".join(pending_texts + ((existing_text,) if existing_text else ())) if isinstance(msg, dict): - cast(dict[str, Any], msg)["reasoning_content"] = combined # cast-ok: mutable reasoning carrier + cast(dict[str, object], msg)["reasoning_content"] = combined # cast-ok: mutable reasoning carrier else: setattr(msg, "reasoning_content", combined) # noqa: B010 # attribute name is fixed, not dynamic if pending_blocks: @@ -685,7 +710,7 @@ class LiteLLMCompletionResponsesConfig: pending_blocks + (_thinking_blocks(msg) or ()) ) if isinstance(msg, dict): - cast(dict[str, Any], msg)["thinking_blocks"] = replayed # cast-ok: mutable reasoning carrier + cast(dict[str, object], msg)["thinking_blocks"] = replayed # cast-ok: mutable reasoning carrier else: setattr(msg, "thinking_blocks", replayed) # noqa: B010 # attribute name is fixed, not dynamic @@ -1034,7 +1059,7 @@ class LiteLLMCompletionResponsesConfig: def _add_tool_call_to_assistant(assistant_message: object, tool_call_chunk: ChatCompletionToolCallChunk) -> None: """Add a tool_call to an assistant message.""" if isinstance(assistant_message, dict): - prev_assistant_dict: Final = cast(dict[str, Any], assistant_message) + prev_assistant_dict: Final = cast(dict[str, object], assistant_message) if "tool_calls" not in prev_assistant_dict: prev_assistant_dict["tool_calls"] = [] tool_calls_list: Final = prev_assistant_dict["tool_calls"] @@ -1119,7 +1144,7 @@ class LiteLLMCompletionResponsesConfig: # Type-safe way to set tool_call_id on tool message if isinstance(message, dict): # Cast to dict to allow setting tool_call_id - message_dict = cast(dict[str, Any], message) + message_dict = cast(dict[str, object], message) message_dict["tool_call_id"] = tool_call_id elif hasattr(message, "tool_call_id"): setattr(message, "tool_call_id", tool_call_id) @@ -1171,7 +1196,7 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def _transform_responses_api_input_item_to_chat_completion_message( - input_item: Any, + input_item: Mapping[str, object], replay_reasoning: bool = False, ) -> list[AllMessageValues | GenericChatCompletionMessage | ChatCompletionResponseMessage]: """ @@ -1199,7 +1224,9 @@ class LiteLLMCompletionResponsesConfig: elif LiteLLMCompletionResponsesConfig._is_input_item_function_call(input_item): # handle function call input items return LiteLLMCompletionResponsesConfig._transform_responses_api_function_call_to_chat_completion_message( - function_call=input_item + function_call=cast( # cast-ok: callee coerces every field it reads with `or ""` / str() + Mapping[str, str], input_item + ) ) elif input_item.get("type") == "reasoning": # A ResponseReasoningItemParam carries the prior-turn chain-of-thought. @@ -1224,7 +1251,7 @@ class LiteLLMCompletionResponsesConfig: return [] # mutable-ok: empty drop result return [ # mutable-ok: single message result GenericChatCompletionMessage( - role=input_item.get("role") or "user", + role=_input_item_role(input_item), content=LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( inspectable ), @@ -1252,7 +1279,7 @@ class LiteLLMCompletionResponsesConfig: return [] return [ GenericChatCompletionMessage( - role=input_item.get("role") or "user", + role=_input_item_role(input_item), content=LiteLLMCompletionResponsesConfig._transform_responses_api_content_to_chat_completion_content( content ), @@ -1339,7 +1366,7 @@ class LiteLLMCompletionResponsesConfig: if not isinstance(encrypted_content, str) or not encrypted_content.strip(): return None try: - decoded: Final[object] = json.loads(encrypted_content) + decoded: Final[object] = cast(object, json.loads(encrypted_content)) # cast-ok: json.loads returns Any except ValueError: return None if not isinstance(decoded, list): @@ -1406,7 +1433,7 @@ class LiteLLMCompletionResponsesConfig: def _normalize_function_call_output_to_tool_content( output: object, - ) -> Any: + ) -> str | list[ChatCompletionTextObject | ChatCompletionImageObject]: """ Normalize Responses API function_call_output.output into a shape that downstream chat adapters (esp. Gemini) can reliably consume. @@ -1428,7 +1455,7 @@ class LiteLLMCompletionResponsesConfig: # Some adapters represent tool output as a list of "input_*" parts if isinstance(output, list): - normalized_blocks: Final[list[dict[str, object]]] = [] + normalized_blocks: Final[list[ChatCompletionTextObject | ChatCompletionImageObject]] = [] text_acc: Final[list[str]] = [] for part in output: if not isinstance(part, dict): @@ -1899,7 +1926,7 @@ class LiteLLMCompletionResponsesConfig: result.append(tool) continue if tool.get("type") == "function": - fn = cast(dict[str, Any], tool.get("function") or {}) + fn = cast(_ToolFunctionDefinition, tool.get("function") or {}) parameters = dict(fn.get("parameters", {}) or {}) if not parameters or "type" not in parameters: parameters["type"] = "object" @@ -2095,7 +2122,7 @@ class LiteLLMCompletionResponsesConfig: @staticmethod def convert_response_function_tool_call_to_chat_completion_tool_call( - tool_call_item: Any, + tool_call_item: object, index: int = 0, ) -> dict[str, object]: """ @@ -2108,24 +2135,25 @@ class LiteLLMCompletionResponsesConfig: Returns: Dictionary in ChatCompletionToolCallChunk format """ + item: Final = cast( # cast-ok: duck-typed tool call item, .get access guarded by hasattr below + _ResponsesToolCallItem, tool_call_item + ) # Extract provider_specific_fields if present - provider_specific_fields = getattr(tool_call_item, "provider_specific_fields", None) + provider_specific_fields: object = getattr(tool_call_item, "provider_specific_fields", None) if provider_specific_fields and not isinstance(provider_specific_fields, dict): - provider_specific_fields = ( - dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {} - ) - elif hasattr(tool_call_item, "get") and callable(tool_call_item.get): - provider_fields: Final = tool_call_item.get("provider_specific_fields") + provider_specific_fields = _attribute_fields(provider_specific_fields) + elif hasattr(tool_call_item, "get") and callable(item.get): + provider_fields: Final = item.get("provider_specific_fields") if provider_fields: provider_specific_fields = ( - provider_fields + cast("dict[str, object]", provider_fields) # cast-ok: passed through as-is, keys unvalidated if isinstance(provider_fields, dict) - else (dict(provider_fields) if hasattr(provider_fields, "__dict__") else {}) + else _attribute_fields(provider_fields) ) function_dict: Final[dict[str, object]] = { - "name": tool_call_item.name, - "arguments": tool_call_item.arguments, + "name": item.name, + "arguments": item.arguments, } if provider_specific_fields: @@ -2306,7 +2334,7 @@ class LiteLLMCompletionResponsesConfig: """ output_items: Final[list] = [] for choice in chat_completion_response.choices or []: - message = getattr(choice, "message", None) + message: object = getattr(choice, "message", None) if not message: continue psf = getattr(message, "provider_specific_fields", None) @@ -2338,7 +2366,7 @@ class LiteLLMCompletionResponsesConfig: for choice in choices: if hasattr(choice, "message") and choice.message: message = choice.message - reasoning_content = getattr(message, "reasoning_content", None) or "" + reasoning_content: str = getattr(message, "reasoning_content", None) or "" encrypted_content = LiteLLMCompletionResponsesConfig._encode_thinking_blocks(message) if reasoning_content or encrypted_content: # Only check the first choice for reasoning content diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 34058e8eca7..d6ebc44ac52 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -697,7 +697,7 @@ def _apply_managed_file_id_mapping( tools = cast( Iterable[ToolParam] | None, update_responses_tools_with_model_file_ids( - tools=cast(list[dict[str, Any]] | None, tools), + tools=cast(list[dict[str, object]] | None, tools), model_id=model_info_id, model_file_id_mapping=model_file_id_mapping, ), @@ -734,7 +734,7 @@ def _responses_try_dispatch_mcp_gateway( extra_body: dict[str, object] | None, timeout: float | httpx.Timeout | None, custom_llm_provider: str | None, - kwargs: dict[str, Any], + kwargs: dict[str, object], _is_async: bool, ) -> Any | None: """Return a response when MCP gateway handles the call; otherwise None.""" diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 2a0406f9a4d..a75b3768636 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -1,7 +1,9 @@ """Helpers for handling MCP-aware `/chat/completions` requests.""" import logging -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Final, cast + +from typing_extensions import TypedDict, Unpack from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, @@ -14,6 +16,10 @@ if TYPE_CHECKING: from litellm.proxy._types import UserAPIKeyAuth +class _MCPCompletionKwargs(TypedDict, total=False, extra_items=object): + """Extra keywords forwarded verbatim to ``litellm.acompletion``, which owns their contract.""" + + def _add_mcp_metadata_to_response( response: ModelResponse | CustomStreamWrapper, openai_tools: list | None, @@ -79,7 +85,7 @@ async def acompletion_with_mcp( model: str, messages: list, tools: list | None = None, - **kwargs: Any, + **kwargs: Unpack[_MCPCompletionKwargs], # kwargs-ok: forwarded verbatim to litellm.acompletion, which owns them ) -> ModelResponse | CustomStreamWrapper: """ Async completion with MCP integration. @@ -126,7 +132,7 @@ async def acompletion_with_mcp( ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - litellm_trace_id=kwargs.get("litellm_trace_id"), + litellm_trace_id=context.litellm_trace_id, mcp_auth_header=mcp_auth_header, mcp_server_auth_headers=mcp_server_auth_headers, request_tags=request_tags, @@ -168,7 +174,7 @@ async def acompletion_with_mcp( return response # For auto-execute: handle streaming vs non-streaming differently - stream: Final[bool] = kwargs.get("stream", False) + stream: Final[object] = kwargs.get("stream", False) mock_tool_calls: Final = base_call_args.pop("mock_tool_calls", None) if stream: @@ -490,8 +496,8 @@ async def acompletion_with_mcp( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, - litellm_call_id=kwargs.get("litellm_call_id"), - litellm_trace_id=kwargs.get("litellm_trace_id"), + litellm_call_id=context.litellm_call_id, + litellm_trace_id=context.litellm_trace_id, openai_tools=openai_tools, base_call_args=base_call_args, request_tags=request_tags, @@ -604,8 +610,8 @@ async def acompletion_with_mcp( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, - litellm_call_id=kwargs.get("litellm_call_id"), - litellm_trace_id=kwargs.get("litellm_trace_id"), + litellm_call_id=context.litellm_call_id, + litellm_trace_id=context.litellm_trace_id, request_tags=request_tags, ) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index c7471518398..8f5dc926c68 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -22,6 +22,8 @@ from litellm.types.llms.openai import ( ) if TYPE_CHECKING: + from collections.abc import AsyncIterator, Iterator + from mcp.types import Tool as MCPTool from litellm.proxy._types import UserAPIKeyAuth @@ -511,7 +513,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if self.base_iterator: if hasattr(self.base_iterator, "__anext__"): try: - chunk: Final[ResponsesAPIStreamingResponse] = await cast(Any, self.base_iterator).__anext__() + chunk: Final[ResponsesAPIStreamingResponse] = await cast( # cast-ok: hasattr __anext__ checked + "AsyncIterator[ResponsesAPIStreamingResponse]", self.base_iterator + ).__anext__() # Capture the response ID from the first event to ensure consistency if self._cached_response_id is None and hasattr(chunk, "response"): @@ -569,7 +573,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if not self.base_iterator or not hasattr(self.base_iterator, "__anext__"): raise StopAsyncIteration - chunk: Final[ResponsesAPIStreamingResponse] = await cast(Any, self.base_iterator).__anext__() + chunk: Final[ResponsesAPIStreamingResponse] = await cast( # cast-ok: hasattr __anext__ checked above + "AsyncIterator[ResponsesAPIStreamingResponse]", self.base_iterator + ).__anext__() if self._cached_response_id is None and hasattr(chunk, "response"): new_response: Final[ResponsesAPIResponse | None] = getattr(chunk, "response", None) @@ -834,7 +840,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if not self.is_async: try: if self.base_iterator and hasattr(self.base_iterator, "__next__"): - return next(cast(Any, self.base_iterator)) + return next( + cast("Iterator[ResponsesAPIStreamingResponse]", self.base_iterator) # cast-ok: hasattr-checked + ) else: raise StopIteration except StopIteration: diff --git a/litellm/responses/mcp/request_context.py b/litellm/responses/mcp/request_context.py index 0689c041a95..22869dcd502 100644 --- a/litellm/responses/mcp/request_context.py +++ b/litellm/responses/mcp/request_context.py @@ -10,14 +10,25 @@ still executes the tool, just with no credentials. from collections.abc import Iterable, Mapping, Sequence from dataclasses import dataclass -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final + +from typing_extensions import NotRequired, ReadOnly, TypedDict + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + + +class _AuthCarryingMetadata(TypedDict): + """The one key this module reads out of a request's ``metadata`` / ``litellm_metadata``.""" + + user_api_key_auth: ReadOnly[NotRequired["UserAPIKeyAuth | None"]] @dataclass(frozen=True, slots=True) class MCPRequestContext: """Everything a gateway handler must forward to MCP tool listing and execution.""" - user_api_key_auth: Any # any-ok: UserAPIKeyAuth is proxy-only; importing it here would create a cycle + user_api_key_auth: "UserAPIKeyAuth | None" mcp_auth_header: str | None = None mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = None oauth2_headers: Mapping[str, str] | None = None @@ -30,7 +41,7 @@ class MCPRequestContext: def resolve( cls, kwargs: Mapping[str, Any], - tools: Iterable[Any] | None, + tools: Iterable[object] | None, ) -> "MCPRequestContext": """ Build the context from a gateway handler's kwargs. @@ -44,9 +55,9 @@ class MCPRequestContext: ) from litellm.responses.utils import ResponsesAPIRequestUtils - litellm_metadata: Final = kwargs.get("litellm_metadata") or {} - metadata: Final = kwargs.get("metadata") or {} - user_api_key_auth: Final = ( + litellm_metadata: Final[_AuthCarryingMetadata] = kwargs.get("litellm_metadata") or {} + metadata: Final[_AuthCarryingMetadata] = kwargs.get("metadata") or {} + user_api_key_auth: Final[UserAPIKeyAuth | None] = ( kwargs.get("user_api_key_auth") or litellm_metadata.get("user_api_key_auth") or metadata.get("user_api_key_auth") diff --git a/litellm/responses/sse_output_recovery.py b/litellm/responses/sse_output_recovery.py index 208dec10c62..adc6a30319c 100644 --- a/litellm/responses/sse_output_recovery.py +++ b/litellm/responses/sse_output_recovery.py @@ -8,14 +8,17 @@ caller automatically applies to all of them. """ import json -from typing import Any, Final +from collections.abc import Mapping +from typing import Final, SupportsInt, TypeAlias, cast # noqa: TID251 # int() re-checks the cast below at runtime from litellm.constants import STREAM_SSE_DONE_STRING _MAX_CONTENT_INDEX: Final = 1024 +_ConvertibleToInt: TypeAlias = SupportsInt | str -def parse_sse_json_chunk(chunk: str) -> dict[str, Any] | None: + +def parse_sse_json_chunk(chunk: str) -> dict[str, object] | None: """Parse a single raw SSE line into a JSON object dict. Returns ``None`` for empty lines, ``event:`` lines, ``[DONE]`` markers, @@ -30,7 +33,7 @@ def parse_sse_json_chunk(chunk: str) -> dict[str, Any] | None: if not stripped_chunk or stripped_chunk == STREAM_SSE_DONE_STRING or stripped_chunk.startswith("event:"): return None try: - parsed_chunk: Final = json.loads(stripped_chunk) + parsed_chunk: Final[object] = json.loads(stripped_chunk) except json.JSONDecodeError: return None if not isinstance(parsed_chunk, dict): @@ -38,9 +41,19 @@ def parse_sse_json_chunk(chunk: str) -> dict[str, Any] | None: return parsed_chunk +def _chunk_index(parsed_chunk: Mapping[str, object], key: str, fallback: int) -> int: + raw_index: Final = parsed_chunk.get(key) + if raw_index is None: + return fallback + try: + return int(cast(_ConvertibleToInt, raw_index)) # cast-ok: int() raises TypeError otherwise, caught below + except (TypeError, ValueError): + return fallback + + def record_output_item_chunk( - parsed_chunk: dict[str, Any], - output_items: dict[int, dict[str, Any]], + parsed_chunk: Mapping[str, object], + output_items: dict[int, dict[str, object]], ) -> None: """Record an OUTPUT_ITEM_DONE chunk into ``output_items`` keyed by ``output_index`` (falling back to the next free slot when missing). @@ -48,20 +61,14 @@ def record_output_item_chunk( item: Final = parsed_chunk.get("item") if not isinstance(item, dict): return - try: - output_index_raw: Final = parsed_chunk.get("output_index") - if output_index_raw is None: - raise ValueError("missing output_index") - output_index = int(output_index_raw) - except (TypeError, ValueError): - output_index = len(output_items) + output_index: Final = _chunk_index(parsed_chunk, "output_index", len(output_items)) output_items[output_index] = item def record_output_text_chunk( - parsed_chunk: dict[str, Any], - output_items: dict[int, dict[str, Any]], - text_only_items: dict[int, dict[str, Any]], + parsed_chunk: Mapping[str, object], + output_items: Mapping[int, dict[str, object]], + text_only_items: dict[int, dict[str, object]], ) -> None: """Record an OUTPUT_TEXT_DONE chunk as a synthetic message item in ``text_only_items``. Real OUTPUT_ITEM_DONE events already captured in @@ -71,13 +78,7 @@ def record_output_text_chunk( if not isinstance(text, str): return - try: - output_index_raw: Final = parsed_chunk.get("output_index") - if output_index_raw is None: - raise ValueError("missing output_index") - output_index = int(output_index_raw) - except (TypeError, ValueError): - output_index = len(text_only_items) + output_index: Final = _chunk_index(parsed_chunk, "output_index", len(text_only_items)) if output_index in output_items: return @@ -97,13 +98,7 @@ def record_output_text_chunk( if not isinstance(content, list): return - try: - content_index_raw: Final = parsed_chunk.get("content_index") - if content_index_raw is None: - raise ValueError("missing content_index") - content_index = int(content_index_raw) - except (TypeError, ValueError): - content_index = len(content) + content_index: Final = _chunk_index(parsed_chunk, "content_index", len(content)) if content_index < 0 or content_index > _MAX_CONTENT_INDEX: return diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index a6924c1d87a..5c0d6fc536e 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -49,6 +49,21 @@ if TYPE_CHECKING: ResponsesClientWebSocket, ) + class _StreamCachingHandler(Protocol): + """The ``_llm_caching_handler`` attached to a logging object, as this module uses it.""" + + original_function: Callable[..., object] + + def _should_store_result_in_cache( + self, original_function: Callable[..., object], kwargs: Mapping[str, object] + ) -> bool: ... + + class PiiUnmaskingGuardrailCallback(PresidioGuardrailCallback, Protocol): + """Guardrail callback that can also reverse its own masking, selected by + ``llm_http_handler`` on exactly this attribute.""" + + def _unmask_pii_text(self, text: str, pii_tokens: Mapping[str, str]) -> str: ... + class ProjectQuotaCallback(Protocol): async def enforce_project_io_token_quota_for_frame( @@ -84,6 +99,11 @@ def _load_json_object(payload: str | bytes) -> dict[str, object]: return json.loads(payload) +def _load_json_value(payload: str | bytes) -> object: + """Parse a JSON payload whose top-level shape the caller narrows itself.""" + return json.loads(payload) + + def _model_id_from_metadata(litellm_metadata: dict[str, object] | None) -> str | None: model_info: Final = litellm_metadata.get("model_info") if litellm_metadata else None model_id: Final = model_info.get("id") if _is_json_object(model_info) else None @@ -243,10 +263,10 @@ class BaseResponsesAPIStreamingIterator: try: # Parse the JSON chunk - parsed_chunk: Final = json.loads(chunk) + parsed_chunk: Final = _load_json_value(chunk) # Format as ResponsesAPIStreamingResponse - if isinstance(parsed_chunk, dict): + if _is_json_object(parsed_chunk): if self.responses_api_provider_config is None: raise ValueError("responses_api_provider_config is required to process live streaming chunks") openai_responses_api_chunk: Final = self.responses_api_provider_config.transform_streaming_response( @@ -529,7 +549,7 @@ class BaseResponsesAPIStreamingIterator: if response_obj is None: return - caching_handler: Final = getattr(self.logging_obj, "_llm_caching_handler", None) + caching_handler: Final[_StreamCachingHandler | None] = getattr(self.logging_obj, "_llm_caching_handler", None) if caching_handler is None: return @@ -547,7 +567,7 @@ class BaseResponsesAPIStreamingIterator: if preset_cache_key is not None: request_kwargs["cache_key"] = preset_cache_key - if not caching_handler._should_store_result_in_cache( + if not caching_handler._should_store_result_in_cache( # pyright: ignore[reportPrivateUsage] # no public API original_function=caching_handler.original_function, kwargs=request_kwargs, ): @@ -1401,7 +1421,7 @@ async def _enforce_frame_project_quota( if not quota_callbacks: return try: - msg_obj = json.loads(raw_message) + msg_obj: Final = _load_json_value(raw_message) except (json.JSONDecodeError, TypeError): return if not _is_json_object(msg_obj) or msg_obj.get("type") != "response.create": @@ -1451,7 +1471,7 @@ class ResponsesWebSocketStreaming: user_api_key_dict: UserAPIKeyAuth | None = None, request_data: dict[str, object] | None = None, first_message: str | None = None, - guardrail_callbacks: list[Any] | None = None, + guardrail_callbacks: list[PiiUnmaskingGuardrailCallback] | None = None, output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, quota_callbacks: Sequence[ProjectQuotaCallback] | None = None, authorized_model: str | None = None, @@ -1464,7 +1484,7 @@ class ResponsesWebSocketStreaming: self.messages: list[dict[str, object]] = [] self.input_messages: list[dict[str, object]] = [] self.first_message = first_message - self.guardrail_callbacks: list[Any] = guardrail_callbacks or [] + self.guardrail_callbacks: list[PiiUnmaskingGuardrailCallback] = guardrail_callbacks or [] self.output_guardrail_callbacks: list[PresidioGuardrailCallback] = output_guardrail_callbacks or [] self.quota_callbacks: tuple[ProjectQuotaCallback, ...] = tuple(quota_callbacks) if quota_callbacks else () # Model name authorized at connection time; enforced on every @@ -1780,7 +1800,9 @@ class ResponsesWebSocketStreaming: continue text = content_block.get("text") if isinstance(text, str): - unmasked = cb._unmask_pii_text(text, pii_tokens) + unmasked = cb._unmask_pii_text( # pyright: ignore[reportPrivateUsage] # no public unmasker + text, pii_tokens + ) if unmasked != text: content_block["text"] = unmasked modified = True @@ -1789,7 +1811,9 @@ class ResponsesWebSocketStreaming: if event_type in self._DELTA_EVENT_TYPES: delta: Final = evt_obj.get("delta") if isinstance(delta, str): - unmasked = cb._unmask_pii_text(delta, pii_tokens) + unmasked = cb._unmask_pii_text( # pyright: ignore[reportPrivateUsage] # no public unmasker + delta, pii_tokens + ) if unmasked != delta: evt_obj["delta"] = unmasked return json.dumps(evt_obj) diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 716a815547d..0ff6bc8a7d2 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -1,15 +1,17 @@ import base64 import re -from collections.abc import Iterable, Mapping +from collections.abc import Iterable, Mapping, Sequence from typing import Any, Final, Optional, Union, cast, get_type_hints, overload from pydantic import BaseModel +from typing_extensions import TypeIs # noqa: TID251 # narrows untyped wire payloads without a runtime conversion import litellm from litellm._logging import verbose_logger from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.types.llms.openai import ( AllMessageValues, + OutputTokensDetails, ResponseAPIUsage, ResponseInputParam, ResponsesAPIOptionalRequestParams, @@ -26,6 +28,16 @@ from litellm.types.utils import ( ) +def _is_object_sequence(value: object) -> TypeIs[Sequence[object]]: # guard-ok: a list is a Sequence of anything + return isinstance(value, list) + + +def _is_object_dict( + value: object, +) -> TypeIs[dict[str, object]]: # guard-ok: wire dicts have str keys # mutable-ok: callers rewrite ids in place + return isinstance(value, dict) + + def normalize_responses_api_stream_options( stream_options: object, ) -> ResponsesAPIStreamOptions | None: @@ -703,12 +715,12 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_ids_in_annotations( - annotations: Any, + annotations: object, custom_llm_provider: str | None, model_id: str | None, ) -> None: """Encode ``container_id`` on each annotation (e.g. ``container_file_citation``).""" - if not annotations or not isinstance(annotations, list): + if not annotations or not _is_object_sequence(annotations): return for ann in annotations: ResponsesAPIRequestUtils._encode_container_id_on_output_item( @@ -719,16 +731,16 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_ids_in_message_content( - content: Any, + content: object, custom_llm_provider: str | None, model_id: str | None, ) -> None: """Walk message ``content`` parts and encode citation ``container_id`` values.""" if not content: return - if isinstance(content, list): + if _is_object_sequence(content): for part in content: - if isinstance(part, dict): + if _is_object_dict(part): ResponsesAPIRequestUtils._encode_container_ids_in_annotations( part.get("annotations"), custom_llm_provider, @@ -743,7 +755,7 @@ class ResponsesAPIRequestUtils: @staticmethod def _encode_container_id_on_output_item( - item: Any, + item: object, custom_llm_provider: str | None, model_id: str | None, ) -> None: @@ -770,14 +782,14 @@ class ResponsesAPIRequestUtils: container_id=container_id, ) - if isinstance(item, dict): + if _is_object_dict(item): cid: Final = item.get("container_id") if isinstance(cid, str): enc = _maybe_encode(cid) if enc is not None: - item["container_id"] = enc + item["container_id"] = enc # rebind-ok: this helper's contract is to rewrite the item in place nested: Final = item.get("code_interpreter_call") - if isinstance(nested, dict): + if _is_object_dict(nested): nc: Final = nested.get("container_id") if isinstance(nc, str): enc = _maybe_encode(nc) @@ -803,7 +815,7 @@ class ResponsesAPIRequestUtils: exc_info=True, ) - nested_obj: Final = getattr(item, "code_interpreter_call", None) + nested_obj: Final[object] = getattr(item, "code_interpreter_call", None) if nested_obj is not None: ResponsesAPIRequestUtils._encode_container_id_on_output_item( nested_obj, @@ -820,24 +832,24 @@ class ResponsesAPIRequestUtils: @staticmethod def _collect_container_ids_from_annotations( - annotations: Any, + annotations: object, collected: set[str], ) -> None: - if not annotations or not isinstance(annotations, list): + if not annotations or not _is_object_sequence(annotations): return for ann in annotations: ResponsesAPIRequestUtils._collect_container_ids_from_output_item(ann, collected) @staticmethod def _collect_container_ids_from_message_content( - content: Any, + content: object, collected: set[str], ) -> None: if not content: return - if isinstance(content, list): + if _is_object_sequence(content): for part in content: - if isinstance(part, dict): + if _is_object_dict(part): ResponsesAPIRequestUtils._collect_container_ids_from_annotations( part.get("annotations"), collected, @@ -850,19 +862,19 @@ class ResponsesAPIRequestUtils: @staticmethod def _collect_container_ids_from_output_item( - item: Any, + item: object, collected: set[str], ) -> None: """Collect managed or raw ``container_id`` values from one output item.""" if item is None: return - if isinstance(item, dict): + if _is_object_dict(item): cid: Final = item.get("container_id") if isinstance(cid, str) and cid: collected.add(cid) nested: Final = item.get("code_interpreter_call") - if isinstance(nested, dict): + if _is_object_dict(nested): nc: Final = nested.get("container_id") if isinstance(nc, str) and nc: collected.add(nc) @@ -877,7 +889,7 @@ class ResponsesAPIRequestUtils: if isinstance(cid_attr, str) and cid_attr: collected.add(cid_attr) - nested_obj: Final = getattr(item, "code_interpreter_call", None) + nested_obj: Final[object] = getattr(item, "code_interpreter_call", None) if nested_obj is not None: ResponsesAPIRequestUtils._collect_container_ids_from_output_item(nested_obj, collected) @@ -1108,7 +1120,9 @@ class ResponseAPILoggingUtils: cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None), ) completion_tokens_details: CompletionTokensDetailsWrapper | None = None - output_tokens_details: Final = getattr(response_api_usage, "output_tokens_details", None) + output_tokens_details: Final[OutputTokensDetails | None] = getattr( + response_api_usage, "output_tokens_details", None + ) if output_tokens_details: completion_tokens_details = CompletionTokensDetailsWrapper( reasoning_tokens=getattr(output_tokens_details, "reasoning_tokens", None), diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index bd9a7bff101..22d27bc3266 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -1,7 +1,14 @@ # litellm/proxy/vector_stores/vector_store_registry.py import json +from collections.abc import Mapping from datetime import datetime, timezone -from typing import TYPE_CHECKING, Any, Final, get_args +from typing import ( + TYPE_CHECKING, + Any, # noqa: TID251 # untyped non_default_params dict is the only source of the unknown key type + Final, + cast, # noqa: TID251 # untyped non_default_params dict is the only source of the unknown key type + get_args, +) from litellm._logging import verbose_logger from litellm.litellm_core_utils.core_helpers import remove_items_at_indices @@ -336,7 +343,9 @@ class VectorStoreRegistry: try: # Check if it still exists in database db_vector_store = await ManagedVectorStoresRepository(prisma_client).table.find_unique( - where={"vector_store_id": vector_store_id} + where=cast( # cast-ok: every value is already an object, only the popped id is stub-untyped + "Mapping[str, object]", {"vector_store_id": vector_store_id} + ) ) if db_vector_store is None: # Vector store was deleted from database, remove from cache diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 03318718fb5..d5cddbb2a20 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,6 +1,6 @@ { "ANN001": { - "limit": 3018 + "limit": 3016 }, "ANN002": { "limit": 71 @@ -9,7 +9,7 @@ "limit": 827 }, "ANN201": { - "limit": 2016 + "limit": 2013 }, "ANN202": { "limit": 852 @@ -24,7 +24,7 @@ "limit": 133 }, "ANN401": { - "limit": 1188 + "limit": 1157 }, "ASYNC230": { "limit": 11 @@ -33,13 +33,13 @@ "limit": 2 }, "B006": { - "limit": 177 + "limit": 176 }, "B008": { "limit": 503 }, "B009": { - "limit": 59 + "limit": 58 }, "B010": { "limit": 190 @@ -231,7 +231,7 @@ "limit": 5 }, "TID251": { - "limit": 1212 + "limit": 1201 }, "TRY002": { "limit": 524 diff --git a/tests/proxy_unit_tests/test_jwt_key_mapping.py b/tests/proxy_unit_tests/test_jwt_key_mapping.py index 65d075b7f99..4b50f83e9eb 100644 --- a/tests/proxy_unit_tests/test_jwt_key_mapping.py +++ b/tests/proxy_unit_tests/test_jwt_key_mapping.py @@ -416,6 +416,34 @@ async def test_update_returns_404_when_not_found(): assert exc_info.value.status_code == 404 +@pytest.mark.asyncio +async def test_update_returns_404_when_row_deleted_before_write(): + """A mapping deleted between the read and the write must 404, not 500. + + Prisma's update returns None when the row is gone, and the endpoint used to + dereference it for the cache key. + """ + from litellm.proxy._types import UpdateJWTKeyMappingRequest + + mock_prisma = _mock_prisma() + mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = _mock_mapping() + mock_prisma.db.litellm_jwtkeymapping.update.return_value = None + mock_cache = AsyncMock() + + data = UpdateJWTKeyMappingRequest(id="mapping-1", description="test") + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), # test-quality-ok: proxy_server module global is the endpoint's only injection point + ): + with pytest.raises(HTTPException) as exc_info: + await update_jwt_key_mapping( + data=data, user_api_key_dict=_make_admin_auth() + ) + assert exc_info.value.status_code == 404 + assert exc_info.value.detail == "Mapping not found" + + @pytest.mark.asyncio async def test_info_returns_404_when_not_found(): """Getting info for non-existent mapping should return 404.""" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py index a87f5384c6f..7ce62fdf648 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py @@ -428,3 +428,60 @@ async def test_migrate_legacy_grant_ids_no_ops_without_config_agents(): assert await registry.migrate_legacy_grant_ids(table=table) == GrantMigrationResult(rewritten=0, missed=0) table.find_many.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_agent_in_db_raises_when_row_deleted_mid_update(): + """Prisma's update returns None when the row vanished between read and write. Without a + guard the code dereferences None and reports an opaque AttributeError instead of the id.""" + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) + + with pytest.raises(Exception, match="Error updating agent in DB") as exc_info: + await registry.update_agent_in_db( + agent_id="agent-123", + agent={ + "agent_name": "Updated Agent", + "agent_card_params": _sample_agent_card_params(), + "litellm_params": {}, + }, + prisma_client=mock_prisma, + updated_by="test-user", + ) + + assert str(exc_info.value) == "Error updating agent in DB: Agent not found, passed agent_id=agent-123" + + +@pytest.mark.asyncio +async def test_patch_agent_in_db_raises_when_row_deleted_mid_update(): + """Same race on PATCH: the existing row is read, then deleted before the update lands.""" + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.find_unique = AsyncMock( + return_value={"agent_id": "agent-123", "agent_name": "Old Agent", "object_permission_id": None} + ) + mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None) + + with pytest.raises(Exception, match="Error patching agent in DB") as exc_info: + await registry.patch_agent_in_db( + agent_id="agent-123", + agent={"agent_name": "Patched Agent"}, + prisma_client=mock_prisma, + updated_by="test-user", + ) + + assert str(exc_info.value) == "Error patching agent in DB: Agent not found, passed agent_id=agent-123" + + +@pytest.mark.asyncio +async def test_delete_agent_from_db_raises_when_row_already_gone(): + """Prisma's delete returns None for a missing row, which dict() cannot consume.""" + registry: Final = AgentRegistry() + mock_prisma: Final = MagicMock() + mock_prisma.db.litellm_agentstable.delete = AsyncMock(return_value=None) + + with pytest.raises(Exception, match="Error deleting agent from DB") as exc_info: + await registry.delete_agent_from_db(agent_id="agent-123", prisma_client=mock_prisma) + + assert str(exc_info.value) == "Error deleting agent from DB: Agent not found, passed agent_id=agent-123" diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py index d1ccd86044c..99b09f48fe5 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_claude_code_marketplace.py @@ -231,6 +231,30 @@ async def test_update_plugin_db_error_maps_to_structured_500(): assert "connection lost" in exc_info.value.detail["error"] +@pytest.mark.asyncio +async def test_update_plugin_deleted_mid_update_returns_404(): + """A concurrent delete between the find_unique pre-check and the update makes prisma's + update return None; that must surface the same 404 as a plain miss, not an AttributeError.""" + name = "my-monorepo-plugin" + await register_plugin( + request=RegisterPluginRequest(name=name, source=_GIT_SUBDIR_SOURCE, version="1.0.0"), + user_api_key_dict=_USER, + ) + + table = litellm.proxy.proxy_server.prisma_client.db.litellm_claudecodeplugintable + table.update = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await update_plugin( + plugin_name=name, + request=UpdatePluginRequest(source={"source": "github", "repo": "org/replacement"}), + user_api_key_dict=_USER, + ) + + assert exc_info.value.status_code == 404 + assert exc_info.value.detail == {"error": f"Plugin '{name}' not found"} + + @pytest.mark.asyncio async def test_get_marketplace_skips_plugin_with_null_manifest(): await register_plugin( diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 25c177a308d..c575265ed91 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -695,7 +695,7 @@ def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock "object_permission_id": None, "object_permission": None, "litellm_budget_table": None, - "dict": lambda self=None: { + "model_dump": lambda self=None: { "spend": 25.0, "user_id": "enduser-implicit", "blocked": False, diff --git a/tests/test_litellm/proxy/db/mcp_server/test_db.py b/tests/test_litellm/proxy/db/mcp_server/test_db.py index aa40ec0d76c..e2440e49f19 100644 --- a/tests/test_litellm/proxy/db/mcp_server/test_db.py +++ b/tests/test_litellm/proxy/db/mcp_server/test_db.py @@ -4,7 +4,11 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from litellm.proxy._experimental.mcp_server.db import get_mcp_servers_by_team +from litellm.proxy._experimental.mcp_server.db import ( + approve_mcp_server, + get_mcp_servers_by_team, + reject_mcp_server, +) def _prisma_client_returning(team_record: object) -> MagicMock: @@ -38,3 +42,30 @@ async def test_fetch_mcp_servers_by_team(team_record, expected): where={"team_id": "team-123"}, include={"object_permission": True}, ) + + +def _prisma_client_with_missing_mcp_server_row() -> MagicMock: + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.update = AsyncMock(return_value=None) + return prisma_client + + +@pytest.mark.asyncio +async def test_approve_mcp_server_raises_value_error_when_row_missing(): + prisma_client = _prisma_client_with_missing_mcp_server_row() + + with pytest.raises(ValueError, match=r"^MCP server not found, passed server_id=server-gone$"): + await approve_mcp_server(prisma_client, "server-gone", touched_by="admin") + + +@pytest.mark.asyncio +async def test_reject_mcp_server_raises_value_error_when_row_missing(): + prisma_client = _prisma_client_with_missing_mcp_server_row() + + with pytest.raises(ValueError, match=r"^MCP server not found, passed server_id=server-gone$"): + await reject_mcp_server( + prisma_client, + "server-gone", + touched_by="admin", + review_notes="spam", + ) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 729dbce6b9a..26b3890464e 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -1,8 +1,11 @@ +from unittest.mock import AsyncMock, MagicMock + import pytest from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy.guardrails.guardrail_registry import ( get_guardrail_initializer_from_hooks, + GuardrailRegistry, InMemoryGuardrailHandler, ) from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams @@ -657,3 +660,22 @@ class TestScanOnlyToolResultsInitRefusal: "scan_only_tool_results": True, }, ) + + +@pytest.mark.asyncio +async def test_update_guardrail_in_db_raises_when_row_missing(): + prisma_client = MagicMock() + prisma_client.db.litellm_guardrailstable.update = AsyncMock(return_value=None) + + with pytest.raises( + Exception, + match=r"^Error updating guardrail in DB: Guardrail not found, passed guardrail_id=missing-guardrail$", + ): + await GuardrailRegistry().update_guardrail_in_db( + guardrail_id="missing-guardrail", + guardrail=Guardrail( + guardrail_name="missing-guardrail", + litellm_params=LitellmParams(guardrail="bedrock", mode="pre_call"), + ), + prisma_client=prisma_client, + ) 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 5f6c1a2375b..6d8f7f4ccdb 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 @@ -5446,3 +5446,46 @@ async def test_handle_group_membership_changes_already_in_team_is_noop(mocker): ) assert mock_team_member_add.await_count == 2 + + +@pytest.mark.asyncio +async def test_patch_group_404s_when_team_deleted_mid_request(mocker): + """A group deleted between the existence check and the write must 404. + + Prisma returns None from both the update and the refresh reads once the row is + gone, and patch_group used to dereference that None while building the response. + """ + group_id = "team-gone" + + snapshot_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[Member(user_id="zed", role="user")], + metadata={"externalId": "grp-ext"}, + ) + + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="replace", path="displayName", value="Renamed")], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=[snapshot_team, None, None]) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=None) + + mocker.patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + mocker.patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + with pytest.raises(ProxyException) as exc_info: + await patch_group(group_id=group_id, patch_ops=patch_ops) + + assert exc_info.value.code == "404" + assert f"Group not found with ID: {group_id}" in exc_info.value.message diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index f86e17c61b0..8bce967b316 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2398,7 +2398,7 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker): mock_user_row.user_id = "admin-creator" mock_user_row.user_email = "admin@example.com" mock_user_row.teams = [] - mock_user_row.json.return_value = "{}" + mock_user_row.model_dump_json.return_value = "{}" mock_user_row.model_dump.return_value = { "user_id": "admin-creator", "user_email": "admin@example.com", diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 0c615cbaa32..a37c4f72b3d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -7766,6 +7766,45 @@ async def test_validate_key_list_check_key_hash_not_found(): assert "Key Hash not found" in exc_info.value.message +@pytest.mark.asyncio +async def test_validate_key_list_check_key_hash_row_missing(): + """A key_hash with no row reaches the same 'Key Hash not found' 403 as a failed + lookup, instead of blowing up inside the ownership check on a None row.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable( + user_id="test-user", + user_email="test@example.com", + teams=[], + organization_memberships=[], + ) + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=None + ) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + api_key="sk-caller", + ) + + with pytest.raises(ProxyException) as exc_info: + await validate_key_list_check( + user_api_key_dict=user_api_key_dict, + user_id=None, + team_id=None, + organization_id=None, + key_alias=None, + key_hash="hash-of-a-deleted-key", + prisma_client=mock_prisma_client, + ) + + assert exc_info.value.code == "403" or exc_info.value.code == 403 + assert exc_info.value.param == "key_hash" + assert "Key Hash not found" in exc_info.value.message + + @pytest.mark.asyncio async def test_validate_key_list_check_proxy_admin_viewer_skips_db_lookup(): """proxy_admin_viewer takes the same unscoped read fast-path as proxy_admin, so no diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 097230108d4..eeb1b2d50e6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -3312,6 +3312,61 @@ class TestPatchModelBlockedAuthGate: mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() +class TestPatchModelRowDeletedBeforeWrite: + """A row deleted between the read and the update makes prisma's `update` + return None. That must surface patch_model's own 404 not-found contract, + not a 500 from dereferencing the missing row.""" + + @pytest.mark.asyncio + async def test_patch_model_404s_when_update_returns_none(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + from litellm.proxy.proxy_server import ProxyException + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + existing_row = MagicMock() + existing_row.litellm_params = {"model": "openai/gpt-4o-mini"} + existing_row.model_dump.return_value = { + "model_name": "gpt-4o-mini", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": "m1"}, + } + existing_row.model_dump_json.return_value = "{}" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=existing_row + ) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=None) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( # test-quality-ok: stubs the auth gate so the test exercises the not-found branch under test + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( # test-quality-ok: stubs the cache write so the test observes only the DB result handling + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock( + return_value=ReconcileOutcome(still_desired=None, live_after=None) + ), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await patch_model( + model_id="m1", + patch_data=updateDeployment(blocked=True), + user_api_key_dict=admin, + ) + + assert exc_info.value.code == "404" + assert exc_info.value.message == "Model m1 not found on proxy." + + class TestWriteSurfacesReloadDrop: """A model-write endpoint may report success only if every row it wrote is, after the reload it triggered, live in this pod's router or deliberately environment-inactive.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index a62c98e56a7..e2d89a660c2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -1037,3 +1037,29 @@ async def test_get_organization_daily_activity_non_admin_without_org_admin_role_ assert get_daily_activity_mock.call_args.kwargs["entity_id"] == [] assert org_table_find_many.call_args.kwargs["where"] == {"organization_id": {"in": []}} + + +@pytest.mark.asyncio +async def test_find_member_if_email_missing_row_raises_documented_400(): + """A user_email lookup that matches nothing returns None instead of raising, so the + only failure the surrounding try/except models is never entered. Without an explicit + None guard the next line dereferences None and /organization/member_add answers with + an AttributeError-driven 500 rather than the documented 400. + """ + from litellm.proxy.management_endpoints.organization_endpoints import ( + find_member_if_email, + ) + + prisma_client = AsyncMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await find_member_if_email("missing@example.com", prisma_client) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == { + "error": ( + "Unique user not found for user_email=missing@example.com. Potential duplicate OR " + "non-existent user_email in LiteLLM_UserTable. Use 'user_id' instead." + ) + } diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py index c2610d88927..08e931e6405 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py @@ -1169,6 +1169,44 @@ async def test_delete_team_callback_404s_for_unknown_team(): mock_prisma.db.litellm_teamtable.update.assert_not_called() +@pytest.mark.asyncio +async def test_add_team_callbacks_rejects_team_deleted_before_write(): + """A team deleted between the existence check and the write must be rejected. + + Prisma's update returns None for a row that is gone, and add_team_callbacks + used to hand that None to the cache refresh and report success with a null + body. The rejection reuses this endpoint's own missing-team contract, so a + caller sees the same 400 whether the team vanished before or after the read. + """ + mock_prisma = _patch_prisma(_team_row(team_id="team-1", metadata={})) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=None) + + data = AddTeamCallback( + callback_name="langfuse", + callback_type="success", + callback_vars={ + "langfuse_public_key": "pk-demo", + "langfuse_secret_key": "sk-demo", + }, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.master_key", None), # test-quality-ok: proxy_server module global is the endpoint's only injection point + ): + with pytest.raises(HTTPException) as exc: + await add_team_callbacks( + data=data, + http_request=MagicMock(spec=Request), + team_id="team-1", + user_api_key_dict=_admin_auth(), + ) + + mock_prisma.db.litellm_teamtable.update.assert_called_once() + assert exc.value.status_code == 400 + assert exc.value.detail == {"error": "Team id = team-1 does not exist. Please use a different team id."} + + @pytest.mark.asyncio async def test_delete_team_callback_keeps_last_removal_from_reviving_legacy_shape(): """Removing the last entry must leave metadata["logging"] present and empty. diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index f6d74a189bc..fbe856f4adf 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -3,7 +3,7 @@ import json from contextlib import asynccontextmanager from datetime import datetime, timezone from types import SimpleNamespace -from typing import Optional, cast +from typing import Final, Optional, cast from unittest.mock import AsyncMock, MagicMock, call, patch import pytest @@ -2078,6 +2078,100 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name): assert update_call_kwargs.get("include", {}).get("object_permission") is True +@pytest.mark.asyncio +@pytest.mark.parametrize( + "endpoint_name", + ["team_model_add", "team_model_delete", "update_team_member_permissions"], +) +async def test_team_write_404s_when_row_vanishes_before_update(endpoint_name): + """A team deleted between the read and the write must 404. + + Prisma's `update` returns None when no row matches `where`, and the team + row can be deleted between the read these endpoints do first and the + update that follows it. Without the guard, `team_model_add` / + `team_model_delete` hand that None to `_refresh_cached_team` (which + reads `team_row.team_id`) and `/team/permissions_update` returns None + out of a route declared to return a team, so a plain race turns into a + 500 instead of the 404 every other not-found path in this file raises. + """ + from unittest.mock import AsyncMock, MagicMock, Mock, patch + + from fastapi import Request + + from litellm.proxy._types import ( + LitellmUserRoles, + TeamModelAddRequest, + TeamModelDeleteRequest, + UserAPIKeyAuth, + ) + from litellm.proxy.management_endpoints.team_endpoints import ( + team_model_add, + team_model_delete, + update_team_member_permissions, + ) + + mock_request = Mock(spec=Request) + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" + ) + + existing_team = MagicMock() + existing_team.team_id = "team-1234" + existing_team.model_dump.return_value = { + "team_id": "team-1234", + "models": ["bedrock-claude-sonnet-4", "openai/*"], + "team_member_permissions": [], + "spend": 0.0, + } + + call_endpoint_under_test: Final = { + "team_model_add": lambda: team_model_add( + data=TeamModelAddRequest(team_id="team-1234", models=["team-byok-1"]), + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ), + "team_model_delete": lambda: team_model_delete( + data=TeamModelDeleteRequest(team_id="team-1234", models=["openai/*"]), + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ), + "update_team_member_permissions": lambda: update_team_member_permissions( + data=UpdateTeamMemberPermissionsRequest( + team_id="team-1234", + team_member_permissions=["/key/generate"], + ), + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ), + }[endpoint_name] + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.user_api_key_cache"), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.proxy.proxy_server.proxy_logging_obj"), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( # test-quality-ok: stubs the cache write so the test observes only the DB result handling + "litellm.proxy.management_endpoints.team_endpoints._cache_team_object", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=existing_team, + ), + ): + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=existing_team + ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=None) + mock_prisma_client.db.execute_raw = AsyncMock(return_value=None) + + with pytest.raises(HTTPException) as exc_info: + await call_endpoint_under_test() + + assert exc_info.value.status_code == 404 + assert exc_info.value.detail == {"error": "Team not found, passed team_id=team-1234"} + + @pytest.mark.asyncio async def test_update_team_team_member_budget_not_passed_to_db( disable_audit_logging_for_mocked_team, diff --git a/tests/test_litellm/proxy/memory/test_memory_endpoints.py b/tests/test_litellm/proxy/memory/test_memory_endpoints.py index 3d99a600a73..da23e362bae 100644 --- a/tests/test_litellm/proxy/memory/test_memory_endpoints.py +++ b/tests/test_litellm/proxy/memory/test_memory_endpoints.py @@ -7,6 +7,7 @@ We patch the endpoint module's `_require_prisma` helper so we never need the real proxy_server import chain (which pulls heavy optional deps). """ +import json from datetime import datetime, timezone from typing import Any, Dict, List, Optional from unittest.mock import MagicMock, patch @@ -176,14 +177,36 @@ class _InMemoryTeamTable: return None -def _make_team(team_id: str, *, admin_user_ids: List[str]) -> MagicMock: - """Build a team-row stub with `members_with_roles` shaped like Prisma.""" - members = [MagicMock(user_id=uid, role="admin") for uid in admin_user_ids] - team = MagicMock() - team.team_id = team_id - team.organization_id = None # skip org-admin path in tests - team.members_with_roles = members - return team +def _make_team(team_id: str, *, admin_user_ids: List[str]) -> Any: + """Build a real Prisma team row. + + `members_with_roles` is a JSON column, so Prisma deserializes it into plain + dicts, not `Member` objects. A stub that hands back attribute-style members + would let the router read `member.role` off something Prisma never returns. + """ + from prisma import models as prisma_models + + now = datetime.now(timezone.utc) + return prisma_models.LiteLLM_TeamTable( + team_id=team_id, + organization_id=None, + members_with_roles=json.dumps([{"user_id": uid, "role": "admin"} for uid in admin_user_ids]), + metadata="{}", + models=[], + blocked=False, + created_at=now, + updated_at=now, + spend=0.0, + model_spend="{}", + model_max_budget="{}", + admins=[], + members=[], + team_member_permissions=[], + access_group_ids=[], + policies=[], + default_team_member_models=[], + allow_team_guardrail_config=False, + ) def _make_prisma() -> MagicMock: @@ -653,6 +676,39 @@ class TestMemoryEndpoints: assert resp.json()["value"] == "new" assert len(table.rows) == 1 + def test_put_memory_row_deleted_mid_update_returns_404(self): + """ + A concurrent DELETE landing between the visibility read and the write + makes Prisma's `update` return None. That must surface the same 404 the + read path uses, not an AttributeError bubbling out as an unhandled 500. + """ + table = self.prisma.db.litellm_memorytable + table.rows.append( + _make_row( + memory_id="m1", + key="notes", + value="old", + user_id="user-a", + team_id="team-a", + ) + ) + + async def vanished(*_args, **_kwargs): + return None + + original_update = table.update + table.update = vanished + + client = _make_client(_user_auth("user-a", "team-a")) + try: + with _patch_prisma(self.prisma): + resp = client.put("/v1/memory/notes", json={"value": "new"}) + finally: + table.update = original_update + + assert resp.status_code == 404, resp.text + assert resp.json()["detail"] == "Memory with key 'notes' not found" + def test_put_memory_explicit_null_metadata_clears_field(self): """ prisma-client-python can't write a true SQL NULL to a `Json?` column diff --git a/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py b/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py index 4fb8e54e68d..3e8e1e9dff8 100644 --- a/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py +++ b/tests/test_litellm/proxy/prompts/test_prompt_endpoints_crud.py @@ -191,3 +191,58 @@ async def test_get_prompt_info_by_base_id(): response.prompt_spec.prompt_id == "test_prompt" ) # Should return base ID in spec response assert response.prompt_spec.version == 3 # Should identify it as version 3 + + +@pytest.mark.asyncio +async def test_patch_prompt_row_deleted_mid_update_returns_404(): + """ + A concurrent delete between the version lookup and the write makes Prisma's + `update` return None. That must reuse the endpoint's existing not-found 404 + contract rather than blowing up into an opaque 500. + """ + from fastapi import HTTPException + + from litellm.proxy.prompts.prompt_endpoints import PatchPromptRequest, patch_prompt + + mock_user_auth = UserAPIKeyAuth( + api_key="sk-1234", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + target_row = MagicMock() + target_row.id = "row-1" + target_row.version = 1 + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_prompttable.find_many = AsyncMock( + return_value=[target_row] + ) + mock_prisma_client.db.litellm_prompttable.update = AsyncMock(return_value=None) + + existing_prompt = PromptSpec( + prompt_id="test_prompt.v1", + litellm_params=PromptLiteLLMParams( + prompt_id="test_prompt", prompt_integration="dotprompt" + ), + prompt_info=PromptInfo(prompt_type="db"), + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch( # test-quality-ok: stubs the collaborator so the test pins the endpoint's own error contract + "litellm.proxy.prompts.prompt_registry.IN_MEMORY_PROMPT_REGISTRY" + ) as mock_registry, + ): + mock_registry.get_prompt_by_id.return_value = existing_prompt + + with pytest.raises(HTTPException) as exc_info: + await patch_prompt( + prompt_id="test_prompt", + request=PatchPromptRequest(prompt_info=PromptInfo(prompt_type="db")), + user_api_key_dict=mock_user_auth, + ) + + assert exc_info.value.status_code == 404 + assert ( + exc_info.value.detail + == "Prompt with ID test_prompt not found in environment development" + ) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 31d2a6cef98..d12684beb92 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1850,6 +1850,39 @@ def test_add_team_models_to_all_models_excludes_other_teams_byok_with_shared_nam assert result == {"model-a-id": {"team-a"}} +@pytest.mark.asyncio +async def test_non_admin_all_models_raises_400_when_user_row_missing(): + """ + Regression test: a key whose user row no longer exists made find_unique return + None, and _check_if_model_is_team_model then dereferenced it + (`model_team_id in user_row.teams`) and raised AttributeError, surfacing as a + 500. The miss must reuse the 400 "User not found" contract the neighbouring + except-branch already raises. + """ + from fastapi import HTTPException + + from litellm.proxy.proxy_server import non_admin_all_models + + prisma_client = MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + llm_router = MagicMock() + llm_router.get_model_list.return_value = [ + {"model_info": {"id": "gpt-4-model-1", "team_id": "team-a"}}, + ] + + with pytest.raises(HTTPException) as exc_info: + await non_admin_all_models( + all_models=[], + llm_router=llm_router, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", user_id="deleted-user"), + prisma_client=prisma_client, + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == {"error": "User not found"} + + @pytest.mark.asyncio async def test_apply_search_filter_matches_team_public_model_name(): """ diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 905928428b7..eae6f90863a 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -2924,6 +2924,57 @@ class TestUpdateVectorStoreAccessControlAndRedaction: assert params["api_key"] == REDACTED_BY_LITELM_STRING assert params["api_base"] == "https://api.openai.com/v1" + @pytest.mark.asyncio + async def test_update_row_deleted_mid_update_returns_404(self): + """A concurrent delete between the authorization read and the write makes Prisma's + ``update`` return None. That must reuse the not-found 404 contract instead of + turning an AttributeError into an opaque 500.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.vector_store_endpoints.management_endpoints import ( + update_vector_store, + ) + from litellm.types.vector_stores import VectorStoreUpdateRequest + + existing_row = MagicMock() + existing_row.model_dump = MagicMock( + return_value={"vector_store_id": "vs_owned", "team_id": "team-A"} + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock( + return_value=existing_row + ) + mock_prisma_client.db.litellm_managedvectorstorestable.update = AsyncMock( + return_value=None + ) + + with ( + patch( # test-quality-ok: stubs the auth gate so the test exercises the not-found branch under test + "litellm.proxy.vector_store_endpoints.management_endpoints.check_feature_access_for_user", + new_callable=AsyncMock, + ), + patch( # test-quality-ok: stubs the auth gate so the test exercises the not-found branch under test + "litellm.proxy.vector_store_endpoints.management_endpoints._check_vector_store_access", + new_callable=AsyncMock, + return_value=True, + ), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), # test-quality-ok: proxy_server module global is the endpoint's only injection point + patch("litellm.vector_store_registry", None), # test-quality-ok: litellm module global is the only injection point for the registry + ): + with pytest.raises(HTTPException) as exc_info: + await update_vector_store( + data=VectorStoreUpdateRequest( + vector_store_id="vs_owned", + vector_store_description="new desc", + ), + user_api_key_dict=UserAPIKeyAuth(user_id="owner", team_id="team-A"), + ) + + assert exc_info.value.status_code == 404 + assert exc_info.value.detail == "Vector store with ID vs_owned not found" + class TestAzureAIDocumentWritePassthroughPermission: """Regression tests for the Azure AI Search passthrough write mapping. diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 627811a7f1d..72fdfbef8c5 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,9 +1,9 @@ { "LIT001": { - "limit": 22805 + "limit": 22749 }, "LIT002": { - "limit": 26873 + "limit": 26867 }, "LIT003": { "limit": 269 @@ -15,22 +15,22 @@ "limit": 0 }, "LIT006": { - "limit": 1069 + "limit": 1066 }, "LIT007": { "limit": 0 }, "LIT008": { - "limit": 950 + "limit": 948 }, "LIT009": { "limit": 0 }, "LIT010": { - "limit": 16673 + "limit": 16655 }, "LIT011": { - "limit": 5588 + "limit": 5586 }, "LIT012": { "limit": 4510 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 87f050e417f..c0509ec3358 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26685,7 +26685,7 @@ export interface components { * Admins * @default [] */ - admins: unknown[]; + admins: string[]; /** * Allow Team Guardrail Config * @default false @@ -26725,7 +26725,7 @@ export interface components { * Members * @default [] */ - members: unknown[]; + members: string[]; /** * Members With Roles * @default [] @@ -26755,7 +26755,7 @@ export interface components { * Models * @default [] */ - models: unknown[]; + models: string[]; object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null; /** Object Permission Id */ object_permission_id?: string | null; @@ -27961,7 +27961,7 @@ export interface components { * Admins * @default [] */ - admins: unknown[]; + admins: string[]; /** * Allow Team Guardrail Config * @default false @@ -27991,7 +27991,7 @@ export interface components { * Members * @default [] */ - members: unknown[]; + members: string[]; /** * Members With Roles * @default [] @@ -28021,7 +28021,7 @@ export interface components { * Models * @default [] */ - models: unknown[]; + models: string[]; object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null; /** Object Permission Id */ object_permission_id?: string | null; @@ -30065,7 +30065,7 @@ export interface components { * Admins * @default [] */ - admins: unknown[]; + admins: string[]; /** Allowed Passthrough Routes */ allowed_passthrough_routes?: unknown[] | null; /** Allowed Vector Store Indexes */ @@ -30109,7 +30109,7 @@ export interface components { * Members * @default [] */ - members: unknown[]; + members: string[]; /** * Members With Roles * @default [] @@ -30135,7 +30135,7 @@ export interface components { * Models * @default [] */ - models: unknown[]; + models: string[]; object_permission?: components["schemas"]["LiteLLM_ObjectPermissionBase"] | null; /** Organization Id */ organization_id?: string | null; @@ -33913,7 +33913,7 @@ export interface components { * Admins * @default [] */ - admins: unknown[]; + admins: string[]; /** * Allow Team Guardrail Config * @default false @@ -33943,7 +33943,7 @@ export interface components { * Members * @default [] */ - members: unknown[]; + members: string[]; /** * Members With Roles * @default [] @@ -33973,7 +33973,7 @@ export interface components { * Models * @default [] */ - models: unknown[]; + models: string[]; object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null; /** Object Permission Id */ object_permission_id?: string | null; @@ -34043,7 +34043,7 @@ export interface components { * Admins * @default [] */ - admins: unknown[]; + admins: string[]; /** * Allow Team Guardrail Config * @default false @@ -34078,7 +34078,7 @@ export interface components { * Members * @default [] */ - members: unknown[]; + members: string[]; /** * Members Count * @default 0 @@ -34113,7 +34113,7 @@ export interface components { * Models * @default [] */ - models: unknown[]; + models: string[]; object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null; /** Object Permission Id */ object_permission_id?: string | null; From a18dfb2a9bf57d43f714b6b3485efa296f5c966e Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Tue, 25 Aug 2026 10:58:51 -0400 Subject: [PATCH 038/281] fix(redis): redact provider objects in debug logs --- litellm/_redis.py | 16 ++++++++++++++-- tests/test_litellm/test_redis.py | 20 ++++++++++++++++++++ 2 files changed, 34 insertions(+), 2 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 6ff4c292c47..c33ae45e988 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -764,8 +764,20 @@ def get_redis_connection_pool( return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs) +def _redis_kwargs_for_logging(redis_kwargs: dict) -> dict: + return { + key: "" + if key == "credential_provider" and value is not None + else "" + if key == "redis_connect_func" and value is not None + else value + for key, value in redis_kwargs.items() + } + + def _pretty_print_redis_config(redis_kwargs: dict) -> None: """Pretty print the Redis configuration using rich with sensitive data masking""" + redis_kwargs_for_logging: Final = _redis_kwargs_for_logging(redis_kwargs) try: import logging @@ -783,7 +795,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: masker = SensitiveDataMasker() # Mask sensitive data in redis_kwargs - masked_redis_kwargs = masker.mask_dict(redis_kwargs) + masked_redis_kwargs = masker.mask_dict(redis_kwargs_for_logging) # Create main panel title title: Final = Text("Redis Configuration", style="bold blue") @@ -846,7 +858,7 @@ def _pretty_print_redis_config(redis_kwargs: dict) -> None: except ImportError: # Fallback to simple logging if rich is not available masker = SensitiveDataMasker() - masked_redis_kwargs = masker.mask_dict(redis_kwargs) + masked_redis_kwargs = masker.mask_dict(redis_kwargs_for_logging) verbose_logger.info("Redis configuration: %s", masked_redis_kwargs) except Exception as e: verbose_logger.error("Error pretty printing Redis configuration: %s", e) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index ed2045ba76f..0961357731d 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -15,6 +15,7 @@ from litellm._redis import ( _get_redis_env_kwarg_mapping, _get_redis_kwargs, _get_redis_url_kwargs, + _pretty_print_redis_config, get_redis_async_client, get_redis_client, get_redis_connection_pool, @@ -415,6 +416,25 @@ def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache): assert first_key != second_cache._get_async_client_cache_key() +def test_pretty_print_never_expands_credential_provider(capsys): + secret = "aaaa-UNIQUE-SENTINEL-bbbb" + + with patch("litellm._redis.verbose_logger.isEnabledFor", return_value=True): + _pretty_print_redis_config( + redis_kwargs={ + "host": "redis-host", + "port": 6379, + "credential_provider": _HostileCredentialProvider(secret), + } + ) + + output = capsys.readouterr().out + assert secret not in output + assert "UNIQUE" not in output + assert "_payload" not in output + assert "credential_provider" in output + + def test_redis_cache_key_does_not_serialize_connect_func(): def connect(connection): return None From f304b2ba7b18b6352c804f56ac0443e3b57ba0cc Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Tue, 25 Aug 2026 11:08:07 -0400 Subject: [PATCH 039/281] fix(redis): satisfy lint budget for log helper --- litellm/_redis.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index c33ae45e988..4cf903c4a2e 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -12,7 +12,7 @@ import json # s/o [@Frank Colson](https://www.linkedin.com/in/frank-colson-422b9b183/) for this redis implementation import os -from collections.abc import Callable +from collections.abc import Callable, Mapping from typing import Final from urllib.parse import urlsplit, urlunsplit @@ -764,7 +764,7 @@ def get_redis_connection_pool( return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs) -def _redis_kwargs_for_logging(redis_kwargs: dict) -> dict: +def _redis_kwargs_for_logging(redis_kwargs: Mapping[str, object]) -> Mapping[str, object]: return { key: "" if key == "credential_provider" and value is not None From 6dff830343b257e11365ec816a660b0eac83aaa5 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Tue, 25 Aug 2026 11:18:23 -0400 Subject: [PATCH 040/281] test(redis): explain debug logger patch --- tests/test_litellm/test_redis.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 0961357731d..70b972d3259 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -419,7 +419,9 @@ def test_redis_cache_key_does_not_inspect_provider(clear_llm_client_cache): def test_pretty_print_never_expands_credential_provider(capsys): secret = "aaaa-UNIQUE-SENTINEL-bbbb" - with patch("litellm._redis.verbose_logger.isEnabledFor", return_value=True): + with patch( # test-quality-ok: enable the debug-only printer without changing process-wide logger state + "litellm._redis.verbose_logger.isEnabledFor", return_value=True + ): _pretty_print_redis_config( redis_kwargs={ "host": "redis-host", From 1e2645203b348082cb0536bdca5d9b6152231f97 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 19:17:14 -0400 Subject: [PATCH 041/281] fix(anthropic-responses): default structured output strict to caller value Read strict from the caller's output_format/output_config.format instead of hardcoding true, defaulting to false to match OpenAI's API default. Explicit true/false values are preserved and output_format still takes precedence over output_config.format. --- .../responses_adapters/transformation.py | 2 +- .../test_responses_adapters_transformation.py | 56 +++++++++++++++++-- 2 files changed, 51 insertions(+), 7 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 6d47d0de19f..01917cd9a59 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -520,7 +520,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: "type": "json_schema", "name": "structured_output", "schema": schema, - "strict": True, + "strict": bool(output_format.get("strict")), } } diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py index 8225e7cff39..a7efff6aa33 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py @@ -132,14 +132,14 @@ class TestOutputConfigStructuredOutput: } def test_output_config_format_json_schema_converted(self): - """output_config.format.json_schema is converted to OpenAI text.format.""" + """output_config.format.json_schema is converted to OpenAI text.format, defaulting strict to False.""" req = _make_request(output_config={"format": {"type": "json_schema", "schema": self._SCHEMA}}) kwargs = _ADAPTER.translate_request(req) assert "text" in kwargs fmt = kwargs["text"]["format"] assert fmt["type"] == "json_schema" assert fmt["schema"] == self._SCHEMA - assert fmt["strict"] is True + assert fmt["strict"] is False assert fmt["name"] == "structured_output" def test_output_config_without_format_does_not_set_text(self): @@ -149,21 +149,65 @@ class TestOutputConfigStructuredOutput: assert "text" not in kwargs def test_output_format_still_works(self): - """The original output_format field still takes precedence when present.""" + """The original output_format field still takes precedence when present, defaulting strict to False.""" req = _make_request(output_format={"type": "json_schema", "schema": self._SCHEMA}) kwargs = _ADAPTER.translate_request(req) assert "text" in kwargs assert kwargs["text"]["format"]["type"] == "json_schema" + assert kwargs["text"]["format"]["strict"] is False + + def test_output_format_explicit_strict_false_is_preserved(self): + """output_format with an explicit strict=False is preserved as False.""" + req = _make_request(output_format={"type": "json_schema", "schema": self._SCHEMA, "strict": False}) + kwargs = _ADAPTER.translate_request(req) + assert kwargs["text"]["format"]["strict"] is False + + def test_output_format_explicit_strict_true_is_preserved(self): + """output_format with an explicit strict=True is preserved as True.""" + req = _make_request(output_format={"type": "json_schema", "schema": self._SCHEMA, "strict": True}) + kwargs = _ADAPTER.translate_request(req) + assert kwargs["text"]["format"]["strict"] is True def test_output_format_takes_precedence_over_output_config(self): - """output_format takes precedence over output_config.format.""" + """output_format takes precedence over output_config.format, for both schema and strict.""" other_schema = {"type": "object", "properties": {"id": {"type": "integer"}}} req = _make_request( - output_format={"type": "json_schema", "schema": self._SCHEMA}, - output_config={"format": {"type": "json_schema", "schema": other_schema}}, + output_format={"type": "json_schema", "schema": self._SCHEMA, "strict": False}, + output_config={"format": {"type": "json_schema", "schema": other_schema, "strict": True}}, ) kwargs = _ADAPTER.translate_request(req) assert kwargs["text"]["format"]["schema"] == self._SCHEMA + assert kwargs["text"]["format"]["strict"] is False + + def test_optional_property_stays_out_of_required_list(self): + """A property absent from required must stay absent from required in the translated schema.""" + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "nickname": {"type": "string"}, + }, + "required": ["name"], + "additionalProperties": False, + } + req = _make_request(output_format={"type": "json_schema", "schema": schema}) + kwargs = _ADAPTER.translate_request(req) + fmt_schema = kwargs["text"]["format"]["schema"] + assert fmt_schema["required"] == ["name"] + assert "nickname" not in fmt_schema["required"] + assert fmt_schema["additionalProperties"] is False + + def test_translate_request_does_not_mutate_input_schema(self): + """translate_request must not mutate the caller's output_format or schema dicts.""" + schema = {"type": "object", "properties": {"x": {"type": "number"}}, "required": ["x"]} + output_format = {"type": "json_schema", "schema": schema, "strict": False} + req = _make_request(output_format=output_format) + snapshot = json.loads(json.dumps(output_format)) + + _ADAPTER.translate_request(req) + + assert output_format == snapshot + assert req["output_format"] == snapshot # --------------------------------------------------------------------------- From 690656e2b3804ea25aaca3c9ad3828dd834a66bb Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 19:26:22 -0400 Subject: [PATCH 042/281] fix(anthropic-responses): preserve nested strict setting --- .../responses_adapters/transformation.py | 2 +- .../test_responses_adapters_transformation.py | 8 ++++++++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 01917cd9a59..5ed8f26afca 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -520,7 +520,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: "type": "json_schema", "name": "structured_output", "schema": schema, - "strict": bool(output_format.get("strict")), + "strict": output_format.get("strict", False), } } diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py index a7efff6aa33..4057706f297 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/responses_adapters/test_responses_adapters_transformation.py @@ -142,6 +142,14 @@ class TestOutputConfigStructuredOutput: assert fmt["strict"] is False assert fmt["name"] == "structured_output" + def test_output_config_format_explicit_strict_true_is_preserved(self): + """Nested output_config.format with explicit strict=True is preserved.""" + req = _make_request( + output_config={"format": {"type": "json_schema", "schema": self._SCHEMA, "strict": True}} + ) + kwargs = _ADAPTER.translate_request(req) + assert kwargs["text"]["format"]["strict"] is True + def test_output_config_without_format_does_not_set_text(self): """output_config with only non-format keys doesn't produce text.format.""" req = _make_request(output_config={"effort": "high"}) From 482e712da183d9d75c2d9a4caa7bfdbae248be59 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Tue, 25 Aug 2026 12:29:18 -0400 Subject: [PATCH 043/281] fix(anthropic-responses): type structured output strictness --- litellm/types/llms/anthropic.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index cc6eccbf3e0..4ce04dd0d69 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -36,6 +36,7 @@ AnthropicInputSchema = TypedDict( class AnthropicOutputSchema(TypedDict, total=False): type: Required[Literal["json_schema"]] schema: Required[dict] + strict: ReadOnly[bool] class AnthropicOutputConfig(TypedDict, total=False): From 44c7cb20aee5832734f626ad522c867b851fde40 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:10:07 -0700 Subject: [PATCH 044/281] feat(models): add missing Together AI serverless models to the cost map Backfill 21 serverless chat models, the multilingual-e5 embedding model, and Llama-Guard-4-12B from the live Together catalog with per-token pricing and capability flags. Mark 25 delisted together_ai entries with their documented deprecation_date and point superseded models at a live successor via metadata. Reprice Llama-3.3-70B-Instruct-Turbo to Together's current rate. --- ...odel_prices_and_context_window_backup.json | 349 +++++++++++++++++- model_prices_and_context_window.json | 349 +++++++++++++++++- .../test_together_ai_model_metadata.py | 149 ++++++++ 3 files changed, 843 insertions(+), 4 deletions(-) create mode 100644 tests/test_litellm/test_together_ai_model_metadata.py diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9f953e11df1..9b2b910e007 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -37886,6 +37886,7 @@ "output_cost_per_token": 1e-07 }, "together_ai/Qwen/Qwen2.5-72B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -37902,6 +37903,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { + "deprecation_date": "2026-07-10", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 262000, @@ -37914,6 +37916,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-Thinking-2507": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 256000, @@ -37926,6 +37929,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-fp8-tput": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 40000, @@ -37937,6 +37941,7 @@ "supports_tool_choice": false }, "together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "deprecation_date": "2026-06-04", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 256000, @@ -37949,11 +37954,15 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-R1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 3e-06, "litellm_provider": "together_ai", "max_input_tokens": 128000, "max_output_tokens": 20480, "max_tokens": 20480, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 7e-06, "supports_function_calling": true, @@ -37962,6 +37971,7 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-R1-0528-tput": { + "deprecation_date": "2026-02-03", "input_cost_per_token": 5.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 128000, @@ -37979,6 +37989,9 @@ "max_input_tokens": 65536, "max_output_tokens": 8192, "max_tokens": 8192, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 1.25e-06, "supports_function_calling": true, @@ -37987,9 +38000,13 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-V3.1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_tokens": 16384, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 1.7e-06, "source": "https://www.together.ai/models/deepseek-v3-1", @@ -38001,6 +38018,7 @@ "max_output_tokens": 16384 }, "together_ai/meta-llama/Llama-3.2-3B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38009,16 +38027,21 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo": { - "input_cost_per_token": 8.8e-07, + "input_cost_per_token": 1.04e-06, "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 8.8e-07, + "output_cost_per_token": 1.04e-06, + "source": "https://docs.together.ai/docs/serverless-models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true }, "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo-Free": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 0, "litellm_provider": "together_ai", "mode": "chat", @@ -38029,6 +38052,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8": { + "deprecation_date": "2026-03-31", "input_cost_per_token": 2.7e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38039,6 +38063,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-4-Scout-17B-16E-Instruct": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38049,6 +38074,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 3.5e-06, "litellm_provider": "together_ai", "mode": "chat", @@ -38059,6 +38085,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38069,6 +38096,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38079,6 +38107,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mistral-7B-Instruct-v0.1": { + "deprecation_date": "2025-11-13", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38087,6 +38116,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { + "deprecation_date": "2026-04-02", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38094,6 +38124,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38106,6 +38137,9 @@ "together_ai/moonshotai/Kimi-K2-Instruct": { "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://www.together.ai/models/kimi-k2-instruct", @@ -38149,6 +38183,7 @@ "supports_tool_choice": true }, "together_ai/zai-org/GLM-4.5-Air-FP8": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 128000, @@ -38166,6 +38201,9 @@ "max_input_tokens": 200000, "max_output_tokens": 200000, "max_tokens": 200000, + "metadata": { + "successor": "together_ai/zai-org/GLM-5.2" + }, "mode": "chat", "output_cost_per_token": 2.2e-06, "source": "https://www.together.ai/models/glm-4-6", @@ -38175,11 +38213,15 @@ "supports_tool_choice": true }, "together_ai/zai-org/GLM-4.7": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 200000, "max_output_tokens": 200000, "max_tokens": 200000, + "metadata": { + "successor": "together_ai/zai-org/GLM-5.2" + }, "mode": "chat", "output_cost_per_token": 2e-06, "source": "https://www.together.ai/models/glm-4-7", @@ -38189,11 +38231,15 @@ "supports_tool_choice": true }, "together_ai/moonshotai/Kimi-K2.5": { + "deprecation_date": "2026-05-21", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 256000, "max_output_tokens": 256000, "max_tokens": 256000, + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 2.8e-06, "source": "https://www.together.ai/models/kimi-k2-5", @@ -38203,9 +38249,13 @@ "supports_reasoning": true }, "together_ai/moonshotai/Kimi-K2-Instruct-0905": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://www.together.ai/models/kimi-k2-0905", @@ -38214,9 +38264,13 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/Qwen/Qwen3.7-Plus" + }, "mode": "chat", "output_cost_per_token": 1.5e-06, "source": "https://www.together.ai/models/qwen3-next-80b-a3b-instruct", @@ -38226,9 +38280,13 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/Qwen/Qwen3.6-Plus" + }, "mode": "chat", "output_cost_per_token": 1.5e-06, "source": "https://www.together.ai/models/qwen3-next-80b-a3b-thinking", @@ -38238,6 +38296,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3.5-397B-A17B": { + "deprecation_date": "2026-06-29", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -38249,6 +38308,292 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "together_ai/MiniMaxAI/MiniMax-M3": { + "input_cost_per_token": 3e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/Prism-ML/Ternary-Bonsai-27B": { + "input_cost_per_token": 0.0, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 0.0, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.5-9B": { + "input_cost_per_token": 1.7e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/Qwen/Qwen3.6-Plus": { + "input_cost_per_token": 5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 3e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_reasoning": true + }, + "together_ai/Qwen/Qwen3.7-Max": { + "input_cost_per_token": 1.25e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 3.75e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.7-Plus": { + "input_cost_per_token": 3.2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.28e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.8-2.4T-A95B": { + "input_cost_per_token": 2.5e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1010000, + "max_output_tokens": 1010000, + "max_tokens": 1010000, + "mode": "chat", + "output_cost_per_token": 6.25e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/arize-ai/qwen-2-1.5b-instruct": { + "input_cost_per_token": 1e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { + "input_cost_per_token": 1.4e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/deepseek-ai/DeepSeek-V4-Pro": { + "input_cost_per_token": 1.74e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/deepseek-ai/DeepSeek-V4-Pro-0813": { + "input_cost_per_token": 1.32e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 3.96e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/google/gemma-3n-E4B-it": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.2e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/google/gemma-4-31B-it": { + "input_cost_per_token": 3.9e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 9.7e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/intfloat/multilingual-e5-large-instruct": { + "input_cost_per_token": 2e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 514, + "max_tokens": 514, + "mode": "embedding", + "output_cost_per_token": 2e-08, + "output_vector_size": 1024, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/meta-llama/Llama-Guard-4-12B": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/meta-models/Muse-Glimmer-30B": { + "input_cost_per_token": 3.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/moonshotai/Kimi-K2.7-Code": { + "input_cost_per_token": 9.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/moonshotai/Kimi-K3": { + "input_cost_per_token": 3e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/nvidia/nemotron-3-ultra-550b-a55b": { + "input_cost_per_token": 6e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 512288, + "max_output_tokens": 512288, + "max_tokens": 512288, + "mode": "chat", + "output_cost_per_token": 3.6e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/pearl-ai/gemma-4-31b-it": { + "input_cost_per_token": 2.8e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8.6e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/thinkingmachines/Inkling": { + "input_cost_per_token": 1e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 4.05e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/thinkingmachines/Inkling-Small": { + "input_cost_per_token": 5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/zai-org/GLM-5.2": { + "input_cost_per_token": 1.4e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048575, + "max_output_tokens": 1048575, + "max_tokens": 1048575, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "tts-1": { "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9f953e11df1..9b2b910e007 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -37886,6 +37886,7 @@ "output_cost_per_token": 1e-07 }, "together_ai/Qwen/Qwen2.5-72B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -37902,6 +37903,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { + "deprecation_date": "2026-07-10", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 262000, @@ -37914,6 +37916,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-Thinking-2507": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 256000, @@ -37926,6 +37929,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-235B-A22B-fp8-tput": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 40000, @@ -37937,6 +37941,7 @@ "supports_tool_choice": false }, "together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": { + "deprecation_date": "2026-06-04", "input_cost_per_token": 2e-06, "litellm_provider": "together_ai", "max_input_tokens": 256000, @@ -37949,11 +37954,15 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-R1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 3e-06, "litellm_provider": "together_ai", "max_input_tokens": 128000, "max_output_tokens": 20480, "max_tokens": 20480, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 7e-06, "supports_function_calling": true, @@ -37962,6 +37971,7 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-R1-0528-tput": { + "deprecation_date": "2026-02-03", "input_cost_per_token": 5.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 128000, @@ -37979,6 +37989,9 @@ "max_input_tokens": 65536, "max_output_tokens": 8192, "max_tokens": 8192, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 1.25e-06, "supports_function_calling": true, @@ -37987,9 +38000,13 @@ "supports_tool_choice": true }, "together_ai/deepseek-ai/DeepSeek-V3.1": { + "deprecation_date": "2026-05-14", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_tokens": 16384, + "metadata": { + "successor": "together_ai/deepseek-ai/DeepSeek-V4-Pro" + }, "mode": "chat", "output_cost_per_token": 1.7e-06, "source": "https://www.together.ai/models/deepseek-v3-1", @@ -38001,6 +38018,7 @@ "max_output_tokens": 16384 }, "together_ai/meta-llama/Llama-3.2-3B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38009,16 +38027,21 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo": { - "input_cost_per_token": 8.8e-07, + "input_cost_per_token": 1.04e-06, "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 8.8e-07, + "output_cost_per_token": 1.04e-06, + "source": "https://docs.together.ai/docs/serverless-models", "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, "supports_tool_choice": true }, "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo-Free": { + "deprecation_date": "2025-11-13", "input_cost_per_token": 0, "litellm_provider": "together_ai", "mode": "chat", @@ -38029,6 +38052,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8": { + "deprecation_date": "2026-03-31", "input_cost_per_token": 2.7e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38039,6 +38063,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Llama-4-Scout-17B-16E-Instruct": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38049,6 +38074,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo": { + "deprecation_date": "2026-02-06", "input_cost_per_token": 3.5e-06, "litellm_provider": "together_ai", "mode": "chat", @@ -38059,6 +38085,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 8.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38069,6 +38096,7 @@ "supports_tool_choice": true }, "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1.8e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38079,6 +38107,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mistral-7B-Instruct-v0.1": { + "deprecation_date": "2025-11-13", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38087,6 +38116,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { + "deprecation_date": "2026-04-02", "litellm_provider": "together_ai", "mode": "chat", "supports_function_calling": true, @@ -38094,6 +38124,7 @@ "supports_tool_choice": true }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { + "deprecation_date": "2026-04-16", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "mode": "chat", @@ -38106,6 +38137,9 @@ "together_ai/moonshotai/Kimi-K2-Instruct": { "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://www.together.ai/models/kimi-k2-instruct", @@ -38149,6 +38183,7 @@ "supports_tool_choice": true }, "together_ai/zai-org/GLM-4.5-Air-FP8": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 2e-07, "litellm_provider": "together_ai", "max_input_tokens": 128000, @@ -38166,6 +38201,9 @@ "max_input_tokens": 200000, "max_output_tokens": 200000, "max_tokens": 200000, + "metadata": { + "successor": "together_ai/zai-org/GLM-5.2" + }, "mode": "chat", "output_cost_per_token": 2.2e-06, "source": "https://www.together.ai/models/glm-4-6", @@ -38175,11 +38213,15 @@ "supports_tool_choice": true }, "together_ai/zai-org/GLM-4.7": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 200000, "max_output_tokens": 200000, "max_tokens": 200000, + "metadata": { + "successor": "together_ai/zai-org/GLM-5.2" + }, "mode": "chat", "output_cost_per_token": 2e-06, "source": "https://www.together.ai/models/glm-4-7", @@ -38189,11 +38231,15 @@ "supports_tool_choice": true }, "together_ai/moonshotai/Kimi-K2.5": { + "deprecation_date": "2026-05-21", "input_cost_per_token": 5e-07, "litellm_provider": "together_ai", "max_input_tokens": 256000, "max_output_tokens": 256000, "max_tokens": 256000, + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 2.8e-06, "source": "https://www.together.ai/models/kimi-k2-5", @@ -38203,9 +38249,13 @@ "supports_reasoning": true }, "together_ai/moonshotai/Kimi-K2-Instruct-0905": { + "deprecation_date": "2026-03-06", "input_cost_per_token": 1e-06, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/moonshotai/Kimi-K3" + }, "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://www.together.ai/models/kimi-k2-0905", @@ -38214,9 +38264,13 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": { + "deprecation_date": "2026-04-02", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/Qwen/Qwen3.7-Plus" + }, "mode": "chat", "output_cost_per_token": 1.5e-06, "source": "https://www.together.ai/models/qwen3-next-80b-a3b-instruct", @@ -38226,9 +38280,13 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": { + "deprecation_date": "2026-02-25", "input_cost_per_token": 1.5e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, + "metadata": { + "successor": "together_ai/Qwen/Qwen3.6-Plus" + }, "mode": "chat", "output_cost_per_token": 1.5e-06, "source": "https://www.together.ai/models/qwen3-next-80b-a3b-thinking", @@ -38238,6 +38296,7 @@ "supports_tool_choice": true }, "together_ai/Qwen/Qwen3.5-397B-A17B": { + "deprecation_date": "2026-06-29", "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", "max_input_tokens": 262144, @@ -38249,6 +38308,292 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "together_ai/MiniMaxAI/MiniMax-M3": { + "input_cost_per_token": 3e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/Prism-ML/Ternary-Bonsai-27B": { + "input_cost_per_token": 0.0, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 0.0, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.5-9B": { + "input_cost_per_token": 1.7e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/Qwen/Qwen3.6-Plus": { + "input_cost_per_token": 5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 3e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_reasoning": true + }, + "together_ai/Qwen/Qwen3.7-Max": { + "input_cost_per_token": 1.25e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 3.75e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.7-Plus": { + "input_cost_per_token": 3.2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1000000, + "max_output_tokens": 1000000, + "max_tokens": 1000000, + "mode": "chat", + "output_cost_per_token": 1.28e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/Qwen/Qwen3.8-2.4T-A95B": { + "input_cost_per_token": 2.5e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1010000, + "max_output_tokens": 1010000, + "max_tokens": 1010000, + "mode": "chat", + "output_cost_per_token": 6.25e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/arize-ai/qwen-2-1.5b-instruct": { + "input_cost_per_token": 1e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731": { + "input_cost_per_token": 1.4e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/deepseek-ai/DeepSeek-V4-Pro": { + "input_cost_per_token": 1.74e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/deepseek-ai/DeepSeek-V4-Pro-0813": { + "input_cost_per_token": 1.32e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 3.96e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/google/gemma-3n-E4B-it": { + "input_cost_per_token": 6e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.2e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/google/gemma-4-31B-it": { + "input_cost_per_token": 3.9e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 9.7e-07, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/intfloat/multilingual-e5-large-instruct": { + "input_cost_per_token": 2e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 514, + "max_tokens": 514, + "mode": "embedding", + "output_cost_per_token": 2e-08, + "output_vector_size": 1024, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/meta-llama/Llama-Guard-4-12B": { + "input_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 2e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/meta-models/Muse-Glimmer-30B": { + "input_cost_per_token": 3.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/moonshotai/Kimi-K2.7-Code": { + "input_cost_per_token": 9.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/moonshotai/Kimi-K3": { + "input_cost_per_token": 3e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "together_ai/nvidia/nemotron-3-ultra-550b-a55b": { + "input_cost_per_token": 6e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 512288, + "max_output_tokens": 512288, + "max_tokens": 512288, + "mode": "chat", + "output_cost_per_token": 3.6e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/pearl-ai/gemma-4-31b-it": { + "input_cost_per_token": 2.8e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8.6e-07, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/thinkingmachines/Inkling": { + "input_cost_per_token": 1e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 4.05e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "together_ai/thinkingmachines/Inkling-Small": { + "input_cost_per_token": 5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 524288, + "max_output_tokens": 524288, + "max_tokens": 524288, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.together.ai/docs/serverless-models" + }, + "together_ai/zai-org/GLM-5.2": { + "input_cost_per_token": 1.4e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 1048575, + "max_output_tokens": 1048575, + "max_tokens": 1048575, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.together.ai/docs/serverless-models", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "tts-1": { "input_cost_per_character": 1.5e-05, "litellm_provider": "openai", diff --git a/tests/test_litellm/test_together_ai_model_metadata.py b/tests/test_litellm/test_together_ai_model_metadata.py new file mode 100644 index 00000000000..7d7712f3c56 --- /dev/null +++ b/tests/test_litellm/test_together_ai_model_metadata.py @@ -0,0 +1,149 @@ +import json +from pathlib import Path +from typing import Final + +import pytest + +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + +REPO_ROOT: Final = Path(__file__).parents[2] + +SERVERLESS_CHAT_MODELS: Final = ( + "together_ai/moonshotai/Kimi-K3", + "together_ai/zai-org/GLM-5.2", + "together_ai/deepseek-ai/DeepSeek-V4-Pro", + "together_ai/deepseek-ai/DeepSeek-V4-Pro-0813", + "together_ai/deepseek-ai/DeepSeek-V4-Flash-0731", + "together_ai/moonshotai/Kimi-K2.7-Code", + "together_ai/MiniMaxAI/MiniMax-M3", + "together_ai/thinkingmachines/Inkling", + "together_ai/thinkingmachines/Inkling-Small", + "together_ai/Qwen/Qwen3.8-2.4T-A95B", + "together_ai/Qwen/Qwen3.7-Max", + "together_ai/Qwen/Qwen3.7-Plus", + "together_ai/Qwen/Qwen3.6-Plus", + "together_ai/Qwen/Qwen3.5-9B", + "together_ai/nvidia/nemotron-3-ultra-550b-a55b", + "together_ai/meta-models/Muse-Glimmer-30B", + "together_ai/google/gemma-4-31B-it", + "together_ai/pearl-ai/gemma-4-31b-it", + "together_ai/google/gemma-3n-E4B-it", + "together_ai/arize-ai/qwen-2-1.5b-instruct", + "together_ai/Prism-ML/Ternary-Bonsai-27B", + "together_ai/meta-llama/Llama-Guard-4-12B", + "together_ai/openai/gpt-oss-120b", + "together_ai/openai/gpt-oss-20b", + "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo", +) + +DEPRECATED_MODELS: Final = { + "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": "2026-07-10", + "together_ai/Qwen/Qwen3.5-397B-A17B": "2026-06-29", + "together_ai/Qwen/Qwen3-Coder-480B-A35B-Instruct-FP8": "2026-06-04", + "together_ai/moonshotai/Kimi-K2.5": "2026-05-21", + "together_ai/deepseek-ai/DeepSeek-R1": "2026-05-14", + "together_ai/deepseek-ai/DeepSeek-V3.1": "2026-05-14", + "together_ai/Qwen/Qwen3-235B-A22B-Thinking-2507": "2026-04-16", + "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": "2026-04-16", + "together_ai/zai-org/GLM-4.5-Air-FP8": "2026-04-02", + "together_ai/zai-org/GLM-4.7": "2026-04-02", + "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": "2026-04-02", + "together_ai/Qwen/Qwen3-Next-80B-A3B-Instruct": "2026-04-02", + "together_ai/meta-llama/Llama-4-Maverick-17B-128E-Instruct-FP8": "2026-03-31", + "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo": "2026-03-06", + "together_ai/moonshotai/Kimi-K2-Instruct-0905": "2026-03-06", + "together_ai/meta-llama/Llama-3.2-3B-Instruct-Turbo": "2026-03-06", + "together_ai/Qwen/Qwen3-Next-80B-A3B-Thinking": "2026-02-25", + "together_ai/meta-llama/Meta-Llama-3.1-70B-Instruct-Turbo": "2026-02-25", + "together_ai/Qwen/Qwen3-235B-A22B-fp8-tput": "2026-02-06", + "together_ai/meta-llama/Llama-4-Scout-17B-16E-Instruct": "2026-02-06", + "together_ai/Qwen/Qwen2.5-72B-Instruct-Turbo": "2026-02-06", + "together_ai/meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo": "2026-02-06", + "together_ai/deepseek-ai/DeepSeek-R1-0528-tput": "2026-02-03", + "together_ai/mistralai/Mistral-7B-Instruct-v0.1": "2025-11-13", + "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo-Free": "2025-11-13", +} + + +@pytest.fixture(scope="module") +def cost_map() -> dict: + with open(REPO_ROOT / "model_prices_and_context_window.json") as f: + return json.load(f) + + +@pytest.mark.parametrize("model", SERVERLESS_CHAT_MODELS) +def test_together_serverless_chat_model_is_mapped(cost_map: dict, model: str): + info = cost_map.get(model) + assert info is not None, f"{model} missing from model_prices_and_context_window.json" + assert info["litellm_provider"] == "together_ai" + assert info["mode"] == "chat" + assert info["input_cost_per_token"] >= 0 + assert info["output_cost_per_token"] >= info["input_cost_per_token"] + assert "deprecation_date" not in info + + routed_model, provider, _, _ = get_llm_provider(model=model) + assert routed_model == model.removeprefix("together_ai/") + assert provider == "together_ai" + + +def test_together_kimi_k3_pricing_and_capabilities(cost_map: dict): + info = cost_map["together_ai/moonshotai/Kimi-K3"] + assert info["input_cost_per_token"] == 3e-06 + assert info["output_cost_per_token"] == 1.5e-05 + assert info["max_input_tokens"] == 1048576 + assert info["supports_function_calling"] is True + assert info["supports_tool_choice"] is True + assert info["supports_response_schema"] is True + assert info["supports_vision"] is True + assert info["supports_reasoning"] is True + + +def test_together_glm_52_pricing(cost_map: dict): + info = cost_map["together_ai/zai-org/GLM-5.2"] + assert info["input_cost_per_token"] == 1.4e-06 + assert info["output_cost_per_token"] == 4.4e-06 + assert info["supports_function_calling"] is True + assert info["supports_reasoning"] is True + + +def test_together_multilingual_e5_embedding_entry(cost_map: dict): + info = cost_map["together_ai/intfloat/multilingual-e5-large-instruct"] + assert info["mode"] == "embedding" + assert info["input_cost_per_token"] == 2e-08 + assert info["max_input_tokens"] == 514 + assert info["output_vector_size"] == 1024 + + +def test_together_llama_33_70b_repriced_to_current_together_rate(cost_map: dict): + info = cost_map["together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo"] + assert info["input_cost_per_token"] == 1.04e-06 + assert info["output_cost_per_token"] == 1.04e-06 + assert info["max_input_tokens"] == 131072 + + +@pytest.mark.parametrize("model", sorted(DEPRECATED_MODELS)) +def test_together_deprecated_model_carries_deprecation_date(cost_map: dict, model: str): + info = cost_map.get(model) + assert info is not None, f"{model} missing from model_prices_and_context_window.json" + assert info.get("deprecation_date") == DEPRECATED_MODELS[model] + + +def test_together_successor_metadata_points_at_live_models(cost_map: dict): + successors = { + model: info["metadata"]["successor"] + for model, info in cost_map.items() + if model.startswith("together_ai/") and "successor" in info.get("metadata", {}) + } + assert len(successors) >= 10 + for model, successor in successors.items(): + target = cost_map.get(successor) + assert target is not None, f"{model} names successor {successor} that is not in the map" + assert "deprecation_date" not in target, f"{model} names deprecated successor {successor}" + + +def test_together_backup_cost_map_in_sync(cost_map: dict): + with open(REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json") as f: + backup = json.load(f) + together_main = {k: v for k, v in cost_map.items() if k.startswith("together_ai/")} + together_backup = {k: v for k, v in backup.items() if k.startswith("together_ai/")} + assert together_backup == together_main From 4b5e3db8906625ba2128d8702d37e8d4ee95995e Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Tue, 25 Aug 2026 10:15:29 -0700 Subject: [PATCH 045/281] test(e2e): cover the Bedrock provider-feature cells customers run Adds live e2e coverage for the Bedrock combinations behind recent customer incidents: llm_provider-* response-header forwarding on /chat/completions (nonstream and stream), regional us.anthropic.* inference-profile ids over the invoke route, and the Admin UI Test Connection probe for a responses-mode Bedrock Mantle deployment. Registers the matching cells in the coverage registry and publishes the provider x feature matrix table in its README. --- tests/e2e/coverage_registry/README.md | 18 ++ .../coverage_registry/llm_conversational.yaml | 4 + tests/e2e/coverage_registry/mgmt.yaml | 1 + tests/e2e/coverage_registry/schema.py | 1 + .../test_bedrock_provider_matrix_e2e.py | 157 ++++++++++++++++++ tests/e2e/management/management_client.py | 13 ++ .../test_model_test_connection_e2e.py | 42 +++++ tests/e2e/models.py | 20 +++ 8 files changed, 256 insertions(+) create mode 100644 tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py create mode 100644 tests/e2e/management/test_model_test_connection_e2e.py diff --git a/tests/e2e/coverage_registry/README.md b/tests/e2e/coverage_registry/README.md index 5627c88dee4..da6aee84cc4 100644 --- a/tests/e2e/coverage_registry/README.md +++ b/tests/e2e/coverage_registry/README.md @@ -77,6 +77,24 @@ Strict mode exits non-zero on `@pytest.mark.covers(...)` ids that are not checke the registry. Add `--fail-on-collection-errors` when the job should also fail on pytest collection errors. +## Provider x feature matrix: customer-run Bedrock combinations + +The provider and feature combinations customers actually run get explicit cells, expanded +here as incidents surface new ones. The current Bedrock set, seeded from a customer's +production shape (regional `us.anthropic.*` inference-profile ids over both chat routes, +provider response headers for AWS-side correlation, and the Test Connection probe for a +responses-mode Bedrock Mantle deployment): + +| Cell | Feature | Covering test | +|------|---------|---------------| +| `llm.chat_completions.bedrock_converse.basic.nonstream.works` | regional `us.` id, Converse | `llm_translation/test_chat_completions_regression_e2e.py` | +| `llm.chat_completions.bedrock_converse.basic.stream.works` | regional `us.` id, Converse stream | `llm_translation/test_chat_completions_regression_e2e.py` | +| `llm.chat_completions.bedrock_invoke.basic.nonstream.works` | regional `us.` id, Invoke | `llm_translation/test_bedrock_provider_matrix_e2e.py` | +| `llm.chat_completions.bedrock_invoke.basic.stream.works` | regional `us.` id, Invoke stream | `llm_translation/test_bedrock_provider_matrix_e2e.py` | +| `llm.chat_completions.bedrock_converse.response_headers.nonstream.works` | `llm_provider-*` headers | `llm_translation/test_bedrock_provider_matrix_e2e.py` | +| `llm.chat_completions.bedrock_converse.response_headers.stream.works` | `llm_provider-*` headers, stream | `llm_translation/test_bedrock_provider_matrix_e2e.py` | +| `mgmt.model.test_connection.happy_path` | Test Connection, Bedrock Mantle | `management/test_model_test_connection_e2e.py` | + ## Status: this is a draft for review The cells were enumerated from the codebase and the tiers are a first proposal. Known diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 1d4e1e028ca..d13f17e7eb6 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -29,6 +29,10 @@ - {id: llm.chat_completions.bedrock_converse.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Bedrock vision (Anthropic/Nova)"} - {id: llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic-on-Bedrock caching"} - {id: llm.chat_completions.bedrock_converse.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic thinking on Bedrock"} +- {id: llm.chat_completions.bedrock_converse.response_headers.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: response_headers, streaming: nonstream, assertions: [works], source: "llms/bedrock/chat/converse_handler.py:248", rationale: "Bedrock request ids must surface as llm_provider-* response headers on /chat/completions so callers can correlate calls with AWS-side logs (#37003)", fail_before_fix: proven} +- {id: llm.chat_completions.bedrock_converse.response_headers.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_converse, capability: response_headers, streaming: stream, assertions: [works], source: "llms/bedrock/chat/converse_handler.py:154", rationale: "The llm_provider-* headers must also surface on streaming /chat/completions, where CustomStreamWrapper carries them instead of the nonstream setter"} +- {id: llm.chat_completions.bedrock_invoke.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "Regional inference-profile ids (us.anthropic.*) over the invoke route, the deployment shape behind a customer timeout report on v1.90.0"} +- {id: llm.chat_completions.bedrock_invoke.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: bedrock_invoke, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming with regional inference-profile ids over the invoke route"} - {id: llm.chat_completions.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Vertex AI"} - {id: llm.chat_completions.gemini.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini OpenAI-compatible chat translation"} - {id: llm.chat_completions.gemini.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini chat cost lands in SpendLogs"} diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index d8788d7fcb0..1e6de0c3d6a 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -75,3 +75,4 @@ - {id: mgmt.workflow.list.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "workflow_management_endpoints.py", rationale: "Workflow tracking (smoke)"} - {id: mgmt.credential_migration.check.happy_path, module: mgmt, tier: P2, surface: api, assertions: [happy_path], source: "key_management_endpoints.py:4252", rationale: "Encryption migration (smoke)"} - {id: mgmt.credential.new.serves_request, module: mgmt, tier: P1, surface: api, assertions: [serves_request], source: "credential_endpoints/endpoints.py:42", rationale: "Stored credential resolves into a deployment and serves a live /messages request"} +- {id: mgmt.model.test_connection.happy_path, module: mgmt, tier: P0, surface: api, assertions: [happy_path], source: "_health_endpoints.py:1785", rationale: "Test Connection for a responses-mode Bedrock Mantle deployment reaches the live provider and reports success; this exact shape 500ed on an acompletion partial before v1.91.0", fail_before_fix: proven} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index a5c723f8965..03d15f532b8 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -71,6 +71,7 @@ LlmCapability = Literal[ "pdf_input", "prompt_cache_1h", "prompt_cache_5m", + "response_headers", "service_tier", "structured_output", "thinking", diff --git a/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py b/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py new file mode 100644 index 00000000000..3c6aaa75ab3 --- /dev/null +++ b/tests/e2e/llm_translation/test_bedrock_provider_matrix_e2e.py @@ -0,0 +1,157 @@ +"""Live e2e for the Bedrock cells of the provider-feature matrix: provider +response headers on /chat/completions and regional inference-profile model ids +(us.anthropic.*) over the invoke route. + +Header forwarding is the #37003 contract: the proxy surfaces Bedrock's response +headers prefixed llm_provider- (llm_provider-x-amzn-requestid above all) so a +caller can hand AWS support the request id behind a completion. Regional +inference-profile ids are the deployment shape most Bedrock customers run; a +v1.90.0 regression timed them out, and the Converse route keeps them covered in +test_chat_completions_regression_e2e.py, so the invoke route carries its own +rows here. +""" + +from __future__ import annotations + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import StreamingResponse, unwrap +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody +from passthrough_client import PassthroughClient + +pytestmark = pytest.mark.e2e + +CONVERSE_REGIONAL_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +INVOKE_REGIONAL_BACKEND = "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0" +PROVIDER_HEADER_PREFIX = "llm_provider-" +BEDROCK_REQUEST_ID_HEADER = "llm_provider-x-amzn-requestid" + + +class _StreamDelta(BaseModel): + content: str | None = None + + +class _StreamChoice(BaseModel): + delta: _StreamDelta = _StreamDelta() + + +class _StreamChunk(BaseModel): + choices: list[_StreamChoice] = [] + + +def _streamed_text(events: list[str]) -> str: + chunks = [_StreamChunk.model_validate_json(event) for event in events] + return "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) + + +def _assert_streamed_completion(result: StreamingResponse) -> None: + assert result.ok and result.is_streaming, f"stream was not established: {result}" + assert result.stream_error is None, f"stream carried an error event: {result.stream_error}" + assert len(result.stream_events) > 1, f"stream did not deliver multiple data events: {result}" + assert _streamed_text(result.stream_events).strip(), ( + f"stream completed with no content deltas: {result.stream_events[:3]}" + ) + + +def _assert_request_id_header(result: StreamingResponse) -> None: + forwarded = [name for name in result.headers if name.startswith(PROVIDER_HEADER_PREFIX)] + assert result.headers.get(BEDROCK_REQUEST_ID_HEADER), ( + f"missing {BEDROCK_REQUEST_ID_HEADER}; forwarded provider headers: {forwarded}" + ) + + +def _assert_completion(response: ChatResponse) -> None: + assert response.choices, f"completion returned no choices: {response}" + message = response.choices[0].message + content = (message.content if message else None) or "" + assert content.strip(), f"completion carried no content: {response}" + + +def _register_bedrock_model( + client: PassthroughClient, resources: ResourceManager, prefix: str, backend: str +) -> str: + model = f"{prefix}-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody( + model=backend, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + return model + + +def _prompt() -> list[ChatMessage]: + return [ChatMessage(role="user", content="reply with one word")] + + +class TestBedrockResponseHeaders: + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.response_headers.nonstream.works", + exercised_on=[], + ) + def test_bedrock_request_id_header_surfaces( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_bedrock_model(client, resources, "e2e-bedrock-headers", CONVERSE_REGIONAL_BACKEND) + key = resources.key() + + result = client.proxy.transport.send( + "/chat/completions", + headers=client.proxy.transport.bearer(key), + json=ChatBody(model=model, messages=_prompt(), max_tokens=64), + ) + + assert result.ok, f"chat call failed: {result.status_code} {result.body[:300]}" + _assert_request_id_header(result) + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.response_headers.stream.works", + exercised_on=[], + ) + def test_bedrock_request_id_header_surfaces_on_stream( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_bedrock_model( + client, resources, "e2e-bedrock-headers-stream", CONVERSE_REGIONAL_BACKEND + ) + key = resources.key() + + result = client.proxy.chat_stream( + key, ChatBody(model=model, messages=_prompt(), stream=True, max_tokens=64) + ) + + _assert_streamed_completion(result) + _assert_request_id_header(result) + + +class TestBedrockInvokeRegionalModelIds: + @pytest.mark.covers("llm.chat_completions.bedrock_invoke.basic.nonstream.works", exercised_on=[]) + def test_invoke_regional_id_completes( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_bedrock_model(client, resources, "e2e-bedrock-invoke", INVOKE_REGIONAL_BACKEND) + key = resources.key() + + response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_prompt(), max_tokens=64))) + + _assert_completion(response) + + @pytest.mark.covers("llm.chat_completions.bedrock_invoke.basic.stream.works", exercised_on=[]) + def test_invoke_regional_id_streams( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = _register_bedrock_model(client, resources, "e2e-bedrock-invoke-stream", INVOKE_REGIONAL_BACKEND) + key = resources.key() + + result = client.proxy.chat_stream( + key, ChatBody(model=model, messages=_prompt(), stream=True, max_tokens=64) + ) + + _assert_streamed_completion(result) diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index cdc31aeea79..b2bd41e19ba 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -14,6 +14,8 @@ from e2e_http import NoBody, ProbeResult, Result, StreamingResponse, Success, Un from models import ( ChatBody, ChatMessage, + ConnectionTestBody, + ConnectionTestResponse, CustomerDeleteBody, CustomerInfoParams, CustomerNewBody, @@ -118,6 +120,17 @@ class ManagementClient: ) ) + def connection_test(self, body: ConnectionTestBody) -> Result[ConnectionTestResponse]: + """POST /health/test_connection, the call behind the Admin UI's Test + Connection button, probing the live provider with the supplied params.""" + return self.proxy.transport.post( + "/health/test_connection", + headers=self.proxy.transport.master, + json=body, + response_type=ConnectionTestResponse, + timeout=120.0, + ) + def block_key(self, key: str) -> None: _ = unwrap( self.proxy.transport.post( diff --git a/tests/e2e/management/test_model_test_connection_e2e.py b/tests/e2e/management/test_model_test_connection_e2e.py new file mode 100644 index 00000000000..a1f714df4c8 --- /dev/null +++ b/tests/e2e/management/test_model_test_connection_e2e.py @@ -0,0 +1,42 @@ +"""Live e2e for POST /health/test_connection, the API behind the Admin UI's +Test Connection button on the add-model form. + +The covered cell is a responses-mode Bedrock Mantle deployment: exactly this +shape 500ed on a functools.partial acompletion conflict before v1.91.0 while +every chat-mode probe stayed green, so the happy path asserts a real success +verdict from the live provider rather than just a 200 envelope. The region is a +literal because the endpoint rejects request-supplied os.environ/ references; +credentials fall through to the proxy's own environment (bearer token locally, +pod identity in CI). +""" + +from __future__ import annotations + +import pytest + +from e2e_http import unwrap +from management_client import ManagementClient +from models import ConnectionTestBody, LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + +MANTLE_RESPONSES_BACKEND = "bedrock_mantle/openai.gpt-5.6-luna" +MANTLE_REGION = "us-east-1" + + +class TestModelTestConnection: + @pytest.mark.covers("mgmt.model.test_connection.happy_path") + def test_bedrock_mantle_responses_connection_succeeds(self, client: ManagementClient) -> None: + response = unwrap( + client.connection_test( + ConnectionTestBody( + litellm_params=LiteLLMParamsBody( + model=MANTLE_RESPONSES_BACKEND, aws_region_name=MANTLE_REGION + ), + mode="responses", + ) + ) + ) + + error = response.result.error if response.result else None + assert response.status == "success", f"test_connection reported an error: {error}" diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 5e2cb90958e..e6bde9770a6 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -820,6 +820,26 @@ class ModelDeleteBody(BaseModel): id: str +class ConnectionTestBody(BaseModel): + """POST /health/test_connection body, the API behind the Admin UI's Test + Connection button: the deployment params as typed into the add-model form and + the health-check mode picking which endpoint the probe calls. The endpoint + rejects `os.environ/` references, so credentials are either literal values or + omitted to fall through to the proxy's own environment.""" + + litellm_params: LiteLLMParamsBody + mode: Literal["chat", "completion", "embedding", "responses"] + + +class ConnectionTestResult(BaseModel): + error: str | None = None + + +class ConnectionTestResponse(BaseModel): + status: Literal["success", "error"] + result: ConnectionTestResult | None = None + + class CredentialCreateBody(BaseModel): credential_name: str credential_values: dict[str, str] From db8c49305ead5151c2c9b7d9bdfb5cf388b1c140 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:24:28 -0700 Subject: [PATCH 046/281] test: type the cost map fixture instead of bare dict --- .../test_together_ai_model_metadata.py | 38 ++++++++++++------- 1 file changed, 25 insertions(+), 13 deletions(-) diff --git a/tests/test_litellm/test_together_ai_model_metadata.py b/tests/test_litellm/test_together_ai_model_metadata.py index 7d7712f3c56..5a0aadf4737 100644 --- a/tests/test_litellm/test_together_ai_model_metadata.py +++ b/tests/test_litellm/test_together_ai_model_metadata.py @@ -3,11 +3,15 @@ from pathlib import Path from typing import Final import pytest +from pydantic import TypeAdapter from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider REPO_ROOT: Final = Path(__file__).parents[2] +CostMap = dict[str, dict[str, object]] +COST_MAP_ADAPTER: Final = TypeAdapter(CostMap) + SERVERLESS_CHAT_MODELS: Final = ( "together_ai/moonshotai/Kimi-K3", "together_ai/zai-org/GLM-5.2", @@ -66,13 +70,13 @@ DEPRECATED_MODELS: Final = { @pytest.fixture(scope="module") -def cost_map() -> dict: +def cost_map() -> CostMap: with open(REPO_ROOT / "model_prices_and_context_window.json") as f: - return json.load(f) + return COST_MAP_ADAPTER.validate_python(json.load(f)) @pytest.mark.parametrize("model", SERVERLESS_CHAT_MODELS) -def test_together_serverless_chat_model_is_mapped(cost_map: dict, model: str): +def test_together_serverless_chat_model_is_mapped(cost_map: CostMap, model: str): info = cost_map.get(model) assert info is not None, f"{model} missing from model_prices_and_context_window.json" assert info["litellm_provider"] == "together_ai" @@ -86,7 +90,7 @@ def test_together_serverless_chat_model_is_mapped(cost_map: dict, model: str): assert provider == "together_ai" -def test_together_kimi_k3_pricing_and_capabilities(cost_map: dict): +def test_together_kimi_k3_pricing_and_capabilities(cost_map: CostMap): info = cost_map["together_ai/moonshotai/Kimi-K3"] assert info["input_cost_per_token"] == 3e-06 assert info["output_cost_per_token"] == 1.5e-05 @@ -98,7 +102,7 @@ def test_together_kimi_k3_pricing_and_capabilities(cost_map: dict): assert info["supports_reasoning"] is True -def test_together_glm_52_pricing(cost_map: dict): +def test_together_glm_52_pricing(cost_map: CostMap): info = cost_map["together_ai/zai-org/GLM-5.2"] assert info["input_cost_per_token"] == 1.4e-06 assert info["output_cost_per_token"] == 4.4e-06 @@ -106,7 +110,7 @@ def test_together_glm_52_pricing(cost_map: dict): assert info["supports_reasoning"] is True -def test_together_multilingual_e5_embedding_entry(cost_map: dict): +def test_together_multilingual_e5_embedding_entry(cost_map: CostMap): info = cost_map["together_ai/intfloat/multilingual-e5-large-instruct"] assert info["mode"] == "embedding" assert info["input_cost_per_token"] == 2e-08 @@ -114,7 +118,7 @@ def test_together_multilingual_e5_embedding_entry(cost_map: dict): assert info["output_vector_size"] == 1024 -def test_together_llama_33_70b_repriced_to_current_together_rate(cost_map: dict): +def test_together_llama_33_70b_repriced_to_current_together_rate(cost_map: CostMap): info = cost_map["together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo"] assert info["input_cost_per_token"] == 1.04e-06 assert info["output_cost_per_token"] == 1.04e-06 @@ -122,17 +126,25 @@ def test_together_llama_33_70b_repriced_to_current_together_rate(cost_map: dict) @pytest.mark.parametrize("model", sorted(DEPRECATED_MODELS)) -def test_together_deprecated_model_carries_deprecation_date(cost_map: dict, model: str): +def test_together_deprecated_model_carries_deprecation_date(cost_map: CostMap, model: str): info = cost_map.get(model) assert info is not None, f"{model} missing from model_prices_and_context_window.json" assert info.get("deprecation_date") == DEPRECATED_MODELS[model] -def test_together_successor_metadata_points_at_live_models(cost_map: dict): +def _successor(info: dict[str, object]) -> str | None: + metadata = info.get("metadata") + if not isinstance(metadata, dict): + return None + successor = metadata.get("successor") + return successor if isinstance(successor, str) else None + + +def test_together_successor_metadata_points_at_live_models(cost_map: CostMap): successors = { - model: info["metadata"]["successor"] + model: successor for model, info in cost_map.items() - if model.startswith("together_ai/") and "successor" in info.get("metadata", {}) + if model.startswith("together_ai/") and (successor := _successor(info)) is not None } assert len(successors) >= 10 for model, successor in successors.items(): @@ -141,9 +153,9 @@ def test_together_successor_metadata_points_at_live_models(cost_map: dict): assert "deprecation_date" not in target, f"{model} names deprecated successor {successor}" -def test_together_backup_cost_map_in_sync(cost_map: dict): +def test_together_backup_cost_map_in_sync(cost_map: CostMap): with open(REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json") as f: - backup = json.load(f) + backup = COST_MAP_ADAPTER.validate_python(json.load(f)) together_main = {k: v for k, v in cost_map.items() if k.startswith("together_ai/")} together_backup = {k: v for k, v in backup.items() if k.startswith("together_ai/")} assert together_backup == together_main From 5470645f87d1f9a6367121cfa99b8672e57ca350 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 17:32:21 +0000 Subject: [PATCH 047/281] refactor(azure/realtime): keep auth header build within lint budgets after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/azure/realtime/handler.py | 10 +++--- litellm/realtime_api/main.py | 33 ++++++++++++------- .../realtime/test_azure_realtime_handler.py | 4 ++- 3 files changed, 30 insertions(+), 17 deletions(-) diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index 70e39f47d63..88492ef996e 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -4,6 +4,8 @@ This file contains the calling Azure OpenAI's `/openai/realtime` endpoint. This requires websockets, and is currently only supported on LiteLLM Proxy. """ +from collections.abc import Mapping +from types import MappingProxyType from typing import Any, Final, cast from litellm._logging import _redact_string, verbose_proxy_logger @@ -31,15 +33,15 @@ async def forward_messages(client_ws: Any, backend_ws: Any): class AzureOpenAIRealtime(AzureChatCompletion): @staticmethod - def get_auth_headers(api_key: str | None, azure_ad_token: str | None) -> dict[str, str]: + def get_auth_headers(api_key: str | None, azure_ad_token: str | None) -> Mapping[str, str]: """ Build the websocket handshake auth headers, preferring a static api-key and falling back to an Azure AD (Entra ID) bearer token. Never sends both. """ if api_key: - return {"api-key": api_key} + return MappingProxyType({"api-key": api_key}) if azure_ad_token: - return {"Authorization": f"Bearer {azure_ad_token}"} + return MappingProxyType({"Authorization": f"Bearer {azure_ad_token}"}) raise ValueError( "Missing Azure credentials for the realtime endpoint. Set an api_key, or configure Azure AD auth " "(azure_ad_token, tenant_id/client_id/client_secret, or a managed identity)" @@ -132,7 +134,7 @@ class AzureOpenAIRealtime(AzureChatCompletion): query_params=query_params, ) - auth_headers = self.get_auth_headers(api_key=api_key, azure_ad_token=azure_ad_token) + auth_headers: Final = self.get_auth_headers(api_key=api_key, azure_ad_token=azure_ad_token) try: ssl_context: Final = get_shared_realtime_ssl_context() diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 8933d5e4506..e5f6c8328f4 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -2,6 +2,8 @@ import asyncio import os +from collections.abc import Mapping +from types import MappingProxyType from typing import Any, Final, Literal, cast import litellm @@ -45,6 +47,7 @@ bedrock_realtime: Final = BedrockRealtime() xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() base_llm_http_handler = BaseLLMHTTPHandler() +_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({}) def _with_resolved_session_model(session: dict[str, Any], model_name: str) -> dict[str, Any]: @@ -412,10 +415,8 @@ async def _arealtime( if realtime_protocol is None and (query_params or {}).get("intent") == "transcription": realtime_protocol = "GA" realtime_protocol = realtime_protocol or "beta" - resolved_azure_ad_token = ( - None - if api_key - else get_azure_ad_token(GenericLiteLLMParams(**{**kwargs, "azure_ad_token": azure_ad_token})) + resolved_azure_ad_token: Final = ( + None if api_key else get_azure_ad_token(GenericLiteLLMParams(**kwargs, azure_ad_token=azure_ad_token)) ) await azure_realtime.async_realtime( model=model, @@ -556,6 +557,17 @@ async def _arealtime( raise ValueError(f"Unsupported model: {model}") +def _realtime_health_check_auth_headers( + custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any] +) -> Mapping[str, str | None]: + if custom_llm_provider != "azure": + return MappingProxyType({"api-key": api_key}) + return azure_realtime.get_auth_headers( + api_key=api_key, + azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))), + ) + + async def _realtime_health_check( model: str, custom_llm_provider: str, @@ -584,7 +596,11 @@ async def _realtime_health_check( import websockets url: str | None = None - auth_headers: dict[str, str | None] = {"api-key": api_key} + auth_headers: Final = _realtime_health_check_auth_headers( + custom_llm_provider=custom_llm_provider, + api_key=api_key, + model_params=model_params or _EMPTY_MODEL_PARAMS, + ) if custom_llm_provider == "azure": url = azure_realtime._construct_url( api_base=api_base or "", @@ -592,13 +608,6 @@ async def _realtime_health_check( api_version=api_version or "2024-10-01-preview", realtime_protocol=realtime_protocol, ) - azure_litellm_params = GenericLiteLLMParams(**(model_params or {})) - auth_headers = dict( - azure_realtime.get_auth_headers( - api_key=api_key, - azure_ad_token=(None if api_key else get_azure_ad_token(azure_litellm_params)), - ) - ) elif custom_llm_provider == "openai": url = openai_realtime._construct_url( api_base=api_base or "https://api.openai.com/", diff --git a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py index 570843da7a8..7d24e604569 100644 --- a/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/test_litellm/llms/azure/realtime/test_azure_realtime_handler.py @@ -590,7 +590,9 @@ async def test_async_realtime_uses_bearer_token_when_no_api_key(): "websockets.connect", return_value=_DummyAsyncContextManager(mock_backend_ws), ) as mock_ws_connect, - patch("litellm.llms.azure.realtime.handler.RealTimeStreaming") as mock_realtime_streaming, + patch( # test-quality-ok: handler owns the streaming loop, only the handshake headers are under test + "litellm.llms.azure.realtime.handler.RealTimeStreaming" + ) as mock_realtime_streaming, ): mock_realtime_streaming.return_value.bidirectional_forward = AsyncMock() From 90f9a8bfda7aff99e1c969f573931cc86eee2103 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 11:05:38 -0700 Subject: [PATCH 048/281] test(e2e): retry timeout-shaped Mantle test_connection probes The endpoint answers a probe that exceeds HEALTH_CHECK_TIMEOUT_SECONDS with HTTP 200 and an in-body "Timeout exceeded", which the harness's status-code rerun policy cannot see. The suite's parallel Bedrock load can push a Mantle probe past that cap transiently, so only that exact error is retried, three bounded attempts with visible prints; any other error verdict still fails immediately. --- .../test_model_test_connection_e2e.py | 51 ++++++++++++++----- 1 file changed, 38 insertions(+), 13 deletions(-) diff --git a/tests/e2e/management/test_model_test_connection_e2e.py b/tests/e2e/management/test_model_test_connection_e2e.py index a1f714df4c8..25b0b4f24e6 100644 --- a/tests/e2e/management/test_model_test_connection_e2e.py +++ b/tests/e2e/management/test_model_test_connection_e2e.py @@ -8,35 +8,60 @@ verdict from the live provider rather than just a 200 envelope. The region is a literal because the endpoint rejects request-supplied os.environ/ references; credentials fall through to the proxy's own environment (bearer token locally, pod identity in CI). + +The endpoint caps every probe at HEALTH_CHECK_TIMEOUT_SECONDS and answers a +timed-out probe with HTTP 200 and an in-body "Timeout exceeded", which the +harness's status-code retry policy cannot see. A Mantle probe can hit that cap +transiently while the rest of the suite saturates the same AWS account, so only +that exact error is retried here; any other error verdict fails immediately. """ from __future__ import annotations +import time + import pytest from e2e_http import unwrap from management_client import ManagementClient -from models import ConnectionTestBody, LiteLLMParamsBody +from models import ConnectionTestBody, ConnectionTestResponse, LiteLLMParamsBody pytestmark = pytest.mark.e2e MANTLE_RESPONSES_BACKEND = "bedrock_mantle/openai.gpt-5.6-luna" MANTLE_REGION = "us-east-1" +PROBE_TIMEOUT_ERROR = "Timeout exceeded" +PROBE_ATTEMPTS = 3 +PROBE_RETRY_SLEEP_SECONDS = 30 + + +def _probe_mantle(client: ManagementClient) -> ConnectionTestResponse: + return unwrap( + client.connection_test( + ConnectionTestBody( + litellm_params=LiteLLMParamsBody( + model=MANTLE_RESPONSES_BACKEND, aws_region_name=MANTLE_REGION + ), + mode="responses", + ) + ) + ) class TestModelTestConnection: @pytest.mark.covers("mgmt.model.test_connection.happy_path") def test_bedrock_mantle_responses_connection_succeeds(self, client: ManagementClient) -> None: - response = unwrap( - client.connection_test( - ConnectionTestBody( - litellm_params=LiteLLMParamsBody( - model=MANTLE_RESPONSES_BACKEND, aws_region_name=MANTLE_REGION - ), - mode="responses", + for attempt in range(1, PROBE_ATTEMPTS + 1): + response = _probe_mantle(client) + if response.status == "success": + return + error = response.result.error if response.result else None + assert error == PROBE_TIMEOUT_ERROR, f"test_connection reported an error: {error}" + if attempt < PROBE_ATTEMPTS: + print( + f"test_connection probe timed out; retry {attempt}/{PROBE_ATTEMPTS - 1}" + f" in {PROBE_RETRY_SLEEP_SECONDS}s", + flush=True, ) - ) - ) - - error = response.result.error if response.result else None - assert response.status == "success", f"test_connection reported an error: {error}" + time.sleep(PROBE_RETRY_SLEEP_SECONDS) + pytest.fail(f"test_connection timed out on all {PROBE_ATTEMPTS} attempts") From e14f485827ab78808f3a0e6bafda3ddfd1da2474 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 18:11:42 +0000 Subject: [PATCH 049/281] fix(anthropic): raise missing-credential error on /v1/messages passthrough Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../messages/transformation.py | 14 +++++- .../anthropic/test_anthropic_common_utils.py | 44 +++++++++++++++++++ 2 files changed, 56 insertions(+), 2 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 032bf0130ce..75146922a39 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -8,6 +8,7 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, ) +from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import verbose_logger from litellm.llms.base_llm.anthropic_messages.transformation import ( @@ -309,8 +310,17 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): if "x-api-key" not in headers and "authorization" not in headers: auth_header: Final = AnthropicModelInfo.get_auth_header(api_key) - if auth_header is not None: - headers.update(auth_header) + if auth_header is None: + raise AuthenticationError( + message=( + "Missing Anthropic API Key - A call is being made to anthropic but no key is set " + "either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` " + "or `ANTHROPIC_AUTH_TOKEN` in your environment vars" + ), + llm_provider=self._resolved_provider, + model=model, + ) + headers.update(auth_header) if "anthropic-version" not in headers: headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION if "content-type" not in headers: diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index c27362bf49f..a96c44ac5c0 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -1227,6 +1227,50 @@ class TestPassthroughAuthToken: assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY assert "authorization" not in updated_headers + def test_passthrough_missing_credentials_raises_authentication_error(self): + """Passthrough endpoint should raise locally instead of forwarding an unauthenticated request.""" + from unittest.mock import patch as mock_patch + + import litellm + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + config = AnthropicMessagesConfig() + with mock_patch.dict("os.environ", {}, clear=True): + with pytest.raises(litellm.AuthenticationError, match="Missing Anthropic API Key"): + config.validate_anthropic_messages_environment( + headers={}, + model="claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + def test_passthrough_client_x_api_key_header_is_kept(self): + """A client-forwarded x-api-key header should satisfy validation without env credentials.""" + from unittest.mock import patch as mock_patch + + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + config = AnthropicMessagesConfig() + with mock_patch.dict("os.environ", {}, clear=True): + updated_headers, _ = config.validate_anthropic_messages_environment( + headers={"x-api-key": FAKE_REGULAR_KEY}, + model="claude-sonnet-4-5-20250929", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY + def test_passthrough_get_complete_url_honours_base_url_env(self): """get_complete_url should use ANTHROPIC_BASE_URL when api_base is None.""" from unittest.mock import patch as mock_patch From 0fc042e37613b8e9b72c354c18d99e4b5fc2dddb Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 25 Aug 2026 18:28:52 +0000 Subject: [PATCH 050/281] test(anthropic): pass explicit api_key where passthrough env validation now raises Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llms/anthropic/chat/test_anthropic_chat_transformation.py | 1 + .../messages/test_anthropic_messages_speed.py | 2 ++ 2 files changed, 3 insertions(+) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 4f340ee0f3f..0933638635b 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -1050,6 +1050,7 @@ def test_anthropic_messages_validate_adds_beta_header(): messages=[{"role": "user", "content": [{"type": "text", "text": "Hi"}]}], optional_params={"context_management": _sample_context_management_payload()}, litellm_params={}, + api_key="fake-anthropic-key", ) assert headers["anthropic-beta"] == "context-management-2025-06-27" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py index 6900f1062bf..efd49962ac8 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_messages_speed.py @@ -28,6 +28,7 @@ def test_messages_drop_params_strips_speed_for_unsupported_models(): messages=[{"role": "user", "content": "Hello"}], optional_params=dict(optional_params), litellm_params={}, + api_key="fake-anthropic-key", ) result = config.transform_anthropic_messages_request( model="claude-sonnet-4-6", @@ -60,6 +61,7 @@ def test_messages_drop_params_keeps_speed_for_supporting_models(): messages=[{"role": "user", "content": "Hello"}], optional_params=dict(optional_params), litellm_params={}, + api_key="fake-anthropic-key", ) result = config.transform_anthropic_messages_request( model="claude-opus-4-6", From 62ec3b61167d780b93839a7b03be30f2305f7e86 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 11:44:06 -0700 Subject: [PATCH 051/281] fix(together_ai): route chat completions through a dedicated TogetherAIChatConfig --- litellm/__init__.py | 3 + litellm/_lazy_imports_registry.py | 5 + .../get_supported_openai_params.py | 2 +- litellm/llms/together_ai/chat.py | 58 ----- litellm/llms/together_ai/chat/__init__.py | 3 + .../llms/together_ai/chat/transformation.py | 49 ++++ litellm/main.py | 62 ++++- litellm/utils.py | 4 +- .../test_together_ai_chat_transformation.py | 232 ++++++++++++++++++ 9 files changed, 348 insertions(+), 70 deletions(-) delete mode 100644 litellm/llms/together_ai/chat.py create mode 100644 litellm/llms/together_ai/chat/__init__.py create mode 100644 litellm/llms/together_ai/chat/transformation.py create mode 100644 tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index ee2c551481c..39556d1f04d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1628,6 +1628,9 @@ if TYPE_CHECKING: AmazonMantleMessagesConfig as AmazonMantleMessagesConfig, ) from .llms.together_ai.chat import TogetherAIConfig as TogetherAIConfig + from .llms.together_ai.chat.transformation import ( + TogetherAIChatConfig as TogetherAIChatConfig, + ) from .llms.nlp_cloud.chat.handler import NLPCloudConfig as NLPCloudConfig from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig as VertexGeminiConfig, diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index c34c9eefe85..1c833256598 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -177,6 +177,7 @@ LLM_CONFIG_NAMES: Final = ( "AmazonAnthropicClaudeMessagesConfig", "AmazonMantleMessagesConfig", "TogetherAIConfig", + "TogetherAIChatConfig", "NLPCloudConfig", "VertexGeminiConfig", "GoogleAIStudioGeminiConfig", @@ -741,6 +742,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { "AmazonMantleMessagesConfig", ), "TogetherAIConfig": (".llms.together_ai.chat", "TogetherAIConfig"), + "TogetherAIChatConfig": ( + ".llms.together_ai.chat.transformation", + "TogetherAIChatConfig", + ), "NLPCloudConfig": (".llms.nlp_cloud.chat.handler", "NLPCloudConfig"), "VertexGeminiConfig": ( ".llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 72f36661f4c..7a16ffe4d85 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -172,7 +172,7 @@ def get_supported_openai_params( if request_type == "embeddings": return litellm.JinaAIEmbeddingConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "together_ai": - return litellm.TogetherAIConfig().get_supported_openai_params(model=model) + return litellm.TogetherAIChatConfig().get_supported_openai_params(model=model) elif custom_llm_provider == "databricks": if request_type == "chat_completion": return litellm.DatabricksConfig().get_supported_openai_params(model=model) diff --git a/litellm/llms/together_ai/chat.py b/litellm/llms/together_ai/chat.py deleted file mode 100644 index 58d47e45faa..00000000000 --- a/litellm/llms/together_ai/chat.py +++ /dev/null @@ -1,58 +0,0 @@ -""" -Support for OpenAI's `/v1/chat/completions` endpoint. - -Calls done in OpenAI/openai.py as TogetherAI is openai-compatible. - -Docs: https://docs.together.ai/reference/completions-1 -""" - -from typing import Final - -from litellm._logging import verbose_logger -from litellm.utils import supports_function_calling - -from ..openai.chat.gpt_transformation import OpenAIGPTConfig - - -class TogetherAIConfig(OpenAIGPTConfig): - def get_supported_openai_params(self, model: str) -> list: - """ - Only some together models support response_format / tool calling - - Docs: https://docs.together.ai/docs/json-mode - """ - # Use supports_function_calling() — which reads _get_model_info_helper - # directly — instead of get_model_info(). get_model_info() calls - # get_supported_openai_params() as its first step, which routes back - # into this method for together_ai models, creating a recursion that - # only terminates when Python's recursion limit or the "not mapped" - # exception in _get_model_info_helper is hit (~332 deep calls). - supports_fc: bool | None = None - try: - supports_fc = supports_function_calling(model, custom_llm_provider="together_ai") - except Exception as e: - verbose_logger.debug("Error getting supported openai params: %s", e) - - optional_params: Final = super().get_supported_openai_params(model) - if supports_fc is not True: - verbose_logger.debug( - "Only some together models support function calling/response_format. Docs - https://docs.together.ai/docs/function-calling" - ) - optional_params.remove("tools") - optional_params.remove("tool_choice") - optional_params.remove("function_call") - optional_params.remove("response_format") - return optional_params - - def map_openai_params( - self, - non_default_params: dict, - optional_params: dict, - model: str, - drop_params: bool, - ) -> dict: - mapped_openai_params: Final = super().map_openai_params(non_default_params, optional_params, model, drop_params) - - if "response_format" in mapped_openai_params and mapped_openai_params["response_format"] == {"type": "text"}: - mapped_openai_params.pop("response_format") - return mapped_openai_params diff --git a/litellm/llms/together_ai/chat/__init__.py b/litellm/llms/together_ai/chat/__init__.py new file mode 100644 index 00000000000..f260d9126d7 --- /dev/null +++ b/litellm/llms/together_ai/chat/__init__.py @@ -0,0 +1,3 @@ +from .transformation import TogetherAIChatConfig as TogetherAIChatConfig + +TogetherAIConfig = TogetherAIChatConfig diff --git a/litellm/llms/together_ai/chat/transformation.py b/litellm/llms/together_ai/chat/transformation.py new file mode 100644 index 00000000000..eb0954bceef --- /dev/null +++ b/litellm/llms/together_ai/chat/transformation.py @@ -0,0 +1,49 @@ +""" +Translates from OpenAI's `/v1/chat/completions` to Together AI's `/v1/chat/completions`. + +Docs: https://docs.together.ai/docs/chat-overview +""" + +from types import MappingProxyType +from typing import Final + +from litellm._logging import verbose_logger +from litellm.utils import supports_function_calling + +from ...openai.chat.gpt_transformation import OpenAIGPTConfig + +FUNCTION_CALLING_ONLY_PARAMS: Final = ("tools", "tool_choice", "function_call", "response_format") +PLAIN_TEXT_RESPONSE_FORMAT: Final = MappingProxyType({"type": "text"}) + + +class TogetherAIChatConfig(OpenAIGPTConfig): + def get_supported_openai_params(self, model: str) -> list: + supports_fc: bool | None = None + try: + supports_fc = supports_function_calling(model, custom_llm_provider="together_ai") + except Exception as e: + verbose_logger.debug("Error getting supported openai params: %s", e) + + supported_params: Final = super().get_supported_openai_params(model) + if supports_fc is True: + return supported_params + verbose_logger.debug( + "Only some together models support function calling/response_format. Docs - https://docs.together.ai/docs/function-calling" + ) + for param in FUNCTION_CALLING_ONLY_PARAMS: + if param in supported_params: + supported_params.remove(param) + return supported_params + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + mapped_openai_params: Final = super().map_openai_params(non_default_params, optional_params, model, drop_params) + + if mapped_openai_params.get("response_format") == PLAIN_TEXT_RESPONSE_FORMAT: + mapped_openai_params.pop("response_format") + return mapped_openai_params diff --git a/litellm/main.py b/litellm/main.py index d3967473f99..8ddcef2b9fa 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -24,6 +24,7 @@ from concurrent import futures from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from copy import deepcopy from functools import partial +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, Union, cast, get_args from litellm._logging import _redact_string @@ -1811,6 +1812,56 @@ def _complete_fireworks_ai( return response +def _complete_together_ai(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: + acompletion: Final = ctx.acompletion + api_base: Final = ctx.api_base + api_key: Final = ctx.api_key + client: Final = _dispatch_client_http(ctx) + custom_llm_provider: Final = ctx.custom_llm_provider + headers: Final = ctx.headers + litellm_params: Final = ctx.litellm_params + logging: Final = ctx.logging + messages: Final = ctx.messages + model: Final = ctx.model + model_response: Final = ctx.model_response + optional_params: Final = ctx.optional_params + provider_config: Final = ctx.provider_config + shared_session: Final = ctx.shared_session + stream: Final = ctx.stream + timeout: Final = ctx.timeout + + try: + response: Final = base_llm_http_handler.completion( + model=model, + messages=messages, + headers=headers, + model_response=model_response, + api_key=api_key, + api_base=api_base, + acompletion=acompletion, + logging_obj=logging, + optional_params=optional_params, + litellm_params=litellm_params, + shared_session=shared_session, + timeout=timeout, + client=client, + custom_llm_provider=custom_llm_provider, + encoding=_get_encoding(), + stream=stream, + provider_config=provider_config, + ) + except Exception as e: + logging.post_call( + input=messages, + api_key=api_key, + original_response=str(e), + additional_args=MappingProxyType({"headers": headers}), + ) + raise + + return response + + def _complete_heroku(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: acompletion: Final = ctx.acompletion api_base: Final = ctx.api_base @@ -5600,6 +5651,8 @@ def completion( elif custom_llm_provider == "fireworks_ai": ## COMPLETION CALL response = _complete_fireworks_ai(_dispatch_ctx) + elif custom_llm_provider == "together_ai": + response = _complete_together_ai(_dispatch_ctx) elif custom_llm_provider == "heroku": response = _complete_heroku(_dispatch_ctx) @@ -5649,7 +5702,6 @@ def completion( or custom_llm_provider == "volcengine" or custom_llm_provider == "anyscale" or custom_llm_provider == "openai" - or custom_llm_provider == "together_ai" or custom_llm_provider == "nebius" or custom_llm_provider == "wandb" or custom_llm_provider == "clarifai" @@ -5699,14 +5751,6 @@ def completion( response = _complete_openrouter(_dispatch_ctx) elif custom_llm_provider == "vercel_ai_gateway": response = _complete_vercel_ai_gateway(_dispatch_ctx) - elif ( - custom_llm_provider == "together_ai" - or ("togethercomputer" in model) - or (model in litellm.together_ai_models) - ): - """ - Deprecated. We now do together ai calls via the openai client - https://docs.together.ai/docs/openai-api-compatibility - """ elif custom_llm_provider == "palm": raise ValueError( "Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en" diff --git a/litellm/utils.py b/litellm/utils.py index 012e8785321..43f146d52d5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4130,7 +4130,7 @@ def get_optional_params( drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False), ) elif custom_llm_provider == "together_ai": - optional_params = litellm.TogetherAIConfig().map_openai_params( + optional_params = litellm.TogetherAIChatConfig().map_openai_params( non_default_params=non_default_params, optional_params=optional_params, model=model, @@ -7898,7 +7898,7 @@ class ProviderConfigManager: LlmProviders.GALADRIEL: (lambda: litellm.GaladrielChatConfig(), False), LlmProviders.REPLICATE: (lambda: litellm.ReplicateConfig(), False), LlmProviders.HUGGINGFACE: (lambda: litellm.HuggingFaceChatConfig(), False), - LlmProviders.TOGETHER_AI: (lambda: litellm.TogetherAIConfig(), False), + LlmProviders.TOGETHER_AI: (lambda: litellm.TogetherAIChatConfig(), False), LlmProviders.OPENROUTER: (lambda: litellm.OpenrouterConfig(), False), LlmProviders.VERCEL_AI_GATEWAY: ( lambda: litellm.VercelAIGatewayConfig(), diff --git a/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py b/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py new file mode 100644 index 00000000000..6216d3bf225 --- /dev/null +++ b/tests/test_litellm/llms/together_ai/chat/test_together_ai_chat_transformation.py @@ -0,0 +1,232 @@ +import json +from unittest.mock import MagicMock + +import httpx +import pytest + +import litellm +from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj +from litellm.llms.openai.chat.gpt_transformation import ( + OpenAIChatCompletionStreamingHandler, +) +from litellm.llms.together_ai.chat.transformation import TogetherAIChatConfig +from litellm.types.utils import LlmProviders, ModelResponse + +TOOL_CALLING_MODEL = "openai/gpt-oss-20b" +REASONING_MODEL = "deepseek-ai/DeepSeek-V3.1" +PLAIN_MODEL = "Qwen/Qwen3-235B-A22B-fp8-tput" +UNMAPPED_MODEL = "MiniMaxAI/MiniMax-M3" + +FUNCTION_CALLING_PARAMS = ("tools", "tool_choice", "function_call", "response_format") + + +@pytest.fixture(autouse=True) +def force_local_model_cost(monkeypatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map + + monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url)) + + +def test_supported_params_tool_calling_model(): + supported = TogetherAIChatConfig().get_supported_openai_params(model=TOOL_CALLING_MODEL) + + for param in FUNCTION_CALLING_PARAMS: + assert param in supported + + +def test_supported_params_plain_model(): + supported = TogetherAIChatConfig().get_supported_openai_params(model=PLAIN_MODEL) + + for param in FUNCTION_CALLING_PARAMS: + assert param not in supported + assert "temperature" in supported + assert "max_tokens" in supported + + +def test_supported_params_unmapped_model_treated_as_plain(): + supported = TogetherAIChatConfig().get_supported_openai_params(model=UNMAPPED_MODEL) + + for param in FUNCTION_CALLING_PARAMS: + assert param not in supported + assert "stream" in supported + + +def test_map_openai_params_tool_calling_model_passes_tools(): + tools = [{"type": "function", "function": {"name": "get_weather", "parameters": {}}}] + + mapped = TogetherAIChatConfig().map_openai_params( + non_default_params={"tools": tools, "tool_choice": "auto"}, + optional_params={}, + model=TOOL_CALLING_MODEL, + drop_params=False, + ) + + assert mapped["tools"] == tools + assert mapped["tool_choice"] == "auto" + + +def test_map_openai_params_reasoning_model_passes_sampling_params(): + mapped = TogetherAIChatConfig().map_openai_params( + non_default_params={"temperature": 0.2, "max_tokens": 512}, + optional_params={}, + model=REASONING_MODEL, + drop_params=False, + ) + + assert mapped["temperature"] == 0.2 + assert mapped["max_tokens"] == 512 + + +def test_map_openai_params_drops_text_response_format(): + mapped = TogetherAIChatConfig().map_openai_params( + non_default_params={"response_format": {"type": "text"}, "temperature": 0.5}, + optional_params={}, + model=REASONING_MODEL, + drop_params=False, + ) + + assert "response_format" not in mapped + assert mapped["temperature"] == 0.5 + + +def test_map_openai_params_keeps_json_response_format(): + response_format = {"type": "json_object"} + + mapped = TogetherAIChatConfig().map_openai_params( + non_default_params={"response_format": response_format}, + optional_params={}, + model=TOOL_CALLING_MODEL, + drop_params=False, + ) + + assert mapped["response_format"] == response_format + + +def _transform_response(message: dict) -> ModelResponse: + raw_response_json = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1234567890, + "model": REASONING_MODEL, + "choices": [{"index": 0, "message": message, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + } + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = raw_response_json + mock_response.text = json.dumps(raw_response_json) + mock_response.headers = {} + logging_obj = MagicMock(spec=LiteLLMLoggingObj) + logging_obj.post_call = MagicMock() + logging_obj.model_call_details = {} + + return TogetherAIChatConfig().transform_response( + model=REASONING_MODEL, + raw_response=mock_response, + model_response=ModelResponse(), + logging_obj=logging_obj, + request_data={}, + messages=[{"role": "user", "content": "What is 2+2?"}], + optional_params={}, + litellm_params={}, + encoding=None, + api_key="test-key", + json_mode=False, + ) + + +def test_transform_response_maps_reasoning_to_reasoning_content(): + result = _transform_response( + {"role": "assistant", "content": "4", "reasoning": "2+2 equals 4"} + ) + + assert result.choices[0].message.content == "4" + assert result.choices[0].message.reasoning_content == "2+2 equals 4" + + +def test_transform_response_preserves_reasoning_content_field(): + result = _transform_response( + {"role": "assistant", "content": "4", "reasoning_content": "adding 2 and 2"} + ) + + assert result.choices[0].message.reasoning_content == "adding 2 and 2" + + +def test_streaming_chunk_maps_delta_reasoning_to_reasoning_content(): + iterator = TogetherAIChatConfig().get_model_response_iterator( + streaming_response=iter(()), sync_stream=True + ) + assert isinstance(iterator, OpenAIChatCompletionStreamingHandler) + + parsed = iterator.chunk_parser( + { + "id": "chunk-1", + "created": 1234567890, + "model": REASONING_MODEL, + "choices": [{"index": 0, "delta": {"reasoning": "thinking about 2+2"}}], + } + ) + + assert parsed.choices[0]["delta"]["reasoning_content"] == "thinking about 2+2" + + +def test_together_ai_config_alias_points_at_chat_config(): + assert litellm.TogetherAIConfig is litellm.TogetherAIChatConfig + config = litellm.TogetherAIConfig(max_tokens=10) + assert isinstance(config, TogetherAIChatConfig) + + +def test_provider_config_manager_returns_together_chat_config(): + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_chat_config( + model=REASONING_MODEL, provider=LlmProviders.TOGETHER_AI + ) + + assert isinstance(config, TogetherAIChatConfig) + + +def test_completion_routes_through_together_chat_config(): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + captured_requests = [] + + def respond(request: httpx.Request) -> httpx.Response: + captured_requests.append(request) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-together", + "object": "chat.completion", + "created": 1234567890, + "model": REASONING_MODEL, + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "4", + "reasoning": "2+2 equals 4", + }, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, + }, + ) + + client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))) + + response = litellm.completion( + model=f"together_ai/{REASONING_MODEL}", + messages=[{"role": "user", "content": "What is 2+2?"}], + api_key="fake-key", + client=client, + ) + + request = captured_requests[0] + assert str(request.url) == "https://api.together.ai/v1/chat/completions" + assert request.headers["authorization"] == "Bearer fake-key" + assert json.loads(request.content)["model"] == REASONING_MODEL + assert response.choices[0].message.content == "4" + assert response.choices[0].message.reasoning_content == "2+2 equals 4" From 32ebfba5ed7810ead375c613ee2419e167ba831c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:03:51 -0700 Subject: [PATCH 052/281] refactor(together_ai): build the trimmed supported-params list without mutating the inherited list --- litellm/llms/together_ai/chat/transformation.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/litellm/llms/together_ai/chat/transformation.py b/litellm/llms/together_ai/chat/transformation.py index eb0954bceef..88fd79f2366 100644 --- a/litellm/llms/together_ai/chat/transformation.py +++ b/litellm/llms/together_ai/chat/transformation.py @@ -30,10 +30,9 @@ class TogetherAIChatConfig(OpenAIGPTConfig): verbose_logger.debug( "Only some together models support function calling/response_format. Docs - https://docs.together.ai/docs/function-calling" ) - for param in FUNCTION_CALLING_ONLY_PARAMS: - if param in supported_params: - supported_params.remove(param) - return supported_params + return [ # mutable-ok: the inherited contract returns a plain list; building fresh avoids mutating the base class's value + param for param in supported_params if param not in FUNCTION_CALLING_ONLY_PARAMS + ] def map_openai_params( self, From 17845b4fb01b8ec3d0c90254bece1b802807f77c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:20:51 -0700 Subject: [PATCH 053/281] fix(anthropic): translate tool_result document blocks in the /v1/messages bridge --- .../adapters/transformation.py | 6 +- ...al_pass_through_adapters_transformation.py | 71 +++++++++++++++++++ 2 files changed, 74 insertions(+), 3 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 7c89da81fe6..109017bda27 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -434,7 +434,7 @@ class LiteLLMAnthropicMessagesAdapter: content_items = list(content.get("content", [])) # Single-item text keeps the backward-compatible string format; a single - # image becomes a structured image_url part + # image or document becomes a structured image_url part if len(content_items) == 1: c = content_items[0] if isinstance(c, str): @@ -454,7 +454,7 @@ class LiteLLMAnthropicMessagesAdapter: ) self._add_cache_control_if_applicable(content, tool_result, model) tool_message_list.append(tool_result) - elif c.get("type") == "image": + elif c.get("type") in ("image", "document"): image_part = self._tool_result_image_part(c.get("source")) tool_result = ChatCompletionToolMessage( role="tool", @@ -482,7 +482,7 @@ class LiteLLMAnthropicMessagesAdapter: text=c.get("text", ""), ) ) - elif c.get("type") == "image": + elif c.get("type") in ("image", "document"): image_part = self._tool_result_image_part(c.get("source")) if image_part: combined_content_parts.append(image_part) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index d0169963962..a7fbd069e61 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1,3 +1,4 @@ +import base64 from typing import Any, cast import pytest @@ -11,6 +12,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( ) from litellm.litellm_core_utils.prompt_templates.factory import ( THOUGHT_SIGNATURE_SEPARATOR, + _bedrock_converse_messages_pt, ) from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( OPENAI_MAX_TOOL_NAME_LENGTH, @@ -3872,6 +3874,75 @@ def test_tool_result_plain_text_unchanged_by_openai_transform(): assert _image_urls_in_user_messages(result) == [] +TOOL_RESULT_PDF_B64 = base64.b64encode(b"%PDF-1.4 minimal regression fixture").decode() + + +def _base64_pdf_block(): + return { + "type": "document", + "source": {"type": "base64", "media_type": "application/pdf", "data": TOOL_RESULT_PDF_B64}, + } + + +def test_tool_result_single_document_kept_as_pdf_data_url(): + adapter = LiteLLMAnthropicMessagesAdapter() + translated = adapter.translate_anthropic_messages_to_openai( + messages=[ + _anthropic_tool_use_turn("toolu_01"), + _anthropic_tool_result_turn({"toolu_01": [_base64_pdf_block()]}), + ] + ) + + tool_messages = [m for m in translated if m.get("role") == "tool"] + assert len(tool_messages) == 1 + assert tool_messages[0]["content"] == [ + { + "type": "image_url", + "image_url": {"url": f"data:application/pdf;base64,{TOOL_RESULT_PDF_B64}"}, + } + ] + + +def test_tool_result_text_and_document_reach_bedrock_converse_tool_result(): + """Claude Code >= 2.1.245 sends Read-tool PDF output as a document block inside + tool_result; dropping it left bedrock converse models blind to the PDF content.""" + adapter = LiteLLMAnthropicMessagesAdapter() + translated = adapter.translate_anthropic_messages_to_openai( + messages=[ + AnthropicMessagesUserMessageParam(role="user", content="Read pong.pdf"), + _anthropic_tool_use_turn("toolu_01"), + _anthropic_tool_result_turn( + { + "toolu_01": [ + {"type": "text", "text": "PDF file read: pong.pdf (579 bytes)"}, + _base64_pdf_block(), + ] + } + ), + ] + ) + + converse_messages = _bedrock_converse_messages_pt( + messages=translated, + model="anthropic.claude-haiku-4-5-20251001-v1:0", + llm_provider="bedrock_converse", + ) + + tool_results = [ + block["toolResult"] + for message in converse_messages + for block in message["content"] + if "toolResult" in block + ] + assert len(tool_results) == 1 + documents = [part["document"] for part in tool_results[0]["content"] if "document" in part] + assert len(documents) == 1 + assert documents[0]["format"] == "pdf" + assert documents[0]["source"]["bytes"] == TOOL_RESULT_PDF_B64 + texts = [part["text"] for part in tool_results[0]["content"] if "text" in part] + assert texts == ["PDF file read: pong.pdf (579 bytes)"] + + def test_translate_anthropic_to_openai_carries_prompt_cache_breakpoint_on_system_and_user_blocks(): explicit = {"mode": "explicit"} openai_request, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( From 2bd2c1393cb231f6572b4e99f5cff427bcb66562 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:31:43 -0700 Subject: [PATCH 054/281] docs(pr-template): split Caveats bullets into severity tiers and call for plain engineering language --- .github/pull_request_template.md | 16 ++++++++++++++-- CLAUDE.md | 1 + 2 files changed, 15 insertions(+), 2 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index 4e428d8cebf..bcdb228746a 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -1,7 +1,10 @@ + + ## TLDR - + Problem this solves: @@ -112,6 +115,15 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac ## QA runbook diff --git a/CLAUDE.md b/CLAUDE.md index b3383b4a895..03053b8392c 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -44,6 +44,7 @@ If you ever make public-facing PR descriptions, comments, issues, commit message - don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose - don't add a trailing "." at the end of paragraphs (just like this file). That means every paragraph, not just the last one (of the markdown file, PR description, GitHub comment, etc.). Rule of thumb: if you're adding new line(s) before the next sentence, don't add a "." - don't use →. Instead, prefer not to use arrows, and if need be, use -> instead +- do use plain, simple, everyday engineering language: the common phrase engineers actually say over rare compact phrasing, in grammatically complete sentences. When structure genuinely helps the reader, prefer nested bullets (any depth is fine) over one dense line. This applies to all human-facing text: discussion posts, release notes, and docs included Don't hesitate to use values in .env to get needed API keys and other secrets, as long as you never add them to conversation history, commit them, or include them in GitHub issues / PRs From d749b186de18b1861a996044eca01156e1ae2a4e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 25 Aug 2026 12:34:31 -0700 Subject: [PATCH 055/281] docs(pr-template): make intent the severe-vs-high discriminator --- .github/pull_request_template.md | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index bcdb228746a..8a19547cb34 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -116,10 +116,12 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac -### Final Attestation +## Final Attestation - [ ] The tests check the right things, including the edge cases, and regressions in the respective real-world customer use-cases are not possible after this PR From 27ca05a70759d8c2e78b2e4c0bc08aa526640ed2 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 25 Aug 2026 13:01:51 -0700 Subject: [PATCH 057/281] fix(ui): read reasoning tokens from Responses API output_tokens_details (#37952) --- .../src/components/llm_calls/responses_api.test.tsx | 12 ++++++++++++ .../src/components/llm_calls/responses_api.tsx | 6 ++++-- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/llm_calls/responses_api.test.tsx b/ui/litellm-dashboard/src/components/llm_calls/responses_api.test.tsx index 065eaaf3632..033813397fc 100644 --- a/ui/litellm-dashboard/src/components/llm_calls/responses_api.test.tsx +++ b/ui/litellm-dashboard/src/components/llm_calls/responses_api.test.tsx @@ -400,4 +400,16 @@ describe("responses_api prompt cache usage", () => { expect(usageData).not.toHaveProperty("cacheCreationTokens"); expect(usageData.promptTokens).toBe(5000); }); + + it("surfaces reasoning tokens from Responses-shape output_tokens_details", async () => { + await expect(captureUsage({ output_tokens_details: { reasoning_tokens: 42 } })).resolves.toMatchObject({ + reasoningTokens: 42, + }); + }); + + it("falls back to completion_tokens_details reasoning tokens when output_tokens_details is absent", async () => { + await expect(captureUsage({ completion_tokens_details: { reasoning_tokens: 17 } })).resolves.toMatchObject({ + reasoningTokens: 17, + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/llm_calls/responses_api.tsx b/ui/litellm-dashboard/src/components/llm_calls/responses_api.tsx index 94e8cb46765..8d71a4e29a8 100644 --- a/ui/litellm-dashboard/src/components/llm_calls/responses_api.tsx +++ b/ui/litellm-dashboard/src/components/llm_calls/responses_api.tsx @@ -295,8 +295,10 @@ export async function makeOpenAIResponsesRequest( }; // Add reasoning tokens if available - if (usage.completion_tokens_details?.reasoning_tokens) { - usageData.reasoningTokens = usage.completion_tokens_details.reasoning_tokens; + const reasoningTokens = + usage.output_tokens_details?.reasoning_tokens ?? usage.completion_tokens_details?.reasoning_tokens; + if (reasoningTokens) { + usageData.reasoningTokens = reasoningTokens; } if (usage.cost !== undefined && usage.cost !== null) { From 104fe73113cafa167a11edfa7c268b1d17b5ca29 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 25 Aug 2026 13:07:09 -0700 Subject: [PATCH 058/281] fix(dashboard): don't show a stale provider prompt-cache chip on a response-cache hit (#37951) * fix(dashboard): don't show a stale provider prompt-cache chip on a response-cache hit The playground's non-streaming chat completion and responses paths replayed a cache hit's original usage payload verbatim, so ResponseMetrics kept rendering the provider's prompt-cache-write/read chips using token counts from the original request. Detect the hit via the x-litellm-cache-key response header and render a Response Cache indicator instead. * fix(dashboard): expose x-litellm-cache-key through CORS for the playground cache-hit indicator --- litellm/constants.py | 1 + tests/test_litellm/proxy/test_proxy_server.py | 10 + .../chat_ui/ResponseMetrics.test.tsx | 18 ++ .../components/chat_ui/ResponseMetrics.tsx | 20 ++ .../llm_calls/chat_completion.test.tsx | 210 +++++++++++++++--- .../components/llm_calls/chat_completion.tsx | 10 +- .../llm_calls/responses_api.test.tsx | 188 +++++++++++++--- .../components/llm_calls/responses_api.tsx | 12 +- 8 files changed, 408 insertions(+), 61 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 78aba30f9c0..765bbfe1e54 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -147,6 +147,7 @@ LITELLM_UI_ALLOW_HEADERS: Final = [ "x-litellm-adaptive-router-model", "x-litellm-applied-guardrails", "x-litellm-guardrail-scan-id", + "x-litellm-cache-key", ] # Gemini model-specific minimal thinking budget constants diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 3383527e932..b9a31acca96 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -78,6 +78,16 @@ def client_no_auth(): return TestClient(app) +def test_cors_exposes_cache_key_header_to_browser_js(): + from fastapi.middleware.cors import CORSMiddleware + + from litellm.constants import LITELLM_UI_ALLOW_HEADERS + + cors_middleware = next(m for m in app.user_middleware if m.cls is CORSMiddleware) + assert cors_middleware.kwargs["expose_headers"] is LITELLM_UI_ALLOW_HEADERS + assert "x-litellm-cache-key" in cors_middleware.kwargs["expose_headers"] + + def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch): mock_login_result = {"user_id": "test-user"} mock_prisma_client = MagicMock() diff --git a/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.test.tsx b/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.test.tsx index 5afc94eb043..f31e1839739 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.test.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.test.tsx @@ -33,4 +33,22 @@ describe("ResponseMetrics prompt cache chips", () => { expect(screen.queryByText(/Cache Read/)).not.toBeInTheDocument(); expect(screen.queryByText(/Cache Write/)).not.toBeInTheDocument(); }); + + it("shows the response cache indicator instead of the provider cache chips on a response-cache hit", () => { + render( + , + ); + + expect(screen.getByText("Response Cache: Hit")).toBeInTheDocument(); + expect(screen.queryByText(/Cache Read/)).not.toBeInTheDocument(); + expect(screen.queryByText(/Cache Write/)).not.toBeInTheDocument(); + }); + + it("does not show the response cache indicator when the flag is absent", () => { + render(); + + expect(screen.queryByText(/Response Cache/)).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.tsx b/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.tsx index 3e7f2884b23..ec62d0618d7 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ResponseMetrics.tsx @@ -7,12 +7,16 @@ import { DatabaseBackup, DollarSign, Hash, + History, Lightbulb, Wrench, } from "lucide-react"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; import { PROMPT_CACHE_CREATION_TOOLTIP, PROMPT_CACHE_READ_TOOLTIP } from "@/utils/promptCacheUsage"; +const RESPONSE_CACHE_TOOLTIP = + "This response was replayed from LiteLLM's response cache. The request never reached the provider, so it did not read from or write to the provider's own prompt cache."; + export interface TokenUsage { completionTokens?: number; promptTokens?: number; @@ -21,6 +25,7 @@ export interface TokenUsage { cacheReadTokens?: number; cacheCreationTokens?: number; cost?: number; + servedFromResponseCache?: boolean; } interface ResponseMetricsProps { @@ -51,7 +56,22 @@ function MetricItem({ label, tooltip, icon, value }: MetricItemProps) { ); } +function ResponseCacheIndicator() { + return ( +