From 48fe111c12f6599fafb694946d86c9d1e1dd77a7 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 13 Jul 2026 17:47:48 +0000 Subject: [PATCH 001/253] fix(responses): stop managed Responses WS from leaking litellm_params into provider body Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/responses/streaming_iterator.py | 3 +- .../test_responses_websocket_all_providers.py | 53 +++++++++++++++++++ 2 files changed, 54 insertions(+), 2 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index eb78e6f9c8d..616a6659a55 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -2097,8 +2097,7 @@ class ManagedResponsesWebSocketHandler: if "litellm_metadata" not in call_kwargs: call_kwargs["litellm_metadata"] = {} call_kwargs["litellm_metadata"]["proxy_server_request"] = proxy_server_request - call_kwargs.setdefault("litellm_params", {}) - call_kwargs["litellm_params"]["proxy_server_request"] = proxy_server_request + call_kwargs["proxy_server_request"] = proxy_server_request async def _stream_and_forward(self, model: str, call_kwargs: Dict[str, Any]) -> Optional[Dict[str, Any]]: """ diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 4509abc7749..7557757ded1 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -660,6 +660,59 @@ class TestChunkTransformation: assert ManagedResponsesWebSocketHandler._input_to_messages({}) == [] +class TestUpdateProxyRequest: + """Regression tests for ManagedResponsesWebSocketHandler._update_proxy_request. + + The managed WebSocket path calls ``litellm.aresponses(model=..., **call_kwargs)``. + ``litellm_params`` is not a Responses API request field, so passing it as a + top-level kwarg leaks it into the provider request body and providers that + forbid extra inputs (e.g. Anthropic) reject the call with + ``litellm_params: Extra inputs are not permitted``. The request-tracking data + must ride along as ``proxy_server_request`` instead, which litellm consumes + internally and never forwards to the provider. + """ + + def test_does_not_inject_litellm_params_kwarg(self): + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + call_kwargs = { + "input": "hello", + "store": True, + "litellm_metadata": { + "proxy_server_request": {"headers": {}, "body": {}}, + }, + } + + ManagedResponsesWebSocketHandler._update_proxy_request( + call_kwargs, "anthropic/claude-sonnet-4-5" + ) + + assert "litellm_params" not in call_kwargs + assert call_kwargs["proxy_server_request"]["body"]["model"] == ( + "anthropic/claude-sonnet-4-5" + ) + assert call_kwargs["proxy_server_request"]["body"]["input"] == "hello" + + def test_proxy_server_request_matches_metadata(self): + from litellm.responses.streaming_iterator import ( + ManagedResponsesWebSocketHandler, + ) + + call_kwargs = { + "input": "hi", + "litellm_metadata": {"proxy_server_request": {"body": {}}}, + } + + ManagedResponsesWebSocketHandler._update_proxy_request(call_kwargs, "gpt-4o") + + assert ( + call_kwargs["proxy_server_request"] + == call_kwargs["litellm_metadata"]["proxy_server_request"] + ) + + class TestWebSocketEventTypes: """Test that all WebSocket event types are properly handled with dict-based chunks""" From 52f3ff13f0a7da9f5a2cbe7333e9b7dcd8643780 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Thu, 6 Aug 2026 13:26:11 -0700 Subject: [PATCH 002/253] feat(proxy): limit repeated failed Admin UI sign-in attempts The Admin UI sign-in endpoints accept an unbounded number of password attempts. All three call authenticate_user, and none of them keeps any record of how many times a given caller has already been refused, so a misbehaving or misconfigured client can retry indefinitely at full speed. A LoginThrottle is now a required argument to authenticate_user, so the accounting lives at the one function all three endpoints share and a fourth endpoint cannot be added without deciding what to pass. Failures are counted per username and source address over a fixed window and further attempts are refused with 429 and a Retry-After header. The check runs before the database lookup and before the password comparison, so a refused caller does no further work. Only genuine credential rejections count. Configuration errors do not, a refused attempt does not extend the window, and a successful sign-in clears the bucket. The username is case folded because the user lookup is case insensitive, so casing cannot multiply the allowance. Both credential rejections now return one identical message. SSO is unaffected; it never calls this function. max_failed_login_attempts (10) and failed_login_window_seconds (900) are read from config.yaml, with LITELLM_DISABLE_LOGIN_RATE_LIMIT to turn the accounting off. They are deliberately not database backed, so editing YAML always wins and an operator refused by a bad value can recover. --- litellm/proxy/_types.py | 10 + litellm/proxy/auth/login_throttle.py | 177 +++++++ litellm/proxy/auth/login_utils.py | 18 +- litellm/proxy/proxy_server.py | 7 +- .../proxy/auth/test_login_utils.py | 477 ++++++++++++++++++ .../proxy/proxy_server/conftest.py | 25 + .../proxy_server/test_routes_login_sso.py | 92 +++- tests/test_litellm/proxy/test_proxy_server.py | 39 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 10 + 9 files changed, 843 insertions(+), 12 deletions(-) create mode 100644 litellm/proxy/auth/login_throttle.py diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ed35691c6ec..1fd2781bd95 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2525,6 +2525,16 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): description="sends alerts if requests hang for 5min+", ) ui_access_mode: Literal["admin_only", "all"] | None = Field("all", description="Control access to the Proxy UI") + max_failed_login_attempts: int | None = Field( + None, + ge=1, + description="Number of failed Admin UI sign-in attempts allowed for one username from one source address within `failed_login_window_seconds`, before further attempts are refused with 429. Configurable from config.yaml only. Defaults to 10", + ) + failed_login_window_seconds: int | None = Field( + None, + ge=1, + description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Configurable from config.yaml only. Defaults to 900", + ) allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access") reject_clientside_metadata_tags: bool | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py new file mode 100644 index 00000000000..e2befd15a35 --- /dev/null +++ b/litellm/proxy/auth/login_throttle.py @@ -0,0 +1,177 @@ +"""Failed-login accounting for the Admin UI sign-in path. + +Counts failed credential checks per (username, source address) over a fixed window and +denies further attempts with 429 once the count reaches the limit. Built per request by +``LoginThrottle.from_request`` because it carries that request's resolved source address, +and because the coordination cache is assigned at startup and can be reassigned later. +""" + +import hashlib +from collections.abc import Awaitable +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final + +from fastapi import Request + +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache +from litellm.proxy._types import ProxyErrorTypes, ProxyException +from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges, resolve_client_ip +from litellm.proxy.auth.trusted_proxy_utils import TRUSTED_PROXY_RANGES_KEY +from litellm.secret_managers.main import get_secret_bool + +DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS: Final = 10 +DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 900 + +_CACHE_KEY_PREFIX: Final = "login_fail" +_UNKNOWN_SOURCE: Final = "unknown" +_MAX_LOGGED_USERNAME_CHARS: Final = 128 + +_MAX_TRACKED_LOGIN_SOURCES: Final = 10_000 + +_FAILED_LOGIN_CACHE: Final = DualCache( + in_memory_cache=InMemoryCache(max_size_in_memory=_MAX_TRACKED_LOGIN_SOURCES), + default_in_memory_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, +) +_NO_SETTINGS: Final = MappingProxyType({}) + + +def _int_setting(name: str, value: object, default: int, minimum: int) -> int: + if value is None: + return default + if isinstance(value, bool) or not isinstance(value, int) or value < minimum: + verbose_proxy_logger.warning( + "general_settings.%s=%r is not an integer >= %s; using the default of %s", name, value, minimum, default + ) + return default + return value + + +def _as_count(cached: object) -> int: + return int(cached) if isinstance(cached, int | float) and not isinstance(cached, bool) else 0 + + +@dataclass(frozen=True, slots=True) +class LoginThrottle: + """Fixed-window failed-login accounting for one request's source address.""" + + client_ip: str + max_attempts: int + window_seconds: int + cache: DualCache + redis_cache: RedisCache | None = None + enabled: bool = True + + @classmethod + def from_request(cls, request: Request) -> "LoginThrottle": + """Build the throttle for this request from the live proxy settings and caches.""" + from litellm.proxy.proxy_server import general_settings, redis_usage_cache + + settings: Final = general_settings or _NO_SETTINGS + cidrs: Final = normalize_cidr_ranges( + settings.get(TRUSTED_PROXY_RANGES_KEY), setting_name=TRUSTED_PROXY_RANGES_KEY + ) + resolved, _ = resolve_client_ip( + request, TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs) + ) + return cls( + client_ip=resolved or _UNKNOWN_SOURCE, + max_attempts=_int_setting( + "max_failed_login_attempts", + settings.get("max_failed_login_attempts"), + DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS, + 1, + ), + window_seconds=_int_setting( + "failed_login_window_seconds", + settings.get("failed_login_window_seconds"), + DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, + 1, + ), + cache=_FAILED_LOGIN_CACHE, + redis_cache=redis_usage_cache, + enabled=not get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT"), + ) + + @staticmethod + def _loggable(username: str) -> str: + """The username with anything that could forge a log line removed.""" + return "".join(c for c in username if c.isprintable())[:_MAX_LOGGED_USERNAME_CHARS] + + def _key(self, username: str) -> str: + identity: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest() + return f"{_CACHE_KEY_PREFIX}:{identity}:{self.client_ip}" + + async def _outcome(self, work: Awaitable[object]) -> object: + try: + return await work + except Exception as exc: # noqa: BLE001 # an unreachable cache must never deny a valid credential + verbose_proxy_logger.warning("login attempt accounting unavailable: %s", exc) + return None + + async def _failures(self, key: str) -> int: + store: Final = self.cache if self.redis_cache is None else self.redis_cache + return _as_count(await self._outcome(store.async_get_cache(key=key))) + + async def _ensure_expiry(self, key: str) -> None: + """Give the counter an expiry if it somehow has none. + + Redis commits the increment before setting the TTL, so a failure in between can + leave a counter that never expires. Nothing increments the key again once the + limit is reached, so without this the pair would stay refused indefinitely. + """ + redis_cache: Final = self.redis_cache + if redis_cache is None: + return + if isinstance(await self._outcome(redis_cache.async_get_ttl(key)), int): + return + await self._outcome(redis_cache.async_increment(key, 0, ttl=self.window_seconds)) + + async def raise_if_blocked(self, username: str) -> None: + """Deny before the database lookup and before the password comparison.""" + if not self.enabled: + return + key: Final = self._key(username) + if await self._failures(key) < self.max_attempts: + return + await self._ensure_expiry(key) + verbose_proxy_logger.warning( + "Admin UI sign-in attempts exhausted for username=%s source=%s; %s attempts in %ss, retry after %ss", + self._loggable(username), + self.client_ip, + self.max_attempts, + self.window_seconds, + self.window_seconds, + ) + raise ProxyException( + message="Too many failed sign-in attempts. Try again later.", + type=ProxyErrorTypes.auth_error, + param="max_failed_login_attempts", + code=429, + headers={"Retry-After": str(self.window_seconds)}, # mutable-ok: ProxyException coerces header values + ) + + async def record_failure(self, username: str) -> None: + """Count one rejected credential guess against this username and source.""" + if not self.enabled: + return + key: Final = self._key(username) + redis_cache: Final = self.redis_cache + if redis_cache is not None: + await self._outcome(redis_cache.async_increment(key, 1, ttl=self.window_seconds)) + await self._ensure_expiry(key) + return + await self._outcome(self.cache.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) + + async def clear(self, username: str) -> None: + """Drop the bucket after a successful sign-in.""" + if not self.enabled: + return + key: Final = self._key(username) + redis_cache: Final = self.redis_cache + if redis_cache is not None: + await self._outcome(redis_cache.async_delete_cache(key)) + await self._outcome(self.cache.async_delete_cache(key=key)) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index fba95972944..98780a4692a 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -24,6 +24,7 @@ from litellm.proxy._types import ( UpdateUserRequest, UserAPIKeyAuth, ) +from litellm.proxy.auth.login_throttle import LoginThrottle from litellm.proxy.management_endpoints.internal_user_endpoints import user_update from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, @@ -41,6 +42,10 @@ from litellm.repositories.user_repository import UserRepository from litellm.secret_managers.main import get_secret_bool from litellm.types.proxy.ui_sso import ReturnedUITokenObject +INVALID_UI_CREDENTIALS_MESSAGE: Final = ( + "Invalid credentials used to access UI. Check 'UI_USERNAME' and 'UI_PASSWORD', or the password set for your user" +) + async def _rehash_password_if_needed(user_id: str, password: str, stored: str) -> None: """Rehash legacy password (SHA256) to scrypt on successful login.""" @@ -111,6 +116,7 @@ async def authenticate_user( password: str, master_key: str | None, prisma_client: PrismaClient | None, + throttle: LoginThrottle, ) -> LoginResult: """ Authenticate a user and generate an API key for UI access. @@ -139,6 +145,8 @@ async def authenticate_user( code=500, ) + await throttle.raise_if_blocked(username) + ui_username, ui_password = get_ui_credentials(master_key) # Check if we can find the `username` in the db. On the UI, users can enter username=their email @@ -240,6 +248,8 @@ async def authenticate_user( key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user_info) + await throttle.clear(username) + return LoginResult( user_id=user_id, key=key, @@ -294,6 +304,8 @@ async def authenticate_user( key = response["token"] + await throttle.clear(username) + return LoginResult( user_id=user_id, key=key, @@ -302,15 +314,17 @@ async def authenticate_user( login_method="username_password", ) else: + await throttle.record_failure(username) raise ProxyException( - message=f"Invalid credentials used to access UI.\nNot valid credentials for {username}", + message=INVALID_UI_CREDENTIALS_MESSAGE, type=ProxyErrorTypes.auth_error, param="invalid_credentials", code=401, ) else: + await throttle.record_failure(username) raise ProxyException( - message="Invalid credentials used to access UI.\nCheck 'UI_USERNAME', 'UI_PASSWORD' in .env file", + message=INVALID_UI_CREDENTIALS_MESSAGE, type=ProxyErrorTypes.auth_error, param="invalid_credentials", code=401, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9b5d1b9bfea..39a419df73f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -293,6 +293,7 @@ from litellm.proxy.auth.auth_utils import ( ) from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import LicenseCheck +from litellm.proxy.auth.login_throttle import LoginThrottle from litellm.proxy.auth.model_checks import ( expand_wildcard_deployments_for_model_info, get_all_fallbacks, @@ -723,6 +724,7 @@ from fastapi.openapi.docs import get_swagger_ui_html from fastapi.openapi.utils import get_openapi from fastapi.responses import ( FileResponse, + HTMLResponse, JSONResponse, ORJSONResponse, RedirectResponse, @@ -14730,8 +14732,6 @@ async def fallback_login(request: Request): else: redirect_url += "/sso/callback" - from fastapi.responses import HTMLResponse - hide_default_credentials_hint: Final = ( os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true" or general_settings.get("hide_default_credentials_hint", False) is True @@ -14761,6 +14761,7 @@ async def login(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, + throttle=LoginThrottle.from_request(request), ) # Create UI token object @@ -14835,6 +14836,7 @@ async def login_v2(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, + throttle=LoginThrottle.from_request(request), ) returned_ui_token_object: Final = create_ui_token_object( @@ -14905,6 +14907,7 @@ async def login_v3(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, + throttle=LoginThrottle.from_request(request), ) returned_ui_token_object: Final = create_ui_token_object( diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 1c66acf8678..1e3fcaf5d78 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -10,6 +10,16 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest + +def _unlimited_throttle(): + """A throttle wired to a real in-memory store with a limit no test can reach.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.auth.login_throttle import LoginThrottle + + return LoginThrottle(client_ip="1.2.3.4", max_attempts=10_000, window_seconds=900, cache=DualCache()) + + + from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._types import ( LiteLLM_UserTable, @@ -98,6 +108,7 @@ async def test_authenticate_user_admin_login_with_ui_credentials(): password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert isinstance(result, LoginResult) @@ -155,6 +166,7 @@ async def test_authenticate_user_admin_login_with_master_key_as_password(monkeyp password=master_key, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert isinstance(result, LoginResult) @@ -179,6 +191,7 @@ async def test_authenticate_user_invalid_credentials(): password=wrong_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert exc_info.value.type == ProxyErrorTypes.auth_error @@ -197,6 +210,7 @@ async def test_authenticate_user_missing_master_key(): password="password", master_key=None, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert exc_info.value.type == ProxyErrorTypes.auth_error @@ -237,6 +251,7 @@ async def test_authenticate_user_wrong_password(): password=wrong_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert exc_info.value.type == ProxyErrorTypes.auth_error @@ -295,12 +310,14 @@ async def test_authenticate_user_email_case_insensitive_login(): password=correct_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) result_lower = await authenticate_user( username=stored_email, password=correct_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert result_mixed.user_id == result_lower.user_id == "test-user-123" @@ -342,6 +359,7 @@ async def test_authenticate_user_database_required_for_admin(monkeypatch): password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert exc_info.value.type == ProxyErrorTypes.auth_error @@ -393,6 +411,7 @@ async def test_authenticate_user_admin_login_with_non_ascii_characters(): password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert isinstance(result, LoginResult) @@ -469,18 +488,21 @@ async def test_authenticate_user_multiple_logins_generate_unique_tokens(): password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) result2 = await authenticate_user( username=ui_username, password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) result3 = await authenticate_user( username=ui_username, password=ui_password, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) # Each login should return a unique token @@ -538,6 +560,7 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): password=password_with_special_char, master_key=master_key, prisma_client=mock_prisma_client, + throttle=_unlimited_throttle(), ) assert isinstance(result, LoginResult) @@ -598,3 +621,457 @@ class TestEncodeUiSessionJwt: request.cookies = {"token": token} with patch("litellm.proxy.proxy_server.master_key", "sk-master-for-tests"): assert _user_id_from_session_cookie(request) == "cornell-user" + + +# --------------------------------------------------------------------------- +# Failed-login accounting (LIT-5285) +# --------------------------------------------------------------------------- + + +def _throttle(max_attempts: int = 3, window_seconds: int = 900, client_ip: str = "1.2.3.4", cache=None, redis_cache=None): + """A throttle over a real in-memory store, so the tests exercise the true counters.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.auth.login_throttle import LoginThrottle + + return LoginThrottle( + client_ip=client_ip, + max_attempts=max_attempts, + window_seconds=window_seconds, + cache=cache if cache is not None else DualCache(), + redis_cache=redis_cache, + ) + + +async def _guess(throttle, username: str = "admin", password: str = "wrong"): + from litellm.proxy.auth.login_utils import authenticate_user + + return await authenticate_user( + username=username, + password=password, + master_key="sk-master", + prisma_client=None, + throttle=throttle, + ) + + +@pytest.mark.asyncio +async def test_attempts_are_refused_once_the_limit_is_reached(monkeypatch): + """The limit denies further attempts for the window, and the denial carries Retry-After.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=3, window_seconds=77) + + for _ in range(3): + with pytest.raises(ProxyException) as first: + await _guess(throttle) + assert first.value.code == "401" + + with pytest.raises(ProxyException) as blocked: + await _guess(throttle) + assert blocked.value.code == "429" + assert blocked.value.headers.get("Retry-After") == "77" + + +@pytest.mark.asyncio +async def test_a_correct_password_is_refused_while_blocked(monkeypatch): + """The check precedes the credential comparison, so being over the limit wins.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=2) + + for _ in range(2): + with pytest.raises(ProxyException): + await _guess(throttle) + + with pytest.raises(ProxyException) as blocked: + await _guess(throttle, password="right") + assert blocked.value.code == "429" + + +@pytest.mark.asyncio +async def test_a_blocked_attempt_does_not_extend_the_window(monkeypatch): + """Hammering while blocked must not push the counter or refresh its TTL.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=2) + key = throttle._key("admin") + + for _ in range(2): + with pytest.raises(ProxyException): + await _guess(throttle) + counted_at_limit = await throttle._failures(key) + + for _ in range(5): + with pytest.raises(ProxyException): + await _guess(throttle) + + assert await throttle._failures(key) == counted_at_limit == 2 + + +@pytest.mark.asyncio +async def test_a_successful_sign_in_clears_the_bucket(monkeypatch): + """Success resets the budget rather than leaving the operator near the limit.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(max_attempts=3) + + for _ in range(2): + with pytest.raises(ProxyException): + await _guess(throttle) + + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ): + await _guess(throttle, password="right") + + assert await throttle._failures(throttle._key("admin")) == 0 + + +@pytest.mark.asyncio +async def test_a_configuration_error_never_counts(monkeypatch): + """A 500 from an unset master key is not a guess and must not consume the budget.""" + from litellm.proxy._types import ProxyException + + throttle = _throttle(max_attempts=2) + for _ in range(5): + with pytest.raises(ProxyException) as exc: + await authenticate_user( + username="admin", password="x", master_key=None, prisma_client=None, throttle=throttle + ) + assert exc.value.code == "500" + + assert await throttle._failures(throttle._key("admin")) == 0 + + +@pytest.mark.asyncio +async def test_the_username_is_case_folded_into_one_bucket(monkeypatch): + """The DB lookup is case-insensitive, so casing must not multiply the budget.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=4) + + for name in ("admin@corp.com", "ADMIN@corp.com", "Admin@corp.com", "aDmIn@corp.com"): + with pytest.raises(ProxyException) as exc: + await _guess(throttle, username=name) + assert exc.value.code == "401" + + with pytest.raises(ProxyException) as blocked: + await _guess(throttle, username="admin@CORP.com") + assert blocked.value.code == "429" + + +@pytest.mark.asyncio +async def test_a_different_username_from_the_same_source_is_unaffected(monkeypatch): + """The bucket is the pair, so one username's failures do not block another.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=2) + + for _ in range(3): + with pytest.raises(ProxyException): + await _guess(throttle, username="admin") + + with pytest.raises(ProxyException) as other: + await _guess(throttle, username="someone-else@example.com") + assert other.value.code == "401", "a second username must still reach the credential check" + + +@pytest.mark.asyncio +async def test_both_credential_rejections_are_indistinguishable(monkeypatch): + """One message for the known and the unknown username, so responses do not enumerate.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + + with pytest.raises(ProxyException) as unknown: + await _guess(_throttle(max_attempts=99), username="nobody@example.com") + + fake_user = MagicMock() + fake_user.user_id = "u-1" + fake_user.user_email = "known@example.com" + fake_user.user_role = "internal_user" + fake_user.password = "scrypt:fake" + repo = MagicMock() + repo.return_value.table.find_first = AsyncMock(return_value=fake_user) + with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( + "litellm.proxy.auth.login_utils.verify_password", return_value=False + ): + with pytest.raises(ProxyException) as known: + await authenticate_user( + username="known@example.com", + password="wrong", + master_key="sk-master", + prisma_client=MagicMock(), + throttle=_throttle(max_attempts=99), + ) + + assert unknown.value.message == known.value.message + assert "known@example.com" not in unknown.value.message + known.value.message + + +@pytest.mark.asyncio +async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypatch): + """That 401 is deterministic and guards no secret, so counting it would only let + someone burn a passwordless account's bucket.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=2) + + passwordless = MagicMock() + passwordless.user_id = "u-2" + passwordless.user_email = "nopass@example.com" + passwordless.user_role = "internal_user" + passwordless.password = None + repo = MagicMock() + repo.return_value.table.find_first = AsyncMock(return_value=passwordless) + + with patch("litellm.proxy.auth.login_utils.UserRepository", repo): + for _ in range(5): + with pytest.raises(ProxyException) as exc: + await authenticate_user( + username="nopass@example.com", + password="x", + master_key="sk-master", + prisma_client=MagicMock(), + throttle=throttle, + ) + assert exc.value.code == "401" + + assert await throttle._failures(throttle._key("nopass@example.com")) == 0 + + +@pytest.mark.asyncio +async def test_a_wrong_password_for_a_known_user_also_counts(monkeypatch): + """The database-user branch must charge the bucket too, not just the unknown-user branch.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=3) + + known = MagicMock() + known.user_id = "u-1" + known.user_email = "known@example.com" + known.user_role = "internal_user" + known.password = "scrypt:stored" + repo = MagicMock() + repo.return_value.table.find_first = AsyncMock(return_value=known) + + async def _attempt(): + return await authenticate_user( + username="known@example.com", + password="wrong", + master_key="sk-master", + prisma_client=MagicMock(), + throttle=throttle, + ) + + with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( + "litellm.proxy.auth.login_utils.verify_password", return_value=False + ): + for _ in range(3): + with pytest.raises(ProxyException) as rejected: + await _attempt() + assert rejected.value.code == "401" + + with pytest.raises(ProxyException) as blocked: + await _attempt() + assert blocked.value.code == "429" + + +@pytest.mark.asyncio +async def test_two_source_addresses_do_not_share_a_bucket(monkeypatch): + """The key is the pair, so one address exhausting its budget must not block another. + + Dropping the address from the key would turn this into the username-only counter the + design rejects, where anyone can lock a named admin out from anywhere. + """ + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + shared_store = DualCache() + attacker = _throttle(max_attempts=2, client_ip="203.0.113.9", cache=shared_store) + operator = _throttle(max_attempts=2, client_ip="198.51.100.7", cache=shared_store) + + for _ in range(3): + with pytest.raises(ProxyException): + await _guess(attacker, username="admin") + + with pytest.raises(ProxyException) as blocked: + await _guess(attacker, username="admin") + assert blocked.value.code == "429" + + with pytest.raises(ProxyException) as unaffected: + await _guess(operator, username="admin") + assert unaffected.value.code == "401", "the real operator must still reach the credential check" + + +class _NoExpiryRedis: + """Redis that stores the counter but never records an expiry for it. + + Models the window between INCRBYFLOAT committing and the TTL call failing. + """ + + def __init__(self): + self.values: dict = {} + self.expiry_repairs = 0 + + async def async_get_cache(self, key, **kwargs): + return self.values.get(key) + + async def async_increment(self, key, value, ttl=None, **kwargs): + if int(value) == 0: + self.expiry_repairs += 1 + self.values[key] = self.values.get(key, 0) + int(value) + return self.values[key] + + async def async_get_ttl(self, key): + return None + + async def async_delete_cache(self, key): + self.values.pop(key, None) + + +@pytest.mark.asyncio +async def test_a_counter_left_without_an_expiry_is_repaired(monkeypatch): + """Regression: a counter with no TTL would refuse the pair forever. + + Redis commits the increment before setting the expiry, and nothing increments the key + again once the limit is reached, so a TTL that never landed is never repaired on its + own and the username and source pair stays refused with no way back. + """ + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + redis = _NoExpiryRedis() + throttle = _throttle(max_attempts=2, redis_cache=redis) + + for _ in range(2): + with pytest.raises(ProxyException): + await _guess(throttle) + + assert redis.expiry_repairs >= 1, "each recorded failure must leave the counter with an expiry" + + repairs_before_block = redis.expiry_repairs + with pytest.raises(ProxyException) as blocked: + await _guess(throttle) + assert blocked.value.code == "429" + assert redis.expiry_repairs > repairs_before_block, "the refusal path must repair a missing expiry too" + + +@pytest.mark.asyncio +async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): + """Regression: throttle entries must not evict cached credentials. + + user_api_key_cache holds at most 200 in-memory entries and evicts the soonest to + expire first, so parking 900s sign-in counters there let a stream of made-up usernames + push out the much shorter lived credential entries, sending every ordinary API request + back to the database. + """ + from litellm.proxy import proxy_server as ps + from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX, LoginThrottle + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setattr(ps, "redis_usage_cache", None) + + auth_cache_keys_before = set(ps.user_api_key_cache.in_memory_cache.cache_dict) + + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "1.2.3.4" + throttle = LoginThrottle.from_request(request) + + for i in range(25): + with pytest.raises(Exception): + await _guess(throttle, username=f"made-up-{i}@example.com") + + added = set(ps.user_api_key_cache.in_memory_cache.cache_dict) - auth_cache_keys_before + assert not [k for k in added if str(k).startswith(_CACHE_KEY_PREFIX)], ( + "sign-in counters must live in their own cache, not the key-authentication cache" + ) + + +@pytest.mark.asyncio +async def test_a_refused_username_cannot_forge_log_lines(monkeypatch): + """The username reaches a warning log, so it must not carry newlines or control bytes.""" + import logging + + from litellm._logging import verbose_proxy_logger + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=1) + forged = "victim@example.com\nWARNING: sign-in succeeded for attacker\x00" + + with pytest.raises(ProxyException): + await _guess(throttle, username=forged) + + records: list[logging.LogRecord] = [] + handler = logging.Handler() + handler.emit = records.append + verbose_proxy_logger.addHandler(handler) + try: + with pytest.raises(ProxyException) as blocked: + await _guess(throttle, username=forged) + finally: + verbose_proxy_logger.removeHandler(handler) + + assert blocked.value.code == "429" + emitted = [r.getMessage() for r in records if "sign-in attempts exhausted" in r.getMessage()] + assert emitted, "the refusal must be logged" + assert "\n" not in emitted[0] and "\x00" not in emitted[0] + assert "victim@example.com" in emitted[0] + + +@pytest.mark.asyncio +async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): + """Regression: the in-memory tier must hold more counters than a spray can create. + + The default in-memory cache keeps 200 entries and evicts the soonest to expire, and + every counter shares one window, so eviction was effectively oldest-first. A few + hundred made-up usernames therefore pushed out the attacker's own counter and handed + back a fresh allowance against the real account. + """ + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.login_throttle import _FAILED_LOGIN_CACHE, _MAX_TRACKED_LOGIN_SOURCES + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + assert _MAX_TRACKED_LOGIN_SOURCES >= 10_000 + assert _FAILED_LOGIN_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_SOURCES + + throttle = _throttle(max_attempts=3, cache=_FAILED_LOGIN_CACHE, client_ip="10.9.9.9") + victim = "spray-victim@corp.com" + for _ in range(3): + with pytest.raises(ProxyException): + await _guess(throttle, username=victim) + + for i in range(500): + await throttle.record_failure(f"spray-filler-{i}@corp.com") + + assert await throttle._failures(throttle._key(victim)) == 3, "the counter must survive a spray" + with pytest.raises(ProxyException) as blocked: + await _guess(throttle, username=victim) + assert blocked.value.code == "429" diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index c545965f9a9..cdfd3549544 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -511,3 +511,28 @@ def make_key( max_budget=max_budget, **kwargs, ) + + +@pytest.fixture(autouse=True) +def reset_login_throttle(monkeypatch): + """Clear the Admin UI failed-login counters between tests. + + `client` is session scoped and the counters live in the shared `user_api_key_cache` + with a 900s window, so without this any test that fails a sign-in enough times would + start returning 429 from unrelated tests later in the same process. Only the throttle's + own keys are removed, so nothing else in that cache is disturbed. + """ + from litellm.proxy import proxy_server as ps + from litellm.proxy.auth.login_throttle import _FAILED_LOGIN_CACHE, _CACHE_KEY_PREFIX + + def _drop_throttle_keys() -> None: + in_memory = getattr(_FAILED_LOGIN_CACHE, "in_memory_cache", None) + cache_dict = getattr(in_memory, "cache_dict", None) + if isinstance(cache_dict, dict): + for key in [k for k in cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)]: + cache_dict.pop(key, None) + + monkeypatch.setattr(ps, "redis_usage_cache", None) + _drop_throttle_keys() + yield _drop_throttle_keys + _drop_throttle_keys() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index af37dbe85fe..faeb7750e03 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -10,6 +10,7 @@ Routes covered: from __future__ import annotations +from concurrent.futures import ThreadPoolExecutor from unittest.mock import AsyncMock, MagicMock import pytest @@ -29,7 +30,7 @@ def _install_login_mocks(monkeypatch, raise_on_auth: bool = False) -> None: """ from litellm.proxy import proxy_server as ps - async def _fake_auth(username, password, master_key, prisma_client): + async def _fake_auth(username, password, master_key, prisma_client, throttle=None): if raise_on_auth: raise Exception("boom-auth-failure") fake = MagicMock() @@ -460,3 +461,92 @@ def test_login_form_ignores_open_redirect_return_to(client, monkeypatch): location = response.headers.get("location", "") assert "evil.example.com" not in location assert "/ui" in location # dashboard fallback + + +# --------------------------------------------------------------------------- +# Failed-login accounting across the login routes (LIT-5285) +# --------------------------------------------------------------------------- + + +def _install_real_auth(monkeypatch, **settings): + """Run the real authenticate_user so the throttle inside it is exercised. + + prisma_client stays None, so every guess falls through to the credential rejection. + """ + from litellm.proxy import proxy_server as ps + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right-password") + monkeypatch.setattr(ps, "master_key", "sk-test-master") + monkeypatch.setattr(ps, "prisma_client", None) + monkeypatch.setattr(ps, "premium_user", False) + monkeypatch.setattr(ps, "general_settings", dict(settings)) + + +def _form_login(client, username="admin", password="wrong"): + return client.post( + "/login", data={"username": username, "password": password}, follow_redirects=False + ).status_code + + +def _json_login(client, path, username="admin", password="wrong"): + return client.post(path, json={"username": username, "password": password}).status_code + + +def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset_login_throttle): + """The endpoint is not part of the key, so spending the budget on one route blocks the rest. + + Partitioning the counter per endpoint would silently triple the real allowance. + """ + _install_real_auth( + monkeypatch, + max_failed_login_attempts=10, + control_plane_url="https://cp.example.com", + ) + + assert [_form_login(client) for _ in range(5)] == [401] * 5 + assert [_json_login(client, "/v2/login") for _ in range(5)] == [401] * 5 + + assert _json_login(client, "/v3/login") == 429, "the eleventh attempt must be refused on a third route" + + +def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_login_throttle): + """The database lookup is case-insensitive, so casing must not partition the counter.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=10) + + assert [_json_login(client, "/v2/login", username="admin@corp.com") for _ in range(5)] == [401] * 5 + assert [_json_login(client, "/v2/login", username="ADMIN@corp.com") for _ in range(5)] == [401] * 5 + + assert _json_login(client, "/v2/login", username="Admin@corp.com") == 429 + + +def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_throttle): + """The 429 tells the caller how long the window has left.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=2, failed_login_window_seconds=77) + + assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] + + refused = client.post("/v2/login", json={"username": "admin", "password": "wrong"}) + assert refused.status_code == 429 + assert refused.headers.get("retry-after") == "77" + + +def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle): + """The bucket is the username and source pair, so one account cannot block another.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=2) + + for _ in range(3): + _json_login(client, "/v2/login", username="admin") + + assert _json_login(client, "/v2/login", username="someone-else@example.com") == 401 + + +def test_sign_in_succeeds_again_once_the_budget_is_restored(client, monkeypatch, reset_login_throttle): + """A cleared bucket lets the same username straight back in.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=2) + + assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] + assert _json_login(client, "/v2/login") == 429 + + reset_login_throttle() + assert _json_login(client, "/v2/login") == 401 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index afc42e8db45..ac2e84b8474 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -19,7 +19,6 @@ from fastapi import FastAPI from fastapi.staticfiles import StaticFiles from fastapi.testclient import TestClient - import litellm import litellm.proxy.proxy_server as proxy_server_module from litellm.caching.caching import RedisCache @@ -27,6 +26,7 @@ from litellm.caching.redis_cluster_cache import RedisClusterCache from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded from litellm.caching.dual_cache import DualCache from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.login_throttle import LoginThrottle from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.proxy_server import app, initialize from litellm.utils import _invalidate_model_cost_lowercase_map @@ -124,12 +124,13 @@ def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch): } assert response.cookies.get("token") == "signed-token" - mock_authenticate_user.assert_awaited_once_with( - username="alice", - password="secret", - master_key="test-master-key", - prisma_client=mock_prisma_client, - ) + mock_authenticate_user.assert_awaited_once() + auth_kwargs = mock_authenticate_user.call_args.kwargs + assert auth_kwargs["username"] == "alice" + assert auth_kwargs["password"] == "secret" + assert auth_kwargs["master_key"] == "test-master-key" + assert auth_kwargs["prisma_client"] is mock_prisma_client + assert isinstance(auth_kwargs["throttle"], LoginThrottle), "the endpoint must thread a throttle through" mock_create_ui_token_object.assert_called_once_with( login_result=mock_login_result, general_settings={}, @@ -11424,3 +11425,27 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the assert real_spend_counter_cache.in_memory_cache.get_cache(key=marker_key) == 0.0, ( "the in-flight DB read clobbered the post-reset floor marker with the stale pre-reset value" ) + + +@pytest.mark.asyncio +async def test_login_throttle_settings_are_not_overridable_from_the_database(): + """LIT-5285: the sign-in limits stay config.yaml only. + + _update_general_settings copies an allowlist of keys out of the DB row. Adding these + to it would let a stored value outrank config.yaml, so an operator refused by a bad + value could not fix it by editing YAML and restarting. + """ + import litellm.proxy.proxy_server as ps + from litellm.proxy.proxy_server import ProxyConfig + + original = dict(ps.general_settings) + try: + ps.general_settings.clear() + await ProxyConfig()._update_general_settings( + db_general_settings={"max_failed_login_attempts": 999, "failed_login_window_seconds": 1} + ) + assert "max_failed_login_attempts" not in ps.general_settings + assert "failed_login_window_seconds" not in ps.general_settings + finally: + ps.general_settings.clear() + ps.general_settings.update(original) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index e0b8cf19159..bb2841b6388 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24150,6 +24150,11 @@ export interface components { * @default false */ enable_public_model_hub: boolean; + /** + * Failed Login Window Seconds + * @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Configurable from config.yaml only. Defaults to 900 + */ + failed_login_window_seconds?: number | null; /** * Forward Client Headers To Llm Api * @description If True, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription. @@ -24194,6 +24199,11 @@ export interface components { * @description max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider */ max_batch_file_size_mb?: number | null; + /** + * Max Failed Login Attempts + * @description Number of failed Admin UI sign-in attempts allowed for one username from one source address within `failed_login_window_seconds`, before further attempts are refused with 429. Configurable from config.yaml only. Defaults to 10 + */ + max_failed_login_attempts?: number | null; /** * Max Parallel Requests * @description maximum parallel requests for each api key From a39efb1a1d59fbf122e8ddfc5929f256bda9f9c0 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Wed, 26 Aug 2026 10:24:32 -0700 Subject: [PATCH 003/253] test(proxy): clear failed-login TTLs when resetting the throttle between tests --- tests/test_litellm/proxy/proxy_server/conftest.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index cdfd3549544..22f2e5b3feb 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -527,10 +527,17 @@ def reset_login_throttle(monkeypatch): def _drop_throttle_keys() -> None: in_memory = getattr(_FAILED_LOGIN_CACHE, "in_memory_cache", None) - cache_dict = getattr(in_memory, "cache_dict", None) - if isinstance(cache_dict, dict): - for key in [k for k in cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)]: - cache_dict.pop(key, None) + if in_memory is None: + return + tracked = tuple( + key + for store in (getattr(in_memory, "cache_dict", None), getattr(in_memory, "ttl_dict", None)) + if isinstance(store, dict) + for key in tuple(store) + if str(key).startswith(_CACHE_KEY_PREFIX) + ) + for key in tracked: + in_memory.delete_cache(key) monkeypatch.setattr(ps, "redis_usage_cache", None) _drop_throttle_keys() From 2d33c949cb15a6007793bd2aeed42ec52d3d9fce Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 13:03:43 -0700 Subject: [PATCH 004/253] feat(proxy): harden Admin UI login throttling --- litellm/proxy/_types.py | 7 +- litellm/proxy/auth/login_throttle.py | 236 ++++++++++---- litellm/proxy/auth/login_utils.py | 17 +- litellm/proxy/proxy_cli.py | 4 + litellm/proxy/proxy_server.py | 43 ++- .../proxy/auth/test_login_utils.py | 294 ++++++++++++++++-- .../proxy/proxy_server/conftest.py | 45 ++- .../proxy_server/test_routes_login_sso.py | 41 ++- tests/test_litellm/proxy/test_proxy_server.py | 7 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 7 +- 10 files changed, 581 insertions(+), 120 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 1fd2781bd95..0071ad4603d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2528,7 +2528,12 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): max_failed_login_attempts: int | None = Field( None, ge=1, - description="Number of failed Admin UI sign-in attempts allowed for one username from one source address within `failed_login_window_seconds`, before further attempts are refused with 429. Configurable from config.yaml only. Defaults to 10", + description="Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Configurable from config.yaml only. Defaults to 50", + ) + max_failed_login_attempts_per_source: int | None = Field( + None, + ge=1, + description="Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Configurable from config.yaml only. Defaults to 250", ) failed_login_window_seconds: int | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index e2befd15a35..ed53d9e7bc7 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -1,16 +1,20 @@ """Failed-login accounting for the Admin UI sign-in path. -Counts failed credential checks per (username, source address) over a fixed window and -denies further attempts with 429 once the count reaches the limit. Built per request by -``LoginThrottle.from_request`` because it carries that request's resolved source address, -and because the coordination cache is assigned at startup and can be reassigned later. +Counts failed credential checks over a fixed window against two independent keys, the +username on its own and the source address on its own, so that one username attacked from +many sources and one source spraying many usernames are both counted. Repeated failures +are answered slowly, doubling from one second, and refused with 429 once either counter +reaches its limit. Built per request by ``LoginThrottle.from_request`` because it carries +that request's resolved source address, and because the coordination cache is assigned at +startup and can be reassigned later. """ +import asyncio import hashlib from collections.abc import Awaitable from dataclasses import dataclass from types import MappingProxyType -from typing import Final +from typing import Final, NamedTuple, NoReturn from fastapi import Request @@ -23,21 +27,53 @@ from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges from litellm.proxy.auth.trusted_proxy_utils import TRUSTED_PROXY_RANGES_KEY from litellm.secret_managers.main import get_secret_bool -DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS: Final = 10 +DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS: Final = 50 +DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 250 DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 900 +USERNAME_DELAY_ONSET: Final = 3 +SOURCE_DELAY_ONSET: Final = 25 +FIRST_DELAY_SECONDS: Final = 1.0 +MAX_DELAY_SECONDS: Final = 30.0 +MAX_CONCURRENT_DELAYS_PER_SOURCE: Final = 5 + +_MAX_DELAY_DOUBLINGS: Final = 16 + _CACHE_KEY_PREFIX: Final = "login_fail" _UNKNOWN_SOURCE: Final = "unknown" _MAX_LOGGED_USERNAME_CHARS: Final = 128 +_MAX_TRACKED_LOGIN_USERNAMES: Final = 10_000 _MAX_TRACKED_LOGIN_SOURCES: Final = 10_000 -_FAILED_LOGIN_CACHE: Final = DualCache( - in_memory_cache=InMemoryCache(max_size_in_memory=_MAX_TRACKED_LOGIN_SOURCES), - default_in_memory_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, -) + +def _bounded_store(max_entries: int) -> DualCache: + return DualCache( + in_memory_cache=InMemoryCache(max_size_in_memory=max_entries), + default_in_memory_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, + ) + + +# Separate stores: eviction is earliest-expiring-first, so in one shared store a spray of +# fresh usernames would evict the source counter that is meant to stop that same spray. +_FAILED_LOGIN_USERNAME_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_USERNAMES) +_FAILED_LOGIN_SOURCE_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_SOURCES) _NO_SETTINGS: Final = MappingProxyType({}) +_DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} + + +async def _sleep(seconds: float) -> None: + """The wait a rejected sign-in is held for. Replaced in tests so the suite pays no wall clock.""" + await asyncio.sleep(seconds) + + +class FailureCounts(NamedTuple): + """Failures recorded so far in this window against each of the two keys.""" + + username: int + source: int + def _int_setting(name: str, value: object, default: int, minimum: int) -> int: if value is None: @@ -56,12 +92,14 @@ def _as_count(cached: object) -> int: @dataclass(frozen=True, slots=True) class LoginThrottle: - """Fixed-window failed-login accounting for one request's source address.""" + """Fixed-window failed-login accounting for one request's username and source address.""" client_ip: str max_attempts: int + max_attempts_per_source: int window_seconds: int - cache: DualCache + username_cache: DualCache + source_cache: DualCache redis_cache: RedisCache | None = None enabled: bool = True @@ -85,13 +123,20 @@ class LoginThrottle: DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS, 1, ), + max_attempts_per_source=_int_setting( + "max_failed_login_attempts_per_source", + settings.get("max_failed_login_attempts_per_source"), + DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE, + 1, + ), window_seconds=_int_setting( "failed_login_window_seconds", settings.get("failed_login_window_seconds"), DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, 1, ), - cache=_FAILED_LOGIN_CACHE, + username_cache=_FAILED_LOGIN_USERNAME_CACHE, + source_cache=_FAILED_LOGIN_SOURCE_CACHE, redis_cache=redis_usage_cache, enabled=not get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT"), ) @@ -101,9 +146,13 @@ class LoginThrottle: """The username with anything that could forge a log line removed.""" return "".join(c for c in username if c.isprintable())[:_MAX_LOGGED_USERNAME_CHARS] - def _key(self, username: str) -> str: + @staticmethod + def _username_key(username: str) -> str: identity: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest() - return f"{_CACHE_KEY_PREFIX}:{identity}:{self.client_ip}" + return f"{_CACHE_KEY_PREFIX}:user:{identity}" + + def _source_key(self) -> str: + return f"{_CACHE_KEY_PREFIX}:source:{self.client_ip}" async def _outcome(self, work: Awaitable[object]) -> object: try: @@ -112,66 +161,147 @@ class LoginThrottle: verbose_proxy_logger.warning("login attempt accounting unavailable: %s", exc) return None - async def _failures(self, key: str) -> int: - store: Final = self.cache if self.redis_cache is None else self.redis_cache - return _as_count(await self._outcome(store.async_get_cache(key=key))) + async def _failures(self, store: DualCache, key: str) -> int: + """The larger of the shared and the process-local count, so a Redis outage degrades + to per-worker accounting instead of switching the control off.""" + local: Final = _as_count(await self._outcome(store.async_get_cache(key=key))) + redis_cache: Final = self.redis_cache + if redis_cache is None: + return local + return max(local, _as_count(await self._outcome(redis_cache.async_get_cache(key)))) - async def _ensure_expiry(self, key: str) -> None: - """Give the counter an expiry if it somehow has none. + async def _remaining_window(self, key: str) -> int: + """Seconds until this counter expires, repairing a counter left without an expiry. Redis commits the increment before setting the TTL, so a failure in between can leave a counter that never expires. Nothing increments the key again once the - limit is reached, so without this the pair would stay refused indefinitely. + limit is reached, so without the repair the key would stay refused indefinitely. """ redis_cache: Final = self.redis_cache if redis_cache is None: - return - if isinstance(await self._outcome(redis_cache.async_get_ttl(key)), int): - return + return self.window_seconds + ttl: Final = await self._outcome(redis_cache.async_get_ttl(key)) + if isinstance(ttl, int) and ttl > 0: + return min(ttl, self.window_seconds) await self._outcome(redis_cache.async_increment(key, 0, ttl=self.window_seconds)) + return self.window_seconds - async def raise_if_blocked(self, username: str) -> None: - """Deny before the database lookup and before the password comparison.""" - if not self.enabled: - return - key: Final = self._key(username) - if await self._failures(key) < self.max_attempts: - return - await self._ensure_expiry(key) - verbose_proxy_logger.warning( - "Admin UI sign-in attempts exhausted for username=%s source=%s; %s attempts in %ss, retry after %ss", - self._loggable(username), - self.client_ip, - self.max_attempts, - self.window_seconds, - self.window_seconds, - ) - raise ProxyException( + def _refused(self, retry_after: int, param: str) -> ProxyException: + return ProxyException( message="Too many failed sign-in attempts. Try again later.", type=ProxyErrorTypes.auth_error, - param="max_failed_login_attempts", + param=param, code=429, - headers={"Retry-After": str(self.window_seconds)}, # mutable-ok: ProxyException coerces header values + headers={"Retry-After": str(retry_after)}, # mutable-ok: ProxyException coerces header values ) - async def record_failure(self, username: str) -> None: - """Count one rejected credential guess against this username and source.""" + async def _refuse(self, key: str, scope: str, param: str, username: str, failures: int, limit: int) -> NoReturn: + retry_after: Final = await self._remaining_window(key) + verbose_proxy_logger.warning( + "Admin UI sign-in attempts exhausted for %s; username=%s source=%s failures=%s limit=%s window=%ss", + scope, + self._loggable(username), + self.client_ip, + failures, + limit, + self.window_seconds, + ) + raise self._refused(retry_after, param) + + async def raise_if_blocked(self, username: str) -> None: + """Refuse before the database lookup and before the invite-link password hash.""" if not self.enabled: return - key: Final = self._key(username) + username_key: Final = self._username_key(username) + source_key: Final = self._source_key() + username_failures: Final = await self._failures(self.username_cache, username_key) + if username_failures >= self.max_attempts: + await self._refuse( + username_key, "username", "max_failed_login_attempts", username, username_failures, self.max_attempts + ) + source_failures: Final = await self._failures(self.source_cache, source_key) + if source_failures >= self.max_attempts_per_source: + await self._refuse( + source_key, + "source address", + "max_failed_login_attempts_per_source", + username, + source_failures, + self.max_attempts_per_source, + ) + + async def _bump(self, store: DualCache, key: str) -> int: + local: Final = _as_count( + await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) + ) redis_cache: Final = self.redis_cache - if redis_cache is not None: - await self._outcome(redis_cache.async_increment(key, 1, ttl=self.window_seconds)) - await self._ensure_expiry(key) + if redis_cache is None: + return local + shared: Final = _as_count(await self._outcome(redis_cache.async_increment(key, 1, ttl=self.window_seconds))) + await self._remaining_window(key) + return max(local, shared) + + async def record_failure(self, username: str) -> FailureCounts: + """Count one rejected credential guess against this username and against this source.""" + if not self.enabled: + return FailureCounts(username=0, source=0) + return FailureCounts( + username=await self._bump(self.username_cache, self._username_key(username)), + source=await self._bump(self.source_cache, self._source_key()), + ) + + @staticmethod + def delay_seconds(counts: FailureCounts) -> float: + """Seconds to hold a rejected attempt for, doubling per failure past whichever onset is further along.""" + steps: Final = min( + max(counts.username - USERNAME_DELAY_ONSET, counts.source - SOURCE_DELAY_ONSET), + _MAX_DELAY_DOUBLINGS, + ) + if steps < 0: + return 0.0 + return min(FIRST_DELAY_SECONDS * float(2**steps), MAX_DELAY_SECONDS) + + async def delay_for(self, username: str, counts: FailureCounts) -> None: + """Hold this rejected attempt open before answering it, so guessing costs wall-clock time. + + Only ever reached once the credentials are known to be wrong, so a valid password is + never delayed. Sources are capped at ``MAX_CONCURRENT_DELAYS_PER_SOURCE`` held + connections; over that, the attempt is refused immediately instead of parking a socket. + """ + if not self.enabled: return - await self._outcome(self.cache.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) + delay: Final = self.delay_seconds(counts) + if delay <= 0: + return + in_flight: Final = _DELAYS_IN_FLIGHT.get(self.client_ip, 0) + if in_flight >= MAX_CONCURRENT_DELAYS_PER_SOURCE: + verbose_proxy_logger.warning( + "Admin UI sign-in attempts held concurrently exhausted; username=%s source=%s in_flight=%s", + self._loggable(username), + self.client_ip, + in_flight, + ) + raise self._refused(int(MAX_DELAY_SECONDS), "concurrent_failed_logins") + _DELAYS_IN_FLIGHT[self.client_ip] = in_flight + 1 + try: + await _sleep(delay) + finally: + remaining: Final = _DELAYS_IN_FLIGHT.get(self.client_ip, 1) - 1 + if remaining > 0: + _DELAYS_IN_FLIGHT[self.client_ip] = remaining + else: + _DELAYS_IN_FLIGHT.pop(self.client_ip, None) async def clear(self, username: str) -> None: - """Drop the bucket after a successful sign-in.""" + """Drop the username counter after a successful sign-in. + + The source counter is left alone. It is shared by every account behind that address, + so one success there says nothing about the other attempts it is counting. + """ if not self.enabled: return - key: Final = self._key(username) + key: Final = self._username_key(username) redis_cache: Final = self.redis_cache if redis_cache is not None: await self._outcome(redis_cache.async_delete_cache(key)) - await self._outcome(self.cache.async_delete_cache(key=key)) + await self._outcome(self.username_cache.async_delete_cache(key=key)) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 98780a4692a..8835227cd1c 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -145,10 +145,15 @@ async def authenticate_user( code=500, ) - await throttle.raise_if_blocked(username) - ui_username, ui_password = get_ui_credentials(master_key) + admin_credentials_match: Final = secrets.compare_digest( + username.encode("utf-8"), ui_username.encode("utf-8") + ) and secrets.compare_digest(password.encode("utf-8"), ui_password.encode("utf-8")) + + if not admin_credentials_match: + await throttle.raise_if_blocked(username) + # Check if we can find the `username` in the db. On the UI, users can enter username=their email _user_row: LiteLLM_UserTable | None = None user_role: ( @@ -174,9 +179,7 @@ async def authenticate_user( - Login with UI_USERNAME and UI_PASSWORD - Login with Invite Link `user_email` and `password` combination """ - if secrets.compare_digest(username.encode("utf-8"), ui_username.encode("utf-8")) and secrets.compare_digest( - password.encode("utf-8"), ui_password.encode("utf-8") - ): + if admin_credentials_match: # Non SSO -> If user is using UI_USERNAME and UI_PASSWORD they are Proxy admin user_role = LitellmUserRoles.PROXY_ADMIN user_id = LITELLM_PROXY_ADMIN_NAME @@ -314,7 +317,7 @@ async def authenticate_user( login_method="username_password", ) else: - await throttle.record_failure(username) + await throttle.delay_for(username, await throttle.record_failure(username)) raise ProxyException( message=INVALID_UI_CREDENTIALS_MESSAGE, type=ProxyErrorTypes.auth_error, @@ -322,7 +325,7 @@ async def authenticate_user( code=401, ) else: - await throttle.record_failure(username) + await throttle.delay_for(username, await throttle.record_failure(username)) raise ProxyException( message=INVALID_UI_CREDENTIALS_MESSAGE, type=ProxyErrorTypes.auth_error, diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 0449802abae..a2e3a649586 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1358,6 +1358,10 @@ def run_server( # DO NOT DELETE - enables global variables to work across files from litellm.proxy.proxy_server import app + # Write the resolved --num_workers back to its env var so worker processes can read + # the fleet size at startup (the failed-login accounting warning keys off it) + os.environ["NUM_WORKERS"] = str(num_workers) + # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( num_workers=num_workers, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 39a419df73f..c6a1dc4cdb7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2264,6 +2264,7 @@ user_custom_key_generate = None # Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'. _pkce_no_redis_warning_emitted: bool = False _cp_no_redis_warning_emitted: bool = False +_login_throttle_no_redis_warning_emitted: bool = False user_custom_key_update = None user_custom_sso = None user_custom_ui_sso_sign_in_handler = None @@ -5317,6 +5318,21 @@ class ProxyConfig: "or ensure sticky sessions for single-instance deployments." ) + ### FAILED-LOGIN ACCOUNTING MULTI-INSTANCE PREREQUISITE CHECK ### + # Failed Admin UI sign-in counters live in redis_usage_cache when available so a + # brute-force run is counted once across workers instead of once per worker. + if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: + global _login_throttle_no_redis_warning_emitted + if not _login_throttle_no_redis_warning_emitted: + _login_throttle_no_redis_warning_emitted = True + verbose_proxy_logger.warning( + "Running %s workers but Redis is not configured for LiteLLM caching. " + "Failed Admin UI sign-in attempts are counted per worker, so an attacker " + "gets max_failed_login_attempts guesses per worker instead of overall. " + "Configure Redis via the 'cache' section in your proxy config.", + os.getenv("NUM_WORKERS", "1"), + ) + ### STORE MODEL IN DB ### feature flag for `/model/new` store_model_in_db = general_settings.get("store_model_in_db", False) if store_model_in_db is None: @@ -14756,13 +14772,26 @@ async def login(request: Request): password: Final = str(form.get("password")) # Authenticate user and get login result - login_result: Final = await authenticate_user( - username=username, - password=password, - master_key=master_key, - prisma_client=prisma_client, - throttle=LoginThrottle.from_request(request), - ) + try: + login_result: Final = await authenticate_user( + username=username, + password=password, + master_key=master_key, + prisma_client=prisma_client, + throttle=LoginThrottle.from_request(request), + ) + except ProxyException as exc: + if int(exc.code) != status.HTTP_429_TOO_MANY_REQUESTS: + raise + retry_after: Final = exc.headers.get("Retry-After", "30") + return HTMLResponse( + content=( + "

Too many sign-in attempts

" + f"

Try again in about {retry_after} seconds

" + ), + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + headers=exc.headers, + ) # Create UI token object returned_ui_token_object: Final = create_ui_token_object( diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 1e3fcaf5d78..1b85c9f4046 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -6,17 +6,48 @@ to login_utils.py for better reusability. """ import os +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +class _RecordedSleeps: + """A sleep that records what it was asked to wait for instead of waiting.""" + + def __init__(self): + self.seconds: list[float] = [] + + async def __call__(self, seconds: float) -> None: + self.seconds.append(seconds) + + +@pytest.fixture(autouse=True) +def login_delays(monkeypatch): + """Replace the failed-login wait, so the suite pays no wall clock and can read it back.""" + from litellm.proxy.auth import login_throttle + + recorded = _RecordedSleeps() + monkeypatch.setattr(login_throttle, "_sleep", recorded) + login_throttle._DELAYS_IN_FLIGHT.clear() + yield recorded + login_throttle._DELAYS_IN_FLIGHT.clear() + + def _unlimited_throttle(): """A throttle wired to a real in-memory store with a limit no test can reach.""" from litellm.caching.dual_cache import DualCache from litellm.proxy.auth.login_throttle import LoginThrottle - return LoginThrottle(client_ip="1.2.3.4", max_attempts=10_000, window_seconds=900, cache=DualCache()) + store: Final = DualCache() + return LoginThrottle( + client_ip="1.2.3.4", + max_attempts=10_000, + max_attempts_per_source=10_000, + window_seconds=900, + username_cache=store, + source_cache=store, + ) @@ -628,16 +659,26 @@ class TestEncodeUiSessionJwt: # --------------------------------------------------------------------------- -def _throttle(max_attempts: int = 3, window_seconds: int = 900, client_ip: str = "1.2.3.4", cache=None, redis_cache=None): +def _throttle( + max_attempts: int = 3, + window_seconds: int = 900, + client_ip: str = "1.2.3.4", + cache=None, + redis_cache=None, + max_attempts_per_source: int = 10_000, +): """A throttle over a real in-memory store, so the tests exercise the true counters.""" from litellm.caching.dual_cache import DualCache from litellm.proxy.auth.login_throttle import LoginThrottle + store: Final = cache if cache is not None else DualCache() return LoginThrottle( client_ip=client_ip, max_attempts=max_attempts, + max_attempts_per_source=max_attempts_per_source, window_seconds=window_seconds, - cache=cache if cache is not None else DualCache(), + username_cache=store, + source_cache=store, redis_cache=redis_cache, ) @@ -675,21 +716,32 @@ async def test_attempts_are_refused_once_the_limit_is_reached(monkeypatch): @pytest.mark.asyncio -async def test_a_correct_password_is_refused_while_blocked(monkeypatch): - """The check precedes the credential comparison, so being over the limit wins.""" +async def test_a_correct_admin_password_is_accepted_while_blocked(monkeypatch): + """The configured admin credentials are compared before the gate, so the operator gets in. + + A throttle that refuses a valid password hands anyone who can reach the login form a + denial of service against the one account that can fix it. + """ from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") throttle = _throttle(max_attempts=2) for _ in range(2): with pytest.raises(ProxyException): await _guess(throttle) - with pytest.raises(ProxyException) as blocked: - await _guess(throttle, password="right") - assert blocked.value.code == "429" + with pytest.raises(ProxyException) as still_blocked: + await _guess(throttle) + assert still_blocked.value.code == "429", "a wrong password is still refused" + + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ): + result = await _guess(throttle, password="right") + assert result.key == "sk-ui" @pytest.mark.asyncio @@ -700,18 +752,18 @@ async def test_a_blocked_attempt_does_not_extend_the_window(monkeypatch): monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") throttle = _throttle(max_attempts=2) - key = throttle._key("admin") + key = throttle._username_key("admin") for _ in range(2): with pytest.raises(ProxyException): await _guess(throttle) - counted_at_limit = await throttle._failures(key) + counted_at_limit = await throttle._failures(throttle.username_cache, key) for _ in range(5): with pytest.raises(ProxyException): await _guess(throttle) - assert await throttle._failures(key) == counted_at_limit == 2 + assert await throttle._failures(throttle.username_cache, key) == counted_at_limit == 2 @pytest.mark.asyncio @@ -733,7 +785,7 @@ async def test_a_successful_sign_in_clears_the_bucket(monkeypatch): ): await _guess(throttle, password="right") - assert await throttle._failures(throttle._key("admin")) == 0 + assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 @pytest.mark.asyncio @@ -749,7 +801,7 @@ async def test_a_configuration_error_never_counts(monkeypatch): ) assert exc.value.code == "500" - assert await throttle._failures(throttle._key("admin")) == 0 + assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 @pytest.mark.asyncio @@ -773,7 +825,7 @@ async def test_the_username_is_case_folded_into_one_bucket(monkeypatch): @pytest.mark.asyncio async def test_a_different_username_from_the_same_source_is_unaffected(monkeypatch): - """The bucket is the pair, so one username's failures do not block another.""" + """The counters are independent, so one username's failures do not exhaust another's.""" from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") @@ -853,7 +905,7 @@ async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypat ) assert exc.value.code == "401" - assert await throttle._failures(throttle._key("nopass@example.com")) == 0 + assert await throttle._failures(throttle.username_cache, throttle._username_key("nopass@example.com")) == 0 @pytest.mark.asyncio @@ -896,11 +948,40 @@ async def test_a_wrong_password_for_a_known_user_also_counts(monkeypatch): @pytest.mark.asyncio -async def test_two_source_addresses_do_not_share_a_bucket(monkeypatch): - """The key is the pair, so one address exhausting its budget must not block another. +async def test_one_source_exhausting_its_own_budget_does_not_refuse_another_source(monkeypatch): + """The source counter is per address, so a noisy office does not take its neighbour down.""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import ProxyException - Dropping the address from the key would turn this into the username-only counter the - design rejects, where anyone can lock a named admin out from anywhere. + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + shared_store = DualCache() + attacker = _throttle( + max_attempts=10_000, max_attempts_per_source=2, client_ip="203.0.113.9", cache=shared_store + ) + operator = _throttle( + max_attempts=10_000, max_attempts_per_source=2, client_ip="198.51.100.7", cache=shared_store + ) + + for i in range(2): + with pytest.raises(ProxyException): + await _guess(attacker, username=f"target-{i}@corp.com") + + with pytest.raises(ProxyException) as blocked: + await _guess(attacker, username="target-2@corp.com") + assert blocked.value.code == "429" + + with pytest.raises(ProxyException) as unaffected: + await _guess(operator, username="target-3@corp.com") + assert unaffected.value.code == "401", "the other address must still reach the credential check" + + +@pytest.mark.asyncio +async def test_a_username_exhausted_from_one_source_is_refused_from_another(monkeypatch): + """The username counter carries no address, so spreading the guesses buys nothing. + + The pair key this replaced reset the budget for every new address, which is exactly the + shape of a credential-stuffing run from a proxy pool. """ from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ProxyException @@ -908,20 +989,154 @@ async def test_two_source_addresses_do_not_share_a_bucket(monkeypatch): monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") shared_store = DualCache() - attacker = _throttle(max_attempts=2, client_ip="203.0.113.9", cache=shared_store) - operator = _throttle(max_attempts=2, client_ip="198.51.100.7", cache=shared_store) + first_hop = _throttle(max_attempts=2, client_ip="203.0.113.9", cache=shared_store) + second_hop = _throttle(max_attempts=2, client_ip="198.51.100.7", cache=shared_store) - for _ in range(3): + for _ in range(2): with pytest.raises(ProxyException): - await _guess(attacker, username="admin") + await _guess(first_hop, username="victim@corp.com") + + with pytest.raises(ProxyException) as rotated: + await _guess(second_hop, username="victim@corp.com") + assert rotated.value.code == "429" + + +@pytest.mark.asyncio +async def test_a_source_wide_spray_is_counted_even_though_each_username_is_fresh(monkeypatch): + """One guess against each of many usernames never trips a username counter, only the source one.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=10_000, max_attempts_per_source=6, client_ip="203.0.113.11") + + for i in range(6): + with pytest.raises(ProxyException) as rejected: + await _guess(throttle, username=f"sprayed-{i}@corp.com") + assert rejected.value.code == "401" with pytest.raises(ProxyException) as blocked: - await _guess(attacker, username="admin") + await _guess(throttle, username="sprayed-7@corp.com") assert blocked.value.code == "429" + assert await throttle._failures(throttle.username_cache, throttle._username_key("sprayed-7@corp.com")) == 0 - with pytest.raises(ProxyException) as unaffected: - await _guess(operator, username="admin") - assert unaffected.value.code == "401", "the real operator must still reach the credential check" + +@pytest.mark.asyncio +async def test_a_successful_sign_in_leaves_the_source_counter_alone(monkeypatch): + """One account's success says nothing about the other attempts the address is making.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(max_attempts=10) + + for _ in range(2): + with pytest.raises(ProxyException): + await _guess(throttle) + + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ): + await _guess(throttle, password="right") + + assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 + assert await throttle._failures(throttle.source_cache, throttle._source_key()) == 2 + + +@pytest.mark.asyncio +async def test_the_delay_doubles_from_one_second_and_is_capped(monkeypatch, login_delays): + """Guessing has to cost wall clock, and the cost has to stop short of an unbounded hang.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.login_throttle import MAX_DELAY_SECONDS + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=10_000) + + for _ in range(9): + with pytest.raises(ProxyException): + await _guess(throttle) + + assert login_delays.seconds == [1.0, 2.0, 4.0, 8.0, 16.0, 30.0, 30.0], ( + "the first two failures answer immediately, then the wait doubles up to the cap" + ) + assert max(login_delays.seconds) == MAX_DELAY_SECONDS + + +@pytest.mark.asyncio +async def test_the_delay_tracks_whichever_counter_is_further_past_its_onset(monkeypatch): + """A source deep into a spray must not be answered instantly just because the username is fresh.""" + from litellm.proxy.auth.login_throttle import FailureCounts, LoginThrottle + + assert LoginThrottle.delay_seconds(FailureCounts(username=1, source=1)) == 0.0 + assert LoginThrottle.delay_seconds(FailureCounts(username=2, source=24)) == 0.0 + assert LoginThrottle.delay_seconds(FailureCounts(username=3, source=1)) == 1.0 + assert LoginThrottle.delay_seconds(FailureCounts(username=1, source=25)) == 1.0 + assert LoginThrottle.delay_seconds(FailureCounts(username=4, source=28)) == 8.0 + + +@pytest.mark.asyncio +async def test_held_attempts_from_one_source_are_capped(monkeypatch): + """Holding a rejected attempt open must not let one address park unlimited sockets.""" + import asyncio + + from litellm.proxy._types import ProxyException + from litellm.proxy.auth import login_throttle as lt + from litellm.proxy.auth.login_throttle import MAX_CONCURRENT_DELAYS_PER_SOURCE + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + release = asyncio.Event() + + async def _park(_seconds: float) -> None: + await release.wait() + + monkeypatch.setattr(lt, "_sleep", _park) + throttle = _throttle(max_attempts=10_000, client_ip="203.0.113.44") + await throttle.record_failure("admin") + await throttle.record_failure("admin") + + held = [asyncio.create_task(_guess(throttle)) for _ in range(MAX_CONCURRENT_DELAYS_PER_SOURCE)] + for _ in range(1000): + if lt._DELAYS_IN_FLIGHT.get("203.0.113.44") == MAX_CONCURRENT_DELAYS_PER_SOURCE: + break + await asyncio.sleep(0) + assert lt._DELAYS_IN_FLIGHT.get("203.0.113.44") == MAX_CONCURRENT_DELAYS_PER_SOURCE + + try: + with pytest.raises(ProxyException) as over_cap: + await _guess(throttle) + assert over_cap.value.code == "429" + assert over_cap.value.headers.get("Retry-After") == "30" + finally: + release.set() + for task in held: + with pytest.raises(ProxyException): + await task + + with pytest.raises(ProxyException) as after_drain: + await _guess(throttle) + assert after_drain.value.code == "401", "the cap must release once the held attempts answer" + + +@pytest.mark.asyncio +async def test_disabling_the_control_removes_the_delay_as_well(monkeypatch, login_delays): + """The escape hatch has to turn off the whole control, not only the refusal.""" + import dataclasses + + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = dataclasses.replace(_throttle(max_attempts=2), enabled=False) + + for _ in range(6): + with pytest.raises(ProxyException) as rejected: + await _guess(throttle) + assert rejected.value.code == "401" + + assert login_delays.seconds == [] class _NoExpiryRedis: @@ -1003,7 +1218,7 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): throttle = LoginThrottle.from_request(request) for i in range(25): - with pytest.raises(Exception): + with pytest.raises(ProxyException, match="Invalid credentials"): await _guess(throttle, username=f"made-up-{i}@example.com") added = set(ps.user_api_key_cache.in_memory_cache.cache_dict) - auth_cache_keys_before @@ -1055,14 +1270,29 @@ async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): back a fresh allowance against the real account. """ from litellm.proxy._types import ProxyException - from litellm.proxy.auth.login_throttle import _FAILED_LOGIN_CACHE, _MAX_TRACKED_LOGIN_SOURCES + from litellm.proxy.auth.login_throttle import ( + LoginThrottle, + _FAILED_LOGIN_SOURCE_CACHE, + _FAILED_LOGIN_USERNAME_CACHE, + _MAX_TRACKED_LOGIN_SOURCES, + _MAX_TRACKED_LOGIN_USERNAMES, + ) monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") assert _MAX_TRACKED_LOGIN_SOURCES >= 10_000 - assert _FAILED_LOGIN_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_SOURCES + assert _MAX_TRACKED_LOGIN_USERNAMES >= 10_000 + assert _FAILED_LOGIN_SOURCE_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_SOURCES + assert _FAILED_LOGIN_USERNAME_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_USERNAMES - throttle = _throttle(max_attempts=3, cache=_FAILED_LOGIN_CACHE, client_ip="10.9.9.9") + throttle = LoginThrottle( + client_ip="10.9.9.9", + max_attempts=3, + max_attempts_per_source=10_000, + window_seconds=900, + username_cache=_FAILED_LOGIN_USERNAME_CACHE, + source_cache=_FAILED_LOGIN_SOURCE_CACHE, + ) victim = "spray-victim@corp.com" for _ in range(3): with pytest.raises(ProxyException): @@ -1071,7 +1301,7 @@ async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): for i in range(500): await throttle.record_failure(f"spray-filler-{i}@corp.com") - assert await throttle._failures(throttle._key(victim)) == 3, "the counter must survive a spray" + assert await throttle._failures(throttle.username_cache, throttle._username_key(victim)) == 3, "the counter must survive a spray" with pytest.raises(ProxyException) as blocked: await _guess(throttle, username=victim) assert blocked.value.code == "429" diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index 22f2e5b3feb..56aa87f1e50 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -517,27 +517,38 @@ def make_key( def reset_login_throttle(monkeypatch): """Clear the Admin UI failed-login counters between tests. - `client` is session scoped and the counters live in the shared `user_api_key_cache` - with a 900s window, so without this any test that fails a sign-in enough times would - start returning 429 from unrelated tests later in the same process. Only the throttle's - own keys are removed, so nothing else in that cache is disturbed. + `client` is session scoped and the counters live in shared module stores with a 900s + window, so without this a failed sign-in test could return 429 in unrelated tests later. + Only the throttle's own keys are removed, so other cache entries remain untouched. """ from litellm.proxy import proxy_server as ps - from litellm.proxy.auth.login_throttle import _FAILED_LOGIN_CACHE, _CACHE_KEY_PREFIX + from litellm.proxy.auth import login_throttle + from litellm.proxy.auth.login_throttle import ( + _CACHE_KEY_PREFIX, + _FAILED_LOGIN_SOURCE_CACHE, + _FAILED_LOGIN_USERNAME_CACHE, + ) + + async def _no_delay(_seconds: float) -> None: + """The escalating wait on a rejected sign-in, replaced so the route tests stay fast.""" + + monkeypatch.setattr(login_throttle, "_sleep", _no_delay) def _drop_throttle_keys() -> None: - in_memory = getattr(_FAILED_LOGIN_CACHE, "in_memory_cache", None) - if in_memory is None: - return - tracked = tuple( - key - for store in (getattr(in_memory, "cache_dict", None), getattr(in_memory, "ttl_dict", None)) - if isinstance(store, dict) - for key in tuple(store) - if str(key).startswith(_CACHE_KEY_PREFIX) - ) - for key in tracked: - in_memory.delete_cache(key) + login_throttle._DELAYS_IN_FLIGHT.clear() + for cache in (_FAILED_LOGIN_USERNAME_CACHE, _FAILED_LOGIN_SOURCE_CACHE): + in_memory = getattr(cache, "in_memory_cache", None) + if in_memory is None: + continue + tracked = tuple( + key + for store in (getattr(in_memory, "cache_dict", None), getattr(in_memory, "ttl_dict", None)) + if isinstance(store, dict) + for key in tuple(store) + if str(key).startswith(_CACHE_KEY_PREFIX) + ) + for key in tracked: + in_memory.delete_cache(key) monkeypatch.setattr(ps, "redis_usage_cache", None) _drop_throttle_keys() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index faeb7750e03..c8185cdb886 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -531,8 +531,21 @@ def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_ assert refused.headers.get("retry-after") == "77" +def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, reset_login_throttle): + """The no-JavaScript form must render a wait page when its POST is throttled.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=2, failed_login_window_seconds=77) + + assert [_form_login(client) for _ in range(2)] == [401, 401] + + refused = client.post("/login", data={"username": "admin", "password": "wrong"}) + assert refused.status_code == 429 + assert refused.headers.get("content-type", "").startswith("text/html") + assert "Try again in about 77 seconds" in refused.text + assert refused.headers.get("retry-after") == "77" + + def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle): - """The bucket is the username and source pair, so one account cannot block another.""" + """The username counter carries no address, so one account exhausting it cannot block another.""" _install_real_auth(monkeypatch, max_failed_login_attempts=2) for _ in range(3): @@ -541,6 +554,32 @@ def test_a_second_username_from_the_same_source_still_gets_through(client, monke assert _json_login(client, "/v2/login", username="someone-else@example.com") == 401 +def test_a_spray_across_usernames_is_refused_on_the_source_counter(client, monkeypatch, reset_login_throttle): + """A fresh username per guess keeps every username counter at one, so the address is what stops it.""" + _install_real_auth(monkeypatch, max_failed_login_attempts=100, max_failed_login_attempts_per_source=4) + + sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(4)] + assert sprayed == [401] * 4 + + assert _json_login(client, "/v2/login", username="sprayed-5@corp.com") == 429 + + +def test_the_configured_admin_password_still_signs_in_while_refused(client, monkeypatch, reset_login_throttle): + """The operator must never be locked out of the console by traffic aimed at it.""" + from unittest.mock import AsyncMock, patch + + _install_real_auth(monkeypatch, max_failed_login_attempts=2) + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + + assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] + assert _json_login(client, "/v2/login") == 429 + + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ): + assert _json_login(client, "/v2/login", password="right-password") == 200 + + def test_sign_in_succeeds_again_once_the_budget_is_restored(client, monkeypatch, reset_login_throttle): """A cleared bucket lets the same username straight back in.""" _install_real_auth(monkeypatch, max_failed_login_attempts=2) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index ac2e84b8474..17563f60a4f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11442,9 +11442,14 @@ async def test_login_throttle_settings_are_not_overridable_from_the_database(): try: ps.general_settings.clear() await ProxyConfig()._update_general_settings( - db_general_settings={"max_failed_login_attempts": 999, "failed_login_window_seconds": 1} + db_general_settings={ + "max_failed_login_attempts": 999, + "max_failed_login_attempts_per_source": 999, + "failed_login_window_seconds": 1, + } ) assert "max_failed_login_attempts" not in ps.general_settings + assert "max_failed_login_attempts_per_source" not in ps.general_settings assert "failed_login_window_seconds" not in ps.general_settings finally: ps.general_settings.clear() diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index bb2841b6388..9d66320041f 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24201,9 +24201,14 @@ export interface components { max_batch_file_size_mb?: number | null; /** * Max Failed Login Attempts - * @description Number of failed Admin UI sign-in attempts allowed for one username from one source address within `failed_login_window_seconds`, before further attempts are refused with 429. Configurable from config.yaml only. Defaults to 10 + * @description Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Configurable from config.yaml only. Defaults to 50 */ max_failed_login_attempts?: number | null; + /** + * Max Failed Login Attempts Per Source + * @description Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Configurable from config.yaml only. Defaults to 250 + */ + max_failed_login_attempts_per_source?: number | null; /** * Max Parallel Requests * @description maximum parallel requests for each api key From fb7ca1eda16ef666035b22b6c003bfb11f9ed6d4 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 13:21:44 -0700 Subject: [PATCH 005/253] docs(proxy): clarify login throttle configuration --- litellm/proxy/_types.py | 6 +++--- tests/test_litellm/proxy/test_proxy_server.py | 10 +++++----- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 +++--- 3 files changed, 11 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0071ad4603d..756a878f7f8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2528,17 +2528,17 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): max_failed_login_attempts: int | None = Field( None, ge=1, - description="Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Configurable from config.yaml only. Defaults to 50", + description="Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Set under `general_settings` in config.yaml. Defaults to 50", ) max_failed_login_attempts_per_source: int | None = Field( None, ge=1, - description="Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Configurable from config.yaml only. Defaults to 250", + description="Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Set under `general_settings` in config.yaml. Defaults to 250", ) failed_login_window_seconds: int | None = Field( None, ge=1, - description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Configurable from config.yaml only. Defaults to 900", + description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 900", ) allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access") reject_clientside_metadata_tags: bool | None = Field( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 17563f60a4f..a39fec7f523 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11428,12 +11428,12 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the @pytest.mark.asyncio -async def test_login_throttle_settings_are_not_overridable_from_the_database(): - """LIT-5285: the sign-in limits stay config.yaml only. +async def test_login_throttle_settings_are_not_hot_applied_from_the_database(): + """LIT-5285: a stored sign-in limit does not take effect on a live worker. - _update_general_settings copies an allowlist of keys out of the DB row. Adding these - to it would let a stored value outrank config.yaml, so an operator refused by a bad - value could not fix it by editing YAML and restarting. + _update_general_settings copies an allowlist of keys out of the DB row on every config + poll. Adding these to it would let a stored value outrank config.yaml without a restart, + so an operator locked out by a bad value could not fix it by editing YAML and restarting. """ import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ProxyConfig diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 9d66320041f..62ba7ce8b95 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24152,7 +24152,7 @@ export interface components { enable_public_model_hub: boolean; /** * Failed Login Window Seconds - * @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Configurable from config.yaml only. Defaults to 900 + * @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 900 */ failed_login_window_seconds?: number | null; /** @@ -24201,12 +24201,12 @@ export interface components { max_batch_file_size_mb?: number | null; /** * Max Failed Login Attempts - * @description Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Configurable from config.yaml only. Defaults to 50 + * @description Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Set under `general_settings` in config.yaml. Defaults to 50 */ max_failed_login_attempts?: number | null; /** * Max Failed Login Attempts Per Source - * @description Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Configurable from config.yaml only. Defaults to 250 + * @description Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Set under `general_settings` in config.yaml. Defaults to 250 */ max_failed_login_attempts_per_source?: number | null; /** From 65141a5fd89769dc1ec64873637b5992563f9c41 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 13:39:19 -0700 Subject: [PATCH 006/253] fix(proxy): keep the login throttle inside the type budget and fail safe on secret errors --- litellm/proxy/auth/login_throttle.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index ed53d9e7bc7..91e165006c3 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -60,7 +60,7 @@ _FAILED_LOGIN_USERNAME_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_USERNAME _FAILED_LOGIN_SOURCE_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_SOURCES) _NO_SETTINGS: Final = MappingProxyType({}) -_DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} +_DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} # mutable-ok: per-source slots taken and released around each held delay async def _sleep(seconds: float) -> None: @@ -138,7 +138,7 @@ class LoginThrottle: username_cache=_FAILED_LOGIN_USERNAME_CACHE, source_cache=_FAILED_LOGIN_SOURCE_CACHE, redis_cache=redis_usage_cache, - enabled=not get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT"), + enabled=not get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", False), ) @staticmethod From e9166019a523204e6964c976eb76791ffcb6e47a Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 14:02:29 -0700 Subject: [PATCH 007/253] fix(proxy): warn about per-worker sign-in counters without a module global --- litellm/proxy/auth/login_throttle.py | 13 +++++++++++++ litellm/proxy/proxy_server.py | 14 ++------------ 2 files changed, 15 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 91e165006c3..b84f944907b 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -13,6 +13,7 @@ import asyncio import hashlib from collections.abc import Awaitable from dataclasses import dataclass +from functools import cache from types import MappingProxyType from typing import Final, NamedTuple, NoReturn @@ -63,6 +64,18 @@ _NO_SETTINGS: Final = MappingProxyType({}) _DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} # mutable-ok: per-source slots taken and released around each held delay +@cache +def warn_login_counters_are_per_worker(num_workers: str) -> None: + """Warn once per process that failed sign-in counters are not shared across workers.""" + verbose_proxy_logger.warning( + "Running %s workers but Redis is not configured for LiteLLM caching. " + "Failed Admin UI sign-in attempts are counted per worker, so an attacker " + "gets max_failed_login_attempts guesses per worker instead of overall. " + "Configure Redis via the 'cache' section in your proxy config.", + num_workers, + ) + + async def _sleep(seconds: float) -> None: """The wait a rejected sign-in is held for. Replaced in tests so the suite pays no wall clock.""" await asyncio.sleep(seconds) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d9d647c5f18..56a9ba390c4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -299,7 +299,7 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.fallback_model_access import router_fallback_access_check from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import LicenseCheck -from litellm.proxy.auth.login_throttle import LoginThrottle +from litellm.proxy.auth.login_throttle import LoginThrottle, warn_login_counters_are_per_worker from litellm.proxy.auth.model_checks import ( expand_wildcard_deployments_for_model_info, get_all_fallbacks, @@ -2288,7 +2288,6 @@ user_custom_key_generate = None # Tests that need to reset it can patch 'litellm.proxy.proxy_server._pkce_no_redis_warning_emitted'. _pkce_no_redis_warning_emitted: bool = False _cp_no_redis_warning_emitted: bool = False -_login_throttle_no_redis_warning_emitted: bool = False user_custom_key_update = None user_custom_sso = None user_custom_ui_sso_sign_in_handler = None @@ -5512,16 +5511,7 @@ class ProxyConfig: # Failed Admin UI sign-in counters live in redis_usage_cache when available so a # brute-force run is counted once across workers instead of once per worker. if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: - global _login_throttle_no_redis_warning_emitted - if not _login_throttle_no_redis_warning_emitted: - _login_throttle_no_redis_warning_emitted = True - verbose_proxy_logger.warning( - "Running %s workers but Redis is not configured for LiteLLM caching. " - "Failed Admin UI sign-in attempts are counted per worker, so an attacker " - "gets max_failed_login_attempts guesses per worker instead of overall. " - "Configure Redis via the 'cache' section in your proxy config.", - os.getenv("NUM_WORKERS", "1"), - ) + warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) ### STORE MODEL IN DB ### feature flag for `/model/new` store_model_in_db = general_settings.get("store_model_in_db", False) From 429a4f213547d8a442913ebd7e83f573011224e7 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 14:35:05 -0700 Subject: [PATCH 008/253] fix: honor environment login throttle settings --- litellm/proxy/auth/login_throttle.py | 11 +++- .../proxy/auth/test_login_utils.py | 63 +++++++++++++++++-- .../proxy_server/test_routes_login_sso.py | 2 +- 3 files changed, 67 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index b84f944907b..4290cdbd62c 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -91,12 +91,19 @@ class FailureCounts(NamedTuple): def _int_setting(name: str, value: object, default: int, minimum: int) -> int: if value is None: return default - if isinstance(value, bool) or not isinstance(value, int) or value < minimum: + if isinstance(value, str): + try: + parsed: Final = int(value.strip()) + except ValueError: + parsed = value + else: + parsed = value + if isinstance(parsed, bool) or not isinstance(parsed, int) or parsed < minimum: verbose_proxy_logger.warning( "general_settings.%s=%r is not an integer >= %s; using the default of %s", name, value, minimum, default ) return default - return value + return parsed def _as_count(cached: object) -> int: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 1b85c9f4046..969103add6f 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -737,7 +737,7 @@ async def test_a_correct_admin_password_is_accepted_while_blocked(monkeypatch): await _guess(throttle) assert still_blocked.value.code == "429", "a wrong password is still refused" - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ): result = await _guess(throttle, password="right") @@ -780,7 +780,7 @@ async def test_a_successful_sign_in_clears_the_bucket(monkeypatch): with pytest.raises(ProxyException): await _guess(throttle) - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ): await _guess(throttle, password="right") @@ -859,7 +859,7 @@ async def test_both_credential_rejections_are_indistinguishable(monkeypatch): fake_user.password = "scrypt:fake" repo = MagicMock() repo.return_value.table.find_first = AsyncMock(return_value=fake_user) - with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( + with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( # test-quality-ok: reaches the known-DB-user branch without a database "litellm.proxy.auth.login_utils.verify_password", return_value=False ): with pytest.raises(ProxyException) as known: @@ -893,7 +893,7 @@ async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypat repo = MagicMock() repo.return_value.table.find_first = AsyncMock(return_value=passwordless) - with patch("litellm.proxy.auth.login_utils.UserRepository", repo): + with patch("litellm.proxy.auth.login_utils.UserRepository", repo): # test-quality-ok: reaches the passwordless-DB-user branch without a database for _ in range(5): with pytest.raises(ProxyException) as exc: await authenticate_user( @@ -934,7 +934,7 @@ async def test_a_wrong_password_for_a_known_user_also_counts(monkeypatch): throttle=throttle, ) - with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( + with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( # test-quality-ok: reaches the known-DB-user branch without a database "litellm.proxy.auth.login_utils.verify_password", return_value=False ): for _ in range(3): @@ -1035,7 +1035,7 @@ async def test_a_successful_sign_in_leaves_the_source_counter_alone(monkeypatch) with pytest.raises(ProxyException): await _guess(throttle) - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ): await _guess(throttle, password="right") @@ -1227,6 +1227,57 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): ) +def test_settings_that_arrive_as_environment_strings_are_honored(monkeypatch): + """An `os.environ/VAR` reference in general_settings resolves to a string, not an int. + + Regression: a digit string fell back to the default with only a log line, so an operator + tightening the limits through environment substitution silently kept the stock ceilings. + """ + from litellm.proxy import proxy_server as ps + from litellm.proxy.auth.login_throttle import LoginThrottle + + monkeypatch.setattr( + ps, + "general_settings", + { + "max_failed_login_attempts": "7", + "max_failed_login_attempts_per_source": " 70 ", + "failed_login_window_seconds": "not-a-number", + }, + ) + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "1.2.3.4" + + throttle = LoginThrottle.from_request(request) + + assert throttle.max_attempts == 7 + assert throttle.max_attempts_per_source == 70 + assert throttle.window_seconds == 900, "garbage still falls back to the default" + + +def test_a_negative_or_boolean_setting_falls_back_to_the_default(monkeypatch): + """A limit below one would refuse everyone; a bool is a typo, not a count.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy.auth.login_throttle import LoginThrottle + + monkeypatch.setattr( + ps, + "general_settings", + {"max_failed_login_attempts": "-7", "max_failed_login_attempts_per_source": True}, + ) + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "1.2.3.4" + + throttle = LoginThrottle.from_request(request) + + assert throttle.max_attempts == 50 + assert throttle.max_attempts_per_source == 250 + + @pytest.mark.asyncio async def test_a_refused_username_cannot_forge_log_lines(monkeypatch): """The username reaches a warning log, so it must not carry newlines or control bytes.""" diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index c8185cdb886..22d96fdc254 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -574,7 +574,7 @@ def test_the_configured_admin_password_still_signs_in_while_refused(client, monk assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] assert _json_login(client, "/v2/login") == 429 - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ): assert _json_login(client, "/v2/login", password="right-password") == 200 From 53ed7391ebcc0bb330ca12d06913f4fda6635fc8 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Tue, 1 Sep 2026 14:44:28 -0700 Subject: [PATCH 009/253] fix: satisfy login setting type checks --- litellm/proxy/auth/login_throttle.py | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 4290cdbd62c..1a978d9c212 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -88,16 +88,19 @@ class FailureCounts(NamedTuple): source: int +def _parse_int_setting(value: object) -> object: + if not isinstance(value, str): + return value + try: + return int(value.strip()) + except ValueError: + return value + + def _int_setting(name: str, value: object, default: int, minimum: int) -> int: if value is None: return default - if isinstance(value, str): - try: - parsed: Final = int(value.strip()) - except ValueError: - parsed = value - else: - parsed = value + parsed: Final = _parse_int_setting(value) if isinstance(parsed, bool) or not isinstance(parsed, int) or parsed < minimum: verbose_proxy_logger.warning( "general_settings.%s=%r is not an integer >= %s; using the default of %s", name, value, minimum, default From b40fc0ac22745924f2ff6bd073220f96d113cbce Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 9 Sep 2026 17:51:31 -0700 Subject: [PATCH 010/253] refactor(bedrock): resolve AWS credentials from one typed auth struct Every Bedrock and SageMaker call site hand-copied the same nine aws_* kwargs into BaseAWSLLM.get_credentials, so each new auth param has to be threaded into a dozen places and any site that misses one silently assumes the role with the wrong parameters. Introduce AwsAuthParams, a frozen pydantic model whose fields are exactly the credential-shaped params get_credentials accepts, plus resolve_credentials on BaseAWSLLM and pop_aws_auth_params for the call sites that must strip the keys out of optional_params. Deriving AWS_AUTH_PARAM_KEYS from the model's fields means the mirror list in common_utils can no longer drift from the struct. Behavior is unchanged: the same values reach STS from the same call sites. Dropping any one field from the resolver fails one of the new tests. Claude-Session: https://claude.ai/code/session_01E6zsK1DBcXfbetkgX86fw2 --- litellm/llms/bedrock/base_aws_llm.py | 81 +++++-------- litellm/llms/bedrock/batches/handler.py | 28 +---- litellm/llms/bedrock/chat/converse_handler.py | 31 +---- litellm/llms/bedrock/common_utils.py | 28 +---- litellm/llms/bedrock/embed/embedding.py | 36 ++---- litellm/llms/bedrock/files/handler.py | 15 +-- litellm/llms/bedrock/files/transformation.py | 42 +------ litellm/llms/bedrock/realtime/handler.py | 8 +- litellm/llms/sagemaker/chat/handler.py | 29 +---- litellm/llms/sagemaker/completion/handler.py | 29 +---- litellm/types/llms/bedrock.py | 20 ++++ .../llms/bedrock/test_base_aws_llm.py | 110 ++++++++++++++++++ .../types/llms/test_types_llms_bedrock.py | 46 ++++++++ 13 files changed, 248 insertions(+), 255 deletions(-) create mode 100644 tests/test_litellm/types/llms/test_types_llms_bedrock.py diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 96804a1fa62..609a6efbe95 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -6,11 +6,12 @@ import json import os import re import urllib.parse -from collections.abc import Callable, Mapping +from collections.abc import Callable, Mapping, MutableMapping from concurrent.futures import ThreadPoolExecutor from datetime import datetime from functools import partial from threading import Lock +from types import MappingProxyType from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload import httpx @@ -31,6 +32,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.aws_partition import contains_bedrock_arn, get_aws_dns_suffix from litellm.litellm_core_utils.dd_tracing import tracer from litellm.secret_managers.main import get_secret, get_secret_str +from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams if TYPE_CHECKING: from botocore.awsrequest import AWSPreparedRequest @@ -53,6 +55,14 @@ _STS_REGION_FROM_ENDPOINT_PATTERN: Final = re.compile( SIGV4_COMPUTED_HEADERS: Final = frozenset({"authorization", "x-amz-date", "x-amz-security-token", "date"}) +def pop_aws_auth_params( + optional_params: MutableMapping[str, object], # mutable-ok: pops the aws_* keys out of the caller's mapping +) -> AwsAuthParams: + return AwsAuthParams.model_validate( + MappingProxyType({key: optional_params.pop(key, None) for key in AWS_AUTH_PARAM_KEYS}) + ) + + class BedrockRequestTarget(BaseModel): aws_region_name: str aws_bedrock_runtime_endpoint: str | None @@ -379,6 +389,20 @@ class BaseAWSLLM(SignsRequestsWithAWS): else: return self._get_or_set_cached_credentials(args, self._auth_with_env_vars) + def resolve_credentials(self, auth_params: AwsAuthParams, aws_region_name: str | None) -> Credentials: + return self.get_credentials( + aws_access_key_id=auth_params.aws_access_key_id, + aws_secret_access_key=auth_params.aws_secret_access_key, + aws_session_token=auth_params.aws_session_token, + aws_region_name=aws_region_name, + aws_session_name=auth_params.aws_session_name, + aws_profile_name=auth_params.aws_profile_name, + aws_role_name=auth_params.aws_role_name, + aws_web_identity_token=auth_params.aws_web_identity_token, + aws_sts_endpoint=auth_params.aws_sts_endpoint, + aws_external_id=auth_params.aws_external_id, + ) + def _get_aws_region_from_model_arn(self, model: str | None) -> str | None: try: # First check if the string contains the expected prefix @@ -1453,22 +1477,10 @@ class BaseAWSLLM(SignsRequestsWithAWS): from botocore.credentials import Credentials except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None) - aws_session_token: Final = optional_params.pop("aws_session_token", None) aws_region_name: Final = self._get_aws_region_name(optional_params, model) optional_params.pop("aws_region_name", None) - aws_role_name: Final = optional_params.pop("aws_role_name", None) - aws_session_name: Final = optional_params.pop("aws_session_name", None) - aws_profile_name: Final = optional_params.pop("aws_profile_name", None) - aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None) - aws_bedrock_runtime_endpoint: Final = optional_params.pop( - "aws_bedrock_runtime_endpoint", None - ) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_external_id: Final = optional_params.pop("aws_external_id", None) + auth_params: Final = pop_aws_auth_params(optional_params) + aws_bedrock_runtime_endpoint: Final = optional_params.pop("aws_bedrock_runtime_endpoint", None) if bearer_token is not None: return BearerRequestTarget( @@ -1476,18 +1488,7 @@ class BaseAWSLLM(SignsRequestsWithAWS): aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, ) - credentials: Final[Credentials] = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name) return Boto3CredentialsInfo( credentials=credentials, aws_region_name=aws_region_name, @@ -1621,31 +1622,9 @@ class BaseAWSLLM(SignsRequestsWithAWS): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.get("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.get("aws_access_key_id", None) - aws_session_token: Final = optional_params.get("aws_session_token", None) - aws_role_name: Final = optional_params.get("aws_role_name", None) - aws_session_name: Final = optional_params.get("aws_session_name", None) - aws_profile_name: Final = optional_params.get("aws_profile_name", None) - aws_web_identity_token: Final = optional_params.get("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.get("aws_sts_endpoint", None) - aws_external_id: Final = optional_params.get("aws_external_id", None) + auth_params: Final = AwsAuthParams.model_validate(optional_params) aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model=model) - - credentials: Final[Credentials] = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name) sigv4: Final = SigV4Auth(credentials, service_name, aws_region_name) headers = headers or {} diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index 4b500897642..fe8575d323e 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -6,6 +6,7 @@ from openai.types.batch import BatchRequestCounts from openai.types.batch import Metadata as OpenAIBatchMetadata from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.types.llms.bedrock import AwsAuthParams from litellm.types.utils import LiteLLMBatch if TYPE_CHECKING: @@ -128,11 +129,10 @@ class BedrockBatchesHandler: from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig - creds: Final = BedrockBatchesConfig().get_credentials( + auth_params: Final = AwsAuthParams( aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, aws_session_token=aws_session_token, - aws_region_name=region, aws_session_name=aws_session_name, aws_profile_name=aws_profile_name, aws_role_name=aws_role_name, @@ -140,6 +140,7 @@ class BedrockBatchesHandler: aws_sts_endpoint=aws_sts_endpoint, aws_external_id=aws_external_id, ) + creds: Final = BedrockBatchesConfig().resolve_credentials(auth_params, region) client: Final = boto3.client( "bedrock", @@ -154,15 +155,7 @@ class BedrockBatchesHandler: batch_id=batch_id, aws_region_name=region, logging_obj=logging_obj, - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, + **auth_params.model_dump(), ) try: @@ -306,18 +299,7 @@ class BedrockBatchesHandler: # BaseAWSLLM) lazily to avoid a circular import at module load. from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig - creds: Final = BedrockBatchesConfig().get_credentials( - aws_access_key_id=kwargs.get("aws_access_key_id"), - aws_secret_access_key=kwargs.get("aws_secret_access_key"), - aws_session_token=kwargs.get("aws_session_token"), - aws_region_name=region, - aws_session_name=kwargs.get("aws_session_name"), - aws_profile_name=kwargs.get("aws_profile_name"), - aws_role_name=kwargs.get("aws_role_name"), - aws_web_identity_token=kwargs.get("aws_web_identity_token"), - aws_sts_endpoint=kwargs.get("aws_sts_endpoint"), - aws_external_id=kwargs.get("aws_external_id"), - ) + creds: Final = BedrockBatchesConfig().resolve_credentials(AwsAuthParams.model_validate(kwargs), region) client: Final = boto3.client( "bedrock", diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 6d48ff3f07c..666230076e1 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -21,7 +21,7 @@ from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper -from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing +from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -343,20 +343,8 @@ class BedrockConverseLLM(BaseAWSLLM): model_id=unencoded_model_id, ) - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None) - aws_session_token: Final = optional_params.pop("aws_session_token", None) - aws_role_name: Final = optional_params.pop("aws_role_name", None) - aws_session_name: Final = optional_params.pop("aws_session_name", None) - aws_profile_name: Final = optional_params.pop("aws_profile_name", None) - aws_bedrock_runtime_endpoint: Final = optional_params.pop( - "aws_bedrock_runtime_endpoint", None - ) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None) - aws_external_id: Final = optional_params.pop("aws_external_id", None) + auth_params: Final = pop_aws_auth_params(optional_params) + aws_bedrock_runtime_endpoint: Final = optional_params.pop("aws_bedrock_runtime_endpoint", None) optional_params.pop("aws_region_name", None) litellm_params["aws_region_name"] = aws_region_name # [DO NOT DELETE] important for async calls @@ -364,18 +352,7 @@ class BedrockConverseLLM(BaseAWSLLM): credentials: Final[Credentials | None] = ( None if bedrock_bearer_token(api_key) is not None - else self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + else self.resolve_credentials(auth_params, aws_region_name) ) ### SET RUNTIME ENDPOINT ### diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index be4f0f32689..8bbb9ad3723 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -28,6 +28,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import ( from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.secret_managers.main import get_secret, get_secret_str +from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams if TYPE_CHECKING: from litellm.types.llms.openai import AllMessageValues @@ -82,18 +83,7 @@ class BedrockError(BaseLLMException): ) -_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = ( - "aws_access_key_id", - "aws_secret_access_key", - "aws_session_token", - "aws_region_name", - "aws_session_name", - "aws_profile_name", - "aws_role_name", - "aws_web_identity_token", - "aws_sts_endpoint", - "aws_external_id", -) +_BEDROCK_AWS_AUTH_PARAMETER_KEYS: Final[tuple[str, ...]] = (*AWS_AUTH_PARAM_KEYS, "aws_region_name") def merge_bedrock_aws_request_params( @@ -1650,19 +1640,9 @@ class CommonBatchFilesUtils: except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - # Get AWS credentials using existing methods aws_region_name: Final = self._base_aws._get_aws_region_name(optional_params=optional_params, model="") - credentials: Final = self._base_aws.get_credentials( - aws_access_key_id=optional_params.get("aws_access_key_id"), - aws_secret_access_key=optional_params.get("aws_secret_access_key"), - aws_session_token=optional_params.get("aws_session_token"), - aws_region_name=aws_region_name, - aws_session_name=optional_params.get("aws_session_name"), - aws_profile_name=optional_params.get("aws_profile_name"), - aws_role_name=optional_params.get("aws_role_name"), - aws_web_identity_token=optional_params.get("aws_web_identity_token"), - aws_sts_endpoint=optional_params.get("aws_sts_endpoint"), - aws_external_id=optional_params.get("aws_external_id"), + credentials: Final = self._base_aws.resolve_credentials( + AwsAuthParams.model_validate(optional_params), aws_region_name ) # Prepare the request data diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index be766eaedd0..46d7b1ef9e7 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -26,7 +26,14 @@ from litellm.types.llms.bedrock import ( ) from litellm.types.utils import EmbeddingResponse, LlmProviders -from ..base_aws_llm import AWSPreparedRequest, BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing +from ..base_aws_llm import ( + AWSPreparedRequest, + BaseAWSLLM, + Credentials, + bedrock_bearer_token, + pop_aws_auth_params, + run_aws_signing, +) from ..common_utils import BedrockError from .amazon_nova_transformation import AmazonNovaEmbeddingConfig from .amazon_titan_g1_transformation import AmazonTitanG1Config @@ -75,18 +82,8 @@ class BedrockEmbedding(BaseAWSLLM): optional_params: dict, bearer_token: str | None = None, ) -> tuple[Credentials | None, str]: - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None) - aws_session_token: Final = optional_params.pop("aws_session_token", None) + auth_params: Final = pop_aws_auth_params(optional_params) aws_region_name = optional_params.pop("aws_region_name", None) - aws_role_name: Final = optional_params.pop("aws_role_name", None) - aws_session_name: Final = optional_params.pop("aws_session_name", None) - aws_profile_name: Final = optional_params.pop("aws_profile_name", None) - aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None) - aws_external_id: Final = optional_params.pop("aws_external_id", None) ### SET REGION NAME ### if aws_region_name is None: @@ -104,20 +101,7 @@ class BedrockEmbedding(BaseAWSLLM): aws_region_name = "us-west-2" credentials: Final[Credentials | None] = ( - None - if bearer_token is not None - else self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + None if bearer_token is not None else self.resolve_credentials(auth_params, aws_region_name) ) return credentials, aws_region_name diff --git a/litellm/llms/bedrock/files/handler.py b/litellm/llms/bedrock/files/handler.py index e74c3802d20..0b75474ba1b 100644 --- a/litellm/llms/bedrock/files/handler.py +++ b/litellm/llms/bedrock/files/handler.py @@ -11,6 +11,7 @@ from litellm.litellm_core_utils.cloud_storage_security import ( validate_managed_cloud_file_id, ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.bedrock import AwsAuthParams from litellm.types.llms.openai import ( FileContentRequest, HttpxBinaryResponseContent, @@ -101,19 +102,9 @@ class BedrockFilesHandler(BaseAWSLLM): allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(optional_params), ) - # Get AWS credentials aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model="") - credentials: Final[Credentials] = self.get_credentials( - aws_access_key_id=optional_params.get("aws_access_key_id"), - aws_secret_access_key=optional_params.get("aws_secret_access_key"), - aws_session_token=optional_params.get("aws_session_token"), - aws_region_name=aws_region_name, - aws_session_name=optional_params.get("aws_session_name"), - aws_profile_name=optional_params.get("aws_profile_name"), - aws_role_name=optional_params.get("aws_role_name"), - aws_web_identity_token=optional_params.get("aws_web_identity_token"), - aws_sts_endpoint=optional_params.get("aws_sts_endpoint"), - aws_external_id=optional_params.get("aws_external_id"), + credentials: Final[Credentials] = self.resolve_credentials( + AwsAuthParams.model_validate(optional_params), aws_region_name ) # Create S3 client diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 9875ac2b9c3..cccac11dc36 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -41,7 +41,7 @@ from litellm.llms.base_llm.files.transformation import ( BaseFilesConfig, LiteLLMLoggingObj, ) -from litellm.types.llms.bedrock import BedrockBatchRecordKind +from litellm.types.llms.bedrock import AwsAuthParams, BedrockBatchRecordKind from litellm.types.llms.openai import ( AllMessageValues, CreateFileRequest, @@ -133,21 +133,10 @@ def _responses_request_adapter() -> TypeAdapter[ResponsesAPIOptionalRequestParam return TypeAdapter(ResponsesAPIOptionalRequestParams) -class _BedrockS3RequestParams(BaseModel): +class _BedrockS3RequestParams(AwsAuthParams): """Typed view of the credential/region params the S3 GetObject path reads.""" - model_config = ConfigDict(extra="ignore") - - aws_access_key_id: str | None = None - aws_secret_access_key: str | None = None - aws_session_token: str | None = None aws_region_name: str | None = None - aws_session_name: str | None = None - aws_profile_name: str | None = None - aws_role_name: str | None = None - aws_web_identity_token: str | None = None - aws_sts_endpoint: str | None = None - aws_external_id: str | None = None s3_region_name: str | None = None s3_endpoint_url: str | None = None @@ -1019,20 +1008,8 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - # Get AWS credentials using existing methods aws_region_name: Final = self._get_aws_region_name(optional_params=optional_params, model="") - credentials: Final = self.get_credentials( - aws_access_key_id=optional_params.get("aws_access_key_id"), - aws_secret_access_key=optional_params.get("aws_secret_access_key"), - aws_session_token=optional_params.get("aws_session_token"), - aws_region_name=aws_region_name, - aws_session_name=optional_params.get("aws_session_name"), - aws_profile_name=optional_params.get("aws_profile_name"), - aws_role_name=optional_params.get("aws_role_name"), - aws_web_identity_token=optional_params.get("aws_web_identity_token"), - aws_sts_endpoint=optional_params.get("aws_sts_endpoint"), - aws_external_id=optional_params.get("aws_external_id"), - ) + credentials: Final = self.resolve_credentials(AwsAuthParams.model_validate(optional_params), aws_region_name) # Calculate SHA256 hash of the content (REQUIRED for S3) content_hash: Final = hashlib.sha256(content.encode("utf-8")).hexdigest() @@ -1296,18 +1273,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - credentials: Final = self.get_credentials( # any-ok: boto3 Credentials is untyped - aws_access_key_id=request_params.aws_access_key_id, - aws_secret_access_key=request_params.aws_secret_access_key, - aws_session_token=request_params.aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=request_params.aws_session_name, - aws_profile_name=request_params.aws_profile_name, - aws_role_name=request_params.aws_role_name, - aws_web_identity_token=request_params.aws_web_identity_token, - aws_sts_endpoint=request_params.aws_sts_endpoint, - aws_external_id=request_params.aws_external_id, - ) + credentials: Final = self.resolve_credentials(request_params, aws_region_name) empty_body_hash: Final = hashlib.sha256(b"").hexdigest() aws_request: Final = AWSRequest( # any-ok: botocore AWSRequest is untyped diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index ca2370303f2..43f34da9d21 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -18,6 +18,7 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.realtime_streaming import DefaultLoggedRealTimeEventTypes +from litellm.types.llms.bedrock import AwsAuthParams from litellm.types.llms.openai import OpenAIRealtimeEvents from litellm.types.realtime import RealtimeResponseTransformInput @@ -149,12 +150,10 @@ class BedrockRealtime(BaseAWSLLM): verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model) - credentials: Final = await run_aws_signing( - self.get_credentials, + auth_params: Final = AwsAuthParams( aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, aws_session_token=aws_session_token, - aws_region_name=aws_region_name, aws_session_name=aws_session_name, aws_profile_name=aws_profile_name, aws_role_name=aws_role_name, @@ -162,7 +161,8 @@ class BedrockRealtime(BaseAWSLLM): aws_sts_endpoint=aws_sts_endpoint, aws_external_id=aws_external_id, ) - if credentials is None: + credentials: Final = await run_aws_signing(self.resolve_credentials, auth_params, aws_region_name) + if credentials is None: # pyright: ignore[reportUnnecessaryComparison] # boto3.Session() env fallback yields None raise BedrockError( status_code=401, message=( diff --git a/litellm/llms/sagemaker/chat/handler.py b/litellm/llms/sagemaker/chat/handler.py index 3f62b7276df..10be9ef384c 100644 --- a/litellm/llms/sagemaker/chat/handler.py +++ b/litellm/llms/sagemaker/chat/handler.py @@ -6,7 +6,7 @@ from typing import Final import httpx from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix -from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.utils import ModelResponse, get_secret @@ -23,19 +23,9 @@ class SagemakerChatHandler(BaseAWSLLM): from botocore.credentials import Credentials except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None) - aws_session_token: Final = optional_params.pop("aws_session_token", None) + auth_params: Final = pop_aws_auth_params(optional_params) aws_region_name = optional_params.pop("aws_region_name", None) - aws_role_name: Final = optional_params.pop("aws_role_name", None) - aws_session_name: Final = optional_params.pop("aws_session_name", None) - aws_profile_name: Final = optional_params.pop("aws_profile_name", None) - optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None) - aws_external_id: Final = optional_params.pop("aws_external_id", None) + optional_params.pop("aws_bedrock_runtime_endpoint", None) ### SET REGION NAME ### if aws_region_name is None: @@ -52,18 +42,7 @@ class SagemakerChatHandler(BaseAWSLLM): if aws_region_name is None: aws_region_name = "us-west-2" - credentials: Final[Credentials] = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name) return credentials, aws_region_name def _prepare_request( diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index fb8074d3682..3e110a869bc 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -10,7 +10,7 @@ from litellm._logging import verbose_logger from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, @@ -46,19 +46,9 @@ class SagemakerLLM(BaseAWSLLM): from botocore.credentials import Credentials except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - ## CREDENTIALS ## - # pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them - aws_secret_access_key: Final = optional_params.pop("aws_secret_access_key", None) - aws_access_key_id: Final = optional_params.pop("aws_access_key_id", None) - aws_session_token: Final = optional_params.pop("aws_session_token", None) + auth_params: Final = pop_aws_auth_params(optional_params) aws_region_name = optional_params.pop("aws_region_name", None) - aws_role_name: Final = optional_params.pop("aws_role_name", None) - aws_session_name: Final = optional_params.pop("aws_session_name", None) - aws_profile_name: Final = optional_params.pop("aws_profile_name", None) - optional_params.pop("aws_bedrock_runtime_endpoint", None) # https://bedrock-runtime.{region_name}.amazonaws.com - aws_web_identity_token: Final = optional_params.pop("aws_web_identity_token", None) - aws_sts_endpoint: Final = optional_params.pop("aws_sts_endpoint", None) - aws_external_id: Final = optional_params.pop("aws_external_id", None) + optional_params.pop("aws_bedrock_runtime_endpoint", None) ### SET REGION NAME ### if aws_region_name is None: @@ -75,18 +65,7 @@ class SagemakerLLM(BaseAWSLLM): if aws_region_name is None: aws_region_name = "us-west-2" - credentials: Final[Credentials] = self.get_credentials( - aws_access_key_id=aws_access_key_id, - aws_secret_access_key=aws_secret_access_key, - aws_session_token=aws_session_token, - aws_region_name=aws_region_name, - aws_session_name=aws_session_name, - aws_profile_name=aws_profile_name, - aws_role_name=aws_role_name, - aws_web_identity_token=aws_web_identity_token, - aws_sts_endpoint=aws_sts_endpoint, - aws_external_id=aws_external_id, - ) + credentials: Final[Credentials] = self.resolve_credentials(auth_params, aws_region_name) return credentials, aws_region_name def _prepare_request( diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 9f93886a9c6..9a655a01c74 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -3,6 +3,7 @@ from collections.abc import Sequence from enum import Enum from typing import TYPE_CHECKING, Any, Final, Literal, TypeAlias +from pydantic import BaseModel, ConfigDict from typing_extensions import ReadOnly, Required, TypedDict, override from .openai import ChatCompletionToolCallChunk @@ -1107,6 +1108,25 @@ class BedrockTag(TypedDict): value: str +class AwsAuthParams(BaseModel): + """Every credential-shaped aws_* param BaseAWSLLM.get_credentials accepts; region is resolved separately.""" + + model_config = ConfigDict(frozen=True, extra="ignore") + + aws_access_key_id: str | None = None + aws_secret_access_key: str | None = None + aws_session_token: str | None = None + aws_session_name: str | None = None + aws_profile_name: str | None = None + aws_role_name: str | None = None + aws_web_identity_token: str | None = None + aws_sts_endpoint: str | None = None + aws_external_id: str | None = None + + +AWS_AUTH_PARAM_KEYS: Final[tuple[str, ...]] = tuple(AwsAuthParams.model_fields) + + class BedrockCreateBatchRequest(TypedDict, total=False): """ Request structure for creating a Bedrock batch inference job. diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index 2c7e476e9a8..a08165e855b 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -3278,3 +3278,113 @@ def test_run_aws_signing_leaves_the_default_executor_free_for_other_providers(): other_provider, signing_thread = asyncio.run(scenario()) assert other_provider != signing_thread assert signing_thread.startswith("aws-signing") + + +def _recording_boto3_client(recorded: Dict[str, Any]): + """boto3.client replacement that records the STS client kwargs and the assume-role params.""" + + def _client(service_name, **client_kwargs): + recorded["client_kwargs"] = client_kwargs + sts = MagicMock() + + def _assume(**params): + recorded["assume_role"] = params + return { + "Credentials": { + "AccessKeyId": "ASIAASSUMED", + "SecretAccessKey": "assumed-secret", + "SessionToken": "assumed-token", + "Expiration": datetime.now(timezone.utc) + timedelta(minutes=30), + } + } + + def _assume_web_identity(**params): + recorded["assume_role_with_web_identity"] = params + return { + "Credentials": { + "AccessKeyId": "ASIAWEBIDENTITY", + "SecretAccessKey": "assumed-secret", + "SessionToken": "assumed-token", + "Expiration": datetime.now(timezone.utc) + timedelta(minutes=30), + }, + "PackedPolicySize": 10, + } + + sts.assume_role.side_effect = _assume + sts.assume_role_with_web_identity.side_effect = _assume_web_identity + return sts + + return _client + + +def test_resolve_credentials_forwards_static_keys_role_session_and_external_id(): + """Every field the role-assumption route reads must reach STS, so a dropped struct field fails here.""" + from litellm.types.llms.bedrock import AwsAuthParams + + auth_params = AwsAuthParams( + aws_access_key_id="AKIACALLER", + aws_secret_access_key="caller-secret", + aws_session_token="caller-token", + aws_role_name="arn:aws:iam::123456789012:role/litellm-target", + aws_session_name="litellm-session", + aws_external_id="litellm-external-id", + aws_sts_endpoint="https://custom-sts.example", + ) + recorded: Dict[str, Any] = {} + + with ( + patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), + patch("boto3.client", side_effect=_recording_boto3_client(recorded)), + ): + credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1") + + assert recorded["client_kwargs"]["aws_access_key_id"] == "AKIACALLER" + assert recorded["client_kwargs"]["aws_secret_access_key"] == "caller-secret" + assert recorded["client_kwargs"]["aws_session_token"] == "caller-token" + assert recorded["client_kwargs"]["endpoint_url"] == "https://custom-sts.example" + assert recorded["assume_role"]["RoleArn"] == "arn:aws:iam::123456789012:role/litellm-target" + assert recorded["assume_role"]["RoleSessionName"] == "litellm-session" + assert recorded["assume_role"]["ExternalId"] == "litellm-external-id" + assert credentials.access_key == "ASIAASSUMED" + + +def test_resolve_credentials_forwards_web_identity_token(): + """A struct carrying a web-identity token must take the web-identity route, not plain role assumption.""" + from litellm.types.llms.bedrock import AwsAuthParams + + auth_params = AwsAuthParams( + aws_web_identity_token="unresolvable-oidc-token", + aws_role_name="arn:aws:iam::123456789012:role/litellm-wif", + aws_session_name="litellm-wif-session", + ) + recorded: Dict[str, Any] = {} + + with ( + patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), + patch("boto3.client", side_effect=_recording_boto3_client(recorded)), + ): + with pytest.raises(AwsAuthError) as exc: + BaseAWSLLM().resolve_credentials(auth_params, "us-east-1") + + assert exc.value.status_code == 401 + assert "assume_role" not in recorded + + +def test_resolve_credentials_forwards_profile_name(): + """The profile route must receive the struct's profile name rather than the ambient session.""" + from litellm.types.llms.bedrock import AwsAuthParams + + auth_params = AwsAuthParams(aws_profile_name="litellm-qa-profile") + session_instance = MagicMock() + session_instance.get_credentials.return_value = Credentials( + access_key="AKIAPROFILE", secret_key="profile-secret", token=None + ) + + with ( + patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), + patch("boto3.Session", return_value=session_instance) as mock_session_cls, + ): + credentials = BaseAWSLLM().resolve_credentials(auth_params, "us-east-1") + + assert mock_session_cls.call_args.kwargs["profile_name"] == "litellm-qa-profile" + assert credentials.access_key == "AKIAPROFILE" diff --git a/tests/test_litellm/types/llms/test_types_llms_bedrock.py b/tests/test_litellm/types/llms/test_types_llms_bedrock.py new file mode 100644 index 00000000000..a5ad882e775 --- /dev/null +++ b/tests/test_litellm/types/llms/test_types_llms_bedrock.py @@ -0,0 +1,46 @@ +import pytest +from pydantic import ValidationError + +from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams + + +def test_model_validate_keeps_auth_params_and_ignores_request_params(): + auth_params = AwsAuthParams.model_validate( + { + "aws_role_name": "arn:aws:iam::999999999999:role/litellm-role", + "aws_session_name": "litellm-session", + "aws_external_id": "litellm-external-id", + "aws_region_name": "us-west-2", + "aws_bedrock_runtime_endpoint": "https://bedrock.example.com", + "model": "anthropic.claude-haiku-4-5-20251001-v1:0", + "temperature": 0.1, + "messages": [{"role": "user", "content": "hi"}], + } + ) + + assert auth_params.aws_role_name == "arn:aws:iam::999999999999:role/litellm-role" + assert auth_params.aws_session_name == "litellm-session" + assert auth_params.aws_external_id == "litellm-external-id" + assert auth_params.aws_access_key_id is None + assert set(auth_params.model_dump()) == set(AWS_AUTH_PARAM_KEYS) + assert not set(AWS_AUTH_PARAM_KEYS) & {"aws_region_name", "aws_bedrock_runtime_endpoint", "model", "temperature"} + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("aws_role_name", 1234), + ("aws_session_name", ["litellm-session"]), + ("aws_external_id", {"id": "x"}), + ], +) +def test_model_validate_rejects_non_string_credentials(field, value): + with pytest.raises(ValidationError): + AwsAuthParams.model_validate({field: value}) + + +def test_frozen_struct_rejects_field_assignment(): + auth_params = AwsAuthParams(aws_role_name="arn:aws:iam::999999999999:role/litellm-role") + + with pytest.raises(ValidationError): + auth_params.aws_role_name = "arn:aws:iam::999999999999:role/other-role" From 36b346d31ae2da6c1ffcfde835a27b833b583623 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 01:02:44 +0000 Subject: [PATCH 011/253] refactor(proxy): keep authenticate_user within the C901 budget after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_utils.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 5b0d5e6edd7..d0bb9e3087d 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -98,6 +98,14 @@ def _matches_env_credentials(username: str, password: str, master_key: str | Non ) +def _admin_credentials_match( + username: str, password: str, master_key: str, general_settings: Mapping[str, object] +) -> bool: + return general_settings.get("disable_env_credential_login") is not True and _matches_env_credentials( + username, password, master_key + ) + + def _invalid_credentials_message(general_settings: Mapping[str, object]) -> str: """One rejection message for unknown usernames and wrong passwords alike, so neither can be enumerated.""" if is_env_credential_login_enabled(general_settings): @@ -209,9 +217,7 @@ async def authenticate_user( code=500, ) - admin_credentials_match: Final = general_settings.get("disable_env_credential_login") is not True and ( - _matches_env_credentials(username, password, master_key) - ) + admin_credentials_match: Final = _admin_credentials_match(username, password, master_key, general_settings) if not admin_credentials_match: await throttle.raise_if_blocked(username) @@ -247,12 +253,7 @@ async def authenticate_user( user_id = LITELLM_PROXY_ADMIN_NAME # we want the key created to have PROXY_ADMIN_PERMISSIONS - key_user_id = LITELLM_PROXY_ADMIN_NAME - if ( - os.getenv("PROXY_ADMIN_ID", None) is not None and os.environ["PROXY_ADMIN_ID"] == user_id - ) or user_id == LITELLM_PROXY_ADMIN_NAME: - # checks if user is admin - key_user_id = os.getenv("PROXY_ADMIN_ID", LITELLM_PROXY_ADMIN_NAME) + key_user_id: Final = os.getenv("PROXY_ADMIN_ID", LITELLM_PROXY_ADMIN_NAME) # Admin is Authe'd in - generate key for the UI to access Proxy From 82902e83c2d7a8f909ee8ee49a5123b0618de4de Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 07:54:19 +0000 Subject: [PATCH 012/253] fix(proxy): write failed-login counters and their expiry in one Redis call Use RedisCache.async_increment_with_floor (a single Lua INCRBY + EXPIRE) for the shared login counters instead of the two-step INCRBYFLOAT then EXPIRE, so a counter can never be committed to Redis without its expiry. The repair in _remaining_window now only covers expiries stripped out of band (PERSIST, a restore) and uses the same atomic call Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 15 ++++--- .../proxy/auth/test_login_utils.py | 43 +++++++++++-------- 2 files changed, 33 insertions(+), 25 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 1a978d9c212..926b7a0b9ce 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -194,11 +194,11 @@ class LoginThrottle: return max(local, _as_count(await self._outcome(redis_cache.async_get_cache(key)))) async def _remaining_window(self, key: str) -> int: - """Seconds until this counter expires, repairing a counter left without an expiry. + """Seconds until this counter expires. - Redis commits the increment before setting the TTL, so a failure in between can - leave a counter that never expires. Nothing increments the key again once the - limit is reached, so without the repair the key would stay refused indefinitely. + Counters are only ever written together with their expiry, so a counter without one + was stripped out of band (PERSIST, a restore). It is given the full window again, + since nothing increments a key once the limit is reached. """ redis_cache: Final = self.redis_cache if redis_cache is None: @@ -206,7 +206,7 @@ class LoginThrottle: ttl: Final = await self._outcome(redis_cache.async_get_ttl(key)) if isinstance(ttl, int) and ttl > 0: return min(ttl, self.window_seconds) - await self._outcome(redis_cache.async_increment(key, 0, ttl=self.window_seconds)) + await self._outcome(redis_cache.async_increment_with_floor(key, 0, self.window_seconds)) return self.window_seconds def _refused(self, retry_after: int, param: str) -> ProxyException: @@ -260,8 +260,9 @@ class LoginThrottle: redis_cache: Final = self.redis_cache if redis_cache is None: return local - shared: Final = _as_count(await self._outcome(redis_cache.async_increment(key, 1, ttl=self.window_seconds))) - await self._remaining_window(key) + shared: Final = _as_count( + await self._outcome(redis_cache.async_increment_with_floor(key, 1, self.window_seconds)) + ) return max(local, shared) async def record_failure(self, username: str) -> FailureCounts: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 73154c5ecb0..f5c8b02b834 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1142,58 +1142,65 @@ async def test_disabling_the_control_removes_the_delay_as_well(monkeypatch, logi assert login_delays.seconds == [] -class _NoExpiryRedis: - """Redis that stores the counter but never records an expiry for it. +class _FakeRedis: + """Redis whose only counter write is the atomic INCRBY-plus-EXPIRE Lua call. - Models the window between INCRBYFLOAT committing and the TTL call failing. + `async_increment` is deliberately absent: a two-step increment would fail the test + with AttributeError, because Redis could then commit a count without its expiry. """ def __init__(self): self.values: dict = {} - self.expiry_repairs = 0 + self.ttls: dict = {} async def async_get_cache(self, key, **kwargs): return self.values.get(key) - async def async_increment(self, key, value, ttl=None, **kwargs): - if int(value) == 0: - self.expiry_repairs += 1 - self.values[key] = self.values.get(key, 0) + int(value) + async def async_increment_with_floor(self, key, value, ttl): + self.values[key] = self.values.get(key, 0) + value + self.ttls.setdefault(key, ttl) return self.values[key] async def async_get_ttl(self, key): - return None + return self.ttls.get(key) async def async_delete_cache(self, key): self.values.pop(key, None) + self.ttls.pop(key, None) + + def persist(self): + self.ttls.clear() @pytest.mark.asyncio -async def test_a_counter_left_without_an_expiry_is_repaired(monkeypatch): +async def test_counters_are_written_with_their_expiry_and_re_armed_if_stripped(monkeypatch): """Regression: a counter with no TTL would refuse the pair forever. - Redis commits the increment before setting the expiry, and nothing increments the key - again once the limit is reached, so a TTL that never landed is never repaired on its - own and the username and source pair stays refused with no way back. + Nothing increments a key once the limit is reached, so a counter that ever exists + without an expiry stays refused with no way back. Every write must therefore carry the + expiry, and a refusal that finds it stripped (PERSIST) must put the window back. """ from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - redis = _NoExpiryRedis() - throttle = _throttle(max_attempts=2, redis_cache=redis) + redis = _FakeRedis() + throttle = _throttle(max_attempts=2, window_seconds=77, redis_cache=redis) for _ in range(2): with pytest.raises(ProxyException): await _guess(throttle) - assert redis.expiry_repairs >= 1, "each recorded failure must leave the counter with an expiry" + assert redis.values, "failures must land in the shared counter" + assert set(redis.ttls) == set(redis.values), "no counter may exist without its expiry" + assert set(redis.ttls.values()) == {77} - repairs_before_block = redis.expiry_repairs + redis.persist() with pytest.raises(ProxyException) as blocked: await _guess(throttle) assert blocked.value.code == "429" - assert redis.expiry_repairs > repairs_before_block, "the refusal path must repair a missing expiry too" + assert blocked.value.headers.get("Retry-After") == "77" + assert set(redis.ttls) >= {k for k in redis.values if ":user:" in k}, "the refusal must re-arm a stripped expiry" @pytest.mark.asyncio From ec9a926e84df7e477e16664d6f517dd34612a8a7 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 08:34:11 +0000 Subject: [PATCH 013/253] fix(proxy): resolve the login rate limit kill switch once per process Reading LITELLM_DISABLE_LOGIN_RATE_LIMIT through get_secret_bool on every unauthenticated sign-in attempt meant a hosted secret manager in read mode was queried once per password guess, before any counter was checked Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 8 +++++- .../proxy/auth/test_login_utils.py | 25 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 926b7a0b9ce..11290ef6d40 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -76,6 +76,12 @@ def warn_login_counters_are_per_worker(num_workers: str) -> None: ) +@cache +def _rate_limit_disabled() -> bool: + """Resolved once per process so an unauthenticated flood never reaches the secret manager.""" + return bool(get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", False)) + + async def _sleep(seconds: float) -> None: """The wait a rejected sign-in is held for. Replaced in tests so the suite pays no wall clock.""" await asyncio.sleep(seconds) @@ -161,7 +167,7 @@ class LoginThrottle: username_cache=_FAILED_LOGIN_USERNAME_CACHE, source_cache=_FAILED_LOGIN_SOURCE_CACHE, redis_cache=redis_usage_cache, - enabled=not get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", False), + enabled=not _rate_limit_disabled(), ) @staticmethod diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index f5c8b02b834..8afac8a622d 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1267,6 +1267,31 @@ def test_settings_that_arrive_as_environment_strings_are_honored(monkeypatch): assert throttle.window_seconds == 900, "garbage still falls back to the default" +def test_the_disable_flag_is_read_once_not_per_login_attempt(monkeypatch): + """Regression: the kill switch was read through the secret manager on every unauthenticated request. + + With a hosted secret manager in read mode that is a synchronous network call per guess, so a + flood of wrong passwords could exhaust the secret manager even after the source was refused. + """ + from litellm.proxy import proxy_server as ps + from litellm.proxy.auth import login_throttle + + reads: Final[list[str]] = [] # mutable-ok: test-only call recorder + monkeypatch.setattr(login_throttle, "get_secret_bool", lambda name, default: reads.append(name) or default) + login_throttle._rate_limit_disabled.cache_clear() + monkeypatch.setattr(ps, "general_settings", {}) + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "1.2.3.4" + + for _ in range(50): + assert login_throttle.LoginThrottle.from_request(request).enabled is True + + login_throttle._rate_limit_disabled.cache_clear() + assert reads == ["LITELLM_DISABLE_LOGIN_RATE_LIMIT"] + + def test_a_negative_or_boolean_setting_falls_back_to_the_default(monkeypatch): """A limit below one would refuse everyone; a bool is a typo, not a count.""" from litellm.proxy import proxy_server as ps From 9c541d9ce3ecf31ca8673ebc4472250f477b55ae Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 15:27:33 +0000 Subject: [PATCH 014/253] fix(ui): show internal user email in logs table and log detail drawer Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../(dashboard)/hooks/users/useUsers.test.ts | 61 ++++++++++++++++++- .../app/(dashboard)/hooks/users/useUsers.ts | 18 ++++++ .../LogDetailContent.integration.test.tsx | 23 +++++++ .../LogDetailsDrawer/LogDetailContent.tsx | 23 ++++++- .../LogDetailsDrawer.test.tsx | 39 +++++++++++- .../LogDetailsDrawer/LogDetailsDrawer.tsx | 3 + .../components/view_logs/RequestLogsTable.tsx | 10 ++- .../RequestLogsTableColumns.test.tsx | 21 +++++++ .../view_logs/RequestLogsTableColumns.tsx | 23 ++++++- 9 files changed, 212 insertions(+), 9 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts index dd7209140a9..2e8471ba84f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts @@ -2,7 +2,7 @@ import { describe, it, expect, vi, beforeEach } from "vitest"; import { renderHook, waitFor } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import React, { ReactNode } from "react"; -import { useInfiniteUsers, useUserLookup } from "./useUsers"; +import { useInfiniteUsers, useUserEmailLookup, useUserLookup } from "./useUsers"; import { userListCall } from "@/components/networking"; import type { UserListResponse } from "@/components/networking"; @@ -335,3 +335,62 @@ describe("useUserLookup", () => { expect(userListCall).not.toHaveBeenCalled(); }); }); + +describe("useUserEmailLookup", () => { + let queryClient: QueryClient; + + beforeEach(() => { + queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + vi.clearAllMocks(); + mockUseAuthorized.mockReturnValue(DEFAULT_AUTH); + }); + + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + + it("fetches the distinct ids in one call and maps each id to its email", async () => { + const response = buildUserListResponse(1, 1, 2); + vi.mocked(userListCall).mockResolvedValue(response); + + const { result } = renderHook(() => useUserEmailLookup(["user-1-1", "user-1-0", "user-1-1", ""]), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(userListCall).toHaveBeenCalledTimes(1); + expect(userListCall).toHaveBeenCalledWith("test-access-token", ["user-1-0", "user-1-1"], 1, 2); + expect(result.current.data).toEqual({ + "user-1-0": "user-1-0@example.com", + "user-1-1": "user-1-1@example.com", + }); + }); + + it("omits users that have no email so callers fall back to the id", async () => { + const response = buildUserListResponse(1, 1, 2); + vi.mocked(userListCall).mockResolvedValue({ + ...response, + users: [{ ...response.users[0], user_email: "" }, response.users[1]], + }); + + const { result } = renderHook(() => useUserEmailLookup(["user-1-0", "user-1-1"]), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(result.current.data).toEqual({ "user-1-1": "user-1-1@example.com" }); + }); + + it("does not query with no ids", async () => { + const { result } = renderHook(() => useUserEmailLookup([]), { wrapper }); + + await new Promise((resolve) => setTimeout(resolve, 20)); + expect(result.current.fetchStatus).toBe("idle"); + expect(userListCall).not.toHaveBeenCalled(); + }); + + it("does not query for a non-admin role", async () => { + mockUseAuthorized.mockReturnValue({ ...DEFAULT_AUTH, userRole: "Internal User" }); + + const { result } = renderHook(() => useUserEmailLookup(["user-1-0"]), { wrapper }); + + await new Promise((resolve) => setTimeout(resolve, 20)); + expect(result.current.fetchStatus).toBe("idle"); + expect(userListCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts index 011e43777b5..4a28e7ff3f2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts @@ -38,6 +38,24 @@ export const useInfiniteUsers = (pageSize: number = DEFAULT_PAGE_SIZE, searchEma }); }; +const USER_LIST_MAX_PAGE_SIZE = 100; + +export const useUserEmailLookup = (userIds: readonly string[]) => { + const { accessToken, userRole } = useAuthorized(); + const distinctIds = Array.from(new Set(userIds.filter((id) => id !== ""))).sort(); + return useQuery>({ + queryKey: userLookupKeys.list({ filters: { ids: distinctIds.join(",") } }), + queryFn: async () => { + const ids = distinctIds.slice(0, USER_LIST_MAX_PAGE_SIZE); + const response = await userListCall(accessToken!, ids, 1, ids.length); + return Object.fromEntries( + response.users.filter((user) => Boolean(user.user_email)).map((user) => [user.user_id, user.user_email]), + ); + }, + enabled: Boolean(accessToken) && distinctIds.length > 0 && all_admin_roles.includes(userRole!), + }); +}; + export const useUserLookup = (userId: string | null) => { const { accessToken, userRole } = useAuthorized(); return useQuery({ diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx index 721525268b3..793bfa7c246 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.integration.test.tsx @@ -56,6 +56,29 @@ describe("LogDetailContent", () => { expect(screen.getByText("completion")).toBeInTheDocument(); }); + it("shows the requesting user's email and id in Request Details when the email is resolved", () => { + render( + , + ); + + expect(screen.getByText("User")).toBeInTheDocument(); + expect(screen.getByText("alice@example.com")).toBeInTheDocument(); + expect(screen.getByText("106514937785257944828")).toBeInTheDocument(); + }); + + it("falls back to the user id in Request Details when no email is resolved", () => { + render(); + + expect(screen.getByText("User")).toBeInTheDocument(); + expect(screen.getByText("106514937785257944828")).toBeInTheDocument(); + }); + + it("omits the User row when the log has no internal user", () => { + render(); + + expect(screen.queryByText("User")).not.toBeInTheDocument(); + }); + it("should display error alert when request has failed", () => { render( {logEntry.model} {logEntry.custom_llm_provider || "-"} {logEntry.call_type} + {logEntry.user && ( + + + + )} @@ -333,6 +344,16 @@ function TagsSection({ tags }: { tags: Record }) { ); } +function UserIdentity({ userId, email }: { userId: string; email?: string }) { + if (!email || email === userId) return ; + return ( + + {email} + + + ); +} + function GuardrailLabel({ label, maskedCount }: { label: string; maskedCount: number }) { const handleClick = () => { const el = document.getElementById("guardrail-section"); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx index b96e5279722..f5bf7cda951 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx @@ -14,8 +14,13 @@ vi.mock("@/app/(dashboard)/hooks/logDetails/useLogDetails", () => ({ useLogDetails: () => ({ data: null, isLoading: false }), })); +const mockUseUserLookup = vi.fn(() => ({ data: undefined })); +vi.mock("@/app/(dashboard)/hooks/users/useUsers", () => ({ + useUserLookup: (userId: string | null) => mockUseUserLookup(userId), +})); + vi.mock("./LogDetailContent", () => ({ - LogDetailContent: () => null, + LogDetailContent: ({ userEmail }: { userEmail?: string }) => user-email:{userEmail ?? "none"}, GuardrailJumpLink: () => null, })); @@ -124,6 +129,38 @@ describe("LogDetailsDrawer session sidebar sorting", () => { }); }); +describe("LogDetailsDrawer internal user email", () => { + const renderSingleLog = (user: string | undefined) => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + {}} + logEntry={makeLog({ request_id: "single", user })} + accessToken="token" + /> + , + ); + }; + + it("looks up the log's internal user and hands the resolved email to the detail content", () => { + mockUseUserLookup.mockReturnValue({ data: { user_id: "u-1", user_email: "alice@example.com" } }); + renderSingleLog("u-1"); + + expect(mockUseUserLookup).toHaveBeenCalledWith("u-1"); + expect(screen.getByText("user-email:alice@example.com")).toBeInTheDocument(); + }); + + it("skips the lookup and passes no email when the log has no internal user", () => { + mockUseUserLookup.mockReturnValue({ data: undefined }); + renderSingleLog(undefined); + + expect(mockUseUserLookup).toHaveBeenCalledWith(null); + expect(screen.getByText("user-email:none")).toBeInTheDocument(); + }); +}); + describe("LogDetailsDrawer session sidebar auto-router icon", () => { const routedSessionLogs = [ makeLog({ request_id: "routed", model: "claude-opus-4-8", model_group: "smart-router" }), diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index ddd0a650c04..dc9207d59eb 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -17,6 +17,7 @@ import { getSpendString } from "@/utils/dataUtils"; import { normalizeGuardrailEntries, sortSessionLogs, SessionLogSortMode } from "./utils"; import { DRAWER_WIDTH } from "./constants"; import { useLogDetails } from "@/app/(dashboard)/hooks/logDetails/useLogDetails"; +import { useUserLookup } from "@/app/(dashboard)/hooks/users/useUsers"; export interface LogDetailsDrawerProps { open: boolean; @@ -245,6 +246,7 @@ export function LogDetailsDrawer({ const logDetails = useLogDetails(currentLog?.request_id, startTime, open && !!currentLog?.request_id); const detailsData = logDetails.data as any; const isLoadingDetails = logDetails.isLoading; + const { data: logUser } = useUserLookup(open && currentLog?.user ? currentLog.user : null); // Build an enriched log entry that merges lazy-loaded details. // The list endpoint may already include messages/response when store_prompts_in_spend_logs is enabled, @@ -465,6 +467,7 @@ export function LogDetailsDrawer({ logEntry={enrichedLog} isLoadingDetails={isLoadingDetails} accessToken={accessToken ?? null} + userEmail={logUser?.user_email || undefined} /> diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx index c3d204e1ae3..5a9bac428c8 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTable.tsx @@ -4,6 +4,7 @@ import type { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } fr import { ScrollText } from "lucide-react"; import { useMemo, useState, type ReactNode } from "react"; +import { useUserEmailLookup } from "@/app/(dashboard)/hooks/users/useUsers"; import { DataTable, DataTableFilterDrawer, DataTableToolbar } from "@/components/shared/DataTable"; import type { Team } from "../key_team_helpers/key_list"; @@ -73,10 +74,13 @@ export function RequestLogsTable({ }: RequestLogsTableProps) { const [filtersOpen, setFiltersOpen] = useState(false); + const userIds = useMemo(() => data.flatMap((log) => (log.user ? [log.user] : [])), [data]); + const { data: emailByUserId } = useUserEmailLookup(userIds); + const columns = useMemo(() => { - const deps = { onKeyHashClick, onSessionClick }; - return getRequestLogsTableColumns(deps); - }, [onKeyHashClick, onSessionClick]); + const resolveUserEmail = (userId: string) => emailByUserId?.[userId]; + return getRequestLogsTableColumns({ onKeyHashClick, onSessionClick, resolveUserEmail }); + }, [onKeyHashClick, onSessionClick, emailByUserId]); const isFiltered = columnFilters.length > 0 || searchValue !== ""; diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx index 9f0e659cb1f..08853814269 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.test.tsx @@ -75,6 +75,27 @@ describe("Cost column", () => { }); }); +describe("Internal User column", () => { + const emailById: Record = { "106514937785257944828": "alice@example.com" }; + const deps = { ...noopDeps, resolveUserEmail: (userId: string) => emailById[userId] }; + + it("shows the user's email instead of the raw id, with both in the tooltip", async () => { + const user = userEvent.setup(); + renderRows([logEntry({ request_id: "req-known-user", user: "106514937785257944828" })], deps); + + const emailCell = screen.getByText("alice@example.com"); + expect(screen.queryByText("106514937785257944828")).not.toBeInTheDocument(); + await user.hover(emailCell); + expect(await screen.findByText("alice@example.com (106514937785257944828)")).toBeInTheDocument(); + }); + + it("falls back to the raw id when no email is known for the user", () => { + renderRows([logEntry({ request_id: "req-unknown-user", user: "unknown-user-id" })], deps); + + expect(screen.getByText("unknown-user-id")).toBeInTheDocument(); + }); +}); + describe("Tokens column", () => { const sessionRow: Partial = { request_id: "req-session-tokens", diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx index 1ec1087a1a4..dd83ad6eb05 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsTableColumns.tsx @@ -15,6 +15,7 @@ import { AgentBadge, AgentIcon, BatchBadge, LlmBadge, McpBadge, SparkleIcon, Wre export interface RequestLogsTableColumnsDeps { onKeyHashClick: (keyHash: string) => void; onSessionClick: (log: LogEntry) => void; + resolveUserEmail?: (userId: string) => string | undefined; } const readMetaString = (metadata: Record | undefined, key: string): string | undefined => { @@ -32,14 +33,25 @@ const readMcpLogoUrl = (metadata: Record | undefined): string | const getLogoUrl = (row: LogEntry, provider: string): string => readMcpLogoUrl(row.metadata) ?? (provider ? getProviderLogoAndName(provider).logo : ""); -function TruncatedText({ value }: { value: string | undefined }) { +function TruncatedText({ value, tooltip }: { value: string | undefined; tooltip?: string }) { const display = value ?? "-"; - return {display}} />; + return ( + {display}} + /> + ); +} + +function UserCell({ userId, email }: { userId: string | undefined; email: string | undefined }) { + if (!userId || !email || email === userId) return ; + return ; } export const getRequestLogsTableColumns = ({ onKeyHashClick, onSessionClick, + resolveUserEmail = () => undefined, }: RequestLogsTableColumnsDeps): ColumnDef[] => [ { id: "startTime", @@ -313,7 +325,12 @@ export const getRequestLogsTableColumns = ({ header: "Internal User", size: 150, enableSorting: false, - cell: ({ row }) => , + cell: ({ row }) => ( + + ), }, { id: "end_user", From bde96e3197bf1d1d9af07fa238f63649b543adf0 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 15:36:29 +0000 Subject: [PATCH 015/253] fix(ui): keep user id boundaries in email lookup query key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/app/(dashboard)/hooks/users/useUsers.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts index 4a28e7ff3f2..84aeba90ae2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts @@ -44,7 +44,7 @@ export const useUserEmailLookup = (userIds: readonly string[]) => { const { accessToken, userRole } = useAuthorized(); const distinctIds = Array.from(new Set(userIds.filter((id) => id !== ""))).sort(); return useQuery>({ - queryKey: userLookupKeys.list({ filters: { ids: distinctIds.join(",") } }), + queryKey: userLookupKeys.list({ filters: { ids: JSON.stringify(distinctIds) } }), queryFn: async () => { const ids = distinctIds.slice(0, USER_LIST_MAX_PAGE_SIZE); const response = await userListCall(accessToken!, ids, 1, ids.length); From 3630642110e0055040f5c5666cb2ff1c195d519e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 16:21:20 +0000 Subject: [PATCH 016/253] fix(ui): allow Org Admin session role to resolve user emails in logs Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/app/(dashboard)/hooks/users/useUsers.test.ts | 12 +++++++++++- .../src/app/(dashboard)/hooks/users/useUsers.ts | 8 ++++---- ui/litellm-dashboard/src/utils/roles.ts | 6 ++++++ 3 files changed, 21 insertions(+), 5 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts index 2e8471ba84f..f49c728446b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.test.ts @@ -235,7 +235,7 @@ describe("useInfiniteUsers", () => { }); it("should execute query for each admin role", async () => { - const adminRoles = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer", "org_admin"]; + const adminRoles = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer", "org_admin", "Org Admin"]; for (const role of adminRoles) { vi.clearAllMocks(); @@ -384,6 +384,16 @@ describe("useUserEmailLookup", () => { expect(userListCall).not.toHaveBeenCalled(); }); + it("queries for the formatted Org Admin session role", async () => { + mockUseAuthorized.mockReturnValue({ ...DEFAULT_AUTH, userRole: "Org Admin" }); + vi.mocked(userListCall).mockResolvedValue(buildUserListResponse(1, 1, 1)); + + const { result } = renderHook(() => useUserEmailLookup(["user-1-0"]), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(result.current.data).toEqual({ "user-1-0": "user-1-0@example.com" }); + }); + it("does not query for a non-admin role", async () => { mockUseAuthorized.mockReturnValue({ ...DEFAULT_AUTH, userRole: "Internal User" }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts index 84aeba90ae2..3b7f9fbeb02 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/users/useUsers.ts @@ -1,7 +1,7 @@ import { userListCall, UserInfo, UserListResponse } from "@/components/networking"; import { useInfiniteQuery, useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; -import { all_admin_roles } from "@/utils/roles"; +import { canListUsers } from "@/utils/roles"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; const infiniteUsersKeys = createQueryKeys("infiniteUsers"); @@ -34,7 +34,7 @@ export const useInfiniteUsers = (pageSize: number = DEFAULT_PAGE_SIZE, searchEma } return undefined; }, - enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), + enabled: Boolean(accessToken) && canListUsers(userRole), }); }; @@ -52,7 +52,7 @@ export const useUserEmailLookup = (userIds: readonly string[]) => { response.users.filter((user) => Boolean(user.user_email)).map((user) => [user.user_id, user.user_email]), ); }, - enabled: Boolean(accessToken) && distinctIds.length > 0 && all_admin_roles.includes(userRole!), + enabled: Boolean(accessToken) && distinctIds.length > 0 && canListUsers(userRole), }); }; @@ -64,6 +64,6 @@ export const useUserLookup = (userId: string | null) => { const response = await userListCall(accessToken!, [userId!], 1, 1); return response.users.find((user) => user.user_id === userId) ?? null; }, - enabled: Boolean(accessToken) && Boolean(userId) && all_admin_roles.includes(userRole!), + enabled: Boolean(accessToken) && Boolean(userId) && canListUsers(userRole), }); }; diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 62a5f02cc39..1cb7e75c19b 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -28,6 +28,12 @@ export const isAdminRole = (role: string): boolean => { return all_admin_roles.includes(role); }; +// /user/list admits proxy admins and org admins; the session role for the latter is the formatted +// "Org Admin", which all_admin_roles does not carry +const rolesAllowedToListUsers: string[] = [...all_admin_roles, "Org Admin"]; + +export const canListUsers = (role: string | null): boolean => rolesAllowedToListUsers.includes(role ?? ""); + export const isProxyAdminRole = (role: string): boolean => { return role === "proxy_admin" || role === "Admin"; }; From b7596d6fba97f712a2dca3ead9ed280af3654292 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 17:16:30 +0000 Subject: [PATCH 017/253] refactor(ui): drop redundant comment on canListUsers role list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/src/utils/roles.ts | 2 -- 1 file changed, 2 deletions(-) diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 1cb7e75c19b..066a83992b4 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -28,8 +28,6 @@ export const isAdminRole = (role: string): boolean => { return all_admin_roles.includes(role); }; -// /user/list admits proxy admins and org admins; the session role for the latter is the formatted -// "Org Admin", which all_admin_roles does not carry const rolesAllowedToListUsers: string[] = [...all_admin_roles, "Org Admin"]; export const canListUsers = (role: string | null): boolean => rolesAllowedToListUsers.includes(role ?? ""); From fcca3c239e1683a6f1a960e87fff2e7ce34eaffc Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 22:50:52 +0000 Subject: [PATCH 018/253] fix(proxy): count failed sign-ins in Redis alone while it answers Every worker spends one shared budget and a successful sign-in clears it for all of them. This worker's own counter is only consulted while Redis raises, so an outage degrades to per-worker accounting instead of switching the control off Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 48 ++++++------ .../proxy/auth/test_login_utils.py | 77 +++++++++++++++++++ 2 files changed, 101 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 11290ef6d40..0f5b7e64a57 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -11,7 +11,7 @@ startup and can be reassigned later. import asyncio import hashlib -from collections.abc import Awaitable +from collections.abc import Awaitable, Callable from dataclasses import dataclass from functools import cache from types import MappingProxyType @@ -60,6 +60,7 @@ def _bounded_store(max_entries: int) -> DualCache: _FAILED_LOGIN_USERNAME_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_USERNAMES) _FAILED_LOGIN_SOURCE_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_SOURCES) _NO_SETTINGS: Final = MappingProxyType({}) +_UNAVAILABLE: Final = object() _DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} # mutable-ok: per-source slots taken and released around each held delay @@ -188,16 +189,22 @@ class LoginThrottle: return await work except Exception as exc: # noqa: BLE001 # an unreachable cache must never deny a valid credential verbose_proxy_logger.warning("login attempt accounting unavailable: %s", exc) - return None + return _UNAVAILABLE - async def _failures(self, store: DualCache, key: str) -> int: - """The larger of the shared and the process-local count, so a Redis outage degrades - to per-worker accounting instead of switching the control off.""" - local: Final = _as_count(await self._outcome(store.async_get_cache(key=key))) + async def _shared(self, work: Callable[[RedisCache], Awaitable[object]]) -> object: + """The Redis result, or ``_UNAVAILABLE`` when Redis is not configured or the call raised.""" redis_cache: Final = self.redis_cache if redis_cache is None: - return local - return max(local, _as_count(await self._outcome(redis_cache.async_get_cache(key)))) + return _UNAVAILABLE + return await self._outcome(work(redis_cache)) + + async def _failures(self, store: DualCache, key: str) -> int: + """The shared count while Redis answers, so every worker sees one budget; this worker's + own count only while it does not, so an outage degrades to per-worker accounting.""" + shared: Final = await self._shared(lambda redis_cache: redis_cache.async_get_cache(key)) + if shared is not _UNAVAILABLE: + return _as_count(shared) + return _as_count(await self._outcome(store.async_get_cache(key=key))) async def _remaining_window(self, key: str) -> int: """Seconds until this counter expires. @@ -206,13 +213,12 @@ class LoginThrottle: was stripped out of band (PERSIST, a restore). It is given the full window again, since nothing increments a key once the limit is reached. """ - redis_cache: Final = self.redis_cache - if redis_cache is None: + if self.redis_cache is None: return self.window_seconds - ttl: Final = await self._outcome(redis_cache.async_get_ttl(key)) + ttl: Final = await self._shared(lambda redis_cache: redis_cache.async_get_ttl(key)) if isinstance(ttl, int) and ttl > 0: return min(ttl, self.window_seconds) - await self._outcome(redis_cache.async_increment_with_floor(key, 0, self.window_seconds)) + await self._shared(lambda redis_cache: redis_cache.async_increment_with_floor(key, 0, self.window_seconds)) return self.window_seconds def _refused(self, retry_after: int, param: str) -> ProxyException: @@ -260,16 +266,12 @@ class LoginThrottle: ) async def _bump(self, store: DualCache, key: str) -> int: - local: Final = _as_count( - await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) + shared: Final = await self._shared( + lambda redis_cache: redis_cache.async_increment_with_floor(key, 1, self.window_seconds) ) - redis_cache: Final = self.redis_cache - if redis_cache is None: - return local - shared: Final = _as_count( - await self._outcome(redis_cache.async_increment_with_floor(key, 1, self.window_seconds)) - ) - return max(local, shared) + if shared is not _UNAVAILABLE: + return _as_count(shared) + return _as_count(await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds))) async def record_failure(self, username: str) -> FailureCounts: """Count one rejected credential guess against this username and against this source.""" @@ -331,7 +333,5 @@ class LoginThrottle: if not self.enabled: return key: Final = self._username_key(username) - redis_cache: Final = self.redis_cache - if redis_cache is not None: - await self._outcome(redis_cache.async_delete_cache(key)) + await self._shared(lambda redis_cache: redis_cache.async_delete_cache(key)) await self._outcome(self.username_cache.async_delete_cache(key=key)) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 8afac8a622d..70e7125485c 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1203,6 +1203,83 @@ async def test_counters_are_written_with_their_expiry_and_re_armed_if_stripped(m assert set(redis.ttls) >= {k for k in redis.values if ":user:" in k}, "the refusal must re-arm a stripped expiry" +class _DownRedis(_FakeRedis): + """Redis whose every call raises, as during an outage or an open circuit breaker.""" + + async def async_get_cache(self, key, **kwargs): + raise ConnectionError("redis is down") + + async def async_increment_with_floor(self, key, value, ttl): + raise ConnectionError("redis is down") + + async def async_get_ttl(self, key): + raise ConnectionError("redis is down") + + async def async_delete_cache(self, key): + raise ConnectionError("redis is down") + + +@pytest.mark.asyncio +async def test_redis_is_the_only_counter_while_it_answers(monkeypatch): + """Regression: every worker must spend the same budget, and a success must clear it for all. + + Counting in this worker's memory as well as in Redis let the two drift apart: a worker + whose Redis write failed kept its own count while the others gave the attacker fresh + guesses, and a stale local count outlived the shared clear after a correct password. + """ + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + redis = _FakeRedis() + first_worker_store = DualCache() + second_worker_store = DualCache() + first_worker = _throttle(max_attempts=2, cache=first_worker_store, redis_cache=redis) + second_worker = _throttle(max_attempts=2, cache=second_worker_store, redis_cache=redis) + + for _ in range(2): + with pytest.raises(ProxyException, match="Invalid credentials"): + await _guess(first_worker) + + assert not [k for k in first_worker_store.in_memory_cache.cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)], ( + "with Redis answering, no worker may keep a counter of its own" + ) + with pytest.raises(ProxyException) as blocked: + await _guess(second_worker) + assert blocked.value.code == "429", "the second worker must see the budget the first one spent" + + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ): + await _guess(second_worker, password="right") + + assert not [k for k in redis.values if ":user:" in k], "a success must clear the shared username counter" + with pytest.raises(ProxyException, match="Invalid credentials"): + await _guess(first_worker) + + +@pytest.mark.asyncio +async def test_a_redis_outage_falls_back_to_this_workers_own_counter(monkeypatch): + """With Redis raising, guesses are still counted and refused, per worker, instead of unbounded.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(max_attempts=2, redis_cache=_DownRedis()) + + for _ in range(2): + with pytest.raises(ProxyException, match="Invalid credentials"): + await _guess(throttle) + + with pytest.raises(ProxyException) as blocked: + await _guess(throttle) + assert blocked.value.code == "429" + assert blocked.value.headers.get("Retry-After") == "900" + + @pytest.mark.asyncio async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): """Regression: throttle entries must not evict cached credentials. From 383edfe9532dfbff6827e579fca5e97d0feb8e23 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 23:25:33 +0000 Subject: [PATCH 019/253] fix(proxy): keep counting the failed sign-ins Redis missed once it answers again A guess is recorded in exactly one place, Redis or this worker's own store when Redis refused it, so the count is the sum of the two. Redis is read through async_batch_get_counts, which raises on failure, instead of async_get_cache, which swallows it into None and read as an empty counter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 26 ++++++--- .../proxy/auth/test_login_utils.py | 56 ++++++++++++++++++- 2 files changed, 72 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 0f5b7e64a57..ca73bb4f26b 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -199,12 +199,18 @@ class LoginThrottle: return await self._outcome(work(redis_cache)) async def _failures(self, store: DualCache, key: str) -> int: - """The shared count while Redis answers, so every worker sees one budget; this worker's - own count only while it does not, so an outage degrades to per-worker accounting.""" - shared: Final = await self._shared(lambda redis_cache: redis_cache.async_get_cache(key)) - if shared is not _UNAVAILABLE: - return _as_count(shared) - return _as_count(await self._outcome(store.async_get_cache(key=key))) + """The shared count plus this worker's own. + + A failure is written to exactly one of the two: Redis, or this worker's store when Redis + refused it. So the local store is empty while Redis is healthy, and once Redis answers + again the guesses it missed still count. Read through ``async_batch_get_counts`` because + ``async_get_cache`` turns a failed GET into ``None``, which would pass as an empty counter. + """ + local: Final = _as_count(await self._outcome(store.async_get_cache(key=key))) + shared: Final = await self._shared(lambda redis_cache: redis_cache.async_batch_get_counts([key])) + if not isinstance(shared, tuple): + return local + return _as_count(shared[0]) + local async def _remaining_window(self, key: str) -> int: """Seconds until this counter expires. @@ -269,9 +275,11 @@ class LoginThrottle: shared: Final = await self._shared( lambda redis_cache: redis_cache.async_increment_with_floor(key, 1, self.window_seconds) ) - if shared is not _UNAVAILABLE: - return _as_count(shared) - return _as_count(await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds))) + if shared is _UNAVAILABLE: + return _as_count( + await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) + ) + return _as_count(shared) + _as_count(await self._outcome(store.async_get_cache(key=key))) async def record_failure(self, username: str) -> FailureCounts: """Count one rejected credential guess against this username and against this source.""" diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 70e7125485c..14f8dd19683 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1156,6 +1156,9 @@ class _FakeRedis: async def async_get_cache(self, key, **kwargs): return self.values.get(key) + async def async_batch_get_counts(self, key_list): + return tuple(self.values.get(key) for key in key_list) + async def async_increment_with_floor(self, key, value, ttl): self.values[key] = self.values.get(key, 0) + value self.ttls.setdefault(key, ttl) @@ -1204,9 +1207,16 @@ async def test_counters_are_written_with_their_expiry_and_re_armed_if_stripped(m class _DownRedis(_FakeRedis): - """Redis whose every call raises, as during an outage or an open circuit breaker.""" + """Redis whose every call fails, as during an outage or an open circuit breaker. + + `async_get_cache` returns None rather than raising, as the real one does: it swallows the + error, so a failed GET is indistinguishable from an empty key to anyone reading through it. + """ async def async_get_cache(self, key, **kwargs): + return None + + async def async_batch_get_counts(self, key_list): raise ConnectionError("redis is down") async def async_increment_with_floor(self, key, value, ttl): @@ -1280,6 +1290,50 @@ async def test_a_redis_outage_falls_back_to_this_workers_own_counter(monkeypatch assert blocked.value.headers.get("Retry-After") == "900" +class _WriteRefusingRedis(_FakeRedis): + """Redis that answers reads but raises on writes until `recover()` is called.""" + + def __init__(self): + super().__init__() + self.writable = False + + def recover(self): + self.writable = True + + async def async_increment_with_floor(self, key, value, ttl): + if not self.writable: + raise ConnectionError("redis write failed") + return await super().async_increment_with_floor(key, value, ttl) + + +@pytest.mark.asyncio +async def test_failures_redis_refused_still_count_once_redis_recovers(monkeypatch): + """Regression: a guess Redis could not record must not be forgotten when Redis comes back. + + Such a guess lands in this worker's own store. Reading only Redis afterwards handed the + attacker that guess again, so the budget was the limit plus however many writes failed. + """ + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + redis = _WriteRefusingRedis() + throttle = _throttle(max_attempts=2, redis_cache=redis) + + with pytest.raises(ProxyException, match="Invalid credentials"): + await _guess(throttle) + assert not redis.values, "the refused write must not have reached Redis" + + redis.recover() + with pytest.raises(ProxyException, match="Invalid credentials"): + await _guess(throttle) + assert [v for k, v in redis.values.items() if ":user:" in k] == [1], "only the recorded guess is in Redis" + + with pytest.raises(ProxyException) as blocked: + await _guess(throttle) + assert blocked.value.code == "429", "the guess Redis missed and the one it took must add up to the limit" + + @pytest.mark.asyncio async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): """Regression: throttle entries must not evict cached credentials. From 04b7b716561551814b641bd8ce0fa67b4cddad1f Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 23:37:22 +0000 Subject: [PATCH 020/253] refactor(proxy): read the shared sign-in counter through a tuple of keys Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_cache.py | 4 ++-- litellm/proxy/auth/login_throttle.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 6b93529e456..f37ca8a23a1 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -791,7 +791,7 @@ class RedisCache(BaseCache): return _LUA_COUNT.validate_python(count) @_redis_circuit_breaker_guard_sync - def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]: + def batch_get_counts(self, key_list: Sequence[str]) -> tuple[int | None, ...]: """Read integer counters for ``key_list``, in order, raising when Redis cannot answer. ``batch_get_cache`` swallows every failure and returns an empty dict, which the caller @@ -802,7 +802,7 @@ class RedisCache(BaseCache): return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys)) @_redis_circuit_breaker_guard - async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]: + async def async_batch_get_counts(self, key_list: Sequence[str]) -> tuple[int | None, ...]: """Async twin of ``batch_get_counts``, raising on failure the same way.""" namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list] return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys)) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index ca73bb4f26b..5a56628bed2 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -207,7 +207,7 @@ class LoginThrottle: ``async_get_cache`` turns a failed GET into ``None``, which would pass as an empty counter. """ local: Final = _as_count(await self._outcome(store.async_get_cache(key=key))) - shared: Final = await self._shared(lambda redis_cache: redis_cache.async_batch_get_counts([key])) + shared: Final = await self._shared(lambda redis_cache: redis_cache.async_batch_get_counts((key,))) if not isinstance(shared, tuple): return local return _as_count(shared[0]) + local From 49d49b70508a3e15f360a5da4e754001e8eb36eb Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 23:48:32 +0000 Subject: [PATCH 021/253] refactor(caching): spell out the key collections batch_get_counts accepts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_cache.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index f37ca8a23a1..f1b723de625 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -791,7 +791,7 @@ class RedisCache(BaseCache): return _LUA_COUNT.validate_python(count) @_redis_circuit_breaker_guard_sync - def batch_get_counts(self, key_list: Sequence[str]) -> tuple[int | None, ...]: + def batch_get_counts(self, key_list: list[str] | tuple[str, ...]) -> tuple[int | None, ...]: """Read integer counters for ``key_list``, in order, raising when Redis cannot answer. ``batch_get_cache`` swallows every failure and returns an empty dict, which the caller @@ -802,7 +802,7 @@ class RedisCache(BaseCache): return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys)) @_redis_circuit_breaker_guard - async def async_batch_get_counts(self, key_list: Sequence[str]) -> tuple[int | None, ...]: + async def async_batch_get_counts(self, key_list: list[str] | tuple[str, ...]) -> tuple[int | None, ...]: """Async twin of ``batch_get_counts``, raising on failure the same way.""" namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list] return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys)) From 97dbd2dfbf3de0386efbe77057bbf75bc1aa1336 Mon Sep 17 00:00:00 2001 From: tusharjamunkar Date: Sat, 12 Sep 2026 22:16:19 +0530 Subject: [PATCH 022/253] fix(gemini): preserve candidates with finishReason and no content (#40477) --- litellm/litellm_core_utils/core_helpers.py | 1 + .../adapters/transformation.py | 2 + .../vertex_and_google_ai_studio_gemini.py | 30 ++-- .../transformation.py | 20 ++- ...test_vertex_and_google_ai_studio_gemini.py | 146 ++++++++++++++++++ 5 files changed, 184 insertions(+), 15 deletions(-) diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index aa7d6ca1699..66180b165f8 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -224,6 +224,7 @@ _FINISH_REASON_MAP: Final[dict[str, OpenAIChatCompletionFinishReason]] = { "IMAGE_PROHIBITED_CONTENT": "content_filter", "TOO_MANY_TOOL_CALLS": "stop", "MALFORMED_RESPONSE": "stop", + "NO_IMAGE": "content_filter", # Zhipu GLM "network_error": "stop", "sensitive": "content_filter", diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 8ff9f2e0679..f4cc569bcef 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -1367,6 +1367,8 @@ class LiteLLMAnthropicMessagesAdapter: return "max_tokens" elif openai_finish_reason == "tool_calls": return "tool_use" + elif openai_finish_reason in ["content_filter", "refusal"]: + return "refusal" return "end_turn" @staticmethod 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 d113b2b4f6b..01d1f063b1b 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 @@ -1340,6 +1340,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "IMAGE_PROHIBITED_CONTENT", "TOO_MANY_TOOL_CALLS", "MALFORMED_RESPONSE", + "NO_IMAGE", } ) @@ -2224,22 +2225,23 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): grounding_metadata: Final[list[dict]] = [] url_context_metadata: Final[list[dict]] = [] - image_response: list[ImageURLListItem] | None = None safety_ratings: Final[list] = [] citation_metadata: Final[list] = [] - chat_completion_message: Final[ChatCompletionResponseMessage] = {"role": "assistant"} - chat_completion_logprobs: ChoiceLogprobs | None = None - tools: list[ChatCompletionToolCallChunk] | None = [] - functions: ChatCompletionToolCallFunctionChunk | None = None - thinking_blocks: list[ChatCompletionThinkingBlock] | None = None - reasoning_content: str | None = None - thought_signatures: Sequence[str] | None = None - server_side_tool_invocations: list[dict[str, object]] | None = None for idx, candidate in enumerate(_candidates): - if "content" not in candidate: + if "content" not in candidate and "finishReason" not in candidate: continue + image_response: list[ImageURLListItem] | None = None + chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} + chat_completion_logprobs: ChoiceLogprobs | None = None + tools: list[ChatCompletionToolCallChunk] | None = [] + functions: ChatCompletionToolCallFunctionChunk | None = None + thinking_blocks: list[ChatCompletionThinkingBlock] | None = None + reasoning_content: str | None = None + thought_signatures: Sequence[str] | None = None + server_side_tool_invocations: list[dict[str, object]] | None = None + # Extract metadata using helper function ( candidate_grounding_metadata, @@ -2253,7 +2255,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): safety_ratings.extend(candidate_safety_ratings) citation_metadata.extend(candidate_citation_metadata) - if "parts" in candidate["content"]: + if "content" in candidate and candidate["content"] and "parts" in candidate["content"]: ( content, reasoning_content, @@ -2348,6 +2350,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): tool_invocation_fields["server_side_tool_invocations"] = server_side_tool_invocations chat_completion_message["provider_specific_fields"] = tool_invocation_fields + if candidate.get("finishReason"): + finish_reason_fields = chat_completion_message.get("provider_specific_fields") or {} + finish_reason_fields["native_finish_reason"] = candidate.get("finishReason") + chat_completion_message["provider_specific_fields"] = finish_reason_fields + if isinstance(model_response, ModelResponseStream): choice = VertexGeminiConfig._create_streaming_choice( chat_completion_message=chat_completion_message, @@ -2368,6 +2375,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): message=chat_completion_message, logprobs=chat_completion_logprobs, enhancements=None, + provider_specific_fields=chat_completion_message.get("provider_specific_fields"), ) model_response.choices.append(choice) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index fca5b0d11cf..13119085e46 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2272,13 +2272,27 @@ class LiteLLMCompletionResponsesConfig: if choices and len(choices) > 0: finish_reason = choices[0].finish_reason + status: Final[ResponsesAPIStatus] = ( + LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( + finish_reason + ) + ) + incomplete_details = getattr(chat_completion_response, "incomplete_details", None) + if incomplete_details is None and status == "incomplete": + from openai.types.responses.response import IncompleteDetails + + if finish_reason == "length": + incomplete_details = IncompleteDetails(reason="max_output_tokens") + elif finish_reason in ["content_filter", "refusal"]: + incomplete_details = IncompleteDetails(reason="content_filter") + responses_api_response: Final[ResponsesAPIResponse] = ResponsesAPIResponse( id=chat_completion_response.id, created_at=chat_completion_response.created, model=chat_completion_response.model, object="response", error=getattr(chat_completion_response, "error", None), - incomplete_details=getattr(chat_completion_response, "incomplete_details", None), + incomplete_details=incomplete_details, instructions=getattr(chat_completion_response, "instructions", None), metadata=getattr(chat_completion_response, "metadata", {}), output=LiteLLMCompletionResponsesConfig._transform_chat_completion_choices_to_responses_output( @@ -2296,9 +2310,7 @@ class LiteLLMCompletionResponsesConfig: max_output_tokens=getattr(chat_completion_response, "max_output_tokens", None), previous_response_id=getattr(chat_completion_response, "previous_response_id", None), reasoning=None, - status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( - finish_reason - ), + status=status, text={}, truncation=getattr(chat_completion_response, "truncation", None), usage=LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( 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 101f6e6fa5d..a048ad4f171 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 @@ -5836,3 +5836,149 @@ def test_supported_reasoning_efforts_still_map(model): drop_params=False, ) assert "thinkingConfig" in result + + +def test_gemini_candidate_with_finish_reason_no_content_chat_completion(): + config = VertexGeminiConfig() + completion_response = { + "candidates": [ + { + "finishReason": "NO_IMAGE", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 19, + "candidatesTokenCount": 0, + "totalTokenCount": 19, + }, + } + model_response = ModelResponse() + logging_obj = MagicMock() + raw_response = MagicMock() + raw_response.headers = {} + + resp = config._transform_google_generate_content_to_openai_model_response( + completion_response=completion_response, + model_response=model_response, + model="gemini-2.5-flash-image", + logging_obj=logging_obj, + raw_response=raw_response, + ) + assert len(resp.choices) == 1 + assert resp.choices[0].finish_reason == "content_filter" + assert resp.choices[0].message.content is None + assert resp.choices[0].provider_specific_fields["native_finish_reason"] == "NO_IMAGE" + + +def test_gemini_candidate_with_finish_reason_no_content_anthropic_messages(): + from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( + LiteLLMAnthropicMessagesAdapter, + ) + + config = VertexGeminiConfig() + completion_response = { + "candidates": [ + { + "finishReason": "NO_IMAGE", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 19, + "candidatesTokenCount": 0, + "totalTokenCount": 19, + }, + } + resp = config._transform_google_generate_content_to_openai_model_response( + completion_response=completion_response, + model_response=ModelResponse(), + model="gemini-2.5-flash-image", + logging_obj=MagicMock(), + raw_response=MagicMock(headers={}), + ) + + adapter = LiteLLMAnthropicMessagesAdapter() + anthropic_resp = adapter.translate_openai_response_to_anthropic( + response=resp, + tool_name_mapping={}, + ) + assert anthropic_resp["stop_reason"] == "refusal" + assert anthropic_resp["content"] == [] + + +def test_gemini_candidate_with_finish_reason_no_content_responses_api(): + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + config = VertexGeminiConfig() + completion_response = { + "candidates": [ + { + "finishReason": "NO_IMAGE", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 19, + "candidatesTokenCount": 0, + "totalTokenCount": 19, + }, + } + resp = config._transform_google_generate_content_to_openai_model_response( + completion_response=completion_response, + model_response=ModelResponse(), + model="gemini-2.5-flash-image", + logging_obj=MagicMock(), + raw_response=MagicMock(headers={}), + ) + + responses_resp = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="Generate picture", + responses_api_request={}, + chat_completion_response=resp, + ) + assert responses_resp.status == "incomplete" + assert responses_resp.incomplete_details is not None + assert responses_resp.incomplete_details.reason == "content_filter" + + +def test_gemini_candidate_other_finish_reasons_no_content(): + from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( + LiteLLMAnthropicMessagesAdapter, + ) + from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, + ) + + config = VertexGeminiConfig() + max_tokens_response = { + "candidates": [{"finishReason": "MAX_TOKENS", "index": 0}], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 50, "totalTokenCount": 60}, + } + resp_length = config._transform_google_generate_content_to_openai_model_response( + completion_response=max_tokens_response, + model_response=ModelResponse(), + model="gemini-2.5-flash", + logging_obj=MagicMock(), + raw_response=MagicMock(headers={}), + ) + assert len(resp_length.choices) == 1 + assert resp_length.choices[0].finish_reason == "length" + assert resp_length.choices[0].provider_specific_fields["native_finish_reason"] == "MAX_TOKENS" + + anthropic_length = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic( + response=resp_length, + tool_name_mapping={}, + ) + assert anthropic_length["stop_reason"] == "max_tokens" + + responses_length = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="thinking request", + responses_api_request={}, + chat_completion_response=resp_length, + ) + assert responses_length.status == "incomplete" + assert responses_length.incomplete_details.reason == "max_output_tokens" + From cd66b34b45dace036126f4ea809df132407d907c Mon Sep 17 00:00:00 2001 From: tusharjamunkar Date: Sat, 12 Sep 2026 22:48:08 +0530 Subject: [PATCH 023/253] style(responses): apply ruff formatting to transformation.py --- .../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 13119085e46..10e56e85bff 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2273,9 +2273,7 @@ class LiteLLMCompletionResponsesConfig: finish_reason = choices[0].finish_reason status: Final[ResponsesAPIStatus] = ( - LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( - finish_reason - ) + LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status(finish_reason) ) incomplete_details = getattr(chat_completion_response, "incomplete_details", None) if incomplete_details is None and status == "incomplete": From e54399eff7a13e649b6a353486922166b1bf5688 Mon Sep 17 00:00:00 2001 From: tusharjamunkar Date: Sat, 12 Sep 2026 23:07:31 +0530 Subject: [PATCH 024/253] test: add direct coverage for content_filter and refusal in anthropic and responses adapters --- ...al_pass_through_adapters_transformation.py | 40 ++++++++++++ .../test_litellm_completion_responses.py | 62 +++++++++++++++++++ 2 files changed, 102 insertions(+) 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 03b9840b1c3..0e09e27f4db 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 @@ -102,6 +102,46 @@ def test_translate_chat_length_takes_precedence_over_refusal(): assert result.get("stop_details") is None +def test_translate_chat_content_filter_to_anthropic_response(): + response = ModelResponse( + id="chatcmpl-content-filter", + model="openai-model", + choices=[ + Choices( + index=0, + finish_reason="content_filter", + message=Message(content=None, role="assistant"), + ) + ], + usage=Usage(prompt_tokens=1, completion_tokens=0, total_tokens=1), + ) + + result = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(response) + + assert result["content"] == [] + assert result["stop_reason"] == "refusal" + + +def test_translate_chat_refusal_finish_reason_to_anthropic_response(): + response = ModelResponse( + id="chatcmpl-refusal-reason", + model="openai-model", + choices=[ + Choices( + index=0, + finish_reason="refusal", + message=Message(content=None, role="assistant"), + ) + ], + usage=Usage(prompt_tokens=1, completion_tokens=0, total_tokens=1), + ) + + result = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(response) + + assert result["content"] == [] + assert result["stop_reason"] == "refusal" + + def test_translate_streaming_openai_chunk_to_anthropic_content_block(): choices = [ StreamingChoices( diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 46249e50572..f2ebbf316c3 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -4246,3 +4246,65 @@ class TestStreamingSnapshotItemIds: reasoning_items = _bridged_output_items(completed_event.response, "reasoning") assert len(reasoning_items) == 1 assert reasoning_items[0].id == streamed_event.item_id + + +def test_transform_chat_completion_response_incomplete_details(): + from openai.types.responses.response import IncompleteDetails + + resp_length = ModelResponse( + id="resp-length", + choices=[Choices(index=0, finish_reason="length", message=Message(content="cutoff", role="assistant"))], + model="gpt-4o", + ) + result_length = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="test prompt", + responses_api_request={}, + chat_completion_response=resp_length, + ) + assert result_length.status == "incomplete" + assert result_length.incomplete_details is not None + assert result_length.incomplete_details.reason == "max_output_tokens" + + resp_filter = ModelResponse( + id="resp-filter", + choices=[Choices(index=0, finish_reason="content_filter", message=Message(content=None, role="assistant"))], + model="gpt-4o", + ) + result_filter = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="test prompt", + responses_api_request={}, + chat_completion_response=resp_filter, + ) + assert result_filter.status == "incomplete" + assert result_filter.incomplete_details is not None + assert result_filter.incomplete_details.reason == "content_filter" + + resp_refusal = ModelResponse( + id="resp-refusal", + choices=[Choices(index=0, finish_reason="refusal", message=Message(content=None, role="assistant"))], + model="gpt-4o", + ) + result_refusal = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="test prompt", + responses_api_request={}, + chat_completion_response=resp_refusal, + ) + assert result_refusal.status == "incomplete" + assert result_refusal.incomplete_details is not None + assert result_refusal.incomplete_details.reason == "content_filter" + + existing_details = IncompleteDetails(reason="content_filter") + resp_existing = ModelResponse( + id="resp-existing", + choices=[Choices(index=0, finish_reason="length", message=Message(content="cutoff", role="assistant"))], + model="gpt-4o", + ) + resp_existing.incomplete_details = existing_details + result_existing = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( + request_input="test prompt", + responses_api_request={}, + chat_completion_response=resp_existing, + ) + assert result_existing.status == "incomplete" + assert result_existing.incomplete_details == existing_details + From c3a7b7c3eeb750e4fc1c7479fc502a2d143e2d2f Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 13 Sep 2026 04:36:41 +0000 Subject: [PATCH 025/253] fix(proxy): warn about per-worker login counters even without general_settings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_server.py | 13 +++++----- tests/test_litellm/proxy/test_proxy_server.py | 24 +++++++++++++++++++ 2 files changed, 31 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 24532cc051c..138657fd5c3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -5748,6 +5748,13 @@ class ProxyConfig: general_settings = config.get("general_settings", {}) if general_settings is None: general_settings = {} + + ### FAILED-LOGIN ACCOUNTING MULTI-INSTANCE PREREQUISITE CHECK ### + # Failed Admin UI sign-in counters live in redis_usage_cache when available so a + # brute-force run is counted once across workers instead of once per worker. + if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: + warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) + _bg_hc_model_groups: Final = parse_background_health_check_model_groups(general_settings) _enable_hc_routing = False _hc_staleness = None @@ -5844,12 +5851,6 @@ class ProxyConfig: "or ensure sticky sessions for single-instance deployments." ) - ### FAILED-LOGIN ACCOUNTING MULTI-INSTANCE PREREQUISITE CHECK ### - # Failed Admin UI sign-in counters live in redis_usage_cache when available so a - # brute-force run is counted once across workers instead of once per worker. - if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: - warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) - ### STORE MODEL IN DB ### feature flag for `/model/new` store_model_in_db = general_settings.get("store_model_in_db", False) if store_model_in_db is None: diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 2f8d5cbc43a..9a41b3a7006 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3411,6 +3411,30 @@ async def test_load_config_user_url_validation_handles_null_and_string_false(tmp assert litellm.user_url_validation is False +@pytest.mark.asyncio +async def test_load_config_warns_per_worker_login_counters_without_general_settings(tmp_path, monkeypatch, caplog): + """Regression: the failed-login throttle is on by default, so a multi-worker proxy with no + Redis must hear that its counters are per worker even when the config has no general_settings.""" + import logging + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.auth.login_throttle import warn_login_counters_are_per_worker + from litellm.proxy.proxy_server import ProxyConfig + + for redis_var in ("REDIS_HOST", "REDIS_URL", "REDIS_CLUSTER_NODES", "REDIS_SENTINEL_NODES"): + monkeypatch.delenv(redis_var, raising=False) + monkeypatch.setenv("NUM_WORKERS", "4") + monkeypatch.setattr(proxy_server, "redis_usage_cache", None) + warn_login_counters_are_per_worker.cache_clear() + config_file = tmp_path / "config.yaml" + config_file.write_text("model_list: []\n") + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + + assert "Running 4 workers but Redis is not configured" in caplog.text + + @pytest.mark.asyncio async def test_load_environment_variables_direct_and_os_environ(): """ From 9ca6307b9d8fc60f557838ac6dc1b58d41f4c98d Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 13 Sep 2026 08:28:06 +0000 Subject: [PATCH 026/253] chore(ui): regenerate schema.d.ts after merging main Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 -- 1 file changed, 2 deletions(-) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f1ef4351a8b..08f8f773a87 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16781,7 +16781,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -16887,7 +16886,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) From ce82033f62c76a9d85332a171457a2b1fbbde757 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 13 Sep 2026 09:00:20 +0000 Subject: [PATCH 027/253] refactor(proxy): inject settings and Redis cache into LoginThrottle.from_request Removes the runtime import of proxy_server from login_throttle so the throttle module no longer participates in the import cycle CodeQL flagged (py/cyclic-import). Callers pass general_settings and redis_usage_cache explicitly; behavior is unchanged. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 12 +++--- litellm/proxy/proxy_server.py | 6 +-- .../proxy/auth/test_login_utils.py | 43 ++++++++----------- 3 files changed, 27 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 5a56628bed2..b6ec37afc0e 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -11,7 +11,7 @@ startup and can be reassigned later. import asyncio import hashlib -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass from functools import cache from types import MappingProxyType @@ -134,10 +134,10 @@ class LoginThrottle: enabled: bool = True @classmethod - def from_request(cls, request: Request) -> "LoginThrottle": - """Build the throttle for this request from the live proxy settings and caches.""" - from litellm.proxy.proxy_server import general_settings, redis_usage_cache - + def from_request( + cls, request: Request, general_settings: Mapping[str, object] | None, redis_cache: RedisCache | None + ) -> "LoginThrottle": + """Build the throttle for this request from the proxy's general_settings and shared Redis cache.""" settings: Final = general_settings or _NO_SETTINGS cidrs: Final = normalize_cidr_ranges( settings.get(TRUSTED_PROXY_RANGES_KEY), setting_name=TRUSTED_PROXY_RANGES_KEY @@ -167,7 +167,7 @@ class LoginThrottle: ), username_cache=_FAILED_LOGIN_USERNAME_CACHE, source_cache=_FAILED_LOGIN_SOURCE_CACHE, - redis_cache=redis_usage_cache, + redis_cache=redis_cache, enabled=not _rate_limit_disabled(), ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2e4417db696..37e7496610c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15853,7 +15853,7 @@ async def login(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, - throttle=LoginThrottle.from_request(request), + throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache), general_settings=general_settings, ) except ProxyException as exc: @@ -15946,7 +15946,7 @@ async def login_v2(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, - throttle=LoginThrottle.from_request(request), + throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache), general_settings=general_settings, ) @@ -16018,7 +16018,7 @@ async def login_v3(request: Request): password=password, master_key=master_key, prisma_client=prisma_client, - throttle=LoginThrottle.from_request(request), + throttle=LoginThrottle.from_request(request, general_settings, redis_usage_cache), general_settings=general_settings, ) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 14f8dd19683..7f5cd3e808f 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1348,7 +1348,6 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - monkeypatch.setattr(ps, "redis_usage_cache", None) auth_cache_keys_before = set(ps.user_api_key_cache.in_memory_cache.cache_dict) @@ -1356,7 +1355,7 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): request.headers = {} request.client = MagicMock() request.client.host = "1.2.3.4" - throttle = LoginThrottle.from_request(request) + throttle = LoginThrottle.from_request(request, general_settings={}, redis_cache=None) for i in range(25): with pytest.raises(ProxyException, match="Invalid credentials"): @@ -1368,30 +1367,28 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): ) -def test_settings_that_arrive_as_environment_strings_are_honored(monkeypatch): +def test_settings_that_arrive_as_environment_strings_are_honored(): """An `os.environ/VAR` reference in general_settings resolves to a string, not an int. Regression: a digit string fell back to the default with only a log line, so an operator tightening the limits through environment substitution silently kept the stock ceilings. """ - from litellm.proxy import proxy_server as ps from litellm.proxy.auth.login_throttle import LoginThrottle - monkeypatch.setattr( - ps, - "general_settings", - { - "max_failed_login_attempts": "7", - "max_failed_login_attempts_per_source": " 70 ", - "failed_login_window_seconds": "not-a-number", - }, - ) request = MagicMock() request.headers = {} request.client = MagicMock() request.client.host = "1.2.3.4" - throttle = LoginThrottle.from_request(request) + throttle = LoginThrottle.from_request( + request, + general_settings={ + "max_failed_login_attempts": "7", + "max_failed_login_attempts_per_source": " 70 ", + "failed_login_window_seconds": "not-a-number", + }, + redis_cache=None, + ) assert throttle.max_attempts == 7 assert throttle.max_attempts_per_source == 70 @@ -1404,41 +1401,37 @@ def test_the_disable_flag_is_read_once_not_per_login_attempt(monkeypatch): With a hosted secret manager in read mode that is a synchronous network call per guess, so a flood of wrong passwords could exhaust the secret manager even after the source was refused. """ - from litellm.proxy import proxy_server as ps from litellm.proxy.auth import login_throttle reads: Final[list[str]] = [] # mutable-ok: test-only call recorder monkeypatch.setattr(login_throttle, "get_secret_bool", lambda name, default: reads.append(name) or default) login_throttle._rate_limit_disabled.cache_clear() - monkeypatch.setattr(ps, "general_settings", {}) request = MagicMock() request.headers = {} request.client = MagicMock() request.client.host = "1.2.3.4" for _ in range(50): - assert login_throttle.LoginThrottle.from_request(request).enabled is True + assert login_throttle.LoginThrottle.from_request(request, general_settings={}, redis_cache=None).enabled is True login_throttle._rate_limit_disabled.cache_clear() assert reads == ["LITELLM_DISABLE_LOGIN_RATE_LIMIT"] -def test_a_negative_or_boolean_setting_falls_back_to_the_default(monkeypatch): +def test_a_negative_or_boolean_setting_falls_back_to_the_default(): """A limit below one would refuse everyone; a bool is a typo, not a count.""" - from litellm.proxy import proxy_server as ps from litellm.proxy.auth.login_throttle import LoginThrottle - monkeypatch.setattr( - ps, - "general_settings", - {"max_failed_login_attempts": "-7", "max_failed_login_attempts_per_source": True}, - ) request = MagicMock() request.headers = {} request.client = MagicMock() request.client.host = "1.2.3.4" - throttle = LoginThrottle.from_request(request) + throttle = LoginThrottle.from_request( + request, + general_settings={"max_failed_login_attempts": "-7", "max_failed_login_attempts_per_source": True}, + redis_cache=None, + ) assert throttle.max_attempts == 50 assert throttle.max_attempts_per_source == 250 From 05fe17e027325bfaa73086450240e2cc4a41700d Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 13 Sep 2026 09:02:45 +0000 Subject: [PATCH 028/253] test(proxy): drop the section banner comment from the login tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/auth/test_login_utils.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 7f5cd3e808f..dadce1cd6c5 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -657,11 +657,6 @@ class TestEncodeUiSessionJwt: assert _user_id_from_session_cookie(request) == "cornell-user" -# --------------------------------------------------------------------------- -# Failed-login accounting (LIT-5285) -# --------------------------------------------------------------------------- - - def _throttle( max_attempts: int = 3, window_seconds: int = 900, From 685c6542985a6ae8cdb817f13a6a183a4111092b Mon Sep 17 00:00:00 2001 From: yucheng Date: Sun, 13 Sep 2026 09:18:39 +0000 Subject: [PATCH 029/253] refactor(proxy): assert separate login counter stores in the spray regression test instead of a comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 2 -- tests/test_litellm/proxy/auth/test_login_utils.py | 4 +++- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index b6ec37afc0e..50333ffad05 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -55,8 +55,6 @@ def _bounded_store(max_entries: int) -> DualCache: ) -# Separate stores: eviction is earliest-expiring-first, so in one shared store a spray of -# fresh usernames would evict the source counter that is meant to stop that same spray. _FAILED_LOGIN_USERNAME_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_USERNAMES) _FAILED_LOGIN_SOURCE_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_SOURCES) _NO_SETTINGS: Final = MappingProxyType({}) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index dadce1cd6c5..c53476f5148 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1472,7 +1472,8 @@ async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): The default in-memory cache keeps 200 entries and evicts the soonest to expire, and every counter shares one window, so eviction was effectively oldest-first. A few hundred made-up usernames therefore pushed out the attacker's own counter and handed - back a fresh allowance against the real account. + back a fresh allowance against the real account. Username and source counters must also + live in separate stores, or the same spray evicts the source counter meant to stop it. """ from litellm.proxy._types import ProxyException from litellm.proxy.auth.login_throttle import ( @@ -1489,6 +1490,7 @@ async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): assert _MAX_TRACKED_LOGIN_USERNAMES >= 10_000 assert _FAILED_LOGIN_SOURCE_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_SOURCES assert _FAILED_LOGIN_USERNAME_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_USERNAMES + assert _FAILED_LOGIN_SOURCE_CACHE.in_memory_cache is not _FAILED_LOGIN_USERNAME_CACHE.in_memory_cache throttle = LoginThrottle( client_ip="10.9.9.9", From cd8887d72cef4581519914d99be1da153fbe8ee8 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 09:50:11 +0000 Subject: [PATCH 030/253] fix(mistral): accept reasoning_effort on all models and drop client_metadata for Codex compatibility Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/mistral/chat/transformation.py | 14 +++++--- .../test_mistral_chat_transformation.py | 33 +++++++++++++++++-- 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index a76a8a3e98c..50aefcdc918 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -99,11 +99,11 @@ class MistralConfig(OpenAIGPTConfig): "stop", "response_format", "parallel_tool_calls", + "reasoning_effort", ] - # Add reasoning support for magistral models if "magistral" in model.lower(): - supported_params.extend(["thinking", "reasoning_effort"]) + supported_params.append("thinking") return supported_params @@ -171,9 +171,11 @@ class MistralConfig(OpenAIGPTConfig): optional_params["extra_body"] = {"random_seed": value} if param == "response_format": optional_params["response_format"] = value - if param == "reasoning_effort" and "magistral" in model.lower(): - # Flag that we need to add reasoning system prompt - optional_params["_add_reasoning_prompt"] = True + if param == "reasoning_effort": + if "magistral" in model.lower(): + optional_params["_add_reasoning_prompt"] = True + else: + optional_params["reasoning_effort"] = value if param == "thinking" and "magistral" in model.lower(): # Flag that we need to add reasoning system prompt optional_params["_add_reasoning_prompt"] = True @@ -534,6 +536,8 @@ class MistralConfig(OpenAIGPTConfig): if "magistral" in model.lower() and optional_params.get("_add_reasoning_prompt", False): messages = self._add_reasoning_system_prompt_if_needed(messages, optional_params) + optional_params.pop("client_metadata", None) + # Call parent transform_request which handles _transform_messages return super().transform_request( model=model, diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index 15694d9f218..edfaf352e1f 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -51,11 +51,11 @@ class TestMistralReasoningSupport: assert "reasoning_effort" in supported_params assert "thinking" in supported_params - # Test non-magistral model doesn't include reasoning parameters + # Non-magistral models accept reasoning_effort (forwarded verbatim) but not thinking supported_params_normal = mistral_config.get_supported_openai_params( "mistral/mistral-large-latest" ) - assert "reasoning_effort" not in supported_params_normal + assert "reasoning_effort" in supported_params_normal assert "thinking" not in supported_params_normal def test_map_openai_params_reasoning_effort(self): @@ -73,7 +73,7 @@ class TestMistralReasoningSupport: assert result.get("_add_reasoning_prompt") is True - # Test reasoning_effort ignored for non-magistral model + # Test reasoning_effort forwarded verbatim for non-magistral model optional_params_normal = {} result_normal = mistral_config.map_openai_params( non_default_params={"reasoning_effort": "low"}, @@ -83,6 +83,33 @@ class TestMistralReasoningSupport: ) assert "_add_reasoning_prompt" not in result_normal + assert result_normal["reasoning_effort"] == "low" + + def test_reasoning_effort_not_unsupported_for_non_magistral(self): + """Codex sends reasoning_effort to every model; Mistral must not raise UnsupportedParamsError.""" + import litellm + + optional_params = litellm.get_optional_params( + model="mistral-medium-latest", + custom_llm_provider="mistral", + reasoning_effort="medium", + ) + assert optional_params["reasoning_effort"] == "medium" + + def test_client_metadata_stripped_from_request(self): + """client_metadata passed by Codex must not reach Mistral, whose schema rejects unknown fields.""" + mistral_config = MistralConfig() + + request = mistral_config.transform_request( + model="mistral-medium-latest", + messages=[{"role": "user", "content": "hi"}], + optional_params={"client_metadata": {"originator": "codex_cli_rs"}, "temperature": 0.2}, + litellm_params={}, + headers={}, + ) + + assert "client_metadata" not in request + assert request["temperature"] == 0.2 def test_map_openai_params_thinking(self): """Test that thinking parameter is properly mapped for magistral models.""" From a8305129a7ff0b8389411b27935008536d698aec Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 09:56:08 +0000 Subject: [PATCH 031/253] refactor(mistral): keep map_openai_params under the complexity ceiling Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/mistral/chat/transformation.py | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 50aefcdc918..970da0582ae 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -171,12 +171,9 @@ class MistralConfig(OpenAIGPTConfig): optional_params["extra_body"] = {"random_seed": value} if param == "response_format": optional_params["response_format"] = value - if param == "reasoning_effort": - if "magistral" in model.lower(): - optional_params["_add_reasoning_prompt"] = True - else: - optional_params["reasoning_effort"] = value - if param == "thinking" and "magistral" in model.lower(): + if param == "reasoning_effort" and "magistral" not in model.lower(): + optional_params["reasoning_effort"] = value + if param in ("reasoning_effort", "thinking") and "magistral" in model.lower(): # Flag that we need to add reasoning system prompt optional_params["_add_reasoning_prompt"] = True if param == "parallel_tool_calls": From dd209ba97b3730523b0a3b3a0c00e84acbb89a9a Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 22:25:53 +0000 Subject: [PATCH 032/253] fix(bedrock): carry s3_endpoint_url and s3_region_name into file content downloads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/batches/batch_utils.py | 2 ++ litellm/litellm_core_utils/get_litellm_params.py | 2 ++ litellm/types/utils.py | 1 + tests/test_litellm/batches/test_batch_utils.py | 2 ++ .../litellm_core_utils/test_get_litellm_params.py | 12 ++++++++++++ .../files/test_bedrock_files_transformation.py | 13 +++++++++++++ 6 files changed, 32 insertions(+) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 87f8fd3946e..0ec39ebf2b2 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -530,6 +530,8 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: "vertex_credentials", "gcs_bucket_name", "bucket_name", + "s3_endpoint_url", + "s3_region_name", "timeout", "max_retries", "_litellm_internal_model_credentials", diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index edd2e88f95c..a70dce89680 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -43,6 +43,8 @@ OPTIONAL_KWARGS_KEYS: Final = ( "timeout", "gcs_bucket_name", "bucket_name", + "s3_endpoint_url", + "s3_region_name", "vertex_credentials", "vertex_project", "vertex_location", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 1d73542c9bb..ee8a3956be8 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3721,6 +3721,7 @@ bedrock_batch_litellm_params: Final = ( "aws_batch_role_arn", "s3_bucket_name", "s3_region_name", + "s3_endpoint_url", "s3_output_bucket_name", "bedrock_tags", ) diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 768ea332677..8c8e0621b07 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -278,6 +278,8 @@ def test_extract_credentials_all_supported_keys(): "vertex_credentials", "gcs_bucket_name", "bucket_name", + "s3_endpoint_url", + "s3_region_name", "timeout", "max_retries", } diff --git a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py index f026ff57719..a34bc2af59d 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py @@ -55,6 +55,18 @@ class TestGetLitellmParamsKwargsExtraction: assert result["timeout"] == 30 assert result["rpm"] == 100 + def test_s3_endpoint_kwargs_are_extracted_when_provided(self): + result = get_litellm_params( + s3_endpoint_url="https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com", + s3_region_name="us-east-1", + ) + assert result["s3_endpoint_url"] == "https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com" + assert result["s3_region_name"] == "us-east-1" + + result_without_s3_kwargs = get_litellm_params() + assert "s3_endpoint_url" not in result_without_s3_kwargs + assert "s3_region_name" not in result_without_s3_kwargs + def test_subset_of_kwargs_only_includes_provided(self): """Only provided kwargs appear, others remain absent.""" result = get_litellm_params(azure_ad_token="token123") diff --git a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py index c609455f3d8..d5bfe5cdfc7 100644 --- a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py @@ -2271,6 +2271,19 @@ class TestBedrockFileContentTransformation: authorization = litellm_params[S3_SIGNED_REQUEST_HEADERS_PARAM]["Authorization"] assert "/eu-west-1/s3/aws4_request" in authorization + def test_s3_request_target_uses_configured_endpoint_url(self): + from litellm.litellm_core_utils.get_litellm_params import get_litellm_params + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + lp = get_litellm_params( + aws_region_name="us-east-1", + s3_endpoint_url="https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com", + ) + + assert BedrockFilesConfig()._s3_request_target( + optional_params={}, litellm_params=lp + ).endpoint_url == "https://bucket.vpce-abc.s3.us-east-1.vpce.amazonaws.com" + def test_validate_environment_merges_and_pops_signed_get_headers(self): from litellm.llms.bedrock.files.transformation import ( S3_SIGNED_REQUEST_HEADERS_PARAM, From 817c396383e56db8055b9b5601771107ca201bf7 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 22:39:09 +0000 Subject: [PATCH 033/253] fix(bedrock): preserve S3 endpoint in credential snapshots Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/types/router.py | 1 + 1 file changed, 1 insertion(+) diff --git a/litellm/types/router.py b/litellm/types/router.py index 0aefc07ae4b..1ce86479f34 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -299,6 +299,7 @@ class CredentialLiteLLMParams(BaseModel): aws_bedrock_runtime_endpoint: str | None = None aws_bedrock_project_id: str | None = None s3_bucket_name: str | None = None + s3_endpoint_url: str | None = None s3_region_name: str | None = None s3_encryption_key_id: str | None = None aws_batch_role_arn: str | None = None From 3b620c65d25ea15ed5d955ae6b31dbb697fa789f Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 14 Sep 2026 22:43:27 +0000 Subject: [PATCH 034/253] chore: sync schema.d.ts with proxy OpenAPI spec Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 0b0e3e18215..6e471d42e49 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29899,6 +29899,8 @@ export interface components { s3_bucket_name?: string | null; /** S3 Encryption Key Id */ s3_encryption_key_id?: string | null; + /** S3 Endpoint Url */ + s3_endpoint_url?: string | null; /** S3 Output Bucket Name */ s3_output_bucket_name?: string | null; /** S3 Region Name */ @@ -40113,6 +40115,8 @@ export interface components { s3_bucket_name?: string | null; /** S3 Encryption Key Id */ s3_encryption_key_id?: string | null; + /** S3 Endpoint Url */ + s3_endpoint_url?: string | null; /** S3 Output Bucket Name */ s3_output_bucket_name?: string | null; /** S3 Region Name */ From c19a19999f86e3b1615e444714cdcf29be2f4349 Mon Sep 17 00:00:00 2001 From: Chloe Lu Date: Tue, 15 Sep 2026 15:34:24 +0800 Subject: [PATCH 035/253] fix(anthropic): register thinking-binding-controls-2026-08-01 in beta headers config Anthropic's preserved-thinking controls (`thinking.block_binding`, Claude Fable 5.1) are only accepted alongside the beta header `thinking-binding-controls-2026-08-01`. The proxy forwards the body field untouched but `filter_and_transform_beta_headers` drops the header because it has no entry in `anthropic_beta_headers_config.json`, so Bedrock and Vertex reject the request with "thinking.adaptive.block_binding: Extra inputs are not permitted". Map the header for anthropic, bedrock, bedrock_converse, vertex_ai and databricks (same beta name on all of them per Anthropic's docs). azure_ai is left null pending verification on Foundry. --- litellm/anthropic_beta_headers_config.json | 6 ++++++ .../test_anthropic_beta_headers_filtering.py | 16 ++++++++++++++++ 2 files changed, 22 insertions(+) diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 3f6817f6e35..38fb9c9462f 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -27,6 +27,7 @@ "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", "token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19", "web-fetch-2025-09-10": "web-fetch-2025-09-10", "web-search-2025-03-05": "web-search-2025-03-05" @@ -57,6 +58,7 @@ "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": null, "token-efficient-tools-2025-02-19": null, "web-fetch-2025-09-10": "web-fetch-2025-09-10", "web-search-2025-03-05": "web-search-2025-03-05" @@ -87,6 +89,7 @@ "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": null, "web-fetch-2025-09-10": null, @@ -118,6 +121,7 @@ "structured-outputs-2025-11-13": null, "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", "web-fetch-2025-09-10": null, @@ -149,6 +153,7 @@ "structured-outputs-2025-11-13": null, "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", "web-fetch-2025-09-10": null, @@ -181,6 +186,7 @@ "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", "text_editor_20241022": null, "text_editor_20250124": null, + "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", "token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19", "web-fetch-2025-09-10": "web-fetch-2025-09-10", "web-search-2025-03-05": "web-search-2025-03-05" diff --git a/tests/test_litellm/test_anthropic_beta_headers_filtering.py b/tests/test_litellm/test_anthropic_beta_headers_filtering.py index 59bab22de74..3c967283abf 100644 --- a/tests/test_litellm/test_anthropic_beta_headers_filtering.py +++ b/tests/test_litellm/test_anthropic_beta_headers_filtering.py @@ -426,6 +426,22 @@ class TestAnthropicBetaHeadersFiltering: assert filtered == ["fine-grained-tool-streaming-2025-05-14"] + @pytest.mark.parametrize( + "provider", ["anthropic", "bedrock", "bedrock_converse", "vertex_ai", "databricks"] + ) + def test_thinking_binding_controls_forwarded(self, provider): + """`thinking.block_binding` (preserved thinking, Claude Fable 5.1) is only + accepted alongside thinking-binding-controls-2026-08-01. The body field is + forwarded untouched, so stripping the header (previously unknown, hence + dropped) makes Bedrock and Vertex reject the request with + "thinking.adaptive.block_binding: Extra inputs are not permitted".""" + filtered = filter_and_transform_beta_headers( + beta_headers=["thinking-binding-controls-2026-08-01"], + provider=provider, + ) + + assert filtered == ["thinking-binding-controls-2026-08-01"] + def test_null_value_headers_filtered(self): """Test that headers with null values are always filtered out.""" for provider in [ From 163c0f3aee99ba61f12317453cd76741b9f1558b Mon Sep 17 00:00:00 2001 From: clonylu Date: Tue, 15 Sep 2026 15:49:06 +0800 Subject: [PATCH 036/253] fix(router): honor stream_timeout on the SDK-native passthrough route Anthropic /v1/messages and Bedrock /converse resolve their upstream timeout through resolve_llm_passthrough_timeout, which only reads timeout / request_timeout and then falls back to the 600s pass_through default. A stream_timeout set on the deployment or in router_settings was never consulted on that route, while /chat/completions honors it through Router._get_stream_timeout. For a streaming call the resolver now checks stream_timeout at each level before the non-stream key (kwargs -> litellm_params -> router), mirroring _get_stream_timeout; non-streaming resolution is unchanged. The router passes its stream_timeout alongside the explicit timeout. --- litellm/passthrough/timeout_utils.py | 23 ++++++-- litellm/router.py | 4 ++ .../test_pass_through_endpoints.py | 58 +++++++++++++++++++ tests/test_litellm/test_router.py | 58 +++++++++++++++++++ 4 files changed, 139 insertions(+), 4 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index 39127d19183..fb649a9eeaf 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -34,22 +34,37 @@ def resolve_llm_passthrough_timeout( kwargs: dict | None = None, litellm_params: dict | None = None, router_timeout: float | None = None, + router_stream_timeout: float | None = None, ) -> float: """ - Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse). + Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse, + Anthropic /v1/messages). - Precedence: kwargs timeout/request_timeout -> litellm_params timeout/request_timeout - -> router_timeout -> general_settings.pass_through_request_timeout -> 600s default. + Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params + timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout + -> 600s default. + + Streaming (``kwargs["stream"]`` truthy) additionally consults ``stream_timeout`` at each + level before the non-streaming key, matching ``Router._get_stream_timeout`` on the + completion route: kwargs stream_timeout -> kwargs timeout/request_timeout -> + litellm_params stream_timeout -> litellm_params timeout/request_timeout -> + router_stream_timeout -> router_timeout -> pass_through_request_timeout -> 600s. """ kwargs = kwargs or {} litellm_params = litellm_params or {} + is_stream: Final[bool] = bool(kwargs.get("stream", False)) + keys: Final[tuple[str, ...]] = ( + ("stream_timeout", "timeout", "request_timeout") if is_stream else ("timeout", "request_timeout") + ) for source in (kwargs, litellm_params): - for key in ("timeout", "request_timeout"): + for key in keys: val = source.get(key) if val is not None: return float(val) + if is_stream and router_stream_timeout is not None: + return float(router_stream_timeout) if router_timeout is not None: return float(router_timeout) diff --git a/litellm/router.py b/litellm/router.py index 1665583386f..7ce7ba30502 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3879,10 +3879,14 @@ class Router: _router_timeout: Final = ( float(self._explicit_timeout) if isinstance(self._explicit_timeout, (int, float)) else None ) + _router_stream_timeout: Final = ( + float(self.stream_timeout) if isinstance(self.stream_timeout, (int, float)) else None + ) kwargs["timeout"] = resolve_llm_passthrough_timeout( kwargs=kwargs, litellm_params=deployment["litellm_params"], router_timeout=_router_timeout, + router_stream_timeout=_router_stream_timeout, ) else: kwargs["timeout"] = self._get_timeout(kwargs=kwargs, data=deployment["litellm_params"]) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index d57bed430c1..d697a114613 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1119,6 +1119,64 @@ def test_resolve_llm_passthrough_timeout_precedence(): assert resolve_llm_passthrough_timeout() == 6.0 +def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): + # streaming: stream_timeout wins at each level, then falls through to the non-stream keys + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True, "stream_timeout": 1800, "timeout": 45}, + ) + == 1800.0 + ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True}, + litellm_params={"stream_timeout": 1800, "timeout": 90}, + ) + == 1800.0 + ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True}, + litellm_params={"timeout": 90}, + router_stream_timeout=1800, + ) + == 90.0 + ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True}, + router_timeout=120, + router_stream_timeout=1800, + ) + == 1800.0 + ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True}, + router_timeout=120, + ) + == 120.0 + ) + + # non-streaming: stream_timeout is ignored everywhere + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": False, "stream_timeout": 1800}, + litellm_params={"stream_timeout": 1800, "timeout": 90}, + router_stream_timeout=1800, + ) + == 90.0 + ) + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert ( + resolve_llm_passthrough_timeout( + litellm_params={"stream_timeout": 1800}, + router_stream_timeout=1800, + ) + == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + ) + + @pytest.mark.asyncio async def test_pass_through_request_uses_resolved_timeout(): with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 3cabcd71627..70ce182c028 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -5480,6 +5480,64 @@ def test_update_kwargs_with_deployment_uses_pass_through_request_timeout(): assert kwargs["timeout"] == 6.0 +def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): + """ + The SDK-native passthrough route (anthropic /v1/messages, bedrock /converse) resolves + its upstream timeout separately from the completion route. A streaming call must get + stream_timeout (deployment litellm_params first, then router_settings), while a + non-streaming call on the same deployment keeps the non-stream resolution. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "anthropic-with-stream-timeout", + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "api_key": "fake-key", + "stream_timeout": 1800, + }, + }, + { + "model_name": "anthropic-router-default", + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "api_key": "fake-key", + }, + }, + ], + stream_timeout=900, + ) + per_deployment, router_default = router.model_list + + with patch( + "litellm.proxy.proxy_server.general_settings", + {"pass_through_request_timeout": 6}, + ): + kwargs: dict = {"stream": True} + router._update_kwargs_with_deployment( + deployment=per_deployment, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 1800.0 + + kwargs = {"stream": True} + router._update_kwargs_with_deployment( + deployment=router_default, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 900.0 + + kwargs = {"stream": False} + router._update_kwargs_with_deployment( + deployment=per_deployment, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 6.0 + + @pytest.mark.asyncio async def test_router_acompletion_with_unknown_model_and_default_fallback(): """ From efb2bcd87fae4c6a78bb562cbcde98778b967a79 Mon Sep 17 00:00:00 2001 From: clonylu Date: Tue, 15 Sep 2026 16:11:55 +0800 Subject: [PATCH 037/253] test(router): cover passthrough stream_timeout without patching proxy globals --- .../test_pass_through_endpoints.py | 14 ++--- tests/test_litellm/test_router.py | 56 ++++++++++--------- 2 files changed, 38 insertions(+), 32 deletions(-) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index d697a114613..7f8663ea860 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1167,14 +1167,14 @@ def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): ) == 90.0 ) - with patch("litellm.proxy.proxy_server.general_settings", {}): - assert ( - resolve_llm_passthrough_timeout( - litellm_params={"stream_timeout": 1800}, - router_stream_timeout=1800, - ) - == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + assert ( + resolve_llm_passthrough_timeout( + litellm_params={"stream_timeout": 1800}, + router_timeout=120, + router_stream_timeout=1800, ) + == 120.0 + ) @pytest.mark.asyncio diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 70ce182c028..4c39f8ba4e4 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -5494,6 +5494,7 @@ def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): "litellm_params": { "model": "anthropic/claude-sonnet-4-5", "api_key": "fake-key", + "timeout": 60, "stream_timeout": 1800, }, }, @@ -5505,37 +5506,42 @@ def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): }, }, ], + timeout=120, stream_timeout=900, ) per_deployment, router_default = router.model_list - with patch( - "litellm.proxy.proxy_server.general_settings", - {"pass_through_request_timeout": 6}, - ): - kwargs: dict = {"stream": True} - router._update_kwargs_with_deployment( - deployment=per_deployment, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 1800.0 + kwargs: dict = {"stream": True} + router._update_kwargs_with_deployment( + deployment=per_deployment, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 1800.0 - kwargs = {"stream": True} - router._update_kwargs_with_deployment( - deployment=router_default, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 900.0 + kwargs = {"stream": True} + router._update_kwargs_with_deployment( + deployment=router_default, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 900.0 - kwargs = {"stream": False} - router._update_kwargs_with_deployment( - deployment=per_deployment, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 6.0 + kwargs = {"stream": False} + router._update_kwargs_with_deployment( + deployment=per_deployment, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 60.0 + + kwargs = {"stream": False} + router._update_kwargs_with_deployment( + deployment=router_default, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + assert kwargs["timeout"] == 120.0 @pytest.mark.asyncio From 911f66aff69fabb6666bde3f54db70960cb04b56 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 15 Sep 2026 12:12:57 -0500 Subject: [PATCH 038/253] feat(azure_ai): support FLUX.2 flex images --- litellm/images/main.py | 9 +- litellm/images/utils.py | 7 +- .../litellm_core_utils/llm_cost_calc/utils.py | 3 + .../image_edit/flux2_transformation.py | 64 ++++--- .../image_generation/cost_calculator.py | 34 +++- .../image_generation/flux_transformation.py | 91 +++++++-- ...odel_prices_and_context_window_backup.json | 19 ++ litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/types/llms/openai.py | 4 + model_prices_and_context_window.json | 19 ++ ...test_azure_ai_image_edit_transformation.py | 116 ++++++++++++ .../test_azure_ai_flux2_image_generation.py | 172 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 +- 13 files changed, 491 insertions(+), 53 deletions(-) create mode 100644 tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py diff --git a/litellm/images/main.py b/litellm/images/main.py index 6a94e7c8df2..81547a153c3 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -846,7 +846,12 @@ def image_edit( local_vars.update(kwargs) # Get ImageEditOptionalRequestParams with only valid parameters image_edit_optional_params: Final[ImageEditOptionalRequestParams] = ( - _get_ImageEditRequestUtils().get_requested_image_edit_optional_param(local_vars) + _get_ImageEditRequestUtils().get_requested_image_edit_optional_param( + local_vars, + provider_supported_params=frozenset( + image_edit_provider_config.get_supported_openai_params(model) + ).intersection(non_default_params), + ) ) # Get optional parameters for the responses API image_edit_request_params: Final[dict] = _get_ImageEditRequestUtils().get_optional_params_image_edit( @@ -857,7 +862,7 @@ def image_edit( additional_drop_params=kwargs.get("additional_drop_params"), ) - if ( + if image_edit_provider_config.use_multipart_form_data() and ( custom_llm_provider == "openai" or custom_llm_provider == "azure" or custom_llm_provider in litellm.openai_compatible_providers diff --git a/litellm/images/utils.py b/litellm/images/utils.py index 49b70870de6..24454954714 100644 --- a/litellm/images/utils.py +++ b/litellm/images/utils.py @@ -1,4 +1,4 @@ -from collections.abc import Mapping +from collections.abc import Collection, Mapping from io import BufferedReader, BytesIO from typing import Any, Final, cast, get_type_hints @@ -63,6 +63,7 @@ class ImageEditRequestUtils: @staticmethod def get_requested_image_edit_optional_param( params: Mapping[str, object], + provider_supported_params: Collection[str] = (), ) -> ImageEditOptionalRequestParams: """ Filter parameters to only include those defined in ImageEditOptionalRequestParams. @@ -73,7 +74,9 @@ class ImageEditRequestUtils: Returns: ImageEditOptionalRequestParams instance with only the valid parameters """ - valid_keys: Final = get_type_hints(ImageEditOptionalRequestParams).keys() + valid_keys: Final = frozenset(get_type_hints(ImageEditOptionalRequestParams)) | frozenset( + provider_supported_params + ) filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None} return cast(ImageEditOptionalRequestParams, filtered_params) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index baa9aab1087..0a0e92ff3a3 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -1853,6 +1853,9 @@ class CostCalculatorUtils: return azure_ai_image_cost_calculator( model=model, image_response=completion_response, + size=resolved_size, + n=resolved_n, + optional_params=optional_params, ) elif custom_llm_provider == litellm.LlmProviders.FAL_AI.value: from litellm.llms.fal_ai.cost_calculator import ( diff --git a/litellm/llms/azure_ai/image_edit/flux2_transformation.py b/litellm/llms/azure_ai/image_edit/flux2_transformation.py index a09a80985b7..aa8905e5601 100644 --- a/litellm/llms/azure_ai/image_edit/flux2_transformation.py +++ b/litellm/llms/azure_ai/image_edit/flux2_transformation.py @@ -1,5 +1,7 @@ import base64 +from collections.abc import Mapping, Sequence from io import BufferedReader +from types import MappingProxyType from typing import Any, Final from httpx._types import RequestFiles @@ -24,7 +26,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): Azure AI Foundry FLUX 2 image edit config Supports FLUX 2 models (e.g., flux.2-pro) for image editing. - Uses the same /providers/blackforestlabs/v1/flux-2-pro endpoint as image generation, + Uses the model-specific /providers/blackforestlabs/v1/flux-2-* endpoint as image generation, with the image passed as base64 in JSON body. """ @@ -33,11 +35,17 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): FLUX 2 supports a subset of OpenAI image edit params """ return [ - "prompt", - "image", - "model", "n", "size", + "width", + "height", + "num_images", + "seed", + "safety_tolerance", + "output_format", + "aspect_ratio", + "guidance", + "steps", ] def map_openai_params( @@ -50,14 +58,14 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): Map OpenAI params to FLUX 2 params. FLUX 2 uses the same param names as OpenAI for supported params. """ - mapped_params: Final[dict[str, Any]] = {} - supported_params: Final = self.get_supported_openai_params(model) - - for key, value in dict(image_edit_optional_params).items(): - if key in supported_params and value is not None: - mapped_params[key] = value - - return mapped_params + return AzureFoundryFluxImageGenerationConfig().map_openai_params( + non_default_params=MappingProxyType( + {key: value for key, value in image_edit_optional_params.items() if value is not None} + ), + optional_params=MappingProxyType({}), + model=model, + drop_params=drop_params, + ) def use_multipart_form_data(self) -> bool: """FLUX 2 uses JSON requests, not multipart/form-data.""" @@ -90,7 +98,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): self, model: str, prompt: str | None, - image: FileTypes | None, + image: FileTypes | Sequence[FileTypes] | None, image_edit_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -107,29 +115,29 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): if image is None: raise ValueError("FLUX 2 image edit requires an image.") - image_b64: Final = self._convert_image_to_base64(image) + images: Final = tuple(image) if isinstance(image, list) else (image,) + if not images: + raise ValueError("FLUX 2 image edit requires at least one image.") + max_reference_images: Final = 10 if "flex" in model.lower() else 8 + if len(images) > max_reference_images: + raise ValueError(f"{model} supports at most {max_reference_images} reference images.") - # Build request body with required params + reference_images: Final[Mapping[str, str]] = MappingProxyType( + { + "input_image" if index == 1 else f"input_image_{index}": self._convert_image_to_base64(reference_image) + for index, reference_image in enumerate(images, start=1) + } + ) request_body: Final[dict[str, Any]] = { "prompt": prompt, - "image": image_b64, "model": model, + **reference_images, + **image_edit_optional_request_params, } - - # Add mapped optional params (already filtered by map_openai_params) - request_body.update(image_edit_optional_request_params) - - # Return JSON body and empty files list (FLUX 2 doesn't use multipart) return request_body, [] def _convert_image_to_base64(self, image: Any) -> str: """Convert image file to base64 string""" - # Handle list of images (take first one) - if isinstance(image, list): - if len(image) == 0: - raise ValueError("Empty image list provided") - image = image[0] - if isinstance(image, BufferedReader): image_bytes = image.read() image.seek(0) # Reset file pointer for potential reuse @@ -151,7 +159,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): """ Constructs a complete URL for Azure AI Foundry FLUX 2 image edits. - Uses the same /providers/blackforestlabs/v1/flux-2-pro endpoint as image generation. + Uses the same model-specific BFL provider endpoint as image generation. """ api_base = AzureFoundryModelInfo.get_api_base(api_base) diff --git a/litellm/llms/azure_ai/image_generation/cost_calculator.py b/litellm/llms/azure_ai/image_generation/cost_calculator.py index 106c7e42b83..086293f26c0 100644 --- a/litellm/llms/azure_ai/image_generation/cost_calculator.py +++ b/litellm/llms/azure_ai/image_generation/cost_calculator.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Any, Final import litellm @@ -10,6 +11,9 @@ from litellm.types.utils import ImageResponse def cost_calculator( model: str, image_response: Any, + size: str | None = None, + n: int | None = None, + optional_params: Mapping[str, object] | None = None, ) -> float: """ Azure AI image generation cost calculator @@ -28,10 +32,32 @@ def cost_calculator( if token_based_cost is not None: return token_based_cost + num_images: Final = n if n is not None else len(image_response.data or ()) output_cost_per_image: Final[float] = _model_info.get("output_cost_per_image") or 0.0 - num_images: int = 0 - if image_response.data: - num_images = len(image_response.data) - return output_cost_per_image * num_images + if output_cost_per_image: + return output_cost_per_image * num_images + + model_cost: Final = litellm.model_cost[_model_info["key"]] + input_cost_per_pixel: Final[float] = model_cost.get("input_cost_per_pixel") or 0.0 + if input_cost_per_pixel: + from litellm.cost_calculator import default_image_cost_calculator + + cost_model: Final = ( + model if model.startswith(f"{litellm.LlmProviders.AZURE_AI.value}/") else f"azure_ai/{model}" + ) + width: Final = optional_params.get("width") if optional_params else None + height: Final = optional_params.get("height") if optional_params else None + pixel_size: Final = ( + f"{width}x{height}" + if type(width) is int and type(height) is int and width > 0 and height > 0 + else size or image_response.size + ) + return default_image_cost_calculator( + model=cost_model, + custom_llm_provider=litellm.LlmProviders.AZURE_AI.value, + size=pixel_size, + n=num_images, + ) + return 0.0 raise ValueError(f"image_response must be of type ImageResponse got type={type(image_response)}") diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index 65b5a35af52..b10e8a6f35e 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -1,18 +1,13 @@ +from collections.abc import Mapping +from types import MappingProxyType from typing import Final from litellm.llms.openai.image_generation import GPTImageGenerationConfig +from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): - """ - Azure Foundry flux image generation config - - From manual testing it follows the gpt-image-1 image generation config - - (Azure Foundry does not have any docs on supported params at the time of writing) - - From our test suite - following GPTImageGenerationConfig is working for this model - """ + """Azure Foundry BFL API configuration for FLUX image generation.""" @staticmethod def get_flux2_image_generation_url( @@ -25,11 +20,11 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): FLUX 2 models on Azure AI use a different URL pattern than standard Azure OpenAI: - Standard: /openai/deployments/{model}/images/generations - - FLUX 2: /providers/blackforestlabs/v1/flux-2-pro + - FLUX 2: /providers/blackforestlabs/v1/{model-path} Args: api_base: Base URL (e.g., https://litellm-ci-cd-prod.services.ai.azure.com) - model: Model name (e.g., flux.2-pro) + model: Model name (e.g., FLUX.2-flex or FLUX.2-pro) api_version: API version (e.g., preview) Returns: @@ -47,9 +42,8 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): return api_base return f"{api_base}?api-version={api_version}" - # Construct the FLUX 2 provider path - # Model name flux.2-pro maps to endpoint flux-2-pro - return f"{api_base}/providers/blackforestlabs/v1/flux-2-pro?api-version={api_version}" + provider_model_path: Final = AzureFoundryFluxImageGenerationConfig.get_flux2_provider_model_path(model) + return f"{api_base}/providers/blackforestlabs/v1/{provider_model_path}?api-version={api_version}" @staticmethod def is_flux2_model(model: str) -> bool: @@ -64,3 +58,72 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): """ model_lower: Final = model.lower().replace(".", "-").replace("_", "-") return "flux-2" in model_lower or "flux2" in model_lower + + @staticmethod + def get_flux2_provider_model_path(model: str) -> str: + normalized_model: Final = model.lower().replace(".", "-").replace("_", "-") + return "flux-2-flex" if "flux-2-flex" in normalized_model else "flux-2-pro" + + def get_supported_openai_params( # mutable-ok: inherited config contract returns a list + self, model: str + ) -> list[OpenAIImageGenerationOptionalParams]: + if not self.is_flux2_model(model): + return super().get_supported_openai_params(model) + return [ # mutable-ok: BaseImageGenerationConfig requires a list + "n", + "size", + "output_format", + "seed", + "safety_tolerance", + "aspect_ratio", + "width", + "height", + "num_images", + "guidance", + "steps", + ] + + @staticmethod + def _map_parameter(name: str, value: object) -> tuple[tuple[str, object], ...]: + if name == "n": + return (("num_images", value),) + if name != "size": + return ((name, value),) + + try: + width, height = (int(dimension) for dimension in str(value).lower().split("x")) + except (TypeError, ValueError): + raise ValueError(f"Invalid size format '{value}'. Expected 'WxH', for example '1024x1024'.") + return (("width", width), ("height", height)) + + def map_openai_params( + self, + non_default_params: Mapping[str, object], + optional_params: Mapping[str, object], + model: str, + drop_params: bool, + ) -> dict: # mutable-ok: inherited config contract returns a dict + if not self.is_flux2_model(model): + return super().map_openai_params( + non_default_params=dict(non_default_params), + optional_params=dict(optional_params), + model=model, + drop_params=drop_params, + ) + supported_params: Final = self.get_supported_openai_params(model) + unsupported_params: Final = tuple(name for name in non_default_params if name not in supported_params) + if unsupported_params and not drop_params: + raise ValueError( + f"Parameters {unsupported_params} are not supported for model {model}. " + f"Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters." + ) + + mapped_params: Final[Mapping[str, object]] = MappingProxyType( + { + mapped_name: mapped_value + for name, value in non_default_params.items() + if name in supported_params + for mapped_name, mapped_value in self._map_parameter(name, value) + } + ) + return {**optional_params, **mapped_params} # mutable-ok: inherited config contract returns a dict diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9f91cf82f41..bcb2dfc53e8 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -9900,6 +9900,25 @@ "/v1/images/generations" ] }, + "azure_ai/FLUX.2-flex": { + "input_cost_per_pixel": 5e-08, + "litellm_provider": "azure_ai", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "image_generation", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/black-forest-labs/", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "azure_ai/FW-DeepSeek-V3.2": { "deprecation_date": "2027-07-01", "cache_read_input_token_cost": 3.1e-07, diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fb2f014e3d8..81889e4da0c 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -18974,7 +18974,7 @@ } } }, - "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" }, "500": { "content": { diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index e3eac9b9205..710b34116e5 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1160,6 +1160,10 @@ OpenAIImageGenerationOptionalParams = Literal[ "image_url", "image_prompt_strength", "aspect_ratio", + "width", + "height", + "guidance", + "steps", "imageConfig", ] diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9f91cf82f41..bcb2dfc53e8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -9900,6 +9900,25 @@ "/v1/images/generations" ] }, + "azure_ai/FLUX.2-flex": { + "input_cost_per_pixel": 5e-08, + "litellm_provider": "azure_ai", + "max_input_tokens": 32000, + "max_tokens": 32000, + "mode": "image_generation", + "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/black-forest-labs/", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "azure_ai/FW-DeepSeek-V3.2": { "deprecation_date": "2027-07-01", "cache_read_input_token_cost": 3.1e-07, diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py index b6cb7ea9b54..94b23c5f1c2 100644 --- a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py @@ -1,12 +1,20 @@ +import base64 +import json +from collections.abc import Mapping +from typing import Final +import httpx +import pytest import litellm +from litellm.images.utils import ImageEditRequestUtils from litellm.llms.azure_ai.image_edit.flux2_transformation import ( AzureFoundryFlux2ImageEditConfig, ) from litellm.llms.azure_ai.image_edit.transformation import ( AzureFoundryFluxImageEditConfig, ) +from litellm.llms.custom_httpx.http_handler import HTTPHandler def test_azure_ai_validate_environment(): @@ -60,3 +68,111 @@ def test_flux2_validate_environment_with_entra_token(monkeypatch): assert headers["Authorization"] == "Bearer entra-token" assert headers["Content-Type"] == "application/json" + + +def test_flux2_image_edit_maps_openai_and_provider_parameters(): + config = AzureFoundryFlux2ImageEditConfig() + requested_params = ImageEditRequestUtils.get_requested_image_edit_optional_param( + { + "n": 2, + "size": "1536x1024", + "guidance": 4.5, + "steps": 32, + "unrelated": "discarded", + }, + provider_supported_params=config.get_supported_openai_params("FLUX.2-flex"), + ) + mapped_params = config.map_openai_params( + image_edit_optional_params=requested_params, + model="FLUX.2-flex", + drop_params=False, + ) + + assert mapped_params == { + "num_images": 2, + "width": 1536, + "height": 1024, + "guidance": 4.5, + "steps": 32, + } + + +@pytest.mark.parametrize( + ("model", "max_reference_images"), + [ + ("FLUX.2-flex", 10), + ("FLUX.2-pro", 8), + ], +) +def test_flux2_image_edit_uses_all_reference_fields(model: str, max_reference_images: int): + images = [f"image-{index}".encode() for index in range(1, max_reference_images + 1)] + request, files = AzureFoundryFlux2ImageEditConfig().transform_image_edit_request( + model=model, + prompt="Blend every reference", + image=images, + image_edit_optional_request_params={"guidance": 4.5, "steps": 20}, + litellm_params={}, + headers={}, + ) + + assert files == [] + assert request["input_image"] == base64.b64encode(images[0]).decode() + assert request[f"input_image_{max_reference_images}"] == base64.b64encode(images[-1]).decode() + assert "input_image_1" not in request + assert "image" not in request + assert len([key for key in request if key.startswith("input_image")]) == max_reference_images + assert request["guidance"] == 4.5 + assert request["steps"] == 20 + + +@pytest.mark.parametrize( + ("model", "reference_images"), + [ + ("FLUX.2-flex", 11), + ("FLUX.2-pro", 9), + ], +) +def test_flux2_image_edit_rejects_too_many_references(model: str, reference_images: int): + with pytest.raises(ValueError, match=f"at most {reference_images - 1} reference images"): + AzureFoundryFlux2ImageEditConfig().transform_image_edit_request( + model=model, + prompt="Blend every reference", + image=[b"image"] * reference_images, + image_edit_optional_request_params={}, + litellm_params={}, + headers={}, + ) + + +@pytest.mark.parametrize("dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024})) +@pytest.mark.usefixtures("local_model_cost_map") +def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[str, int | str]): + def respond(request: httpx.Request) -> httpx.Response: + body: Final = json.loads(request.content) + assert body == { + "model": "FLUX.2-flex", + "prompt": "Add a hat", + "input_image": base64.b64encode(b"image").decode(), + "num_images": 2, + "width": 2048, + "height": 1024, + "guidance": 4.5, + "steps": 32, + } + return httpx.Response(200, json={"data": [{"b64_json": "aW1n"}, {"b64_json": "aW1n"}]}) + + client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))) + response: Final = litellm.image_edit( + model="azure_ai/FLUX.2-flex", + image=b"image", + prompt="Add a hat", + api_key="test-key", + api_base="https://example.services.ai.azure.com", + client=client, + n=2, + guidance=4.5, + steps=32, + **dimensions, + ) + + assert response._hidden_params["response_cost"] == pytest.approx(5e-08 * 2048 * 1024 * 2) diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py new file mode 100644 index 00000000000..07026cf4309 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py @@ -0,0 +1,172 @@ +from collections.abc import Mapping +from typing import Final +from unittest.mock import MagicMock + +import httpx +import pytest + +import litellm +from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap +from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils +from litellm.llms.azure.azure import AzureChatCompletion +from litellm.llms.azure.image_generation import get_azure_image_generation_config +from litellm.llms.azure.image_generation.http_utils import azure_deployment_image_generation_json_body +from litellm.llms.azure_ai.image_generation.flux_transformation import ( + AzureFoundryFluxImageGenerationConfig, +) +from litellm.types.utils import ImageObject, ImageResponse +from litellm.utils import _invalidate_model_cost_lowercase_map + + +@pytest.fixture(autouse=True) +def use_local_model_cost_map(monkeypatch): + monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map()) + litellm.get_model_info.cache_clear() + _invalidate_model_cost_lowercase_map() + yield + litellm.get_model_info.cache_clear() + _invalidate_model_cost_lowercase_map() + + +@pytest.mark.parametrize( + ("model", "provider_path"), + [ + ("FLUX.2-flex", "flux-2-flex"), + ("FLUX.2-pro", "flux-2-pro"), + ], +) +def test_flux2_uses_model_specific_provider_url(model: str, provider_path: str): + url = AzureChatCompletion().create_azure_base_url( + azure_client_params={ + "azure_endpoint": "https://example.services.ai.azure.com/", + "api_version": "preview", + }, + model=model, + ) + + assert ( + url == f"https://example.services.ai.azure.com/providers/blackforestlabs/v1/{provider_path}?api-version=preview" + ) + + +def test_flux2_flex_maps_openai_and_provider_parameters(): + config = AzureFoundryFluxImageGenerationConfig() + mapped_params = config.map_openai_params( + non_default_params={ + "n": 2, + "size": "1536x1024", + "guidance": 4.5, + "steps": 32, + "output_format": "jpeg", + }, + optional_params={}, + model="FLUX.2-flex", + drop_params=False, + ) + url = config.get_flux2_image_generation_url( + api_base="https://example.services.ai.azure.com", + model="FLUX.2-flex", + api_version="preview", + ) + request = azure_deployment_image_generation_json_body( + api_base=url, + data={"model": "FLUX.2-flex", "prompt": "A red fox", **mapped_params}, + deployment_name="FLUX.2-flex", + ) + + assert request == { + "model": "FLUX.2-flex", + "prompt": "A red fox", + "num_images": 2, + "width": 1536, + "height": 1024, + "guidance": 4.5, + "steps": 32, + "output_format": "jpeg", + } + + +def test_flux2_flex_rejects_invalid_size(): + with pytest.raises(ValueError, match="Expected 'WxH'"): + AzureFoundryFluxImageGenerationConfig().map_openai_params( + non_default_params={"size": "large"}, + optional_params={}, + model="FLUX.2-flex", + drop_params=False, + ) + + +def test_flux2_flex_model_info(): + model_info = litellm.get_model_info( + model="FLUX.2-flex", + custom_llm_provider="azure_ai", + ) + catalog_info = litellm.model_cost["azure_ai/FLUX.2-flex"] + + assert model_info["mode"] == "image_generation" + assert model_info["max_input_tokens"] == 32000 + assert model_info["max_tokens"] == 32000 + assert model_info["supported_endpoints"] == ["/v1/images/generations", "/v1/images/edits"] + assert catalog_info["input_cost_per_pixel"] == 5e-08 + assert catalog_info["supported_modalities"] == ["text", "image"] + assert catalog_info["supported_output_modalities"] == ["image"] + + +def test_flux2_flex_cost_uses_generated_megapixels(): + response = ImageResponse( + data=[ + ImageObject(url="https://example.com/one.png"), + ImageObject(url="https://example.com/two.png"), + ] + ) + + cost = CostCalculatorUtils.route_image_generation_cost_calculator( + model="FLUX.2-flex", + completion_response=response, + custom_llm_provider="azure_ai", + size="2048x1024", + call_type="image_generation", + ) + + assert cost == pytest.approx(5e-08 * 2048 * 1024 * 2) + + +@pytest.mark.parametrize("model", ("FLUX-1.1-pro", "FLUX.1-Kontext-pro")) +def test_flux1_preserves_existing_openai_parameters(model: str): + params: Final = {"n": 2, "size": "1536x1024", "quality": "high", "user": "test-user"} + + mapped: Final = AzureFoundryFluxImageGenerationConfig().map_openai_params( + non_default_params=params, + optional_params={}, + model=model, + drop_params=False, + ) + + assert mapped == params + + +@pytest.mark.parametrize("dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024})) +def test_flux2_cost_uses_mapped_dimensions_after_response_transformation(dimensions: Mapping[str, int | str]): + params: Final = AzureFoundryFluxImageGenerationConfig().map_openai_params( + non_default_params={"n": 2, **dimensions}, + optional_params={}, + model="FLUX.2-flex", + drop_params=False, + ) + response: Final = get_azure_image_generation_config("FLUX.2-flex").transform_image_generation_response( + model="FLUX.2-flex", + raw_response=httpx.Response(200, json={"data": [{"b64_json": "aW1n"}, {"b64_json": "aW1n"}]}), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={"prompt": "A red fox", **params}, + optional_params=params, + litellm_params={}, + encoding=None, + ) + + assert litellm.completion_cost( + model="azure_ai/FLUX.2-flex", + completion_response=response, + optional_params=params, + call_type="image_generation", + ) == pytest.approx(5e-08 * 2048 * 1024 * 2) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b26f5e25b6f..6b7337ea8d1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -35343,7 +35343,7 @@ export interface components { default_model?: string | null; /** * Deployment Affinity - * @description When True and a session_id is resolvable on the request, pin the deployment chosen inside each routed model group and reuse it whenever the session returns to that group, without pinning which group the session routes to. Independent of session_affinity, which pins the model group instead (and always carries this deployment pin with it): with session_affinity off, every turn is still classified on its own merits while a session that escalates to a stronger tier and comes back still lands on the deployment it used before, which is what keeps a provider prompt cache warm. Pins are held per model group, so switching tiers does not disturb the pin left behind in the previous group. On by default because re-shuffling a conversation across deployments of the same model discards that cache for no benefit; set False to keep every turn load-balanced across the group, which is what a deployment set with tight per-deployment rate limits wants. Inert when no session_id is resolvable, since there is nothing to key a pin on, and suppressed when plugins are configured, for the same reason session_affinity is. + * @description When True and a client session_id is resolvable, reuse the session's chosen model for each classified tier and its deployment within each model group. With session_affinity off, every turn is still classified: moving to another tier leaves the previous tier's model pin intact for a later return. Pins yield to current candidate, context, modality, and availability constraints. Adaptive selection chooses the initial model from its eligible pool, then reuses that choice per tier. This reduces avoidable provider prompt-cache misses; it does not guarantee cache hits. Set False to select models and load-balance deployments on every turn, unless session_affinity or user_turn classification requires a pin. Inert without a client session_id and suppressed when plugins are configured. * @default true */ deployment_affinity: boolean; @@ -35487,7 +35487,7 @@ export interface components { session_affinity: boolean; /** * Session Affinity Ttl Seconds - * @description TTL for the session affinity pin; refreshed on every cache hit. Bounds both the session_affinity model pin and the deployment_affinity deployment pin, so it measures idle time for the session's routing decisions rather than total session length + * @description TTL for the session affinity pin; refreshed on every cache hit. Bounds both the session_affinity model pin and the deployment_affinity per-tier model and deployment pins, so it measures idle time for the session's routing decisions rather than total session length * @default 3600 */ session_affinity_ttl_seconds: number; From ff1e2a02ba9b7bafee87896edcbf935749bf59c3 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 15 Sep 2026 12:33:58 -0500 Subject: [PATCH 039/253] fix(proxy): preserve CI-compatible OpenAPI snapshot formatting --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 81889e4da0c..fb2f014e3d8 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -18974,7 +18974,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " }, "500": { "content": { From 3821e5ace617f43122e757ec0d922c6f4a8c9b6a Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 15 Sep 2026 12:43:24 -0500 Subject: [PATCH 040/253] fix(azure_ai): specify FLUX parameter mapping return type --- litellm/llms/azure_ai/image_generation/flux_transformation.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index b10e8a6f35e..b9d5e11ff2a 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -102,7 +102,7 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): optional_params: Mapping[str, object], model: str, drop_params: bool, - ) -> dict: # mutable-ok: inherited config contract returns a dict + ) -> dict[str, object]: # mutable-ok: inherited config contract returns a dict if not self.is_flux2_model(model): return super().map_openai_params( non_default_params=dict(non_default_params), From 4595b4f62f209d3d71734e8f1f2692f9a726e9a1 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Tue, 15 Sep 2026 13:01:12 -0500 Subject: [PATCH 041/253] fix(azure-ai): coerce FLUX controls and preserve response dimensions --- .../image_generation/flux_transformation.py | 5 +++++ .../image_generation/gpt_transformation.py | 9 ++++++++- .../test_azure_ai_image_edit_transformation.py | 6 +++--- .../test_azure_ai_flux2_image_generation.py | 18 ++++++++++++++++++ 4 files changed, 34 insertions(+), 4 deletions(-) diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index b9d5e11ff2a..997205b1fc3 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -85,6 +85,11 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): @staticmethod def _map_parameter(name: str, value: object) -> tuple[tuple[str, object], ...]: + if isinstance(value, str): + if name in ("n", "num_images", "width", "height", "steps", "seed", "safety_tolerance"): + return (("num_images" if name == "n" else name, int(value)),) + if name == "guidance": + return ((name, float(value)),) if name == "n": return (("num_images", value),) if name != "size": diff --git a/litellm/llms/openai/image_generation/gpt_transformation.py b/litellm/llms/openai/image_generation/gpt_transformation.py index 090b2eba387..8dc4d8953ea 100644 --- a/litellm/llms/openai/image_generation/gpt_transformation.py +++ b/litellm/llms/openai/image_generation/gpt_transformation.py @@ -82,7 +82,14 @@ class GPTImageGenerationConfig(BaseImageGenerationConfig): ) # set optional params - image_response.size = image_response.size or optional_params.get("size", "1024x1024") + width: Final = optional_params.get("width") + height: Final = optional_params.get("height") + requested_size: Final = ( + f"{width}x{height}" + if isinstance(width, int) and isinstance(height, int) + else optional_params.get("size", "1024x1024") + ) + image_response.size = image_response.size or requested_size image_response.quality = image_response.quality or optional_params.get("quality", "high") image_response.output_format = image_response.output_format or optional_params.get("output_format", "png") diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py index 94b23c5f1c2..d74afa88a6b 100644 --- a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py @@ -144,7 +144,7 @@ def test_flux2_image_edit_rejects_too_many_references(model: str, reference_imag ) -@pytest.mark.parametrize("dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024})) +@pytest.mark.parametrize("dimensions", ({"size": "2048x1024"}, {"width": 2048, "height": 1024}, {"width": "2048", "height": "1024"})) @pytest.mark.usefixtures("local_model_cost_map") def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[str, int | str]): def respond(request: httpx.Request) -> httpx.Response: @@ -170,8 +170,8 @@ def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[ api_base="https://example.services.ai.azure.com", client=client, n=2, - guidance=4.5, - steps=32, + guidance="4.5", + steps="32", **dimensions, ) diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py index 07026cf4309..c1d46b7e919 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py @@ -170,3 +170,21 @@ def test_flux2_cost_uses_mapped_dimensions_after_response_transformation(dimensi optional_params=params, call_type="image_generation", ) == pytest.approx(5e-08 * 2048 * 1024 * 2) + + +def test_flux2_response_preserves_mapped_dimensions(): + config = AzureFoundryFluxImageGenerationConfig() + params = config.map_openai_params( + non_default_params={"size": "2048x1024"}, optional_params={}, model="FLUX.2-flex", drop_params=False + ) + response = config.transform_image_generation_response( + model="FLUX.2-flex", + raw_response=httpx.Response(200, json={"data": [{"b64_json": "aW1n"}]}), + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={"prompt": "A landscape"}, + optional_params=params, + litellm_params={}, + encoding=None, + ) + assert response.size == "2048x1024" From 92e55b3b2262c40e41178435453edf3802819229 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 21:17:00 +0000 Subject: [PATCH 042/253] perf(proxy): split aggregated usage query into key-free rollups and bounded top-N keys Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 4 + .../common_daily_activity.py | 241 +++++++----- .../common_daily_activity.py | 5 + .../test_common_daily_activity.py | 355 +++++++++++++++--- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 5 files changed, 479 insertions(+), 131 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index ba5ec73d435..36bc16a0ae3 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2027,6 +2027,10 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: " PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__" PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job" PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900 +# Per-api_key rollups on the aggregated usage endpoint cover only the top N keys +# by spend so the result set stops growing with key count. Totals and the +# model/provider/endpoint rollups still cover every key. +USAGE_TOP_API_KEYS_LIMIT: Final[int] = 100 # Furthest back the catch-up pass looks for unpriced PTU days when a deployment # declares no ptu_effective_from, bounding the scan for an open-ended window. PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90 diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 44ed0017e42..4aafd416263 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -9,7 +9,7 @@ from fastapi import HTTPException, status from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger -from litellm.constants import PTU_SENTINEL_API_KEY +from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT from litellm.proxy._types import CommonProxyErrors from litellm.proxy.spend_tracking.key_metadata_recovery import ( attach_user_emails, @@ -169,6 +169,32 @@ class _EntityRollupRow(_GroupingSetsRow): api_key_rolled: int +class _AggregatedQueryKwargs(TypedDict): + """Filter arguments shared by the three aggregated SQL builders.""" + + table_name: ReadOnly[str] + entity_id_field: ReadOnly[str] + entity_id: ReadOnly[str | list[str] | None] + start_date: ReadOnly[str] + end_date: ReadOnly[str] + model: ReadOnly[str | None] + api_key: ReadOnly[str | list[str] | None] + exclude_entity_ids: ReadOnly[list[str] | None] + timezone_offset_minutes: ReadOnly[int | None] + include_current_utc_day: ReadOnly[bool] + + +_SqlQuery = tuple[str, list[str]] + + +async def _query_raw_optional( + prisma_client: PrismaClient, query: _SqlQuery | None +) -> list[dict[str, object]] | None: # mutable-ok: prisma query_raw return shape + if query is None: + return None + return await prisma_client.db.query_raw(query[0], *query[1]) + + def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float: """Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled. @@ -689,6 +715,27 @@ def _ptu_flat_cost_select(table_name: str) -> str: return "0::float AS ptu_flat_cost" +def _rollup_metric_select(table_name: str) -> str: + return f""" + SUM(spend)::float AS spend, + {_ptu_flat_cost_select(table_name)}, + SUM(prompt_tokens)::bigint AS prompt_tokens, + SUM(completion_tokens)::bigint AS completion_tokens, + SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, + SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, + SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, + SUM(compression_savings_spend)::float AS compression_savings_spend, + SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, + SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, + SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, + SUM(api_requests)::bigint AS api_requests, + SUM(successful_requests)::bigint AS successful_requests, + SUM(failed_requests)::bigint AS failed_requests""" + + +_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)" + + def _build_aggregated_sql_query( *, table_name: str, @@ -702,12 +749,16 @@ def _build_aggregated_sql_query( timezone_offset_minutes: int | None = None, include_current_utc_day: bool = False, ) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params - """Build a parameterized SQL GROUP BY query for aggregated daily activity. + """Build the key-free GROUPING SETS query for aggregated daily activity. - Groups by (date, api_key, model, model_group, custom_llm_provider, - mcp_namespaced_tool_name, endpoint) with SUMs on all metric columns. - The entity_id column is intentionally omitted from GROUP BY to collapse - rows across entities — this is where the biggest row reduction comes from. + Emits the grand total, per-date totals and the per-(date, model), model_group, + provider, mcp tool and endpoint rollups. api_key is never a grouping column here, + so the row count is bounded by dates x distinct models/providers/endpoints and + does not grow with the number of keys. Per-key rollups come from + _build_top_api_keys_sql_query. Both queries emit the same 7-bit group_level + bitmask (date, api_key, model, model_group, provider, mcp, endpoint); this one + hard-codes the api_key bit to "rolled up" so the dispatcher can consume the two + result sets as one stream. Returns: Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw(). @@ -730,14 +781,6 @@ def _build_aggregated_sql_query( exclude_entity_ids=exclude_entity_ids, ) - # Postgres computes every rollup level the response needs — per-date - # totals, per-(date, model), per-(date, model, api_key), per-provider, - # etc. — in a single pass via GROUPING SETS. The GROUPING() bitmask - # encodes which level a row belongs to so Python can dispatch rows - # straight into their buckets without re-summing. The leaf grouping - # is omitted on purpose: nothing in the response shape needs it once - # all the rollups are present. - # # TODO: drop the successful_requests/failed_requests aggregates (and the # total_successful_requests metadata they feed) once the admin UI reads SGR # only from LiteLLM_DailyGatewayRequests. The remaining spend, token and @@ -745,44 +788,25 @@ def _build_aggregated_sql_query( sql_query: Final = f""" SELECT date, - api_key, + NULL::text AS api_key, model, - COALESCE(NULLIF(model_group, ''), model) AS model_group, + {_MODEL_GROUP_EXPR} AS model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint, - GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model), - custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level, - SUM(spend)::float AS spend, - {_ptu_flat_cost_select(table_name)}, - SUM(prompt_tokens)::bigint AS prompt_tokens, - SUM(completion_tokens)::bigint AS completion_tokens, - SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, - SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, - SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, - SUM(compression_savings_spend)::float AS compression_savings_spend, - SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, - SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, - SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, - SUM(api_requests)::bigint AS api_requests, - SUM(successful_requests)::bigint AS successful_requests, - SUM(failed_requests)::bigint AS failed_requests + (GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT} + | GROUPING(model, {_MODEL_GROUP_EXPR}, + custom_llm_provider, mcp_namespaced_tool_name, + endpoint) AS group_level,{_rollup_metric_select(table_name)} FROM "{pg_table}" WHERE {where_clause} GROUP BY GROUPING SETS ( (date), - (date, api_key), (date, model), - (date, model, api_key), - (date, COALESCE(NULLIF(model_group, ''), model)), - (date, COALESCE(NULLIF(model_group, ''), model), api_key), + (date, {_MODEL_GROUP_EXPR}), (date, custom_llm_provider), - (date, custom_llm_provider, api_key), (date, mcp_namespaced_tool_name), - (date, mcp_namespaced_tool_name, api_key), (date, endpoint), - (date, endpoint, api_key), () ) """ @@ -790,6 +814,80 @@ def _build_aggregated_sql_query( return sql_query, sql_params +def _build_top_api_keys_sql_query( + *, + table_name: str, + entity_id_field: str, + entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path + start_date: str, + end_date: str, + model: str | None, + api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path + exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path + timezone_offset_minutes: int | None = None, + include_current_utc_day: bool = False, +) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params + """Per-key companion to _build_aggregated_sql_query. + + Ranks keys by spend over the same WHERE clause, keeps the top + USAGE_TOP_API_KEYS_LIMIT (ties broken by api_key so the set is stable across + refreshes) and emits the six (date, , api_key) rollups for those keys + only. The PTU flat-cost sentinel never ranks, so it cannot occupy a visible slot. + """ + pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name) + if pg_table is None: + raise ValueError(f"Unknown table name: {table_name}") + + adjusted_start, adjusted_end = _adjust_dates_for_timezone( + start_date, end_date, timezone_offset_minutes, include_current_utc_day + ) + + where_clause, where_params = _build_aggregated_where_clause( + entity_id_field=entity_id_field, + entity_id=entity_id, + adjusted_start=adjusted_start, + adjusted_end=adjusted_end, + model=model, + api_key=api_key, + exclude_entity_ids=exclude_entity_ids, + ) + sentinel_param: Final = f"${len(where_params) + 1}" + + sql_query: Final = f""" + WITH top_api_keys AS ( + SELECT api_key + FROM "{pg_table}" + WHERE {where_clause} AND api_key <> {sentinel_param} + GROUP BY api_key + ORDER BY SUM(spend) DESC, api_key + LIMIT {USAGE_TOP_API_KEYS_LIMIT} + ) + SELECT + date, + api_key, + model, + {_MODEL_GROUP_EXPR} AS model_group, + custom_llm_provider, + mcp_namespaced_tool_name, + endpoint, + GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR}, + custom_llm_provider, mcp_namespaced_tool_name, + endpoint) AS group_level,{_rollup_metric_select(table_name)} + FROM "{pg_table}" + WHERE {where_clause} AND api_key IN (SELECT api_key FROM top_api_keys) + GROUP BY GROUPING SETS ( + (date, api_key), + (date, model, api_key), + (date, {_MODEL_GROUP_EXPR}, api_key), + (date, custom_llm_provider, api_key), + (date, mcp_namespaced_tool_name, api_key), + (date, endpoint, api_key) + ) + """ + + return sql_query, [*where_params, PTU_SENTINEL_API_KEY] + + def _build_entity_rollup_sql_query( *, table_name: str, @@ -832,21 +930,7 @@ def _build_entity_rollup_sql_query( "{entity_id_field}" AS entity_id, date, api_key, - GROUPING(api_key) AS api_key_rolled, - SUM(spend)::float AS spend, - {_ptu_flat_cost_select(table_name)}, - SUM(prompt_tokens)::bigint AS prompt_tokens, - SUM(completion_tokens)::bigint AS completion_tokens, - SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, - SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, - SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, - SUM(compression_savings_spend)::float AS compression_savings_spend, - SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, - SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, - SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, - SUM(api_requests)::bigint AS api_requests, - SUM(successful_requests)::bigint AS successful_requests, - SUM(failed_requests)::bigint AS failed_requests + GROUPING(api_key) AS api_key_rolled,{_rollup_metric_select(table_name)} FROM "{pg_table}" WHERE {where_clause} GROUP BY GROUPING SETS ( @@ -948,6 +1032,7 @@ async def _aggregate_spend_records( # current grouping set's key), 0 when the column is part of the key. _GROUP_GRAND_TOTAL: Final = 127 # 0b1111111 — all rolled up _GROUP_DATE: Final = 63 # 0b0111111 — only date kept +_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000 — api_key position in the 7-bit mask _GROUP_DATE_API_KEY: Final = 31 # 0b0011111 _GROUP_DATE_MODEL: Final = 47 # 0b0101111 _GROUP_DATE_MODEL_API_KEY: Final = 15 # 0b0001111 @@ -1311,9 +1396,11 @@ async def get_daily_activity_aggregated( ) -> SpendAnalyticsPaginatedResponse: """Aggregated variant that returns the full result set (no pagination). - Uses SQL GROUP BY to aggregate rows in the database rather than fetching - all individual rows into Python. This collapses rows across entities - (users/teams/orgs), reducing ~150k rows to ~2-3k grouped rows. + Runs two GROUPING SETS queries in parallel: a key-free one for totals and the + model/provider/mcp/endpoint rollups (row count independent of key cardinality) + and a bounded one for the per-key rollups of the top USAGE_TOP_API_KEYS_LIMIT + keys by spend. breakdown.api_keys and every api_key_breakdown therefore list at + most that many keys, while the totals and the key-free rollups cover every key. include_entity_breakdown runs a small companion rollup query and folds `breakdown.entities` onto the response, as entity-scoped views like Team Usage need. @@ -1333,7 +1420,7 @@ async def get_daily_activity_aggregated( ) try: - sql_query, sql_params = _build_aggregated_sql_query( + query_kwargs: Final = _AggregatedQueryKwargs( table_name=table_name, entity_id_field=entity_id_field, entity_id=entity_id, @@ -1345,36 +1432,17 @@ async def get_daily_activity_aggregated( timezone_offset_minutes=timezone_offset_minutes, include_current_utc_day=include_current_utc_day, ) + key_free_sql, key_free_params = _build_aggregated_sql_query(**query_kwargs) + top_keys_sql, top_keys_params = _build_top_api_keys_sql_query(**query_kwargs) + entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None - entity_query: Final = ( - _build_entity_rollup_sql_query( - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - start_date=start_date, - end_date=end_date, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - timezone_offset_minutes=timezone_offset_minutes, - include_current_utc_day=include_current_utc_day, - ) - if include_entity_breakdown - else None + raw_key_free_rows, raw_top_key_rows, raw_entity_rows = await asyncio.gather( + prisma_client.db.query_raw(key_free_sql, *key_free_params), + prisma_client.db.query_raw(top_keys_sql, *top_keys_params), + _query_raw_optional(prisma_client, entity_query), ) - # Execute the GROUPING SETS query (one row per rollup level), alongside - # the per-entity companion rollup when the caller wants entities. - raw_rows, raw_entity_rows = ( - await asyncio.gather( - prisma_client.db.query_raw(sql_query, *sql_params), - prisma_client.db.query_raw(entity_query[0], *entity_query[1]), - ) - if entity_query is not None - else (await prisma_client.db.query_raw(sql_query, *sql_params), None) - ) - - records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or [])] + records: Final = [_GroupingSetsRow(**row) for row in (*(raw_key_free_rows or ()), *(raw_top_key_rows or ()))] # The grouping-sets dispatcher places each row directly in its bucket # using the row's GROUPING() bitmask. No Python-side summing needed. @@ -1426,6 +1494,7 @@ async def get_daily_activity_aggregated( page=1, total_pages=1, has_more=False, + api_key_limit=USAGE_TOP_API_KEYS_LIMIT, ), ) diff --git a/litellm/types/proxy/management_endpoints/common_daily_activity.py b/litellm/types/proxy/management_endpoints/common_daily_activity.py index 090e5c42376..f58d40dfa9c 100644 --- a/litellm/types/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/types/proxy/management_endpoints/common_daily_activity.py @@ -96,6 +96,11 @@ class DailySpendMetadata(BaseModel): page: int = Field(default=1) total_pages: int = Field(default=1) has_more: bool = Field(default=False) + api_key_limit: int | None = Field( + default=None, + description="When set, api_keys and every api_key_breakdown list at most this many keys, " + "ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.", + ) class SpendAnalyticsPaginatedResponse(BaseModel): diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 71896a18f48..dec98a1da27 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1,17 +1,24 @@ +import re +from collections.abc import Sequence from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock +import psycopg import pytest +from psycopg.rows import dict_row +from pytest_postgresql import factories from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR +from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT from litellm.proxy.management_endpoints.common_daily_activity import ( _adjust_dates_for_timezone, _build_aggregated_sql_query, _build_entity_rollup_sql_query, + _build_top_api_keys_sql_query, _is_user_agent_tag, _record_to_spend_metrics, get_api_key_metadata, @@ -159,7 +166,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "autorouter_savings_spend": 0.0, "failed_requests": 0, } - mock_rows = [ + key_free_rows = [ # (date, endpoint) — rolls up across api_keys and models { **base, @@ -185,31 +192,6 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "api_requests": 1, "successful_requests": 1, }, - # (date, endpoint, api_key) — populates the per-key sub-bucket - { - **base, - "date": "2024-01-01", - "endpoint": "/v1/chat/completions", - "api_key": "key-1", - "group_level": 30, - "spend": 15.0, - "prompt_tokens": 150, - "completion_tokens": 75, - "api_requests": 2, - "successful_requests": 2, - }, - { - **base, - "date": "2024-01-01", - "endpoint": "/v1/embeddings", - "api_key": "key-2", - "group_level": 30, - "spend": 3.0, - "prompt_tokens": 30, - "completion_tokens": 0, - "api_requests": 1, - "successful_requests": 1, - }, # (date) — per-date totals { **base, @@ -237,8 +219,35 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "successful_requests": 3, }, ] + top_key_rows = [ + # (date, endpoint, api_key) — populates the per-key sub-bucket + { + **base, + "date": "2024-01-01", + "endpoint": "/v1/chat/completions", + "api_key": "key-1", + "group_level": 30, + "spend": 15.0, + "prompt_tokens": 150, + "completion_tokens": 75, + "api_requests": 2, + "successful_requests": 2, + }, + { + **base, + "date": "2024-01-01", + "endpoint": "/v1/embeddings", + "api_key": "key-2", + "group_level": 30, + "spend": 3.0, + "prompt_tokens": 30, + "completion_tokens": 0, + "api_requests": 1, + "successful_requests": 1, + }, + ] - mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) + mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows]) mock_prisma.db.litellm_verificationtoken = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -284,8 +293,11 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): assert "key-2" in embeddings_endpoint.api_key_breakdown assert embeddings_endpoint.api_key_breakdown["key-2"].metrics.spend == 3.0 - # Verify query_raw was called (not find_many) - mock_prisma.db.query_raw.assert_called_once() + # One key-free rollup query plus one bounded per-key query, no find_many + assert mock_prisma.db.query_raw.call_count == 2 + key_free_sql, top_keys_sql = (call.args[0] for call in mock_prisma.db.query_raw.call_args_list) + assert "top_api_keys" not in key_free_sql + assert "WITH top_api_keys AS" in top_keys_sql @pytest.mark.asyncio @@ -812,7 +824,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "autorouter_savings_spend": 0.0, "failed_requests": 0, } - mock_rows = [ + key_free_rows = [ { **base, "date": "2024-01-01", @@ -825,6 +837,8 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "api_requests": 1, "successful_requests": 1, }, + ] + top_key_rows = [ { **base, "date": "2024-01-01", @@ -839,7 +853,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): }, ] - mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) + mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows]) # Active table returns nothing for this key mock_prisma.db.litellm_verificationtoken = MagicMock() @@ -1240,17 +1254,61 @@ class TestBuildAggregatedSqlQuery: normalized = " ".join(sql.split()) fallback = "COALESCE(NULLIF(model_group, ''), model)" assert f"{fallback} AS model_group" in normalized - assert ( - f"GROUPING(date, api_key, model, {fallback}, " - "custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized - ) - assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized + assert f"GROUPING(model, {fallback}, custom_llm_provider, mcp_namespaced_tool_name, endpoint)" in normalized + assert f"(date, {fallback})," in normalized assert "(date, model_group)" not in normalized assert "COALESCE(model_group, model)" not in normalized + def test_key_free_query_never_groups_by_api_key(self): + """The main rollup query must not emit one row per key, that is what blew up + the query engine at 3k+ keys. Every grouping set stays key-free and api_key + is projected as a NULL literal so the dispatcher's row shape is unchanged.""" + sql, _ = _build_aggregated_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + start_date="2026-07-01", + end_date="2026-07-01", + model=None, + api_key=None, + ) + + normalized = " ".join(sql.split()) + grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1] + assert "api_key" not in grouping_block + assert "NULL::text AS api_key" in normalized + + def test_top_api_keys_query_ranks_keys_deterministically_and_shares_filters(self): + sql, params = _build_top_api_keys_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id="user-1", + start_date="2026-05-29", + end_date="2026-06-02", + model="bedrock/global.anthropic.claude-opus-4-8", + api_key="sk-test", + timezone_offset_minutes=-330, + ) + + normalized = " ".join(sql.split()) + assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in normalized + assert "api_key IN (SELECT api_key FROM top_api_keys)" in normalized + assert "api_key <> $6" in normalized + grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1] + assert grouping_block.count(", api_key)") == 6 + assert grouping_block.count("(date") == 6 + assert params == [ + "2026-05-29", + "2026-06-02", + "user-1", + "bedrock/global.anthropic.claude-opus-4-8", + "sk-test", + PTU_SENTINEL_API_KEY, + ] + class TestAggregatedEmptyEntityFilter: - _BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query) + _BUILDERS: Final = (_build_aggregated_sql_query, _build_top_api_keys_sql_query, _build_entity_rollup_sql_query) @pytest.mark.parametrize("build", _BUILDERS) def test_empty_entity_list_emits_no_degenerate_in_clause(self, build): @@ -1267,7 +1325,8 @@ class TestAggregatedEmptyEntityFilter: normalized = " ".join(sql.split()) assert "IN ()" not in normalized assert '"team_id" IN' not in normalized - assert params == ["2026-08-01", "2026-08-19"] + sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else [] + assert params == ["2026-08-01", "2026-08-19", *sentinel_params] @pytest.mark.parametrize("build", _BUILDERS) def test_empty_entity_list_matches_nothing_rather_than_everything(self, build): @@ -1298,7 +1357,8 @@ class TestAggregatedEmptyEntityFilter: normalized = " ".join(sql.split()) assert '"team_id" IN ($3, $4)' in normalized assert "FALSE" not in normalized - assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta"] + sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else [] + assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta", *sentinel_params] @pytest.mark.asyncio @@ -1313,7 +1373,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() - mock_rows = [ + key_free_rows = [ { "date": None, "api_key": None, @@ -1338,7 +1398,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): "failed_requests": None, } ] - mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) + mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, []]) result = await get_daily_activity_aggregated( prisma_client=mock_prisma, @@ -1365,6 +1425,211 @@ async def test_get_daily_activity_aggregated_empty_result_set(): assert result.metadata.total_compression_saved_tokens == 0 +_aggregated_postgresql_proc: Final = factories.postgresql_proc() +_aggregated_postgresql: Final = factories.postgresql("_aggregated_postgresql_proc") + +_DAILY_USER_SPEND_DDL: Final = """ + CREATE TABLE "LiteLLM_DailyUserSpend" ( + id TEXT PRIMARY KEY, + user_id TEXT, + date TEXT NOT NULL, + api_key TEXT NOT NULL, + model TEXT, + model_group TEXT, + custom_llm_provider TEXT, + mcp_namespaced_tool_name TEXT, + endpoint TEXT, + prompt_tokens BIGINT DEFAULT 0, + completion_tokens BIGINT DEFAULT 0, + cache_read_input_tokens BIGINT DEFAULT 0, + cache_creation_input_tokens BIGINT DEFAULT 0, + compression_saved_tokens BIGINT DEFAULT 0, + compression_savings_spend DOUBLE PRECISION DEFAULT 0, + prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + autorouter_savings_spend DOUBLE PRECISION DEFAULT 0, + spend DOUBLE PRECISION DEFAULT 0, + api_requests BIGINT DEFAULT 0, + successful_requests BIGINT DEFAULT 0, + failed_requests BIGINT DEFAULT 0 + ) +""" + + +def _seed_daily_user_spend(conn: psycopg.Connection, rows: Sequence[tuple[object, ...]]) -> None: + with conn.cursor() as cur: + cur.execute(_DAILY_USER_SPEND_DDL) + cur.executemany( + """ + INSERT INTO "LiteLLM_DailyUserSpend" + (id, user_id, date, api_key, model, model_group, custom_llm_provider, + endpoint, prompt_tokens, spend, api_requests, successful_requests) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """, + rows, + ) + conn.commit() + + +def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]): + """Run the proxy's $N-parameterized SQL through psycopg, recording each result size.""" + + async def query_raw(sql: str, *params: str) -> list[dict[str, object]]: + converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql) + with conn.cursor(row_factory=dict_row) as cur: + cur.execute( + converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query + {f"p{i}": v for i, v in enumerate(params, start=1)}, + ) + rows: Final = cur.fetchall() + row_counts.append(len(rows)) + return rows + + return query_raw + + +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_bounds_api_key_rollups( + _aggregated_postgresql: psycopg.Connection, +): + """Run both GROUPING SETS queries against real Postgres with more keys than the cap. + + key-004 and key-005 tie on spend exactly at the USAGE_TOP_API_KEYS_LIMIT + cutoff; the api_key tiebreaker must keep key-004 and drop key-005. The PTU + sentinel outspends every key but must not take a slot. Excluded keys and the + sentinel still count toward the totals and the model rollup, which come from + the key-free query. + """ + n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 5 + key_rows: Final = [ + ( + f"row-{i:03d}", + f"user-{i:03d}", + "2026-06-01", + f"key-{i:03d}", + "gpt-5", + "", + "openai", + "/v1/chat/completions", + 10, + 6.0 if i == 4 else float(i + 1), + 1, + 1, + ) + for i in range(n_keys) + ] + sentinel_row: Final = ( + "row-ptu", + None, + "2026-06-01", + PTU_SENTINEL_API_KEY, + "gpt-5", + "", + "azure", + None, + 0, + 1000.0, + 0, + 0, + ) + _seed_daily_user_spend(_aggregated_postgresql, [*key_rows, sentinel_row]) + key_spend: Final = sum(6.0 if i == 4 else float(i + 1) for i in range(n_keys)) + + row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity_aggregated( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + api_key=None, + ) + + # Key-free query: (), (date), (date, model), (date, model_group), two providers, + # one mcp NULL bucket, endpoint plus its NULL bucket = 9 rows regardless of key count. + # Top-key query: six per-key grouping sets, each capped at the limit. + assert row_counts == [9, 6 * USAGE_TOP_API_KEYS_LIMIT] + + assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0) + assert result.metadata.total_api_requests == n_keys + assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT + + expected_top: Final = {f"key-{i:03d}" for i in range(6, n_keys)} | {"key-004"} + day: Final = result.results[0] + assert day.metrics.spend == pytest.approx(key_spend + 1000.0) + assert set(day.breakdown.api_keys) == expected_top + assert day.breakdown.api_keys["key-004"].metrics.spend == 6.0 + assert "key-005" not in day.breakdown.api_keys + assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys + + assert day.breakdown.models["gpt-5"].metrics.spend == pytest.approx(key_spend + 1000.0) + assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == expected_top + assert day.breakdown.providers["openai"].metrics.spend == pytest.approx(key_spend) + assert set(day.breakdown.providers["openai"].api_key_breakdown) == expected_top + assert day.breakdown.endpoints["/v1/chat/completions"].metrics.api_requests == n_keys + + +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_queries( + _aggregated_postgresql: psycopg.Connection, +): + """An explicit api_key filter must scope the key-free totals and the per-key + rollups to that key alone, so the two result sets never disagree.""" + rows: Final = [ + ( + f"row-{i}", + f"user-{i}", + "2026-06-01", + f"key-{i}", + "gpt-5", + "", + "openai", + "/v1/chat/completions", + 10, + float(i + 1), + 1, + 1, + ) + for i in range(3) + ] + _seed_daily_user_spend(_aggregated_postgresql, rows) + + row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity_aggregated( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + api_key="key-1", + ) + + assert result.metadata.total_spend == 2.0 + day: Final = result.results[0] + assert set(day.breakdown.api_keys) == {"key-1"} + assert day.breakdown.api_keys["key-1"].metrics.spend == 2.0 + assert day.breakdown.models["gpt-5"].metrics.spend == 2.0 + assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"} + + def _no_spend_record(): """A rollup row for a key with no spend, where SUM() returns NULL (None).""" return SimpleNamespace( @@ -2128,8 +2393,8 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): {**base, "date": None, "group_level": 127, "spend": 18.0}, {**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0}, {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0}, - {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}, ] + top_key_rows = [{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}] entity_base = { key: value for key, value in base.items() @@ -2156,7 +2421,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): }, ] - mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, entity_rows]) + mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, top_key_rows, entity_rows]) mock_prisma.db.litellm_verificationtoken = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -2173,9 +2438,9 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): include_entity_breakdown=True, ) - assert mock_prisma.db.query_raw.call_count == 2 + assert mock_prisma.db.query_raw.call_count == 3 main_sql = mock_prisma.db.query_raw.call_args_list[0][0][0] - entity_sql = mock_prisma.db.query_raw.call_args_list[1][0][0] + entity_sql = mock_prisma.db.query_raw.call_args_list[2][0][0] assert "entity_id" not in main_sql assert '"team_id" AS entity_id' in entity_sql assert '(date, "team_id"),' in entity_sql diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 4c16a8613eb..f9b234e0298 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27430,6 +27430,11 @@ export interface components { }; /** DailySpendMetadata */ DailySpendMetadata: { + /** + * Api Key Limit + * @description When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key. + */ + api_key_limit?: number | null; /** * Has More * @default false From a0311dddf773c9558b4da9732a4a15c9646e9336 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 21:28:09 +0000 Subject: [PATCH 043/253] chore(proxy): regenerate lazy OpenAPI snapshot for api_key_limit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index f3b579d22c7..a3f525f0773 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3050,6 +3050,18 @@ }, "DailySpendMetadata": { "properties": { + "api_key_limit": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "description": "When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.", + "title": "Api Key Limit" + }, "has_more": { "default": false, "title": "Has More", From f94a40f841c3dcf4498d4cb7acef7b482cf91dc5 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 21:52:37 +0000 Subject: [PATCH 044/253] perf(proxy): serve key-free rollups and top-N keys from one UNION ALL statement Both arms now run in a single query_raw call so totals and per-key breakdowns come from the same snapshot. USAGE_TOP_API_KEYS_LIMIT can be raised via env for deployments that need every key in the response. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 5 +- .../common_daily_activity.py | 134 ++++++------------ .../test_common_daily_activity.py | 98 ++++++------- 3 files changed, 87 insertions(+), 150 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 36bc16a0ae3..565c6433c6e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2027,10 +2027,7 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: " PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__" PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job" PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900 -# Per-api_key rollups on the aggregated usage endpoint cover only the top N keys -# by spend so the result set stops growing with key count. Totals and the -# model/provider/endpoint rollups still cover every key. -USAGE_TOP_API_KEYS_LIMIT: Final[int] = 100 +USAGE_TOP_API_KEYS_LIMIT: Final[int] = int(os.getenv("USAGE_TOP_API_KEYS_LIMIT", "100")) # Furthest back the catch-up pass looks for unpriced PTU days when a deployment # declares no ptu_effective_from, bounding the scan for an open-ended window. PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90 diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 4aafd416263..95a89460d90 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -170,8 +170,6 @@ class _EntityRollupRow(_GroupingSetsRow): class _AggregatedQueryKwargs(TypedDict): - """Filter arguments shared by the three aggregated SQL builders.""" - table_name: ReadOnly[str] entity_id_field: ReadOnly[str] entity_id: ReadOnly[str | list[str] | None] @@ -749,16 +747,14 @@ def _build_aggregated_sql_query( timezone_offset_minutes: int | None = None, include_current_utc_day: bool = False, ) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params - """Build the key-free GROUPING SETS query for aggregated daily activity. + """Build the GROUPING SETS query for aggregated daily activity. - Emits the grand total, per-date totals and the per-(date, model), model_group, - provider, mcp tool and endpoint rollups. api_key is never a grouping column here, - so the row count is bounded by dates x distinct models/providers/endpoints and - does not grow with the number of keys. Per-key rollups come from - _build_top_api_keys_sql_query. Both queries emit the same 7-bit group_level - bitmask (date, api_key, model, model_group, provider, mcp, endpoint); this one - hard-codes the api_key bit to "rolled up" so the dispatcher can consume the two - result sets as one stream. + One statement, two UNION ALL arms over the same WHERE clause. The first arm is + key-free: grand total, per-date totals and the (date, model / model_group / + provider / mcp / endpoint) rollups, so its row count never grows with the number + of keys. The second arm emits the (date, , api_key) rollups for the + USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit + group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint). Returns: Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw(). @@ -771,77 +767,6 @@ def _build_aggregated_sql_query( start_date, end_date, timezone_offset_minutes, include_current_utc_day ) - where_clause, sql_params = _build_aggregated_where_clause( - entity_id_field=entity_id_field, - entity_id=entity_id, - adjusted_start=adjusted_start, - adjusted_end=adjusted_end, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - ) - - # TODO: drop the successful_requests/failed_requests aggregates (and the - # total_successful_requests metadata they feed) once the admin UI reads SGR - # only from LiteLLM_DailyGatewayRequests. The remaining spend, token and - # api_requests rollups are still served from here. - sql_query: Final = f""" - SELECT - date, - NULL::text AS api_key, - model, - {_MODEL_GROUP_EXPR} AS model_group, - custom_llm_provider, - mcp_namespaced_tool_name, - endpoint, - (GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT} - | GROUPING(model, {_MODEL_GROUP_EXPR}, - custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level,{_rollup_metric_select(table_name)} - FROM "{pg_table}" - WHERE {where_clause} - GROUP BY GROUPING SETS ( - (date), - (date, model), - (date, {_MODEL_GROUP_EXPR}), - (date, custom_llm_provider), - (date, mcp_namespaced_tool_name), - (date, endpoint), - () - ) - """ - - return sql_query, sql_params - - -def _build_top_api_keys_sql_query( - *, - table_name: str, - entity_id_field: str, - entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - start_date: str, - end_date: str, - model: str | None, - api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path - timezone_offset_minutes: int | None = None, - include_current_utc_day: bool = False, -) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params - """Per-key companion to _build_aggregated_sql_query. - - Ranks keys by spend over the same WHERE clause, keeps the top - USAGE_TOP_API_KEYS_LIMIT (ties broken by api_key so the set is stable across - refreshes) and emits the six (date, , api_key) rollups for those keys - only. The PTU flat-cost sentinel never ranks, so it cannot occupy a visible slot. - """ - pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name) - if pg_table is None: - raise ValueError(f"Unknown table name: {table_name}") - - adjusted_start, adjusted_end = _adjust_dates_for_timezone( - start_date, end_date, timezone_offset_minutes, include_current_utc_day - ) - where_clause, where_params = _build_aggregated_where_clause( entity_id_field=entity_id_field, entity_id=entity_id, @@ -852,9 +777,38 @@ def _build_top_api_keys_sql_query( exclude_entity_ids=exclude_entity_ids, ) sentinel_param: Final = f"${len(where_params) + 1}" + metric_select: Final = _rollup_metric_select(table_name) + # TODO: drop the successful_requests/failed_requests aggregates (and the + # total_successful_requests metadata they feed) once the admin UI reads SGR + # only from LiteLLM_DailyGatewayRequests. The remaining spend, token and + # api_requests rollups are still served from here. sql_query: Final = f""" - WITH top_api_keys AS ( + (SELECT + date, + NULL::text AS api_key, + model, + {_MODEL_GROUP_EXPR} AS model_group, + custom_llm_provider, + mcp_namespaced_tool_name, + endpoint, + (GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT} + | GROUPING(model, {_MODEL_GROUP_EXPR}, + custom_llm_provider, mcp_namespaced_tool_name, + endpoint) AS group_level,{metric_select} + FROM "{pg_table}" + WHERE {where_clause} + GROUP BY GROUPING SETS ( + (date), + (date, model), + (date, {_MODEL_GROUP_EXPR}), + (date, custom_llm_provider), + (date, mcp_namespaced_tool_name), + (date, endpoint), + () + )) + UNION ALL + (WITH top_api_keys AS ( SELECT api_key FROM "{pg_table}" WHERE {where_clause} AND api_key <> {sentinel_param} @@ -872,7 +826,7 @@ def _build_top_api_keys_sql_query( endpoint, GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR}, custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level,{_rollup_metric_select(table_name)} + endpoint) AS group_level,{metric_select} FROM "{pg_table}" WHERE {where_clause} AND api_key IN (SELECT api_key FROM top_api_keys) GROUP BY GROUPING SETS ( @@ -882,7 +836,7 @@ def _build_top_api_keys_sql_query( (date, custom_llm_provider, api_key), (date, mcp_namespaced_tool_name, api_key), (date, endpoint, api_key) - ) + )) """ return sql_query, [*where_params, PTU_SENTINEL_API_KEY] @@ -1432,17 +1386,15 @@ async def get_daily_activity_aggregated( timezone_offset_minutes=timezone_offset_minutes, include_current_utc_day=include_current_utc_day, ) - key_free_sql, key_free_params = _build_aggregated_sql_query(**query_kwargs) - top_keys_sql, top_keys_params = _build_top_api_keys_sql_query(**query_kwargs) + sql_query, sql_params = _build_aggregated_sql_query(**query_kwargs) entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None - raw_key_free_rows, raw_top_key_rows, raw_entity_rows = await asyncio.gather( - prisma_client.db.query_raw(key_free_sql, *key_free_params), - prisma_client.db.query_raw(top_keys_sql, *top_keys_params), + raw_rows, raw_entity_rows = await asyncio.gather( + prisma_client.db.query_raw(sql_query, *sql_params), _query_raw_optional(prisma_client, entity_query), ) - records: Final = [_GroupingSetsRow(**row) for row in (*(raw_key_free_rows or ()), *(raw_top_key_rows or ()))] + records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or ())] # The grouping-sets dispatcher places each row directly in its bucket # using the row's GROUPING() bitmask. No Python-side summing needed. diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index dec98a1da27..5ff3f89343b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -18,7 +18,6 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( _adjust_dates_for_timezone, _build_aggregated_sql_query, _build_entity_rollup_sql_query, - _build_top_api_keys_sql_query, _is_user_agent_tag, _record_to_spend_metrics, get_api_key_metadata, @@ -166,7 +165,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "autorouter_savings_spend": 0.0, "failed_requests": 0, } - key_free_rows = [ + mock_rows = [ # (date, endpoint) — rolls up across api_keys and models { **base, @@ -218,8 +217,6 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "api_requests": 3, "successful_requests": 3, }, - ] - top_key_rows = [ # (date, endpoint, api_key) — populates the per-key sub-bucket { **base, @@ -247,7 +244,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): }, ] - mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows]) + mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) mock_prisma.db.litellm_verificationtoken = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -293,11 +290,8 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): assert "key-2" in embeddings_endpoint.api_key_breakdown assert embeddings_endpoint.api_key_breakdown["key-2"].metrics.spend == 3.0 - # One key-free rollup query plus one bounded per-key query, no find_many - assert mock_prisma.db.query_raw.call_count == 2 - key_free_sql, top_keys_sql = (call.args[0] for call in mock_prisma.db.query_raw.call_args_list) - assert "top_api_keys" not in key_free_sql - assert "WITH top_api_keys AS" in top_keys_sql + # Verify query_raw was called (not find_many) + mock_prisma.db.query_raw.assert_called_once() @pytest.mark.asyncio @@ -484,9 +478,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash( return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com")] ) mock_prisma.db.query_raw = AsyncMock( - return_value=[ - {"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"} - ] + return_value=[{"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"}] ) result = await get_api_key_metadata( @@ -824,7 +816,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "autorouter_savings_spend": 0.0, "failed_requests": 0, } - key_free_rows = [ + mock_rows = [ { **base, "date": "2024-01-01", @@ -837,8 +829,6 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "api_requests": 1, "successful_requests": 1, }, - ] - top_key_rows = [ { **base, "date": "2024-01-01", @@ -853,7 +843,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): }, ] - mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, top_key_rows]) + mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) # Active table returns nothing for this key mock_prisma.db.litellm_verificationtoken = MagicMock() @@ -1226,6 +1216,7 @@ class TestBuildAggregatedSqlQuery: "user-1", "bedrock/global.anthropic.claude-opus-4-8", "sk-test", + PTU_SENTINEL_API_KEY, ] assert "model = $4" in sql assert "api_key = $5" in sql @@ -1259,9 +1250,9 @@ class TestBuildAggregatedSqlQuery: assert "(date, model_group)" not in normalized assert "COALESCE(model_group, model)" not in normalized - def test_key_free_query_never_groups_by_api_key(self): - """The main rollup query must not emit one row per key, that is what blew up - the query engine at 3k+ keys. Every grouping set stays key-free and api_key + def test_totals_arm_never_groups_by_api_key(self): + """The totals arm must not emit one row per key, that is what blew up the + query engine at 3k+ keys. Every grouping set there stays key-free and api_key is projected as a NULL literal so the dispatcher's row shape is unchanged.""" sql, _ = _build_aggregated_sql_query( table_name="litellm_dailyuserspend", @@ -1273,13 +1264,15 @@ class TestBuildAggregatedSqlQuery: api_key=None, ) - normalized = " ".join(sql.split()) - grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1] + totals_arm, _ = " ".join(sql.split()).split("UNION ALL") + grouping_block = totals_arm.split("GROUP BY GROUPING SETS (", 1)[1] assert "api_key" not in grouping_block - assert "NULL::text AS api_key" in normalized + assert "NULL::text AS api_key" in totals_arm - def test_top_api_keys_query_ranks_keys_deterministically_and_shares_filters(self): - sql, params = _build_top_api_keys_sql_query( + def test_per_key_arm_ranks_keys_deterministically_and_shares_filters(self): + """Both arms sit in one statement so totals and per-key rows come from the + same snapshot, and the per-key arm reuses the caller's filter params.""" + sql, params = _build_aggregated_sql_query( table_name="litellm_dailyuserspend", entity_id_field="user_id", entity_id="user-1", @@ -1290,25 +1283,20 @@ class TestBuildAggregatedSqlQuery: timezone_offset_minutes=-330, ) - normalized = " ".join(sql.split()) - assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in normalized - assert "api_key IN (SELECT api_key FROM top_api_keys)" in normalized - assert "api_key <> $6" in normalized - grouping_block = normalized.split("GROUP BY GROUPING SETS (", 1)[1] + totals_arm, per_key_arm = " ".join(sql.split()).split("UNION ALL") + assert "top_api_keys" not in totals_arm + assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in per_key_arm + assert "api_key IN (SELECT api_key FROM top_api_keys)" in per_key_arm + assert "api_key <> $6" in per_key_arm + assert per_key_arm.count("model = $4 AND api_key = $5") == 2 + grouping_block = per_key_arm.split("GROUP BY GROUPING SETS (", 1)[1] assert grouping_block.count(", api_key)") == 6 assert grouping_block.count("(date") == 6 - assert params == [ - "2026-05-29", - "2026-06-02", - "user-1", - "bedrock/global.anthropic.claude-opus-4-8", - "sk-test", - PTU_SENTINEL_API_KEY, - ] + assert params[-1] == PTU_SENTINEL_API_KEY class TestAggregatedEmptyEntityFilter: - _BUILDERS: Final = (_build_aggregated_sql_query, _build_top_api_keys_sql_query, _build_entity_rollup_sql_query) + _BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query) @pytest.mark.parametrize("build", _BUILDERS) def test_empty_entity_list_emits_no_degenerate_in_clause(self, build): @@ -1325,7 +1313,7 @@ class TestAggregatedEmptyEntityFilter: normalized = " ".join(sql.split()) assert "IN ()" not in normalized assert '"team_id" IN' not in normalized - sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else [] + sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else [] assert params == ["2026-08-01", "2026-08-19", *sentinel_params] @pytest.mark.parametrize("build", _BUILDERS) @@ -1357,7 +1345,7 @@ class TestAggregatedEmptyEntityFilter: normalized = " ".join(sql.split()) assert '"team_id" IN ($3, $4)' in normalized assert "FALSE" not in normalized - sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_top_api_keys_sql_query else [] + sentinel_params = [PTU_SENTINEL_API_KEY] if build is _build_aggregated_sql_query else [] assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta", *sentinel_params] @@ -1373,7 +1361,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): mock_prisma = MagicMock() mock_prisma.db = MagicMock() - key_free_rows = [ + mock_rows = [ { "date": None, "api_key": None, @@ -1398,7 +1386,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): "failed_requests": None, } ] - mock_prisma.db.query_raw = AsyncMock(side_effect=[key_free_rows, []]) + mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) result = await get_daily_activity_aggregated( prisma_client=mock_prisma, @@ -1492,13 +1480,13 @@ def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]): async def test_get_daily_activity_aggregated_bounds_api_key_rollups( _aggregated_postgresql: psycopg.Connection, ): - """Run both GROUPING SETS queries against real Postgres with more keys than the cap. + """Run the GROUPING SETS statement against real Postgres with more keys than the cap. key-004 and key-005 tie on spend exactly at the USAGE_TOP_API_KEYS_LIMIT cutoff; the api_key tiebreaker must keep key-004 and drop key-005. The PTU sentinel outspends every key but must not take a slot. Excluded keys and the sentinel still count toward the totals and the model rollup, which come from - the key-free query. + the key-free arm. """ n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 5 key_rows: Final = [ @@ -1554,10 +1542,10 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups( api_key=None, ) - # Key-free query: (), (date), (date, model), (date, model_group), two providers, + # Key-free arm: (), (date), (date, model), (date, model_group), two providers, # one mcp NULL bucket, endpoint plus its NULL bucket = 9 rows regardless of key count. - # Top-key query: six per-key grouping sets, each capped at the limit. - assert row_counts == [9, 6 * USAGE_TOP_API_KEYS_LIMIT] + # Per-key arm: six per-key grouping sets, each capped at the limit. + assert row_counts == [9 + 6 * USAGE_TOP_API_KEYS_LIMIT] assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0) assert result.metadata.total_api_requests == n_keys @@ -1579,11 +1567,11 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups( @pytest.mark.asyncio -async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_queries( +async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both_arms( _aggregated_postgresql: psycopg.Connection, ): """An explicit api_key filter must scope the key-free totals and the per-key - rollups to that key alone, so the two result sets never disagree.""" + rollups to that key alone, so the two arms never disagree.""" rows: Final = [ ( f"row-{i}", @@ -2358,7 +2346,7 @@ def test_entity_rollup_sql_query_and_api_key_list_filter(): api_key=[], ) assert "FALSE" in empty_sql - assert empty_params == ["2024-01-01", "2024-01-31"] + assert empty_params == ["2024-01-01", "2024-01-31", PTU_SENTINEL_API_KEY] @pytest.mark.asyncio @@ -2393,8 +2381,8 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): {**base, "date": None, "group_level": 127, "spend": 18.0}, {**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0}, {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0}, + {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}, ] - top_key_rows = [{**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}] entity_base = { key: value for key, value in base.items() @@ -2421,7 +2409,7 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): }, ] - mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, top_key_rows, entity_rows]) + mock_prisma.db.query_raw = AsyncMock(side_effect=[main_rows, entity_rows]) mock_prisma.db.litellm_verificationtoken = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) @@ -2438,9 +2426,9 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): include_entity_breakdown=True, ) - assert mock_prisma.db.query_raw.call_count == 3 + assert mock_prisma.db.query_raw.call_count == 2 main_sql = mock_prisma.db.query_raw.call_args_list[0][0][0] - entity_sql = mock_prisma.db.query_raw.call_args_list[2][0][0] + entity_sql = mock_prisma.db.query_raw.call_args_list[1][0][0] assert "entity_id" not in main_sql assert '"team_id" AS entity_id' in entity_sql assert '(date, "team_id"),' in entity_sql From c3e937b84544902b2de41a2748def7803b8f3368 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 22:11:20 +0000 Subject: [PATCH 045/253] docs(proxy): describe the single UNION ALL aggregate statement in the endpoint docstring Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/common_daily_activity.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 95a89460d90..8a3ba196ab2 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -1350,11 +1350,12 @@ async def get_daily_activity_aggregated( ) -> SpendAnalyticsPaginatedResponse: """Aggregated variant that returns the full result set (no pagination). - Runs two GROUPING SETS queries in parallel: a key-free one for totals and the - model/provider/mcp/endpoint rollups (row count independent of key cardinality) - and a bounded one for the per-key rollups of the top USAGE_TOP_API_KEYS_LIMIT - keys by spend. breakdown.api_keys and every api_key_breakdown therefore list at - most that many keys, while the totals and the key-free rollups cover every key. + Runs one GROUPING SETS statement with two UNION ALL arms: a key-free one for totals + and the model/provider/mcp/endpoint rollups (row count independent of key + cardinality) and a bounded one for the per-key rollups of the top + USAGE_TOP_API_KEYS_LIMIT keys by spend. breakdown.api_keys and every + api_key_breakdown therefore list at most that many keys, while the totals and the + key-free rollups cover every key. include_entity_breakdown runs a small companion rollup query and folds `breakdown.entities` onto the response, as entity-scoped views like Team Usage need. From c8a2d8c3496ba643aa930b16982382d9a5d76d8c Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 23:00:36 +0000 Subject: [PATCH 046/253] feat(proxy): add LiteLLM_DailyGlobalSpend key-free rollup for the usage dashboard Adds a daily spend table without api_key or user_id, written atomically alongside LiteLLM_DailyUserSpend from the batched writer, reconciled from history by a scheduled job that advances a marker in LiteLLM_Config, and read by the key-free arm of the aggregated usage query once the marker covers the requested range. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../migration.sql | 33 ++ .../litellm_proxy_extras/schema.prisma | 29 ++ litellm/constants.py | 3 + litellm/proxy/db/daily_spend_bulk_upsert.py | 98 +++-- litellm/proxy/db/db_spend_update_writer.py | 7 +- .../common_daily_activity.py | 38 +- litellm/proxy/proxy_server.py | 43 ++ litellm/proxy/schema.prisma | 29 ++ .../daily_global_spend_rollup.py | 235 +++++++++++ schema.prisma | 29 ++ .../proxy/db/test_daily_spend_bulk_upsert.py | 150 +++++++ .../proxy/db/test_db_spend_update_writer.py | 101 ++++- .../test_common_daily_activity.py | 154 ++++++- .../proxy/proxy_server/test_lifecycle.py | 48 +++ .../test_daily_global_spend_rollup.py | 382 ++++++++++++++++++ 15 files changed, 1347 insertions(+), 32 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql create mode 100644 litellm/proxy/spend_tracking/daily_global_spend_rollup.py create mode 100644 tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql new file mode 100644 index 00000000000..1d6cdea0c7b --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql @@ -0,0 +1,33 @@ +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_DailyGlobalSpend" ( + "id" TEXT NOT NULL, + "date" TEXT NOT NULL, + "model" TEXT, + "model_group" TEXT, + "custom_llm_provider" TEXT, + "mcp_namespaced_tool_name" TEXT, + "endpoint" TEXT, + "prompt_tokens" BIGINT NOT NULL DEFAULT 0, + "completion_tokens" BIGINT NOT NULL DEFAULT 0, + "cache_read_input_tokens" BIGINT NOT NULL DEFAULT 0, + "cache_creation_input_tokens" BIGINT NOT NULL DEFAULT 0, + "compression_saved_tokens" BIGINT NOT NULL DEFAULT 0, + "compression_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "prompt_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "gateway_injected_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "autorouter_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "api_requests" BIGINT NOT NULL DEFAULT 0, + "successful_requests" BIGINT NOT NULL DEFAULT 0, + "failed_requests" BIGINT NOT NULL DEFAULT 0, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_DailyGlobalSpend_pkey" PRIMARY KEY ("id") +); + +-- CreateIndex +CREATE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_idx" ON "LiteLLM_DailyGlobalSpend"("date"); + +-- CreateIndex +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_model_model_group_custom_llm__key" ON "LiteLLM_DailyGlobalSpend"("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 8072df5aa5b..5d433e916d6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -809,6 +809,35 @@ model LiteLLM_DailyUserSpend { @@index([endpoint]) } +// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view +model LiteLLM_DailyGlobalSpend { + id String @id @default(uuid()) + date String + model String? + model_group String? + custom_llm_provider String? + mcp_namespaced_tool_name String? + endpoint String? + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + cache_read_input_tokens BigInt @default(0) + cache_creation_input_tokens BigInt @default(0) + compression_saved_tokens BigInt @default(0) + compression_savings_spend Float @default(0.0) + prompt_caching_savings_spend Float @default(0.0) + gateway_injected_caching_savings_spend Float @default(0.0) + autorouter_savings_spend Float @default(0.0) + spend Float @default(0.0) + api_requests BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint]) + @@index([date]) +} + // Track daily organization spend metrics per model and key model LiteLLM_DailyOrganizationSpend { id String @id @default(uuid()) diff --git a/litellm/constants.py b/litellm/constants.py index 565c6433c6e..1bf3a150aeb 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2034,6 +2034,9 @@ PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90 # Deployments named in the lapsed-window alert before it is truncated, so a fleet-wide # expiry cannot produce an alert too large for the channel delivering it. PTU_LAPSED_ALERT_LIMIT: Final[int] = 10 +DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job" +DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600 +DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through" # Slack allowed when deciding a sentinel row is stale. The row's updated_at and the # run's cutoff are stamped by different hosts, so clock skew between them must not let # one run delete a charge another just wrote. A stale row is hours old and a concurrent diff --git a/litellm/proxy/db/daily_spend_bulk_upsert.py b/litellm/proxy/db/daily_spend_bulk_upsert.py index a143643577e..c83043101eb 100644 --- a/litellm/proxy/db/daily_spend_bulk_upsert.py +++ b/litellm/proxy/db/daily_spend_bulk_upsert.py @@ -25,29 +25,41 @@ SpendRow = Mapping[str, object] @dataclass(frozen=True, slots=True) class DailySpendTable: - """The physical table behind one entity's daily rollup.""" + """A daily rollup table and the unique constraint its upserts arbitrate on.""" name: str - entity_id_column: str + key_columns: tuple[str, ...] carries_request_id: bool = False -DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType( - { - "user": DailySpendTable(name="LiteLLM_DailyUserSpend", entity_id_column="user_id"), - "team": DailySpendTable(name="LiteLLM_DailyTeamSpend", entity_id_column="team_id"), - "org": DailySpendTable(name="LiteLLM_DailyOrganizationSpend", entity_id_column="organization_id"), - "end_user": DailySpendTable(name="LiteLLM_DailyEndUserSpend", entity_id_column="end_user_id"), - "agent": DailySpendTable(name="LiteLLM_DailyAgentSpend", entity_id_column="agent_id"), - "tag": DailySpendTable(name="LiteLLM_DailyTagSpend", entity_id_column="tag", carries_request_id=True), - } -) - # The unique constraint's columns after the entity id, in constraint order. A NULL can # never match itself in a unique index, so every one of these is normalized to '': the # conflict target has to be NULL-free or the row is re-inserted on every single flush. _KEY_COLUMNS: Final = ("date", "api_key", "model", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") + +def _entity_table(name: str, entity_id_column: str, carries_request_id: bool = False) -> DailySpendTable: + return DailySpendTable( + name=name, key_columns=(entity_id_column, *_KEY_COLUMNS), carries_request_id=carries_request_id + ) + + +DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType( + { + "user": _entity_table("LiteLLM_DailyUserSpend", "user_id"), + "team": _entity_table("LiteLLM_DailyTeamSpend", "team_id"), + "org": _entity_table("LiteLLM_DailyOrganizationSpend", "organization_id"), + "end_user": _entity_table("LiteLLM_DailyEndUserSpend", "end_user_id"), + "agent": _entity_table("LiteLLM_DailyAgentSpend", "agent_id"), + "tag": _entity_table("LiteLLM_DailyTagSpend", "tag", carries_request_id=True), + } +) + +GLOBAL_SPEND_TABLE: Final = DailySpendTable( + name="LiteLLM_DailyGlobalSpend", + key_columns=("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint"), +) + _COUNTER_COLUMNS: Final = ( "prompt_tokens", "completion_tokens", @@ -92,7 +104,7 @@ def _as_float(value: object) -> float: def conflict_key(table: DailySpendTable, transaction: SpendRow) -> tuple[str, ...]: """The tuple the database arbitrates the upsert on, normalized free of NULLs.""" - return tuple(_as_text(transaction.get(column)) for column in (table.entity_id_column, *_KEY_COLUMNS)) + return tuple(_as_text(transaction.get(column)) for column in table.key_columns) def _merge(group: Sequence[SpendRow]) -> SpendRow: @@ -130,7 +142,11 @@ def _row_params( return ( str(uuid.uuid4()), *key, - None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")), + *( + () + if "model_group" in table.key_columns + else (None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")),) + ), *(_as_int(transaction.get(column)) for column in _COUNTER_COLUMNS), *(_as_float(transaction.get(column)) for column in _SPEND_COLUMNS), *((None if request_id is None else _as_text(request_id),) if table.carries_request_id else ()), @@ -140,26 +156,25 @@ def _row_params( def _insert_columns(table: DailySpendTable) -> tuple[str, ...]: return ( "id", - table.entity_id_column, - *_KEY_COLUMNS, - "model_group", + *table.key_columns, + *(() if "model_group" in table.key_columns else ("model_group",)), *_COUNTER_COLUMNS, *_SPEND_COLUMNS, *(("request_id",) if table.carries_request_id else ()), ) -def build_bulk_upsert( +def _upsert_statement( table: DailySpendTable, batch: Sequence[tuple[tuple[str, ...], SpendRow]], -) -> tuple[str, tuple[SqlValue, ...]]: - """The single statement writing one merged batch, plus its positional arguments.""" + first_param: int, +) -> str: columns: Final = _insert_columns(table) quoted_table: Final = f'"{table.name}"' rows: Final = ", ".join( "(" + ", ".join( - f"${row_index * len(columns) + offset + 1}::{_CASTS.get(column, 'text')}" + f"${first_param + row_index * len(columns) + offset}::{_CASTS.get(column, 'text')}" for offset, column in enumerate(columns) ) + ", (NOW() AT TIME ZONE 'UTC'))" @@ -176,11 +191,44 @@ def build_bulk_upsert( if table.carries_request_id else "" ) - sql: Final = ( + return ( f'INSERT INTO {quoted_table} ({_quoted(columns)}, "updated_at")\n' f"VALUES {rows}\n" - f"ON CONFLICT ({_quoted((table.entity_id_column, *_KEY_COLUMNS))}) DO UPDATE SET\n" + f"ON CONFLICT ({_quoted(table.key_columns)}) DO UPDATE SET\n" f" {increments}{request_id_update},\n" f" \"updated_at\" = (NOW() AT TIME ZONE 'UTC')" ) - return sql, tuple(value for key, transaction in batch for value in _row_params(table, key, transaction)) + + +def _params(table: DailySpendTable, batch: Sequence[tuple[tuple[str, ...], SpendRow]]) -> tuple[SqlValue, ...]: + return tuple(value for key, transaction in batch for value in _row_params(table, key, transaction)) + + +def build_bulk_upsert( + table: DailySpendTable, + batch: Sequence[tuple[tuple[str, ...], SpendRow]], +) -> tuple[str, tuple[SqlValue, ...]]: + """The single statement writing one merged batch, plus its positional arguments.""" + return _upsert_statement(table, batch, first_param=1), _params(table, batch) + + +def build_bulk_upsert_with_global_rollup( + table: DailySpendTable, + batch: Sequence[tuple[tuple[str, ...], SpendRow]], +) -> tuple[str, tuple[SqlValue, ...]]: + """One statement writing a batch to its table and, atomically, its key-free rollup + to ``LiteLLM_DailyGlobalSpend``. + + A data-modifying CTE runs both inserts in the same snapshot and transaction, so a + batch that lands in one table lands in both and a retried deadlock replays both. + Postgres does not order the CTE against the main statement, so two writers can still + deadlock across the tables; the caller's deadlock retry covers that, and each insert + takes its own rows in key order so same-table lock order stays deterministic. + """ + global_batch: Final = merge_by_conflict_key(GLOBAL_SPEND_TABLE, tuple(row for _, row in batch)) + entity_params: Final = _params(table, batch) + sql: Final = ( + f"WITH entity_rows AS (\n{_upsert_statement(table, batch, first_param=1)}\nRETURNING 1)\n" + f"{_upsert_statement(GLOBAL_SPEND_TABLE, global_batch, first_param=len(entity_params) + 1)}" + ) + return sql, (*entity_params, *_params(GLOBAL_SPEND_TABLE, global_batch)) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index eaa03c5d7f7..d5c839a9be8 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -46,6 +46,7 @@ from litellm.proxy._types import ( from litellm.proxy.db.daily_spend_bulk_upsert import ( DAILY_SPEND_TABLES, build_bulk_upsert, + build_bulk_upsert_with_global_rollup, merge_by_conflict_key, ) from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( @@ -1939,7 +1940,11 @@ class DBSpendUpdateWriter: merged_batch = merge_by_conflict_key( table=table, transactions=tuple(transactions_to_process.values()) ) - sql, params = build_bulk_upsert(table=table, batch=merged_batch) + sql, params = ( + build_bulk_upsert_with_global_rollup(table=table, batch=merged_batch) + if entity_type == "user" + else build_bulk_upsert(table=table, batch=merged_batch) + ) await prisma_client.db.execute_raw(sql, *params) except Exception as batch_error: # Log detailed error information for debugging batch upsert failures diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 8a3ba196ab2..f1d78dca201 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -11,6 +11,8 @@ from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT from litellm.proxy._types import CommonProxyErrors +from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE +from litellm.proxy.spend_tracking.daily_global_spend_rollup import reconciled_through from litellm.proxy.spend_tracking.key_metadata_recovery import ( attach_user_emails, recover_double_hashed_key_metadata, @@ -734,6 +736,30 @@ def _rollup_metric_select(table_name: str) -> str: _MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)" +async def key_free_source_table(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None: + """The table the key-free arm reads from, when the global rollup can answer instead of the per-key table. + + Only an unfiltered read of the user table has the same rows as ``LiteLLM_DailyGlobalSpend``, + and only through the day the reconcile marker has reached: the writer keeps that day + current, later days are covered once the next run advances the marker. + """ + if query["table_name"] != "litellm_dailyuserspend": + return None + if query["entity_id"] is not None or query["api_key"] is not None or query["exclude_entity_ids"]: + return None + _, adjusted_end = _adjust_dates_for_timezone( + query["start_date"], query["end_date"], query["timezone_offset_minutes"], query["include_current_utc_day"] + ) + try: + marker: Final = await reconciled_through(prisma_client) + except Exception as exc: # noqa: BLE001 # the per-key table is always a correct answer, so never fail the read + verbose_proxy_logger.warning("Could not read the daily global spend marker, using the per-key table: %s", exc) + return None + if marker is None or adjusted_end > marker: + return None + return GLOBAL_SPEND_TABLE.name + + def _build_aggregated_sql_query( *, table_name: str, @@ -746,13 +772,16 @@ def _build_aggregated_sql_query( exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path timezone_offset_minutes: int | None = None, include_current_utc_day: bool = False, + key_free_table: str | None = None, ) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params """Build the GROUPING SETS query for aggregated daily activity. One statement, two UNION ALL arms over the same WHERE clause. The first arm is key-free: grand total, per-date totals and the (date, model / model_group / provider / mcp / endpoint) rollups, so its row count never grows with the number - of keys. The second arm emits the (date, , api_key) rollups for the + of keys; it reads ``key_free_table`` when given (the global rollup, whose row count + never grew with the number of keys to begin with) and the entity table otherwise. + The second arm emits the (date, , api_key) rollups for the USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint). @@ -778,6 +807,7 @@ def _build_aggregated_sql_query( ) sentinel_param: Final = f"${len(where_params) + 1}" metric_select: Final = _rollup_metric_select(table_name) + key_free_source: Final = key_free_table or pg_table # TODO: drop the successful_requests/failed_requests aggregates (and the # total_successful_requests metadata they feed) once the admin UI reads SGR @@ -796,7 +826,7 @@ def _build_aggregated_sql_query( | GROUPING(model, {_MODEL_GROUP_EXPR}, custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level,{metric_select} - FROM "{pg_table}" + FROM "{key_free_source}" WHERE {where_clause} GROUP BY GROUPING SETS ( (date), @@ -1387,7 +1417,9 @@ async def get_daily_activity_aggregated( timezone_offset_minutes=timezone_offset_minutes, include_current_utc_day=include_current_utc_day, ) - sql_query, sql_params = _build_aggregated_sql_query(**query_kwargs) + sql_query, sql_params = _build_aggregated_sql_query( + **query_kwargs, key_free_table=await key_free_source_table(prisma_client, query_kwargs) + ) entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None raw_rows, raw_entity_rows = await asyncio.gather( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f63e088ebf7..ed8f6886734 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -259,6 +259,7 @@ from litellm.constants import ( APSCHEDULER_MISFIRE_GRACE_TIME, APSCHEDULER_REPLACE_EXISTING, CLI_SSO_SESSION_TTL_SECONDS, + DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, DAYS_IN_A_MONTH, DEFAULT_HEALTH_CHECK_INTERVAL, DEFAULT_MODEL_CREATED_AT_TIME, @@ -662,6 +663,9 @@ from litellm.proxy.route_priority import hot_routes_first from litellm.proxy.search_endpoints.endpoints import router as search_router from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start +from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( + run_scheduled_daily_global_spend_reconcile, +) from litellm.proxy.spend_tracking.spend_counter_batch import ( PendingSpendIncrement, active_spend_counter_batch, @@ -9970,6 +9974,12 @@ class ProxyStartupEvent: await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler) + cls._initialize_daily_global_spend_reconcile_job( + scheduler=scheduler, + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma_client, + ) + ### PTU DAILY ROLLUP ### from litellm.proxy.spend_tracking.ptu_feature_flag import ( is_ptu_cost_attribution_enabled, @@ -10311,6 +10321,39 @@ class ProxyStartupEvent: "LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_ENABLED=true to enable)" ) + @classmethod + def _initialize_daily_global_spend_reconcile_job( + cls, + scheduler: AsyncIOScheduler, + proxy_logging_obj: ProxyLogging, + prisma_client: PrismaClient, + ) -> None: + async def alert(message: str) -> None: + await proxy_logging_obj.alerting_handler( + message=message, + level="High", + alert_type=AlertType.failed_tracking_spend, + ) + + async def reconcile() -> None: + await run_scheduled_daily_global_spend_reconcile( + prisma_client, + pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager, + alert=alert, + ) + + scheduler.add_job( + reconcile, + "cron", + hour=0, + minute=30, + timezone="UTC", + id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2), + ) + @classmethod async def _initialize_slack_alerting_jobs( cls, diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 8072df5aa5b..5d433e916d6 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -809,6 +809,35 @@ model LiteLLM_DailyUserSpend { @@index([endpoint]) } +// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view +model LiteLLM_DailyGlobalSpend { + id String @id @default(uuid()) + date String + model String? + model_group String? + custom_llm_provider String? + mcp_namespaced_tool_name String? + endpoint String? + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + cache_read_input_tokens BigInt @default(0) + cache_creation_input_tokens BigInt @default(0) + compression_saved_tokens BigInt @default(0) + compression_savings_spend Float @default(0.0) + prompt_caching_savings_spend Float @default(0.0) + gateway_injected_caching_savings_spend Float @default(0.0) + autorouter_savings_spend Float @default(0.0) + spend Float @default(0.0) + api_requests BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint]) + @@index([date]) +} + // Track daily organization spend metrics per model and key model LiteLLM_DailyOrganizationSpend { id String @id @default(uuid()) diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py new file mode 100644 index 00000000000..9d344421332 --- /dev/null +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -0,0 +1,235 @@ +"""Reconcile ``LiteLLM_DailyGlobalSpend`` from ``LiteLLM_DailyUserSpend``, one day per transaction. + +The spend writer keeps both tables in step from the moment it is deployed; this job rolls up +the days before that and records how far it has reached in ``LiteLLM_Config`` so usage reads +know when the global table can answer for a date range. It runs as a background cron, never +in a Prisma migration, since on a large deployment the aggregate is minutes of work. +""" + +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from datetime import date, datetime, timedelta, timezone +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel, ConfigDict, ValidationError + +from litellm._logging import verbose_proxy_logger +from litellm.constants import ( + DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, + DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS, + DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, +) +from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE +from litellm.repositories.config_repository import ConfigRepository + +if TYPE_CHECKING: + from litellm.caching.redis_cache import RedisCache + from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager + from litellm.proxy.utils import PrismaClient + +_DAY_TRANSACTION_TIMEOUT: Final = timedelta(minutes=10) +_REPLAY_DAYS: Final = 1 +_METRIC_COLUMNS: Final = ( + "prompt_tokens", + "completion_tokens", + "cache_read_input_tokens", + "cache_creation_input_tokens", + "compression_saved_tokens", + "api_requests", + "successful_requests", + "failed_requests", + "compression_savings_spend", + "prompt_caching_savings_spend", + "gateway_injected_caching_savings_spend", + "autorouter_savings_spend", + "spend", +) + + +def _quoted(columns: tuple[str, ...]) -> str: + return ", ".join(f'"{column}"' for column in columns) + + +def _reconcile_day_sql() -> str: + key_columns: Final = GLOBAL_SPEND_TABLE.key_columns + normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in key_columns) + sums: Final = ", ".join(f'SUM("{column}")' for column in _METRIC_COLUMNS) + overwrite: Final = ", ".join(f'"{column}" = EXCLUDED."{column}"' for column in _METRIC_COLUMNS) + return ( + f'INSERT INTO "{GLOBAL_SPEND_TABLE.name}" ("id", {_quoted(key_columns)}, {_quoted(_METRIC_COLUMNS)}, ' + '"updated_at")\n' + f"SELECT gen_random_uuid()::text, {normalized_keys}, {sums}, (NOW() AT TIME ZONE 'UTC')\n" + 'FROM "LiteLLM_DailyUserSpend" WHERE "date" = $1\n' + f"GROUP BY {normalized_keys}\n" + f"ON CONFLICT ({_quoted(key_columns)}) DO UPDATE SET {overwrite}, " + "\"updated_at\" = (NOW() AT TIME ZONE 'UTC')" + ) + + +RECONCILE_DAY_SQL: Final = _reconcile_day_sql() +_LOCK_GLOBAL_TABLE_SQL: Final = f'LOCK TABLE "{GLOBAL_SPEND_TABLE.name}" IN EXCLUSIVE MODE' +_PENDING_DAYS_SQL: Final = ( + 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" >= $1 AND "date" <= $2 ORDER BY "date"' +) + + +class ReconciledThrough(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + reconciled_through: str + + +class _MarkerRow(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore", from_attributes=True) + + param_value: object = None + + +class _DateRow(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + date: str + + +@dataclass(frozen=True, slots=True) +class ReconcileResult: + days_reconciled: tuple[str, ...] + reconciled_through: str | None + failed_day: str | None = None + + +def _marker_from_param_value(value: object) -> str | None: + try: + parsed: Final = ( + ReconciledThrough.model_validate_json(value) + if isinstance(value, str) + else ReconciledThrough.model_validate(value) + ) + except ValidationError: + return None + return parsed.reconciled_through + + +async def reconciled_through(prisma_client: "PrismaClient") -> str | None: + """The last UTC day ``LiteLLM_DailyGlobalSpend`` is known to cover, or None before the first run.""" + from litellm.proxy.utils import get_config_param + + row: Final = await get_config_param(prisma_client, DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + return None if row is None else _marker_from_param_value(_MarkerRow.model_validate(row).param_value) + + +async def _record_reconciled_through(prisma_client: "PrismaClient", day: str) -> None: + from litellm.proxy.utils import invalidate_config_param + + await ConfigRepository(prisma_client).set_param( + DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, ReconciledThrough(reconciled_through=day).model_dump_json() + ) + await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + +def _first_pending_day(marker: str | None) -> str: + if marker is None: + return "" + return (date.fromisoformat(marker) - timedelta(days=_REPLAY_DAYS)).isoformat() + + +async def pending_days(prisma_client: "PrismaClient", today: date) -> tuple[str, ...]: + """Every UTC day through today still to roll up, oldest first; the marker day and the one + before it are replayed so rows flushed by a pre-writer pod during a rolling deploy are folded in.""" + marker: Final = await reconciled_through(prisma_client) + rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), today.isoformat()) + return tuple(sorted({*(_DateRow.model_validate(row).date for row in rows), today.isoformat()})) + + +async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None: + """Rewrite one day of the global table from the per-key sums; the table lock keeps the + writer's increments out between the aggregate and the overwrite so none are lost.""" + async with prisma_client.db.tx(timeout=_DAY_TRANSACTION_TIMEOUT) as transaction: + await transaction.execute_raw(_LOCK_GLOBAL_TABLE_SQL) + await transaction.execute_raw(RECONCILE_DAY_SQL, day) + + +async def run_daily_global_spend_reconcile( + prisma_client: "PrismaClient", + today: date | None = None, +) -> ReconcileResult: + """Roll up every pending day, advancing the marker after each; a failing day stops the run + with the marker on the last good day so the next run resumes there.""" + effective_today: Final = today or datetime.now(timezone.utc).date() + days: Final = await pending_days(prisma_client, effective_today) + done: Final = await _reconcile_until_failure(prisma_client, days) + failed: Final = days[len(done)] if len(done) < len(days) else None + marker: Final = done[-1] if done else await reconciled_through(prisma_client) + return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=failed) + + +async def _reconcile_until_failure(prisma_client: "PrismaClient", days: tuple[str, ...]) -> tuple[str, ...]: + for index, day in enumerate(days): + if not await _reconcile_and_record(prisma_client, day): + return days[:index] + return days + + +async def _reconcile_and_record(prisma_client: "PrismaClient", day: str) -> bool: + try: + await reconcile_day(prisma_client, day) + await _record_reconciled_through(prisma_client, day) + except Exception as exc: # noqa: BLE001 # one bad day must not lose the days already done + verbose_proxy_logger.exception("Daily global spend reconcile: day %s failed: %s", day, exc) + return False + return True + + +async def run_scheduled_daily_global_spend_reconcile( + prisma_client: "PrismaClient", + pod_lock_manager: "PodLockManager | None" = None, + alert: Callable[[str], Awaitable[None]] | None = None, + today: date | None = None, +) -> ReconcileResult | None: + """Run the reconcile under a cross-pod lock so one proxy does the work; the lock only saves + effort (each day is an idempotent rewrite), so an unreachable Redis runs unguarded rather than skipping.""" + redis_cache: Final = None if pod_lock_manager is None else pod_lock_manager.redis_cache + if pod_lock_manager is None or redis_cache is None: + return await _run_and_alert(prisma_client, alert=alert, today=today) + + acquired: Final = await pod_lock_manager.acquire_lock( + cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, ttl=DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS + ) + if not acquired and await _lock_is_held(pod_lock_manager, redis_cache): + verbose_proxy_logger.info("Daily global spend reconcile: another pod holds the lock, skipping this run") + return None + try: + return await _run_and_alert(prisma_client, alert=alert, today=today) + finally: + if acquired: + await pod_lock_manager.release_lock(cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID) + + +async def _lock_is_held(pod_lock_manager: "PodLockManager", redis_cache: "RedisCache") -> bool: + try: + lock_key: Final = pod_lock_manager.get_redis_lock_key(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID) + return bool(await redis_cache.async_get_cache(lock_key)) + except Exception as exc: # noqa: BLE001 # an unreadable lock must not skip the run + verbose_proxy_logger.warning("Daily global spend reconcile: could not read the lock: %s", exc) + return False + + +async def _run_and_alert( + prisma_client: "PrismaClient", + *, + alert: Callable[[str], Awaitable[None]] | None, + today: date | None, +) -> ReconcileResult: + result: Final = await run_daily_global_spend_reconcile(prisma_client, today=today) + if result.days_reconciled: + verbose_proxy_logger.info( + "Daily global spend reconcile: rolled up %d day(s), reconciled through %s", + len(result.days_reconciled), + result.reconciled_through, + ) + if result.failed_day is not None and alert is not None: + await alert( + f"Daily global spend reconcile stopped at {result.failed_day}; usage totals keep reading the per-key " + f"table for ranges past {result.reconciled_through or 'the beginning'} until the next run succeeds." + ) + return result diff --git a/schema.prisma b/schema.prisma index 8072df5aa5b..5d433e916d6 100644 --- a/schema.prisma +++ b/schema.prisma @@ -809,6 +809,35 @@ model LiteLLM_DailyUserSpend { @@index([endpoint]) } +// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view +model LiteLLM_DailyGlobalSpend { + id String @id @default(uuid()) + date String + model String? + model_group String? + custom_llm_provider String? + mcp_namespaced_tool_name String? + endpoint String? + prompt_tokens BigInt @default(0) + completion_tokens BigInt @default(0) + cache_read_input_tokens BigInt @default(0) + cache_creation_input_tokens BigInt @default(0) + compression_saved_tokens BigInt @default(0) + compression_savings_spend Float @default(0.0) + prompt_caching_savings_spend Float @default(0.0) + gateway_injected_caching_savings_spend Float @default(0.0) + autorouter_savings_spend Float @default(0.0) + spend Float @default(0.0) + api_requests BigInt @default(0) + successful_requests BigInt @default(0) + failed_requests BigInt @default(0) + created_at DateTime @default(now()) + updated_at DateTime @updatedAt + + @@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint]) + @@index([date]) +} + // Track daily organization spend metrics per model and key model LiteLLM_DailyOrganizationSpend { id String @id @default(uuid()) diff --git a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py index c1efb3e7220..cc443a2cfe5 100644 --- a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py +++ b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py @@ -1,12 +1,19 @@ """Tests for the single-statement daily spend upsert (LIT-5291).""" +import pathlib import re +from typing import Final +import psycopg import pytest +from psycopg.rows import dict_row +from pytest_postgresql import factories from litellm.proxy.db.daily_spend_bulk_upsert import ( DAILY_SPEND_TABLES, + GLOBAL_SPEND_TABLE, build_bulk_upsert, + build_bulk_upsert_with_global_rollup, conflict_key, merge_by_conflict_key, ) @@ -185,3 +192,146 @@ async def test_writer_survives_a_transaction_whose_key_columns_are_null(): _, params = prisma_client.db.statements[0] assert None not in params[:9] assert transactions == {} + + +def user_txn(**overrides): + txn = {**tag_txn(), "user_id": "u-1", **overrides} + del txn["tag"] + del txn["request_id"] + return txn + + +def _bound_rows(insert_sql: str, params: tuple[object, ...]) -> list[dict[str, object]]: + """Each VALUES row of one INSERT as a column -> bound value mapping, consuming params in order.""" + header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", insert_sql) + assert header is not None, insert_sql + columns = [c.strip('"') for c in header.group(1).split(", ") if c != '"updated_at"'] + row_count = insert_sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") + return [dict(zip(columns, params[i * len(columns) : (i + 1) * len(columns)])) for i in range(row_count)] + + +def test_global_rollup_folds_every_key_and_user_into_one_row_per_dimension_tuple(): + """The global table has no api_key or user_id, so a batch spread over many keys and + users must collapse to one row per (date, model, group, provider, mcp, endpoint).""" + batch = merge_by_conflict_key( + USER_TABLE, + tuple(user_txn(user_id=f"u-{i}", api_key=f"sk-{i}", spend=1.0, api_requests=1) for i in range(5)) + + (user_txn(user_id="u-0", api_key="sk-0", model="claude", spend=10.0, api_requests=3),), + ) + + sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) + + entity_insert, global_insert = sql.split("RETURNING 1)") + entity_rows = _bound_rows(entity_insert, params) + global_rows = _bound_rows(global_insert, params[len(entity_rows) * len(entity_rows[0]) :]) + assert len(entity_rows) == 6 + assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in global_insert + assert [(r["model"], r["spend"], r["api_requests"]) for r in global_rows] == [ + ("claude", 10.0, 3), + ("gpt-4o-mini", 5.0, 5), + ] + assert all("api_key" not in r and "user_id" not in r for r in global_rows) + conflict = re.search(r"ON CONFLICT \(([^)]*)\)", global_insert) + assert conflict is not None + assert conflict.group(1) == ", ".join(f'"{c}"' for c in GLOBAL_SPEND_TABLE.key_columns) + + +def test_global_rollup_params_follow_the_entity_params_in_one_placeholder_sequence(): + """Both inserts bind from one flat tuple, so the global arm's placeholders must start + exactly where the entity arm's stop or every value lands one column off.""" + batch = merge_by_conflict_key(USER_TABLE, (user_txn(),)) + + sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) + + placeholders = [int(n) for n in re.findall(r"\$(\d+)::", sql)] + assert placeholders == list(range(1, len(params) + 1)) + + +_bulk_upsert_postgresql_proc: Final = factories.postgresql_proc() +_bulk_upsert_postgresql: Final = factories.postgresql("_bulk_upsert_postgresql_proc") + +_MIGRATIONS_DIR: Final = ( + pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations" +) +_GLOBAL_SPEND_MIGRATION: Final = _MIGRATIONS_DIR / "20260915000000_add_daily_global_spend" / "migration.sql" + +_DAILY_USER_SPEND_DDL: Final = """ + CREATE TABLE "LiteLLM_DailyUserSpend" ( + id TEXT PRIMARY KEY, + user_id TEXT, + date TEXT NOT NULL, + api_key TEXT NOT NULL, + model TEXT, + model_group TEXT, + custom_llm_provider TEXT, + mcp_namespaced_tool_name TEXT, + endpoint TEXT, + prompt_tokens BIGINT DEFAULT 0, + completion_tokens BIGINT DEFAULT 0, + cache_read_input_tokens BIGINT DEFAULT 0, + cache_creation_input_tokens BIGINT DEFAULT 0, + compression_saved_tokens BIGINT DEFAULT 0, + compression_savings_spend DOUBLE PRECISION DEFAULT 0, + prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + autorouter_savings_spend DOUBLE PRECISION DEFAULT 0, + spend DOUBLE PRECISION DEFAULT 0, + api_requests BIGINT DEFAULT 0, + successful_requests BIGINT DEFAULT 0, + failed_requests BIGINT DEFAULT 0, + created_at TIMESTAMP DEFAULT now(), + updated_at TIMESTAMP, + UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint) + ) +""" + + +def _execute_dollar_sql(conn: psycopg.Connection, sql: str, params: tuple[object, ...]) -> None: + converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql) + conn.execute( + converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query + {f"p{i}": v for i, v in enumerate(params, start=1)}, + ) + conn.commit() + + +def test_global_rollup_equals_the_per_key_sums_after_repeated_flushes(_bulk_upsert_postgresql: psycopg.Connection): + """Against real Postgres and the shipped migration: two flushes of a mixed batch leave + the global table exactly equal to the per-key table summed over user and key, with the + NULL and '' spellings of a dimension folded into one row.""" + conn: Final = _bulk_upsert_postgresql + conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal + conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal + conn.commit() + + batch = merge_by_conflict_key( + USER_TABLE, + ( + user_txn(user_id="u-1", api_key="sk-1", spend=1.0, prompt_tokens=10), + user_txn(user_id="u-2", api_key="sk-2", spend=2.0, prompt_tokens=20), + user_txn(user_id="u-1", api_key="sk-3", model=None, custom_llm_provider=None, spend=4.0), + user_txn(user_id="u-3", api_key="sk-4", model="", custom_llm_provider="", spend=8.0), + ), + ) + sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) + _execute_dollar_sql(conn, sql, params) + _execute_dollar_sql(conn, sql, params) + + with conn.cursor(row_factory=dict_row) as cur: + global_rows = cur.execute( + 'SELECT model, spend, prompt_tokens, api_requests FROM "LiteLLM_DailyGlobalSpend" ORDER BY model' + ).fetchall() + per_key = cur.execute( + """ + SELECT COALESCE(model, '') AS model, SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, + SUM(api_requests) AS api_requests + FROM "LiteLLM_DailyUserSpend" GROUP BY COALESCE(model, '') ORDER BY 1 + """ + ).fetchall() + + assert [row["model"] for row in global_rows] == ["", "gpt-4o-mini"] + assert [(r["model"], r["spend"], int(r["prompt_tokens"]), int(r["api_requests"])) for r in global_rows] == [ + (r["model"], float(r["spend"]), int(r["prompt_tokens"]), int(r["api_requests"])) for r in per_key + ] + assert global_rows[0]["spend"] == pytest.approx(24.0) + assert global_rows[1]["spend"] == pytest.approx(6.0) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 5e977712a1e..d8a9013398e 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -254,14 +254,19 @@ class _RecordingPrisma: def _row_values(statement: Statement, column: str) -> list[object]: - """Every row's value for one column, read out of the flat parameter tuple.""" + """Every row's value for one column of the first INSERT, read out of the flat parameter tuple. + + The user-table statement chains a global rollup INSERT after its own, so the row count + comes from the first INSERT's VALUES rather than from the parameter count. + """ sql, params = statement header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", sql) assert header is not None, sql columns = header.group(1).split(", ") stride = len(columns) - 1 # updated_at is inlined, not bound offset = columns.index(f'"{column}"') - return [params[row * stride + offset] for row in range(len(params) // stride)] + rows = sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") + return [params[row * stride + offset] for row in range(rows)] @pytest.mark.asyncio @@ -1463,6 +1468,98 @@ async def test_update_daily_spend_keeps_failed_transactions_for_retry(): assert daily_spend_transactions == expected +def _entity_txn(entity_field: str, entity_id: str, api_key: str) -> dict[str, object]: + txn = _daily_txn() + del txn["user_id"] + return {**txn, entity_field: entity_id, "api_key": api_key} + + +@pytest.mark.asyncio +async def test_user_flush_writes_the_global_rollup_in_the_same_statement(): + """The user flush is the one place per-key spend becomes key-free spend, so a batch spread + over many keys must land in LiteLLM_DailyGlobalSpend as one row in the same statement. + A separate statement would let a crash between the two leave the tables out of sync.""" + prisma_client = _RecordingPrisma() + txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(4)} + + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions=txns, + entity_type="user", + entity_id_field="user_id", + ) + + assert len(prisma_client.db.statements) == 1 + sql, params = prisma_client.db.statements[0] + assert sql.count('INSERT INTO "LiteLLM_DailyUserSpend"') == 1 + assert sql.count('INSERT INTO "LiteLLM_DailyGlobalSpend"') == 1 + assert sql.index('"LiteLLM_DailyUserSpend"') < sql.index('"LiteLLM_DailyGlobalSpend"') + global_insert = sql.split('INSERT INTO "LiteLLM_DailyGlobalSpend"', 1)[1] + assert global_insert.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") == 1 + assert "api_key" not in global_insert + assert params.count(0.4) == 1 + assert txns == {} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("entity_type", "entity_field"), + [ + ("team", "team_id"), + ("org", "organization_id"), + ("tag", "tag"), + ("end_user", "end_user_id"), + ("agent", "agent_id"), + ], +) +async def test_other_entity_flushes_leave_the_global_table_alone(entity_type, entity_field): + """Every entity table sees the same request, so writing the rollup from more than one of + them would count each request once per entity type.""" + prisma_client = _RecordingPrisma() + txn = _entity_txn(entity_field, "e-1", "sk-1") + if entity_type == "tag": + txn["request_id"] = "req-1" + + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions={"k": txn}, + entity_type=entity_type, + entity_id_field=entity_field, + ) + + (sql, _params) = prisma_client.db.statements[0] + assert "LiteLLM_DailyGlobalSpend" not in sql + + +@pytest.mark.asyncio +async def test_a_failed_chained_user_flush_keeps_every_transaction_for_retry(): + def raise_outage(): + raise ValueError("simulated database outage") + + prisma_client = _RecordingPrisma(execute_raw=raise_outage) + txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(3)} + expected = dict(txns) + mock_proxy_logging = MagicMock() + mock_proxy_logging.failure_handler = AsyncMock() + + with pytest.raises(ValueError, match="simulated database outage"): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=mock_proxy_logging, + daily_spend_transactions=txns, + entity_type="user", + entity_id_field="user_id", + ) + + assert txns == expected + assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in prisma_client.db.statements[0][0] + + @pytest.mark.asyncio async def test_commit_key_spend_updates_includes_last_active(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 5ff3f89343b..5c74facae6a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -13,7 +13,13 @@ from pytest_postgresql import factories from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR -from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT +import pathlib + +from litellm.constants import ( + DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, + PTU_SENTINEL_API_KEY, + USAGE_TOP_API_KEYS_LIMIT, +) from litellm.proxy.management_endpoints.common_daily_activity import ( _adjust_dates_for_timezone, _build_aggregated_sql_query, @@ -23,8 +29,11 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( get_api_key_metadata, get_daily_activity, get_daily_activity_aggregated, + key_free_source_table, update_metrics, ) +from litellm.proxy.spend_tracking.daily_global_spend_rollup import RECONCILE_DAY_SQL +from litellm.proxy.utils import evict_config_param from litellm.types.proxy.management_endpoints.common_daily_activity import ( DailySpendMetadata, SpendMetrics, @@ -1618,6 +1627,149 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"} +def _prisma_with_marker(marker: str | None) -> MagicMock: + prisma = MagicMock() + prisma.db = MagicMock() + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + row = None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}') + prisma.get_generic_data = AsyncMock(return_value=row) + return prisma + + +def _unfiltered_user_query(**overrides): + return { + "table_name": "litellm_dailyuserspend", + "entity_id_field": "user_id", + "entity_id": None, + "start_date": "2026-06-01", + "end_date": "2026-06-02", + "model": None, + "api_key": None, + "exclude_entity_ids": None, + "timezone_offset_minutes": None, + "include_current_utc_day": False, + **overrides, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("marker", "overrides", "expected"), + [ + ("2026-06-02", {}, "LiteLLM_DailyGlobalSpend"), + ("2026-06-02", {"model": "gpt-5"}, "LiteLLM_DailyGlobalSpend"), + ("2026-06-01", {}, None), + (None, {}, None), + ("2026-06-02", {"api_key": "sk-1"}, None), + ("2026-06-02", {"api_key": []}, None), + ("2026-06-02", {"entity_id": "u-1"}, None), + ("2026-06-02", {"exclude_entity_ids": ["u-1"]}, None), + ("2026-06-02", {"table_name": "litellm_dailyteamspend", "entity_id_field": "team_id"}, None), + ], +) +async def test_key_free_source_table_routes_only_unfiltered_user_reads_within_the_marker(marker, overrides, expected): + """Anything that filters by key or entity has no counterpart in the global table, and a + range the reconcile has not reached must stay on the per-key table.""" + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + prisma = _prisma_with_marker(marker) + + assert await key_free_source_table(prisma, _unfiltered_user_query(**overrides)) == expected + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + +@pytest.mark.asyncio +async def test_key_free_source_table_judges_the_timezone_extended_end_not_the_requested_one(): + """A caller west of UTC asking through their local today gets today's UTC bucket added to + the range; the marker must cover that extended day, not just the requested end.""" + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + today_utc: Final = datetime.now(timezone.utc).date() + yesterday: Final = (today_utc - timedelta(days=1)).isoformat() + query: Final = _unfiltered_user_query( + start_date=yesterday, end_date=yesterday, timezone_offset_minutes=24 * 60, include_current_utc_day=True + ) + + assert await key_free_source_table(_prisma_with_marker(yesterday), query) is None + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + assert await key_free_source_table(_prisma_with_marker(today_utc.isoformat()), query) == "LiteLLM_DailyGlobalSpend" + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + +_GLOBAL_SPEND_MIGRATION: Final = ( + pathlib.Path(__file__).resolve().parents[4] + / "litellm-proxy-extras" + / "litellm_proxy_extras" + / "migrations" + / "20260915000000_add_daily_global_spend" + / "migration.sql" +) + + +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_free_arm( + _aggregated_postgresql: psycopg.Connection, +): + """With the range reconciled, the key-free arm reads LiteLLM_DailyGlobalSpend while the + per-key arm stays on the user table, and the response is identical to the all-per-key + read: same totals, same rollups, same top keys.""" + n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 3 + rows: Final = [ + ( + f"row-{day}-{i:03d}", + f"user-{i % 7}", + day, + f"key-{i:03d}", + "gpt-5" if i % 2 else "claude", + "" if i % 3 else "gpt-5", + "openai" if i % 2 else None, + "/v1/chat/completions" if i % 5 else None, + 10, + float(i + 1), + 1, + 1, + ) + for day in ("2026-06-01", "2026-06-02") + for i in range(n_keys) + ] + _seed_daily_user_spend(_aggregated_postgresql, rows) + with _aggregated_postgresql.cursor() as cur: + cur.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal + for day in ("2026-06-01", "2026-06-02"): + cur.execute( + re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg + {"p1": day}, + ) + _aggregated_postgresql.commit() + + async def read(marker: str | None, sql_seen: list[str]): + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + prisma = _prisma_with_marker(marker) + run_query = _psycopg_query_raw(_aggregated_postgresql, []) + + async def query_raw(sql: str, *params: str): + sql_seen.append(sql) + return await run_query(sql, *params) + + prisma.db.query_raw = query_raw + return await get_daily_activity_aggregated( + prisma_client=prisma, + entity_metadata_field=None, + **_unfiltered_user_query(), + ) + + per_key_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim + global_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim + from_per_key = await read(None, per_key_sql) + from_global = await read("2026-06-02", global_sql) + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + assert per_key_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 0 + assert global_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 1 + assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 2 + assert from_global.model_dump() == from_per_key.model_dump() + assert from_global.metadata.total_spend == pytest.approx(2 * sum(float(i + 1) for i in range(n_keys))) + assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT + assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"} def _no_spend_record(): """A rollup row for a key with no spend, where SUM() returns NULL (None).""" return SimpleNamespace( diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index deb7289d2d1..ee72e98ffa9 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -1042,6 +1042,54 @@ async def test_spend_report_locks_are_never_released(): proxy_logging_obj.db_spend_update_writer.pod_lock_manager.release_lock.assert_not_awaited() +def _init_daily_global_spend_reconcile_job() -> tuple[MagicMock, MagicMock, MagicMock]: + scheduler = MagicMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.alerting_handler = AsyncMock() + prisma_client = MagicMock() + ProxyStartupEvent._initialize_daily_global_spend_reconcile_job( + scheduler=scheduler, + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma_client, + ) + return scheduler, proxy_logging_obj, prisma_client + + +def test_daily_global_spend_reconcile_job_is_scheduled_nightly_with_an_immediate_catch_up_run(): + """Startup schedules the LiteLLM_DailyGlobalSpend backfill a couple of minutes out, so a + fresh deploy switches usage reads to the global table without waiting for the nightly + run, and replaces any previous registration of the same job id.""" + from datetime import datetime, timedelta, timezone + + from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID + + scheduler, _, _ = _init_daily_global_spend_reconcile_job() + + (call,) = scheduler.add_job.call_args_list + assert call.kwargs["id"] == DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID + assert call.kwargs["replace_existing"] is True + assert call.args[1:] == ("cron",) + assert (call.kwargs["hour"], call.kwargs["minute"], call.kwargs["timezone"]) == (0, 30, "UTC") + assert timedelta(0) < call.kwargs["next_run_time"] - datetime.now(timezone.utc) <= timedelta(minutes=2) + + +@pytest.mark.asyncio +async def test_daily_global_spend_reconcile_job_runs_under_the_pod_lock_and_alerts_through_the_proxy(monkeypatch): + scheduler, proxy_logging_obj, prisma_client = _init_daily_global_spend_reconcile_job() + run = AsyncMock() + monkeypatch.setattr(ps, "run_scheduled_daily_global_spend_reconcile", run) + + await scheduler.add_job.call_args.args[0]() + + run.assert_awaited_once() + assert run.await_args.args == (prisma_client,) + assert run.await_args.kwargs["pod_lock_manager"] is proxy_logging_obj.db_spend_update_writer.pod_lock_manager + await run.await_args.kwargs["alert"]("day 2026-09-01 failed") + proxy_logging_obj.alerting_handler.assert_awaited_once() + assert proxy_logging_obj.alerting_handler.await_args.kwargs["message"] == "day 2026-09-01 failed" + assert proxy_logging_obj.alerting_handler.await_args.kwargs["level"] == "High" + + @pytest.mark.asyncio async def test_prometheus_fallback_stats_job_skipped_when_another_pod_holds_the_lock(monkeypatch): """The boot-time send goes through the same gate, so a losing pod sends nothing at all: diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py new file mode 100644 index 00000000000..13dc757cbbd --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -0,0 +1,382 @@ +"""Tests for the LiteLLM_DailyGlobalSpend reconcile job (LIT-7818).""" + +import pathlib +import re +from contextlib import asynccontextmanager +from datetime import date +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import psycopg +import pytest +from psycopg.rows import dict_row +from pytest_postgresql import factories + +from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM +from litellm.proxy.db.daily_spend_bulk_upsert import ( + DAILY_SPEND_TABLES, + build_bulk_upsert_with_global_rollup, + merge_by_conflict_key, +) +from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( + RECONCILE_DAY_SQL, + reconciled_through, + run_daily_global_spend_reconcile, + run_scheduled_daily_global_spend_reconcile, +) +from litellm.proxy.utils import evict_config_param + +USER_TABLE: Final = DAILY_SPEND_TABLES["user"] +TODAY: Final = date(2026, 9, 15) + + +class _FakeConfigRow: + def __init__(self, param_name: str, param_value: object) -> None: + self.param_name = param_name + self.param_value = param_value + + +class _FakeConfigTable: + def __init__(self) -> None: + self.rows: dict[str, object] = {} + + async def upsert(self, *, where: dict[str, str], data: dict[str, dict[str, str]]) -> _FakeConfigRow: + self.rows[where["param_name"]] = data["update"]["param_value"] + return _FakeConfigRow(where["param_name"], data["update"]["param_value"]) + + +class _FakeTransaction: + def __init__(self, prisma: "_FakePrisma") -> None: + self._prisma = prisma + + async def execute_raw(self, sql: str, *params: str) -> int: + if "LOCK TABLE" in sql: + self._prisma.locks_taken += 1 + return 0 + (day,) = params + if day in self._prisma.failing_days: + raise RuntimeError(f"day {day} exploded") + self._prisma.reconciled.append(day) + return 1 + + +class _FakeDb: + def __init__(self, prisma: "_FakePrisma") -> None: + self._prisma = prisma + self.litellm_config = _FakeConfigTable() + + async def query_raw(self, sql: str, *params: str) -> list[dict[str, str]]: + first, last = params + return [{"date": d} for d in sorted(self._prisma.user_days) if first <= d <= last] + + @asynccontextmanager + async def tx(self, timeout: object): + yield _FakeTransaction(self._prisma) + + +class _FakePrisma: + """Enough of PrismaClient for the reconcile: per-key dates, a config table, and a transaction.""" + + def __init__(self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset()) -> None: + self.user_days = user_days + self.failing_days = failing_days + self.reconciled: list[str] = [] + self.locks_taken = 0 + self.db = _FakeDb(self) + + async def get_generic_data(self, key: str, value: str, table_name: str) -> _FakeConfigRow | None: + stored = self.db.litellm_config.rows.get(value) + return None if stored is None else _FakeConfigRow(value, stored) + + +@pytest.fixture(autouse=True) +async def _fresh_marker_cache(): + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + yield + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + +@pytest.mark.asyncio +async def test_first_run_rolls_up_every_historical_day_and_today_then_marks_today(): + """Before any marker exists, every day with per-key rows is rolled up, plus today even + with no rows yet, so reads for ranges ending today can switch to the global table.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14")) + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15") + assert result.failed_day is None + assert result.reconciled_through == "2026-09-15" + assert await reconciled_through(prisma) == "2026-09-15" + assert prisma.locks_taken == 4 + + +@pytest.mark.asyncio +async def test_later_run_replays_the_marker_day_and_the_day_before_only(): + """Days older than marker-1 are settled; the marker day and its predecessor are replayed so + rows a pre-writer pod flushed around midnight during a rolling deploy get folded in.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14")) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13)) + prisma.reconciled.clear() + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14", "2026-09-15") + assert "2026-09-01" not in prisma.reconciled + assert await reconciled_through(prisma) == "2026-09-15" + + +@pytest.mark.asyncio +async def test_a_failing_day_stops_the_run_and_leaves_the_marker_on_the_last_good_day(): + """The marker may never claim a day that was not rewritten: reads past it would then trust + a global table missing that day's spend.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"})) + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-01",) + assert result.failed_day == "2026-09-02" + assert result.reconciled_through == "2026-09-01" + assert prisma.reconciled == ["2026-09-01"] + assert await reconciled_through(prisma) == "2026-09-01" + + +@pytest.mark.asyncio +async def test_the_next_run_resumes_from_the_failed_day(): + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"})) + await run_daily_global_spend_reconcile(prisma, today=TODAY) + prisma.failing_days = frozenset() + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03", "2026-09-15") + assert await reconciled_through(prisma) == "2026-09-15" + + +@pytest.mark.asyncio +async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts(): + """A pre-writer pod flushing rows for the day before the marker is exactly the replay case; + when that replay fails the marker must stay put and the operator must hear about it.""" + prisma = _FakePrisma(user_days=("2026-09-13",)) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13)) + prisma.user_days = ("2026-09-12", "2026-09-13") + prisma.failing_days = frozenset({"2026-09-12"}) + alert = AsyncMock() + + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert, today=TODAY) + + assert result is not None + assert result.days_reconciled == () + assert result.failed_day == "2026-09-12" + assert result.reconciled_through == "2026-09-13" + alert.assert_awaited_once() + assert "2026-09-12" in alert.await_args.args[0] + + +@pytest.mark.asyncio +async def test_a_clean_run_does_not_alert(): + prisma = _FakePrisma(user_days=("2026-09-13",)) + alert = AsyncMock() + + await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert, today=TODAY) + + alert.assert_not_awaited() + + +def _pod_lock(acquired: bool) -> MagicMock: + lock = MagicMock() + lock.redis_cache = MagicMock() + lock.redis_cache.async_get_cache = AsyncMock(return_value="other-pod") + lock.get_redis_lock_key = MagicMock(return_value="lock-key") + lock.acquire_lock = AsyncMock(return_value=acquired) + lock.release_lock = AsyncMock() + return lock + + +@pytest.mark.asyncio +async def test_scheduled_run_skips_when_another_pod_holds_the_lock(): + prisma = _FakePrisma(user_days=("2026-09-13",)) + lock = _pod_lock(acquired=False) + + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + + assert result is None + assert prisma.reconciled == [] + lock.release_lock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_scheduled_run_runs_and_releases_the_lock_when_it_wins(): + prisma = _FakePrisma(user_days=("2026-09-13",)) + lock = _pod_lock(acquired=True) + + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + + assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15") + lock.release_lock.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_scheduled_run_proceeds_when_the_lock_cannot_be_acquired_or_read(): + """A Redis outage must not stall the backfill: the day rewrite is idempotent, so running + twice is only wasted effort while skipping forever leaves usage on the slow path.""" + prisma = _FakePrisma(user_days=("2026-09-13",)) + lock = _pod_lock(acquired=False) + lock.redis_cache.async_get_cache = AsyncMock(side_effect=ConnectionError("redis down")) + + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + + assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15") + lock.release_lock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_marker_is_read_back_from_the_json_string_the_config_table_stores(): + prisma = _FakePrisma(user_days=()) + prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"reconciled_through": "2026-09-10"}' + + assert await reconciled_through(prisma) == "2026-09-10" + + +@pytest.mark.asyncio +async def test_an_unparseable_marker_reads_as_never_reconciled(): + prisma = _FakePrisma(user_days=()) + prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"something_else": 1}' + + assert await reconciled_through(prisma) is None + + +_rollup_postgresql_proc: Final = factories.postgresql_proc() +_rollup_postgresql: Final = factories.postgresql("_rollup_postgresql_proc") + +_MIGRATIONS_DIR: Final = ( + pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations" +) +_GLOBAL_SPEND_MIGRATION: Final = _MIGRATIONS_DIR / "20260915000000_add_daily_global_spend" / "migration.sql" + +_DAILY_USER_SPEND_DDL: Final = """ + CREATE TABLE "LiteLLM_DailyUserSpend" ( + id TEXT PRIMARY KEY, + user_id TEXT, + date TEXT NOT NULL, + api_key TEXT NOT NULL, + model TEXT, + model_group TEXT, + custom_llm_provider TEXT, + mcp_namespaced_tool_name TEXT, + endpoint TEXT, + prompt_tokens BIGINT DEFAULT 0, + completion_tokens BIGINT DEFAULT 0, + cache_read_input_tokens BIGINT DEFAULT 0, + cache_creation_input_tokens BIGINT DEFAULT 0, + compression_saved_tokens BIGINT DEFAULT 0, + compression_savings_spend DOUBLE PRECISION DEFAULT 0, + prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0, + autorouter_savings_spend DOUBLE PRECISION DEFAULT 0, + spend DOUBLE PRECISION DEFAULT 0, + api_requests BIGINT DEFAULT 0, + successful_requests BIGINT DEFAULT 0, + failed_requests BIGINT DEFAULT 0, + created_at TIMESTAMP DEFAULT now(), + updated_at TIMESTAMP, + UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint) + ) +""" + +_PER_KEY_SUMS_SQL: Final = """ + SELECT COALESCE(model, '') AS model, COALESCE(model_group, '') AS model_group, + COALESCE(custom_llm_provider, '') AS custom_llm_provider, + SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, SUM(api_requests) AS api_requests + FROM "LiteLLM_DailyUserSpend" WHERE date = %s + GROUP BY 1, 2, 3 ORDER BY 1, 2, 3 +""" +_GLOBAL_ROWS_SQL: Final = """ + SELECT model, model_group, custom_llm_provider, spend, prompt_tokens, api_requests + FROM "LiteLLM_DailyGlobalSpend" WHERE date = %s ORDER BY 1, 2, 3 +""" + + +def _execute_dollar_sql(conn: psycopg.Connection, sql: str, params: tuple[object, ...]) -> None: + converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql) + conn.execute( + converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query + {f"p{i}": v for i, v in enumerate(params, start=1)}, + ) + conn.commit() + + +def _user_txn(**overrides): + return { + "user_id": "u-1", + "date": "2026-09-14", + "api_key": "sk-1", + "model": "gpt-5", + "model_group": "gpt-5", + "custom_llm_provider": "openai", + "mcp_namespaced_tool_name": "", + "endpoint": "/chat/completions", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 1.0, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + **overrides, + } + + +def _normalized(rows: list[dict[str, object]]) -> list[tuple[object, ...]]: + return [ + ( + r["model"], + r["model_group"], + r["custom_llm_provider"], + float(r["spend"]), + int(r["prompt_tokens"]), + int(r["api_requests"]), + ) # pyright: ignore[reportArgumentType] # dict_row values are untyped + for r in rows + ] + + +def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_postgresql: psycopg.Connection): + """Against real Postgres and the shipped migration: rows the writer never saw (a + pre-writer pod's flush, NULL and '' dimension spellings) end up folded into the global + day, running the day twice changes nothing, and other days are left alone.""" + conn: Final = _rollup_postgresql + conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal + conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal + conn.commit() + + written_batch = merge_by_conflict_key( + USER_TABLE, + (_user_txn(api_key="sk-1", spend=1.0), _user_txn(api_key="sk-2", user_id="u-2", spend=2.0, prompt_tokens=20)), + ) + _execute_dollar_sql(conn, *build_bulk_upsert_with_global_rollup(USER_TABLE, written_batch)) + + conn.execute( + """ + INSERT INTO "LiteLLM_DailyUserSpend" + (id, user_id, date, api_key, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, + endpoint, prompt_tokens, spend, api_requests) + VALUES + ('legacy-1', 'u-9', '2026-09-14', 'sk-9', 'gpt-5', NULL, 'openai', NULL, NULL, 5, 4.0, 1), + ('legacy-2', 'u-9', '2026-09-14', 'sk-9', 'gpt-5', '', 'openai', '', '', 5, 8.0, 1), + ('legacy-3', 'u-9', '2026-09-13', 'sk-9', 'claude', '', 'anthropic', '', '', 7, 16.0, 1) + """ + ) + conn.commit() + + _execute_dollar_sql(conn, RECONCILE_DAY_SQL, ("2026-09-14",)) + _execute_dollar_sql(conn, RECONCILE_DAY_SQL, ("2026-09-14",)) + + with conn.cursor(row_factory=dict_row) as cur: + global_rows = cur.execute(_GLOBAL_ROWS_SQL, ("2026-09-14",)).fetchall() + per_key = cur.execute(_PER_KEY_SUMS_SQL, ("2026-09-14",)).fetchall() + untouched = cur.execute(_GLOBAL_ROWS_SQL, ("2026-09-13",)).fetchall() + + assert _normalized(global_rows) == _normalized(per_key) + assert sum(float(r["spend"]) for r in global_rows) == pytest.approx(15.0) # pyright: ignore[reportArgumentType] # dict_row values are untyped + assert [(r["model"], r["model_group"]) for r in global_rows] == [("gpt-5", ""), ("gpt-5", "gpt-5")] + assert untouched == [] From ad8de0e1927c18d5d14c92939bbd531c54573874 Mon Sep 17 00:00:00 2001 From: yassin Date: Tue, 15 Sep 2026 23:42:19 +0000 Subject: [PATCH 047/253] fix(proxy): roll up only closed days into LiteLLM_DailyGlobalSpend and split the key-free read at the marker The write path no longer dual-writes the global table. The cron rolls up closed UTC days only, so a pod still flushing the current day can never leave the global table short. The key-free arm reads days through the marker from the global table and later days from LiteLLM_DailyUserSpend in one UNION ALL, and the marker comes from the config cache rather than a per-request database lookup. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/daily_spend_bulk_upsert.py | 98 +++--------- litellm/proxy/db/db_spend_update_writer.py | 7 +- .../common_daily_activity.py | 79 ++++++--- .../daily_global_spend_rollup.py | 43 ++--- .../proxy/db/test_daily_spend_bulk_upsert.py | 150 ------------------ .../proxy/db/test_db_spend_update_writer.py | 101 +----------- .../test_common_daily_activity.py | 90 ++++++----- .../test_daily_global_spend_rollup.py | 92 +++++------ 8 files changed, 204 insertions(+), 456 deletions(-) diff --git a/litellm/proxy/db/daily_spend_bulk_upsert.py b/litellm/proxy/db/daily_spend_bulk_upsert.py index c83043101eb..a143643577e 100644 --- a/litellm/proxy/db/daily_spend_bulk_upsert.py +++ b/litellm/proxy/db/daily_spend_bulk_upsert.py @@ -25,41 +25,29 @@ SpendRow = Mapping[str, object] @dataclass(frozen=True, slots=True) class DailySpendTable: - """A daily rollup table and the unique constraint its upserts arbitrate on.""" + """The physical table behind one entity's daily rollup.""" name: str - key_columns: tuple[str, ...] + entity_id_column: str carries_request_id: bool = False +DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType( + { + "user": DailySpendTable(name="LiteLLM_DailyUserSpend", entity_id_column="user_id"), + "team": DailySpendTable(name="LiteLLM_DailyTeamSpend", entity_id_column="team_id"), + "org": DailySpendTable(name="LiteLLM_DailyOrganizationSpend", entity_id_column="organization_id"), + "end_user": DailySpendTable(name="LiteLLM_DailyEndUserSpend", entity_id_column="end_user_id"), + "agent": DailySpendTable(name="LiteLLM_DailyAgentSpend", entity_id_column="agent_id"), + "tag": DailySpendTable(name="LiteLLM_DailyTagSpend", entity_id_column="tag", carries_request_id=True), + } +) + # The unique constraint's columns after the entity id, in constraint order. A NULL can # never match itself in a unique index, so every one of these is normalized to '': the # conflict target has to be NULL-free or the row is re-inserted on every single flush. _KEY_COLUMNS: Final = ("date", "api_key", "model", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") - -def _entity_table(name: str, entity_id_column: str, carries_request_id: bool = False) -> DailySpendTable: - return DailySpendTable( - name=name, key_columns=(entity_id_column, *_KEY_COLUMNS), carries_request_id=carries_request_id - ) - - -DAILY_SPEND_TABLES: Final[Mapping[DailySpendEntity, DailySpendTable]] = MappingProxyType( - { - "user": _entity_table("LiteLLM_DailyUserSpend", "user_id"), - "team": _entity_table("LiteLLM_DailyTeamSpend", "team_id"), - "org": _entity_table("LiteLLM_DailyOrganizationSpend", "organization_id"), - "end_user": _entity_table("LiteLLM_DailyEndUserSpend", "end_user_id"), - "agent": _entity_table("LiteLLM_DailyAgentSpend", "agent_id"), - "tag": _entity_table("LiteLLM_DailyTagSpend", "tag", carries_request_id=True), - } -) - -GLOBAL_SPEND_TABLE: Final = DailySpendTable( - name="LiteLLM_DailyGlobalSpend", - key_columns=("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint"), -) - _COUNTER_COLUMNS: Final = ( "prompt_tokens", "completion_tokens", @@ -104,7 +92,7 @@ def _as_float(value: object) -> float: def conflict_key(table: DailySpendTable, transaction: SpendRow) -> tuple[str, ...]: """The tuple the database arbitrates the upsert on, normalized free of NULLs.""" - return tuple(_as_text(transaction.get(column)) for column in table.key_columns) + return tuple(_as_text(transaction.get(column)) for column in (table.entity_id_column, *_KEY_COLUMNS)) def _merge(group: Sequence[SpendRow]) -> SpendRow: @@ -142,11 +130,7 @@ def _row_params( return ( str(uuid.uuid4()), *key, - *( - () - if "model_group" in table.key_columns - else (None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")),) - ), + None if transaction.get("model_group") is None else _as_text(transaction.get("model_group")), *(_as_int(transaction.get(column)) for column in _COUNTER_COLUMNS), *(_as_float(transaction.get(column)) for column in _SPEND_COLUMNS), *((None if request_id is None else _as_text(request_id),) if table.carries_request_id else ()), @@ -156,25 +140,26 @@ def _row_params( def _insert_columns(table: DailySpendTable) -> tuple[str, ...]: return ( "id", - *table.key_columns, - *(() if "model_group" in table.key_columns else ("model_group",)), + table.entity_id_column, + *_KEY_COLUMNS, + "model_group", *_COUNTER_COLUMNS, *_SPEND_COLUMNS, *(("request_id",) if table.carries_request_id else ()), ) -def _upsert_statement( +def build_bulk_upsert( table: DailySpendTable, batch: Sequence[tuple[tuple[str, ...], SpendRow]], - first_param: int, -) -> str: +) -> tuple[str, tuple[SqlValue, ...]]: + """The single statement writing one merged batch, plus its positional arguments.""" columns: Final = _insert_columns(table) quoted_table: Final = f'"{table.name}"' rows: Final = ", ".join( "(" + ", ".join( - f"${first_param + row_index * len(columns) + offset}::{_CASTS.get(column, 'text')}" + f"${row_index * len(columns) + offset + 1}::{_CASTS.get(column, 'text')}" for offset, column in enumerate(columns) ) + ", (NOW() AT TIME ZONE 'UTC'))" @@ -191,44 +176,11 @@ def _upsert_statement( if table.carries_request_id else "" ) - return ( + sql: Final = ( f'INSERT INTO {quoted_table} ({_quoted(columns)}, "updated_at")\n' f"VALUES {rows}\n" - f"ON CONFLICT ({_quoted(table.key_columns)}) DO UPDATE SET\n" + f"ON CONFLICT ({_quoted((table.entity_id_column, *_KEY_COLUMNS))}) DO UPDATE SET\n" f" {increments}{request_id_update},\n" f" \"updated_at\" = (NOW() AT TIME ZONE 'UTC')" ) - - -def _params(table: DailySpendTable, batch: Sequence[tuple[tuple[str, ...], SpendRow]]) -> tuple[SqlValue, ...]: - return tuple(value for key, transaction in batch for value in _row_params(table, key, transaction)) - - -def build_bulk_upsert( - table: DailySpendTable, - batch: Sequence[tuple[tuple[str, ...], SpendRow]], -) -> tuple[str, tuple[SqlValue, ...]]: - """The single statement writing one merged batch, plus its positional arguments.""" - return _upsert_statement(table, batch, first_param=1), _params(table, batch) - - -def build_bulk_upsert_with_global_rollup( - table: DailySpendTable, - batch: Sequence[tuple[tuple[str, ...], SpendRow]], -) -> tuple[str, tuple[SqlValue, ...]]: - """One statement writing a batch to its table and, atomically, its key-free rollup - to ``LiteLLM_DailyGlobalSpend``. - - A data-modifying CTE runs both inserts in the same snapshot and transaction, so a - batch that lands in one table lands in both and a retried deadlock replays both. - Postgres does not order the CTE against the main statement, so two writers can still - deadlock across the tables; the caller's deadlock retry covers that, and each insert - takes its own rows in key order so same-table lock order stays deterministic. - """ - global_batch: Final = merge_by_conflict_key(GLOBAL_SPEND_TABLE, tuple(row for _, row in batch)) - entity_params: Final = _params(table, batch) - sql: Final = ( - f"WITH entity_rows AS (\n{_upsert_statement(table, batch, first_param=1)}\nRETURNING 1)\n" - f"{_upsert_statement(GLOBAL_SPEND_TABLE, global_batch, first_param=len(entity_params) + 1)}" - ) - return sql, (*entity_params, *_params(GLOBAL_SPEND_TABLE, global_batch)) + return sql, tuple(value for key, transaction in batch for value in _row_params(table, key, transaction)) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index d5c839a9be8..eaa03c5d7f7 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -46,7 +46,6 @@ from litellm.proxy._types import ( from litellm.proxy.db.daily_spend_bulk_upsert import ( DAILY_SPEND_TABLES, build_bulk_upsert, - build_bulk_upsert_with_global_rollup, merge_by_conflict_key, ) from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( @@ -1940,11 +1939,7 @@ class DBSpendUpdateWriter: merged_batch = merge_by_conflict_key( table=table, transactions=tuple(transactions_to_process.values()) ) - sql, params = ( - build_bulk_upsert_with_global_rollup(table=table, batch=merged_batch) - if entity_type == "user" - else build_bulk_upsert(table=table, batch=merged_batch) - ) + sql, params = build_bulk_upsert(table=table, batch=merged_batch) await prisma_client.db.execute_raw(sql, *params) except Exception as batch_error: # Log detailed error information for debugging batch upsert failures diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index f1d78dca201..b90f874c04c 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -11,8 +11,7 @@ from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT from litellm.proxy._types import CommonProxyErrors -from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE -from litellm.proxy.spend_tracking.daily_global_spend_rollup import reconciled_through +from litellm.proxy.spend_tracking.daily_global_spend_rollup import GLOBAL_SPEND_TABLE_NAME, reconciled_through from litellm.proxy.spend_tracking.key_metadata_recovery import ( attach_user_emails, recover_double_hashed_key_metadata, @@ -736,28 +735,62 @@ def _rollup_metric_select(table_name: str) -> str: _MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)" -async def key_free_source_table(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None: - """The table the key-free arm reads from, when the global rollup can answer instead of the per-key table. +_KEY_FREE_SOURCE_COLUMNS: Final = ( + "date", + "model", + "model_group", + "custom_llm_provider", + "mcp_namespaced_tool_name", + "endpoint", + "spend", + "prompt_tokens", + "completion_tokens", + "cache_read_input_tokens", + "cache_creation_input_tokens", + "compression_saved_tokens", + "compression_savings_spend", + "prompt_caching_savings_spend", + "gateway_injected_caching_savings_spend", + "autorouter_savings_spend", + "api_requests", + "successful_requests", + "failed_requests", +) - Only an unfiltered read of the user table has the same rows as ``LiteLLM_DailyGlobalSpend``, - and only through the day the reconcile marker has reached: the writer keeps that day - current, later days are covered once the next run advances the marker. + +async def global_rollup_reconciled_through(prisma_client: PrismaClient, query: _AggregatedQueryKwargs) -> str | None: + """The last day ``LiteLLM_DailyGlobalSpend`` can answer the key-free arm for, or None to + read it all from the per-key table. + + Only an unfiltered read of the user table sums to the same rows as the global table. The + marker read is served from the config cache, so this is not a database round trip per request. """ if query["table_name"] != "litellm_dailyuserspend": return None if query["entity_id"] is not None or query["api_key"] is not None or query["exclude_entity_ids"]: return None - _, adjusted_end = _adjust_dates_for_timezone( - query["start_date"], query["end_date"], query["timezone_offset_minutes"], query["include_current_utc_day"] - ) try: - marker: Final = await reconciled_through(prisma_client) + return await reconciled_through(prisma_client) except Exception as exc: # noqa: BLE001 # the per-key table is always a correct answer, so never fail the read verbose_proxy_logger.warning("Could not read the daily global spend marker, using the per-key table: %s", exc) return None - if marker is None or adjusted_end > marker: - return None - return GLOBAL_SPEND_TABLE.name + + +def _key_free_source(pg_table: str, where_clause: str, marker_param: str | None) -> str: + """The relation the key-free arm aggregates: the per-key table alone, or the global rollup + for days through the marker plus the per-key table for the days still open after it.""" + if marker_param is None: + return f'"{pg_table}"\n WHERE {where_clause}' + columns: Final = ", ".join(_KEY_FREE_SOURCE_COLUMNS) + return f"""( + SELECT {columns} + FROM "{GLOBAL_SPEND_TABLE_NAME}" + WHERE {where_clause} AND date <= {marker_param} + UNION ALL + SELECT {columns} + FROM "{pg_table}" + WHERE {where_clause} AND date > {marker_param} + ) AS key_free_source""" def _build_aggregated_sql_query( @@ -772,15 +805,16 @@ def _build_aggregated_sql_query( exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path timezone_offset_minutes: int | None = None, include_current_utc_day: bool = False, - key_free_table: str | None = None, + global_rollup_through: str | None = None, ) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params """Build the GROUPING SETS query for aggregated daily activity. One statement, two UNION ALL arms over the same WHERE clause. The first arm is key-free: grand total, per-date totals and the (date, model / model_group / provider / mcp / endpoint) rollups, so its row count never grows with the number - of keys; it reads ``key_free_table`` when given (the global rollup, whose row count - never grew with the number of keys to begin with) and the entity table otherwise. + of keys. With ``global_rollup_through`` it reads days through that marker from + ``LiteLLM_DailyGlobalSpend`` (whose row count never grew with the number of keys to + begin with) and only the days after it from the per-key table. The second arm emits the (date, , api_key) rollups for the USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint). @@ -806,8 +840,8 @@ def _build_aggregated_sql_query( exclude_entity_ids=exclude_entity_ids, ) sentinel_param: Final = f"${len(where_params) + 1}" + marker_param: Final = None if global_rollup_through is None else f"${len(where_params) + 2}" metric_select: Final = _rollup_metric_select(table_name) - key_free_source: Final = key_free_table or pg_table # TODO: drop the successful_requests/failed_requests aggregates (and the # total_successful_requests metadata they feed) once the admin UI reads SGR @@ -826,8 +860,7 @@ def _build_aggregated_sql_query( | GROUPING(model, {_MODEL_GROUP_EXPR}, custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level,{metric_select} - FROM "{key_free_source}" - WHERE {where_clause} + FROM {_key_free_source(pg_table, where_clause, marker_param)} GROUP BY GROUPING SETS ( (date), (date, model), @@ -869,7 +902,8 @@ def _build_aggregated_sql_query( )) """ - return sql_query, [*where_params, PTU_SENTINEL_API_KEY] + marker_params: Final = () if global_rollup_through is None else (global_rollup_through,) + return sql_query, [*where_params, PTU_SENTINEL_API_KEY, *marker_params] def _build_entity_rollup_sql_query( @@ -1418,7 +1452,8 @@ async def get_daily_activity_aggregated( include_current_utc_day=include_current_utc_day, ) sql_query, sql_params = _build_aggregated_sql_query( - **query_kwargs, key_free_table=await key_free_source_table(prisma_client, query_kwargs) + **query_kwargs, + global_rollup_through=await global_rollup_reconciled_through(prisma_client, query_kwargs), ) entity_query: Final = _build_entity_rollup_sql_query(**query_kwargs) if include_entity_breakdown else None diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index 9d344421332..a9fb7669785 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -1,9 +1,10 @@ -"""Reconcile ``LiteLLM_DailyGlobalSpend`` from ``LiteLLM_DailyUserSpend``, one day per transaction. +"""Roll closed UTC days of ``LiteLLM_DailyUserSpend`` up into ``LiteLLM_DailyGlobalSpend``. -The spend writer keeps both tables in step from the moment it is deployed; this job rolls up -the days before that and records how far it has reached in ``LiteLLM_Config`` so usage reads -know when the global table can answer for a date range. It runs as a background cron, never -in a Prisma migration, since on a large deployment the aggregate is minutes of work. +Only days that are over get rolled up, so a pod still flushing per-key spend for the current +day can never leave the global table short; usage reads serve days through the recorded +marker from the global table and later days live from the per-key table. The marker lives in +``LiteLLM_Config``. This runs as a background cron, never in a Prisma migration, since on a +large deployment the first backfill is minutes of work. """ from collections.abc import Awaitable, Callable @@ -19,7 +20,6 @@ from litellm.constants import ( DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS, DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, ) -from litellm.proxy.db.daily_spend_bulk_upsert import GLOBAL_SPEND_TABLE from litellm.repositories.config_repository import ConfigRepository if TYPE_CHECKING: @@ -27,8 +27,11 @@ if TYPE_CHECKING: from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.proxy.utils import PrismaClient -_DAY_TRANSACTION_TIMEOUT: Final = timedelta(minutes=10) _REPLAY_DAYS: Final = 1 +GLOBAL_SPEND_TABLE_NAME: Final = "LiteLLM_DailyGlobalSpend" +# The unique constraint, in constraint order. NULL never matches itself in a unique index, so +# every column is normalized to '' or the same group would be inserted again on every run. +_KEY_COLUMNS: Final = ("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") _METRIC_COLUMNS: Final = ( "prompt_tokens", "completion_tokens", @@ -51,23 +54,21 @@ def _quoted(columns: tuple[str, ...]) -> str: def _reconcile_day_sql() -> str: - key_columns: Final = GLOBAL_SPEND_TABLE.key_columns - normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in key_columns) + normalized_keys: Final = ", ".join(f"COALESCE(\"{column}\", '')" for column in _KEY_COLUMNS) sums: Final = ", ".join(f'SUM("{column}")' for column in _METRIC_COLUMNS) overwrite: Final = ", ".join(f'"{column}" = EXCLUDED."{column}"' for column in _METRIC_COLUMNS) return ( - f'INSERT INTO "{GLOBAL_SPEND_TABLE.name}" ("id", {_quoted(key_columns)}, {_quoted(_METRIC_COLUMNS)}, ' + f'INSERT INTO "{GLOBAL_SPEND_TABLE_NAME}" ("id", {_quoted(_KEY_COLUMNS)}, {_quoted(_METRIC_COLUMNS)}, ' '"updated_at")\n' f"SELECT gen_random_uuid()::text, {normalized_keys}, {sums}, (NOW() AT TIME ZONE 'UTC')\n" 'FROM "LiteLLM_DailyUserSpend" WHERE "date" = $1\n' f"GROUP BY {normalized_keys}\n" - f"ON CONFLICT ({_quoted(key_columns)}) DO UPDATE SET {overwrite}, " + f"ON CONFLICT ({_quoted(_KEY_COLUMNS)}) DO UPDATE SET {overwrite}, " "\"updated_at\" = (NOW() AT TIME ZONE 'UTC')" ) RECONCILE_DAY_SQL: Final = _reconcile_day_sql() -_LOCK_GLOBAL_TABLE_SQL: Final = f'LOCK TABLE "{GLOBAL_SPEND_TABLE.name}" IN EXCLUSIVE MODE' _PENDING_DAYS_SQL: Final = ( 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" >= $1 AND "date" <= $2 ORDER BY "date"' ) @@ -134,19 +135,19 @@ def _first_pending_day(marker: str | None) -> str: async def pending_days(prisma_client: "PrismaClient", today: date) -> tuple[str, ...]: - """Every UTC day through today still to roll up, oldest first; the marker day and the one - before it are replayed so rows flushed by a pre-writer pod during a rolling deploy are folded in.""" + """Every closed UTC day (strictly before today) still to roll up, oldest first. The marker + day and the one before it are replayed so per-key rows that landed after their day was + rolled up (a flush straddling midnight, a late retry) are folded in.""" marker: Final = await reconciled_through(prisma_client) - rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), today.isoformat()) - return tuple(sorted({*(_DateRow.model_validate(row).date for row in rows), today.isoformat()})) + last_closed_day: Final = (today - timedelta(days=1)).isoformat() + rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), last_closed_day) + return tuple(_DateRow.model_validate(row).date for row in rows) async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None: - """Rewrite one day of the global table from the per-key sums; the table lock keeps the - writer's increments out between the aggregate and the overwrite so none are lost.""" - async with prisma_client.db.tx(timeout=_DAY_TRANSACTION_TIMEOUT) as transaction: - await transaction.execute_raw(_LOCK_GLOBAL_TABLE_SQL) - await transaction.execute_raw(RECONCILE_DAY_SQL, day) + """Rewrite one day of the global table from the per-key sums. Idempotent: a rerun + overwrites every group with the same totals.""" + await prisma_client.db.execute_raw(RECONCILE_DAY_SQL, day) async def run_daily_global_spend_reconcile( diff --git a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py index cc443a2cfe5..c1efb3e7220 100644 --- a/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py +++ b/tests/test_litellm/proxy/db/test_daily_spend_bulk_upsert.py @@ -1,19 +1,12 @@ """Tests for the single-statement daily spend upsert (LIT-5291).""" -import pathlib import re -from typing import Final -import psycopg import pytest -from psycopg.rows import dict_row -from pytest_postgresql import factories from litellm.proxy.db.daily_spend_bulk_upsert import ( DAILY_SPEND_TABLES, - GLOBAL_SPEND_TABLE, build_bulk_upsert, - build_bulk_upsert_with_global_rollup, conflict_key, merge_by_conflict_key, ) @@ -192,146 +185,3 @@ async def test_writer_survives_a_transaction_whose_key_columns_are_null(): _, params = prisma_client.db.statements[0] assert None not in params[:9] assert transactions == {} - - -def user_txn(**overrides): - txn = {**tag_txn(), "user_id": "u-1", **overrides} - del txn["tag"] - del txn["request_id"] - return txn - - -def _bound_rows(insert_sql: str, params: tuple[object, ...]) -> list[dict[str, object]]: - """Each VALUES row of one INSERT as a column -> bound value mapping, consuming params in order.""" - header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", insert_sql) - assert header is not None, insert_sql - columns = [c.strip('"') for c in header.group(1).split(", ") if c != '"updated_at"'] - row_count = insert_sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") - return [dict(zip(columns, params[i * len(columns) : (i + 1) * len(columns)])) for i in range(row_count)] - - -def test_global_rollup_folds_every_key_and_user_into_one_row_per_dimension_tuple(): - """The global table has no api_key or user_id, so a batch spread over many keys and - users must collapse to one row per (date, model, group, provider, mcp, endpoint).""" - batch = merge_by_conflict_key( - USER_TABLE, - tuple(user_txn(user_id=f"u-{i}", api_key=f"sk-{i}", spend=1.0, api_requests=1) for i in range(5)) - + (user_txn(user_id="u-0", api_key="sk-0", model="claude", spend=10.0, api_requests=3),), - ) - - sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) - - entity_insert, global_insert = sql.split("RETURNING 1)") - entity_rows = _bound_rows(entity_insert, params) - global_rows = _bound_rows(global_insert, params[len(entity_rows) * len(entity_rows[0]) :]) - assert len(entity_rows) == 6 - assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in global_insert - assert [(r["model"], r["spend"], r["api_requests"]) for r in global_rows] == [ - ("claude", 10.0, 3), - ("gpt-4o-mini", 5.0, 5), - ] - assert all("api_key" not in r and "user_id" not in r for r in global_rows) - conflict = re.search(r"ON CONFLICT \(([^)]*)\)", global_insert) - assert conflict is not None - assert conflict.group(1) == ", ".join(f'"{c}"' for c in GLOBAL_SPEND_TABLE.key_columns) - - -def test_global_rollup_params_follow_the_entity_params_in_one_placeholder_sequence(): - """Both inserts bind from one flat tuple, so the global arm's placeholders must start - exactly where the entity arm's stop or every value lands one column off.""" - batch = merge_by_conflict_key(USER_TABLE, (user_txn(),)) - - sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) - - placeholders = [int(n) for n in re.findall(r"\$(\d+)::", sql)] - assert placeholders == list(range(1, len(params) + 1)) - - -_bulk_upsert_postgresql_proc: Final = factories.postgresql_proc() -_bulk_upsert_postgresql: Final = factories.postgresql("_bulk_upsert_postgresql_proc") - -_MIGRATIONS_DIR: Final = ( - pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations" -) -_GLOBAL_SPEND_MIGRATION: Final = _MIGRATIONS_DIR / "20260915000000_add_daily_global_spend" / "migration.sql" - -_DAILY_USER_SPEND_DDL: Final = """ - CREATE TABLE "LiteLLM_DailyUserSpend" ( - id TEXT PRIMARY KEY, - user_id TEXT, - date TEXT NOT NULL, - api_key TEXT NOT NULL, - model TEXT, - model_group TEXT, - custom_llm_provider TEXT, - mcp_namespaced_tool_name TEXT, - endpoint TEXT, - prompt_tokens BIGINT DEFAULT 0, - completion_tokens BIGINT DEFAULT 0, - cache_read_input_tokens BIGINT DEFAULT 0, - cache_creation_input_tokens BIGINT DEFAULT 0, - compression_saved_tokens BIGINT DEFAULT 0, - compression_savings_spend DOUBLE PRECISION DEFAULT 0, - prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0, - gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0, - autorouter_savings_spend DOUBLE PRECISION DEFAULT 0, - spend DOUBLE PRECISION DEFAULT 0, - api_requests BIGINT DEFAULT 0, - successful_requests BIGINT DEFAULT 0, - failed_requests BIGINT DEFAULT 0, - created_at TIMESTAMP DEFAULT now(), - updated_at TIMESTAMP, - UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint) - ) -""" - - -def _execute_dollar_sql(conn: psycopg.Connection, sql: str, params: tuple[object, ...]) -> None: - converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql) - conn.execute( - converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query - {f"p{i}": v for i, v in enumerate(params, start=1)}, - ) - conn.commit() - - -def test_global_rollup_equals_the_per_key_sums_after_repeated_flushes(_bulk_upsert_postgresql: psycopg.Connection): - """Against real Postgres and the shipped migration: two flushes of a mixed batch leave - the global table exactly equal to the per-key table summed over user and key, with the - NULL and '' spellings of a dimension folded into one row.""" - conn: Final = _bulk_upsert_postgresql - conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal - conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal - conn.commit() - - batch = merge_by_conflict_key( - USER_TABLE, - ( - user_txn(user_id="u-1", api_key="sk-1", spend=1.0, prompt_tokens=10), - user_txn(user_id="u-2", api_key="sk-2", spend=2.0, prompt_tokens=20), - user_txn(user_id="u-1", api_key="sk-3", model=None, custom_llm_provider=None, spend=4.0), - user_txn(user_id="u-3", api_key="sk-4", model="", custom_llm_provider="", spend=8.0), - ), - ) - sql, params = build_bulk_upsert_with_global_rollup(USER_TABLE, batch) - _execute_dollar_sql(conn, sql, params) - _execute_dollar_sql(conn, sql, params) - - with conn.cursor(row_factory=dict_row) as cur: - global_rows = cur.execute( - 'SELECT model, spend, prompt_tokens, api_requests FROM "LiteLLM_DailyGlobalSpend" ORDER BY model' - ).fetchall() - per_key = cur.execute( - """ - SELECT COALESCE(model, '') AS model, SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, - SUM(api_requests) AS api_requests - FROM "LiteLLM_DailyUserSpend" GROUP BY COALESCE(model, '') ORDER BY 1 - """ - ).fetchall() - - assert [row["model"] for row in global_rows] == ["", "gpt-4o-mini"] - assert [(r["model"], r["spend"], int(r["prompt_tokens"]), int(r["api_requests"])) for r in global_rows] == [ - (r["model"], float(r["spend"]), int(r["prompt_tokens"]), int(r["api_requests"])) for r in per_key - ] - assert global_rows[0]["spend"] == pytest.approx(24.0) - assert global_rows[1]["spend"] == pytest.approx(6.0) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index d8a9013398e..5e977712a1e 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -254,19 +254,14 @@ class _RecordingPrisma: def _row_values(statement: Statement, column: str) -> list[object]: - """Every row's value for one column of the first INSERT, read out of the flat parameter tuple. - - The user-table statement chains a global rollup INSERT after its own, so the row count - comes from the first INSERT's VALUES rather than from the parameter count. - """ + """Every row's value for one column, read out of the flat parameter tuple.""" sql, params = statement header = re.search(r"INSERT INTO \"[A-Za-z_]+\" \(([^)]*)\)", sql) assert header is not None, sql columns = header.group(1).split(", ") stride = len(columns) - 1 # updated_at is inlined, not bound offset = columns.index(f'"{column}"') - rows = sql.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") - return [params[row * stride + offset] for row in range(rows)] + return [params[row * stride + offset] for row in range(len(params) // stride)] @pytest.mark.asyncio @@ -1468,98 +1463,6 @@ async def test_update_daily_spend_keeps_failed_transactions_for_retry(): assert daily_spend_transactions == expected -def _entity_txn(entity_field: str, entity_id: str, api_key: str) -> dict[str, object]: - txn = _daily_txn() - del txn["user_id"] - return {**txn, entity_field: entity_id, "api_key": api_key} - - -@pytest.mark.asyncio -async def test_user_flush_writes_the_global_rollup_in_the_same_statement(): - """The user flush is the one place per-key spend becomes key-free spend, so a batch spread - over many keys must land in LiteLLM_DailyGlobalSpend as one row in the same statement. - A separate statement would let a crash between the two leave the tables out of sync.""" - prisma_client = _RecordingPrisma() - txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(4)} - - await DBSpendUpdateWriter._update_daily_spend( - n_retry_times=0, - prisma_client=prisma_client, - proxy_logging_obj=MagicMock(), - daily_spend_transactions=txns, - entity_type="user", - entity_id_field="user_id", - ) - - assert len(prisma_client.db.statements) == 1 - sql, params = prisma_client.db.statements[0] - assert sql.count('INSERT INTO "LiteLLM_DailyUserSpend"') == 1 - assert sql.count('INSERT INTO "LiteLLM_DailyGlobalSpend"') == 1 - assert sql.index('"LiteLLM_DailyUserSpend"') < sql.index('"LiteLLM_DailyGlobalSpend"') - global_insert = sql.split('INSERT INTO "LiteLLM_DailyGlobalSpend"', 1)[1] - assert global_insert.split("ON CONFLICT", 1)[0].count("(NOW() AT TIME ZONE 'UTC'))") == 1 - assert "api_key" not in global_insert - assert params.count(0.4) == 1 - assert txns == {} - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("entity_type", "entity_field"), - [ - ("team", "team_id"), - ("org", "organization_id"), - ("tag", "tag"), - ("end_user", "end_user_id"), - ("agent", "agent_id"), - ], -) -async def test_other_entity_flushes_leave_the_global_table_alone(entity_type, entity_field): - """Every entity table sees the same request, so writing the rollup from more than one of - them would count each request once per entity type.""" - prisma_client = _RecordingPrisma() - txn = _entity_txn(entity_field, "e-1", "sk-1") - if entity_type == "tag": - txn["request_id"] = "req-1" - - await DBSpendUpdateWriter._update_daily_spend( - n_retry_times=0, - prisma_client=prisma_client, - proxy_logging_obj=MagicMock(), - daily_spend_transactions={"k": txn}, - entity_type=entity_type, - entity_id_field=entity_field, - ) - - (sql, _params) = prisma_client.db.statements[0] - assert "LiteLLM_DailyGlobalSpend" not in sql - - -@pytest.mark.asyncio -async def test_a_failed_chained_user_flush_keeps_every_transaction_for_retry(): - def raise_outage(): - raise ValueError("simulated database outage") - - prisma_client = _RecordingPrisma(execute_raw=raise_outage) - txns = {f"k{i}": _entity_txn("user_id", f"user-{i}", f"sk-{i}") for i in range(3)} - expected = dict(txns) - mock_proxy_logging = MagicMock() - mock_proxy_logging.failure_handler = AsyncMock() - - with pytest.raises(ValueError, match="simulated database outage"): - await DBSpendUpdateWriter._update_daily_spend( - n_retry_times=0, - prisma_client=prisma_client, - proxy_logging_obj=mock_proxy_logging, - daily_spend_transactions=txns, - entity_type="user", - entity_id_field="user_id", - ) - - assert txns == expected - assert 'INSERT INTO "LiteLLM_DailyGlobalSpend"' in prisma_client.db.statements[0][0] - - @pytest.mark.asyncio async def test_commit_key_spend_updates_includes_last_active(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 5c74facae6a..12e5fe6af4d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1,3 +1,4 @@ +import pathlib import re from collections.abc import Sequence from datetime import datetime, timedelta, timezone @@ -10,11 +11,6 @@ import pytest from psycopg.rows import dict_row from pytest_postgresql import factories -from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR - - -import pathlib - from litellm.constants import ( DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, PTU_SENTINEL_API_KEY, @@ -29,10 +25,11 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( get_api_key_metadata, get_daily_activity, get_daily_activity_aggregated, - key_free_source_table, + global_rollup_reconciled_through, update_metrics, ) from litellm.proxy.spend_tracking.daily_global_spend_rollup import RECONCILE_DAY_SQL +from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR from litellm.proxy.utils import evict_config_param from litellm.types.proxy.management_endpoints.common_daily_activity import ( DailySpendMetadata, @@ -1632,7 +1629,9 @@ def _prisma_with_marker(marker: str | None) -> MagicMock: prisma.db = MagicMock() prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - row = None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}') + row = ( + None if marker is None else SimpleNamespace(param_name="m", param_value=f'{{"reconciled_through": "{marker}"}}') + ) prisma.get_generic_data = AsyncMock(return_value=row) return prisma @@ -1657,9 +1656,9 @@ def _unfiltered_user_query(**overrides): @pytest.mark.parametrize( ("marker", "overrides", "expected"), [ - ("2026-06-02", {}, "LiteLLM_DailyGlobalSpend"), - ("2026-06-02", {"model": "gpt-5"}, "LiteLLM_DailyGlobalSpend"), - ("2026-06-01", {}, None), + ("2026-06-02", {}, "2026-06-02"), + ("2026-06-02", {"model": "gpt-5"}, "2026-06-02"), + ("2026-05-01", {}, "2026-05-01"), (None, {}, None), ("2026-06-02", {"api_key": "sk-1"}, None), ("2026-06-02", {"api_key": []}, None), @@ -1668,33 +1667,50 @@ def _unfiltered_user_query(**overrides): ("2026-06-02", {"table_name": "litellm_dailyteamspend", "entity_id_field": "team_id"}, None), ], ) -async def test_key_free_source_table_routes_only_unfiltered_user_reads_within_the_marker(marker, overrides, expected): - """Anything that filters by key or entity has no counterpart in the global table, and a - range the reconcile has not reached must stay on the per-key table.""" +async def test_global_rollup_marker_is_used_only_for_unfiltered_user_reads(marker, overrides, expected): + """Anything that filters by key or entity has no counterpart in the global table; the + SQL splits the range at the marker itself, so the marker passes through unchanged.""" await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) prisma = _prisma_with_marker(marker) - assert await key_free_source_table(prisma, _unfiltered_user_query(**overrides)) == expected + assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query(**overrides)) == expected await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) @pytest.mark.asyncio -async def test_key_free_source_table_judges_the_timezone_extended_end_not_the_requested_one(): - """A caller west of UTC asking through their local today gets today's UTC bucket added to - the range; the marker must cover that extended day, not just the requested end.""" +async def test_global_rollup_marker_read_failure_falls_back_to_the_per_key_table(): await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) - today_utc: Final = datetime.now(timezone.utc).date() - yesterday: Final = (today_utc - timedelta(days=1)).isoformat() - query: Final = _unfiltered_user_query( - start_date=yesterday, end_date=yesterday, timezone_offset_minutes=24 * 60, include_current_utc_day=True - ) + prisma = _prisma_with_marker(None) + prisma.get_generic_data = AsyncMock(side_effect=RuntimeError("db down")) - assert await key_free_source_table(_prisma_with_marker(yesterday), query) is None - await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) - assert await key_free_source_table(_prisma_with_marker(today_utc.isoformat()), query) == "LiteLLM_DailyGlobalSpend" + assert await global_rollup_reconciled_through(prisma, _unfiltered_user_query()) is None await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) +def test_aggregated_sql_splits_the_key_free_arm_at_the_marker_and_keeps_the_key_arm_per_key(): + sql, params = _build_aggregated_sql_query(**_unfiltered_user_query(), global_rollup_through="2026-06-01") + marker_param: Final = f"${len(params)}" + + assert params[-1] == "2026-06-01" + assert ( + f'FROM "LiteLLM_DailyGlobalSpend"\n WHERE date >= $1 AND date <= $2 AND date <= {marker_param}' + in sql + ) + assert ( + f'FROM "LiteLLM_DailyUserSpend"\n WHERE date >= $1 AND date <= $2 AND date > {marker_param}' in sql + ) + key_arm: Final = sql.split("UNION ALL\n (WITH top_api_keys")[1] + assert "LiteLLM_DailyGlobalSpend" not in key_arm + assert marker_param not in key_arm + + +def test_aggregated_sql_without_a_marker_reads_the_per_key_table_only(): + sql, params = _build_aggregated_sql_query(**_unfiltered_user_query()) + + assert "LiteLLM_DailyGlobalSpend" not in sql + assert params[-1] == PTU_SENTINEL_API_KEY + + _GLOBAL_SPEND_MIGRATION: Final = ( pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" @@ -1706,12 +1722,12 @@ _GLOBAL_SPEND_MIGRATION: Final = ( @pytest.mark.asyncio -async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_free_arm( +async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_table_and_open_days_live( _aggregated_postgresql: psycopg.Connection, ): - """With the range reconciled, the key-free arm reads LiteLLM_DailyGlobalSpend while the - per-key arm stays on the user table, and the response is identical to the all-per-key - read: same totals, same rollups, same top keys.""" + """Day 1 is rolled up and day 2 is still open (never rolled up), so a marker of day 1 must + give the same response as reading everything per-key: day 1 from the global table, day 2 + live, one grand total across both. The per-key arm stays on the user table throughout.""" n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 3 rows: Final = [ ( @@ -1734,11 +1750,10 @@ async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_ _seed_daily_user_spend(_aggregated_postgresql, rows) with _aggregated_postgresql.cursor() as cur: cur.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal - for day in ("2026-06-01", "2026-06-02"): - cur.execute( - re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg - {"p1": day}, - ) + cur.execute( + re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg + {"p1": "2026-06-01"}, + ) _aggregated_postgresql.commit() async def read(marker: str | None, sql_seen: list[str]): @@ -1760,16 +1775,19 @@ async def test_get_daily_activity_aggregated_reads_the_global_table_for_the_key_ per_key_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim global_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim from_per_key = await read(None, per_key_sql) - from_global = await read("2026-06-02", global_sql) + from_global = await read("2026-06-01", global_sql) await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) assert per_key_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 0 assert global_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 1 - assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 2 + assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 3 assert from_global.model_dump() == from_per_key.model_dump() assert from_global.metadata.total_spend == pytest.approx(2 * sum(float(i + 1) for i in range(n_keys))) + assert {day.date.isoformat() for day in from_global.results} == {"2026-06-01", "2026-06-02"} assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"} + + def _no_spend_record(): """A rollup row for a key with no spend, where SUM() returns NULL (None).""" return SimpleNamespace( diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 13dc757cbbd..11ca72e7b3d 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -2,7 +2,6 @@ import pathlib import re -from contextlib import asynccontextmanager from datetime import date from typing import Final from unittest.mock import AsyncMock, MagicMock @@ -13,11 +12,7 @@ from psycopg.rows import dict_row from pytest_postgresql import factories from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM -from litellm.proxy.db.daily_spend_bulk_upsert import ( - DAILY_SPEND_TABLES, - build_bulk_upsert_with_global_rollup, - merge_by_conflict_key, -) +from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( RECONCILE_DAY_SQL, reconciled_through, @@ -45,21 +40,6 @@ class _FakeConfigTable: return _FakeConfigRow(where["param_name"], data["update"]["param_value"]) -class _FakeTransaction: - def __init__(self, prisma: "_FakePrisma") -> None: - self._prisma = prisma - - async def execute_raw(self, sql: str, *params: str) -> int: - if "LOCK TABLE" in sql: - self._prisma.locks_taken += 1 - return 0 - (day,) = params - if day in self._prisma.failing_days: - raise RuntimeError(f"day {day} exploded") - self._prisma.reconciled.append(day) - return 1 - - class _FakeDb: def __init__(self, prisma: "_FakePrisma") -> None: self._prisma = prisma @@ -69,19 +49,21 @@ class _FakeDb: first, last = params return [{"date": d} for d in sorted(self._prisma.user_days) if first <= d <= last] - @asynccontextmanager - async def tx(self, timeout: object): - yield _FakeTransaction(self._prisma) + async def execute_raw(self, sql: str, *params: str) -> int: + (day,) = params + if day in self._prisma.failing_days: + raise RuntimeError(f"day {day} exploded") + self._prisma.reconciled.append(day) + return 1 class _FakePrisma: - """Enough of PrismaClient for the reconcile: per-key dates, a config table, and a transaction.""" + """Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw.""" def __init__(self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset()) -> None: self.user_days = user_days self.failing_days = failing_days self.reconciled: list[str] = [] - self.locks_taken = 0 self.db = _FakeDb(self) async def get_generic_data(self, key: str, value: str, table_name: str) -> _FakeConfigRow | None: @@ -97,33 +79,45 @@ async def _fresh_marker_cache(): @pytest.mark.asyncio -async def test_first_run_rolls_up_every_historical_day_and_today_then_marks_today(): - """Before any marker exists, every day with per-key rows is rolled up, plus today even - with no rows yet, so reads for ranges ending today can switch to the global table.""" - prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14")) +async def test_first_run_rolls_up_every_closed_day_and_never_today(): + """Before any marker exists every closed day with per-key rows is rolled up. Today is left + out: pods are still flushing it, so it is served live from the per-key table until it closes.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15")) result = await run_daily_global_spend_reconcile(prisma, today=TODAY) - assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15") + assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14") assert result.failed_day is None - assert result.reconciled_through == "2026-09-15" - assert await reconciled_through(prisma) == "2026-09-15" - assert prisma.locks_taken == 4 + assert result.reconciled_through == "2026-09-14" + assert await reconciled_through(prisma) == "2026-09-14" + assert "2026-09-15" not in prisma.reconciled @pytest.mark.asyncio async def test_later_run_replays_the_marker_day_and_the_day_before_only(): """Days older than marker-1 are settled; the marker day and its predecessor are replayed so - rows a pre-writer pod flushed around midnight during a rolling deploy get folded in.""" + per-key rows that landed after their day was rolled up get folded in.""" prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14")) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13)) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) prisma.reconciled.clear() result = await run_daily_global_spend_reconcile(prisma, today=TODAY) - assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14", "2026-09-15") + assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14") assert "2026-09-01" not in prisma.reconciled - assert await reconciled_through(prisma) == "2026-09-15" + assert await reconciled_through(prisma) == "2026-09-14" + + +@pytest.mark.asyncio +async def test_a_run_with_no_new_closed_days_keeps_the_marker(): + prisma = _FakePrisma(user_days=("2026-09-13",)) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma.reconciled.clear() + + result = await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + + assert result.days_reconciled == ("2026-09-13",) + assert result.reconciled_through == "2026-09-13" @pytest.mark.asyncio @@ -149,16 +143,16 @@ async def test_the_next_run_resumes_from_the_failed_day(): result = await run_daily_global_spend_reconcile(prisma, today=TODAY) - assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03", "2026-09-15") - assert await reconciled_through(prisma) == "2026-09-15" + assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03") + assert await reconciled_through(prisma) == "2026-09-03" @pytest.mark.asyncio async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts(): - """A pre-writer pod flushing rows for the day before the marker is exactly the replay case; - when that replay fails the marker must stay put and the operator must hear about it.""" + """A late flush for the day before the marker is exactly the replay case; when that replay + fails the marker must stay put and the operator must hear about it.""" prisma = _FakePrisma(user_days=("2026-09-13",)) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 13)) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) prisma.user_days = ("2026-09-12", "2026-09-13") prisma.failing_days = frozenset({"2026-09-12"}) alert = AsyncMock() @@ -212,7 +206,7 @@ async def test_scheduled_run_runs_and_releases_the_lock_when_it_wins(): result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) - assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15") + assert result is not None and result.days_reconciled == ("2026-09-13",) lock.release_lock.assert_awaited_once() @@ -226,7 +220,7 @@ async def test_scheduled_run_proceeds_when_the_lock_cannot_be_acquired_or_read() result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) - assert result is not None and result.days_reconciled == ("2026-09-13", "2026-09-15") + assert result is not None and result.days_reconciled == ("2026-09-13",) lock.release_lock.assert_not_awaited() @@ -341,9 +335,9 @@ def _normalized(rows: list[dict[str, object]]) -> list[tuple[object, ...]]: def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_postgresql: psycopg.Connection): - """Against real Postgres and the shipped migration: rows the writer never saw (a - pre-writer pod's flush, NULL and '' dimension spellings) end up folded into the global - day, running the day twice changes nothing, and other days are left alone.""" + """Against real Postgres and the shipped migration: writer-shaped rows and legacy rows + (NULL and '' dimension spellings) fold into one global day, running the day twice changes + nothing, and other days are left alone.""" conn: Final = _rollup_postgresql conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal conn.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal @@ -353,7 +347,7 @@ def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_p USER_TABLE, (_user_txn(api_key="sk-1", spend=1.0), _user_txn(api_key="sk-2", user_id="u-2", spend=2.0, prompt_tokens=20)), ) - _execute_dollar_sql(conn, *build_bulk_upsert_with_global_rollup(USER_TABLE, written_batch)) + _execute_dollar_sql(conn, *build_bulk_upsert(USER_TABLE, written_batch)) conn.execute( """ From aa7f1e16b8810a9321fd12539732ff738aeebd38 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 15 Sep 2026 23:53:14 +0000 Subject: [PATCH 048/253] feat(proxy): throttle failed Admin UI sign-ins per source and source/username Replace the username-global lockout with counters keyed by source address and by source/username pair. Each has a fixed counting window (60s) and a separate block TTL (300s). Blocks are soft: a correct password still signs in, wrong passwords from a blocked key take one of 5 held slots per worker and are held 30s before a 429. Once a pair is blocked its failures stop counting against the source. The source scope runs only when trusted_proxy_ranges is set, IPv6 is grouped by /64, and per-source limits accept IP and CIDR overrides with longest-prefix matching. Redis is authoritative through one Lua script per failure, with bounded per-worker fallback when Redis raises. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_cache.py | 4 +- litellm/proxy/_types.py | 23 +- litellm/proxy/auth/login_throttle.py | 578 +++++---- litellm/proxy/auth/login_utils.py | 26 +- litellm/proxy/proxy_server.py | 12 +- .../proxy/auth/test_login_utils.py | 1072 +++++++++-------- .../proxy/proxy_server/conftest.py | 31 +- .../proxy_server/test_routes_login_sso.py | 121 +- tests/test_litellm/proxy/test_proxy_server.py | 34 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 26 +- 10 files changed, 1043 insertions(+), 884 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 3f1bac12563..b4b2b1a334c 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -874,7 +874,7 @@ class RedisCache(BaseCache): return _LUA_COUNT.validate_python(count) @_redis_circuit_breaker_guard_sync - def batch_get_counts(self, key_list: list[str] | tuple[str, ...]) -> tuple[int | None, ...]: + def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]: """Read integer counters for ``key_list``, in order, raising when Redis cannot answer. ``batch_get_cache`` swallows every failure and returns an empty dict, which the caller @@ -885,7 +885,7 @@ class RedisCache(BaseCache): return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys)) @_redis_circuit_breaker_guard - async def async_batch_get_counts(self, key_list: list[str] | tuple[str, ...]) -> tuple[int | None, ...]: + async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]: """Async twin of ``batch_get_counts``, raising on failure the same way.""" namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list] return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys)) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ddcc7b2dece..2537555f316 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2736,20 +2736,29 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): description="sends alerts if requests hang for 5min+", ) ui_access_mode: Literal["admin_only", "all"] | None = Field("all", description="Control access to the Proxy UI") - max_failed_login_attempts: int | None = Field( - None, - ge=1, - description="Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Set under `general_settings` in config.yaml. Defaults to 50", - ) max_failed_login_attempts_per_source: int | None = Field( None, ge=1, - description="Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Set under `general_settings` in config.yaml. Defaults to 250", + description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set, since otherwise every client behind an ingress shares one address. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", + ) + max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field( + None, + description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins. Set under `general_settings` in config.yaml", + ) + max_failed_login_attempts_per_user: int | None = Field( + None, + ge=1, + description="Failed Admin UI sign-in attempts allowed from one source address for one username within `failed_login_window_seconds`. One more blocks that address for that username for `failed_login_block_seconds`, and its further failures stop counting against `max_failed_login_attempts_per_source`, so a script stuck on one account does not block everyone behind the same address. Set under `general_settings` in config.yaml. Defaults to 5", ) failed_login_window_seconds: int | None = Field( None, ge=1, - description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 900", + description="Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 60", + ) + failed_login_block_seconds: int | None = Field( + None, + ge=1, + description="How long a blocked source address, or source address and username, stays blocked. The block is soft: a correct password still signs in, but each attempt from a blocked key takes one of 5 held slots per worker and a wrong password is held for 30 seconds before it is refused with 429. Set under `general_settings` in config.yaml. Defaults to 300", ) allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access") reject_clientside_metadata_tags: bool | None = Field( diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 50333ffad05..fb708ecf872 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -1,142 +1,231 @@ """Failed-login accounting for the Admin UI sign-in path. -Counts failed credential checks over a fixed window against two independent keys, the -username on its own and the source address on its own, so that one username attacked from -many sources and one source spraying many usernames are both counted. Repeated failures -are answered slowly, doubling from one second, and refused with 429 once either counter -reaches its limit. Built per request by ``LoginThrottle.from_request`` because it carries -that request's resolved source address, and because the coordination cache is assigned at -startup and can be reassigned later. +Wrong passwords are counted over a short window per source address and per source-and-username +pair; too many in one window blocks that key for a fixed time. Blocks are soft: a correct password +still signs in, while a wrong one from a blocked key is held open before its 429 and only a few can +be held at once, which bounds how many guesses a blocked key gets checked. A blocked pair stops +counting against its source, so one script stuck on one account does not block the whole office. """ +from __future__ import annotations + import asyncio import hashlib -from collections.abc import Awaitable, Callable, Mapping +import ipaddress +import math +import time +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager from dataclasses import dataclass from functools import cache from types import MappingProxyType -from typing import Final, NamedTuple, NoReturn +from typing import Final, Literal, NamedTuple, NoReturn -from fastapi import Request +from fastapi import Request, status +from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger -from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges, resolve_client_ip -from litellm.proxy.auth.trusted_proxy_utils import TRUSTED_PROXY_RANGES_KEY from litellm.secret_managers.main import get_secret_bool -DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS: Final = 50 -DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 250 -DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 900 +DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 10 +DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER: Final = 5 +DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60 +DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300 -USERNAME_DELAY_ONSET: Final = 3 -SOURCE_DELAY_ONSET: Final = 25 -FIRST_DELAY_SECONDS: Final = 1.0 -MAX_DELAY_SECONDS: Final = 30.0 -MAX_CONCURRENT_DELAYS_PER_SOURCE: Final = 5 +BLOCKED_ATTEMPT_HOLD_SECONDS: Final = 30 +MAX_HELD_ATTEMPTS_PER_KEY: Final = 5 +IPV6_SOURCE_PREFIX_LENGTH: Final = 64 -_MAX_DELAY_DOUBLINGS: Final = 16 +SOURCE_LIMIT_KEY: Final = "max_failed_login_attempts_per_source" +SOURCE_LIMIT_OVERRIDES_KEY: Final = "max_failed_login_attempts_per_source_overrides" +USER_LIMIT_KEY: Final = "max_failed_login_attempts_per_user" +WINDOW_KEY: Final = "failed_login_window_seconds" +BLOCK_KEY: Final = "failed_login_block_seconds" +TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges" _CACHE_KEY_PREFIX: Final = "login_fail" _UNKNOWN_SOURCE: Final = "unknown" -_MAX_LOGGED_USERNAME_CHARS: Final = 128 +_MAX_TRACKED_COUNTERS: Final = 20_000 +_MAX_TRACKED_BLOCKS: Final = 10_000 +_NO_SETTINGS: Final[Mapping[str, object]] = MappingProxyType({}) +_NOT_BLOCKED: Final = (0, 0) +_LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None) +_SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object]) -_MAX_TRACKED_LOGIN_USERNAMES: Final = 10_000 -_MAX_TRACKED_LOGIN_SOURCES: Final = 10_000 +Scope = Literal["user", "source"] + +_BlockTtls = tuple[int, int] +_LUA_BLOCK_TTLS: Final = TypeAdapter[_BlockTtls](_BlockTtls) +_Network = ipaddress.IPv4Network | ipaddress.IPv6Network + +# KEYS: pair counter, pair block, source counter, source block (one cluster slot via the source hash tag) +# ARGV: pair limit, source limit (0 = source scope off), window seconds, block seconds +# Both scripts return {pair block TTL, source block TTL}; 0 or below means not blocked +_BLOCK_TTLS_LUA: Final = "return {redis.call('TTL', KEYS[2]), redis.call('TTL', KEYS[4])}" +_RECORD_FAILURE_LUA: Final = ( + "local function bump(count_key, block_key, limit) " + "local blocked = redis.call('TTL', block_key) " + "if blocked > 0 then return blocked end " + "local count = redis.call('INCR', count_key) " + "if redis.call('TTL', count_key) < 0 then redis.call('EXPIRE', count_key, ARGV[3]) end " + "if count > limit then redis.call('SET', block_key, '1', 'EX', ARGV[4]) return tonumber(ARGV[4]) end " + "return 0 end " + "local user_block = bump(KEYS[1], KEYS[2], tonumber(ARGV[1])) " + "local source_block = 0 " + "if tonumber(ARGV[2]) > 0 and user_block == 0 then " + "source_block = bump(KEYS[3], KEYS[4], tonumber(ARGV[2])) end " + "return {user_block, source_block}" +) + +_COUNTERS: Final = InMemoryCache( + max_size_in_memory=_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS +) +_BLOCKS: Final = InMemoryCache(max_size_in_memory=_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS) +_HELD_ATTEMPTS: Final[dict[str, int]] = {} -def _bounded_store(max_entries: int) -> DualCache: - return DualCache( - in_memory_cache=InMemoryCache(max_size_in_memory=max_entries), - default_in_memory_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, - ) +async def _sleep(seconds: float) -> None: + await asyncio.sleep(seconds) -_FAILED_LOGIN_USERNAME_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_USERNAMES) -_FAILED_LOGIN_SOURCE_CACHE: Final = _bounded_store(_MAX_TRACKED_LOGIN_SOURCES) -_NO_SETTINGS: Final = MappingProxyType({}) -_UNAVAILABLE: Final = object() - -_DELAYS_IN_FLIGHT: Final[dict[str, int]] = {} # mutable-ok: per-source slots taken and released around each held delay +@cache +def _rate_limit_disabled() -> bool: + return get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", default_value=False) is True @cache def warn_login_counters_are_per_worker(num_workers: str) -> None: - """Warn once per process that failed sign-in counters are not shared across workers.""" verbose_proxy_logger.warning( - "Running %s workers but Redis is not configured for LiteLLM caching. " - "Failed Admin UI sign-in attempts are counted per worker, so an attacker " - "gets max_failed_login_attempts guesses per worker instead of overall. " - "Configure Redis via the 'cache' section in your proxy config.", + "Running %s workers but Redis is not configured. Failed Admin UI sign-in attempts are counted " + "per worker, so the effective limits are %s times the configured values. Configure Redis " + "to share one count across workers.", + num_workers, num_workers, ) @cache -def _rate_limit_disabled() -> bool: - """Resolved once per process so an unauthenticated flood never reaches the secret manager.""" - return bool(get_secret_bool("LITELLM_DISABLE_LOGIN_RATE_LIMIT", False)) +def warn_source_login_limit_is_off() -> None: + verbose_proxy_logger.warning( + "%s is not set, so failed Admin UI sign-in attempts are limited per source address and username " + "only. Set it to the address ranges of the proxies in front of LiteLLM to also limit each " + "source address across usernames.", + TRUSTED_PROXY_RANGES_KEY, + ) -async def _sleep(seconds: float) -> None: - """The wait a rejected sign-in is held for. Replaced in tests so the suite pays no wall clock.""" - await asyncio.sleep(seconds) - - -class FailureCounts(NamedTuple): - """Failures recorded so far in this window against each of the two keys.""" - - username: int - source: int - - -def _parse_int_setting(value: object) -> object: - if not isinstance(value, str): - return value - try: - return int(value.strip()) - except ValueError: - return value - - -def _int_setting(name: str, value: object, default: int, minimum: int) -> int: - if value is None: +def _positive_int(raw: object, key: str, default: int) -> int: + if raw is None: return default - parsed: Final = _parse_int_setting(value) - if isinstance(parsed, bool) or not isinstance(parsed, int) or parsed < minimum: + try: + value: Final = int(str(raw)) + except (TypeError, ValueError): + verbose_proxy_logger.warning("Invalid %s value %r; using %s", key, raw, default) + return default + if value < 1: + verbose_proxy_logger.warning("Invalid %s value %s (must be >= 1); using %s", key, value, default) + return default + return value + + +def _int_setting(settings: Mapping[str, object], key: str, default: int) -> int: + return _positive_int(settings.get(key), key, default) + + +def _parse_address(client_ip: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None: + try: + return ipaddress.ip_address(client_ip) + except ValueError: + return None + + +def _parse_network(raw_range: str) -> _Network | None: + try: + return ipaddress.ip_network(raw_range.strip(), strict=False) + except ValueError: verbose_proxy_logger.warning( - "general_settings.%s=%r is not an integer >= %s; using the default of %s", name, value, minimum, default + "Invalid address or range %r in %s; skipping", raw_range, SOURCE_LIMIT_OVERRIDES_KEY + ) + return None + + +def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: + """Failure allowance for this address: the most specific configured range containing it, else the default.""" + default: Final = _int_setting(settings, SOURCE_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE) + raw_overrides: Final = settings.get(SOURCE_LIMIT_OVERRIDES_KEY) + if raw_overrides is None: + return default + try: + overrides: Final = _SOURCE_LIMIT_OVERRIDES.validate_python(raw_overrides) + except ValidationError: + verbose_proxy_logger.warning( + "Invalid %s value; expected a mapping of address or range to limit", SOURCE_LIMIT_OVERRIDES_KEY ) return default - return parsed + address: Final = _parse_address(client_ip) + if address is None: + return default + matches: Final = sorted( + (network.prefixlen, _positive_int(raw_limit, SOURCE_LIMIT_OVERRIDES_KEY, default)) + for raw_range, raw_limit in overrides.items() + if (network := _parse_network(raw_range)) is not None and address in network + ) + return matches[-1][1] if matches else default -def _as_count(cached: object) -> int: - return int(cached) if isinstance(cached, int | float) and not isinstance(cached, bool) else 0 +def source_group(client_ip: str) -> str: + """The bucket an address is counted in: IPv4 as is, IPv6 by its /64, so one prefix holder cannot rotate.""" + address: Final = _parse_address(client_ip) + if address is None: + return client_ip + if isinstance(address, ipaddress.IPv6Address): + mapped: Final = address.ipv4_mapped + if mapped is not None: + return str(mapped) + return str(ipaddress.ip_network((address, IPV6_SOURCE_PREFIX_LENGTH), strict=False)) + return str(address) + + +class _Keys(NamedTuple): + pair_counter: str + pair_block: str + source_counter: str + source_block: str + + +@dataclass(frozen=True, slots=True) +class Block: + scope: Scope + retry_after: int @dataclass(frozen=True, slots=True) class LoginThrottle: - """Fixed-window failed-login accounting for one request's username and source address.""" + """Failed-login limits for one request's source address; ``source_limit`` is None when the + source scope is off because ``trusted_proxy_ranges`` is unset and the peer address is the ingress.""" client_ip: str - max_attempts: int - max_attempts_per_source: int + source_limit: int | None + user_limit: int window_seconds: int - username_cache: DualCache - source_cache: DualCache + block_seconds: int + counters: InMemoryCache + blocks: InMemoryCache redis_cache: RedisCache | None = None enabled: bool = True @classmethod def from_request( - cls, request: Request, general_settings: Mapping[str, object] | None, redis_cache: RedisCache | None - ) -> "LoginThrottle": - """Build the throttle for this request from the proxy's general_settings and shared Redis cache.""" - settings: Final = general_settings or _NO_SETTINGS + cls, + request: Request, + general_settings: Mapping[str, object] | None, + redis_cache: RedisCache | None, + ) -> LoginThrottle: + settings: Final = general_settings if general_settings is not None else _NO_SETTINGS cidrs: Final = normalize_cidr_ranges( settings.get(TRUSTED_PROXY_RANGES_KEY), setting_name=TRUSTED_PROXY_RANGES_KEY ) @@ -145,199 +234,166 @@ class LoginThrottle: ) return cls( client_ip=resolved or _UNKNOWN_SOURCE, - max_attempts=_int_setting( - "max_failed_login_attempts", - settings.get("max_failed_login_attempts"), - DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS, - 1, - ), - max_attempts_per_source=_int_setting( - "max_failed_login_attempts_per_source", - settings.get("max_failed_login_attempts_per_source"), - DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE, - 1, - ), - window_seconds=_int_setting( - "failed_login_window_seconds", - settings.get("failed_login_window_seconds"), - DEFAULT_FAILED_LOGIN_WINDOW_SECONDS, - 1, - ), - username_cache=_FAILED_LOGIN_USERNAME_CACHE, - source_cache=_FAILED_LOGIN_SOURCE_CACHE, + source_limit=_source_limit(settings, resolved) if cidrs and resolved is not None else None, + user_limit=_int_setting(settings, USER_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER), + window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS), + block_seconds=_int_setting(settings, BLOCK_KEY, DEFAULT_FAILED_LOGIN_BLOCK_SECONDS), + counters=_COUNTERS, + blocks=_BLOCKS, redis_cache=redis_cache, enabled=not _rate_limit_disabled(), ) - @staticmethod - def _loggable(username: str) -> str: - """The username with anything that could forge a log line removed.""" - return "".join(c for c in username if c.isprintable())[:_MAX_LOGGED_USERNAME_CHARS] + def _keys(self, username: str) -> _Keys: + group: Final = source_group(self.client_ip) + user: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest() + return _Keys( + pair_counter=f"{_CACHE_KEY_PREFIX}:{{{group}}}:user:{user}", + pair_block=f"{_CACHE_KEY_PREFIX}:{{{group}}}:block:user:{user}", + source_counter=f"{_CACHE_KEY_PREFIX}:{{{group}}}:source", + source_block=f"{_CACHE_KEY_PREFIX}:{{{group}}}:block:source", + ) - @staticmethod - def _username_key(username: str) -> str: - identity: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest() - return f"{_CACHE_KEY_PREFIX}:user:{identity}" - - def _source_key(self) -> str: - return f"{_CACHE_KEY_PREFIX}:source:{self.client_ip}" - - async def _outcome(self, work: Awaitable[object]) -> object: + @asynccontextmanager + async def attempt(self, username: str, *, exempt: bool = False) -> AsyncGenerator[LoginAttempt]: + if not self.enabled or exempt: + yield LoginAttempt(throttle=self, username=username, block=None) + return + keys: Final = self._keys(username) + block: Final = await self._active_block(keys) + if block is None: + yield LoginAttempt(throttle=self, username=username, block=None) + return + slot: Final = keys.pair_block if block.scope == "user" else keys.source_block + held: Final = _HELD_ATTEMPTS.get(slot, 0) + if held >= MAX_HELD_ATTEMPTS_PER_KEY: + verbose_proxy_logger.warning( + "Admin UI sign-in refused: %s attempts already held for a blocked %s; username=%r source=%s", + held, + block.scope, + username, + self.client_ip, + ) + self.refuse(BLOCKED_ATTEMPT_HOLD_SECONDS) + _HELD_ATTEMPTS[slot] = held + 1 try: - return await work - except Exception as exc: # noqa: BLE001 # an unreachable cache must never deny a valid credential - verbose_proxy_logger.warning("login attempt accounting unavailable: %s", exc) - return _UNAVAILABLE + yield LoginAttempt(throttle=self, username=username, block=block) + finally: + remaining: Final = _HELD_ATTEMPTS.get(slot, 1) - 1 + if remaining > 0: + _HELD_ATTEMPTS[slot] = remaining + else: + _HELD_ATTEMPTS.pop(slot, None) - async def _shared(self, work: Callable[[RedisCache], Awaitable[object]]) -> object: - """The Redis result, or ``_UNAVAILABLE`` when Redis is not configured or the call raised.""" - redis_cache: Final = self.redis_cache - if redis_cache is None: - return _UNAVAILABLE - return await self._outcome(work(redis_cache)) + async def _active_block(self, keys: _Keys) -> Block | None: + local: Final = self._local_block_ttls(keys) + shared: Final = await self._shared_block_ttls(keys) + user_ttl: Final = max(local[0], shared[0]) + source_ttl: Final = max(local[1], shared[1]) + if user_ttl > 0: + return Block(scope="user", retry_after=user_ttl) + if self.source_limit is not None and source_ttl > 0: + return Block(scope="source", retry_after=source_ttl) + return None - async def _failures(self, store: DualCache, key: str) -> int: - """The shared count plus this worker's own. - - A failure is written to exactly one of the two: Redis, or this worker's store when Redis - refused it. So the local store is empty while Redis is healthy, and once Redis answers - again the guesses it missed still count. Read through ``async_batch_get_counts`` because - ``async_get_cache`` turns a failed GET into ``None``, which would pass as an empty counter. - """ - local: Final = _as_count(await self._outcome(store.async_get_cache(key=key))) - shared: Final = await self._shared(lambda redis_cache: redis_cache.async_batch_get_counts((key,))) - if not isinstance(shared, tuple): - return local - return _as_count(shared[0]) + local - - async def _remaining_window(self, key: str) -> int: - """Seconds until this counter expires. - - Counters are only ever written together with their expiry, so a counter without one - was stripped out of band (PERSIST, a restore). It is given the full window again, - since nothing increments a key once the limit is reached. - """ + async def _shared_block_ttls(self, keys: _Keys) -> _BlockTtls: if self.redis_cache is None: - return self.window_seconds - ttl: Final = await self._shared(lambda redis_cache: redis_cache.async_get_ttl(key)) - if isinstance(ttl, int) and ttl > 0: - return min(ttl, self.window_seconds) - await self._shared(lambda redis_cache: redis_cache.async_increment_with_floor(key, 0, self.window_seconds)) - return self.window_seconds + return _NOT_BLOCKED + try: + return _LUA_BLOCK_TTLS.validate_python( + await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(list(keys), []) + ) + except Exception as err: + self._warn_redis(err) + return _NOT_BLOCKED - def _refused(self, retry_after: int, param: str) -> ProxyException: - return ProxyException( + def _local_block_ttls(self, keys: _Keys) -> _BlockTtls: + return self._local_block_ttl(keys.pair_block), self._local_block_ttl(keys.source_block) + + def _local_block_ttl(self, block_key: str) -> int: + expires_at: Final = _LOCAL_BLOCK_EXPIRY.validate_python(self.blocks.get_cache(block_key)) + if expires_at is None: + return 0 + return max(math.ceil(expires_at - time.time()), 0) + + async def record_failure(self, username: str) -> _BlockTtls: + keys: Final = self._keys(username) + source_limit: Final = self.source_limit or 0 + if self.redis_cache is not None: + try: + return _LUA_BLOCK_TTLS.validate_python( + await self.redis_cache.async_register_script(_RECORD_FAILURE_LUA)( + list(keys), [self.user_limit, source_limit, self.window_seconds, self.block_seconds] + ) + ) + except Exception as err: + self._warn_redis(err) + user_block: Final = self._local_bump(keys.pair_counter, keys.pair_block, self.user_limit) + if source_limit == 0 or user_block > 0: + return user_block, 0 + return user_block, self._local_bump(keys.source_counter, keys.source_block, source_limit) + + def _local_bump(self, count_key: str, block_key: str, limit: int) -> int: + blocked: Final = self._local_block_ttl(block_key) + if blocked > 0: + return blocked + count: Final = int(self.counters.increment_cache(count_key, 1, ttl=self.window_seconds)) + if count <= limit: + return 0 + self.blocks.set_cache(block_key, time.time() + self.block_seconds, ttl=self.block_seconds) + return self.block_seconds + + async def clear_pair(self, username: str) -> None: + pair_counter: Final = self._keys(username).pair_counter + if self.redis_cache is not None: + try: + await self.redis_cache.async_delete_cache(pair_counter) + except Exception as err: + self._warn_redis(err) + self.counters.delete_cache(pair_counter) + + def _warn_redis(self, err: Exception) -> None: + verbose_proxy_logger.warning( + "Redis failed while counting Admin UI sign-in attempts; using this worker's own counters " + "until it recovers: %s", + err, + ) + + @staticmethod + def refuse(retry_after: int) -> NoReturn: + raise ProxyException( message="Too many failed sign-in attempts. Try again later.", type=ProxyErrorTypes.auth_error, - param=param, - code=429, - headers={"Retry-After": str(retry_after)}, # mutable-ok: ProxyException coerces header values + param="username", + code=status.HTTP_429_TOO_MANY_REQUESTS, + headers={"Retry-After": str(retry_after)}, ) - async def _refuse(self, key: str, scope: str, param: str, username: str, failures: int, limit: int) -> NoReturn: - retry_after: Final = await self._remaining_window(key) + +@dataclass(frozen=True, slots=True) +class LoginAttempt: + throttle: LoginThrottle + username: str + block: Block | None + + async def succeeded(self) -> None: + if not self.throttle.enabled: + return + await self.throttle.clear_pair(self.username) + + async def failed(self) -> None: + if not self.throttle.enabled: + return + if self.block is not None: + await _sleep(BLOCKED_ATTEMPT_HOLD_SECONDS) + self.throttle.refuse(max(self.block.retry_after - BLOCKED_ATTEMPT_HOLD_SECONDS, 1)) + user_block, source_block = await self.throttle.record_failure(self.username) + if user_block == 0 and source_block == 0: + return verbose_proxy_logger.warning( - "Admin UI sign-in attempts exhausted for %s; username=%s source=%s failures=%s limit=%s window=%ss", - scope, - self._loggable(username), - self.client_ip, - failures, - limit, - self.window_seconds, + "Admin UI sign-in blocked for %s seconds after too many failures; scope=%s username=%r source=%s", + user_block or source_block, + "user" if user_block else "source", + self.username, + self.throttle.client_ip, ) - raise self._refused(retry_after, param) - - async def raise_if_blocked(self, username: str) -> None: - """Refuse before the database lookup and before the invite-link password hash.""" - if not self.enabled: - return - username_key: Final = self._username_key(username) - source_key: Final = self._source_key() - username_failures: Final = await self._failures(self.username_cache, username_key) - if username_failures >= self.max_attempts: - await self._refuse( - username_key, "username", "max_failed_login_attempts", username, username_failures, self.max_attempts - ) - source_failures: Final = await self._failures(self.source_cache, source_key) - if source_failures >= self.max_attempts_per_source: - await self._refuse( - source_key, - "source address", - "max_failed_login_attempts_per_source", - username, - source_failures, - self.max_attempts_per_source, - ) - - async def _bump(self, store: DualCache, key: str) -> int: - shared: Final = await self._shared( - lambda redis_cache: redis_cache.async_increment_with_floor(key, 1, self.window_seconds) - ) - if shared is _UNAVAILABLE: - return _as_count( - await self._outcome(store.async_increment_cache(key=key, value=1, ttl=self.window_seconds)) - ) - return _as_count(shared) + _as_count(await self._outcome(store.async_get_cache(key=key))) - - async def record_failure(self, username: str) -> FailureCounts: - """Count one rejected credential guess against this username and against this source.""" - if not self.enabled: - return FailureCounts(username=0, source=0) - return FailureCounts( - username=await self._bump(self.username_cache, self._username_key(username)), - source=await self._bump(self.source_cache, self._source_key()), - ) - - @staticmethod - def delay_seconds(counts: FailureCounts) -> float: - """Seconds to hold a rejected attempt for, doubling per failure past whichever onset is further along.""" - steps: Final = min( - max(counts.username - USERNAME_DELAY_ONSET, counts.source - SOURCE_DELAY_ONSET), - _MAX_DELAY_DOUBLINGS, - ) - if steps < 0: - return 0.0 - return min(FIRST_DELAY_SECONDS * float(2**steps), MAX_DELAY_SECONDS) - - async def delay_for(self, username: str, counts: FailureCounts) -> None: - """Hold this rejected attempt open before answering it, so guessing costs wall-clock time. - - Only ever reached once the credentials are known to be wrong, so a valid password is - never delayed. Sources are capped at ``MAX_CONCURRENT_DELAYS_PER_SOURCE`` held - connections; over that, the attempt is refused immediately instead of parking a socket. - """ - if not self.enabled: - return - delay: Final = self.delay_seconds(counts) - if delay <= 0: - return - in_flight: Final = _DELAYS_IN_FLIGHT.get(self.client_ip, 0) - if in_flight >= MAX_CONCURRENT_DELAYS_PER_SOURCE: - verbose_proxy_logger.warning( - "Admin UI sign-in attempts held concurrently exhausted; username=%s source=%s in_flight=%s", - self._loggable(username), - self.client_ip, - in_flight, - ) - raise self._refused(int(MAX_DELAY_SECONDS), "concurrent_failed_logins") - _DELAYS_IN_FLIGHT[self.client_ip] = in_flight + 1 - try: - await _sleep(delay) - finally: - remaining: Final = _DELAYS_IN_FLIGHT.get(self.client_ip, 1) - 1 - if remaining > 0: - _DELAYS_IN_FLIGHT[self.client_ip] = remaining - else: - _DELAYS_IN_FLIGHT.pop(self.client_ip, None) - - async def clear(self, username: str) -> None: - """Drop the username counter after a successful sign-in. - - The source counter is left alone. It is shared by every account behind that address, - so one success there says nothing about the other attempts it is counting. - """ - if not self.enabled: - return - key: Final = self._username_key(username) - await self._shared(lambda redis_cache: redis_cache.async_delete_cache(key)) - await self._outcome(self.username_cache.async_delete_cache(key=key)) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index b88a7d14ad8..6f52babb255 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -27,7 +27,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.auth_utils import is_sso_provider_fully_configured -from litellm.proxy.auth.login_throttle import LoginThrottle +from litellm.proxy.auth.login_throttle import LoginAttempt, LoginThrottle from litellm.proxy.management_endpoints.internal_user_endpoints import user_update from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, @@ -219,9 +219,21 @@ async def authenticate_user( admin_credentials_match: Final = _admin_credentials_match(username, password, master_key, general_settings) - if not admin_credentials_match: - await throttle.raise_if_blocked(username) + async with throttle.attempt(username, exempt=admin_credentials_match) as attempt: + return await _sign_in( + username, password, master_key, prisma_client, attempt, general_settings, admin_credentials_match + ) + +async def _sign_in( + username: str, + password: str, + master_key: str, + prisma_client: PrismaClient | None, + attempt: LoginAttempt, + general_settings: Mapping[str, object], + admin_credentials_match: bool, +) -> LoginResult: # Check if we can find the `username` in the db. On the UI, users can enter username=their email _user_row: LiteLLM_UserTable | None = None user_role: ( @@ -315,7 +327,7 @@ async def authenticate_user( key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user_info) - await throttle.clear(username) + await attempt.succeeded() return LoginResult( user_id=user_id, @@ -372,7 +384,7 @@ async def authenticate_user( key = response["token"] - await throttle.clear(username) + await attempt.succeeded() return LoginResult( user_id=user_id, @@ -382,7 +394,7 @@ async def authenticate_user( login_method="username_password", ) else: - await throttle.delay_for(username, await throttle.record_failure(username)) + await attempt.failed() raise ProxyException( message=_invalid_credentials_message(general_settings), type=ProxyErrorTypes.auth_error, @@ -390,7 +402,7 @@ async def authenticate_user( code=401, ) else: - await throttle.delay_for(username, await throttle.record_failure(username)) + await attempt.failed() raise ProxyException( message=_invalid_credentials_message(general_settings), type=ProxyErrorTypes.auth_error, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 43013e85e05..5b89beb38d2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -325,7 +325,12 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.fallback_model_access import router_fallback_access_check from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY, LicenseCheck -from litellm.proxy.auth.login_throttle import LoginThrottle, warn_login_counters_are_per_worker +from litellm.proxy.auth.login_throttle import ( + TRUSTED_PROXY_RANGES_KEY, + LoginThrottle, + warn_login_counters_are_per_worker, + warn_source_login_limit_is_off, +) from litellm.proxy.auth.model_checks import ( expand_wildcard_deployments_for_model_info, get_all_fallbacks, @@ -5803,11 +5808,10 @@ class ProxyConfig: if general_settings is None: general_settings = {} - ### FAILED-LOGIN ACCOUNTING MULTI-INSTANCE PREREQUISITE CHECK ### - # Failed Admin UI sign-in counters live in redis_usage_cache when available so a - # brute-force run is counted once across workers instead of once per worker. if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) + if not general_settings.get(TRUSTED_PROXY_RANGES_KEY): + warn_source_login_limit_is_off() _bg_hc_model_groups: Final = parse_background_health_check_model_groups(general_settings) _enable_hc_routing = False diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index c53476f5148..22fcedaa19d 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -25,33 +25,32 @@ class _RecordedSleeps: @pytest.fixture(autouse=True) def login_delays(monkeypatch): - """Replace the failed-login wait, so the suite pays no wall clock and can read it back.""" + """Replace the hold on a blocked wrong password, so the suite pays no wall clock and can read it back.""" from litellm.proxy.auth import login_throttle recorded = _RecordedSleeps() monkeypatch.setattr(login_throttle, "_sleep", recorded) - login_throttle._DELAYS_IN_FLIGHT.clear() + login_throttle._HELD_ATTEMPTS.clear() yield recorded - login_throttle._DELAYS_IN_FLIGHT.clear() + login_throttle._HELD_ATTEMPTS.clear() def _unlimited_throttle(): - """A throttle wired to a real in-memory store with a limit no test can reach.""" - from litellm.caching.dual_cache import DualCache + """A throttle wired to real in-memory stores with limits no test can reach.""" + from litellm.caching.in_memory_cache import InMemoryCache from litellm.proxy.auth.login_throttle import LoginThrottle - store: Final = DualCache() return LoginThrottle( client_ip="1.2.3.4", - max_attempts=10_000, - max_attempts_per_source=10_000, - window_seconds=900, - username_cache=store, - source_cache=store, + source_limit=None, + user_limit=10_000, + window_seconds=60, + block_seconds=300, + counters=InMemoryCache(), + blocks=InMemoryCache(), ) - from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._types import ( LiteLLM_UserTable, @@ -658,29 +657,37 @@ class TestEncodeUiSessionJwt: def _throttle( - max_attempts: int = 3, - window_seconds: int = 900, + user_limit: int = 2, + source_limit: int | None = None, + window_seconds: int = 60, + block_seconds: int = 300, client_ip: str = "1.2.3.4", - cache=None, + stores=None, redis_cache=None, - max_attempts_per_source: int = 10_000, ): - """A throttle over a real in-memory store, so the tests exercise the true counters.""" - from litellm.caching.dual_cache import DualCache + """A throttle over real in-memory stores, so the tests exercise the true counters and blocks.""" + from litellm.caching.in_memory_cache import InMemoryCache from litellm.proxy.auth.login_throttle import LoginThrottle - store: Final = cache if cache is not None else DualCache() + counters, blocks = stores if stores is not None else (InMemoryCache(), InMemoryCache()) return LoginThrottle( client_ip=client_ip, - max_attempts=max_attempts, - max_attempts_per_source=max_attempts_per_source, + source_limit=source_limit, + user_limit=user_limit, window_seconds=window_seconds, - username_cache=store, - source_cache=store, + block_seconds=block_seconds, + counters=counters, + blocks=blocks, redis_cache=redis_cache, ) +def _stores(): + from litellm.caching.in_memory_cache import InMemoryCache + + return InMemoryCache(), InMemoryCache() + + async def _guess(throttle, username: str = "admin", password: str = "wrong"): from litellm.proxy.auth.login_utils import authenticate_user @@ -693,97 +700,376 @@ async def _guess(throttle, username: str = "admin", password: str = "wrong"): ) +async def _fail(throttle, username: str = "admin") -> str: + """One wrong guess; returns the status code it was answered with.""" + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc: + await _guess(throttle, username=username) + return exc.value.code + + +def _known_user(email: str = "known@example.com"): + user = MagicMock() + user.user_id = "u-1" + user.user_email = email + user.user_role = "internal_user" + user.password = "scrypt:stored" + repo = MagicMock() + repo.return_value.table.find_first = AsyncMock(return_value=user) + return repo + + +async def _db_login(throttle, username: str, password: str, *, correct: bool): + """A database user's sign-in with the stored hash faked, so no database or scrypt is needed.""" + from litellm.proxy.auth.login_utils import authenticate_user + + with ( + patch("litellm.proxy.auth.login_utils.UserRepository", _known_user(username)), + patch( # test-quality-ok: reaches the known-DB-user branch without a database + "litellm.proxy.auth.login_utils.verify_password", return_value=correct + ), + patch("litellm.proxy.auth.login_utils._rehash_password_if_needed", new=AsyncMock()), + patch( # test-quality-ok: success mints a UI key; faked so no DB is needed + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ), + ): + return await authenticate_user( + username=username, password=password, master_key="sk-master", prisma_client=MagicMock(), throttle=throttle + ) + + +def _local_count(throttle, key: str) -> int: + return int(throttle.counters.get_cache(key) or 0) + + @pytest.mark.asyncio -async def test_attempts_are_refused_once_the_limit_is_reached(monkeypatch): - """The limit denies further attempts for the window, and the denial carries Retry-After.""" +async def test_too_many_failures_for_one_username_block_that_pair_and_carry_retry_after(monkeypatch): + """One failure past the pair limit blocks the source for that username; the next wrong guess is held + and answered 429 with the block's remaining time, and the counter is not touched by blocked guesses.""" from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=3, window_seconds=77) + throttle = _throttle(user_limit=2, block_seconds=77) + keys = throttle._keys("admin") - for _ in range(3): - with pytest.raises(ProxyException) as first: - await _guess(throttle) - assert first.value.code == "401" + assert [await _fail(throttle) for _ in range(3)] == ["401", "401", "401"], "the limit itself is a plain 401" + assert throttle._local_block_ttl(keys.pair_block) == 77 with pytest.raises(ProxyException) as blocked: await _guess(throttle) assert blocked.value.code == "429" - assert blocked.value.headers.get("Retry-After") == "77" + assert blocked.value.headers.get("Retry-After") == "47", "the 30s hold is taken off the remaining block" + assert _local_count(throttle, keys.pair_counter) == 3, "a blocked guess is not counted again" @pytest.mark.asyncio -async def test_a_correct_admin_password_is_accepted_while_blocked(monkeypatch): - """The configured admin credentials are compared before the gate, so the operator gets in. - - A throttle that refuses a valid password hands anyone who can reach the login form a - denial of service against the one account that can fix it. - """ - from litellm.proxy._types import ProxyException +async def test_a_wrong_password_from_a_blocked_key_is_held_before_it_is_refused(monkeypatch, login_delays): + """The hold is the rate cap: a blocked key gets one verified guess per held slot per 30 seconds.""" + from litellm.proxy.auth.login_throttle import BLOCKED_ATTEMPT_HOLD_SECONDS + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=1) + + assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] + assert login_delays.seconds == [], "an unblocked wrong password is answered at once" + + assert await _fail(throttle) == "429" + assert login_delays.seconds == [BLOCKED_ATTEMPT_HOLD_SECONDS] + + +@pytest.mark.asyncio +async def test_a_correct_password_signs_in_while_its_pair_is_blocked(monkeypatch): + """The block is soft: the real user is still verified and gets in, so nobody can be locked out by + guessing at their account.""" monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - throttle = _throttle(max_attempts=2) + throttle = _throttle(user_limit=1) - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(throttle) + assert [await _fail(throttle, username="user@corp.com") for _ in range(3)] == ["401", "401", "429"] - with pytest.raises(ProxyException) as still_blocked: - await _guess(throttle) - assert still_blocked.value.code == "429", "a wrong password is still refused" - - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed - "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) - ): - result = await _guess(throttle, password="right") + result = await _db_login(throttle, "user@corp.com", "right", correct=True) assert result.key == "sk-ui" @pytest.mark.asyncio -async def test_a_blocked_attempt_does_not_extend_the_window(monkeypatch): - """Hammering while blocked must not push the counter or refresh its TTL.""" - from litellm.proxy._types import ProxyException - +async def test_a_correct_password_signs_in_while_its_source_is_blocked(monkeypatch): + """Same for the source-wide block: it slows guessing from that address, it does not refuse a user.""" monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=2) - key = throttle._username_key("admin") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(user_limit=100, source_limit=2) - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(throttle) - counted_at_limit = await throttle._failures(throttle.username_cache, key) + for i in range(3): + assert await _fail(throttle, username=f"other-{i}@corp.com") == "401" + assert await _fail(throttle, username="other-9@corp.com") == "429", "the source is blocked for everyone" - for _ in range(5): - with pytest.raises(ProxyException): - await _guess(throttle) - - assert await throttle._failures(throttle.username_cache, key) == counted_at_limit == 2 + result = await _db_login(throttle, "user@corp.com", "right", correct=True) + assert result.key == "sk-ui" @pytest.mark.asyncio -async def test_a_successful_sign_in_clears_the_bucket(monkeypatch): - """Success resets the budget rather than leaving the operator near the limit.""" - from litellm.proxy._types import ProxyException +async def test_a_successful_sign_in_clears_the_pair_counter_but_not_the_source_counter(monkeypatch): + """One account's success says nothing about the other guesses the address is making.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(user_limit=5, source_limit=50) + keys = throttle._keys("user@corp.com") + + for _ in range(2): + assert await _fail(throttle, username="user@corp.com") == "401" + assert _local_count(throttle, keys.pair_counter) == 2 + assert _local_count(throttle, keys.source_counter) == 2 + + await _db_login(throttle, "user@corp.com", "right", correct=True) + + assert _local_count(throttle, keys.pair_counter) == 0 + assert _local_count(throttle, keys.source_counter) == 2 + + +@pytest.mark.asyncio +async def test_once_a_pair_is_blocked_its_failures_stop_counting_against_the_source(monkeypatch): + """A script stuck on one account trips the pair block and then leaves the office's shared address alone.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=2, source_limit=4) + keys = throttle._keys("stuck-script@corp.com") + + assert [await _fail(throttle, username="stuck-script@corp.com") for _ in range(3)] == ["401"] * 3 + assert _local_count(throttle, keys.source_counter) == 2, "failures before the pair block count for the source" + + for _ in range(5): + assert await _fail(throttle, username="stuck-script@corp.com") == "429" + assert _local_count(throttle, keys.source_counter) == 2, "blocked-pair failures must not reach the source" + + assert await _fail(throttle, username="colleague@corp.com") == "401", "a colleague still signs in normally" + assert throttle._local_block_ttl(keys.source_block) == 0 + + +@pytest.mark.asyncio +async def test_the_blocking_failure_itself_does_not_count_against_the_source(monkeypatch): + """The guess that installs the pair block is the first one that stops counting, so a pair limit of B + costs the source exactly B, not B plus one.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=2, source_limit=2) + keys = throttle._keys("stuck@corp.com") + + assert [await _fail(throttle, username="stuck@corp.com") for _ in range(3)] == ["401", "401", "401"] + + assert _local_count(throttle, keys.source_counter) == 2 + assert throttle._local_block_ttl(keys.source_block) == 0, "the third guess blocked the pair, not the source" + + +@pytest.mark.asyncio +async def test_too_many_failures_across_usernames_block_the_whole_source(monkeypatch): + """A spray of one guess per username never trips a pair; the source counter is what stops it.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=5, source_limit=3, block_seconds=200) + + assert [await _fail(throttle, username=f"sprayed-{i}@corp.com") for i in range(4)] == ["401"] * 4 + + assert await _fail(throttle, username="sprayed-99@corp.com") == "429" + assert throttle._local_block_ttl(throttle._keys("x").source_block) == 200 + + +@pytest.mark.asyncio +async def test_without_trusted_proxy_ranges_the_source_scope_is_off(monkeypatch): + """Behind an ingress every client shares the peer address, so a source-wide block would block them all. + The pair scope still applies.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + request = MagicMock() + request.headers = {"x-forwarded-for": "203.0.113.9"} + request.client = MagicMock() + request.client.host = "10.0.0.1" + throttle = LoginThrottle.from_request( + request, general_settings={"max_failed_login_attempts_per_source": 1}, redis_cache=None + ) + + assert throttle.source_limit is None + assert throttle.client_ip == "10.0.0.1", "the header is not trusted without a configured proxy range" + assert [await _fail(throttle, username=f"user-{i}@corp.com") for i in range(6)] == ["401"] * 6 + + +@pytest.mark.asyncio +async def test_with_trusted_proxy_ranges_the_source_is_the_forwarded_client(monkeypatch): + """The header is walked right to left past the trusted hops, so a forged left-most entry cannot pick the bucket.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + settings = {"trusted_proxy_ranges": ["10.0.0.0/8"], "max_failed_login_attempts_per_source": 2} + + def _from(peer: str, forwarded: str): + request = MagicMock() + request.headers = {"x-forwarded-for": forwarded} + request.client = MagicMock() + request.client.host = peer + return LoginThrottle.from_request(request, general_settings=settings, redis_cache=None) + + via_proxy = _from("10.0.0.1", "1.1.1.1, 203.0.113.9, 10.0.0.2") + assert via_proxy.client_ip == "203.0.113.9" + assert via_proxy.source_limit == 2 + + direct = _from("198.51.100.7", "203.0.113.9") + assert direct.client_ip == "198.51.100.7", "a peer outside the trusted ranges cannot forward anything" + + +def test_source_overrides_pick_the_most_specific_matching_range(): + """An exact address beats a /16 beats a /8; an address in none of them keeps the default.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + settings = { + "trusted_proxy_ranges": ["10.0.0.0/8"], + "max_failed_login_attempts_per_source": 7, + "max_failed_login_attempts_per_source_overrides": { + "203.0.0.0/8": 100, + "203.0.113.0/24": 200, + "203.0.113.9": 300, + "not-an-address": 999, + "198.51.100.0/24": "not-a-number", + }, + } + + def _limit(client: str) -> int | None: + request = MagicMock() + request.headers = {"x-forwarded-for": client} + request.client = MagicMock() + request.client.host = "10.0.0.1" + return LoginThrottle.from_request(request, general_settings=settings, redis_cache=None).source_limit + + assert _limit("203.0.113.9") == 300 + assert _limit("203.0.113.10") == 200 + assert _limit("203.0.1.1") == 100 + assert _limit("192.0.2.1") == 7 + assert _limit("198.51.100.1") == 7, "a garbage limit falls back to the default rather than a huge or zero budget" + + +def test_ipv6_sources_are_grouped_by_their_64_bit_prefix(): + """A /64 holder has 2^64 addresses; counting each one separately would hand them unlimited fresh buckets.""" + from litellm.proxy.auth.login_throttle import source_group + + assert source_group("2001:db8:1:2::1") == source_group("2001:db8:1:2:ffff:ffff:ffff:ffff") == "2001:db8:1:2::/64" + assert source_group("2001:db8:1:3::1") != source_group("2001:db8:1:2::1") + assert source_group("::ffff:203.0.113.9") == source_group("203.0.113.9") == "203.0.113.9" + assert source_group("unknown") == "unknown" + + +@pytest.mark.asyncio +async def test_two_ipv6_addresses_in_one_64_share_the_source_budget(monkeypatch): + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + stores = _stores() + first = _throttle(user_limit=50, source_limit=2, client_ip="2001:db8:1:2::1", stores=stores) + second = _throttle(user_limit=50, source_limit=2, client_ip="2001:db8:1:2::2", stores=stores) + + assert [await _fail(first, username=f"a-{i}@corp.com") for i in range(3)] == ["401"] * 3 + assert await _fail(second, username="b@corp.com") == "429" + + +@pytest.mark.asyncio +async def test_one_source_being_blocked_does_not_touch_another(monkeypatch): + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + stores = _stores() + attacker = _throttle(user_limit=50, source_limit=2, client_ip="203.0.113.9", stores=stores) + neighbour = _throttle(user_limit=50, source_limit=2, client_ip="198.51.100.7", stores=stores) + + assert [await _fail(attacker, username=f"t-{i}@corp.com") for i in range(3)] == ["401"] * 3 + assert await _fail(attacker, username="t-9@corp.com") == "429" + assert await _fail(neighbour, username="t-9@corp.com") == "401" + + +@pytest.mark.asyncio +async def test_the_same_username_from_another_source_has_its_own_budget(monkeypatch): + """The pair carries the address on purpose: an attacker elsewhere cannot lock a user out of their own office.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + stores = _stores() + attacker = _throttle(user_limit=1, client_ip="203.0.113.9", stores=stores) + office = _throttle(user_limit=1, client_ip="198.51.100.7", stores=stores) + + assert [await _fail(attacker, username="victim@corp.com") for _ in range(3)] == ["401", "401", "429"] + assert await _fail(office, username="victim@corp.com") == "401" + + +@pytest.mark.asyncio +async def test_the_counting_window_is_anchored_at_the_first_failure(monkeypatch): + """Later failures must not push the expiry out, or a slow guesser keeps their own count alive forever.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=50, window_seconds=60) + key = throttle._keys("admin").pair_counter + + await _fail(throttle) + first_expiry = throttle.counters.ttl_dict[key] + for _ in range(3): + await _fail(throttle) + + assert throttle.counters.ttl_dict[key] == first_expiry + + +@pytest.mark.asyncio +async def test_the_block_outlives_the_counting_window(monkeypatch): + """Counters expire after the window and blocks after the block time; the two are separate keys.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=1, window_seconds=10, block_seconds=300) + keys = throttle._keys("admin") + + assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] + + throttle.counters.delete_cache(keys.pair_counter) + + assert await _fail(throttle) == "429", "an expired counter must not lift an active block" + assert 290 <= throttle._local_block_ttl(keys.pair_block) <= 300 + + +@pytest.mark.asyncio +async def test_the_block_time_is_fixed_and_not_refreshed_by_blocked_guesses(monkeypatch): + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + throttle = _throttle(user_limit=1, block_seconds=300) + key = throttle._keys("admin").pair_block + + assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] + installed_at = throttle.blocks.ttl_dict[key] + + for _ in range(4): + assert await _fail(throttle) == "429" + + assert throttle.blocks.ttl_dict[key] == installed_at + + +@pytest.mark.asyncio +async def test_the_configured_admin_credentials_bypass_the_throttle(monkeypatch): + """The only account that can fix a misconfiguration is exempt: no hold, no slot, even while blocked.""" + from litellm.proxy.auth import login_throttle as lt monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - throttle = _throttle(max_attempts=3) + throttle = _throttle(user_limit=1) - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(throttle) + assert [await _fail(throttle) for _ in range(3)] == ["401", "401", "429"] - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed - "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + with ( + patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), + patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ), ): - await _guess(throttle, password="right") - - assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 + result = await _guess(throttle, password="right") + assert result.key == "sk-ui" + assert lt._HELD_ATTEMPTS == {} @pytest.mark.asyncio @@ -791,7 +1077,7 @@ async def test_a_configuration_error_never_counts(monkeypatch): """A 500 from an unset master key is not a guess and must not consume the budget.""" from litellm.proxy._types import ProxyException - throttle = _throttle(max_attempts=2) + throttle = _throttle(user_limit=2) for _ in range(5): with pytest.raises(ProxyException) as exc: await authenticate_user( @@ -799,44 +1085,20 @@ async def test_a_configuration_error_never_counts(monkeypatch): ) assert exc.value.code == "500" - assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 + assert _local_count(throttle, throttle._keys("admin").pair_counter) == 0 @pytest.mark.asyncio -async def test_the_username_is_case_folded_into_one_bucket(monkeypatch): +async def test_the_username_is_case_folded_into_one_pair(monkeypatch): """The DB lookup is case-insensitive, so casing must not multiply the budget.""" - from litellm.proxy._types import ProxyException - monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=4) + throttle = _throttle(user_limit=4) - for name in ("admin@corp.com", "ADMIN@corp.com", "Admin@corp.com", "aDmIn@corp.com"): - with pytest.raises(ProxyException) as exc: - await _guess(throttle, username=name) - assert exc.value.code == "401" + for name in ("admin@corp.com", "ADMIN@corp.com", "Admin@corp.com", "aDmIn@corp.com", "admin@CORP.com"): + assert await _fail(throttle, username=name) == "401" - with pytest.raises(ProxyException) as blocked: - await _guess(throttle, username="admin@CORP.com") - assert blocked.value.code == "429" - - -@pytest.mark.asyncio -async def test_a_different_username_from_the_same_source_is_unaffected(monkeypatch): - """The counters are independent, so one username's failures do not exhaust another's.""" - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=2) - - for _ in range(3): - with pytest.raises(ProxyException): - await _guess(throttle, username="admin") - - with pytest.raises(ProxyException) as other: - await _guess(throttle, username="someone-else@example.com") - assert other.value.code == "401", "a second username must still reach the credential check" + assert await _fail(throttle, username="admin@Corp.com") == "429" @pytest.mark.asyncio @@ -848,26 +1110,9 @@ async def test_both_credential_rejections_are_indistinguishable(monkeypatch): monkeypatch.setenv("UI_PASSWORD", "right") with pytest.raises(ProxyException) as unknown: - await _guess(_throttle(max_attempts=99), username="nobody@example.com") - - fake_user = MagicMock() - fake_user.user_id = "u-1" - fake_user.user_email = "known@example.com" - fake_user.user_role = "internal_user" - fake_user.password = "scrypt:fake" - repo = MagicMock() - repo.return_value.table.find_first = AsyncMock(return_value=fake_user) - with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( # test-quality-ok: reaches the known-DB-user branch without a database - "litellm.proxy.auth.login_utils.verify_password", return_value=False - ): - with pytest.raises(ProxyException) as known: - await authenticate_user( - username="known@example.com", - password="wrong", - master_key="sk-master", - prisma_client=MagicMock(), - throttle=_throttle(max_attempts=99), - ) + await _guess(_throttle(user_limit=99), username="nobody@example.com") + with pytest.raises(ProxyException) as known: + await _db_login(_throttle(user_limit=99), "known@example.com", "wrong", correct=False) assert unknown.value.message == known.value.message assert "known@example.com" not in unknown.value.message + known.value.message @@ -875,13 +1120,12 @@ async def test_both_credential_rejections_are_indistinguishable(monkeypatch): @pytest.mark.asyncio async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypatch): - """That 401 is deterministic and guards no secret, so counting it would only let - someone burn a passwordless account's bucket.""" + """That 401 is deterministic and guards no secret, so counting it would only let someone burn the pair.""" from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=2) + throttle = _throttle(user_limit=2) passwordless = MagicMock() passwordless.user_id = "u-2" @@ -891,7 +1135,9 @@ async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypat repo = MagicMock() repo.return_value.table.find_first = AsyncMock(return_value=passwordless) - with patch("litellm.proxy.auth.login_utils.UserRepository", repo): # test-quality-ok: reaches the passwordless-DB-user branch without a database + with patch( + "litellm.proxy.auth.login_utils.UserRepository", repo + ): # test-quality-ok: reaches the passwordless-DB-user branch without a database for _ in range(5): with pytest.raises(ProxyException) as exc: await authenticate_user( @@ -903,185 +1149,36 @@ async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypat ) assert exc.value.code == "401" - assert await throttle._failures(throttle.username_cache, throttle._username_key("nopass@example.com")) == 0 + assert _local_count(throttle, throttle._keys("nopass@example.com").pair_counter) == 0 @pytest.mark.asyncio async def test_a_wrong_password_for_a_known_user_also_counts(monkeypatch): - """The database-user branch must charge the bucket too, not just the unknown-user branch.""" + """The database-user branch must charge the pair too, not just the unknown-user branch.""" from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=3) + throttle = _throttle(user_limit=2) - known = MagicMock() - known.user_id = "u-1" - known.user_email = "known@example.com" - known.user_role = "internal_user" - known.password = "scrypt:stored" - repo = MagicMock() - repo.return_value.table.find_first = AsyncMock(return_value=known) - - async def _attempt(): - return await authenticate_user( - username="known@example.com", - password="wrong", - master_key="sk-master", - prisma_client=MagicMock(), - throttle=throttle, - ) - - with patch("litellm.proxy.auth.login_utils.UserRepository", repo), patch( # test-quality-ok: reaches the known-DB-user branch without a database - "litellm.proxy.auth.login_utils.verify_password", return_value=False - ): - for _ in range(3): - with pytest.raises(ProxyException) as rejected: - await _attempt() - assert rejected.value.code == "401" - - with pytest.raises(ProxyException) as blocked: - await _attempt() - assert blocked.value.code == "429" - - -@pytest.mark.asyncio -async def test_one_source_exhausting_its_own_budget_does_not_refuse_another_source(monkeypatch): - """The source counter is per address, so a noisy office does not take its neighbour down.""" - from litellm.caching.dual_cache import DualCache - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - shared_store = DualCache() - attacker = _throttle( - max_attempts=10_000, max_attempts_per_source=2, client_ip="203.0.113.9", cache=shared_store - ) - operator = _throttle( - max_attempts=10_000, max_attempts_per_source=2, client_ip="198.51.100.7", cache=shared_store - ) - - for i in range(2): - with pytest.raises(ProxyException): - await _guess(attacker, username=f"target-{i}@corp.com") - - with pytest.raises(ProxyException) as blocked: - await _guess(attacker, username="target-2@corp.com") - assert blocked.value.code == "429" - - with pytest.raises(ProxyException) as unaffected: - await _guess(operator, username="target-3@corp.com") - assert unaffected.value.code == "401", "the other address must still reach the credential check" - - -@pytest.mark.asyncio -async def test_a_username_exhausted_from_one_source_is_refused_from_another(monkeypatch): - """The username counter carries no address, so spreading the guesses buys nothing. - - The pair key this replaced reset the budget for every new address, which is exactly the - shape of a credential-stuffing run from a proxy pool. - """ - from litellm.caching.dual_cache import DualCache - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - shared_store = DualCache() - first_hop = _throttle(max_attempts=2, client_ip="203.0.113.9", cache=shared_store) - second_hop = _throttle(max_attempts=2, client_ip="198.51.100.7", cache=shared_store) - - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(first_hop, username="victim@corp.com") - - with pytest.raises(ProxyException) as rotated: - await _guess(second_hop, username="victim@corp.com") - assert rotated.value.code == "429" - - -@pytest.mark.asyncio -async def test_a_source_wide_spray_is_counted_even_though_each_username_is_fresh(monkeypatch): - """One guess against each of many usernames never trips a username counter, only the source one.""" - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=10_000, max_attempts_per_source=6, client_ip="203.0.113.11") - - for i in range(6): + for _ in range(3): with pytest.raises(ProxyException) as rejected: - await _guess(throttle, username=f"sprayed-{i}@corp.com") + await _db_login(throttle, "known@example.com", "wrong", correct=False) assert rejected.value.code == "401" with pytest.raises(ProxyException) as blocked: - await _guess(throttle, username="sprayed-7@corp.com") + await _db_login(throttle, "known@example.com", "wrong", correct=False) assert blocked.value.code == "429" - assert await throttle._failures(throttle.username_cache, throttle._username_key("sprayed-7@corp.com")) == 0 @pytest.mark.asyncio -async def test_a_successful_sign_in_leaves_the_source_counter_alone(monkeypatch): - """One account's success says nothing about the other attempts the address is making.""" - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - throttle = _throttle(max_attempts=10) - - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(throttle) - - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed - "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) - ): - await _guess(throttle, password="right") - - assert await throttle._failures(throttle.username_cache, throttle._username_key("admin")) == 0 - assert await throttle._failures(throttle.source_cache, throttle._source_key()) == 2 - - -@pytest.mark.asyncio -async def test_the_delay_doubles_from_one_second_and_is_capped(monkeypatch, login_delays): - """Guessing has to cost wall clock, and the cost has to stop short of an unbounded hang.""" - from litellm.proxy._types import ProxyException - from litellm.proxy.auth.login_throttle import MAX_DELAY_SECONDS - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=10_000) - - for _ in range(9): - with pytest.raises(ProxyException): - await _guess(throttle) - - assert login_delays.seconds == [1.0, 2.0, 4.0, 8.0, 16.0, 30.0, 30.0], ( - "the first two failures answer immediately, then the wait doubles up to the cap" - ) - assert max(login_delays.seconds) == MAX_DELAY_SECONDS - - -@pytest.mark.asyncio -async def test_the_delay_tracks_whichever_counter_is_further_past_its_onset(monkeypatch): - """A source deep into a spray must not be answered instantly just because the username is fresh.""" - from litellm.proxy.auth.login_throttle import FailureCounts, LoginThrottle - - assert LoginThrottle.delay_seconds(FailureCounts(username=1, source=1)) == 0.0 - assert LoginThrottle.delay_seconds(FailureCounts(username=2, source=24)) == 0.0 - assert LoginThrottle.delay_seconds(FailureCounts(username=3, source=1)) == 1.0 - assert LoginThrottle.delay_seconds(FailureCounts(username=1, source=25)) == 1.0 - assert LoginThrottle.delay_seconds(FailureCounts(username=4, source=28)) == 8.0 - - -@pytest.mark.asyncio -async def test_held_attempts_from_one_source_are_capped(monkeypatch): - """Holding a rejected attempt open must not let one address park unlimited sockets.""" +async def test_held_attempts_from_one_blocked_key_are_capped(monkeypatch): + """Holding a wrong guess open must not let one blocked key park unlimited sockets in password checks.""" import asyncio from litellm.proxy._types import ProxyException from litellm.proxy.auth import login_throttle as lt - from litellm.proxy.auth.login_throttle import MAX_CONCURRENT_DELAYS_PER_SOURCE + from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") @@ -1091,134 +1188,102 @@ async def test_held_attempts_from_one_source_are_capped(monkeypatch): await release.wait() monkeypatch.setattr(lt, "_sleep", _park) - throttle = _throttle(max_attempts=10_000, client_ip="203.0.113.44") - await throttle.record_failure("admin") - await throttle.record_failure("admin") + throttle = _throttle(user_limit=1, client_ip="203.0.113.44") + slot = throttle._keys("admin").pair_block + assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] - held = [asyncio.create_task(_guess(throttle)) for _ in range(MAX_CONCURRENT_DELAYS_PER_SOURCE)] + held = [asyncio.create_task(_guess(throttle)) for _ in range(MAX_HELD_ATTEMPTS_PER_KEY)] for _ in range(1000): - if lt._DELAYS_IN_FLIGHT.get("203.0.113.44") == MAX_CONCURRENT_DELAYS_PER_SOURCE: + if lt._HELD_ATTEMPTS.get(slot) == MAX_HELD_ATTEMPTS_PER_KEY: break await asyncio.sleep(0) - assert lt._DELAYS_IN_FLIGHT.get("203.0.113.44") == MAX_CONCURRENT_DELAYS_PER_SOURCE + assert lt._HELD_ATTEMPTS.get(slot) == MAX_HELD_ATTEMPTS_PER_KEY try: with pytest.raises(ProxyException) as over_cap: await _guess(throttle) assert over_cap.value.code == "429" assert over_cap.value.headers.get("Retry-After") == "30" + assert await _fail(throttle, username="someone-else@corp.com") == "401", "other keys are not affected" finally: release.set() for task in held: with pytest.raises(ProxyException): await task - with pytest.raises(ProxyException) as after_drain: - await _guess(throttle) - assert after_drain.value.code == "401", "the cap must release once the held attempts answer" + assert lt._HELD_ATTEMPTS.get(slot) is None, "the slots are released once the held attempts answer" @pytest.mark.asyncio -async def test_disabling_the_control_removes_the_delay_as_well(monkeypatch, login_delays): +async def test_disabling_the_control_removes_the_hold_as_well(monkeypatch, login_delays): """The escape hatch has to turn off the whole control, not only the refusal.""" import dataclasses - from litellm.proxy._types import ProxyException - monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = dataclasses.replace(_throttle(max_attempts=2), enabled=False) - - for _ in range(6): - with pytest.raises(ProxyException) as rejected: - await _guess(throttle) - assert rejected.value.code == "401" + throttle = dataclasses.replace(_throttle(user_limit=1), enabled=False) + assert [await _fail(throttle) for _ in range(6)] == ["401"] * 6 assert login_delays.seconds == [] class _FakeRedis: - """Redis whose only counter write is the atomic INCRBY-plus-EXPIRE Lua call. + """Redis whose only writes are the throttle's two scripts, run atomically as one call each. - `async_increment` is deliberately absent: a two-step increment would fail the test - with AttributeError, because Redis could then commit a count without its expiry. + Mirrors the Lua: a blocked key returns its remaining block time and is not counted; a counter + is expired on first write; one over the limit installs the block; a blocked pair stops the + source from being counted. The real scripts are exercised against a live Redis in the PR's + proof, this fake only has to be faithful enough for the worker-sharing tests. """ def __init__(self): self.values: dict = {} self.ttls: dict = {} + self.scripts: list[str] = [] - async def async_get_cache(self, key, **kwargs): - return self.values.get(key) + def async_register_script(self, script: str): + from litellm.proxy.auth import login_throttle as lt - async def async_batch_get_counts(self, key_list): - return tuple(self.values.get(key) for key in key_list) + async def _run(keys, args): + self.scripts.append(script) + if script == lt._BLOCK_TTLS_LUA: + return [self._ttl(keys[1]), self._ttl(keys[3])] + assert script == lt._RECORD_FAILURE_LUA + user_limit, source_limit, window, block = (int(a) for a in args) + user_block = self._bump(keys[0], keys[1], user_limit, window, block) + if source_limit > 0 and user_block == 0: + return [user_block, self._bump(keys[2], keys[3], source_limit, window, block)] + return [user_block, 0] - async def async_increment_with_floor(self, key, value, ttl): - self.values[key] = self.values.get(key, 0) + value - self.ttls.setdefault(key, ttl) - return self.values[key] + return _run - async def async_get_ttl(self, key): - return self.ttls.get(key) + def _ttl(self, key: str) -> int: + return self.ttls.get(key, -2) if key in self.values else -2 + + def _bump(self, count_key: str, block_key: str, limit: int, window: int, block: int) -> int: + if self._ttl(block_key) > 0: + return self._ttl(block_key) + self.values[count_key] = self.values.get(count_key, 0) + 1 + self.ttls.setdefault(count_key, window) + if self.values[count_key] > limit: + self.values[block_key] = 1 + self.ttls[block_key] = block + return block + return 0 async def async_delete_cache(self, key): self.values.pop(key, None) self.ttls.pop(key, None) - def persist(self): - self.ttls.clear() - - -@pytest.mark.asyncio -async def test_counters_are_written_with_their_expiry_and_re_armed_if_stripped(monkeypatch): - """Regression: a counter with no TTL would refuse the pair forever. - - Nothing increments a key once the limit is reached, so a counter that ever exists - without an expiry stays refused with no way back. Every write must therefore carry the - expiry, and a refusal that finds it stripped (PERSIST) must put the window back. - """ - from litellm.proxy._types import ProxyException - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - redis = _FakeRedis() - throttle = _throttle(max_attempts=2, window_seconds=77, redis_cache=redis) - - for _ in range(2): - with pytest.raises(ProxyException): - await _guess(throttle) - - assert redis.values, "failures must land in the shared counter" - assert set(redis.ttls) == set(redis.values), "no counter may exist without its expiry" - assert set(redis.ttls.values()) == {77} - - redis.persist() - with pytest.raises(ProxyException) as blocked: - await _guess(throttle) - assert blocked.value.code == "429" - assert blocked.value.headers.get("Retry-After") == "77" - assert set(redis.ttls) >= {k for k in redis.values if ":user:" in k}, "the refusal must re-arm a stripped expiry" - class _DownRedis(_FakeRedis): - """Redis whose every call fails, as during an outage or an open circuit breaker. + """Redis whose every call fails, as during an outage or an open circuit breaker.""" - `async_get_cache` returns None rather than raising, as the real one does: it swallows the - error, so a failed GET is indistinguishable from an empty key to anyone reading through it. - """ + def async_register_script(self, script: str): + async def _run(keys, args): + raise ConnectionError("redis is down") - async def async_get_cache(self, key, **kwargs): - return None - - async def async_batch_get_counts(self, key_list): - raise ConnectionError("redis is down") - - async def async_increment_with_floor(self, key, value, ttl): - raise ConnectionError("redis is down") - - async def async_get_ttl(self, key): - raise ConnectionError("redis is down") + return _run async def async_delete_cache(self, key): raise ConnectionError("redis is down") @@ -1226,124 +1291,84 @@ class _DownRedis(_FakeRedis): @pytest.mark.asyncio async def test_redis_is_the_only_counter_while_it_answers(monkeypatch): - """Regression: every worker must spend the same budget, and a success must clear it for all. - - Counting in this worker's memory as well as in Redis let the two drift apart: a worker - whose Redis write failed kept its own count while the others gave the attacker fresh - guesses, and a stale local count outlived the shared clear after a correct password. - """ - from litellm.caching.dual_cache import DualCache - from litellm.proxy._types import ProxyException + """Every worker must spend the same budget, see the same block, and a success must clear the pair for all.""" from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") redis = _FakeRedis() - first_worker_store = DualCache() - second_worker_store = DualCache() - first_worker = _throttle(max_attempts=2, cache=first_worker_store, redis_cache=redis) - second_worker = _throttle(max_attempts=2, cache=second_worker_store, redis_cache=redis) + first_worker = _throttle(user_limit=2, stores=_stores(), redis_cache=redis) + second_worker = _throttle(user_limit=2, stores=_stores(), redis_cache=redis) - for _ in range(2): - with pytest.raises(ProxyException, match="Invalid credentials"): - await _guess(first_worker) - - assert not [k for k in first_worker_store.in_memory_cache.cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)], ( + assert [await _fail(first_worker, username="user@corp.com") for _ in range(3)] == ["401"] * 3 + assert not [k for k in first_worker.counters.cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)], ( "with Redis answering, no worker may keep a counter of its own" ) - with pytest.raises(ProxyException) as blocked: - await _guess(second_worker) - assert blocked.value.code == "429", "the second worker must see the budget the first one spent" + assert not first_worker.blocks.cache_dict - monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed - "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) - ): - await _guess(second_worker, password="right") + assert await _fail(second_worker, username="user@corp.com") == "429", "the second worker sees the block" - assert not [k for k in redis.values if ":user:" in k], "a success must clear the shared username counter" - with pytest.raises(ProxyException, match="Invalid credentials"): - await _guess(first_worker) + await _db_login(second_worker, "user@corp.com", "right", correct=True) + + assert not [k for k in redis.values if ":user:" in k and ":block:" not in k], ( + "success clears the shared pair counter" + ) + assert [k for k in redis.values if ":block:user:" in k], "an active block is not lifted by one success" @pytest.mark.asyncio async def test_a_redis_outage_falls_back_to_this_workers_own_counter(monkeypatch): - """With Redis raising, guesses are still counted and refused, per worker, instead of unbounded.""" + """With Redis raising, guesses are still counted and blocked per worker, with a warning, instead of unbounded.""" + import logging + + from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=2, redis_cache=_DownRedis()) + throttle = _throttle(user_limit=2, block_seconds=300, redis_cache=_DownRedis()) - for _ in range(2): - with pytest.raises(ProxyException, match="Invalid credentials"): + records: list[logging.LogRecord] = [] + handler = logging.Handler() + handler.emit = records.append + verbose_proxy_logger.addHandler(handler) + try: + assert [await _fail(throttle) for _ in range(3)] == ["401"] * 3 + with pytest.raises(ProxyException) as blocked: await _guess(throttle) + finally: + verbose_proxy_logger.removeHandler(handler) - with pytest.raises(ProxyException) as blocked: - await _guess(throttle) assert blocked.value.code == "429" - assert blocked.value.headers.get("Retry-After") == "900" - - -class _WriteRefusingRedis(_FakeRedis): - """Redis that answers reads but raises on writes until `recover()` is called.""" - - def __init__(self): - super().__init__() - self.writable = False - - def recover(self): - self.writable = True - - async def async_increment_with_floor(self, key, value, ttl): - if not self.writable: - raise ConnectionError("redis write failed") - return await super().async_increment_with_floor(key, value, ttl) + assert blocked.value.headers.get("Retry-After") == "270" + assert any("Redis failed while counting Admin UI sign-in attempts" in r.getMessage() for r in records) @pytest.mark.asyncio -async def test_failures_redis_refused_still_count_once_redis_recovers(monkeypatch): - """Regression: a guess Redis could not record must not be forgotten when Redis comes back. - - Such a guess lands in this worker's own store. Reading only Redis afterwards handed the - attacker that guess again, so the budget was the limit plus however many writes failed. - """ - from litellm.proxy._types import ProxyException - +async def test_a_failed_redis_delete_still_clears_this_workers_counter(monkeypatch): + """The fail-open tradeoff: when Redis cannot clear the pair, the worker clears what it holds and moves on.""" monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - redis = _WriteRefusingRedis() - throttle = _throttle(max_attempts=2, redis_cache=redis) + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(user_limit=5, redis_cache=_DownRedis()) + key = throttle._keys("user@corp.com").pair_counter - with pytest.raises(ProxyException, match="Invalid credentials"): - await _guess(throttle) - assert not redis.values, "the refused write must not have reached Redis" + assert [await _fail(throttle, username="user@corp.com") for _ in range(2)] == ["401", "401"] + assert _local_count(throttle, key) == 2 - redis.recover() - with pytest.raises(ProxyException, match="Invalid credentials"): - await _guess(throttle) - assert [v for k, v in redis.values.items() if ":user:" in k] == [1], "only the recorded guess is in Redis" - - with pytest.raises(ProxyException) as blocked: - await _guess(throttle) - assert blocked.value.code == "429", "the guess Redis missed and the one it took must add up to the limit" + await _db_login(throttle, "user@corp.com", "right", correct=True) + assert _local_count(throttle, key) == 0 @pytest.mark.asyncio async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): - """Regression: throttle entries must not evict cached credentials. - - user_api_key_cache holds at most 200 in-memory entries and evicts the soonest to - expire first, so parking 900s sign-in counters there let a stream of made-up usernames - push out the much shorter lived credential entries, sending every ordinary API request - back to the database. - """ + """Regression: throttle entries must not evict cached credentials from user_api_key_cache.""" from litellm.proxy import proxy_server as ps from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX, LoginThrottle monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - auth_cache_keys_before = set(ps.user_api_key_cache.in_memory_cache.cache_dict) request = MagicMock() @@ -1353,53 +1378,66 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): throttle = LoginThrottle.from_request(request, general_settings={}, redis_cache=None) for i in range(25): - with pytest.raises(ProxyException, match="Invalid credentials"): - await _guess(throttle, username=f"made-up-{i}@example.com") + assert await _fail(throttle, username=f"made-up-{i}@example.com") == "401" added = set(ps.user_api_key_cache.in_memory_cache.cache_dict) - auth_cache_keys_before - assert not [k for k in added if str(k).startswith(_CACHE_KEY_PREFIX)], ( - "sign-in counters must live in their own cache, not the key-authentication cache" - ) + assert not [k for k in added if str(k).startswith(_CACHE_KEY_PREFIX)] def test_settings_that_arrive_as_environment_strings_are_honored(): - """An `os.environ/VAR` reference in general_settings resolves to a string, not an int. - - Regression: a digit string fell back to the default with only a log line, so an operator - tightening the limits through environment substitution silently kept the stock ceilings. - """ + """An `os.environ/VAR` reference in general_settings resolves to a string, not an int.""" from litellm.proxy.auth.login_throttle import LoginThrottle request = MagicMock() - request.headers = {} + request.headers = {"x-forwarded-for": "203.0.113.9"} request.client = MagicMock() - request.client.host = "1.2.3.4" + request.client.host = "10.0.0.1" throttle = LoginThrottle.from_request( request, general_settings={ - "max_failed_login_attempts": "7", + "trusted_proxy_ranges": "10.0.0.0/8", "max_failed_login_attempts_per_source": " 70 ", + "max_failed_login_attempts_per_user": "7", "failed_login_window_seconds": "not-a-number", + "failed_login_block_seconds": "-5", }, redis_cache=None, ) - assert throttle.max_attempts == 7 - assert throttle.max_attempts_per_source == 70 - assert throttle.window_seconds == 900, "garbage still falls back to the default" + assert throttle.source_limit == 70 + assert throttle.user_limit == 7 + assert throttle.window_seconds == 60, "garbage falls back to the default" + assert throttle.block_seconds == 300, "a value below one would block nothing or forever" + + +def test_the_defaults_are_the_agreed_ones(): + from litellm.proxy.auth.login_throttle import LoginThrottle + + request = MagicMock() + request.headers = {"x-forwarded-for": "203.0.113.9"} + request.client = MagicMock() + request.client.host = "10.0.0.1" + throttle = LoginThrottle.from_request( + request, general_settings={"trusted_proxy_ranges": ["10.0.0.0/8"]}, redis_cache=None + ) + + assert (throttle.source_limit, throttle.user_limit, throttle.window_seconds, throttle.block_seconds) == ( + 10, + 5, + 60, + 300, + ) def test_the_disable_flag_is_read_once_not_per_login_attempt(monkeypatch): - """Regression: the kill switch was read through the secret manager on every unauthenticated request. - - With a hosted secret manager in read mode that is a synchronous network call per guess, so a - flood of wrong passwords could exhaust the secret manager even after the source was refused. - """ + """Regression: the kill switch was read through the secret manager on every unauthenticated request.""" from litellm.proxy.auth import login_throttle reads: Final[list[str]] = [] # mutable-ok: test-only call recorder - monkeypatch.setattr(login_throttle, "get_secret_bool", lambda name, default: reads.append(name) or default) + monkeypatch.setattr( + login_throttle, "get_secret_bool", lambda name, default_value: reads.append(name) or default_value + ) login_throttle._rate_limit_disabled.cache_clear() request = MagicMock() request.headers = {} @@ -1413,105 +1451,71 @@ def test_the_disable_flag_is_read_once_not_per_login_attempt(monkeypatch): assert reads == ["LITELLM_DISABLE_LOGIN_RATE_LIMIT"] -def test_a_negative_or_boolean_setting_falls_back_to_the_default(): - """A limit below one would refuse everyone; a bool is a typo, not a count.""" - from litellm.proxy.auth.login_throttle import LoginThrottle - - request = MagicMock() - request.headers = {} - request.client = MagicMock() - request.client.host = "1.2.3.4" - - throttle = LoginThrottle.from_request( - request, - general_settings={"max_failed_login_attempts": "-7", "max_failed_login_attempts_per_source": True}, - redis_cache=None, - ) - - assert throttle.max_attempts == 50 - assert throttle.max_attempts_per_source == 250 - - @pytest.mark.asyncio -async def test_a_refused_username_cannot_forge_log_lines(monkeypatch): +async def test_a_blocked_username_cannot_forge_log_lines(monkeypatch): """The username reaches a warning log, so it must not carry newlines or control bytes.""" import logging from litellm._logging import verbose_proxy_logger - from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(max_attempts=1) + throttle = _throttle(user_limit=1) forged = "victim@example.com\nWARNING: sign-in succeeded for attacker\x00" - with pytest.raises(ProxyException): - await _guess(throttle, username=forged) + assert await _fail(throttle, username=forged) == "401" records: list[logging.LogRecord] = [] handler = logging.Handler() handler.emit = records.append verbose_proxy_logger.addHandler(handler) try: - with pytest.raises(ProxyException) as blocked: - await _guess(throttle, username=forged) + assert await _fail(throttle, username=forged) == "401" finally: verbose_proxy_logger.removeHandler(handler) - assert blocked.value.code == "429" - emitted = [r.getMessage() for r in records if "sign-in attempts exhausted" in r.getMessage()] - assert emitted, "the refusal must be logged" + emitted = [r.getMessage() for r in records if "Admin UI sign-in blocked" in r.getMessage()] + assert emitted, "installing the block must be logged" assert "\n" not in emitted[0] and "\x00" not in emitted[0] assert "victim@example.com" in emitted[0] @pytest.mark.asyncio -async def test_a_username_spray_cannot_evict_an_existing_counter(monkeypatch): - """Regression: the in-memory tier must hold more counters than a spray can create. - - The default in-memory cache keeps 200 entries and evicts the soonest to expire, and - every counter shares one window, so eviction was effectively oldest-first. A few - hundred made-up usernames therefore pushed out the attacker's own counter and handed - back a fresh allowance against the real account. Username and source counters must also - live in separate stores, or the same spray evicts the source counter meant to stop it. - """ - from litellm.proxy._types import ProxyException +async def test_a_username_spray_cannot_evict_an_active_block(monkeypatch): + """Counters and blocks live in separate bounded stores, so a flood of made-up pairs fills the counter + store while the blocks it already earned stay in force.""" + from litellm.caching.in_memory_cache import InMemoryCache from litellm.proxy.auth.login_throttle import ( + _BLOCKS, + _COUNTERS, + _MAX_TRACKED_BLOCKS, + _MAX_TRACKED_COUNTERS, LoginThrottle, - _FAILED_LOGIN_SOURCE_CACHE, - _FAILED_LOGIN_USERNAME_CACHE, - _MAX_TRACKED_LOGIN_SOURCES, - _MAX_TRACKED_LOGIN_USERNAMES, ) monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - assert _MAX_TRACKED_LOGIN_SOURCES >= 10_000 - assert _MAX_TRACKED_LOGIN_USERNAMES >= 10_000 - assert _FAILED_LOGIN_SOURCE_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_SOURCES - assert _FAILED_LOGIN_USERNAME_CACHE.in_memory_cache.max_size_in_memory == _MAX_TRACKED_LOGIN_USERNAMES - assert _FAILED_LOGIN_SOURCE_CACHE.in_memory_cache is not _FAILED_LOGIN_USERNAME_CACHE.in_memory_cache - + assert _MAX_TRACKED_COUNTERS >= 10_000 and _MAX_TRACKED_BLOCKS >= 10_000 + assert _COUNTERS is not _BLOCKS + counters, blocks = InMemoryCache(max_size_in_memory=50), InMemoryCache(max_size_in_memory=50) throttle = LoginThrottle( client_ip="10.9.9.9", - max_attempts=3, - max_attempts_per_source=10_000, - window_seconds=900, - username_cache=_FAILED_LOGIN_USERNAME_CACHE, - source_cache=_FAILED_LOGIN_SOURCE_CACHE, + source_limit=None, + user_limit=1, + window_seconds=60, + block_seconds=300, + counters=counters, + blocks=blocks, ) victim = "spray-victim@corp.com" - for _ in range(3): - with pytest.raises(ProxyException): - await _guess(throttle, username=victim) + assert [await _fail(throttle, username=victim) for _ in range(2)] == ["401", "401"] - for i in range(500): + for i in range(200): await throttle.record_failure(f"spray-filler-{i}@corp.com") - assert await throttle._failures(throttle.username_cache, throttle._username_key(victim)) == 3, "the counter must survive a spray" - with pytest.raises(ProxyException) as blocked: - await _guess(throttle, username=victim) - assert blocked.value.code == "429" + assert len(counters.cache_dict) <= 50, "the counter store is bounded" + assert counters.get_cache(throttle._keys(victim).pair_counter) is None, "the victim's counter was evicted" + assert await _fail(throttle, username=victim) == "429", "the block survived the spray" def _patch_sso_configured(stack: ExitStack, *, configured: bool) -> None: diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index 56aa87f1e50..a349d378985 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -517,38 +517,25 @@ def make_key( def reset_login_throttle(monkeypatch): """Clear the Admin UI failed-login counters between tests. - `client` is session scoped and the counters live in shared module stores with a 900s - window, so without this a failed sign-in test could return 429 in unrelated tests later. + `client` is session scoped and the counters live in shared module stores with a 300s block + window, so without this a failed sign-in test could block unrelated tests later. Only the throttle's own keys are removed, so other cache entries remain untouched. """ from litellm.proxy import proxy_server as ps from litellm.proxy.auth import login_throttle - from litellm.proxy.auth.login_throttle import ( - _CACHE_KEY_PREFIX, - _FAILED_LOGIN_SOURCE_CACHE, - _FAILED_LOGIN_USERNAME_CACHE, - ) + from litellm.proxy.auth.login_throttle import _BLOCKS, _CACHE_KEY_PREFIX, _COUNTERS async def _no_delay(_seconds: float) -> None: - """The escalating wait on a rejected sign-in, replaced so the route tests stay fast.""" + """The hold on a rejected sign-in from a blocked key, replaced so the route tests stay fast.""" monkeypatch.setattr(login_throttle, "_sleep", _no_delay) def _drop_throttle_keys() -> None: - login_throttle._DELAYS_IN_FLIGHT.clear() - for cache in (_FAILED_LOGIN_USERNAME_CACHE, _FAILED_LOGIN_SOURCE_CACHE): - in_memory = getattr(cache, "in_memory_cache", None) - if in_memory is None: - continue - tracked = tuple( - key - for store in (getattr(in_memory, "cache_dict", None), getattr(in_memory, "ttl_dict", None)) - if isinstance(store, dict) - for key in tuple(store) - if str(key).startswith(_CACHE_KEY_PREFIX) - ) - for key in tracked: - in_memory.delete_cache(key) + login_throttle._HELD_ATTEMPTS.clear() + for store in (_COUNTERS, _BLOCKS): + for key in tuple(store.cache_dict) + tuple(store.ttl_dict): + if key.startswith(_CACHE_KEY_PREFIX): + store.delete_cache(key) monkeypatch.setattr(ps, "redis_usage_cache", None) _drop_throttle_keys() diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index f8382c64e6a..7399e64c421 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -10,11 +10,8 @@ Routes covered: from __future__ import annotations -from concurrent.futures import ThreadPoolExecutor from unittest.mock import AsyncMock, MagicMock -import pytest - from .conftest import normalize # --------------------------------------------------------------------------- @@ -495,15 +492,38 @@ def _install_real_auth(monkeypatch, **settings): def _form_login(client, username="admin", password="wrong"): - return client.post( - "/login", data={"username": username, "password": password}, follow_redirects=False - ).status_code + return client.post("/login", data={"username": username, "password": password}, follow_redirects=False).status_code def _json_login(client, path, username="admin", password="wrong"): return client.post(path, json={"username": username, "password": password}).status_code +def _db_user(monkeypatch, email: str): + """A database user with a stored hash, faked so the route reaches the known-user branch without Postgres.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy import proxy_server as ps + + user = MagicMock() + user.user_id = "u-1" + user.user_email = email + user.user_role = "internal_user" + user.password = "scrypt:stored" + repo = MagicMock() + repo.return_value.table.find_first = AsyncMock(return_value=user) + monkeypatch.setattr(ps, "prisma_client", MagicMock()) + monkeypatch.setattr("litellm.proxy.auth.login_utils.UserRepository", repo) + monkeypatch.setattr("litellm.proxy.auth.login_utils._rehash_password_if_needed", AsyncMock()) + monkeypatch.setattr( + "litellm.proxy.auth.login_utils.verify_password", lambda given, stored: given == "right-db-password" + ) + monkeypatch.setattr( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", AsyncMock(return_value={"token": "sk-ui"}) + ) + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + + def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset_login_throttle): """The endpoint is not part of the key, so spending the budget on one route blocks the rest. @@ -511,92 +531,117 @@ def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset """ _install_real_auth( monkeypatch, - max_failed_login_attempts=10, + max_failed_login_attempts_per_user=10, control_plane_url="https://cp.example.com", ) assert [_form_login(client) for _ in range(5)] == [401] * 5 assert [_json_login(client, "/v2/login") for _ in range(5)] == [401] * 5 - assert _json_login(client, "/v3/login") == 429, "the eleventh attempt must be refused on a third route" + assert _json_login(client, "/v3/login") == 401, "the eleventh failure crosses the limit and installs the block" + assert _json_login(client, "/v3/login") == 429, "the twelfth attempt must be refused on a third route" def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_login_throttle): """The database lookup is case-insensitive, so casing must not partition the counter.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=10) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=3) - assert [_json_login(client, "/v2/login", username="admin@corp.com") for _ in range(5)] == [401] * 5 - assert [_json_login(client, "/v2/login", username="ADMIN@corp.com") for _ in range(5)] == [401] * 5 + assert [_json_login(client, "/v2/login", username="admin@corp.com") for _ in range(2)] == [401] * 2 + assert [_json_login(client, "/v2/login", username="ADMIN@corp.com") for _ in range(2)] == [401] * 2 assert _json_login(client, "/v2/login", username="Admin@corp.com") == 429 def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_throttle): - """The 429 tells the caller how long the window has left.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=2, failed_login_window_seconds=77) + """The 429 tells the caller how long the block has left, after the 30 seconds it was already held.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=77) assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] refused = client.post("/v2/login", json={"username": "admin", "password": "wrong"}) assert refused.status_code == 429 - assert refused.headers.get("retry-after") == "77" + assert refused.headers.get("retry-after") == "47" def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, reset_login_throttle): """The no-JavaScript form must render a wait page when its POST is throttled.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=2, failed_login_window_seconds=77) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=77) assert [_form_login(client) for _ in range(2)] == [401, 401] refused = client.post("/login", data={"username": "admin", "password": "wrong"}) assert refused.status_code == 429 assert refused.headers.get("content-type", "").startswith("text/html") - assert "Try again in about 77 seconds" in refused.text - assert refused.headers.get("retry-after") == "77" + assert "Try again in about 47 seconds" in refused.text + assert refused.headers.get("retry-after") == "47" def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle): - """The username counter carries no address, so one account exhausting it cannot block another.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=2) + """The pair block is per username, so one account's block cannot take the office down with it.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) - for _ in range(3): - _json_login(client, "/v2/login", username="admin") + assert [_json_login(client, "/v2/login", username="admin") for _ in range(3)] == [401, 401, 429] assert _json_login(client, "/v2/login", username="someone-else@example.com") == 401 -def test_a_spray_across_usernames_is_refused_on_the_source_counter(client, monkeypatch, reset_login_throttle): - """A fresh username per guess keeps every username counter at one, so the address is what stops it.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=100, max_failed_login_attempts_per_source=4) +def test_a_spray_across_usernames_is_blocked_on_the_source_when_the_source_is_attributable( + client, monkeypatch, reset_login_throttle +): + """A fresh username per guess keeps every pair at one, so the address is what stops it.""" + _install_real_auth(monkeypatch, trusted_proxy_ranges=["10.0.0.0/8"], max_failed_login_attempts_per_source=4) - sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(4)] - assert sprayed == [401] * 4 + sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(5)] + assert sprayed == [401] * 5 - assert _json_login(client, "/v2/login", username="sprayed-5@corp.com") == 429 + assert _json_login(client, "/v2/login", username="sprayed-6@corp.com") == 429 -def test_the_configured_admin_password_still_signs_in_while_refused(client, monkeypatch, reset_login_throttle): +def test_a_spray_across_usernames_is_not_blocked_without_trusted_proxy_ranges( + client, monkeypatch, reset_login_throttle +): + """Without a configured proxy range the peer address is whoever fronts the proxy, shared by every + client, so a source-wide block would block them all and the source scope stays off.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=4) + + sprayed = [_json_login(client, "/v2/login", username=f"sprayed-{i}@corp.com") for i in range(8)] + assert sprayed == [401] * 8 + + +def test_the_configured_admin_password_still_signs_in_while_blocked(client, monkeypatch, reset_login_throttle): """The operator must never be locked out of the console by traffic aimed at it.""" from unittest.mock import AsyncMock, patch - _install_real_auth(monkeypatch, max_failed_login_attempts=2) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] - assert _json_login(client, "/v2/login") == 429 + assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] - with patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed - "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + with ( + patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), + patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) + ), ): assert _json_login(client, "/v2/login", password="right-password") == 200 -def test_sign_in_succeeds_again_once_the_budget_is_restored(client, monkeypatch, reset_login_throttle): - """A cleared bucket lets the same username straight back in.""" - _install_real_auth(monkeypatch, max_failed_login_attempts=2) +def test_a_database_users_correct_password_signs_in_while_blocked(client, monkeypatch, reset_login_throttle): + """The block is soft: guessing at an account slows the guesser down, it does not lock the owner out.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + _db_user(monkeypatch, "user@corp.com") - assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] - assert _json_login(client, "/v2/login") == 429 + assert [_json_login(client, "/v2/login", username="user@corp.com") for _ in range(3)] == [401, 401, 429] + + assert _json_login(client, "/v2/login", username="user@corp.com", password="right-db-password") == 200 + assert _json_login(client, "/v2/login", username="user@corp.com") == 429, "the block itself is still in force" + + +def test_sign_in_succeeds_again_once_the_block_is_cleared(client, monkeypatch, reset_login_throttle): + """A cleared store lets the same username straight back to a plain credential check.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + + assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] reset_login_throttle() assert _json_login(client, "/v2/login") == 401 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 2833686f115..07696f23da7 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3438,6 +3438,34 @@ async def test_load_config_warns_per_worker_login_counters_without_general_setti assert "Running 4 workers but Redis is not configured" in caplog.text +@pytest.mark.asyncio +async def test_load_config_warns_that_the_source_login_limit_is_off_without_trusted_proxy_ranges( + tmp_path, monkeypatch, caplog +): + """The per-source failed-login limit is skipped when the source cannot be attributed, and the + operator must be told so at startup; a configured range silences it.""" + import logging + + from litellm.proxy.auth.login_throttle import warn_source_login_limit_is_off + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setenv("NUM_WORKERS", "1") + warn_source_login_limit_is_off.cache_clear() + config_file = tmp_path / "config.yaml" + config_file.write_text("model_list: []\n") + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert "trusted_proxy_ranges is not set" in caplog.text + + caplog.clear() + warn_source_login_limit_is_off.cache_clear() + config_file.write_text("model_list: []\ngeneral_settings:\n trusted_proxy_ranges: ['10.0.0.0/8']\n") + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert "trusted_proxy_ranges is not set" not in caplog.text + + @pytest.mark.asyncio async def test_load_environment_variables_direct_and_os_environ(): """ @@ -13331,14 +13359,16 @@ async def test_login_throttle_settings_are_not_hot_applied_from_the_database(): ps.general_settings.clear() await ProxyConfig()._update_general_settings( db_general_settings={ - "max_failed_login_attempts": 999, + "max_failed_login_attempts_per_user": 999, "max_failed_login_attempts_per_source": 999, "failed_login_window_seconds": 1, + "failed_login_block_seconds": 1, } ) - assert "max_failed_login_attempts" not in ps.general_settings + assert "max_failed_login_attempts_per_user" not in ps.general_settings assert "max_failed_login_attempts_per_source" not in ps.general_settings assert "failed_login_window_seconds" not in ps.general_settings + assert "failed_login_block_seconds" not in ps.general_settings finally: ps.general_settings.clear() ps.general_settings.update(original) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 099ce397ffe..6203e79bccd 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26447,9 +26447,14 @@ export interface components { * @description If True, router fallbacks configured in router_settings are only attempted when the calling key (and its team and project) is allowed to call the fallback model; unauthorized fallback targets are skipped and the primary model's error is returned. Default is False. */ enforce_fallback_model_access?: boolean | null; + /** + * Failed Login Block Seconds + * @description How long a blocked source address, or source address and username, stays blocked. The block is soft: a correct password still signs in, but each attempt from a blocked key takes one of 5 held slots per worker and a wrong password is held for 30 seconds before it is refused with 429. Set under `general_settings` in config.yaml. Defaults to 300 + */ + failed_login_block_seconds?: number | null; /** * Failed Login Window Seconds - * @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 900 + * @description Fixed window in seconds over which failed Admin UI sign-in attempts are counted. The window starts at the first failure and is not extended by later ones. Set under `general_settings` in config.yaml. Defaults to 60 */ failed_login_window_seconds?: number | null; /** @@ -26496,16 +26501,23 @@ export interface components { * @description max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider */ max_batch_file_size_mb?: number | null; - /** - * Max Failed Login Attempts - * @description Number of failed Admin UI sign-in attempts allowed for one username, from any source address, within `failed_login_window_seconds`, before further attempts for that username are refused with 429. Attempts are answered with a doubling delay well before this ceiling. Set under `general_settings` in config.yaml. Defaults to 50 - */ - max_failed_login_attempts?: number | null; /** * Max Failed Login Attempts Per Source - * @description Number of failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`, before further attempts from that address are refused with 429. Counted independently of `max_failed_login_attempts`. Set under `general_settings` in config.yaml. Defaults to 250 + * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set, since otherwise every client behind an ingress shares one address. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 */ max_failed_login_attempts_per_source?: number | null; + /** + * Max Failed Login Attempts Per Source Overrides + * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins. Set under `general_settings` in config.yaml + */ + max_failed_login_attempts_per_source_overrides?: { + [key: string]: number; + } | null; + /** + * Max Failed Login Attempts Per User + * @description Failed Admin UI sign-in attempts allowed from one source address for one username within `failed_login_window_seconds`. One more blocks that address for that username for `failed_login_block_seconds`, and its further failures stop counting against `max_failed_login_attempts_per_source`, so a script stuck on one account does not block everyone behind the same address. Set under `general_settings` in config.yaml. Defaults to 5 + */ + max_failed_login_attempts_per_user?: number | null; /** * Max File Size Mb * @description max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider From 7c2c59da903f72d13829c9b2f486b35bef1cb6a6 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 00:13:51 +0000 Subject: [PATCH 049/253] test(proxy): explain the internal patches in the login throttle tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_litellm/proxy/auth/test_login_utils.py | 16 +++++++++++----- .../proxy/proxy_server/test_routes_login_sso.py | 4 +++- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 22fcedaa19d..ac40b5364f9 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -725,11 +725,15 @@ async def _db_login(throttle, username: str, password: str, *, correct: bool): from litellm.proxy.auth.login_utils import authenticate_user with ( - patch("litellm.proxy.auth.login_utils.UserRepository", _known_user(username)), + patch( # test-quality-ok: the user lookup is the database boundary; faked so no DB is needed + "litellm.proxy.auth.login_utils.UserRepository", _known_user(username) + ), patch( # test-quality-ok: reaches the known-DB-user branch without a database "litellm.proxy.auth.login_utils.verify_password", return_value=correct ), - patch("litellm.proxy.auth.login_utils._rehash_password_if_needed", new=AsyncMock()), + patch( # test-quality-ok: the rehash writes to the database; faked so no DB is needed + "litellm.proxy.auth.login_utils._rehash_password_if_needed", new=AsyncMock() + ), patch( # test-quality-ok: success mints a UI key; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ), @@ -1062,7 +1066,9 @@ async def test_the_configured_admin_credentials_bypass_the_throttle(monkeypatch) assert [await _fail(throttle) for _ in range(3)] == ["401", "401", "429"] with ( - patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), + patch( # test-quality-ok: the admin sign-in upserts the admin row; faked so no DB is needed + "litellm.proxy.auth.login_utils.user_update", new=AsyncMock() + ), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ), @@ -1135,9 +1141,9 @@ async def test_a_user_with_no_password_set_does_not_consume_the_budget(monkeypat repo = MagicMock() repo.return_value.table.find_first = AsyncMock(return_value=passwordless) - with patch( + with patch( # test-quality-ok: reaches the passwordless-DB-user branch without a database "litellm.proxy.auth.login_utils.UserRepository", repo - ): # test-quality-ok: reaches the passwordless-DB-user branch without a database + ): for _ in range(5): with pytest.raises(ProxyException) as exc: await authenticate_user( diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index 7399e64c421..82e86086518 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -618,7 +618,9 @@ def test_the_configured_admin_password_still_signs_in_while_blocked(client, monk assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] with ( - patch("litellm.proxy.auth.login_utils.user_update", new=AsyncMock()), + patch( # test-quality-ok: the admin sign-in upserts the admin row; faked so no DB is needed + "litellm.proxy.auth.login_utils.user_update", new=AsyncMock() + ), patch( # test-quality-ok: success mints a UI key and persists the user; faked so no DB is needed "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ), From 84c098df92f8d89ed5d083466ec62347110062e0 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 00:42:03 +0000 Subject: [PATCH 050/253] fix(proxy): fold late-arriving per-key spend into already rolled-up global days The reconcile now records the database clock of the scan behind the last complete run and, on the next run, rewrites every closed day with per-key rows updated since then, however old the day is. Replaying only the marker day and the one before it missed a delayed flush or retry that landed on an older date, and reads through the marker come from the global table alone, so that spend was never counted. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../daily_global_spend_rollup.py | 128 +++++++++++++----- .../test_daily_global_spend_rollup.py | 88 ++++++++++-- 2 files changed, 168 insertions(+), 48 deletions(-) diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index a9fb7669785..73068381dab 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -2,9 +2,12 @@ Only days that are over get rolled up, so a pod still flushing per-key spend for the current day can never leave the global table short; usage reads serve days through the recorded -marker from the global table and later days live from the per-key table. The marker lives in -``LiteLLM_Config``. This runs as a background cron, never in a Prisma migration, since on a -large deployment the first backfill is minutes of work. +marker from the global table and later days live from the per-key table. Per-key rows are +dated by request start, so spend can land on a day that was already rolled up (a flush +straddling midnight, a retry after an outage). Each run therefore also rewrites every closed +day that has rows touched since the previous run's scan, whatever the date. The marker lives +in ``LiteLLM_Config``. This runs as a background cron, never in a Prisma migration, since on +a large deployment the first backfill is minutes of work. """ from collections.abc import Awaitable, Callable @@ -27,7 +30,6 @@ if TYPE_CHECKING: from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.proxy.utils import PrismaClient -_REPLAY_DAYS: Final = 1 GLOBAL_SPEND_TABLE_NAME: Final = "LiteLLM_DailyGlobalSpend" # The unique constraint, in constraint order. NULL never matches itself in a unique index, so # every column is normalized to '' or the same group would be inserted again on every run. @@ -69,15 +71,26 @@ def _reconcile_day_sql() -> str: RECONCILE_DAY_SQL: Final = _reconcile_day_sql() +_DB_NOW_SQL: Final = "SELECT (NOW() AT TIME ZONE 'UTC')::text AS now" +_ALL_CLOSED_DAYS_SQL: Final = 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" <= $1 ORDER BY "date"' +# Pod clocks drift from the database clock and from each other, so rows are picked up from a +# little before the previous scan; rewriting a day twice is idempotent. _PENDING_DAYS_SQL: Final = ( - 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" >= $1 AND "date" <= $2 ORDER BY "date"' + 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" <= $1 ' + 'AND ("date" > $2 OR "updated_at" >= $3::timestamp - INTERVAL \'1 hour\') ' + 'ORDER BY "date"' ) class ReconciledThrough(BaseModel): + """``reconciled_through`` is the last closed UTC day the global table covers. ``scanned_at`` is + the database clock when the scan behind the last fully successful run started: every per-key + row written before it, on any day through the marker, is in the global table.""" + model_config = ConfigDict(frozen=True, extra="ignore") reconciled_through: str + scanned_at: str | None = None class _MarkerRow(BaseModel): @@ -92,6 +105,12 @@ class _DateRow(BaseModel): date: str +class _NowRow(BaseModel): + model_config = ConfigDict(frozen=True, extra="ignore") + + now: str + + @dataclass(frozen=True, slots=True) class ReconcileResult: days_reconciled: tuple[str, ...] @@ -99,49 +118,70 @@ class ReconcileResult: failed_day: str | None = None -def _marker_from_param_value(value: object) -> str | None: +@dataclass(frozen=True, slots=True) +class _PendingScan: + marker: ReconciledThrough | None + scanned_at: str + days: tuple[str, ...] + + +def _marker_from_param_value(value: object) -> ReconciledThrough | None: try: - parsed: Final = ( + return ( ReconciledThrough.model_validate_json(value) if isinstance(value, str) else ReconciledThrough.model_validate(value) ) except ValidationError: return None - return parsed.reconciled_through -async def reconciled_through(prisma_client: "PrismaClient") -> str | None: - """The last UTC day ``LiteLLM_DailyGlobalSpend`` is known to cover, or None before the first run.""" +async def read_marker(prisma_client: "PrismaClient") -> ReconciledThrough | None: from litellm.proxy.utils import get_config_param row: Final = await get_config_param(prisma_client, DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) return None if row is None else _marker_from_param_value(_MarkerRow.model_validate(row).param_value) -async def _record_reconciled_through(prisma_client: "PrismaClient", day: str) -> None: +async def reconciled_through(prisma_client: "PrismaClient") -> str | None: + """The last UTC day ``LiteLLM_DailyGlobalSpend`` is known to cover, or None before the first run.""" + marker: Final = await read_marker(prisma_client) + return None if marker is None else marker.reconciled_through + + +async def _record_marker(prisma_client: "PrismaClient", marker: ReconciledThrough) -> None: from litellm.proxy.utils import invalidate_config_param await ConfigRepository(prisma_client).set_param( - DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, ReconciledThrough(reconciled_through=day).model_dump_json() + DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, marker.model_dump_json() ) await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) -def _first_pending_day(marker: str | None) -> str: - if marker is None: - return "" - return (date.fromisoformat(marker) - timedelta(days=_REPLAY_DAYS)).isoformat() +async def _db_now(prisma_client: "PrismaClient") -> str: + rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL) + return _NowRow.model_validate(rows[0]).now + + +async def _scan_pending(prisma_client: "PrismaClient", today: date) -> _PendingScan: + """Every closed UTC day (strictly before today) still to roll up, oldest first: days past the + marker, plus any day with per-key rows written since the scan behind the marker. Before a + run has fully succeeded there is no such scan, so every closed day is rolled up.""" + marker: Final = await read_marker(prisma_client) + scanned_at: Final = await _db_now(prisma_client) + last_closed_day: Final = (today - timedelta(days=1)).isoformat() + rows: Final = ( + await prisma_client.db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day) + if marker is None or marker.scanned_at is None + else await prisma_client.db.query_raw( + _PENDING_DAYS_SQL, last_closed_day, marker.reconciled_through, marker.scanned_at + ) + ) + return _PendingScan(marker, scanned_at, tuple(_DateRow.model_validate(row).date for row in rows)) async def pending_days(prisma_client: "PrismaClient", today: date) -> tuple[str, ...]: - """Every closed UTC day (strictly before today) still to roll up, oldest first. The marker - day and the one before it are replayed so per-key rows that landed after their day was - rolled up (a flush straddling midnight, a late retry) are folded in.""" - marker: Final = await reconciled_through(prisma_client) - last_closed_day: Final = (today - timedelta(days=1)).isoformat() - rows: Final = await prisma_client.db.query_raw(_PENDING_DAYS_SQL, _first_pending_day(marker), last_closed_day) - return tuple(_DateRow.model_validate(row).date for row in rows) + return (await _scan_pending(prisma_client, today)).days async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None: @@ -155,26 +195,42 @@ async def run_daily_global_spend_reconcile( today: date | None = None, ) -> ReconcileResult: """Roll up every pending day, advancing the marker after each; a failing day stops the run - with the marker on the last good day so the next run resumes there.""" + with the marker on the last good day so the next run resumes there. The scan time is only + recorded once every pending day is done, so late rows a failed run saw are found again.""" effective_today: Final = today or datetime.now(timezone.utc).date() - days: Final = await pending_days(prisma_client, effective_today) - done: Final = await _reconcile_until_failure(prisma_client, days) - failed: Final = days[len(done)] if len(done) < len(days) else None - marker: Final = done[-1] if done else await reconciled_through(prisma_client) - return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=failed) + scan: Final = await _scan_pending(prisma_client, effective_today) + done: Final = await _reconcile_until_failure(prisma_client, scan) + if len(done) < len(scan.days): + marker: Final = await reconciled_through(prisma_client) + return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=scan.days[len(done)]) + if scan.marker is not None or done: + await _record_marker(prisma_client, _advanced(scan.marker, done, scanned_at=scan.scanned_at)) + return ReconcileResult(days_reconciled=done, reconciled_through=await reconciled_through(prisma_client)) -async def _reconcile_until_failure(prisma_client: "PrismaClient", days: tuple[str, ...]) -> tuple[str, ...]: - for index, day in enumerate(days): - if not await _reconcile_and_record(prisma_client, day): - return days[:index] - return days +def _advanced(marker: ReconciledThrough | None, days: tuple[str, ...], *, scanned_at: str | None) -> ReconciledThrough: + """The marker after ``days`` were rewritten: a late old day never moves it back.""" + through: Final = max((marker.reconciled_through if marker is not None else "", *days)) + return ReconciledThrough(reconciled_through=through, scanned_at=scanned_at) -async def _reconcile_and_record(prisma_client: "PrismaClient", day: str) -> bool: +async def _reconcile_until_failure(prisma_client: "PrismaClient", scan: _PendingScan) -> tuple[str, ...]: + for index, day in enumerate(scan.days): + if not await _reconcile_and_record(prisma_client, scan.marker, scan.days[: index + 1]): + return scan.days[:index] + return scan.days + + +async def _reconcile_and_record( + prisma_client: "PrismaClient", marker: ReconciledThrough | None, done_with_this: tuple[str, ...] +) -> bool: + day: Final = done_with_this[-1] try: await reconcile_day(prisma_client, day) - await _record_reconciled_through(prisma_client, day) + await _record_marker( + prisma_client, + _advanced(marker, done_with_this, scanned_at=None if marker is None else marker.scanned_at), + ) except Exception as exc: # noqa: BLE001 # one bad day must not lose the days already done verbose_proxy_logger.exception("Daily global spend reconcile: day %s failed: %s", day, exc) return False diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 11ca72e7b3d..9a098744f08 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -15,6 +15,7 @@ from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( RECONCILE_DAY_SQL, + read_marker, reconciled_through, run_daily_global_spend_reconcile, run_scheduled_daily_global_spend_reconcile, @@ -41,13 +42,25 @@ class _FakeConfigTable: class _FakeDb: + """Per-key rows are ``{date: updated_at}`` with a fake database clock that ticks per query, + so "rows written since the last scan" behaves like Postgres would.""" + def __init__(self, prisma: "_FakePrisma") -> None: self._prisma = prisma self.litellm_config = _FakeConfigTable() async def query_raw(self, sql: str, *params: str) -> list[dict[str, str]]: - first, last = params - return [{"date": d} for d in sorted(self._prisma.user_days) if first <= d <= last] + if sql.startswith("SELECT (NOW()"): + self._prisma.clock += 1 + return [{"now": f"clock-{self._prisma.clock:04d}"}] + rows = self._prisma.user_rows + if len(params) == 1: + (last,) = params + return [{"date": d} for d in sorted(rows) if d <= last] + last, marker, scanned_at = params + return [ + {"date": d} for d, written in sorted(rows.items()) if d <= last and (d > marker or written >= scanned_at) + ] async def execute_raw(self, sql: str, *params: str) -> int: (day,) = params @@ -61,11 +74,17 @@ class _FakePrisma: """Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw.""" def __init__(self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset()) -> None: - self.user_days = user_days + self.clock = 0 + self.user_rows: dict[str, str] = {d: "clock-0000" for d in user_days} self.failing_days = failing_days self.reconciled: list[str] = [] self.db = _FakeDb(self) + def write_late_row(self, day: str) -> None: + """A per-key row for ``day`` lands now, after whatever scans already happened.""" + self.clock += 1 + self.user_rows[day] = f"clock-{self.clock:04d}" + async def get_generic_data(self, key: str, value: str, table_name: str) -> _FakeConfigRow | None: stored = self.db.litellm_config.rows.get(value) return None if stored is None else _FakeConfigRow(value, stored) @@ -94,20 +113,66 @@ async def test_first_run_rolls_up_every_closed_day_and_never_today(): @pytest.mark.asyncio -async def test_later_run_replays_the_marker_day_and_the_day_before_only(): - """Days older than marker-1 are settled; the marker day and its predecessor are replayed so - per-key rows that landed after their day was rolled up get folded in.""" +async def test_later_run_rolls_up_only_new_days_when_nothing_old_changed(): prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14")) await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) prisma.reconciled.clear() result = await run_daily_global_spend_reconcile(prisma, today=TODAY) - assert result.days_reconciled == ("2026-09-12", "2026-09-13", "2026-09-14") - assert "2026-09-01" not in prisma.reconciled + assert result.days_reconciled == ("2026-09-14",) assert await reconciled_through(prisma) == "2026-09-14" +@pytest.mark.asyncio +async def test_spend_landing_on_an_old_rolled_up_day_is_folded_in_by_the_next_run(): + """Per-key rows carry the request start date, so a delayed flush or retry can add spend to a + day far behind the marker. That day is rewritten, and the marker never moves back for it.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-05", "2026-09-13")) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma.reconciled.clear() + prisma.write_late_row("2026-09-01") + prisma.write_late_row("2026-09-03") + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-01", "2026-09-03") + assert "2026-09-05" not in prisma.reconciled + assert await reconciled_through(prisma) == "2026-09-13" + + +@pytest.mark.asyncio +async def test_a_late_row_seen_by_a_failed_run_is_seen_again_by_the_next_one(): + """The scan time only advances when every pending day was rewritten, otherwise a late row + found by the failed run would be counted as handled.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13")) + await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma.write_late_row("2026-09-01") + prisma.failing_days = frozenset({"2026-09-01"}) + failed = await run_daily_global_spend_reconcile(prisma, today=TODAY) + prisma.failing_days = frozenset() + prisma.reconciled.clear() + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert failed.failed_day == "2026-09-01" + assert failed.reconciled_through == "2026-09-13" + assert result.days_reconciled == ("2026-09-01",) + assert result.failed_day is None + + +@pytest.mark.asyncio +async def test_a_marker_without_a_scan_time_rolls_every_closed_day_up_again(): + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13")) + prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"reconciled_through": "2026-09-13"}' + + result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + + assert result.days_reconciled == ("2026-09-01", "2026-09-13") + marker = await read_marker(prisma) + assert marker is not None and marker.reconciled_through == "2026-09-13" and marker.scanned_at is not None + + @pytest.mark.asyncio async def test_a_run_with_no_new_closed_days_keeps_the_marker(): prisma = _FakePrisma(user_days=("2026-09-13",)) @@ -116,7 +181,7 @@ async def test_a_run_with_no_new_closed_days_keeps_the_marker(): result = await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) - assert result.days_reconciled == ("2026-09-13",) + assert result.days_reconciled == () assert result.reconciled_through == "2026-09-13" @@ -149,11 +214,10 @@ async def test_the_next_run_resumes_from_the_failed_day(): @pytest.mark.asyncio async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts(): - """A late flush for the day before the marker is exactly the replay case; when that replay - fails the marker must stay put and the operator must hear about it.""" + """When the rewrite of a late day fails the marker must stay put and the operator must hear about it.""" prisma = _FakePrisma(user_days=("2026-09-13",)) await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) - prisma.user_days = ("2026-09-12", "2026-09-13") + prisma.write_late_row("2026-09-12") prisma.failing_days = frozenset({"2026-09-12"}) alert = AsyncMock() From bc60b49e98f17bbe09a457ba96c6b2d3d0650836 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 00:55:11 +0000 Subject: [PATCH 051/253] fix(proxy): key held sign-in attempts on the source while the source is blocked An active source block now takes precedence over a pair block, so every blocked username behind one blocked source shares the source's five held slots instead of getting five each Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 4 +- .../proxy/auth/test_login_utils.py | 46 +++++++++++++++++++ 2 files changed, 48 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index fb708ecf872..393411d0670 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -290,10 +290,10 @@ class LoginThrottle: shared: Final = await self._shared_block_ttls(keys) user_ttl: Final = max(local[0], shared[0]) source_ttl: Final = max(local[1], shared[1]) - if user_ttl > 0: - return Block(scope="user", retry_after=user_ttl) if self.source_limit is not None and source_ttl > 0: return Block(scope="source", retry_after=source_ttl) + if user_ttl > 0: + return Block(scope="user", retry_after=user_ttl) return None async def _shared_block_ttls(self, keys: _Keys) -> _BlockTtls: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index ac40b5364f9..65c388240e5 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1220,6 +1220,52 @@ async def test_held_attempts_from_one_blocked_key_are_capped(monkeypatch): assert lt._HELD_ATTEMPTS.get(slot) is None, "the slots are released once the held attempts answer" +@pytest.mark.asyncio +async def test_a_blocked_source_shares_one_held_slot_pool_across_its_blocked_usernames(monkeypatch): + """Once the source is blocked, a pair block for a username must not hand that username its own five slots.""" + import asyncio + + from litellm.proxy._types import ProxyException + from litellm.proxy.auth import login_throttle as lt + from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + release = asyncio.Event() + + async def _park(_seconds: float) -> None: + await release.wait() + + monkeypatch.setattr(lt, "_sleep", _park) + throttle = _throttle(user_limit=1, source_limit=3, client_ip="203.0.113.45") + assert [await _fail(throttle) for _ in range(2)] == ["401", "401"], "the admin pair is now blocked" + assert [await _fail(throttle, username=f"spray-{i}@corp.com") for i in range(3)] == ["401"] * 3 + source_slot = throttle._keys("admin").source_block + assert throttle._local_block_ttl(source_slot) > 0, "the source is now blocked as well" + + usernames = ["admin", *(f"fresh-{i}@corp.com" for i in range(MAX_HELD_ATTEMPTS_PER_KEY - 1))] + held = [asyncio.create_task(_guess(throttle, username=name)) for name in usernames] + for _ in range(1000): + if lt._HELD_ATTEMPTS.get(source_slot) == MAX_HELD_ATTEMPTS_PER_KEY: + break + await asyncio.sleep(0) + assert lt._HELD_ATTEMPTS == {source_slot: MAX_HELD_ATTEMPTS_PER_KEY} + + try: + for name in ("admin", "fresh-0@corp.com", "never-seen@corp.com"): + with pytest.raises(ProxyException) as over_cap: + await _guess(throttle, username=name) + assert over_cap.value.code == "429" + assert over_cap.value.headers.get("Retry-After") == "30" + finally: + release.set() + for task in held: + with pytest.raises(ProxyException): + await task + + assert lt._HELD_ATTEMPTS == {} + + @pytest.mark.asyncio async def test_disabling_the_control_removes_the_hold_as_well(monkeypatch, login_delays): """The escape hatch has to turn off the whole control, not only the refusal.""" From 5b5bbac769e548199393e54aacf346953c5c5528 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:16:21 +0000 Subject: [PATCH 052/253] fix(team): link new members to the shared team member budget so /team/update applies to them Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_helpers/utils.py | 23 +-- .../test_management_helpers_utils.py | 135 +++++++++--------- 2 files changed, 82 insertions(+), 76 deletions(-) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index f3bd4b0f6dd..2e7458232cc 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -291,9 +291,9 @@ async def _clone_team_default_budget_for_member( member budget. Returns the new budget_id, or None if the default budget no longer exists in the DB. - Used when adding a new team member without an explicit per-member budget, - so the member starts with the team default's values but gets their own - private budget row (which can be edited independently). + Used when adding a new team member with a per-member ``budget_duration`` + but no other per-member limit, so the member keeps the team default's + values in their own private budget row while the reset window differs. ``budget_duration_override`` replaces the default's reset window for this member while keeping the default's other limits, so an admin can set a @@ -344,14 +344,21 @@ async def _resolve_member_budget_id( """ Resolve the budget a new team member should be linked to. - Explicit per-member limits create a fresh budget. Otherwise the team's - default member budget is cloned (with ``budget_duration`` overriding its - reset window while keeping its other limits). A lone ``budget_duration`` - with no team default creates a window-only budget. With nothing set the - member gets no budget. + Explicit per-member limits create a fresh budget. Otherwise the member is + linked to the team's shared default member budget, so later ``/team/update`` + changes reach them; ``/team/member_update`` clones that row on first write. + A lone ``budget_duration`` clones the default with the reset window + overridden, or creates a window-only budget when there is no team default. + With nothing set the member gets no budget. """ has_explicit_limit: Final = max_budget_in_team is not None or allowed_models is not None + if not has_explicit_limit and default_team_budget_id is not None and budget_duration is None: + default_budget: Final = await _budget_table(prisma_client, tx).find_unique( + where={"budget_id": default_team_budget_id} + ) + return default_team_budget_id if default_budget is not None else None + if not has_explicit_limit and default_team_budget_id is not None: return await _clone_team_default_budget_for_member( prisma_client=prisma_client, diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index a6b1fc32eda..00de5171aa7 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -164,13 +164,16 @@ async def test_management_otel_span_redacts_nested_submission_env_var_secrets( @pytest.mark.asyncio -async def test_add_new_member_clones_default_team_budget_id(): +async def test_add_new_member_links_default_team_budget_id(): """ - Test that add_new_member CLONES the team's default member budget when - max_budget_in_team is None and a default_team_budget_id is provided. + A member added without any per-member limit must be LINKED to the team's + shared default member budget, not given a private copy of it. - Cloning (rather than sharing the same budget row) is what lets admins later - edit one member's budget without mutating every other member's budget. + Linking is what makes a later ``/team/update team_member_budget=...`` + reach existing members: the auth check reads the budget row behind the + membership, so a private clone would freeze the member at the old cap. + Per-member isolation is handled by ``/team/member_update`` cloning the + shared row on first write. """ from litellm.proxy._types import LitellmUserRoles @@ -178,7 +181,6 @@ async def test_add_new_member_clones_default_team_budget_id(): test_user_id = "test_user_123" test_team_id = "test_team_456" test_default_budget_id = "default_budget_789" - test_cloned_budget_id = "cloned_budget_xyz" test_admin_name = "test_admin" new_member = Member(user_id=test_user_id, role="user") @@ -202,36 +204,19 @@ async def test_add_new_member_clones_default_team_budget_id(): return_value=mock_user_response ) - # Mock the default budget row fetched for cloning. mock_default_budget_row = MagicMock() - mock_default_budget_row.model_dump.return_value = { - "budget_id": test_default_budget_id, - "max_budget": 100.0, - "soft_budget": None, - "max_parallel_requests": None, - "tpm_limit": 1000, - "rpm_limit": None, - "model_max_budget": None, - "budget_duration": "1d", - "allowed_models": [], - } + mock_default_budget_row.budget_id = test_default_budget_id mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( return_value=mock_default_budget_row ) - - # Mock the cloned budget row that .create() returns. - mock_cloned_budget_row = MagicMock() - mock_cloned_budget_row.budget_id = test_cloned_budget_id - mock_prisma_client.db.litellm_budgettable.create = AsyncMock( - return_value=mock_cloned_budget_row - ) + mock_prisma_client.db.litellm_budgettable.create = AsyncMock() # Mock the team membership creation mock_team_membership_response = MagicMock() mock_team_membership_response.model_dump.return_value = { "team_id": test_team_id, "user_id": test_user_id, - "budget_id": test_cloned_budget_id, + "budget_id": test_default_budget_id, "litellm_budget_table": None, } mock_prisma_client.db.litellm_teammembership.create = AsyncMock( @@ -251,33 +236,67 @@ async def test_add_new_member_clones_default_team_budget_id(): assert result_user is not None assert result_user.user_id == test_user_id - # Membership should be linked to the new cloned budget, not the shared default. + # Membership points at the shared default row itself. assert result_team_membership is not None - assert result_team_membership.budget_id == test_cloned_budget_id - assert result_team_membership.budget_id != test_default_budget_id + assert result_team_membership.budget_id == test_default_budget_id mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() mock_prisma_client.db.litellm_teammembership.create.assert_called_once() - # The clone must have happened: find_unique on the default, create for the clone. + # The default is only checked for existence; no private budget row is created. mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( where={"budget_id": test_default_budget_id} ) - mock_prisma_client.db.litellm_budgettable.create.assert_called_once() - cloned_create_data = ( - mock_prisma_client.db.litellm_budgettable.create.call_args.kwargs["data"] - ) - # Cloned values from the default budget row - assert cloned_create_data["max_budget"] == 100.0 - assert cloned_create_data["tpm_limit"] == 1000 - assert cloned_create_data["budget_duration"] == "1d" - assert cloned_create_data["created_by"] == user_api_key_dict.user_id + mock_prisma_client.db.litellm_budgettable.create.assert_not_called() team_membership_call_args = ( mock_prisma_client.db.litellm_teammembership.create.call_args ) create_data = team_membership_call_args.kwargs["data"] - assert create_data["budget_id"] == test_cloned_budget_id + assert create_data["budget_id"] == test_default_budget_id + + +@pytest.mark.asyncio +async def test_add_new_member_no_budget_when_default_budget_row_is_missing(): + """If team metadata still names a default member budget whose row was + deleted, the member must get no budget rather than a dangling link that + the membership foreign key would reject.""" + from litellm.proxy._types import LitellmUserRoles + + new_member = Member(user_id="missing-default-user", role="user") + user_api_key_dict = UserAPIKeyAuth( + user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + mock_prisma_client = AsyncMock() + mock_user_response = MagicMock() + mock_user_response.model_dump.return_value = { + "user_id": "missing-default-user", + "user_email": None, + "teams": ["team-md"], + "user_role": "internal_user", + } + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mock_user_response + ) + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_budgettable.create = AsyncMock() + mock_prisma_client.db.litellm_teammembership.create = AsyncMock() + + _, result_team_membership = await add_new_member( + new_member=new_member, + max_budget_in_team=None, + prisma_client=mock_prisma_client, + team_id="team-md", + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name="test_admin", + default_team_budget_id="deleted-default", + ) + + assert result_team_membership is None + mock_prisma_client.db.litellm_budgettable.create.assert_not_called() + mock_prisma_client.db.litellm_teammembership.create.assert_not_called() @pytest.mark.asyncio @@ -636,18 +655,17 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): @pytest.mark.asyncio -async def test_add_new_member_with_user_email_clones_default_budget(): +async def test_add_new_member_with_user_email_links_default_budget(): """ Test add_new_member with user_email instead of user_id and a team default - budget. The default budget should be CLONED into a new private row for - this user, not shared with other members of the team. + budget. The membership must link the shared default row so team-level + budget updates keep applying to this member. """ from litellm.proxy._types import LitellmUserRoles test_user_email = "test@example.com" test_team_id = "test_team_456" test_default_budget_id = "default_budget_789" - test_cloned_budget_id = "cloned_budget_for_email_user" test_admin_name = "test_admin" new_member = Member(user_email=test_user_email, role="user") @@ -669,35 +687,18 @@ async def test_add_new_member_with_user_email_clones_default_budget(): } mock_prisma_client.insert_data = AsyncMock(return_value=mock_user_response) - # Default budget that will be cloned mock_default_budget_row = MagicMock() - mock_default_budget_row.model_dump.return_value = { - "budget_id": test_default_budget_id, - "max_budget": 25.0, - "soft_budget": None, - "max_parallel_requests": None, - "tpm_limit": None, - "rpm_limit": None, - "model_max_budget": None, - "budget_duration": None, - "allowed_models": [], - } + mock_default_budget_row.budget_id = test_default_budget_id mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock( return_value=mock_default_budget_row ) - - # Cloned budget result - mock_cloned_budget_row = MagicMock() - mock_cloned_budget_row.budget_id = test_cloned_budget_id - mock_prisma_client.db.litellm_budgettable.create = AsyncMock( - return_value=mock_cloned_budget_row - ) + mock_prisma_client.db.litellm_budgettable.create = AsyncMock() mock_team_membership_response = MagicMock() mock_team_membership_response.model_dump.return_value = { "team_id": test_team_id, "user_id": "generated_user_id", - "budget_id": test_cloned_budget_id, + "budget_id": test_default_budget_id, "litellm_budget_table": None, } mock_prisma_client.db.litellm_teammembership.create = AsyncMock( @@ -717,9 +718,8 @@ async def test_add_new_member_with_user_email_clones_default_budget(): assert result_user is not None assert result_user.user_email == test_user_email - # Membership should point at the cloned (private) budget, not the shared default. assert result_team_membership is not None - assert result_team_membership.budget_id == test_cloned_budget_id + assert result_team_membership.budget_id == test_default_budget_id mock_prisma_client.get_data.assert_called_once_with( key_val={"user_email": test_user_email}, @@ -733,11 +733,10 @@ async def test_add_new_member_with_user_email_clones_default_budget(): assert insert_data["user_email"] == test_user_email assert insert_data["teams"] == [test_team_id] - # Confirm the clone path ran mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( where={"budget_id": test_default_budget_id} ) - mock_prisma_client.db.litellm_budgettable.create.assert_called_once() + mock_prisma_client.db.litellm_budgettable.create.assert_not_called() @pytest.mark.asyncio From 9f990c4f8694d3f394f27be9dbcd47d8f49c4a5e Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:24:31 +0000 Subject: [PATCH 053/253] fix(router): validate routing_groups at save time and keep invalid DB groups from blocking SSO load Overlapping routing_groups persisted from the Admin UI raised inside Router._init_routing_groups during the DB config reconcile, which skipped loading SSO, guardrails and the other DB-backed settings while leaving the proxy healthy. /config/update now returns 400 for overlapping models, duplicate names, the reserved default name and unknown strategies before writing, the Router builds every group selector before replacing its state so a rejected update keeps the previous groups routing, and the proxy applies routing_groups separately from the other router settings so an already persisted invalid value is logged and skipped instead of aborting the reconcile. The Admin UI modal blocks picking a model another group owns. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_server.py | 30 +++- litellm/router.py | 142 ++++++++---------- litellm/router_utils/routing_groups.py | 95 ++++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 111 ++++++++++++++ .../test_router_routing_groups.py | 113 ++++++++++++++ .../routing_groups/RoutingGroupModal.test.tsx | 14 ++ .../routing_groups/RoutingGroupModal.tsx | 17 ++- .../src/components/routing_groups/index.tsx | 9 +- .../routing_groups/modelOwnership.test.ts | 33 ++++ .../routing_groups/modelOwnership.ts | 18 +++ 10 files changed, 501 insertions(+), 81 deletions(-) create mode 100644 litellm/router_utils/routing_groups.py create mode 100644 ui/litellm-dashboard/src/components/routing_groups/modelOwnership.test.ts create mode 100644 ui/litellm-dashboard/src/components/routing_groups/modelOwnership.ts diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d7964556531..b52d007d9db 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -143,6 +143,7 @@ from litellm.router_utils.auto_router_tuning_baseline import ( snapshot_tuning_baselines, tuning_limit_violation, ) +from litellm.router_utils.routing_groups import parse_routing_groups from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.utils import ( ModelResponse, @@ -759,6 +760,7 @@ from litellm.types.router import ( ClassifierPlugin, DeploymentTypedDict, RouterGeneralSettings, + RoutingGroup, RoutingPlugin, SearchToolTypedDict, updateDeployment, @@ -6896,7 +6898,27 @@ class ProxyConfig: combined_router_settings = db_router_settings.param_value if combined_router_settings: - llm_router.update_settings(**combined_router_settings) + self._apply_router_settings(llm_router, combined_router_settings) + + @staticmethod + def _apply_router_settings(llm_router: Router, router_settings: Mapping[str, object]) -> None: + """ + `routing_groups` is applied on its own so a value persisted before + save-time validation existed cannot abort the reconcile that also loads + SSO, guardrails and the other DB-backed settings. The router keeps the + groups it already holds when the new value is rejected. + """ + llm_router.update_settings(**{k: v for k, v in router_settings.items() if k != "routing_groups"}) + if "routing_groups" not in router_settings: + return + try: + llm_router.update_settings(routing_groups=router_settings["routing_groups"]) + except (TypeError, ValueError) as invalid_groups: + verbose_proxy_logger.error( + "Ignoring invalid router_settings.routing_groups from config/DB, all other router settings still " + "apply. Fix the routing groups in the Admin UI to load them: %s", + invalid_groups, + ) def _add_general_settings_from_db_config( self, config_data: dict, general_settings: dict, proxy_logging_obj: ProxyLogging @@ -16903,6 +16925,12 @@ async def update_config( ) }, ) + try: + parse_routing_groups( + TypeAdapter(list[RoutingGroup] | None).validate_python(raw_router_settings.get("routing_groups")) + ) + except (ValidationError, ValueError) as invalid_groups: + raise HTTPException(status_code=400, detail={"error": str(invalid_groups)}) if prisma_client is None: raise Exception("No DB Connected") diff --git a/litellm/router.py b/litellm/router.py index d531072530b..e31689ef447 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -218,6 +218,7 @@ from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, increment_deployment_successes_for_current_minute, ) +from litellm.router_utils.routing_groups import parse_routing_groups, validate_routing_strategy from litellm.scheduler import FlowItem, Scheduler from litellm.types.llms.openai import ( AllMessageValues, @@ -1244,20 +1245,9 @@ class Router: return strategy.value return strategy - def _validate_routing_strategy(self, routing_strategy: RoutingStrategy | str | None) -> None: - # See: https://github.com/BerriAI/litellm/issues/11330 - valid_strategy_strings: Final = ["simple-shuffle", "lar1"] + [s.value for s in RoutingStrategy] - if routing_strategy is None: - return - is_valid_string: Final = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings - is_valid_enum: Final = isinstance(routing_strategy, RoutingStrategy) - if not is_valid_string and not is_valid_enum: - raise ValueError( - f"Invalid routing_strategy: '{routing_strategy}'. " - f"Valid options: {valid_strategy_strings}. " - f"Check 'router_settings.routing_strategy' in your config.yaml " - f"or the 'routing_strategy' parameter if using the Router SDK directly." - ) + @staticmethod + def _validate_routing_strategy(routing_strategy: RoutingStrategy | str | None) -> None: + validate_routing_strategy(routing_strategy) def _build_strategy_selector( self, @@ -1274,11 +1264,6 @@ class Router: match self._normalize_strategy(strategy): case RoutingStrategy.LEAST_BUSY.value: selector = LeastBusyLoggingHandler(router_cache=self.cache) - if register_callbacks: - if isinstance(litellm.input_callback, list): - litellm.logging_callback_manager.add_litellm_input_callback(selector) - else: - litellm.input_callback = [selector] case RoutingStrategy.USAGE_BASED_ROUTING.value: selector = LowestTPMLoggingHandler( router_cache=self.cache, @@ -1302,11 +1287,21 @@ class Router: case _: pass - if selector is not None and register_callbacks and isinstance(litellm.callbacks, list): - litellm.logging_callback_manager.add_litellm_callback(selector) + if selector is not None and register_callbacks: + self._register_router_selector(selector) return selector + @staticmethod + def _register_router_selector(selector: RouterStrategySelector) -> None: + if isinstance(selector, LeastBusyLoggingHandler): + if isinstance(litellm.input_callback, list): + litellm.logging_callback_manager.add_litellm_input_callback(selector) + else: + litellm.input_callback = [selector] + if isinstance(litellm.callbacks, list): + litellm.logging_callback_manager.add_litellm_callback(selector) + def _unregister_router_selectors(self, selectors: Sequence[object]) -> None: """ Drop router-owned strategy selectors from litellm's global callback @@ -1397,75 +1392,69 @@ class Router: at most one explicit group. Constructs per-group strategy selectors so groups with different `routing_strategy_args` track independent state. + Validation and selector construction run to completion before any + router state changes, so a rejected input raises with the previously + loaded groups still routing. + Models not claimed by any explicit group are served by the implicit `"default"` group, whose selectors are the `self._logger` attributes set up in `routing_strategy_init`. """ - group_selectors: Final[Mapping[str, Mapping[str, RouterStrategySelector]]] = getattr( - self, "_group_selectors", {} - ) - self._unregister_router_selectors([sel for selectors in group_selectors.values() for sel in selectors.values()]) - - self._routing_groups: dict[str, RoutingGroup] = {} - self._model_to_group: dict[str, str] = {} - self._group_selectors: dict[str, dict[str, RouterStrategySelector]] = {} - self._invalidate_model_group_info_cache() - self._invalidate_access_groups_cache() - if not groups_input: + self._replace_routing_groups(()) return - known_model_names: Final = {m.get("model_name") for m in (self.model_list or []) if m.get("model_name")} + known_model_names: Final = frozenset(m["model_name"] for m in (self.model_list or ()) if m.get("model_name")) + groups: Final = parse_routing_groups(groups_input, known_model_names=known_model_names) - seen_group_names: Final[set] = set() - for raw in groups_input: - group = raw if isinstance(raw, RoutingGroup) else RoutingGroup(**raw) - - if not group.group_name: - raise ValueError("routing_groups: group_name must be non-empty.") - if group.group_name == "default": - raise ValueError("routing_groups: 'default' is reserved for the implicit fallback group.") - if group.group_name in known_model_names or group.group_name in (self.model_group_alias or {}): + alias_names: Final = frozenset(self.model_group_alias or ()) + for group in groups: + if group.group_name in known_model_names or group.group_name in alias_names: verbose_router_logger.warning( "routing_groups: group_name '%s' is shadowed by an existing model_name or model_group_alias; " "the group's strategy still applies to its members, but the name is not callable until renamed.", group.group_name, ) - if group.group_name in seen_group_names: - raise ValueError( - f"routing_groups: group names must be unique, duplicate group_name '{group.group_name}'." - ) - seen_group_names.add(group.group_name) - self._validate_routing_strategy(group.routing_strategy) - - for model_name in group.models: - if model_name in self._model_to_group: - raise ValueError( - f"routing_groups: model_name '{model_name}' appears in " - f"both '{self._model_to_group[model_name]}' and " - f"'{group.group_name}'. Each model may belong to at most one group." - ) - if known_model_names and model_name not in known_model_names: - verbose_router_logger.warning( - "routing_groups: model_name '%s' (group '%s') is not in model_list; " - "the group entry will only take effect once a deployment with that " - "model_name is added.", - model_name, - group.group_name, - ) - self._model_to_group[model_name] = group.group_name - - self._routing_groups[group.group_name] = group - - strategy_value = self._normalize_strategy(group.routing_strategy) or "" - group_selector = self._build_strategy_selector( - strategy=group.routing_strategy, - routing_strategy_args=group.routing_strategy_args or {}, + built: Final = tuple( + ( + group, + self._build_strategy_selector( + strategy=group.routing_strategy, + routing_strategy_args=group.routing_strategy_args or {}, + register_callbacks=False, + ), ) - self._group_selectors[group.group_name] = ( - {strategy_value: group_selector} if group_selector is not None else {} + for group in groups + ) + self._replace_routing_groups(built) + + def _replace_routing_groups( + self, + built: tuple[tuple[RoutingGroup, RouterStrategySelector | None], ...], + ) -> None: + previous_selectors: Final[Mapping[str, Mapping[str, RouterStrategySelector]]] = getattr( + self, "_group_selectors", {} + ) + self._unregister_router_selectors( + tuple(sel for selectors in previous_selectors.values() for sel in selectors.values()) + ) + for _, selector in built: + if selector is not None: + self._register_router_selector(selector) + + self._routing_groups: dict[str, RoutingGroup] = {group.group_name: group for group, _ in built} + self._model_to_group: dict[str, str] = { + model_name: group.group_name for group, _ in built for model_name in group.models + } + self._group_selectors: dict[str, dict[str, RouterStrategySelector]] = { + group.group_name: ( + {} if selector is None else {self._normalize_strategy(group.routing_strategy) or "": selector} ) + for group, selector in built + } + self._invalidate_model_group_info_cache() + self._invalidate_access_groups_cache() def get_routing_group(self, model_name: str) -> RoutingGroup | None: """ @@ -12032,7 +12021,6 @@ class Router: _casted_value = int(kwargs[var]) setattr(self, var, _casted_value) elif var == "routing_groups": - self._routing_groups_input = kwargs[var] rebuild_routing_groups = True elif var == "optional_pre_call_checks": self.set_optional_pre_call_checks(kwargs[var]) @@ -12073,7 +12061,9 @@ class Router: self._apply_updated_routing_strategy_args() if rebuild_routing_groups: - self._init_routing_groups(self._routing_groups_input) + routing_groups_input: Final = kwargs.get("routing_groups", self._routing_groups_input) + self._init_routing_groups(routing_groups_input) + self._routing_groups_input = routing_groups_input verbose_router_logger.debug("Updated Router settings: %s", self.get_settings()) def _get_client(self, deployment, kwargs, client_type=None): diff --git a/litellm/router_utils/routing_groups.py b/litellm/router_utils/routing_groups.py new file mode 100644 index 00000000000..c9b3b205bf1 --- /dev/null +++ b/litellm/router_utils/routing_groups.py @@ -0,0 +1,95 @@ +""" +Validation for `router_settings.routing_groups`, shared by the Router and the +proxy's config-update endpoint so a config the UI saves cannot be one the +runtime refuses to load. +""" + +from collections.abc import Sequence +from typing import Final + +from litellm._logging import verbose_router_logger +from litellm.types.router import RoutingGroup, RoutingStrategy + + +def validate_routing_strategy(routing_strategy: RoutingStrategy | str | None) -> None: + """ + Raises `ValueError` unless `routing_strategy` is a known strategy or None. + + See: https://github.com/BerriAI/litellm/issues/11330 + """ + if routing_strategy is None: + return + + valid_strategy_strings: Final = ("simple-shuffle", "lar1", *(s.value for s in RoutingStrategy)) + is_valid_string: Final = isinstance(routing_strategy, str) and routing_strategy in valid_strategy_strings + is_valid_enum: Final = isinstance(routing_strategy, RoutingStrategy) + if not is_valid_string and not is_valid_enum: + raise ValueError( + f"Invalid routing_strategy: '{routing_strategy}'. " + f"Valid options: {list(valid_strategy_strings)}. " + f"Check 'router_settings.routing_strategy' in your config.yaml " + f"or the 'routing_strategy' parameter if using the Router SDK directly." + ) + + +def parse_routing_groups( + groups_input: Sequence[RoutingGroup | dict] | None, + known_model_names: frozenset[str] = frozenset(), +) -> tuple[RoutingGroup, ...]: + """ + Parses and validates `routing_groups`, raising `ValueError` on the first + problem found. Every check runs before the caller mutates any state, so an + invalid update can never leave a router holding a half-applied set of + groups. + """ + if not groups_input: + return () + + groups: Final = tuple(raw if isinstance(raw, RoutingGroup) else RoutingGroup(**raw) for raw in groups_input) + + if any(not group.group_name for group in groups): + raise ValueError("routing_groups: group_name must be non-empty.") + + if any(group.group_name == "default" for group in groups): + raise ValueError("routing_groups: 'default' is reserved for the implicit fallback group.") + + names: Final = tuple(group.group_name for group in groups) + duplicate_names: Final = frozenset(name for name in names if names.count(name) > 1) + if duplicate_names: + raise ValueError(f"routing_groups: group names must be unique, duplicate group_name '{min(duplicate_names)}'.") + + for group in groups: + validate_routing_strategy(group.routing_strategy) + + owners_by_model: Final = tuple( + (model_name, tuple(group.group_name for group in groups if model_name in group.models)) + for model_name in dict.fromkeys(model_name for group in groups for model_name in group.models) + ) + conflicts: Final = tuple( + f"model_name '{model_name}' appears in {' and '.join(repr(owner) for owner in owners)}" + for model_name, owners in owners_by_model + if len(owners) > 1 + ) + if conflicts: + raise ValueError(f"routing_groups: {'; '.join(conflicts)}. Each model may belong to at most one group.") + + unknown_models: Final = ( + tuple( + (model_name, group.group_name) + for group in groups + for model_name in group.models + if model_name not in known_model_names + ) + if known_model_names + else () + ) + for model_name, group_name in unknown_models: + verbose_router_logger.warning( + "routing_groups: model_name '%s' (group '%s') is not in model_list; " + "the group entry will only take effect once a deployment with that " + "model_name is added.", + model_name, + group_name, + ) + + return groups diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 6f55449abab..23aa8e9df55 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4959,6 +4959,71 @@ async def test_add_router_settings_from_db_config_merge_logic(): assert combined_settings["nested_config"] == expected_nested +def _routing_groups_router(): + from litellm import Router + + return Router( + model_list=[ + {"model_name": "m1", "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"}}, + {"model_name": "m2", "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-test"}}, + ], + routing_groups=[{"group_name": "g1", "models": ["m1"], "routing_strategy": "latency-based-routing"}], + ) + + +@pytest.mark.asyncio +async def test_invalid_db_routing_groups_do_not_abort_other_router_settings(): + """Regression: an overlapping routing_groups value persisted in DB used to raise out of + _add_router_settings_from_db_config, which skipped SSO / guardrail loading downstream.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.proxy_server import ProxyConfig + + router = _routing_groups_router() + mock_db_config = MagicMock() + mock_db_config.param_value = { + "num_retries": 7, + "routing_groups": [ + {"group_name": "g1", "models": ["m1"], "routing_strategy": "latency-based-routing"}, + {"group_name": "g2", "models": ["m1"], "routing_strategy": "least-busy"}, + ], + } + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + + await ProxyConfig()._add_router_settings_from_db_config( + config_data={}, llm_router=router, prisma_client=mock_prisma_client + ) + + assert router.num_retries == 7 + assert router._model_to_group == {"m1": "g1"} + assert router._get_routing_context("m1", None)[0] == "latency-based-routing" + + +@pytest.mark.asyncio +async def test_valid_db_routing_groups_still_replace_router_groups(): + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy.proxy_server import ProxyConfig + + router = _routing_groups_router() + mock_db_config = MagicMock() + mock_db_config.param_value = { + "num_retries": 7, + "routing_groups": [{"group_name": "g2", "models": ["m2"], "routing_strategy": "least-busy"}], + } + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + + await ProxyConfig()._add_router_settings_from_db_config( + config_data={}, llm_router=router, prisma_client=mock_prisma_client + ) + + assert router.num_retries == 7 + assert router._model_to_group == {"m2": "g2"} + assert router._get_routing_context("m2", None)[0] == "least-busy" + + @pytest.mark.asyncio async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_config_fallbacks(): """ @@ -9224,6 +9289,52 @@ def test_update_config_writes_only_sent_section(_update_config_setup): restore() +def test_update_config_rejects_overlapping_routing_groups_before_writing(_update_config_setup): + """Regression: overlapping groups were persisted and only failed at router reload, where the + failure took SSO and the other DB-backed settings down with it.""" + existing_groups = [{"group_name": "g1", "models": ["m1"], "routing_strategy": "least-busy"}] + client, prisma, restore = _update_config_setup( + initial_rows={"router_settings": {"num_retries": 2, "routing_groups": existing_groups}} + ) + try: + resp = client.post( + "/config/update", + json={ + "router_settings": { + "routing_groups": [ + *existing_groups, + {"group_name": "g2", "models": ["m1"], "routing_strategy": "latency-based-routing"}, + ] + } + }, + ) + assert resp.status_code == 400 + assert "'m1' appears in 'g1' and 'g2'" in resp.text + assert prisma.db.litellm_config.upsert_calls == [] + assert prisma.db.litellm_config.rows["router_settings"]["routing_groups"] == existing_groups + finally: + restore() + + +def test_update_config_accepts_disjoint_routing_groups(_update_config_setup): + client, prisma, restore = _update_config_setup(initial_rows={"router_settings": {"num_retries": 2}}) + groups = [ + {"group_name": "g1", "models": ["m1"], "routing_strategy": "least-busy"}, + {"group_name": "g2", "models": ["m2"], "routing_strategy": "latency-based-routing"}, + ] + try: + resp = client.post("/config/update", json={"router_settings": {"routing_groups": groups}}) + assert resp.status_code == 200 + stored = prisma.db.litellm_config.rows["router_settings"] + assert stored["num_retries"] == 2 + assert [(g["group_name"], g["models"]) for g in stored["routing_groups"]] == [ + ("g1", ["m1"]), + ("g2", ["m2"]), + ] + finally: + restore() + + def test_update_config_env_var_round_trip_not_double_encrypted(_update_config_setup, monkeypatch): """Endpoint-level regression for the /config/update double-encryption bug. diff --git a/tests/test_litellm/router_strategy/test_router_routing_groups.py b/tests/test_litellm/router_strategy/test_router_routing_groups.py index 25b657b8cd0..99230c1f71c 100644 --- a/tests/test_litellm/router_strategy/test_router_routing_groups.py +++ b/tests/test_litellm/router_strategy/test_router_routing_groups.py @@ -13,6 +13,7 @@ from collections.abc import Callable from unittest.mock import patch import pytest +from pydantic import ValidationError import litellm from litellm import Router @@ -806,6 +807,118 @@ def test_strategy_reinit_unregisters_override_selectors(): assert router._get_override_strategy_selector("latency-based-routing") is router.lowestlatency_logger +def _single_latency_group(): + return [{"group_name": "g1", "models": ["filtered-model"], "routing_strategy": "latency-based-routing"}] + + +def _assert_still_routes_with_original_group(router, selector): + assert list(router._routing_groups) == ["g1"] + assert router._model_to_group == {"filtered-model": "g1"} + assert router._group_selectors["g1"]["latency-based-routing"] is selector + assert router._get_routing_context("filtered-model", None) == ("latency-based-routing", selector) + assert sum(1 for cb in litellm.callbacks if cb is selector) == 1 + + +def test_failed_routing_groups_update_keeps_previous_groups(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router(routing_groups=_single_latency_group()) + selector = router._group_selectors["g1"]["latency-based-routing"] + + with pytest.raises(ValueError, match="appears in"): + router.update_settings( + routing_groups=[ + *_single_latency_group(), + {"group_name": "g2", "models": ["filtered-model"], "routing_strategy": "least-busy"}, + ], + ) + + _assert_still_routes_with_original_group(router, selector) + assert sum(1 for cb in litellm.callbacks if type(cb) is not type(selector)) == 0 + assert litellm.input_callback == [] + + +def test_failed_routing_groups_update_does_not_poison_later_strategy_changes(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router(routing_groups=_single_latency_group()) + + with pytest.raises(ValueError, match="appears in"): + router.update_settings( + routing_groups=[ + *_single_latency_group(), + {"group_name": "g2", "models": ["filtered-model"], "routing_strategy": "least-busy"}, + ], + ) + + router.update_settings(routing_strategy="least-busy") + + assert list(router._routing_groups) == ["g1"] + assert [g["group_name"] for g in router.get_settings()["routing_groups"]] == ["g1"] + + +def test_overlap_error_names_every_conflicting_model(): + with pytest.raises(ValueError, match="appears in") as exc_info: + _build_router( + routing_groups=[ + { + "group_name": "g1", + "models": ["filtered-model", "other-model"], + "routing_strategy": "latency-based-routing", + }, + { + "group_name": "g2", + "models": ["filtered-model", "other-model"], + "routing_strategy": "least-busy", + }, + ], + ) + message = str(exc_info.value) + assert "'filtered-model' appears in 'g1' and 'g2'" in message + assert "'other-model' appears in 'g1' and 'g2'" in message + + +def test_invalid_group_strategy_keeps_previous_groups(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router(routing_groups=_single_latency_group()) + selector = router._group_selectors["g1"]["latency-based-routing"] + + with pytest.raises(ValueError, match="Invalid routing_strategy"): + router.update_settings( + routing_groups=[ + {"group_name": "g2", "models": ["other-model"], "routing_strategy": "not-a-real-strategy"}, + ], + ) + + _assert_still_routes_with_original_group(router, selector) + + +def test_unbuildable_group_selector_keeps_previous_groups(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router(routing_groups=_single_latency_group()) + selector = router._group_selectors["g1"]["latency-based-routing"] + + with pytest.raises(ValidationError, match="ttl"): + router.update_settings( + routing_groups=[ + {"group_name": "g0", "models": ["other-model"], "routing_strategy": "least-busy"}, + *_single_latency_group(), + { + "group_name": "g2", + "models": ["other-model-2"], + "routing_strategy": "latency-based-routing", + "routing_strategy_args": {"ttl": "not-a-number"}, + }, + ], + ) + + _assert_still_routes_with_original_group(router, selector) + assert litellm.callbacks == [selector] + assert litellm.input_callback == [] + + def test_override_selectors_are_not_registered_process_wide(monkeypatch): monkeypatch.setattr(litellm, "callbacks", []) monkeypatch.setattr(litellm, "input_callback", []) diff --git a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.test.tsx b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.test.tsx index 322d90b24c5..e376c551923 100644 --- a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.test.tsx +++ b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.test.tsx @@ -59,6 +59,7 @@ const renderModal = (overrides: Partial { expect(onSubmit.mock.calls[0][0]).toStrictEqual(expected); }); + it("blocks a model another group already claims", async () => { + const user = userEvent.setup(); + const { onSubmit } = renderModal({ groupNameByModel: { "gpt-4o": "cheap" } }); + + await typeName(user, "security"); + await pickModels(user, "gpt-4o"); + await pickStrategy(user, "latency-based-routing"); + await save(user, "Create Group"); + + expect(await screen.findByText(/Already claimed: gpt-4o/)).toBeInTheDocument(); + expect(onSubmit).not.toHaveBeenCalled(); + }); + it("describes the selected strategy", async () => { renderModal(); diff --git a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx index 5865c59d8bd..1057cc6ca16 100644 --- a/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx +++ b/ui/litellm-dashboard/src/components/routing_groups/RoutingGroupModal.tsx @@ -29,6 +29,7 @@ import { toRoutingGroupFormValues, } from "./routingGroupPayload"; import type { RoutingGroup } from "./types"; +import { modelConflictError } from "./modelOwnership"; import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { Button } from "@/components/ui/button"; @@ -40,6 +41,7 @@ interface RoutingGroupModalProps { strategyDescriptions: Record; modelOptions: string[]; existingGroupNames: string[]; + groupNameByModel: Record; onClose: () => void; onSubmit: (group: RoutingGroup) => Promise | void; saving?: boolean; @@ -57,6 +59,7 @@ const RoutingGroupModal: React.FC = ({ strategyDescriptions, modelOptions, existingGroupNames, + groupNameByModel, onClose, onSubmit, saving, @@ -77,12 +80,20 @@ const RoutingGroupModal: React.FC = ({ .min(1, "Group name is required") .max(GROUP_NAME_MAX_LENGTH, `Must be ${GROUP_NAME_MAX_LENGTH} characters or fewer`) .refine((value) => !reservedNames.has(value.toLowerCase()), "A group with this name already exists"), - models: z.array(z.string()).min(1, "Select at least one model"), + models: z + .array(z.string()) + .min(1, "Select at least one model") + .superRefine((models, ctx) => { + const conflict = modelConflictError(models, groupNameByModel); + if (conflict !== null) { + ctx.addIssue({ code: "custom", message: conflict }); + } + }), routing_strategy: z.string().min(1, "Strategy is required"), routing_strategy_args: z.string(), }; return z.object(shape); - }, [reservedNames]); + }, [reservedNames, groupNameByModel]); const form = useZodForm(schema, { defaultValues: toRoutingGroupFormValues(initialValue, availableStrategies) }); @@ -124,7 +135,7 @@ const RoutingGroupModal: React.FC = ({ control={form.control} name="models" label="Models" - description="Models from your model list that this group routes between." + description="Models from your model list that this group routes between. A model can only be in one group." > {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( diff --git a/ui/litellm-dashboard/src/components/routing_groups/index.tsx b/ui/litellm-dashboard/src/components/routing_groups/index.tsx index 7da581be6a9..17329d0b572 100644 --- a/ui/litellm-dashboard/src/components/routing_groups/index.tsx +++ b/ui/litellm-dashboard/src/components/routing_groups/index.tsx @@ -14,6 +14,7 @@ import RoutingGroupsTable from "./RoutingGroupsTable"; import RoutingGroupModal from "./RoutingGroupModal"; import { toast } from "@/lib/toast"; import type { RoutingGroup } from "./types"; +import { groupNameByModel } from "./modelOwnership"; import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; const RoutingGroups: React.FC = () => { @@ -30,7 +31,7 @@ const RoutingGroups: React.FC = () => { const [editingGroup, setEditingGroup] = useState(null); const [deletingGroup, setDeletingGroup] = useState(null); - const groups = data?.routingGroups ?? []; + const groups = useMemo(() => data?.routingGroups ?? [], [data?.routingGroups]); const filteredGroups = useMemo(() => { const q = searchQuery.trim().toLowerCase(); @@ -51,6 +52,11 @@ const RoutingGroups: React.FC = () => { const strategyDescriptions = routerFields?.routing_strategy_descriptions ?? {}; + const ownerByModel = useMemo( + () => groupNameByModel(groups, drawerMode === "edit" ? editingGroup?.group_name : undefined), + [groups, drawerMode, editingGroup], + ); + const modelOptions = useMemo(() => { const records = (modelHub?.data ?? []) as Array<{ model_group?: string }>; const names = records.map((r) => r.model_group).filter((n): n is string => Boolean(n)); @@ -160,6 +166,7 @@ const RoutingGroups: React.FC = () => { strategyDescriptions={strategyDescriptions} modelOptions={modelOptions} existingGroupNames={groups.map((g) => g.group_name)} + groupNameByModel={ownerByModel} onClose={() => setDrawerOpen(false)} onSubmit={handleSubmit} saving={saveMutation.isPending} diff --git a/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.test.ts b/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.test.ts new file mode 100644 index 00000000000..962a593910a --- /dev/null +++ b/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.test.ts @@ -0,0 +1,33 @@ +import { describe, expect, it } from "vitest"; + +import { groupNameByModel, modelConflictError } from "./modelOwnership"; +import type { RoutingGroup } from "./types"; + +const groups: RoutingGroup[] = [ + { group_name: "cheap", models: ["m1", "m2"], routing_strategy: "latency-based-routing" }, + { group_name: "security", models: ["m3"], routing_strategy: "least-busy" }, +]; + +describe("groupNameByModel", () => { + it("maps every claimed model to its owning group", () => { + expect(groupNameByModel(groups)).toEqual({ m1: "cheap", m2: "cheap", m3: "security" }); + }); + + it("excludes the group being edited so its own models stay selectable", () => { + expect(groupNameByModel(groups, "cheap")).toEqual({ m3: "security" }); + }); +}); + +describe("modelConflictError", () => { + it("passes models that no other group claims", () => { + expect(modelConflictError(["m4"], groupNameByModel(groups, "cheap"))).toBeNull(); + expect(modelConflictError(undefined, groupNameByModel(groups))).toBeNull(); + }); + + it("names every model already claimed by another group", () => { + const error = modelConflictError(["m1", "m3", "m4"], groupNameByModel(groups)); + expect(error).toBe( + 'Each model may belong to at most one group. Already claimed: m1 (in "cheap"), m3 (in "security")', + ); + }); +}); diff --git a/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.ts b/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.ts new file mode 100644 index 00000000000..c67ae77d066 --- /dev/null +++ b/ui/litellm-dashboard/src/components/routing_groups/modelOwnership.ts @@ -0,0 +1,18 @@ +import type { RoutingGroup } from "./types"; + +export const groupNameByModel = (groups: RoutingGroup[], excludeGroupName?: string): Record => + Object.fromEntries( + groups + .filter((group) => group.group_name !== excludeGroupName) + .flatMap((group) => group.models.map((model) => [model, group.group_name] as const)), + ); + +export const modelConflictError = ( + models: string[] | undefined, + ownerByModel: Record, +): string | null => { + const conflicts = (models ?? []).filter((model) => ownerByModel[model] !== undefined); + if (conflicts.length === 0) return null; + const detail = conflicts.map((model) => `${model} (in "${ownerByModel[model]}")`).join(", "); + return `Each model may belong to at most one group. Already claimed: ${detail}`; +}; From 3c972cb31f006e13d9e1fbb14054785cc15a6c7d Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 01:33:54 +0000 Subject: [PATCH 054/253] fix(proxy): apply source overrides to IPv4-mapped IPv6 sign-in sources Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 9 +++++---- tests/test_litellm/proxy/auth/test_login_utils.py | 2 ++ 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 393411d0670..4ab57ca461f 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -137,10 +137,14 @@ def _int_setting(settings: Mapping[str, object], key: str, default: int) -> int: def _parse_address(client_ip: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None: + """The address as it is limited and counted: an IPv4-mapped IPv6 address is its IPv4 address.""" try: - return ipaddress.ip_address(client_ip) + address: Final = ipaddress.ip_address(client_ip) except ValueError: return None + if isinstance(address, ipaddress.IPv6Address) and address.ipv4_mapped is not None: + return address.ipv4_mapped + return address def _parse_network(raw_range: str) -> _Network | None: @@ -183,9 +187,6 @@ def source_group(client_ip: str) -> str: if address is None: return client_ip if isinstance(address, ipaddress.IPv6Address): - mapped: Final = address.ipv4_mapped - if mapped is not None: - return str(mapped) return str(ipaddress.ip_network((address, IPV6_SOURCE_PREFIX_LENGTH), strict=False)) return str(address) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 65c388240e5..df5bd29f8ff 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -955,6 +955,8 @@ def test_source_overrides_pick_the_most_specific_matching_range(): assert _limit("203.0.1.1") == 100 assert _limit("192.0.2.1") == 7 assert _limit("198.51.100.1") == 7, "a garbage limit falls back to the default rather than a huge or zero budget" + assert _limit("::ffff:203.0.113.9") == 300, "a mapped address gets the limit of the IPv4 bucket it is counted in" + assert _limit("::ffff:203.0.113.10") == 200 def test_ipv6_sources_are_grouped_by_their_64_bit_prefix(): From 9afac6899538ef27976e2ce54dd3545d11e0cec8 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:43:04 +0000 Subject: [PATCH 055/253] test(router): cover _register_router_selector and _replace_routing_groups directly Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_router_routing_groups.py | 47 +++++++++++++++++++ 1 file changed, 47 insertions(+) diff --git a/tests/test_litellm/router_strategy/test_router_routing_groups.py b/tests/test_litellm/router_strategy/test_router_routing_groups.py index 99230c1f71c..506563a82fb 100644 --- a/tests/test_litellm/router_strategy/test_router_routing_groups.py +++ b/tests/test_litellm/router_strategy/test_router_routing_groups.py @@ -919,6 +919,53 @@ def test_unbuildable_group_selector_keeps_previous_groups(monkeypatch): assert litellm.input_callback == [] +def test_register_router_selector_wires_only_the_hooks_the_strategy_needs(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router() + least_busy = router._build_strategy_selector( + strategy="least-busy", routing_strategy_args={}, register_callbacks=False + ) + latency = router._build_strategy_selector( + strategy="latency-based-routing", routing_strategy_args={}, register_callbacks=False + ) + assert least_busy is not None and latency is not None + assert litellm.callbacks == [] and litellm.input_callback == [] + + router._register_router_selector(least_busy) + router._register_router_selector(latency) + + assert [cb for cb in litellm.callbacks if cb is least_busy or cb is latency] == [least_busy, latency] + assert litellm.input_callback == [least_busy] + + +def test_replace_routing_groups_swaps_state_and_callbacks_in_one_step(monkeypatch): + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "input_callback", []) + router = _build_router(routing_groups=_single_latency_group()) + old_selector = router._group_selectors["g1"]["latency-based-routing"] + new_selector = router._build_strategy_selector( + strategy="least-busy", routing_strategy_args={}, register_callbacks=False + ) + assert new_selector is not None + + router._replace_routing_groups( + ( + (RoutingGroup(group_name="g2", models=["other-model"], routing_strategy="least-busy"), new_selector), + (RoutingGroup(group_name="g3", models=["other-model-2"], routing_strategy="simple-shuffle"), None), + ) + ) + + assert list(router._routing_groups) == ["g2", "g3"] + assert router._model_to_group == {"other-model": "g2", "other-model-2": "g3"} + assert router._group_selectors == {"g2": {"least-busy": new_selector}, "g3": {}} + assert router._get_routing_context("other-model", None) == ("least-busy", new_selector) + assert router._get_routing_context("filtered-model", None)[0] == router.routing_strategy + assert all(cb is not old_selector for cb in litellm.callbacks) + assert sum(1 for cb in litellm.callbacks if cb is new_selector) == 1 + assert litellm.input_callback == [new_selector] + + def test_override_selectors_are_not_registered_process_wide(monkeypatch): monkeypatch.setattr(litellm, "callbacks", []) monkeypatch.setattr(litellm, "input_callback", []) From ab99be9dad023d7d5e25d3e9352f7954357d4963 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:43:21 +0000 Subject: [PATCH 056/253] test(team): cover team_member_budget propagation and per-member isolation end to end Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_management_helpers_utils.py | 172 +++++++++++++++--- 1 file changed, 148 insertions(+), 24 deletions(-) diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 00de5171aa7..7512a3dfae8 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -1,13 +1,17 @@ import json +from collections.abc import Mapping from datetime import datetime, timezone -from litellm._uuid import uuid -from unittest.mock import AsyncMock, MagicMock +from typing import Final +from unittest.mock import AsyncMock, MagicMock, patch import pytest - +import litellm +from litellm._uuid import uuid from litellm.proxy._types import ( + LiteLLM_BudgetTable, LiteLLM_TeamMembership, + LiteLLM_TeamTable, LiteLLM_UserTable, Member, UserAPIKeyAuth, @@ -165,19 +169,8 @@ async def test_management_otel_span_redacts_nested_submission_env_var_secrets( @pytest.mark.asyncio async def test_add_new_member_links_default_team_budget_id(): - """ - A member added without any per-member limit must be LINKED to the team's - shared default member budget, not given a private copy of it. - - Linking is what makes a later ``/team/update team_member_budget=...`` - reach existing members: the auth check reads the budget row behind the - membership, so a private clone would freeze the member at the old cap. - Per-member isolation is handled by ``/team/member_update`` cloning the - shared row on first write. - """ from litellm.proxy._types import LitellmUserRoles - # Setup test data test_user_id = "test_user_123" test_team_id = "test_team_456" test_default_budget_id = "default_budget_789" @@ -236,14 +229,12 @@ async def test_add_new_member_links_default_team_budget_id(): assert result_user is not None assert result_user.user_id == test_user_id - # Membership points at the shared default row itself. assert result_team_membership is not None assert result_team_membership.budget_id == test_default_budget_id mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() mock_prisma_client.db.litellm_teammembership.create.assert_called_once() - # The default is only checked for existence; no private budget row is created. mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( where={"budget_id": test_default_budget_id} ) @@ -258,9 +249,6 @@ async def test_add_new_member_links_default_team_budget_id(): @pytest.mark.asyncio async def test_add_new_member_no_budget_when_default_budget_row_is_missing(): - """If team metadata still names a default member budget whose row was - deleted, the member must get no budget rather than a dangling link that - the membership foreign key would reject.""" from litellm.proxy._types import LitellmUserRoles new_member = Member(user_id="missing-default-user", role="user") @@ -656,11 +644,6 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): @pytest.mark.asyncio async def test_add_new_member_with_user_email_links_default_budget(): - """ - Test add_new_member with user_email instead of user_id and a team default - budget. The membership must link the shared default row so team-level - budget updates keep applying to this member. - """ from litellm.proxy._types import LitellmUserRoles test_user_email = "test@example.com" @@ -739,6 +722,147 @@ async def test_add_new_member_with_user_email_links_default_budget(): mock_prisma_client.db.litellm_budgettable.create.assert_not_called() +class _FakeBudgetTable: + def __init__(self) -> None: + self.rows: dict[str, dict[str, object]] = {} + + def _record(self, budget_id: str) -> LiteLLM_BudgetTable: + row: Final = self.rows[budget_id] + return LiteLLM_BudgetTable(**{k: v for k, v in row.items() if k in LiteLLM_BudgetTable.model_fields}) + + async def create( + self, *, data: Mapping[str, object], include: Mapping[str, bool] | None = None + ) -> LiteLLM_BudgetTable: + budget_id: Final = str(data.get("budget_id") or uuid.uuid4()) + self.rows[budget_id] = {**data, "budget_id": budget_id} + return self._record(budget_id) + + async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_BudgetTable | None: + return self._record(where["budget_id"]) if where["budget_id"] in self.rows else None + + async def update(self, *, where: Mapping[str, str], data: Mapping[str, object]) -> LiteLLM_BudgetTable: + self.rows[where["budget_id"]] = {**self.rows[where["budget_id"]], **data} + return self._record(where["budget_id"]) + + +class _FakeMembershipTable: + def __init__(self, budgets: _FakeBudgetTable) -> None: + self.budgets: Final = budgets + self.budget_ids: dict[tuple[str, str], str | None] = {} + + def membership(self, team_id: str, user_id: str) -> LiteLLM_TeamMembership: + budget_id: Final = self.budget_ids[(team_id, user_id)] + return LiteLLM_TeamMembership( + user_id=user_id, + team_id=team_id, + budget_id=budget_id, + litellm_budget_table=self.budgets._record(budget_id) if budget_id is not None else None, + ) + + async def create(self, *, data: Mapping[str, str], include: Mapping[str, bool]) -> LiteLLM_TeamMembership: + self.budget_ids[(data["team_id"], data["user_id"])] = data["budget_id"] + return self.membership(data["team_id"], data["user_id"]) + + async def upsert(self, *, where: Mapping[str, Mapping[str, str]], data: Mapping[str, Mapping[str, object]]) -> None: + key: Final = where["user_id_team_id"] + connect: Final = data["update"]["litellm_budget_table"] + assert isinstance(connect, dict) + self.budget_ids[(key["team_id"], key["user_id"])] = connect["connect"]["budget_id"] + + +class _FakeUserTable: + async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]) -> LiteLLM_UserTable: + return LiteLLM_UserTable(user_id=where["user_id"], teams=list(data["create"].get("teams", []))) + + async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int: + return 1 + + +class _FakeDb: + def __init__(self) -> None: + self.litellm_budgettable: Final = _FakeBudgetTable() + self.litellm_teammembership: Final = _FakeMembershipTable(self.litellm_budgettable) + self.litellm_usertable: Final = _FakeUserTable() + + +@pytest.mark.asyncio +async def test_team_update_reaches_inherited_members_but_not_overridden_ones(): + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.auth.auth_checks import _check_team_member_budget + from litellm.proxy.management_endpoints.common_utils import _upsert_budget_and_membership + from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler + from litellm.proxy.utils import ProxyLogging + + db: Final = _FakeDb() + prisma_client: Final = MagicMock() + prisma_client.db = db + admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN) + team_id: Final = "team-shared-default" + default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0}) + team: Final = LiteLLM_TeamTable(team_id=team_id, metadata={"team_member_budget_id": default_budget.budget_id}) + + for user_id in ("inherits", "overridden"): + await add_new_member( + new_member=Member(user_id=user_id, role="user"), + max_budget_in_team=None, + prisma_client=prisma_client, + team_id=team_id, + user_api_key_dict=admin, + litellm_proxy_admin_name="admin", + default_team_budget_id=default_budget.budget_id, + ) + + await _upsert_budget_and_membership( + db, + team_id=team_id, + user_id="overridden", + existing_budget_id=default_budget.budget_id, + user_api_key_dict=admin, + budget_patch={"max_budget": 50.0}, + team_default_budget_id=default_budget.budget_id, + ) + assert db.litellm_teammembership.membership(team_id, "inherits").budget_id == default_budget.budget_id + assert db.litellm_teammembership.membership(team_id, "overridden").budget_id != default_budget.budget_id + assert db.litellm_budgettable.rows[default_budget.budget_id]["max_budget"] == 100.0 + + with patch( # test-quality-ok: update_budget reads this module global; no dependency injection seam exists + "litellm.proxy.proxy_server.prisma_client", prisma_client + ): + await TeamMemberBudgetHandler.upsert_team_member_budget_table( + team_table=team, + user_api_key_dict=admin, + updated_kv={}, + team_member_budget=1.0, + ) + + async def spend_from_membership(counter_key: str, fallback_spend: float, max_budget: float | None = None) -> float: + return fallback_spend + + async def check(user_id: str, spend: float) -> None: + membership: Final = db.litellm_teammembership.membership(team_id, user_id).model_copy(update={"spend": spend}) + with patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists + "litellm.proxy.proxy_server.get_current_spend", spend_from_membership + ): + await _check_team_member_budget( + team_object=team, + user_object=LiteLLM_UserTable(user_id=user_id), + valid_token=UserAPIKeyAuth(token="tok", user_id=user_id, team_id=team_id), + prisma_client=prisma_client, + user_api_key_cache=MagicMock(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=None), + team_membership=membership, + team_membership_loaded=True, + ) + + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await check("inherits", spend=2.0) + assert exc_info.value.max_budget == 1.0 + await check("overridden", spend=2.0) + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await check("overridden", spend=60.0) + assert exc_info.value.max_budget == 50.0 + + @pytest.mark.asyncio async def test_attach_object_permission_to_dict_with_object_permission_id(): """ From 1a749d84bdd66706bb41cafdba28e7a8b6a20fa9 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:56:00 +0000 Subject: [PATCH 057/253] fix(proxy): track project spend and enforce project budgets additively Project-scoped keys never wrote spend to LiteLLM_ProjectTable, so /project/info stayed at 0 and project budgets could not block. Wire the PROJECT entity through the spend queue, redis buffer, and db writer, reserve and increment a spend:project counter, reseed it from the project row, reset project spend in the budget cascade, and read the live counter in the project max budget check. Team member budgets keep gating project-scoped keys alongside the project budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 1 + litellm/proxy/auth/auth_checks.py | 34 ++-- .../proxy/common_utils/reset_budget_job.py | 24 +++ .../proxy/common_utils/user_api_key_cache.py | 10 ++ litellm/proxy/db/db_spend_update_writer.py | 79 ++++++++- .../redis_update_buffer.py | 7 + .../spend_update_queue.py | 4 + litellm/proxy/db/spend_counter_reseed.py | 5 + .../proxy/hooks/proxy_track_cost_callback.py | 6 + litellm/proxy/proxy_server.py | 30 ++++ .../spend_tracking/budget_reservation.py | 41 +++++ .../spend_tracking/spend_counter_batch.py | 11 +- litellm/repositories/prisma_protocols.py | 3 + litellm/repositories/unit_of_work.py | 2 + .../proxy/auth/test_auth_checks.py | 61 +++++++ .../common_utils/test_reset_budget_job.py | 30 +++- .../proxy/db/test_db_spend_update_writer.py | 76 +++++++++ .../proxy/db/test_spend_counter_reseed.py | 19 +++ .../hooks/test_proxy_track_cost_callback.py | 1 + .../test_spend_tracking_utils.py | 1 + .../proxy/test_budget_reservation.py | 156 ++++++++++++++++++ .../repositories/test_unit_of_work.py | 3 + 22 files changed, 583 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 321f8190f13..228a91ad446 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -5261,6 +5261,7 @@ class DBSpendUpdateTransactions(TypedDict): team_member_list_transactions: dict[str, float] | None org_list_transactions: dict[str, float] | None org_member_list_transactions: ReadOnly[dict[str, float] | None] + project_list_transactions: ReadOnly[dict[str, float] | None] tag_list_transactions: dict[str, float] | None agent_list_transactions: dict[str, float] | None model_access_group_list_transactions: ReadOnly[dict[str, float] | None] diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e3783c94dc7..ef17913d9ec 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -92,6 +92,8 @@ from litellm.proxy.common_utils.user_api_key_cache import ( model_access_group_registry_cache_key, model_access_group_spend_counter_key, object_permission_cache_key, + project_cache_key, + project_spend_counter_key, tag_cache_key, tag_registry_cache_key, team_membership_auth_cache_key, @@ -5586,16 +5588,22 @@ async def _project_max_budget_check( if project_object.litellm_budget_table is not None: max_budget = project_object.litellm_budget_table.max_budget - if ( - max_budget is not None - and project_object.spend is not None - and math.isfinite(max_budget) - and project_object.spend > max_budget - ): + if max_budget is None or not math.isfinite(max_budget): + return + + from litellm.proxy.proxy_server import get_current_spend + + project_spend: Final = await get_current_spend( + counter_key=project_spend_counter_key(project_object.project_id), + fallback_spend=project_object.spend or 0.0, + max_budget=max_budget, + ) + + if project_spend >= max_budget: if valid_token: call_info: Final = CallInfo( token=valid_token.token, - spend=project_object.spend, + spend=project_spend, max_budget=max_budget, user_id=valid_token.user_id, team_id=valid_token.team_id, @@ -5611,9 +5619,9 @@ async def _project_max_budget_check( ) raise litellm.BudgetExceededError( - current_cost=project_object.spend, + current_cost=project_spend, max_budget=max_budget, - message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_object.spend}, Max budget: {max_budget}", + message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_spend}, Max budget: {max_budget}", entity_type=Litellm_EntityType.PROJECT.value, entity_id=project_object.project_id, ) @@ -5663,10 +5671,6 @@ async def _project_soft_budget_check( ) -def _project_cache_key(project_id: str) -> str: - return f"project_id:{project_id}" - - async def get_project_object( project_id: str, prisma_client: PrismaClient | None, @@ -5684,7 +5688,7 @@ async def get_project_object( return None # Check cache first - cache_key: Final = _project_cache_key(project_id) + cache_key: Final = project_cache_key(project_id) deserialized_project: Final = await user_api_key_cache.async_get_cache( key=cache_key, model_type=LiteLLM_ProjectTableCachedObj, @@ -5726,7 +5730,7 @@ async def delete_cached_project_object( from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast await evict_and_broadcast( - cache_keys=(_project_cache_key(project_id),), + cache_keys=(project_cache_key(project_id),), user_api_key_cache=user_api_key_cache, ) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index acb51e73daf..2baefa89943 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -41,6 +41,8 @@ from litellm.proxy.common_utils.user_api_key_cache import ( end_user_cache_key, model_access_group_cache_key, model_access_group_spend_counter_key, + project_cache_key, + project_spend_counter_key, tag_cache_key, ) from litellm.proxy.db.budget_window_spend_writer import roll_window_spend_row @@ -49,6 +51,7 @@ 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 SpendLinkedTable +from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( EndUserRepository, ModelAccessGroupBudgetRepository, @@ -115,6 +118,11 @@ class _ModelAccessGroupRow(_BudgetLinkedRow, Protocol): def access_group_name(self) -> str: ... +class _ProjectRow(_BudgetLinkedRow, Protocol): + @property + def project_id(self) -> str: ... + + class _EndUserRow(_BudgetLinkedRow, Protocol): @property def user_id(self) -> str: ... @@ -185,6 +193,14 @@ def _model_access_group_cache_keys(row: _ModelAccessGroupRow) -> tuple[str, ...] return (model_access_group_cache_key(row.access_group_name),) +def _project_counter_key(row: _ProjectRow) -> str: + return project_spend_counter_key(row.project_id) + + +def _project_cache_keys(row: _ProjectRow) -> tuple[str, ...]: + return (project_cache_key(row.project_id),) + + def _enduser_counter_key(row: _EndUserRow) -> str: return f"spend:end_user:{row.user_id}" @@ -661,6 +677,11 @@ class ResetBudgetJob: where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE), log_subject="model access groups", ) + projects: Final[tuple[_ProjectRow, ...]] = await self._fetch_linked_rows( + table=ProjectRepository(self.prisma_client).table, + where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE), + log_subject="projects", + ) rollover_caps: Final[Mapping[str, float]] = MappingProxyType( { # mutable-ok: MappingProxyType wraps a one-shot dict comprehension b.budget_id: cap @@ -695,6 +716,7 @@ class ResetBudgetJob: (_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in model_access_groups ), + *((_project_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in projects), *((_enduser_counter_key(row), _enduser_carried_spend(row, rollover_caps)) for row in endusers), ), rollover_caps=rollover_caps, @@ -704,6 +726,7 @@ class ResetBudgetJob: *(key for row in orgs for key in _org_cache_keys(row)), *(key for row in tags for key in _tag_cache_keys(row)), *(key for row in model_access_groups for key in _model_access_group_cache_keys(row)), + *(key for row in projects for key in _project_cache_keys(row)), *(key for row in endusers for key in _enduser_cache_keys(row)), ), ) @@ -731,6 +754,7 @@ class ResetBudgetJob: _queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE) _queue_budget_linked_resets(uow.tags, cascade, extra=_SPENT_ROWS_WHERE) _queue_budget_linked_resets(uow.model_access_groups, cascade, extra=_SPENT_ROWS_WHERE) + _queue_budget_linked_resets(uow.projects, cascade, extra=_SPENT_ROWS_WHERE) _queue_enduser_resets(uow.endusers, cascade) for budget_id, budget_reset_at in cascade.budget_resets: uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 1c7a379897f..2187ed63ea5 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -306,6 +306,16 @@ def model_access_group_spend_counter_key(access_group_name: str) -> str: return f"spend:model_access_group:{access_group_name}" +def project_cache_key(project_id: str) -> str: + """Cache key one project row is stored under; shared by auth, spend tracking and the spend writer.""" + return f"project_id:{project_id}" + + +def project_spend_counter_key(project_id: str) -> str: + """Spend counter key for one project; the reservation, cost callback, auth and reseed paths all read it.""" + return f"spend:project:{project_id}" + + #: Cached under ``end_user_restricted_registry_cache_key`` when the restricted set exceeds #: ``END_USER_RESTRICTED_REGISTRY_MAX_SIZE``: registry unusable, fall back to the per-id fetch. END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL: Final = "__end_user_restricted_registry_overflow__" diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index a90d1351fd7..51d00b9789c 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -44,6 +44,7 @@ from litellm.proxy._types import ( SpendUpdateQueueItem, ToolDiscoveryQueueItem, ) +from litellm.proxy.common_utils.user_api_key_cache import project_cache_key from litellm.proxy.db.daily_spend_bulk_upsert import ( DAILY_SPEND_TABLES, build_bulk_upsert, @@ -116,6 +117,7 @@ class _SpendBatch(Protocol): litellm_teammembership: BatchTable litellm_organizationtable: BatchTable litellm_organizationmembership: BatchTable + litellm_projecttable: BatchTable litellm_tagtable: BatchTable litellm_agentstable: BatchTable litellm_modelaccessgroupbudgettable: BatchTable @@ -254,6 +256,7 @@ class DBSpendUpdateWriter: start_time: datetime | None, end_time: datetime | None, response_cost: float | None, + project_id: str | None = None, ) -> bool: """Record the request's spend, answering whether its cost still needs charging. @@ -335,6 +338,7 @@ class DBSpendUpdateWriter: hashed_token=hashed_token, team_id=team_id, org_id=org_id, + project_id=project_id, end_user_id=end_user_id, prisma_client=prisma_client, litellm_proxy_budget_name=litellm_proxy_budget_name, @@ -631,6 +635,7 @@ class DBSpendUpdateWriter: litellm_proxy_budget_name: str | None, payload: SpendLogsPayload, request_model_access_groups: Sequence[str] = (), + project_id: str | None = None, ): """ Runs all 13 spend-update helpers sequentially inside a single asyncio task. @@ -694,6 +699,18 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) + try: + await self._update_project_db( + response_cost=response_cost, + project_id=project_id, + prisma_client=prisma_client, + ) + except Exception: + verbose_proxy_logger.debug( + "_batch_database_updates: _update_project_db failed: %s", + traceback.format_exc(), + ) + try: await self._update_tag_db( response_cost=response_cost, @@ -956,6 +973,33 @@ class DBSpendUpdateWriter: ) raise e + async def _update_project_db( + self, + response_cost: float | None, + project_id: str | None, + prisma_client: PrismaClient | None, + ): + try: + if project_id is None or prisma_client is None: + return + + await self.spend_update_queue.add_update( + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.PROJECT, + entity_id=project_id, + response_cost=response_cost, + ) + ) + except Exception as e: + spend_log_error( + "Spend tracking - failed to enqueue project spend update. project_id=%s, response_cost=%s - %s", + project_id, + response_cost, + str(e), + exc=e, + ) + raise e + async def _update_agent_db( self, response_cost: float | None, @@ -1193,8 +1237,8 @@ class DBSpendUpdateWriter: if db_spend_update_transactions is not None: verbose_proxy_logger.info( "Spend tracking - committing spend updates from Redis to DB: " - "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, tags=%d, " - "agents=%d, model_access_groups=%d", + "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, " + "projects=%d, tags=%d, agents=%d, model_access_groups=%d", len(db_spend_update_transactions.get("key_list_transactions") or {}), len(db_spend_update_transactions.get("user_list_transactions") or {}), len(db_spend_update_transactions.get("team_list_transactions") or {}), @@ -1202,6 +1246,7 @@ class DBSpendUpdateWriter: len(db_spend_update_transactions.get("end_user_list_transactions") or {}), len(db_spend_update_transactions.get("team_member_list_transactions") or {}), len(db_spend_update_transactions.get("org_member_list_transactions") or {}), + len(db_spend_update_transactions.get("project_list_transactions") or {}), len(db_spend_update_transactions.get("tag_list_transactions") or {}), len(db_spend_update_transactions.get("agent_list_transactions") or {}), len(db_spend_update_transactions.get("model_access_group_list_transactions") or {}), @@ -1762,6 +1807,22 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, ) + ### UPDATE PROJECT TABLE ### + project_list_transactions: Final = db_spend_update_transactions.get("project_list_transactions") + await DBSpendUpdateWriter._update_entity_spend_in_db( + entity_name="Project", + transactions=project_list_transactions, + table_accessor="litellm_projecttable", + where_field="project_id", + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + ) + await DBSpendUpdateWriter._invalidate_project_caches( + project_ids=tuple(project_list_transactions or ()), + proxy_logging_obj=proxy_logging_obj, + ) + ### UPDATE TAG TABLE ### tag_list_transactions: Final = db_spend_update_transactions["tag_list_transactions"] await DBSpendUpdateWriter._update_entity_spend_in_db( @@ -1800,11 +1861,23 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, ) + @staticmethod + async def _invalidate_project_caches(project_ids: Sequence[str], proxy_logging_obj: ProxyLogging | None) -> None: + if not project_ids or proxy_logging_obj is None: + return + user_api_key_cache: Final = proxy_logging_obj.call_details.get("user_api_key_cache") + if user_api_key_cache is None: + return + for project_id in project_ids: + await user_api_key_cache.async_delete_cache(key=project_cache_key(project_id)) + @staticmethod async def _update_entity_spend_in_db( entity_name: str, transactions: dict[str, float] | None, - table_accessor: Literal["litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable"], + table_accessor: Literal[ + "litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable", "litellm_projecttable" + ], where_field: str, n_retry_times: int, prisma_client: PrismaClient, diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 6f49a00b763..cead63795a2 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -70,6 +70,7 @@ _SpendTransactionField: TypeAlias = Literal[ "team_member_list_transactions", "org_list_transactions", "org_member_list_transactions", + "project_list_transactions", "tag_list_transactions", "agent_list_transactions", "model_access_group_list_transactions", @@ -83,6 +84,7 @@ _SPEND_TRANSACTION_FIELDS: Final[tuple[_SpendTransactionField, ...]] = ( "team_member_list_transactions", "org_list_transactions", "org_member_list_transactions", + "project_list_transactions", "tag_list_transactions", "agent_list_transactions", "model_access_group_list_transactions", @@ -418,6 +420,10 @@ class RedisUpdateBuffer: Litellm_EntityType.ORGANIZATION_MEMBER, db_spend_update_transactions.get("org_member_list_transactions"), ), + ( + Litellm_EntityType.PROJECT, + db_spend_update_transactions.get("project_list_transactions"), + ), ( Litellm_EntityType.TAG, db_spend_update_transactions.get("tag_list_transactions"), @@ -885,6 +891,7 @@ class RedisUpdateBuffer: org_member_list_transactions=_merged_entity_transactions( list_of_transactions, "org_member_list_transactions" ), + project_list_transactions=_merged_entity_transactions(list_of_transactions, "project_list_transactions"), tag_list_transactions=_merged_entity_transactions(list_of_transactions, "tag_list_transactions"), agent_list_transactions=_merged_entity_transactions(list_of_transactions, "agent_list_transactions"), model_access_group_list_transactions=_merged_entity_transactions( diff --git a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py index bc068d10daf..2b8535cb113 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_update_queue.py +++ b/litellm/proxy/db/db_transaction_queue/spend_update_queue.py @@ -138,6 +138,7 @@ class SpendUpdateQueue(BaseUpdateQueue): team_member_list_transactions={}, org_list_transactions={}, org_member_list_transactions={}, + project_list_transactions={}, tag_list_transactions={}, agent_list_transactions={}, model_access_group_list_transactions={}, @@ -152,6 +153,7 @@ class SpendUpdateQueue(BaseUpdateQueue): Litellm_EntityType.TEAM_MEMBER: "team_member_list_transactions", Litellm_EntityType.ORGANIZATION: "org_list_transactions", Litellm_EntityType.ORGANIZATION_MEMBER: "org_member_list_transactions", + Litellm_EntityType.PROJECT: "project_list_transactions", Litellm_EntityType.TAG: "tag_list_transactions", Litellm_EntityType.AGENT: "agent_list_transactions", Litellm_EntityType.MODEL_ACCESS_GROUP: "model_access_group_list_transactions", @@ -192,6 +194,8 @@ class SpendUpdateQueue(BaseUpdateQueue): transactions_dict = db_spend_update_transactions["org_list_transactions"] elif dict_key == "org_member_list_transactions": transactions_dict = db_spend_update_transactions["org_member_list_transactions"] + elif dict_key == "project_list_transactions": + transactions_dict = db_spend_update_transactions["project_list_transactions"] elif dict_key == "tag_list_transactions": transactions_dict = db_spend_update_transactions["tag_list_transactions"] elif dict_key == "agent_list_transactions": diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index 89a07234c6c..2dd028454d6 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -26,6 +26,7 @@ from litellm.proxy._types import Litellm_EntityType from litellm.proxy.db.db_lookup_gate import db_lookup_gate from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value from litellm.repositories.organization_repository import OrganizationRepository +from litellm.repositories.project_repository import ProjectRepository from litellm.repositories.table_repositories import ( BudgetWindowSpendRepository, EndUserRepository, @@ -77,6 +78,7 @@ class SpendCounterReseed: spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend spend:user:{user_id} -> LiteLLM_UserTable.spend spend:org:{org_id} -> LiteLLM_OrganizationTable.spend + spend:project:{project_id} -> LiteLLM_ProjectTable.spend End-user and tag spend counters intentionally do not reseed here. Their auth paths already load the corresponding objects via get_end_user_object() @@ -157,6 +159,9 @@ class SpendCounterReseed: row = await OrganizationRepository(prisma_client).table.find_unique( where={"organization_id": org_id} ) + elif counter_key.startswith("spend:project:"): + project_id: Final = counter_key[len("spend:project:") :] + row = await ProjectRepository(prisma_client).table.find_unique(where={"project_id": project_id}) else: return None except Exception: diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 1ae106be390..0c562cf37ef 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -267,6 +267,7 @@ class _ProxyDBLogger(CustomLogger): start_time=actual_start_time, end_time=datetime.now(), org_id=user_api_key_dict.org_id, + project_id=user_api_key_dict.project_id, ) @log_db_metrics @@ -318,6 +319,7 @@ class _ProxyDBLogger(CustomLogger): user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None)) org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None)) + project_id: Final = cast(str | None, metadata.get("user_api_key_project_id", None)) key_alias: Final = cast(str | None, metadata.get("user_api_key_alias", None)) end_user_max_budget: Final = metadata.get("user_api_end_user_max_budget", None) sl_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) @@ -368,6 +370,7 @@ class _ProxyDBLogger(CustomLogger): budget_reservation=budget_reservation, request_tags=tags, model_access_groups=model_access_groups, + project_id=project_id, ) if not charged: return @@ -651,6 +654,7 @@ async def _update_database_and_spend_counters( budget_reservation: dict | None, request_tags: list[str] | None = None, model_access_groups: Sequence[str] | None = None, + project_id: str | None = None, ) -> bool: if budget_reservation is not None: await _reconcile_budget_reservation_before_db_update( @@ -668,6 +672,7 @@ async def _update_database_and_spend_counters( start_time=start_time, end_time=end_time, org_id=org_id, + project_id=project_id, ) except Exception: if budget_reservation is not None: @@ -698,6 +703,7 @@ async def _update_database_and_spend_counters( tags=request_tags, request_started_at=start_time, model_access_groups=model_access_groups, + project_id=project_id, ) except Exception: if budget_reservation is not None: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d7964556531..1e84e5f56e2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -417,6 +417,8 @@ from litellm.proxy.common_utils.user_api_key_cache import ( get_management_object_ttl, model_access_group_cache_key, model_access_group_spend_counter_key, + project_cache_key, + project_spend_counter_key, tag_cache_key, ) from litellm.proxy.config_resolvers import resolve_fields @@ -2780,6 +2782,7 @@ async def increment_spend_counters( tags: list[str] | None = None, request_started_at: datetime | None = None, model_access_groups: Sequence[str] | None = None, + project_id: str | None = None, ): """ Atomically increment spend counters for budget enforcement. @@ -2801,6 +2804,7 @@ async def increment_spend_counters( end_user_id=end_user_id, tags=tags, model_access_groups=model_access_groups, + project_id=project_id, ), ): await _increment_spend_counters_batched( @@ -2814,6 +2818,7 @@ async def increment_spend_counters( tags=tags, request_started_at=request_started_at, model_access_groups=model_access_groups, + project_id=project_id, ) @@ -2828,6 +2833,7 @@ async def _increment_spend_counters_batched( tags: list[str] | None, request_started_at: datetime | None, model_access_groups: Sequence[str] | None, + project_id: str | None = None, ): """Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET.""" reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update( @@ -3028,6 +3034,13 @@ async def _increment_spend_counters_batched( ) if org_id is not None else None, + _prepare_project_spend_increment( + project_id=project_id, + response_cost=cost, + reserved_counter_keys=reserved_counter_keys, + ) + if project_id is not None + else None, ) if coro is not None ) @@ -3180,6 +3193,23 @@ async def _prepare_org_spend_increment( return (pending,) if pending is not None else () +async def _prepare_project_spend_increment( + project_id: str | None, + response_cost: float, + reserved_counter_keys: set[str], +) -> tuple[PendingSpendIncrement, ...]: + if project_id is None: + return () + + pending: Final = await _prepare_unreserved_spend_counter_increment( + counter_key=project_spend_counter_key(project_id), + source_cache_key=project_cache_key(project_id), + increment=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + return (pending,) if pending is not None else () + + async def _prepare_unreserved_spend_counter_increment( counter_key: str, source_cache_key: str | list[str], diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 373f2d0fe36..24b8470145c 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -2,6 +2,7 @@ from __future__ import annotations import asyncio import json +import math from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone @@ -29,6 +30,8 @@ from litellm.proxy.common_utils.user_api_key_cache import ( end_user_cache_key, model_access_group_cache_key, model_access_group_spend_counter_key, + project_cache_key, + project_spend_counter_key, tag_cache_key, team_membership_reservation_cache_key, ) @@ -62,6 +65,7 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = { "Tag": Litellm_EntityType.TAG.value, "Model access group": Litellm_EntityType.MODEL_ACCESS_GROUP.value, "Organization": Litellm_EntityType.ORGANIZATION.value, + "Project": Litellm_EntityType.PROJECT.value, } @@ -542,6 +546,13 @@ async def _get_budget_counters( if org_counter is not None: counters.append(org_counter) + project_counter: Final = await _get_project_budget_counter( + valid_token=valid_token, + user_api_key_cache=user_api_key_cache, + ) + if project_counter is not None: + counters.append(project_counter) + return counters @@ -751,6 +762,36 @@ async def _get_org_budget_counter( ) +async def _get_project_budget_counter( + valid_token: UserAPIKeyAuth, + user_api_key_cache: UserApiKeyCache, +) -> _BudgetCounter | None: + if valid_token.project_id is None: + return None + + source_cache_key: Final = project_cache_key(valid_token.project_id) + project_object: Final = await user_api_key_cache.async_get_cache(key=source_cache_key) + if project_object is None: + return None + + project_budget_table: Final = _get_value(project_object, "litellm_budget_table") + if project_budget_table is None: + return None + + project_max_budget: Final = _to_float(_get_value(project_budget_table, "max_budget")) + if project_max_budget is None or project_max_budget <= 0 or not math.isfinite(project_max_budget): + return None + + return _BudgetCounter( + counter_key=project_spend_counter_key(valid_token.project_id), + source_cache_key=source_cache_key, + max_budget=project_max_budget, + fallback_spend=_to_float(_get_value(project_object, "spend")) or 0.0, + entity_type="Project", + entity_id=valid_token.project_id, + ) + + def _get_budget_limit_counters( entity_prefix: str, entity_type: str, diff --git a/litellm/proxy/spend_tracking/spend_counter_batch.py b/litellm/proxy/spend_tracking/spend_counter_batch.py index 7106d88c655..ddb074ae023 100644 --- a/litellm/proxy/spend_tracking/spend_counter_batch.py +++ b/litellm/proxy/spend_tracking/spend_counter_batch.py @@ -12,7 +12,10 @@ from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm.caching.redis_cache import RedisCache from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.common_utils.user_api_key_cache import model_access_group_spend_counter_key +from litellm.proxy.common_utils.user_api_key_cache import ( + model_access_group_spend_counter_key, + project_spend_counter_key, +) _CounterValues: Final = TypeAdapter(dict[str, float | None]) _NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({}) @@ -154,6 +157,8 @@ def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) yield f"spend:end_user:{end_user_id}" if token.org_id is not None: yield f"spend:org:{token.org_id}" + if token.project_id is not None: + yield project_spend_counter_key(token.project_id) def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]: @@ -168,10 +173,12 @@ def post_call_counter_keys( end_user_id: str | None, tags: Sequence[object] | None, model_access_groups: Sequence[object] | None, + project_id: str | None = None, ) -> frozenset[str]: """Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read.""" entity_keys: Final = admission_counter_keys( - UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id), end_user_id + UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id), + end_user_id, ) tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str)) group_keys: Final = frozenset( diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 93b8c5c7cd7..60c16fbd746 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -152,4 +152,7 @@ class PrismaBatch(Protocol): @property def litellm_modelaccessgroupbudgettable(self) -> BatchTable: ... + @property + def litellm_projecttable(self) -> BatchTable: ... + async def commit(self) -> None: ... diff --git a/litellm/repositories/unit_of_work.py b/litellm/repositories/unit_of_work.py index 0cdce307f9b..c09e5eb75d4 100644 --- a/litellm/repositories/unit_of_work.py +++ b/litellm/repositories/unit_of_work.py @@ -109,6 +109,7 @@ class BudgetCascadeUnitOfWork: organizations: LinkedSpendResetWrites tags: LinkedSpendResetWrites model_access_groups: LinkedSpendResetWrites + projects: LinkedSpendResetWrites endusers: LinkedSpendResetWrites budgets: BudgetWindowWrites @@ -135,6 +136,7 @@ async def budget_cascade_unit_of_work( organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable), tags=LinkedSpendResetWrites(table=batch.litellm_tagtable), model_access_groups=LinkedSpendResetWrites(table=batch.litellm_modelaccessgroupbudgettable), + projects=LinkedSpendResetWrites(table=batch.litellm_projecttable), endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable), budgets=BudgetWindowWrites(table=batch.litellm_budgettable), ) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 8c8b755195f..0d5c0dd5d72 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7138,6 +7138,67 @@ async def test_project_allowlist_enforced_when_key_models_empty(): assert exc_info.value.code == "403" +def _project_with_budget(spend: float, max_budget: float): + from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_ProjectTableCachedObj + + return LiteLLM_ProjectTableCachedObj( + project_id="p-budget", + team_id="t-1", + budget_id="b-1", + spend=spend, + litellm_budget_table=LiteLLM_BudgetTable(budget_id="b-1", max_budget=max_budget), + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "counter_spend, db_spend, blocks", + [ + pytest.param(5.0, 0.0, True, id="counter-at-budget-blocks-despite-stale-db-row"), + pytest.param(4.99, 0.0, False, id="counter-under-budget-admits"), + pytest.param(None, 5.0, True, id="no-counter-falls-back-to-persisted-spend"), + pytest.param(None, 0.0, False, id="no-counter-and-no-persisted-spend-admits"), + ], +) +async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, db_spend, blocks): + """LIT-3269: project budget enforcement must read the cross-pod + ``spend:project:{id}`` counter first and only fall back to the cached row's + spend, matching key/team/org checks. The boundary is inclusive (>=).""" + from litellm.caching.dual_cache import DualCache + from litellm.proxy.auth.auth_checks import _project_max_budget_check + + real_spend_counter_cache = DualCache() + if counter_spend is not None: + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:project:p-budget", value=counter_spend) + valid_token = UserAPIKeyAuth(api_key="hashed-key", project_id="p-budget", team_id="t-1", user_id="u-1") + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + if not blocks: + await _project_max_budget_check( + project_object=_project_with_budget(spend=db_spend, max_budget=5.0), + valid_token=valid_token, + proxy_logging_obj=proxy_logging_obj, + ) + return + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _project_max_budget_check( + project_object=_project_with_budget(spend=db_spend, max_budget=5.0), + valid_token=valid_token, + proxy_logging_obj=proxy_logging_obj, + ) + await asyncio.sleep(0) + + assert exc_info.value.entity_type == Litellm_EntityType.PROJECT.value + assert exc_info.value.entity_id == "p-budget" + assert exc_info.value.current_cost == 5.0 + proxy_logging_obj.budget_alerts.assert_awaited_once() + assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "project_budget" + + def test_is_user_proxy_admin_rejects_view_only_admin(): """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an Admin Viewer answering True here would gain every write route. Read parity for 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 943a6c905c0..8c254686385 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 @@ -78,6 +78,7 @@ class MockBatcher: self.litellm_organizationtable = _Table("org", self) self.litellm_tagtable = _Table("tag", self) self.litellm_modelaccessgroupbudgettable = _Table("model_access_group", self) + self.litellm_projecttable = _Table("project", self) self.litellm_endusertable = _Table("enduser", self) async def commit(self): @@ -93,6 +94,7 @@ class MockDB: self.litellm_organizationtable = MockTable() self.litellm_tagtable = MockTable() self.litellm_modelaccessgroupbudgettable = MockTable() + self.litellm_projecttable = MockTable() self.batch_calls: List[Dict[str, Any]] = [] self.batchers: List[MockBatcher] = [] @@ -1507,13 +1509,19 @@ _INVALIDATION_CASES = [ "spend:model_access_group:gpt-4-group", {"model_access_group:gpt-4-group"}, ), + ( + "litellm_projecttable", + type("Project", (), {"project_id": "proj-1"}), + "spend:project:proj-1", + {"project_id:proj-1"}, + ), ] @pytest.mark.parametrize( "table_attr, linked_row, counter_key, cache_keys", _INVALIDATION_CASES, - ids=["team_membership", "key", "org", "tag", "model_access_group"], + ids=["team_membership", "key", "org", "tag", "model_access_group", "project"], ) def test_budget_table_reset_invalidates_counters_and_management_cache( reset_budget_job, mock_prisma_client, monkeypatch, table_attr, linked_row, counter_key, cache_keys @@ -1657,6 +1665,25 @@ def test_budget_table_reset_invalidates_every_access_group_not_just_the_first( counter_cache.in_memory_cache.delete_cache.assert_any_call(key=f"spend:model_access_group:{name}") +def test_project_reset_zeroes_spend_on_due_tiers(reset_budget_job, mock_prisma_client, monkeypatch): + """A project linked to an expiring budget tier has its spend zeroed in the same cascade transaction.""" + _make_counter_invalidation_job(monkeypatch) + mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-due", budget_duration="7d")] + mock_prisma_client.db.litellm_projecttable.set_find_many_results( + [type("Project", (), {"project_id": "proj-1", "spend": 12.0, "budget_id": "budget-due"})] + ) + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + expected_where = {"budget_id": {"in": ["budget-due"]}, "spend": {"gt": 0}} + assert mock_prisma_client.db.litellm_projecttable.find_many_calls == [{"where": expected_where}] + writes = _batch_writes(mock_prisma_client, "project", op="update_many") + assert len(writes) == 1 + assert writes[0]["where"] == expected_where + assert writes[0]["data"] == {"spend": 0} + assert mock_prisma_client.db.batchers[0].committed is True + + def test_budget_cascade_carries_access_group_overage_when_rollover_enabled( rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch ): @@ -1802,6 +1829,7 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo ("org", "update_many"), ("tag", "update_many"), ("model_access_group", "update_many"), + ("project", "update_many"), ("enduser", "update_many"), ("budget", "update_many"), } diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index c547d06904b..da5879a375a 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1060,6 +1060,82 @@ async def test_batch_database_updates_queues_org_member_spend_for_the_request_us assert transactions["org_member_list_transactions"] == {"organization_id::org1::user_id::u1": 0.1} +@pytest.mark.asyncio +async def test_project_spend_is_persisted_to_project_table_and_project_cache_is_evicted(): + """Regression for LIT-3269: a request made with a project-scoped key must + increment LiteLLM_ProjectTable.spend, otherwise /project/info stays at 0 + and the project budget never blocks. The cached project row is evicted so + the next auth check reads the fresh spend.""" + db_writer: Final = DBSpendUpdateWriter() + await db_writer._batch_database_updates( + response_cost=0.25, + user_id="u1", + hashed_token="t1", + team_id="team-1", + org_id=None, + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25}, + project_id="proj-1", + ) + await db_writer._batch_database_updates( + response_cost=0.5, + user_id="u1", + hashed_token="t1", + team_id="team-1", + org_id=None, + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload={"request_id": "req-2", "model": "gpt-4o-mini", "spend": 0.5}, + project_id="proj-1", + ) + transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + assert transactions["project_list_transactions"] == {"proj-1": 0.75} + assert transactions["team_member_list_transactions"] == {"team_id::team-1::user_id::u1": 0.75} + + mock_batcher: Final = MagicMock() + mock_prisma_client: Final = MagicMock() + mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(mock_batcher)) + user_api_key_cache: Final = MagicMock() + user_api_key_cache.async_delete_cache = AsyncMock() + proxy_logging: Final = MagicMock() + proxy_logging.call_details = {"user_api_key_cache": user_api_key_cache} + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=0, + proxy_logging_obj=proxy_logging, + db_spend_update_transactions=transactions, + ) + + mock_batcher.litellm_projecttable.update_many.assert_called_once_with( + where={"project_id": "proj-1"}, + data={"spend": {"increment": 0.75}}, + ) + user_api_key_cache.async_delete_cache.assert_any_await(key="project_id:proj-1") + + +@pytest.mark.asyncio +async def test_batch_database_updates_without_project_id_touches_no_project_row(): + db_writer: Final = DBSpendUpdateWriter() + await db_writer._batch_database_updates( + response_cost=0.1, + user_id="u1", + hashed_token="t1", + team_id=None, + org_id=None, + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.1}, + ) + transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + + assert transactions["project_list_transactions"] == {} + + @pytest.mark.asyncio async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_id(): """ diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index bca6344b3f7..53e91b8792e 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -67,12 +67,14 @@ class _FakePrismaClient: error: Exception | None = None, end_user_row: SimpleNamespace | None = None, end_user_error: Exception | None = None, + project_row: SimpleNamespace | None = None, ) -> None: self.db = SimpleNamespace( litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error), litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total), litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error), litellm_verificationtoken=_InFlightCountingTable(), + litellm_projecttable=_FakeFindUniqueTable(row=project_row), ) @@ -428,6 +430,23 @@ async def test_from_db_bounds_in_flight_prisma_requests_across_counter_keys(): assert prisma.db.litellm_verificationtoken.max_in_flight == PROXY_DB_LOOKUP_MAX_CONCURRENCY +@pytest.mark.asyncio +async def test_from_db_reseeds_project_counter_from_the_project_row(): + """LIT-3269: a cold ``spend:project:{id}`` counter seeds from LiteLLM_ProjectTable.spend, + so a fresh pod enforces the project budget against persisted spend rather than 0.""" + prisma: Final = _FakePrismaClient(project_row=SimpleNamespace(project_id="proj-1", spend=7.25)) + + assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") == 7.25 + assert prisma.db.litellm_projecttable.where_clauses == [{"project_id": "proj-1"}] + + +@pytest.mark.asyncio +async def test_from_db_returns_none_for_a_missing_project_row(): + prisma: Final = _FakePrismaClient(project_row=None) + + assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") is None + + @pytest.mark.asyncio async def test_from_db_still_never_reads_the_end_user_row(): """A cold end-user counter keeps seeding from the cached end-user object the auth diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index dfc95db3e14..965e134772d 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -586,6 +586,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda tags=["tag-a"], request_started_at=start_time, model_access_groups=("premium",), + project_id=None, ) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 8b105e94d19..834cf8d100d 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -1187,6 +1187,7 @@ async def test_api_key_preserved_through_failure_hook_to_database(): start_time, end_time, org_id, + project_id=None, ): """Mock update_database and capture the payload it creates""" from litellm.proxy.spend_tracking.spend_tracking_utils import ( diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 032722d3259..014f240d9cc 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -22,6 +22,7 @@ from litellm.proxy._types import ( LiteLLM_EndUserTable, Litellm_EntityType, LiteLLM_OrganizationTable, + LiteLLM_ProjectTableCachedObj, LiteLLM_TagTable, LiteLLM_TeamMembership, LiteLLM_TeamTable, @@ -631,6 +632,161 @@ async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_ await release_budget_reservation(reservation) +def _project_scoped_token() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token="key-project-scoped", + spend=0.0, + user_id="user-proj", + team_id="team-proj", + project_id="proj-1", + ) + + +async def _seed_project_scoped_budgets( + key_cache: DualCache, + team_member_spend: float, + team_member_max_budget: float, + project_spend: float, + project_max_budget: float, +) -> None: + await key_cache.async_set_cache( + key="team_membership:user-proj:team-proj", + value=LiteLLM_TeamMembership( + user_id="user-proj", + team_id="team-proj", + spend=team_member_spend, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=team_member_max_budget), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="project_id:proj-1", + value=LiteLLM_ProjectTableCachedObj( + project_id="proj-1", + team_id="team-proj", + budget_id="project-budget-id", + spend=project_spend, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=project_max_budget), + ).model_dump(), + ) + + +@pytest.mark.asyncio +async def test_should_reserve_project_and_team_member_counters_for_project_scoped_key(spend_counter_state): + """LIT-3269: a key carrying user_id, team_id and project_id reserves against + both the team member counter and the project counter; neither replaces the + other. After the call the project counter reflects the real cost once, not + the reservation plus the post-call increment.""" + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + await _seed_project_scoped_budgets( + key_cache, + team_member_spend=0.1, + team_member_max_budget=1.0, + project_spend=0.2, + project_max_budget=1.0, + ) + + estimated = estimate_request_max_cost(request_body=_request_body(), route="/chat/completions", llm_router=None) + assert estimated is not None and estimated > 0 + + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=_project_scoped_token(), + team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None), + user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0), + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-proj:team-proj") == pytest.approx( + 0.1 + estimated + ) + assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") == pytest.approx(0.2 + estimated) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token="key-project-scoped", + team_id="team-proj", + user_id="user-proj", + response_cost=0.05, + budget_reservation=reservation, + project_id="proj-1", + ) + + assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") == pytest.approx(0.25) + assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-proj:team-proj") == pytest.approx(0.15) + + +@pytest.mark.asyncio +async def test_exhausted_team_member_budget_still_blocks_project_scoped_key(spend_counter_state): + """LIT-3269: the project budget is additive. A project with plenty of + headroom must not let a key through once its team member budget is spent.""" + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + await _seed_project_scoped_budgets( + key_cache, + team_member_spend=1.0, + team_member_max_budget=1.0, + project_spend=0.0, + project_max_budget=100.0, + ) + + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=_project_scoped_token(), + team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None), + user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0), + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert "TeamMember=user-proj:team-proj" in str(exc_info.value) + assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") in (None, pytest.approx(0.0)) + + +@pytest.mark.asyncio +async def test_exhausted_project_budget_blocks_project_scoped_key(spend_counter_state): + """LIT-3269: with team member headroom left, the project budget alone blocks the key.""" + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + await _seed_project_scoped_budgets( + key_cache, + team_member_spend=0.0, + team_member_max_budget=100.0, + project_spend=5.0, + project_max_budget=5.0, + ) + + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=_project_scoped_token(), + team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None), + user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0), + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert "Project=proj-1" in str(exc_info.value) + assert exc_info.value.entity_type == Litellm_EntityType.PROJECT.value + assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-proj:team-proj") in ( + None, + pytest.approx(0.0), + ) + + @pytest.mark.asyncio async def test_should_not_reserve_user_budget_counter_for_team_key(spend_counter_state): """The reservation path mirrors the read path: no personal user counter for a team key. diff --git a/tests/test_litellm/repositories/test_unit_of_work.py b/tests/test_litellm/repositories/test_unit_of_work.py index b52b8ced31e..9eff248b917 100644 --- a/tests/test_litellm/repositories/test_unit_of_work.py +++ b/tests/test_litellm/repositories/test_unit_of_work.py @@ -35,6 +35,7 @@ class FakeBatch: self.litellm_organizationtable = FakeBatchTable("litellm_organizationtable", self.calls) self.litellm_tagtable = FakeBatchTable("litellm_tagtable", self.calls) self.litellm_modelaccessgroupbudgettable = FakeBatchTable("litellm_modelaccessgroupbudgettable", self.calls) + self.litellm_projecttable = FakeBatchTable("litellm_projecttable", self.calls) self.litellm_endusertable = FakeBatchTable("litellm_endusertable", self.calls) async def commit(self) -> None: @@ -94,6 +95,7 @@ async def test_budget_cascade_dependents_and_window_advance_share_one_batch(): uow.organizations.queue_spend_zero(where=linked) uow.tags.queue_spend_zero(where=linked) uow.model_access_groups.queue_spend_zero(where=linked) + uow.projects.queue_spend_zero(where=linked) uow.endusers.queue_spend_zero(where={"user_id": {"in": ["enduser-1"]}}) uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=reset_at) assert batch.commit_count == 0 @@ -105,6 +107,7 @@ async def test_budget_cascade_dependents_and_window_advance_share_one_batch(): ("litellm_organizationtable.update_many", linked, {"spend": 0}), ("litellm_tagtable.update_many", linked, {"spend": 0}), ("litellm_modelaccessgroupbudgettable.update_many", linked, {"spend": 0}), + ("litellm_projecttable.update_many", linked, {"spend": 0}), ("litellm_endusertable.update_many", {"user_id": {"in": ["enduser-1"]}}, {"spend": 0}), ("litellm_budgettable.update_many", {"budget_id": "budget-1"}, {"budget_reset_at": reset_at}), ] From 0601d2bb03646c596a680d580b0f9bb5a83ee237 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 01:58:37 +0000 Subject: [PATCH 058/253] feat(proxy): carry response time metrics through LiteLLM_DailyGlobalSpend Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../migration.sql | 2 ++ .../litellm_proxy_extras/schema.prisma | 2 ++ .../management_endpoints/common_daily_activity.py | 2 ++ litellm/proxy/schema.prisma | 2 ++ .../spend_tracking/daily_global_spend_rollup.py | 2 ++ schema.prisma | 2 ++ .../test_common_daily_activity.py | 6 ++++++ .../test_daily_global_spend_rollup.py | 13 +++++++++++-- 8 files changed, 29 insertions(+), 2 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql index 1d6cdea0c7b..d0bc3e159de 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260915000000_add_daily_global_spend/migration.sql @@ -20,6 +20,8 @@ CREATE TABLE IF NOT EXISTS "LiteLLM_DailyGlobalSpend" ( "api_requests" BIGINT NOT NULL DEFAULT 0, "successful_requests" BIGINT NOT NULL DEFAULT 0, "failed_requests" BIGINT NOT NULL DEFAULT 0, + "total_response_time_ms" BIGINT NOT NULL DEFAULT 0, + "timed_requests" BIGINT NOT NULL DEFAULT 0, "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, "updated_at" TIMESTAMP(3) NOT NULL, diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index a73e8774c87..42769c323a9 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -837,6 +837,8 @@ model LiteLLM_DailyGlobalSpend { api_requests BigInt @default(0) successful_requests BigInt @default(0) failed_requests BigInt @default(0) + total_response_time_ms BigInt @default(0) + timed_requests BigInt @default(0) created_at DateTime @default(now()) updated_at DateTime @updatedAt diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 31ec0c51b7d..47465324f42 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -767,6 +767,8 @@ _KEY_FREE_SOURCE_COLUMNS: Final = ( "api_requests", "successful_requests", "failed_requests", + "total_response_time_ms", + "timed_requests", ) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index a73e8774c87..42769c323a9 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -837,6 +837,8 @@ model LiteLLM_DailyGlobalSpend { api_requests BigInt @default(0) successful_requests BigInt @default(0) failed_requests BigInt @default(0) + total_response_time_ms BigInt @default(0) + timed_requests BigInt @default(0) created_at DateTime @default(now()) updated_at DateTime @updatedAt diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index 73068381dab..ba4a4e4e3d6 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -43,6 +43,8 @@ _METRIC_COLUMNS: Final = ( "api_requests", "successful_requests", "failed_requests", + "total_response_time_ms", + "timed_requests", "compression_savings_spend", "prompt_caching_savings_spend", "gateway_injected_caching_savings_spend", diff --git a/schema.prisma b/schema.prisma index a73e8774c87..42769c323a9 100644 --- a/schema.prisma +++ b/schema.prisma @@ -837,6 +837,8 @@ model LiteLLM_DailyGlobalSpend { api_requests BigInt @default(0) successful_requests BigInt @default(0) failed_requests BigInt @default(0) + total_response_time_ms BigInt @default(0) + timed_requests BigInt @default(0) created_at DateTime @default(now()) updated_at DateTime @updatedAt diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 2f25c507e0f..ee9f886acec 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1771,6 +1771,10 @@ async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_ ] _seed_daily_user_spend(_aggregated_postgresql, rows) with _aggregated_postgresql.cursor() as cur: + cur.execute( + 'UPDATE "LiteLLM_DailyUserSpend" SET total_response_time_ms = prompt_tokens * 25, ' + "timed_requests = api_requests" + ) cur.execute(_GLOBAL_SPEND_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # DDL literal cur.execute( re.sub(r"\$(\d+)", r"%(p\1)s", RECONCILE_DAY_SQL), # pyright: ignore[reportArgumentType] # $N -> psycopg @@ -1805,6 +1809,8 @@ async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_ assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 3 assert from_global.model_dump() == from_per_key.model_dump() assert from_global.metadata.total_spend == pytest.approx(2 * sum(float(i + 1) for i in range(n_keys))) + assert from_global.metadata.total_response_time_ms == 2 * n_keys * 10 * 25 + assert from_global.metadata.total_timed_requests == 2 * n_keys assert {day.date.isoformat() for day in from_global.results} == {"2026-06-01", "2026-06-02"} assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"} diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 9a098744f08..b5b78229c11 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -336,6 +336,8 @@ _DAILY_USER_SPEND_DDL: Final = """ api_requests BIGINT DEFAULT 0, successful_requests BIGINT DEFAULT 0, failed_requests BIGINT DEFAULT 0, + total_response_time_ms BIGINT DEFAULT 0, + timed_requests BIGINT DEFAULT 0, created_at TIMESTAMP DEFAULT now(), updated_at TIMESTAMP, UNIQUE (user_id, date, api_key, model, custom_llm_provider, mcp_namespaced_tool_name, endpoint) @@ -345,12 +347,14 @@ _DAILY_USER_SPEND_DDL: Final = """ _PER_KEY_SUMS_SQL: Final = """ SELECT COALESCE(model, '') AS model, COALESCE(model_group, '') AS model_group, COALESCE(custom_llm_provider, '') AS custom_llm_provider, - SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, SUM(api_requests) AS api_requests + SUM(spend) AS spend, SUM(prompt_tokens) AS prompt_tokens, SUM(api_requests) AS api_requests, + SUM(total_response_time_ms) AS total_response_time_ms, SUM(timed_requests) AS timed_requests FROM "LiteLLM_DailyUserSpend" WHERE date = %s GROUP BY 1, 2, 3 ORDER BY 1, 2, 3 """ _GLOBAL_ROWS_SQL: Final = """ - SELECT model, model_group, custom_llm_provider, spend, prompt_tokens, api_requests + SELECT model, model_group, custom_llm_provider, spend, prompt_tokens, api_requests, + total_response_time_ms, timed_requests FROM "LiteLLM_DailyGlobalSpend" WHERE date = %s ORDER BY 1, 2, 3 """ @@ -380,6 +384,8 @@ def _user_txn(**overrides): "api_requests": 1, "successful_requests": 1, "failed_requests": 0, + "total_response_time_ms": 800, + "timed_requests": 1, **overrides, } @@ -393,6 +399,8 @@ def _normalized(rows: list[dict[str, object]]) -> list[tuple[object, ...]]: float(r["spend"]), int(r["prompt_tokens"]), int(r["api_requests"]), + int(r["total_response_time_ms"]), + int(r["timed_requests"]), ) # pyright: ignore[reportArgumentType] # dict_row values are untyped for r in rows ] @@ -436,5 +444,6 @@ def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_p assert _normalized(global_rows) == _normalized(per_key) assert sum(float(r["spend"]) for r in global_rows) == pytest.approx(15.0) # pyright: ignore[reportArgumentType] # dict_row values are untyped + assert sum(int(r["total_response_time_ms"]) for r in global_rows) == 1600 # pyright: ignore[reportArgumentType] # dict_row values are untyped assert [(r["model"], r["model_group"]) for r in global_rows] == [("gpt-5", ""), ("gpt-5", "gpt-5")] assert untouched == [] From 4bca66f30377461ab39aba4c12ae1e3124f1d856 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:59:17 +0000 Subject: [PATCH 059/253] refactor(router): drop explanatory docstrings from routing group helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_server.py | 6 ------ litellm/router.py | 4 ---- litellm/router_utils/routing_groups.py | 17 ----------------- 3 files changed, 27 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b52d007d9db..0a49aef5c7e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6902,12 +6902,6 @@ class ProxyConfig: @staticmethod def _apply_router_settings(llm_router: Router, router_settings: Mapping[str, object]) -> None: - """ - `routing_groups` is applied on its own so a value persisted before - save-time validation existed cannot abort the reconcile that also loads - SSO, guardrails and the other DB-backed settings. The router keeps the - groups it already holds when the new value is rejected. - """ llm_router.update_settings(**{k: v for k, v in router_settings.items() if k != "routing_groups"}) if "routing_groups" not in router_settings: return diff --git a/litellm/router.py b/litellm/router.py index e31689ef447..23f058ef68e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1392,10 +1392,6 @@ class Router: at most one explicit group. Constructs per-group strategy selectors so groups with different `routing_strategy_args` track independent state. - Validation and selector construction run to completion before any - router state changes, so a rejected input raises with the previously - loaded groups still routing. - Models not claimed by any explicit group are served by the implicit `"default"` group, whose selectors are the `self._logger` attributes set up in `routing_strategy_init`. diff --git a/litellm/router_utils/routing_groups.py b/litellm/router_utils/routing_groups.py index c9b3b205bf1..ba65ddf8643 100644 --- a/litellm/router_utils/routing_groups.py +++ b/litellm/router_utils/routing_groups.py @@ -1,9 +1,3 @@ -""" -Validation for `router_settings.routing_groups`, shared by the Router and the -proxy's config-update endpoint so a config the UI saves cannot be one the -runtime refuses to load. -""" - from collections.abc import Sequence from typing import Final @@ -12,11 +6,6 @@ from litellm.types.router import RoutingGroup, RoutingStrategy def validate_routing_strategy(routing_strategy: RoutingStrategy | str | None) -> None: - """ - Raises `ValueError` unless `routing_strategy` is a known strategy or None. - - See: https://github.com/BerriAI/litellm/issues/11330 - """ if routing_strategy is None: return @@ -36,12 +25,6 @@ def parse_routing_groups( groups_input: Sequence[RoutingGroup | dict] | None, known_model_names: frozenset[str] = frozenset(), ) -> tuple[RoutingGroup, ...]: - """ - Parses and validates `routing_groups`, raising `ValueError` on the first - problem found. Every check runs before the caller mutates any state, so an - invalid update can never leave a router holding a half-applied set of - groups. - """ if not groups_input: return () From b00cd15bd73a380c77aad0d04672d27ee4bc13cf Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:07:11 +0000 Subject: [PATCH 060/253] test(proxy): import project_cache_key from user_api_key_cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/test_key_management_endpoints.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) 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 4e70063015d..b57f8b7857d 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 @@ -38,11 +38,10 @@ from litellm.proxy._types import ( from litellm.models.object_permission import LiteLLM_ObjectPermissionTable from litellm.proxy.auth.auth_checks import ( _delete_cache_key_object, - _project_cache_key, jwt_key_mapping_cache_key, ) from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth -from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy.management_endpoints.key_management_endpoints import ( _check_org_key_limits, @@ -18951,7 +18950,7 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read( async def _cache_with_project(project_id: str, project_models: list[str]) -> UserApiKeyCache: user_api_key_cache = UserApiKeyCache() await user_api_key_cache.async_set_cache( - key=_project_cache_key(project_id), + key=project_cache_key(project_id), value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id="team-lit-5823", models=project_models), model_type=LiteLLM_ProjectTableCachedObj, ) From b20f1422eb204c2cb2fba26912a548454cd18b02 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:28:38 +0000 Subject: [PATCH 061/253] fix(proxy): carry project_id through key metadata enrichment and drop docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_utils/user_api_key_cache.py | 2 -- litellm/proxy/hooks/proxy_track_cost_callback.py | 2 ++ tests/test_litellm/proxy/auth/test_auth_checks.py | 3 --- .../proxy/common_utils/test_reset_budget_job.py | 1 - tests/test_litellm/proxy/db/test_db_spend_update_writer.py | 4 ---- tests/test_litellm/proxy/db/test_spend_counter_reseed.py | 2 -- .../proxy/hooks/test_proxy_track_cost_callback.py | 3 +++ tests/test_litellm/proxy/test_budget_reservation.py | 7 ------- 8 files changed, 5 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 2187ed63ea5..0386b58070d 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -307,12 +307,10 @@ def model_access_group_spend_counter_key(access_group_name: str) -> str: def project_cache_key(project_id: str) -> str: - """Cache key one project row is stored under; shared by auth, spend tracking and the spend writer.""" return f"project_id:{project_id}" def project_spend_counter_key(project_id: str) -> str: - """Spend counter key for one project; the reservation, cost callback, auth and reseed paths all read it.""" return f"spend:project:{project_id}" diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 0c562cf37ef..5e525108ade 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -504,6 +504,8 @@ class _ProxyDBLogger(CustomLogger): metadata["user_api_key_team_id"] = key_obj.team_id if metadata.get("user_api_key_org_id") is None: metadata["user_api_key_org_id"] = key_obj.org_id + if metadata.get("user_api_key_project_id") is None: + metadata["user_api_key_project_id"] = key_obj.project_id except Exception: verbose_proxy_logger.debug( "Failed to enrich failure metadata with key info for api_key=%s", diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 0d5c0dd5d72..fe469fc574e 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7161,9 +7161,6 @@ def _project_with_budget(spend: float, max_budget: float): ], ) async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, db_spend, blocks): - """LIT-3269: project budget enforcement must read the cross-pod - ``spend:project:{id}`` counter first and only fall back to the cached row's - spend, matching key/team/org checks. The boundary is inclusive (>=).""" from litellm.caching.dual_cache import DualCache from litellm.proxy.auth.auth_checks import _project_max_budget_check 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 8c254686385..48b06649237 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 @@ -1666,7 +1666,6 @@ def test_budget_table_reset_invalidates_every_access_group_not_just_the_first( def test_project_reset_zeroes_spend_on_due_tiers(reset_budget_job, mock_prisma_client, monkeypatch): - """A project linked to an expiring budget tier has its spend zeroed in the same cascade transaction.""" _make_counter_invalidation_job(monkeypatch) mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-due", budget_duration="7d")] mock_prisma_client.db.litellm_projecttable.set_find_many_results( diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index da5879a375a..60cc742577b 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1062,10 +1062,6 @@ async def test_batch_database_updates_queues_org_member_spend_for_the_request_us @pytest.mark.asyncio async def test_project_spend_is_persisted_to_project_table_and_project_cache_is_evicted(): - """Regression for LIT-3269: a request made with a project-scoped key must - increment LiteLLM_ProjectTable.spend, otherwise /project/info stays at 0 - and the project budget never blocks. The cached project row is evicted so - the next auth check reads the fresh spend.""" db_writer: Final = DBSpendUpdateWriter() await db_writer._batch_database_updates( response_cost=0.25, diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index 53e91b8792e..ff0b67d426b 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -432,8 +432,6 @@ async def test_from_db_bounds_in_flight_prisma_requests_across_counter_keys(): @pytest.mark.asyncio async def test_from_db_reseeds_project_counter_from_the_project_row(): - """LIT-3269: a cold ``spend:project:{id}`` counter seeds from LiteLLM_ProjectTable.spend, - so a fresh pod enforces the project budget against persisted spend rather than 0.""" prisma: Final = _FakePrismaClient(project_row=SimpleNamespace(project_id="proj-1", spend=7.25)) assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") == 7.25 diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 965e134772d..202495517ad 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1372,6 +1372,7 @@ async def test_enrich_failure_metadata_with_full_key_lookup(): mock_key_obj.user_id = "fetched-user-id" mock_key_obj.team_id = "fetched-team-id" mock_key_obj.org_id = "fetched-org-id" + mock_key_obj.project_id = "fetched-project-id" mock_team_obj = MagicMock() mock_team_obj.team_alias = "fetched-team-alias" @@ -1395,12 +1396,14 @@ async def test_enrich_failure_metadata_with_full_key_lookup(): "user_api_key_team_id": None, "user_api_key_team_alias": None, "user_api_key_org_id": None, + "user_api_key_project_id": None, } result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) assert result["user_api_key_alias"] == "fetched-key-alias" assert result["user_api_key_user_id"] == "fetched-user-id" assert result["user_api_key_team_id"] == "fetched-team-id" assert result["user_api_key_org_id"] == "fetched-org-id" + assert result["user_api_key_project_id"] == "fetched-project-id" assert result["user_api_key_team_alias"] == "fetched-team-alias" diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 014f240d9cc..c834ac05f0a 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -672,10 +672,6 @@ async def _seed_project_scoped_budgets( @pytest.mark.asyncio async def test_should_reserve_project_and_team_member_counters_for_project_scoped_key(spend_counter_state): - """LIT-3269: a key carrying user_id, team_id and project_id reserves against - both the team member counter and the project counter; neither replaces the - other. After the call the project counter reflects the real cost once, not - the reservation plus the post-call increment.""" counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) await _seed_project_scoped_budgets( @@ -724,8 +720,6 @@ async def test_should_reserve_project_and_team_member_counters_for_project_scope @pytest.mark.asyncio async def test_exhausted_team_member_budget_still_blocks_project_scoped_key(spend_counter_state): - """LIT-3269: the project budget is additive. A project with plenty of - headroom must not let a key through once its team member budget is spent.""" counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) await _seed_project_scoped_budgets( @@ -755,7 +749,6 @@ async def test_exhausted_team_member_budget_still_blocks_project_scoped_key(spen @pytest.mark.asyncio async def test_exhausted_project_budget_blocks_project_scoped_key(spend_counter_state): - """LIT-3269: with team member headroom left, the project budget alone blocks the key.""" counter_cache, key_cache = spend_counter_state proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) await _seed_project_scoped_budgets( From c03a42a9c8144f8c2346ba8cd9dd342e1fe71149 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 16 Sep 2026 07:17:27 +0000 Subject: [PATCH 062/253] fix(azure): keep api-version query after vector store search path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../openai/vector_stores/transformation.py | 3 +- .../llms/azure/vector_stores/__init__.py | 0 ...test_azure_vector_stores_transformation.py | 20 ++++++++++++ ...est_openai_vector_stores_transformation.py | 32 +++++++++++-------- 4 files changed, 40 insertions(+), 15 deletions(-) create mode 100644 tests/test_litellm/llms/azure/vector_stores/__init__.py create mode 100644 tests/test_litellm/llms/azure/vector_stores/test_azure_vector_stores_transformation.py diff --git a/litellm/llms/openai/vector_stores/transformation.py b/litellm/llms/openai/vector_stores/transformation.py index 125e5168c69..57fe5b04838 100644 --- a/litellm/llms/openai/vector_stores/transformation.py +++ b/litellm/llms/openai/vector_stores/transformation.py @@ -101,7 +101,8 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig): extra_body: dict[str, object] | None = None, ) -> tuple[str, dict]: encoded_vector_store_id: Final = encode_url_path_segment(vector_store_id, field_name="vector_store_id") - url: Final = f"{api_base}/{encoded_vector_store_id}/search" + base_url, query_separator, query_string = api_base.partition("?") + url: Final = f"{base_url}/{encoded_vector_store_id}/search{query_separator}{query_string}" typed_request_body: Final = VectorStoreSearchRequest( query=query, filters=vector_store_search_optional_params.get("filters", None), diff --git a/tests/test_litellm/llms/azure/vector_stores/__init__.py b/tests/test_litellm/llms/azure/vector_stores/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/azure/vector_stores/test_azure_vector_stores_transformation.py b/tests/test_litellm/llms/azure/vector_stores/test_azure_vector_stores_transformation.py new file mode 100644 index 00000000000..59bec08fca6 --- /dev/null +++ b/tests/test_litellm/llms/azure/vector_stores/test_azure_vector_stores_transformation.py @@ -0,0 +1,20 @@ +from litellm.llms.azure.vector_stores.transformation import AzureOpenAIVectorStoreConfig + + +def test_transform_search_vector_store_request_preserves_azure_query_string(): + config = AzureOpenAIVectorStoreConfig() + api_base = config.get_complete_url( + api_base="https://x.openai.azure.com", + litellm_params={"api_version": "2024-10-21"}, + ) + + url, _ = config.transform_search_vector_store_request( + vector_store_id="vs_1", + query="hello", + vector_store_search_optional_params={}, + api_base=api_base, + litellm_logging_obj=None, + litellm_params={"api_version": "2024-10-21"}, + ) + + assert url == "https://x.openai.azure.com/openai/vector_stores/vs_1/search?api-version=2024-10-21" diff --git a/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py b/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py index e7b1aab45b4..ea1f9e87ed8 100644 --- a/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py +++ b/tests/test_litellm/llms/openai/vector_stores/test_openai_vector_stores_transformation.py @@ -7,11 +7,8 @@ from litellm.types.vector_stores import ( class TestOpenAIVectorStoreAPIConfig: - @pytest.mark.parametrize("metadata", [{}, None]) - def test_transform_create_vector_store_request_with_metadata_empty_or_none( - self, metadata - ): + def test_transform_create_vector_store_request_with_metadata_empty_or_none(self, metadata): """ Test transform_create_vector_store_request when metadata is None or empty dict. """ @@ -24,9 +21,7 @@ class TestOpenAIVectorStoreAPIConfig: "metadata": metadata, } - url, request_body = config.transform_create_vector_store_request( - vector_store_create_params, api_base - ) + url, request_body = config.transform_create_vector_store_request(vector_store_create_params, api_base) assert url == api_base assert request_body["name"] == "test-vector-store" @@ -50,9 +45,7 @@ class TestOpenAIVectorStoreAPIConfig: "metadata": large_metadata, } - url, request_body = config.transform_create_vector_store_request( - vector_store_create_params, api_base - ) + url, request_body = config.transform_create_vector_store_request(vector_store_create_params, api_base) assert url == api_base assert request_body["name"] == "test-vector-store" @@ -77,8 +70,19 @@ class TestOpenAIVectorStoreAPIConfig: litellm_params={}, ) - assert ( - url - == "https://api.openai.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search" - ) + assert url == "https://api.openai.com/v1/vector_stores/..%2F..%2Ffiles%3Fx%3D1%23frag/search" assert request_body["query"] == "hello" + + def test_transform_search_vector_store_request_preserves_query_string(self): + config = OpenAIVectorStoreConfig() + + url, _ = config.transform_search_vector_store_request( + vector_store_id="vs_1", + query="hello", + vector_store_search_optional_params={}, + api_base="https://x.openai.azure.com/openai/vector_stores?api-version=2024-10-21", + litellm_logging_obj=None, + litellm_params={}, + ) + + assert url == "https://x.openai.azure.com/openai/vector_stores/vs_1/search?api-version=2024-10-21" From cfd83548b30cf0afdf9a5c9f77a573582dfa7256 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 09:04:37 +0000 Subject: [PATCH 063/253] refactor(proxy): drop narrating docstrings from the aggregated usage query path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/common_daily_activity.py | 14 -------------- 1 file changed, 14 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index cb81d57ac77..5f63f639a11 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -761,13 +761,6 @@ def _build_aggregated_sql_query( ) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params """Build the GROUPING SETS query for aggregated daily activity. - One statement, two UNION ALL arms over the same WHERE clause. The first arm is - key-free: grand total, per-date totals and the (date, model / model_group / - provider / mcp / endpoint) rollups, so its row count never grows with the number - of keys. The second arm emits the (date, , api_key) rollups for the - USAGE_TOP_API_KEYS_LIMIT highest-spend keys only. Both arms share the 7-bit - group_level bitmask (date, api_key, model, model_group, provider, mcp, endpoint). - Returns: Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw(). """ @@ -1366,13 +1359,6 @@ async def get_daily_activity_aggregated( ) -> SpendAnalyticsPaginatedResponse: """Aggregated variant that returns the full result set (no pagination). - Runs one GROUPING SETS statement with two UNION ALL arms: a key-free one for totals - and the model/provider/mcp/endpoint rollups (row count independent of key - cardinality) and a bounded one for the per-key rollups of the top - USAGE_TOP_API_KEYS_LIMIT keys by spend. breakdown.api_keys and every - api_key_breakdown therefore list at most that many keys, while the totals and the - key-free rollups cover every key. - include_entity_breakdown runs a small companion rollup query and folds `breakdown.entities` onto the response, as entity-scoped views like Team Usage need. From f863e746123bea7a6d0d1c730e68db35d8936854 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 09:22:20 +0000 Subject: [PATCH 064/253] test(proxy): assert aggregated usage behavior against Postgres instead of SQL text Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_common_daily_activity.py | 114 +++++++----------- 1 file changed, 41 insertions(+), 73 deletions(-) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 471cabbbb59..e1a4d0d02a8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1239,79 +1239,6 @@ class TestBuildAggregatedSqlQuery: assert "model = $4" in sql assert "api_key = $5" in sql - def test_model_group_rollups_fall_back_to_model_name(self): - """Aggregated model_groups rollups must fall back to model for group-less rows. - - The (date, model_group) grouping level cannot recover the model column - after the fact (it is rolled up), so the fallback has to happen in SQL; - without it, group-less rows silently vanish from the model_groups - breakdown that the usage UI now renders by default. Group-less rows are - stored as empty strings, not NULL (spend_tracking_utils defaults - model_group to ""), so a plain COALESCE is not enough: the fallback must - be NULLIF-wrapped to catch both - """ - sql, _ = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id=None, - start_date="2026-07-01", - end_date="2026-07-01", - model=None, - api_key=None, - ) - - normalized = " ".join(sql.split()) - fallback = "COALESCE(NULLIF(model_group, ''), model)" - assert f"{fallback} AS model_group" in normalized - assert f"GROUPING(model, {fallback}, custom_llm_provider, mcp_namespaced_tool_name, endpoint)" in normalized - assert f"(date, {fallback})," in normalized - assert "(date, model_group)" not in normalized - assert "COALESCE(model_group, model)" not in normalized - - def test_totals_arm_never_groups_by_api_key(self): - """The totals arm must not emit one row per key, that is what blew up the - query engine at 3k+ keys. Every grouping set there stays key-free and api_key - is projected as a NULL literal so the dispatcher's row shape is unchanged.""" - sql, _ = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id=None, - start_date="2026-07-01", - end_date="2026-07-01", - model=None, - api_key=None, - ) - - totals_arm, _ = " ".join(sql.split()).split("UNION ALL") - grouping_block = totals_arm.split("GROUP BY GROUPING SETS (", 1)[1] - assert "api_key" not in grouping_block - assert "NULL::text AS api_key" in totals_arm - - def test_per_key_arm_ranks_keys_deterministically_and_shares_filters(self): - """Both arms sit in one statement so totals and per-key rows come from the - same snapshot, and the per-key arm reuses the caller's filter params.""" - sql, params = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-05-29", - end_date="2026-06-02", - model="bedrock/global.anthropic.claude-opus-4-8", - api_key="sk-test", - timezone_offset_minutes=-330, - ) - - totals_arm, per_key_arm = " ".join(sql.split()).split("UNION ALL") - assert "top_api_keys" not in totals_arm - assert f"ORDER BY SUM(spend) DESC, api_key LIMIT {USAGE_TOP_API_KEYS_LIMIT}" in per_key_arm - assert "api_key IN (SELECT api_key FROM top_api_keys)" in per_key_arm - assert "api_key <> $6" in per_key_arm - assert per_key_arm.count("model = $4 AND api_key = $5") == 2 - grouping_block = per_key_arm.split("GROUP BY GROUPING SETS (", 1)[1] - assert grouping_block.count(", api_key)") == 6 - assert grouping_block.count("(date") == 6 - assert params[-1] == PTU_SENTINEL_API_KEY - class TestAggregatedEmptyEntityFilter: _BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query) @@ -1640,6 +1567,47 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"} +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_model_name( + _aggregated_postgresql: psycopg.Connection, +): + """Rows stored with an empty or NULL model_group must land in the model_groups + breakdown under their model name instead of vanishing from the usage UI.""" + rows: Final = [ + ("row-0", "user-0", "2026-06-01", "key-0", "gpt-5", "gpt-5-eu", "openai", "/v1/chat/completions", 10, 7.0, 1, 1), + ("row-1", "user-1", "2026-06-01", "key-1", "gpt-5", "", "openai", "/v1/chat/completions", 10, 3.0, 1, 1), + ("row-2", "user-2", "2026-06-01", "key-2", "claude-x", None, "anthropic", "/v1/messages", 10, 2.0, 1, 1), + ] + _seed_daily_user_spend(_aggregated_postgresql, rows) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity_aggregated( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + api_key=None, + ) + + breakdown: Final = result.results[0].breakdown + assert set(breakdown.model_groups) == {"gpt-5-eu", "gpt-5", "claude-x"} + assert breakdown.model_groups["gpt-5-eu"].metrics.spend == 7.0 + assert breakdown.model_groups["gpt-5"].metrics.spend == 3.0 + assert breakdown.model_groups["claude-x"].metrics.spend == 2.0 + assert set(breakdown.model_groups["gpt-5"].api_key_breakdown) == {"key-1"} + assert set(breakdown.models) == {"gpt-5", "claude-x"} + assert breakdown.models["gpt-5"].metrics.spend == 10.0 + + def _no_spend_record(): """A rollup row for a key with no spend, where SUM() returns NULL (None).""" return SimpleNamespace( From 61b3611b8c6667c8eae5dadc6036a6d95e897d04 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 09:22:20 +0000 Subject: [PATCH 065/253] fix(ui): block the global usage export when the aggregated key cap is reached Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../components/EntityUsage/EntityUsage.tsx | 2 +- .../_components/components/UsagePageView.tsx | 3 +- .../exportBlockedReason.test.ts | 37 ++++++++++++++++++- .../EntityUsageExport/exportBlockedReason.ts | 18 ++++++++- 4 files changed, 56 insertions(+), 4 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 6687bd4df03..27a4163dff7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -664,7 +664,7 @@ const EntityUsage: React.FC = ({ { key: "endpoints", label: "Endpoint Activity", content: }, ]; - const spendFetchState = { coversRange, cancelled, failed }; + const spendFetchState = { coversRange, cancelled, failed, apiKeyLimitReached: undefined }; return (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index 691c5dc839a..48322477fb4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -30,7 +30,7 @@ import { ActivityMetrics, processActivityData } from "@/components/activity_metr import CloudZeroExportModal from "@/components/cloudzero_export_modal"; import UserDropdown from "@/components/common_components/UserDropdown"; import EntityUsageExportModal from "@/components/EntityUsageExport"; -import { getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; +import { getApiKeyLimitReached, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; import KeyActivityPanel from "@/components/UsagePage/components/KeyActivityPanel"; import { Team } from "@/components/key_team_helpers/key_list"; import { @@ -256,6 +256,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { coversRange: activeAggregated !== null || paginatedResult.coversRange, cancelled: paginatedResult.cancelled, failed: paginatedResult.failed, + apiKeyLimitReached: getApiKeyLimitReached(userSpendData.results, userSpendData.metadata?.api_key_limit), }; const exportBlockedReason = getExportBlockedReason(spendFetchState); diff --git a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts index e39b01a5dea..86783e2e4a0 100644 --- a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts +++ b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts @@ -1,14 +1,24 @@ import { describe, expect, it } from "vitest"; -import { getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason"; +import type { DailyData } from "@/components/UsagePage/types"; + +import { getApiKeyLimitReached, getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason"; const state = (overrides: Partial = {}): UsageFetchState => ({ coversRange: true, cancelled: false, failed: false, + apiKeyLimitReached: undefined, ...overrides, }); +const dayWithKeys = (date: string, ...keys: string[]): DailyData => + ({ + date, + metrics: {}, + breakdown: { api_keys: Object.fromEntries(keys.map((k) => [k, { metrics: {}, metadata: {} }])) }, + }) as unknown as DailyData; + describe("getExportBlockedReason", () => { it("lets the export through once the data on screen covers the range", () => { expect(getExportBlockedReason(state())).toBeUndefined(); @@ -31,4 +41,29 @@ describe("getExportBlockedReason", () => { expect(reason).toMatch(/failed to load/i); expect(reason).not.toMatch(/stopped/i); }); + + it("blocks when the aggregated endpoint hit its key cap, since a per-team CSV would miss keys", () => { + const reason = getExportBlockedReason(state({ apiKeyLimitReached: 100 })); + + expect(reason).toMatch(/100 highest-spend keys/); + expect(reason).toMatch(/USAGE_TOP_API_KEYS_LIMIT/); + }); +}); + +describe("getApiKeyLimitReached", () => { + it("reports the cap once the distinct keys across every day reach it", () => { + const results = [dayWithKeys("2026-06-01", "key-1", "key-2"), dayWithKeys("2026-06-02", "key-2", "key-3")]; + + expect(getApiKeyLimitReached(results, 3)).toBe(3); + }); + + it("stays quiet while fewer keys than the cap came back, which means every key is on screen", () => { + const results = [dayWithKeys("2026-06-01", "key-1", "key-2"), dayWithKeys("2026-06-02", "key-2")]; + + expect(getApiKeyLimitReached(results, 3)).toBeUndefined(); + }); + + it("stays quiet when the response carries no cap, as the paginated fallback does", () => { + expect(getApiKeyLimitReached([dayWithKeys("2026-06-01", "key-1")], undefined)).toBeUndefined(); + }); }); diff --git a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts index 71408ba8f3f..ca756090351 100644 --- a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts +++ b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts @@ -1,13 +1,29 @@ +import type { DailyData } from "@/components/UsagePage/types"; + export interface UsageFetchState { coversRange: boolean; cancelled: boolean; failed: boolean; + apiKeyLimitReached: number | undefined; } -export const getExportBlockedReason = ({ coversRange, cancelled, failed }: UsageFetchState): string | undefined => { +export const getApiKeyLimitReached = (results: DailyData[], apiKeyLimit: unknown): number | undefined => { + if (typeof apiKeyLimit !== "number") return undefined; + const keys = new Set(results.flatMap((day) => Object.keys(day.breakdown.api_keys ?? {}))); + return keys.size >= apiKeyLimit ? apiKeyLimit : undefined; +}; + +export const getExportBlockedReason = ({ + coversRange, + cancelled, + failed, + apiKeyLimitReached, +}: UsageFetchState): string | undefined => { if (failed) return "Some spend data failed to load, so an export would under-report. Reload the page to try again."; if (cancelled) return "Loading was stopped before the whole range arrived, so an export would under-report. Reload the page to load it all."; if (!coversRange) return "Spend data is still loading, so an export would under-report. Wait for it to finish."; + if (apiKeyLimitReached !== undefined) + return `Only the ${apiKeyLimitReached} highest-spend keys were loaded, so a per-team export would under-report. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys.`; return undefined; }; From aec36a2bac12222a681f6552103b6899657abc04 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 09:36:13 +0000 Subject: [PATCH 066/253] style(proxy): drop em dash from api_key bit comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_endpoints/common_daily_activity.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 5f63f639a11..93e6e60c77c 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -991,7 +991,7 @@ async def _aggregate_spend_records( # current grouping set's key), 0 when the column is part of the key. _GROUP_GRAND_TOTAL: Final = 127 # 0b1111111 — all rolled up _GROUP_DATE: Final = 63 # 0b0111111 — only date kept -_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000 — api_key position in the 7-bit mask +_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000 _GROUP_DATE_API_KEY: Final = 31 # 0b0011111 _GROUP_DATE_MODEL: Final = 47 # 0b0101111 _GROUP_DATE_MODEL_API_KEY: Final = 15 # 0b0001111 From 87894f2e7f343a75cc6d4f0977deceb0d564dde9 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 10:02:11 +0000 Subject: [PATCH 067/253] fix(proxy): report total_api_keys so exact-limit key sets are not treated as truncated Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 12 ++++ .../common_daily_activity.py | 15 ++-- .../common_daily_activity.py | 5 ++ .../test_common_daily_activity.py | 68 +++++++++++++++++-- .../components/EntityUsage/EntityUsage.tsx | 2 +- .../_components/components/UsagePageView.tsx | 7 +- .../exportBlockedReason.test.ts | 37 ++++------ .../EntityUsageExport/exportBlockedReason.ts | 20 +++--- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 ++ 9 files changed, 126 insertions(+), 45 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 5ff35747779..fb82a540761 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3072,6 +3072,18 @@ "title": "Page", "type": "integer" }, + "total_api_keys": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "description": "Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key lists are truncated to the highest-spend keys.", + "title": "Total Api Keys" + }, "total_api_requests": { "default": 0, "title": "Total Api Requests", diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 93e6e60c77c..a36b6052981 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -155,6 +155,7 @@ class _GroupingSetsRow(SimpleNamespace): mcp_namespaced_tool_name: str | None endpoint: str | None group_level: int + distinct_api_keys: int | None spend: float | None prompt_tokens: int | None completion_tokens: int | None @@ -800,7 +801,8 @@ def _build_aggregated_sql_query( (GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT} | GROUPING(model, {_MODEL_GROUP_EXPR}, custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level,{metric_select} + endpoint) AS group_level, + NULL::bigint AS distinct_api_keys,{metric_select} FROM "{pg_table}" WHERE {where_clause} GROUP BY GROUPING SETS ( @@ -814,7 +816,7 @@ def _build_aggregated_sql_query( )) UNION ALL (WITH top_api_keys AS ( - SELECT api_key + SELECT api_key, COUNT(*) OVER () AS distinct_api_keys FROM "{pg_table}" WHERE {where_clause} AND api_key <> {sentinel_param} GROUP BY api_key @@ -831,9 +833,10 @@ def _build_aggregated_sql_query( endpoint, GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR}, custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level,{metric_select} - FROM "{pg_table}" - WHERE {where_clause} AND api_key IN (SELECT api_key FROM top_api_keys) + endpoint) AS group_level, + MAX(top_api_keys.distinct_api_keys) AS distinct_api_keys,{metric_select} + FROM "{pg_table}" JOIN top_api_keys USING (api_key) + WHERE {where_clause} GROUP BY GROUPING SETS ( (date, api_key), (date, model, api_key), @@ -1398,6 +1401,7 @@ async def get_daily_activity_aggregated( ) records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or ())] + total_api_keys: Final = next((r.distinct_api_keys for r in records if r.distinct_api_keys is not None), 0) # The grouping-sets dispatcher places each row directly in its bucket # using the row's GROUPING() bitmask. No Python-side summing needed. @@ -1452,6 +1456,7 @@ async def get_daily_activity_aggregated( total_pages=1, has_more=False, api_key_limit=USAGE_TOP_API_KEYS_LIMIT, + total_api_keys=total_api_keys, ), ) diff --git a/litellm/types/proxy/management_endpoints/common_daily_activity.py b/litellm/types/proxy/management_endpoints/common_daily_activity.py index 2893eb30b3c..5d42b1230a0 100644 --- a/litellm/types/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/types/proxy/management_endpoints/common_daily_activity.py @@ -105,6 +105,11 @@ class DailySpendMetadata(BaseModel): description="When set, api_keys and every api_key_breakdown list at most this many keys, " "ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.", ) + total_api_keys: int | None = Field( + default=None, + description="Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key " + "lists are truncated to the highest-spend keys.", + ) class SpendAnalyticsPaginatedResponse(BaseModel): diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index e1a4d0d02a8..e7f8bcc7e4f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -175,6 +175,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/chat/completions", "api_key": None, "group_level": 62, + "distinct_api_keys": None, "spend": 15.0, "prompt_tokens": 150, "completion_tokens": 75, @@ -187,6 +188,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/embeddings", "api_key": None, "group_level": 62, + "distinct_api_keys": None, "spend": 3.0, "prompt_tokens": 30, "completion_tokens": 0, @@ -200,6 +202,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": None, "api_key": None, "group_level": 63, + "distinct_api_keys": None, "spend": 18.0, "prompt_tokens": 180, "completion_tokens": 75, @@ -213,6 +216,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": None, "api_key": None, "group_level": 127, + "distinct_api_keys": None, "spend": 18.0, "prompt_tokens": 180, "completion_tokens": 75, @@ -226,6 +230,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/chat/completions", "api_key": "key-1", "group_level": 30, + "distinct_api_keys": 2, "spend": 15.0, "prompt_tokens": 150, "completion_tokens": 75, @@ -238,6 +243,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/embeddings", "api_key": "key-2", "group_level": 30, + "distinct_api_keys": 2, "spend": 3.0, "prompt_tokens": 30, "completion_tokens": 0, @@ -839,6 +845,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "endpoint": "/v1/chat/completions", "api_key": None, "group_level": 62, + "distinct_api_keys": None, "spend": 10.0, "prompt_tokens": 100, "completion_tokens": 50, @@ -851,6 +858,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "endpoint": "/v1/chat/completions", "api_key": "deleted-key-hash", "group_level": 30, + "distinct_api_keys": 1, "spend": 10.0, "prompt_tokens": 100, "completion_tokens": 50, @@ -1316,6 +1324,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): "mcp_namespaced_tool_name": None, "endpoint": None, "group_level": 127, + "distinct_api_keys": None, "spend": None, "prompt_tokens": None, "completion_tokens": None, @@ -1499,6 +1508,7 @@ async def test_get_daily_activity_aggregated_bounds_api_key_rollups( assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0) assert result.metadata.total_api_requests == n_keys assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT + assert result.metadata.total_api_keys == n_keys expected_top: Final = {f"key-{i:03d}" for i in range(6, n_keys)} | {"key-004"} day: Final = result.results[0] @@ -1560,6 +1570,7 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both ) assert result.metadata.total_spend == 2.0 + assert result.metadata.total_api_keys == 1 day: Final = result.results[0] assert set(day.breakdown.api_keys) == {"key-1"} assert day.breakdown.api_keys["key-1"].metrics.spend == 2.0 @@ -1567,6 +1578,55 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_both assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"} +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_complete( + _aggregated_postgresql: psycopg.Connection, +): + """With exactly USAGE_TOP_API_KEYS_LIMIT keys nothing is dropped, and the + response must say so: total_api_keys equals the limit rather than exceeding it.""" + rows: Final = [ + ( + f"row-{i:03d}", + f"user-{i:03d}", + "2026-06-01", + f"key-{i:03d}", + "gpt-5", + "", + "openai", + "/v1/chat/completions", + 10, + float(i + 1), + 1, + 1, + ) + for i in range(USAGE_TOP_API_KEYS_LIMIT) + ] + _seed_daily_user_spend(_aggregated_postgresql, rows) + + row_counts: Final[list[int]] = [] # mutable-ok: out-param for the query_raw shim + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, row_counts) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity_aggregated( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + api_key=None, + ) + + assert result.metadata.total_api_keys == USAGE_TOP_API_KEYS_LIMIT + assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT + assert len(result.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT + + @pytest.mark.asyncio async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_model_name( _aggregated_postgresql: psycopg.Connection, @@ -2427,10 +2487,10 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): "successful_requests": 0, } main_rows = [ - {**base, "date": None, "group_level": 127, "spend": 18.0}, - {**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0}, - {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0}, - {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}, + {**base, "date": None, "group_level": 127, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "group_level": 63, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "distinct_api_keys": 1, "spend": 12.0}, ] entity_base = { key: value diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 27a4163dff7..001fee4b7bc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -664,7 +664,7 @@ const EntityUsage: React.FC = ({ { key: "endpoints", label: "Endpoint Activity", content: }, ]; - const spendFetchState = { coversRange, cancelled, failed, apiKeyLimitReached: undefined }; + const spendFetchState = { coversRange, cancelled, failed, apiKeyTruncation: undefined }; return (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index 48322477fb4..3a97e54edfc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -30,7 +30,7 @@ import { ActivityMetrics, processActivityData } from "@/components/activity_metr import CloudZeroExportModal from "@/components/cloudzero_export_modal"; import UserDropdown from "@/components/common_components/UserDropdown"; import EntityUsageExportModal from "@/components/EntityUsageExport"; -import { getApiKeyLimitReached, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; +import { getApiKeyTruncation, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; import KeyActivityPanel from "@/components/UsagePage/components/KeyActivityPanel"; import { Team } from "@/components/key_team_helpers/key_list"; import { @@ -256,7 +256,10 @@ const UsagePage: React.FC = ({ teams, organizations }) => { coversRange: activeAggregated !== null || paginatedResult.coversRange, cancelled: paginatedResult.cancelled, failed: paginatedResult.failed, - apiKeyLimitReached: getApiKeyLimitReached(userSpendData.results, userSpendData.metadata?.api_key_limit), + apiKeyTruncation: getApiKeyTruncation( + userSpendData.metadata?.api_key_limit, + userSpendData.metadata?.total_api_keys, + ), }; const exportBlockedReason = getExportBlockedReason(spendFetchState); diff --git a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts index 86783e2e4a0..8491b31f5f9 100644 --- a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts +++ b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.test.ts @@ -1,24 +1,15 @@ import { describe, expect, it } from "vitest"; -import type { DailyData } from "@/components/UsagePage/types"; - -import { getApiKeyLimitReached, getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason"; +import { getApiKeyTruncation, getExportBlockedReason, type UsageFetchState } from "./exportBlockedReason"; const state = (overrides: Partial = {}): UsageFetchState => ({ coversRange: true, cancelled: false, failed: false, - apiKeyLimitReached: undefined, + apiKeyTruncation: undefined, ...overrides, }); -const dayWithKeys = (date: string, ...keys: string[]): DailyData => - ({ - date, - metrics: {}, - breakdown: { api_keys: Object.fromEntries(keys.map((k) => [k, { metrics: {}, metadata: {} }])) }, - }) as unknown as DailyData; - describe("getExportBlockedReason", () => { it("lets the export through once the data on screen covers the range", () => { expect(getExportBlockedReason(state())).toBeUndefined(); @@ -42,28 +33,26 @@ describe("getExportBlockedReason", () => { expect(reason).not.toMatch(/stopped/i); }); - it("blocks when the aggregated endpoint hit its key cap, since a per-team CSV would miss keys", () => { - const reason = getExportBlockedReason(state({ apiKeyLimitReached: 100 })); + it("blocks when the aggregated endpoint dropped keys, since a per-team CSV would miss them", () => { + const reason = getExportBlockedReason(state({ apiKeyTruncation: { limit: 100, total: 3000 } })); - expect(reason).toMatch(/100 highest-spend keys/); + expect(reason).toMatch(/100 highest-spend keys of 3000/); expect(reason).toMatch(/USAGE_TOP_API_KEYS_LIMIT/); }); }); -describe("getApiKeyLimitReached", () => { - it("reports the cap once the distinct keys across every day reach it", () => { - const results = [dayWithKeys("2026-06-01", "key-1", "key-2"), dayWithKeys("2026-06-02", "key-2", "key-3")]; - - expect(getApiKeyLimitReached(results, 3)).toBe(3); +describe("getApiKeyTruncation", () => { + it("reports truncation once the proxy saw more keys than it returned", () => { + expect(getApiKeyTruncation(100, 101)).toEqual({ limit: 100, total: 101 }); }); - it("stays quiet while fewer keys than the cap came back, which means every key is on screen", () => { - const results = [dayWithKeys("2026-06-01", "key-1", "key-2"), dayWithKeys("2026-06-02", "key-2")]; - - expect(getApiKeyLimitReached(results, 3)).toBeUndefined(); + it("stays quiet when exactly the cap exists, since every key is on screen", () => { + expect(getApiKeyTruncation(100, 100)).toBeUndefined(); + expect(getApiKeyTruncation(100, 7)).toBeUndefined(); }); it("stays quiet when the response carries no cap, as the paginated fallback does", () => { - expect(getApiKeyLimitReached([dayWithKeys("2026-06-01", "key-1")], undefined)).toBeUndefined(); + expect(getApiKeyTruncation(undefined, undefined)).toBeUndefined(); + expect(getApiKeyTruncation(100, null)).toBeUndefined(); }); }); diff --git a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts index ca756090351..6c5a5f83231 100644 --- a/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts +++ b/ui/litellm-dashboard/src/components/EntityUsageExport/exportBlockedReason.ts @@ -1,29 +1,31 @@ -import type { DailyData } from "@/components/UsagePage/types"; +export interface ApiKeyTruncation { + limit: number; + total: number; +} export interface UsageFetchState { coversRange: boolean; cancelled: boolean; failed: boolean; - apiKeyLimitReached: number | undefined; + apiKeyTruncation: ApiKeyTruncation | undefined; } -export const getApiKeyLimitReached = (results: DailyData[], apiKeyLimit: unknown): number | undefined => { - if (typeof apiKeyLimit !== "number") return undefined; - const keys = new Set(results.flatMap((day) => Object.keys(day.breakdown.api_keys ?? {}))); - return keys.size >= apiKeyLimit ? apiKeyLimit : undefined; +export const getApiKeyTruncation = (apiKeyLimit: unknown, totalApiKeys: unknown): ApiKeyTruncation | undefined => { + if (typeof apiKeyLimit !== "number" || typeof totalApiKeys !== "number") return undefined; + return totalApiKeys > apiKeyLimit ? { limit: apiKeyLimit, total: totalApiKeys } : undefined; }; export const getExportBlockedReason = ({ coversRange, cancelled, failed, - apiKeyLimitReached, + apiKeyTruncation, }: UsageFetchState): string | undefined => { if (failed) return "Some spend data failed to load, so an export would under-report. Reload the page to try again."; if (cancelled) return "Loading was stopped before the whole range arrived, so an export would under-report. Reload the page to load it all."; if (!coversRange) return "Spend data is still loading, so an export would under-report. Wait for it to finish."; - if (apiKeyLimitReached !== undefined) - return `Only the ${apiKeyLimitReached} highest-spend keys were loaded, so a per-team export would under-report. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys.`; + if (apiKeyTruncation !== undefined) + return `Only the ${apiKeyTruncation.limit} highest-spend keys of ${apiKeyTruncation.total} were loaded, so a per-team export would under-report. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys.`; return undefined; }; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c9e9c550afa..84656e7eee1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27503,6 +27503,11 @@ export interface components { * @default 1 */ page: number; + /** + * Total Api Keys + * @description Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key lists are truncated to the highest-spend keys. + */ + total_api_keys?: number | null; /** * Total Api Requests * @default 0 From 38c1139377279b061b808d12b9b0f8b22c05f339 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 10:13:05 +0000 Subject: [PATCH 068/253] feat(ui): note on the Key Activity tab when only the top-spend keys were loaded Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../usage/_components/components/UsagePageView.tsx | 2 +- .../UsagePage/components/KeyActivityPanel.test.tsx | 10 ++++++++++ .../UsagePage/components/KeyActivityPanel.tsx | 14 +++++++++++++- 3 files changed, 24 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index 3a97e54edfc..a9ab0f17f40 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -908,7 +908,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { - + diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.test.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.test.tsx index 693ac20a360..830139143e9 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.test.tsx @@ -68,4 +68,14 @@ describe("KeyActivityPanel", () => { expect(screen.getByLabelText("Search keys")).toHaveValue(""); expect(screen.getByTestId("rendered-keys")).toHaveTextContent("hash-alicehash-bob"); }); + + it("says how many keys the proxy left out when only the top spenders were loaded", () => { + render(); + expect(screen.getByRole("note")).toHaveTextContent("Only the 2 highest-spend keys of 3,000 are loaded"); + }); + + it("shows no truncation note when every key is loaded", () => { + render(); + expect(screen.queryByRole("note")).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.tsx index 8287a04d0c7..8b2141f8528 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/KeyActivityPanel.tsx @@ -2,6 +2,7 @@ import { Search, X } from "lucide-react"; import React, { useMemo, useState } from "react"; import { ActivityMetrics } from "@/components/activity_metrics"; +import type { ApiKeyTruncation } from "@/components/EntityUsageExport/exportBlockedReason"; import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; import { filterKeyActivity } from "../keyActivityFilter"; @@ -10,9 +11,14 @@ import type { ModelActivityData } from "../types"; interface KeyActivityPanelProps { keyMetrics: Record; hidePromptCachingMetrics?: boolean; + apiKeyTruncation?: ApiKeyTruncation; } -const KeyActivityPanel: React.FC = ({ keyMetrics, hidePromptCachingMetrics = false }) => { +const KeyActivityPanel: React.FC = ({ + keyMetrics, + hidePromptCachingMetrics = false, + apiKeyTruncation, +}) => { const [query, setQuery] = useState(""); const filtered = useMemo(() => filterKeyActivity(keyMetrics, query), [keyMetrics, query]); const totalKeys = Object.keys(keyMetrics).length; @@ -43,6 +49,12 @@ const KeyActivityPanel: React.FC = ({ keyMetrics, hidePro Showing {shownKeys.toLocaleString()} of {totalKeys.toLocaleString()} keys + {apiKeyTruncation !== undefined && ( + + Only the {apiKeyTruncation.limit.toLocaleString()} highest-spend keys of{" "} + {apiKeyTruncation.total.toLocaleString()} are loaded + + )}
{isFiltering && totalKeys > 0 && shownKeys === 0 ? (

From 09fa0833b43820c039f45027ce26663de5a5dc54 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 10:17:46 +0000 Subject: [PATCH 069/253] feat(ui): surface top-key truncation on the team usage view Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../components/EntityUsage/EntityUsage.test.tsx | 17 +++++++++++++++++ .../components/EntityUsage/EntityUsage.tsx | 15 ++++++++++++--- 2 files changed, 29 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 2a6c2ede478..5846a63bc70 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -569,6 +569,23 @@ describe("EntityUsage", () => { expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument(); }); + it("tells the team view how many keys the proxy left out of the per-key lists", async () => { + mockTeamDailyActivityAggregatedCall.mockResolvedValue({ + ...mockSpendData, + metadata: { ...mockSpendData.metadata, api_key_limit: 100, total_api_keys: 3000 }, + }); + render(); + + await waitFor(() => { + expect(mockTeamDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + act(() => { + fireEvent.click(screen.getByText("Key Activity")); + }); + + expect(await screen.findByRole("note")).toHaveTextContent("Only the 100 highest-spend keys of 3,000 are loaded"); + }); + // An inactive tab panel is marked aria-selected="false" by one tab library and hidden by the // other, so treat either as "not on screen" and the assertion holds whichever one is rendering. const isShowing = (element: HTMLElement): boolean => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 001fee4b7bc..27460b21108 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -25,7 +25,7 @@ import TeamMultiSelect from "@/components/common_components/team_multi_select"; import UserDropdown from "@/components/common_components/UserDropdown"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; import { UsageExportHeader } from "@/components/EntityUsageExport"; -import { getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; +import { getApiKeyTruncation, getExportBlockedReason } from "@/components/EntityUsageExport/exportBlockedReason"; import type { EntityType } from "@/components/EntityUsageExport/types"; import { agentDailyActivityCall, @@ -71,6 +71,8 @@ interface EntitySpendData { total_successful_requests: number; total_failed_requests: number; total_tokens: number; + api_key_limit?: number | null; + total_api_keys?: number | null; }; } @@ -160,6 +162,7 @@ const EntityUsage: React.FC = ({ }); const spendData = spendDataRaw as unknown as EntitySpendData; + const apiKeyTruncation = getApiKeyTruncation(spendData.metadata?.api_key_limit, spendData.metadata?.total_api_keys); const { data: agentSpendDataRaw, @@ -659,12 +662,18 @@ const EntityUsage: React.FC = ({ { key: "keys", label: "Key Activity", - content: , + content: ( + + ), }, { key: "endpoints", label: "Endpoint Activity", content: }, ]; - const spendFetchState = { coversRange, cancelled, failed, apiKeyTruncation: undefined }; + const spendFetchState = { coversRange, cancelled, failed, apiKeyTruncation }; return (

From 303a9058cc9db4ef7251c8e7ecbb90a63e17c36e Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 10:21:33 +0000 Subject: [PATCH 070/253] refactor(proxy): type entity rollup rows by their own projection Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../common_daily_activity.py | 25 +++++++++++-------- 1 file changed, 14 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index a36b6052981..a5e292b177d 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -146,16 +146,9 @@ class _AggregatedSpendData(TypedDict): totals: SpendMetrics -class _GroupingSetsRow(SimpleNamespace): +class _RollupMetricsRow(SimpleNamespace): date: str api_key: str | None - model: str | None - model_group: str | None - custom_llm_provider: str | None - mcp_namespaced_tool_name: str | None - endpoint: str | None - group_level: int - distinct_api_keys: int | None spend: float | None prompt_tokens: int | None completion_tokens: int | None @@ -173,7 +166,17 @@ class _GroupingSetsRow(SimpleNamespace): timed_requests: int | None -class _EntityRollupRow(_GroupingSetsRow): +class _GroupingSetsRow(_RollupMetricsRow): + model: str | None + model_group: str | None + custom_llm_provider: str | None + mcp_namespaced_tool_name: str | None + endpoint: str | None + group_level: int + distinct_api_keys: int | None + + +class _EntityRollupRow(_RollupMetricsRow): entity_id: str | None api_key_rolled: int @@ -202,7 +205,7 @@ async def _query_raw_optional( return await prisma_client.db.query_raw(query[0], *query[1]) -def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float: +def _reported_flat_cost(record: DailySpendRecord | _RollupMetricsRow) -> float: """Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled. Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost`` @@ -1008,7 +1011,7 @@ _GROUP_DATE_ENDPOINT: Final = 62 # 0b0111110 _GROUP_DATE_ENDPOINT_API_KEY: Final = 30 # 0b0011110 -def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics: +def _record_to_spend_metrics(record: _RollupMetricsRow) -> SpendMetrics: """Build a SpendMetrics directly from one already-aggregated rollup row. SUM() over zero rows is SQL NULL, so rollup rows (notably the grand-total From b74733d4e6971984d3af6e78c565ff37c47cffe2 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 10:53:04 +0000 Subject: [PATCH 071/253] feat(ui): note on the cache leakage card when only the top-spend keys were loaded Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_components/CacheLeakageCard.test.tsx | 24 +++++++++++++++++++ .../_components/CacheLeakageCard.tsx | 9 ++++++- .../useDailyActivityRange.test.tsx | 17 ++++++++++++- .../_components/useDailyActivityRange.ts | 3 +++ 4 files changed, 51 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx index 54af13d8a90..8cb18344734 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx @@ -178,4 +178,28 @@ describe("CacheLeakageCard", () => { screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."), ).not.toBeInTheDocument(); }); + + it("says which keys are missing from the key ranking when the proxy capped the per-key lists", () => { + const day = dayWithKeys("2026-07-12", { + "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + }); + renderWith([day], { apiKeyTruncation: { limit: 100, total: 3000 } }); + + expect(screen.getByRole("note")).toHaveTextContent( + "Only the 100 highest-spend keys of 3,000 are loaded, so a lower-spend key that leaks more is not listed here.", + ); + + fireEvent.click(screen.getByRole("tab", { name: "By model" })); + + expect(screen.queryByRole("note")).not.toBeInTheDocument(); + }); + + it("keeps the key ranking note off when every key was loaded", () => { + const day = dayWithKeys("2026-07-12", { + "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + }); + renderWith([day]); + + expect(screen.queryByRole("note")).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx index 3f27449ebe1..a0877b04648 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx @@ -81,7 +81,7 @@ const SortableHead = ({ }; const CacheLeakageCard: React.FC = ({ activity }) => { - const { dateValue, onDateChange, results, loading, isFetchingMore } = activity; + const { dateValue, onDateChange, results, loading, isFetchingMore, apiKeyTruncation } = activity; const [dimension, setDimension] = useState("key"); const [sort, setSort] = useState({ column: "potentialSavings", dir: "desc" }); const leakage = useMemo(() => computeCacheLeakage(results, dimension), [results, dimension]); @@ -123,6 +123,13 @@ const CacheLeakageCard: React.FC = ({ activity }) => { + {dimension === "key" && apiKeyTruncation !== undefined && ( +

+ Only the {apiKeyTruncation.limit.toLocaleString()} highest-spend keys of{" "} + {apiKeyTruncation.total.toLocaleString()} are loaded, so a lower-spend key that leaks more is not listed + here. Raise USAGE_TOP_API_KEYS_LIMIT on the proxy to load more keys. +

+ )} {rows.length > 0 && isFetchingMore && (

Data is still loading; rows and totals will update as the rest of the range arrives. diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx index 00902aa9fdd..4059303d5a5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx @@ -4,12 +4,13 @@ import { describe, expect, it, vi } from "vitest"; const mockUsePaginatedDailyActivity = vi.fn(); const mockCancel = vi.fn(); +let mockMetadata: Record = {}; vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({ usePaginatedDailyActivity: (args: unknown) => { mockUsePaginatedDailyActivity(args); return { - data: { results: [] }, + data: { results: [], metadata: mockMetadata }, loading: false, isFetchingMore: false, progress: { currentPage: 4, totalPages: 9 }, @@ -80,4 +81,18 @@ describe("useDailyActivityRange", () => { expect(mockUsePaginatedDailyActivity).toHaveBeenLastCalledWith(expect.objectContaining({ enabled: false })); }); + + it("reports how many keys the proxy left out of the per-key lists", () => { + mockMetadata = { api_key_limit: 100, total_api_keys: 3000 }; + const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); + + expect(result.current.apiKeyTruncation).toEqual({ limit: 100, total: 3000 }); + }); + + it("reports no key truncation when every key fit under the proxy limit", () => { + mockMetadata = { api_key_limit: 100, total_api_keys: 100 }; + const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); + + expect(result.current.apiKeyTruncation).toBeUndefined(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts index 9f793a68bf5..92dd24b8d6d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts @@ -1,6 +1,7 @@ import { useMemo, useState } from "react"; import { userDailyActivityAggregatedCall, userDailyActivityCall } from "@/components/networking"; +import { ApiKeyTruncation, getApiKeyTruncation } from "@/components/EntityUsageExport/exportBlockedReason"; import { DailyData } from "@/components/UsagePage/types"; import { spendScopeUserId } from "@/utils/roles"; import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity"; @@ -22,6 +23,7 @@ export interface DailyActivityRange { cancelled: boolean; failed: boolean; cancel: () => void; + apiKeyTruncation?: ApiKeyTruncation; } /** @@ -78,6 +80,7 @@ export const useScopedDailyActivityRange = ( cancelled, failed, cancel, + apiKeyTruncation: getApiKeyTruncation(data.metadata?.api_key_limit, data.metadata?.total_api_keys), }; }; From 834313af4b188ee5561e8c0e8094fad399e5ac78 Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 12:13:57 +0000 Subject: [PATCH 072/253] test(proxy): assert the global rollup split and scheduler through behavior, not SQL text or add_job arguments Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_common_daily_activity.py | 95 ++++++++++--------- .../proxy/proxy_server/test_lifecycle.py | 24 +++-- 2 files changed, 66 insertions(+), 53 deletions(-) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index ac761315a27..fc3ede88aa9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1647,30 +1647,6 @@ async def test_global_rollup_marker_read_failure_falls_back_to_the_per_key_table await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) -def test_aggregated_sql_splits_the_key_free_arm_at_the_marker_and_keeps_the_key_arm_per_key(): - sql, params = _build_aggregated_sql_query(**_unfiltered_user_query(), global_rollup_through="2026-06-01") - marker_param: Final = f"${len(params)}" - - assert params[-1] == "2026-06-01" - assert ( - f'FROM "LiteLLM_DailyGlobalSpend"\n WHERE date >= $1 AND date <= $2 AND date <= {marker_param}' - in sql - ) - assert ( - f'FROM "LiteLLM_DailyUserSpend"\n WHERE date >= $1 AND date <= $2 AND date > {marker_param}' in sql - ) - key_arm: Final = sql.split("UNION ALL\n (WITH top_api_keys")[1] - assert "LiteLLM_DailyGlobalSpend" not in key_arm - assert marker_param not in key_arm - - -def test_aggregated_sql_without_a_marker_reads_the_per_key_table_only(): - sql, params = _build_aggregated_sql_query(**_unfiltered_user_query()) - - assert "LiteLLM_DailyGlobalSpend" not in sql - assert params[-1] == PTU_SENTINEL_API_KEY - - _GLOBAL_SPEND_MIGRATION: Final = ( pathlib.Path(__file__).resolve().parents[4] / "litellm-proxy-extras" @@ -1687,7 +1663,10 @@ async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_ ): """Day 1 is rolled up and day 2 is still open (never rolled up), so a marker of day 1 must give the same response as reading everything per-key: day 1 from the global table, day 2 - live, one grand total across both. The per-key arm stays on the user table throughout.""" + live, one grand total across both. Per-key rows that land after the rollup then tell the + two sources apart: a late day 1 row is invisible to totals until the next reconcile while a + late day 2 row shows up at once, and both keys rank in the key breakdown, which stays + per-key throughout.""" n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 3 rows: Final = [ ( @@ -1720,39 +1699,56 @@ async def test_get_daily_activity_aggregated_serves_closed_days_from_the_global_ ) _aggregated_postgresql.commit() - async def read(marker: str | None, sql_seen: list[str]): + async def read(marker: str | None): await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) prisma = _prisma_with_marker(marker) - run_query = _psycopg_query_raw(_aggregated_postgresql, []) - - async def query_raw(sql: str, *params: str): - sql_seen.append(sql) - return await run_query(sql, *params) - - prisma.db.query_raw = query_raw + prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, []) return await get_daily_activity_aggregated( prisma_client=prisma, entity_metadata_field=None, **_unfiltered_user_query(), ) - per_key_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim - global_sql: Final[list[str]] = [] # mutable-ok: out-param for the query_raw shim - from_per_key = await read(None, per_key_sql) - from_global = await read("2026-06-01", global_sql) - await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + from_per_key = await read(None) + from_global = await read("2026-06-01") - assert per_key_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 0 - assert global_sql[0].count('FROM "LiteLLM_DailyGlobalSpend"') == 1 - assert global_sql[0].count('FROM "LiteLLM_DailyUserSpend"') == 3 assert from_global.model_dump() == from_per_key.model_dump() - assert from_global.metadata.total_spend == pytest.approx(2 * sum(float(i + 1) for i in range(n_keys))) + seeded_spend: Final = 2 * sum(float(i + 1) for i in range(n_keys)) + assert from_global.metadata.total_spend == pytest.approx(seeded_spend) assert from_global.metadata.total_response_time_ms == 2 * n_keys * 10 * 25 assert from_global.metadata.total_timed_requests == 2 * n_keys assert {day.date.isoformat() for day in from_global.results} == {"2026-06-01", "2026-06-02"} assert len(from_global.results[0].breakdown.api_keys) == USAGE_TOP_API_KEYS_LIMIT assert set(from_global.results[0].breakdown.model_groups) == {"gpt-5", "claude"} + with _aggregated_postgresql.cursor() as cur: + cur.executemany( + """ + INSERT INTO "LiteLLM_DailyUserSpend" + (id, user_id, date, api_key, model, model_group, custom_llm_provider, + endpoint, prompt_tokens, spend, api_requests, successful_requests) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """, + [ + ("late-1", "user-late", "2026-06-01", "key-late-1", "gpt-5", "", "openai", None, 10, 1000.0, 1, 1), + ("late-2", "user-late", "2026-06-02", "key-late-2", "gpt-5", "", "openai", None, 10, 500.0, 1, 1), + ], + ) + _aggregated_postgresql.commit() + + late_per_key = await read(None) + late_global = await read("2026-06-01") + await evict_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + + assert late_per_key.metadata.total_spend == pytest.approx(seeded_spend + 1000.0 + 500.0) + assert late_global.metadata.total_spend == pytest.approx(seeded_spend + 500.0) + by_day: Final = {day.date.isoformat(): day for day in late_global.results} + assert by_day["2026-06-01"].metrics.spend == pytest.approx(seeded_spend / 2) + assert by_day["2026-06-02"].metrics.spend == pytest.approx(seeded_spend / 2 + 500.0) + assert by_day["2026-06-01"].breakdown.api_keys["key-late-1"].metrics.spend == pytest.approx(1000.0) + assert by_day["2026-06-02"].breakdown.api_keys["key-late-2"].metrics.spend == pytest.approx(500.0) + assert late_global.metadata.total_api_keys == n_keys + 2 + @pytest.mark.asyncio async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_complete( @@ -1810,7 +1806,20 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo """Rows stored with an empty or NULL model_group must land in the model_groups breakdown under their model name instead of vanishing from the usage UI.""" rows: Final = [ - ("row-0", "user-0", "2026-06-01", "key-0", "gpt-5", "gpt-5-eu", "openai", "/v1/chat/completions", 10, 7.0, 1, 1), + ( + "row-0", + "user-0", + "2026-06-01", + "key-0", + "gpt-5", + "gpt-5-eu", + "openai", + "/v1/chat/completions", + 10, + 7.0, + 1, + 1, + ), ("row-1", "user-1", "2026-06-01", "key-1", "gpt-5", "", "openai", "/v1/chat/completions", 10, 3.0, 1, 1), ("row-2", "user-2", "2026-06-01", "key-2", "claude-x", None, "anthropic", "/v1/messages", 10, 2.0, 1, 1), ] diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index ee72e98ffa9..6121608b658 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -28,6 +28,7 @@ from typing import List, Optional, Union from unittest.mock import AsyncMock, MagicMock, patch import pytest +from apscheduler.schedulers.asyncio import AsyncIOScheduler from fastapi import FastAPI from pydantic import BaseModel from typing_extensions import TypedDict @@ -1042,8 +1043,8 @@ async def test_spend_report_locks_are_never_released(): proxy_logging_obj.db_spend_update_writer.pod_lock_manager.release_lock.assert_not_awaited() -def _init_daily_global_spend_reconcile_job() -> tuple[MagicMock, MagicMock, MagicMock]: - scheduler = MagicMock() +def _init_daily_global_spend_reconcile_job() -> tuple[AsyncIOScheduler, MagicMock, MagicMock]: + scheduler = AsyncIOScheduler() proxy_logging_obj = MagicMock() proxy_logging_obj.alerting_handler = AsyncMock() prisma_client = MagicMock() @@ -1058,28 +1059,31 @@ def _init_daily_global_spend_reconcile_job() -> tuple[MagicMock, MagicMock, Magi def test_daily_global_spend_reconcile_job_is_scheduled_nightly_with_an_immediate_catch_up_run(): """Startup schedules the LiteLLM_DailyGlobalSpend backfill a couple of minutes out, so a fresh deploy switches usage reads to the global table without waiting for the nightly - run, and replaces any previous registration of the same job id.""" + run, and after that it fires once a day at 00:30 UTC, when the previous UTC day is closed.""" from datetime import datetime, timedelta, timezone from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID scheduler, _, _ = _init_daily_global_spend_reconcile_job() + job = scheduler.get_job(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID) + assert job is not None - (call,) = scheduler.add_job.call_args_list - assert call.kwargs["id"] == DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID - assert call.kwargs["replace_existing"] is True - assert call.args[1:] == ("cron",) - assert (call.kwargs["hour"], call.kwargs["minute"], call.kwargs["timezone"]) == (0, 30, "UTC") - assert timedelta(0) < call.kwargs["next_run_time"] - datetime.now(timezone.utc) <= timedelta(minutes=2) + assert timedelta(0) < job.next_run_time - datetime.now(timezone.utc) <= timedelta(minutes=2) + after_catch_up = datetime(2026, 9, 16, 12, 0, tzinfo=timezone.utc) + assert job.trigger.get_next_fire_time(None, after_catch_up) == datetime(2026, 9, 17, 0, 30, tzinfo=timezone.utc) + just_after_a_run = datetime(2026, 9, 17, 0, 30, 1, tzinfo=timezone.utc) + assert job.trigger.get_next_fire_time(None, just_after_a_run) == datetime(2026, 9, 18, 0, 30, tzinfo=timezone.utc) @pytest.mark.asyncio async def test_daily_global_spend_reconcile_job_runs_under_the_pod_lock_and_alerts_through_the_proxy(monkeypatch): + from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID + scheduler, proxy_logging_obj, prisma_client = _init_daily_global_spend_reconcile_job() run = AsyncMock() monkeypatch.setattr(ps, "run_scheduled_daily_global_spend_reconcile", run) - await scheduler.add_job.call_args.args[0]() + await scheduler.get_job(DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID).func() run.assert_awaited_once() assert run.await_args.args == (prisma_client,) From f25d65940d2a013d1604c729a7dc39836df5e31f Mon Sep 17 00:00:00 2001 From: yassin Date: Wed, 16 Sep 2026 12:39:48 +0000 Subject: [PATCH 073/253] fix(proxy): take the closed-day cutoff for the global spend rollup from the database clock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../daily_global_spend_rollup.py | 43 +++++------- .../test_daily_global_spend_rollup.py | 70 +++++++++++-------- 2 files changed, 59 insertions(+), 54 deletions(-) diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index ba4a4e4e3d6..c135c7d1d9c 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -12,7 +12,7 @@ a large deployment the first backfill is minutes of work. from collections.abc import Awaitable, Callable from dataclasses import dataclass -from datetime import date, datetime, timedelta, timezone +from datetime import date, timedelta from typing import TYPE_CHECKING, Final from pydantic import BaseModel, ConfigDict, ValidationError @@ -73,7 +73,7 @@ def _reconcile_day_sql() -> str: RECONCILE_DAY_SQL: Final = _reconcile_day_sql() -_DB_NOW_SQL: Final = "SELECT (NOW() AT TIME ZONE 'UTC')::text AS now" +_DB_NOW_SQL: Final = "SELECT (NOW() AT TIME ZONE 'UTC')::text AS now, (NOW() AT TIME ZONE 'UTC')::date::text AS today" _ALL_CLOSED_DAYS_SQL: Final = 'SELECT DISTINCT "date" FROM "LiteLLM_DailyUserSpend" WHERE "date" <= $1 ORDER BY "date"' # Pod clocks drift from the database clock and from each other, so rows are picked up from a # little before the previous scan; rewriting a day twice is idempotent. @@ -111,6 +111,7 @@ class _NowRow(BaseModel): model_config = ConfigDict(frozen=True, extra="ignore") now: str + today: str @dataclass(frozen=True, slots=True) @@ -160,18 +161,18 @@ async def _record_marker(prisma_client: "PrismaClient", marker: ReconciledThroug await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) -async def _db_now(prisma_client: "PrismaClient") -> str: +async def _db_now(prisma_client: "PrismaClient") -> _NowRow: rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL) - return _NowRow.model_validate(rows[0]).now + return _NowRow.model_validate(rows[0]) -async def _scan_pending(prisma_client: "PrismaClient", today: date) -> _PendingScan: - """Every closed UTC day (strictly before today) still to roll up, oldest first: days past the - marker, plus any day with per-key rows written since the scan behind the marker. Before a - run has fully succeeded there is no such scan, so every closed day is rolled up.""" +async def _scan_pending(prisma_client: "PrismaClient") -> _PendingScan: + """Every closed UTC day (strictly before the database's today) still to roll up, oldest first: + days past the marker, plus any day with per-key rows written since the scan behind the marker. + Before a run has fully succeeded there is no such scan, so every closed day is rolled up.""" marker: Final = await read_marker(prisma_client) - scanned_at: Final = await _db_now(prisma_client) - last_closed_day: Final = (today - timedelta(days=1)).isoformat() + db_now: Final = await _db_now(prisma_client) + last_closed_day: Final = (date.fromisoformat(db_now.today) - timedelta(days=1)).isoformat() rows: Final = ( await prisma_client.db.query_raw(_ALL_CLOSED_DAYS_SQL, last_closed_day) if marker is None or marker.scanned_at is None @@ -179,11 +180,11 @@ async def _scan_pending(prisma_client: "PrismaClient", today: date) -> _PendingS _PENDING_DAYS_SQL, last_closed_day, marker.reconciled_through, marker.scanned_at ) ) - return _PendingScan(marker, scanned_at, tuple(_DateRow.model_validate(row).date for row in rows)) + return _PendingScan(marker, db_now.now, tuple(_DateRow.model_validate(row).date for row in rows)) -async def pending_days(prisma_client: "PrismaClient", today: date) -> tuple[str, ...]: - return (await _scan_pending(prisma_client, today)).days +async def pending_days(prisma_client: "PrismaClient") -> tuple[str, ...]: + return (await _scan_pending(prisma_client)).days async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None: @@ -192,15 +193,11 @@ async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None: await prisma_client.db.execute_raw(RECONCILE_DAY_SQL, day) -async def run_daily_global_spend_reconcile( - prisma_client: "PrismaClient", - today: date | None = None, -) -> ReconcileResult: +async def run_daily_global_spend_reconcile(prisma_client: "PrismaClient") -> ReconcileResult: """Roll up every pending day, advancing the marker after each; a failing day stops the run with the marker on the last good day so the next run resumes there. The scan time is only recorded once every pending day is done, so late rows a failed run saw are found again.""" - effective_today: Final = today or datetime.now(timezone.utc).date() - scan: Final = await _scan_pending(prisma_client, effective_today) + scan: Final = await _scan_pending(prisma_client) done: Final = await _reconcile_until_failure(prisma_client, scan) if len(done) < len(scan.days): marker: Final = await reconciled_through(prisma_client) @@ -243,13 +240,12 @@ async def run_scheduled_daily_global_spend_reconcile( prisma_client: "PrismaClient", pod_lock_manager: "PodLockManager | None" = None, alert: Callable[[str], Awaitable[None]] | None = None, - today: date | None = None, ) -> ReconcileResult | None: """Run the reconcile under a cross-pod lock so one proxy does the work; the lock only saves effort (each day is an idempotent rewrite), so an unreachable Redis runs unguarded rather than skipping.""" redis_cache: Final = None if pod_lock_manager is None else pod_lock_manager.redis_cache if pod_lock_manager is None or redis_cache is None: - return await _run_and_alert(prisma_client, alert=alert, today=today) + return await _run_and_alert(prisma_client, alert=alert) acquired: Final = await pod_lock_manager.acquire_lock( cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID, ttl=DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS @@ -258,7 +254,7 @@ async def run_scheduled_daily_global_spend_reconcile( verbose_proxy_logger.info("Daily global spend reconcile: another pod holds the lock, skipping this run") return None try: - return await _run_and_alert(prisma_client, alert=alert, today=today) + return await _run_and_alert(prisma_client, alert=alert) finally: if acquired: await pod_lock_manager.release_lock(cronjob_id=DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID) @@ -277,9 +273,8 @@ async def _run_and_alert( prisma_client: "PrismaClient", *, alert: Callable[[str], Awaitable[None]] | None, - today: date | None, ) -> ReconcileResult: - result: Final = await run_daily_global_spend_reconcile(prisma_client, today=today) + result: Final = await run_daily_global_spend_reconcile(prisma_client) if result.days_reconciled: verbose_proxy_logger.info( "Daily global spend reconcile: rolled up %d day(s), reconciled through %s", diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index b5b78229c11..9655953134a 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -43,7 +43,8 @@ class _FakeConfigTable: class _FakeDb: """Per-key rows are ``{date: updated_at}`` with a fake database clock that ticks per query, - so "rows written since the last scan" behaves like Postgres would.""" + so "rows written since the last scan" behaves like Postgres would. The database's own + date decides which day is still open, never the pod's clock.""" def __init__(self, prisma: "_FakePrisma") -> None: self._prisma = prisma @@ -52,7 +53,7 @@ class _FakeDb: async def query_raw(self, sql: str, *params: str) -> list[dict[str, str]]: if sql.startswith("SELECT (NOW()"): self._prisma.clock += 1 - return [{"now": f"clock-{self._prisma.clock:04d}"}] + return [{"now": f"clock-{self._prisma.clock:04d}", "today": self._prisma.today.isoformat()}] rows = self._prisma.user_rows if len(params) == 1: (last,) = params @@ -73,8 +74,11 @@ class _FakeDb: class _FakePrisma: """Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw.""" - def __init__(self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset()) -> None: + def __init__( + self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset(), today: date = TODAY + ) -> None: self.clock = 0 + self.today = today self.user_rows: dict[str, str] = {d: "clock-0000" for d in user_days} self.failing_days = failing_days self.reconciled: list[str] = [] @@ -98,12 +102,14 @@ async def _fresh_marker_cache(): @pytest.mark.asyncio -async def test_first_run_rolls_up_every_closed_day_and_never_today(): +async def test_first_run_rolls_up_every_closed_day_and_never_the_database_s_today(): """Before any marker exists every closed day with per-key rows is rolled up. Today is left - out: pods are still flushing it, so it is served live from the per-key table until it closes.""" + out: pods are still flushing it, so it is served live from the per-key table until it closes. + The database clock says which day that is; a pod booting with its clock a day ahead must not + roll the open day up and mark it reconciled.""" prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-03", "2026-09-14", "2026-09-15")) - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-01", "2026-09-03", "2026-09-14") assert result.failed_day is None @@ -114,11 +120,12 @@ async def test_first_run_rolls_up_every_closed_day_and_never_today(): @pytest.mark.asyncio async def test_later_run_rolls_up_only_new_days_when_nothing_old_changed(): - prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14")) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-12", "2026-09-13", "2026-09-14"), today=date(2026, 9, 14)) + await run_daily_global_spend_reconcile(prisma) prisma.reconciled.clear() + prisma.today = TODAY - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-14",) assert await reconciled_through(prisma) == "2026-09-14" @@ -128,13 +135,14 @@ async def test_later_run_rolls_up_only_new_days_when_nothing_old_changed(): async def test_spend_landing_on_an_old_rolled_up_day_is_folded_in_by_the_next_run(): """Per-key rows carry the request start date, so a delayed flush or retry can add spend to a day far behind the marker. That day is rewritten, and the marker never moves back for it.""" - prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-05", "2026-09-13")) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-05", "2026-09-13"), today=date(2026, 9, 14)) + await run_daily_global_spend_reconcile(prisma) prisma.reconciled.clear() + prisma.today = TODAY prisma.write_late_row("2026-09-01") prisma.write_late_row("2026-09-03") - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-01", "2026-09-03") assert "2026-09-05" not in prisma.reconciled @@ -145,15 +153,16 @@ async def test_spend_landing_on_an_old_rolled_up_day_is_folded_in_by_the_next_ru async def test_a_late_row_seen_by_a_failed_run_is_seen_again_by_the_next_one(): """The scan time only advances when every pending day was rewritten, otherwise a late row found by the failed run would be counted as handled.""" - prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13")) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13"), today=date(2026, 9, 14)) + await run_daily_global_spend_reconcile(prisma) + prisma.today = TODAY prisma.write_late_row("2026-09-01") prisma.failing_days = frozenset({"2026-09-01"}) - failed = await run_daily_global_spend_reconcile(prisma, today=TODAY) + failed = await run_daily_global_spend_reconcile(prisma) prisma.failing_days = frozenset() prisma.reconciled.clear() - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert failed.failed_day == "2026-09-01" assert failed.reconciled_through == "2026-09-13" @@ -166,7 +175,7 @@ async def test_a_marker_without_a_scan_time_rolls_every_closed_day_up_again(): prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-13")) prisma.db.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = '{"reconciled_through": "2026-09-13"}' - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-01", "2026-09-13") marker = await read_marker(prisma) @@ -175,11 +184,11 @@ async def test_a_marker_without_a_scan_time_rolls_every_closed_day_up_again(): @pytest.mark.asyncio async def test_a_run_with_no_new_closed_days_keeps_the_marker(): - prisma = _FakePrisma(user_days=("2026-09-13",)) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma = _FakePrisma(user_days=("2026-09-13",), today=date(2026, 9, 14)) + await run_daily_global_spend_reconcile(prisma) prisma.reconciled.clear() - result = await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == () assert result.reconciled_through == "2026-09-13" @@ -191,7 +200,7 @@ async def test_a_failing_day_stops_the_run_and_leaves_the_marker_on_the_last_goo a global table missing that day's spend.""" prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"})) - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-01",) assert result.failed_day == "2026-09-02" @@ -203,10 +212,10 @@ async def test_a_failing_day_stops_the_run_and_leaves_the_marker_on_the_last_goo @pytest.mark.asyncio async def test_the_next_run_resumes_from_the_failed_day(): prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-02"})) - await run_daily_global_spend_reconcile(prisma, today=TODAY) + await run_daily_global_spend_reconcile(prisma) prisma.failing_days = frozenset() - result = await run_daily_global_spend_reconcile(prisma, today=TODAY) + result = await run_daily_global_spend_reconcile(prisma) assert result.days_reconciled == ("2026-09-01", "2026-09-02", "2026-09-03") assert await reconciled_through(prisma) == "2026-09-03" @@ -215,13 +224,14 @@ async def test_the_next_run_resumes_from_the_failed_day(): @pytest.mark.asyncio async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts(): """When the rewrite of a late day fails the marker must stay put and the operator must hear about it.""" - prisma = _FakePrisma(user_days=("2026-09-13",)) - await run_daily_global_spend_reconcile(prisma, today=date(2026, 9, 14)) + prisma = _FakePrisma(user_days=("2026-09-13",), today=date(2026, 9, 14)) + await run_daily_global_spend_reconcile(prisma) + prisma.today = TODAY prisma.write_late_row("2026-09-12") prisma.failing_days = frozenset({"2026-09-12"}) alert = AsyncMock() - result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert, today=TODAY) + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert) assert result is not None assert result.days_reconciled == () @@ -236,7 +246,7 @@ async def test_a_clean_run_does_not_alert(): prisma = _FakePrisma(user_days=("2026-09-13",)) alert = AsyncMock() - await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert, today=TODAY) + await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=None, alert=alert) alert.assert_not_awaited() @@ -256,7 +266,7 @@ async def test_scheduled_run_skips_when_another_pod_holds_the_lock(): prisma = _FakePrisma(user_days=("2026-09-13",)) lock = _pod_lock(acquired=False) - result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock) assert result is None assert prisma.reconciled == [] @@ -268,7 +278,7 @@ async def test_scheduled_run_runs_and_releases_the_lock_when_it_wins(): prisma = _FakePrisma(user_days=("2026-09-13",)) lock = _pod_lock(acquired=True) - result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock) assert result is not None and result.days_reconciled == ("2026-09-13",) lock.release_lock.assert_awaited_once() @@ -282,7 +292,7 @@ async def test_scheduled_run_proceeds_when_the_lock_cannot_be_acquired_or_read() lock = _pod_lock(acquired=False) lock.redis_cache.async_get_cache = AsyncMock(side_effect=ConnectionError("redis down")) - result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock, today=TODAY) + result = await run_scheduled_daily_global_spend_reconcile(prisma, pod_lock_manager=lock) assert result is not None and result.days_reconciled == ("2026-09-13",) lock.release_lock.assert_not_awaited() From 0a81c6d3a8efadecc6498a13bbdd474374bb490f Mon Sep 17 00:00:00 2001 From: jesus Date: Wed, 16 Sep 2026 19:59:21 +0000 Subject: [PATCH 074/253] fix(proxy): resolve model_group_alias to its target for /v1/models metadata Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/utils.py | 14 +++++-- tests/test_litellm/proxy/test_proxy_utils.py | 40 ++++++++++++++++++++ 2 files changed, 50 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 215fb143f7b..f30a3da9d68 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -192,6 +192,7 @@ from litellm.repositories.user_repository import UserRepository from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) +from litellm.router_utils.common_utils import resolve_model_group_alias from litellm.secret_managers.main import str_to_bool from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES from litellm.types.llms.openai import ResponsesAPIResponse @@ -8177,18 +8178,23 @@ def create_model_info_response( "owned_by": provider, } - listing_info: Final = llm_router.get_model_listing_info(model_id) if llm_router is not None else None + alias_target: Final = ( + resolve_model_group_alias(llm_router.model_group_alias, model_id) if llm_router is not None else None + ) + lookup_model: Final = alias_target if alias_target is not None else model_id + + listing_info: Final = llm_router.get_model_listing_info(lookup_model) if llm_router is not None else None # One entry per distinct model behind the listed name; (None,) when the router knows # nothing about it, so the listed name is resolved on its own as before. deployment_models: Final[tuple[str | None, ...]] = ( listing_info.cost_map_keys if listing_info is not None and listing_info.cost_map_keys else (None,) ) - listed_info: Final = _safe_get_model_info(model_id, get_model_info) + listed_info: Final = _safe_get_model_info(lookup_model, get_model_info) candidate_sets: Final = tuple( _resolve_listing_model_info( deployment_model=deployment_model, - listed_model=model_id, + listed_model=lookup_model, listed_info=listed_info, get_model_info=get_model_info, ) @@ -8219,7 +8225,7 @@ def create_model_info_response( max_output_tokens = listing_info.max_output_tokens if llm_router is not None: - configured_mode: Final = llm_router.get_configured_mode(model_id) + configured_mode: Final = llm_router.get_configured_mode(lookup_model) if isinstance(configured_mode, str): base["mode"] = configured_mode diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 94ccc2762c5..a79e0798e29 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -2236,6 +2236,46 @@ def test_create_model_info_response_resolves_mode_through_deployment_model(): assert response["mode"] == "embedding" +@pytest.mark.parametrize( + "model_group_alias", + [ + {"team-embeddings": "my-embeddings"}, + {"team-embeddings": {"model": "my-embeddings", "hidden": False}}, + ], +) +def test_create_model_info_response_resolves_model_group_alias_to_target(model_group_alias): + """A `model_group_alias` row must report the metadata of the group it points at, + not the cost-map generalization or nothing that the alias name resolves to.""" + from litellm import Router + + saved_model_cost = dict(litellm.model_cost) + try: + router = Router( + model_list=[ + { + "model_name": "my-embeddings", + "litellm_params": {"model": "openai/text-embedding-3-small"}, + } + ], + model_group_alias=model_group_alias, + ) + + alias_response = create_model_info_response( + model_id="team-embeddings", provider="openai", llm_router=router + ) + target_response = create_model_info_response( + model_id="my-embeddings", provider="openai", llm_router=router + ) + finally: + litellm.model_cost.clear() + litellm.model_cost.update(saved_model_cost) + + assert alias_response["id"] == "team-embeddings" + for field in ("mode", "max_input_tokens", "max_output_tokens"): + assert alias_response.get(field) == target_response.get(field) + assert alias_response["mode"] == "embedding" + + @pytest.mark.parametrize( "key_metadata, team_metadata, expected_to_run", [ From 4148bf283c917423b30ba8bf6c5c79b0fc1983e5 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 23:08:03 +0000 Subject: [PATCH 075/253] fix(proxy): catch only Redis failures when falling back to local login counters Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 4ab57ca461f..59d91c12d6e 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -23,10 +23,11 @@ from typing import Final, Literal, NamedTuple, NoReturn from fastapi import Request, status from pydantic import TypeAdapter, ValidationError +from redis.exceptions import RedisError from litellm._logging import verbose_proxy_logger from litellm.caching.in_memory_cache import InMemoryCache -from litellm.caching.redis_cache import RedisCache +from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges, resolve_client_ip from litellm.secret_managers.main import get_secret_bool @@ -53,6 +54,7 @@ _MAX_TRACKED_COUNTERS: Final = 20_000 _MAX_TRACKED_BLOCKS: Final = 10_000 _NO_SETTINGS: Final[Mapping[str, object]] = MappingProxyType({}) _NOT_BLOCKED: Final = (0, 0) +_REDIS_FAILURES: Final = (RedisError, RedisCircuitBreakerOpenError, OSError, asyncio.TimeoutError) _LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None) _SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object]) @@ -304,7 +306,7 @@ class LoginThrottle: return _LUA_BLOCK_TTLS.validate_python( await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(list(keys), []) ) - except Exception as err: + except _REDIS_FAILURES as err: self._warn_redis(err) return _NOT_BLOCKED @@ -327,7 +329,7 @@ class LoginThrottle: list(keys), [self.user_limit, source_limit, self.window_seconds, self.block_seconds] ) ) - except Exception as err: + except _REDIS_FAILURES as err: self._warn_redis(err) user_block: Final = self._local_bump(keys.pair_counter, keys.pair_block, self.user_limit) if source_limit == 0 or user_block > 0: @@ -349,7 +351,7 @@ class LoginThrottle: if self.redis_cache is not None: try: await self.redis_cache.async_delete_cache(pair_counter) - except Exception as err: + except _REDIS_FAILURES as err: self._warn_redis(err) self.counters.delete_cache(pair_counter) From 438b4e6a3f34249b61a83c1d6d959daa64f56e21 Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 23:15:35 +0000 Subject: [PATCH 076/253] fix(proxy): type the login throttle's local store and pass frozen Redis script arguments Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 33 +++++++++++++++++++--------- 1 file changed, 23 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 59d91c12d6e..d17e3fd0d53 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -19,7 +19,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass from functools import cache from types import MappingProxyType -from typing import Final, Literal, NamedTuple, NoReturn +from typing import Final, Literal, NamedTuple, NoReturn, Protocol, TypeAlias from fastapi import Request, status from pydantic import TypeAdapter, ValidationError @@ -58,11 +58,24 @@ _REDIS_FAILURES: Final = (RedisError, RedisCircuitBreakerOpenError, OSError, asy _LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None) _SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object]) -Scope = Literal["user", "source"] +Scope: TypeAlias = Literal["user", "source"] -_BlockTtls = tuple[int, int] +_BlockTtls: TypeAlias = tuple[int, int] _LUA_BLOCK_TTLS: Final = TypeAdapter[_BlockTtls](_BlockTtls) -_Network = ipaddress.IPv4Network | ipaddress.IPv6Network +_Network: TypeAlias = ipaddress.IPv4Network | ipaddress.IPv6Network + + +class LocalStore(Protocol): + """The per-worker store behind the counters and blocks; ``InMemoryCache`` satisfies it.""" + + def get_cache(self, key: str) -> object: ... + + def set_cache(self, key: str, value: float, *, ttl: int) -> None: ... + + def increment_cache(self, key: str, value: float, *, ttl: int) -> float: ... + + def delete_cache(self, key: str) -> None: ... + # KEYS: pair counter, pair block, source counter, source block (one cluster slot via the source hash tag) # ARGV: pair limit, source limit (0 = source scope off), window seconds, block seconds @@ -87,7 +100,7 @@ _COUNTERS: Final = InMemoryCache( max_size_in_memory=_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS ) _BLOCKS: Final = InMemoryCache(max_size_in_memory=_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS) -_HELD_ATTEMPTS: Final[dict[str, int]] = {} +_HELD_ATTEMPTS: Final[dict[str, int]] = {} # mutable-ok: in-flight hold counts rise on entry and fall on exit async def _sleep(seconds: float) -> None: @@ -216,8 +229,8 @@ class LoginThrottle: user_limit: int window_seconds: int block_seconds: int - counters: InMemoryCache - blocks: InMemoryCache + counters: LocalStore + blocks: LocalStore redis_cache: RedisCache | None = None enabled: bool = True @@ -304,7 +317,7 @@ class LoginThrottle: return _NOT_BLOCKED try: return _LUA_BLOCK_TTLS.validate_python( - await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(list(keys), []) + await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(keys, ()) ) except _REDIS_FAILURES as err: self._warn_redis(err) @@ -326,7 +339,7 @@ class LoginThrottle: try: return _LUA_BLOCK_TTLS.validate_python( await self.redis_cache.async_register_script(_RECORD_FAILURE_LUA)( - list(keys), [self.user_limit, source_limit, self.window_seconds, self.block_seconds] + keys, (self.user_limit, source_limit, self.window_seconds, self.block_seconds) ) ) except _REDIS_FAILURES as err: @@ -369,7 +382,7 @@ class LoginThrottle: type=ProxyErrorTypes.auth_error, param="username", code=status.HTTP_429_TOO_MANY_REQUESTS, - headers={"Retry-After": str(retry_after)}, + headers={"Retry-After": str(retry_after)}, # mutable-ok: ProxyException writes into its headers dict ) From 0a8423d77b7fe99e572d8ea5e923fcd05c6985a2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 00:28:39 +0000 Subject: [PATCH 077/253] fix(proxy): keep the sign-in hold pool from refusing a correct password The held-attempt cap ran before the password check, so five parked wrong guesses from a blocked source turned the soft block into a lockout for the real user. The slot is now taken only after a wrong password, and the pool-full refusal carries the block's remaining time as Retry-After Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 49 ++++++++++--------- .../proxy/auth/test_login_utils.py | 44 ++++++++++++++++- 2 files changed, 69 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index d17e3fd0d53..a8899b9c675 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -281,25 +281,7 @@ class LoginThrottle: yield LoginAttempt(throttle=self, username=username, block=None) return slot: Final = keys.pair_block if block.scope == "user" else keys.source_block - held: Final = _HELD_ATTEMPTS.get(slot, 0) - if held >= MAX_HELD_ATTEMPTS_PER_KEY: - verbose_proxy_logger.warning( - "Admin UI sign-in refused: %s attempts already held for a blocked %s; username=%r source=%s", - held, - block.scope, - username, - self.client_ip, - ) - self.refuse(BLOCKED_ATTEMPT_HOLD_SECONDS) - _HELD_ATTEMPTS[slot] = held + 1 - try: - yield LoginAttempt(throttle=self, username=username, block=block) - finally: - remaining: Final = _HELD_ATTEMPTS.get(slot, 1) - 1 - if remaining > 0: - _HELD_ATTEMPTS[slot] = remaining - else: - _HELD_ATTEMPTS.pop(slot, None) + yield LoginAttempt(throttle=self, username=username, block=block, slot=slot) async def _active_block(self, keys: _Keys) -> Block | None: local: Final = self._local_block_ttls(keys) @@ -391,6 +373,7 @@ class LoginAttempt: throttle: LoginThrottle username: str block: Block | None + slot: str | None = None async def succeeded(self) -> None: if not self.throttle.enabled: @@ -400,9 +383,8 @@ class LoginAttempt: async def failed(self) -> None: if not self.throttle.enabled: return - if self.block is not None: - await _sleep(BLOCKED_ATTEMPT_HOLD_SECONDS) - self.throttle.refuse(max(self.block.retry_after - BLOCKED_ATTEMPT_HOLD_SECONDS, 1)) + if self.block is not None and self.slot is not None: + await self._hold_then_refuse(self.block, self.slot) user_block, source_block = await self.throttle.record_failure(self.username) if user_block == 0 and source_block == 0: return @@ -413,3 +395,26 @@ class LoginAttempt: self.username, self.throttle.client_ip, ) + + async def _hold_then_refuse(self, block: Block, slot: str) -> NoReturn: + held: Final = _HELD_ATTEMPTS.get(slot, 0) + if held >= MAX_HELD_ATTEMPTS_PER_KEY: + verbose_proxy_logger.warning( + "Admin UI sign-in refused at once: %s wrong attempts already held for a blocked %s; " + "username=%r source=%s", + held, + block.scope, + self.username, + self.throttle.client_ip, + ) + self.throttle.refuse(block.retry_after) + _HELD_ATTEMPTS[slot] = held + 1 + try: + await _sleep(BLOCKED_ATTEMPT_HOLD_SECONDS) + finally: + remaining: Final = _HELD_ATTEMPTS.get(slot, 1) - 1 + if remaining > 0: + _HELD_ATTEMPTS[slot] = remaining + else: + _HELD_ATTEMPTS.pop(slot, None) + self.throttle.refuse(max(block.retry_after - BLOCKED_ATTEMPT_HOLD_SECONDS, 1)) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index df5bd29f8ff..df35e916aa0 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1211,7 +1211,7 @@ async def test_held_attempts_from_one_blocked_key_are_capped(monkeypatch): with pytest.raises(ProxyException) as over_cap: await _guess(throttle) assert over_cap.value.code == "429" - assert over_cap.value.headers.get("Retry-After") == "30" + assert over_cap.value.headers.get("Retry-After") == "300", "refused at once, for the whole block" assert await _fail(throttle, username="someone-else@corp.com") == "401", "other keys are not affected" finally: release.set() @@ -1222,6 +1222,46 @@ async def test_held_attempts_from_one_blocked_key_are_capped(monkeypatch): assert lt._HELD_ATTEMPTS.get(slot) is None, "the slots are released once the held attempts answer" +@pytest.mark.asyncio +async def test_a_full_hold_pool_still_lets_the_right_password_in(monkeypatch): + """Five parked wrong guesses from the office must not turn the soft block into a lockout for the real user.""" + import asyncio + + from litellm.proxy.auth import login_throttle as lt + from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + release = asyncio.Event() + + async def _park(_seconds: float) -> None: + await release.wait() + + monkeypatch.setattr(lt, "_sleep", _park) + throttle = _throttle(user_limit=1, source_limit=3, client_ip="203.0.113.46") + assert [await _fail(throttle, username=f"spray-{i}@corp.com") for i in range(4)] == ["401"] * 4 + source_slot = throttle._keys("known@example.com").source_block + assert throttle._local_block_ttl(source_slot) > 0, "the source is blocked" + + held = [ + asyncio.create_task(_guess(throttle, username="known@example.com")) for _ in range(MAX_HELD_ATTEMPTS_PER_KEY) + ] + for _ in range(1000): + if lt._HELD_ATTEMPTS.get(source_slot) == MAX_HELD_ATTEMPTS_PER_KEY: + break + await asyncio.sleep(0) + assert lt._HELD_ATTEMPTS == {source_slot: MAX_HELD_ATTEMPTS_PER_KEY} + + try: + signed_in = await _db_login(throttle, "known@example.com", "right", correct=True) + assert signed_in.user_id == "u-1" + finally: + release.set() + for task in held: + with pytest.raises(ProxyException): + await task + + @pytest.mark.asyncio async def test_a_blocked_source_shares_one_held_slot_pool_across_its_blocked_usernames(monkeypatch): """Once the source is blocked, a pair block for a username must not hand that username its own five slots.""" @@ -1258,7 +1298,7 @@ async def test_a_blocked_source_shares_one_held_slot_pool_across_its_blocked_use with pytest.raises(ProxyException) as over_cap: await _guess(throttle, username=name) assert over_cap.value.code == "429" - assert over_cap.value.headers.get("Retry-After") == "30" + assert over_cap.value.headers.get("Retry-After") == "300" finally: release.set() for task in held: From 8a645bcc00dc07e9d4b14b2b735596623e8b8947 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 00:36:12 +0000 Subject: [PATCH 078/253] test(proxy): stub DATABASE_URL in the hold-pool regression test so it passes off the dev box Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/auth/test_login_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index df35e916aa0..d1fdb2d6a70 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1232,6 +1232,7 @@ async def test_a_full_hold_pool_still_lets_the_right_password_in(monkeypatch): monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") release = asyncio.Event() async def _park(_seconds: float) -> None: From f93d80ea840e9b29f17601990a8bccec67f08f1f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 16 Sep 2026 18:02:06 -0700 Subject: [PATCH 079/253] fix(mistral): forward reasoning_effort only on models that accept it --- litellm/llms/mistral/chat/transformation.py | 14 +++--- .../test_mistral_chat_transformation.py | 46 +++++++++++++++---- 2 files changed, 44 insertions(+), 16 deletions(-) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 970da0582ae..807f201a94f 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -24,7 +24,7 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.mistral import MistralThinkingBlock, MistralToolCallMessage from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse, ModelResponseStream -from litellm.utils import convert_to_model_response_object +from litellm.utils import convert_to_model_response_object, supports_reasoning if TYPE_CHECKING: import tiktoken @@ -87,7 +87,9 @@ class MistralConfig(OpenAIGPTConfig): return super().get_config() def get_supported_openai_params(self, model: str) -> list[str]: - supported_params: Final = [ + is_magistral: Final = "magistral" in model.lower() + accepts_reasoning_effort: Final = is_magistral or supports_reasoning(model=model, custom_llm_provider="mistral") + return [ "stream", "temperature", "top_p", @@ -99,14 +101,10 @@ class MistralConfig(OpenAIGPTConfig): "stop", "response_format", "parallel_tool_calls", - "reasoning_effort", + *(("thinking",) if is_magistral else ()), + *(("reasoning_effort",) if accepts_reasoning_effort else ()), ] - if "magistral" in model.lower(): - supported_params.append("thinking") - - return supported_params - def _map_tool_choice(self, tool_choice: str) -> str: if tool_choice == "auto" or tool_choice == "none": return tool_choice diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index edfaf352e1f..57a5f9ef2cd 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -51,11 +51,18 @@ class TestMistralReasoningSupport: assert "reasoning_effort" in supported_params assert "thinking" in supported_params - # Non-magistral models accept reasoning_effort (forwarded verbatim) but not thinking + # Non-magistral reasoning models accept reasoning_effort (forwarded verbatim) but not thinking + supported_params_reasoning = mistral_config.get_supported_openai_params( + "mistral/mistral-medium-latest" + ) + assert "reasoning_effort" in supported_params_reasoning + assert "thinking" not in supported_params_reasoning + + # Models Mistral rejects reasoning_effort on keep it unsupported, so drop_params still drops it supported_params_normal = mistral_config.get_supported_openai_params( "mistral/mistral-large-latest" ) - assert "reasoning_effort" in supported_params_normal + assert "reasoning_effort" not in supported_params_normal assert "thinking" not in supported_params_normal def test_map_openai_params_reasoning_effort(self): @@ -78,23 +85,46 @@ class TestMistralReasoningSupport: result_normal = mistral_config.map_openai_params( non_default_params={"reasoning_effort": "low"}, optional_params=optional_params_normal, - model="mistral/mistral-large-latest", + model="mistral/mistral-medium-latest", drop_params=False, ) assert "_add_reasoning_prompt" not in result_normal assert result_normal["reasoning_effort"] == "low" - def test_reasoning_effort_not_unsupported_for_non_magistral(self): - """Codex sends reasoning_effort to every model; Mistral must not raise UnsupportedParamsError.""" + @pytest.mark.parametrize( + ("model", "reasoning_effort"), + [("mistral-medium-latest", "high"), ("zai-glm-5-2", "xhigh")], + ) + def test_reasoning_effort_forwarded_verbatim_for_reasoning_models(self, model, reasoning_effort): + """Codex sends reasoning_effort to every model; Mistral reasoning models forward it as-is.""" import litellm optional_params = litellm.get_optional_params( - model="mistral-medium-latest", + model=model, custom_llm_provider="mistral", - reasoning_effort="medium", + reasoning_effort=reasoning_effort, ) - assert optional_params["reasoning_effort"] == "medium" + assert optional_params["reasoning_effort"] == reasoning_effort + + def test_reasoning_effort_stays_unsupported_for_non_reasoning_models(self): + """Mistral rejects reasoning_effort on codestral, so drop_params keeps dropping it there.""" + import litellm + + with pytest.raises(litellm.UnsupportedParamsError): + litellm.get_optional_params( + model="codestral-latest", + custom_llm_provider="mistral", + reasoning_effort="high", + ) + + dropped = litellm.get_optional_params( + model="codestral-latest", + custom_llm_provider="mistral", + reasoning_effort="high", + drop_params=True, + ) + assert "reasoning_effort" not in dropped def test_client_metadata_stripped_from_request(self): """client_metadata passed by Codex must not reach Mistral, whose schema rejects unknown fields.""" From b77f866dbb00b41d322d1f0dba80d49e040aa346 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:13:13 +0000 Subject: [PATCH 080/253] refactor(mistral): drop client_metadata without mutating optional_params Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/mistral/chat/transformation.py | 4 ++-- .../llms/mistral/test_mistral_chat_transformation.py | 6 ------ 2 files changed, 2 insertions(+), 8 deletions(-) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 807f201a94f..6316128e6fa 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -531,13 +531,13 @@ class MistralConfig(OpenAIGPTConfig): if "magistral" in model.lower() and optional_params.get("_add_reasoning_prompt", False): messages = self._add_reasoning_system_prompt_if_needed(messages, optional_params) - optional_params.pop("client_metadata", None) + upstream_params: Final = {key: value for key, value in optional_params.items() if key != "client_metadata"} # Call parent transform_request which handles _transform_messages return super().transform_request( model=model, messages=messages, - optional_params=optional_params, + optional_params=upstream_params, litellm_params=litellm_params, headers=headers, ) diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index 57a5f9ef2cd..38639f23050 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -51,14 +51,12 @@ class TestMistralReasoningSupport: assert "reasoning_effort" in supported_params assert "thinking" in supported_params - # Non-magistral reasoning models accept reasoning_effort (forwarded verbatim) but not thinking supported_params_reasoning = mistral_config.get_supported_openai_params( "mistral/mistral-medium-latest" ) assert "reasoning_effort" in supported_params_reasoning assert "thinking" not in supported_params_reasoning - # Models Mistral rejects reasoning_effort on keep it unsupported, so drop_params still drops it supported_params_normal = mistral_config.get_supported_openai_params( "mistral/mistral-large-latest" ) @@ -80,7 +78,6 @@ class TestMistralReasoningSupport: assert result.get("_add_reasoning_prompt") is True - # Test reasoning_effort forwarded verbatim for non-magistral model optional_params_normal = {} result_normal = mistral_config.map_openai_params( non_default_params={"reasoning_effort": "low"}, @@ -97,7 +94,6 @@ class TestMistralReasoningSupport: [("mistral-medium-latest", "high"), ("zai-glm-5-2", "xhigh")], ) def test_reasoning_effort_forwarded_verbatim_for_reasoning_models(self, model, reasoning_effort): - """Codex sends reasoning_effort to every model; Mistral reasoning models forward it as-is.""" import litellm optional_params = litellm.get_optional_params( @@ -108,7 +104,6 @@ class TestMistralReasoningSupport: assert optional_params["reasoning_effort"] == reasoning_effort def test_reasoning_effort_stays_unsupported_for_non_reasoning_models(self): - """Mistral rejects reasoning_effort on codestral, so drop_params keeps dropping it there.""" import litellm with pytest.raises(litellm.UnsupportedParamsError): @@ -127,7 +122,6 @@ class TestMistralReasoningSupport: assert "reasoning_effort" not in dropped def test_client_metadata_stripped_from_request(self): - """client_metadata passed by Codex must not reach Mistral, whose schema rejects unknown fields.""" mistral_config = MistralConfig() request = mistral_config.transform_request( From ed18edbbdd39ce326ed0448250b3ea767c94c4d1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:56:05 +0000 Subject: [PATCH 081/253] chore(proxy): drop a comment that restated the NUM_WORKERS assignment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_cli.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 1ade0652855..9f2e4c9802e 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -1411,8 +1411,6 @@ def run_server( # DO NOT DELETE - enables global variables to work across files from litellm.proxy.proxy_server import app - # Write the resolved --num_workers back to its env var so worker processes can read - # the fleet size at startup (the failed-login accounting warning keys off it) os.environ["NUM_WORKERS"] = str(num_workers) # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups From 2e8dc0a627b552d408910d84628692eca5452be3 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 03:03:19 +0000 Subject: [PATCH 082/253] feat(proxy): add Azure AI Speech pass-through route Adds /azure_speech/{endpoint:path}, an authenticated pass-through for the Azure AI Speech REST APIs: short-audio recognition on .stt.speech.microsoft.com and batch transcription on .api.cognitive.microsoft.com. The proxy resolves the subscription key through PassthroughEndpointRouter (AZURE_SPEECH_API_KEY or an Admin UI credential), picks the host from AZURE_SPEECH_REGION or AZURE_SPEECH_API_BASE, injects Ocp-Apim-Subscription-Key, strips the caller's Authorization and subscription-key headers, forwards the raw audio body byte for byte, and records a zero-cost SpendLogs row tagged azure_speech since the price map has no Azure Speech STT entry Resolves LIT-7939 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- gateway/routes/allowlist.py | 1 + helm/litellm/templates/ingress.yaml | 2 +- litellm/constants.py | 10 + litellm/passthrough/utils.py | 1 + litellm/proxy/_lazy_features.py | 1 + litellm/proxy/_lazy_openapi_snapshot.json | 222 +++++++++++++ litellm/proxy/_types.py | 1 + litellm/proxy/auth/user_api_key_auth.py | 7 + .../proxy/common_utils/http_parsing_utils.py | 13 +- .../llm_passthrough_endpoints.py | 118 +++++++ ...zure_speech_passthrough_logging_handler.py | 84 +++++ .../pass_through_endpoints.py | 3 +- .../pass_through_endpoints/success_handler.py | 24 +- .../provider_create_fields.json | 18 ++ litellm/types/utils.py | 1 + terraform/litellm/aws/locals.tf | 2 +- terraform/litellm/gcp/locals.tf | 2 +- ...est_billable_request_metrics_middleware.py | 5 + ...zure_speech_passthrough_logging_handler.py | 123 ++++++++ .../test_llm_pass_through_endpoints.py | 294 ++++++++++++++++++ .../test_passthrough_endpoint_router.py | 16 + .../src/components/provider_info_helpers.tsx | 4 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 231 ++++++++++++++ 23 files changed, 1177 insertions(+), 6 deletions(-) create mode 100644 litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 099c6d5179f..5b8c44809fe 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -82,6 +82,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/anthropic/", "/azure/", "/azure_ai/", + "/azure_speech/", "/aws/", "/bedrock/", "/comprehendmedical", diff --git a/helm/litellm/templates/ingress.yaml b/helm/litellm/templates/ingress.yaml index d42558b9396..94ac4d2d8b0 100644 --- a/helm/litellm/templates/ingress.yaml +++ b/helm/litellm/templates/ingress.yaml @@ -66,7 +66,7 @@ "/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search" "/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat" "/v1beta" "/interactions" - "/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/comprehendmedical" "/cohere" "/gemini" "/google" + "/anthropic" "/azure" "/azure_ai" "/azure_speech" "/aws" "/bedrock" "/comprehendmedical" "/cohere" "/gemini" "/google" "/vertex_ai" "/vertex-ai" "/assemblyai" "/eu.assemblyai" "/langfuse" "/vllm" "/mistral" "/groq" "/voyage" "/cursor" "/milvus" "/openai_passthrough" "/toolset" diff --git a/litellm/constants.py b/litellm/constants.py index 8409a161800..69de889a326 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1570,6 +1570,16 @@ ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS: Final = { # Works for all LLM pass-through endpoints (Vertex AI, Anthropic, Bedrock, etc.) PASS_THROUGH_HEADER_PREFIX: Final = "x-pass-" +AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech" +AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech" +AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX: Final = "/speech/" +AZURE_SPEECH_BATCH_PATH_PREFIX: Final = "/speechtotext/" +AZURE_SPEECH_STT_DOMAIN: Final = "stt.speech.microsoft.com" +AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN: Final = "api.cognitive.microsoft.com" +AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER: Final = "Ocp-Apim-Subscription-Key" +AZURE_SPEECH_SHORT_AUDIO_MODEL: Final = "short-audio" +AZURE_SPEECH_BATCH_MODEL: Final = "batch-transcription" + BASE_MCP_ROUTE: Final = "/mcp" BATCH_STATUS_POLL_INTERVAL_SECONDS: Final = int(os.getenv("BATCH_STATUS_POLL_INTERVAL_SECONDS", 3600)) # 1 hour diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index 7eb14fcc118..452c9c7de9d 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -17,6 +17,7 @@ _PASS_THROUGH_PROTECTED_HEADERS: Final[frozenset] = frozenset( "api-key", "x-api-key", "x-goog-api-key", + "ocp-apim-subscription-key", "host", "content-length", "accept-encoding", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index faf95397fa5..92cc014967d 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -196,6 +196,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/assemblyai/", "/azure/", "/azure_ai/", + "/azure_speech/", "/bedrock/", "/cohere/", "/comprehendmedical", diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 74f38b3ca6d..dbfbc317d24 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -17133,6 +17133,228 @@ ] } }, + "/azure_speech/{endpoint}": { + "delete": { + "description": "Pass-through for the Azure AI Speech REST APIs (speech to text), e.g.\n`POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US`\nwith the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`.\n\nThe body is forwarded byte for byte and the proxy injects its own\n`Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key\nand is never forwarded.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/azure_speech)", + "operationId": "azure_speech_proxy_route_azure_speech__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Speech Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Pass-through for the Azure AI Speech REST APIs (speech to text), e.g.\n`POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US`\nwith the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`.\n\nThe body is forwarded byte for byte and the proxy injects its own\n`Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key\nand is never forwarded.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/azure_speech)", + "operationId": "azure_speech_proxy_route_azure_speech__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Speech Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Pass-through for the Azure AI Speech REST APIs (speech to text), e.g.\n`POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US`\nwith the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`.\n\nThe body is forwarded byte for byte and the proxy injects its own\n`Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key\nand is never forwarded.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/azure_speech)", + "operationId": "azure_speech_proxy_route_azure_speech__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Speech Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Pass-through for the Azure AI Speech REST APIs (speech to text), e.g.\n`POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US`\nwith the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`.\n\nThe body is forwarded byte for byte and the proxy injects its own\n`Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key\nand is never forwarded.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/azure_speech)", + "operationId": "azure_speech_proxy_route_azure_speech__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Speech Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Pass-through for the Azure AI Speech REST APIs (speech to text), e.g.\n`POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US`\nwith the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`.\n\nThe body is forwarded byte for byte and the proxy injects its own\n`Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key\nand is never forwarded.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/azure_speech)", + "operationId": "azure_speech_proxy_route_azure_speech__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Speech Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, "/bedrock/{endpoint}": { "delete": { "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 680c63393e8..489186a8f69 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -468,6 +468,7 @@ class LiteLLMRoutes(enum.Enum): mapped_pass_through_routes = [ "/bedrock", "/comprehendmedical", + "/azure_speech", "/vertex-ai", "/vertex_ai", "/cohere", diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 4cbd4213463..5f2e0a3c1a8 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -105,6 +105,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, _safe_get_request_query_params, _safe_set_request_parsed_body, + is_opaque_audio_pass_through_request, populate_request_with_path_params, read_raw_json_body, rewrite_request_model, @@ -1354,6 +1355,12 @@ async def _read_request_body_deferring_parse_failure( must run (resolving identity onto the request's trace) before the 400 goes out; the caller re-raises the returned exception once identity is seeded. """ + if is_opaque_audio_pass_through_request( + route=get_request_route(request=request), + content_type=_safe_get_request_headers(request=request).get("content-type", ""), + ): + _safe_set_request_parsed_body(request=request, parsed_body={}) # mutable-ok: the body cache stores a plain dict + return {}, None # mutable-ok: request_data is a plain dict across the whole auth path try: parsed_body: Final = await _read_request_body(request=request) except ProxyException as parse_exception: diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index f5b6a0a766d..29dc36f3dba 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -9,7 +9,12 @@ from fastapi import Request, UploadFile, status from typing_extensions import NotRequired, ReadOnly, Required from litellm._logging import verbose_proxy_logger -from litellm.constants import CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB +from litellm.constants import ( + AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX, + AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, + CLIENT_REQUESTED_MODEL_SCOPE_KEY, + MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB, +) from litellm.proxy._types import ProxyException from litellm.proxy.common_utils.callback_utils import ( get_metadata_variable_name_from_kwargs, @@ -214,6 +219,12 @@ async def _read_request_body(request: Request | None) -> dict: return {} +def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: + return route.startswith( + f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}{AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX}" + ) and _normalize_media_type(content_type).startswith("audio/") + + async def read_raw_json_body(request: Request | None) -> bytes | None: if request is None or _safe_get_request_parsed_body(request=request) is None: return None diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index b9b8cb3a22b..e64eac87a7f 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -30,6 +30,13 @@ from litellm import get_llm_provider from litellm._logging import verbose_proxy_logger from litellm.constants import ( ALLOWED_VERTEX_AI_PASSTHROUGH_HEADERS, + AZURE_SPEECH_BATCH_PATH_PREFIX, + AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN, + AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX, + AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, + AZURE_SPEECH_STT_DOMAIN, + AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix @@ -1316,6 +1323,117 @@ async def comprehend_medical_sdk_proxy_route( ) +AZURE_SPEECH_FORWARDED_REQUEST_HEADERS: Final = ("content-type", "accept") +AZURE_SPEECH_ENDPOINT_FAMILY_DOMAINS: Final = MappingProxyType( + { + AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX: AZURE_SPEECH_STT_DOMAIN, + AZURE_SPEECH_BATCH_PATH_PREFIX: AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN, + } +) + + +def resolve_azure_speech_base_url(endpoint_path: str, api_base: str | None, region: str | None) -> httpx.URL | None: + """ + Azure AI Speech serves the two REST families from different regional hosts: short-audio + recognition under ``{region}.stt.speech.microsoft.com`` and batch transcription under + ``{region}.api.cognitive.microsoft.com``. An operator-configured ``api_base`` (custom + domain or private endpoint) serves both and wins over the region. Returns ``None`` when + the path is outside both families so the operator key is never sent for an unknown API. + """ + domain: Final = next( + ( + family_domain + for family_prefix, family_domain in AZURE_SPEECH_ENDPOINT_FAMILY_DOMAINS.items() + if endpoint_path.startswith(family_prefix) + ), + None, + ) + if domain is None: + return None + if api_base: + return httpx.URL(api_base) + if not region: + return None + return httpx.URL(f"https://{region}.{domain}") + + +@router.api_route( + f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}/{{endpoint:path}}", + methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: fastapi route methods must be a list + tags=["Azure AI Speech Pass-through", "pass-through"], # mutable-ok: fastapi route tags must be a list +) +async def azure_speech_proxy_route( + endpoint: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """ + Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + + The body is forwarded byte for byte and the proxy injects its own + `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + and is never forwarded. + + [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + """ + endpoint_path: Final = httpx.URL(endpoint).path + normalized_endpoint_path: Final = endpoint_path if endpoint_path.startswith("/") else f"/{endpoint_path}" + base_url: Final = resolve_azure_speech_base_url( + endpoint_path=normalized_endpoint_path, + api_base=get_secret_str(secret_name="AZURE_SPEECH_API_BASE"), + region=get_secret_str(secret_name="AZURE_SPEECH_REGION"), + ) + if base_url is None: + raise HTTPException( + status_code=400, + detail=( + f"Unsupported Azure Speech path: {normalized_endpoint_path}. Supported prefixes are " + f"{AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX} and {AZURE_SPEECH_BATCH_PATH_PREFIX}; set " + "AZURE_SPEECH_REGION or AZURE_SPEECH_API_BASE in the proxy environment." + ), + ) + azure_speech_api_key: Final = passthrough_endpoint_router.get_credentials( + custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + region_name=None, + ) + if azure_speech_api_key is None: + raise HTTPException( + status_code=400, + detail="Azure Speech credentials not found. Set AZURE_SPEECH_API_KEY in the proxy environment.", + ) + + target_url: Final = base_url.copy_with( + path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, normalized_endpoint_path) + ) + request_headers: Final = _safe_get_request_headers(request) + upstream_headers: Final = MappingProxyType( + { + header_name: header_value + for header_name, header_value in ( + *( + (header_name, request_headers[header_name]) + for header_name in AZURE_SPEECH_FORWARDED_REQUEST_HEADERS + if header_name in request_headers + ), + (AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, azure_speech_api_key), + ) + } + ) + raw_body: Final = await request.body() + + endpoint_func: Final = create_pass_through_route( + endpoint=endpoint, + target=str(target_url), + custom_headers=upstream_headers, + custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + ) + setattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_body) + return await endpoint_func(request, fastapi_response, user_api_key_dict) + + def _resolve_vertex_model_from_router( model_id: str, llm_router: litellm.Router | None, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py new file mode 100644 index 00000000000..a7084a9545e --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py @@ -0,0 +1,84 @@ +from collections.abc import Mapping +from datetime import datetime +from typing import Final +from urllib.parse import urlparse + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.constants import ( + AZURE_SPEECH_BATCH_MODEL, + AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + AZURE_SPEECH_SHORT_AUDIO_MODEL, + AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, +) +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.litellm_logging import ( + get_standard_logging_object_payload, +) +from litellm.proxy._types import PassThroughEndpointLoggingTypedDict +from litellm.types.utils import StandardPassThroughResponseObject + + +class AzureSpeechPassthroughLoggingHandler: + @staticmethod + def _model_from_url_route(url_route: str) -> str: + path: Final = urlparse(url_route).path + if path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX): + return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_SHORT_AUDIO_MODEL}" + return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_BATCH_MODEL}" + + @staticmethod + def azure_speech_passthrough_handler( + httpx_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + url_route: str, + result: str, + start_time: datetime, + end_time: datetime, + cache_hit: bool, + request_body: Mapping[str, object], + **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler + ) -> PassThroughEndpointLoggingTypedDict: + """ + Records model and provider for an Azure AI Speech REST call. Azure bills per audio + hour after the fact and neither the short-audio response nor the batch job carries + a billable duration this path can trust, so response_cost is recorded as 0.0 rather + than estimated. + """ + try: + model_name: Final = AzureSpeechPassthroughLoggingHandler._model_from_url_route(url_route) + + updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict + **kwargs, + "model": model_name, + "custom_llm_provider": AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + "response_cost": 0.0, + } + logging_obj.model_call_details.update( + model=model_name, + custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + response_cost=0.0, + ) + + standard_logging_object: Final = get_standard_logging_object_payload( + kwargs=updated_kwargs, + init_response_obj=StandardPassThroughResponseObject(response=result), + start_time=start_time, + end_time=end_time, + logging_obj=logging_obj, + status="success", + ) + + handler_payload: Final[PassThroughEndpointLoggingTypedDict] = { + "result": StandardPassThroughResponseObject(response=result), + "kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object}, + } + except Exception as e: # noqa: BLE001 # logging must never fail the forwarded request + verbose_proxy_logger.exception("Error in Azure Speech passthrough logging handler: %s", e) + fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = { + "result": StandardPassThroughResponseObject(response=result), + "kwargs": kwargs, + } + return fallback_payload + return handler_payload diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 685c19062bb..81268a7cf6e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -52,6 +52,7 @@ from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.initialize_dynamic_callback_params import validate_no_callback_env_reference from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.litellm_logging import _get_masked_values from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.base_llm.managed_resources.utils import ( @@ -1023,7 +1024,7 @@ async def pass_through_request( verbose_proxy_logger.debug( "Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n", url, - upstream_headers, + _get_masked_values(upstream_headers), _parsed_body, ) diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 76a471302f4..919de5c1088 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -5,6 +5,7 @@ from urllib.parse import urlparse import httpx +from litellm.constants import AZURE_SPEECH_CUSTOM_LLM_PROVIDER from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import PassThroughEndpointLoggingResultValues from litellm.types.passthrough_endpoints.pass_through_endpoints import ( @@ -256,6 +257,24 @@ class PassThroughEndpointLogging: ) standard_logging_response_object = comprehend_medical_handler_result["result"] # rebind-ok: elif-chain kwargs = comprehend_medical_handler_result["kwargs"] # rebind-ok: elif-chain contract + elif self.is_azure_speech_route(custom_llm_provider): + from .llm_provider_handlers.azure_speech_passthrough_logging_handler import ( + AzureSpeechPassthroughLoggingHandler, + ) + + azure_speech_handler_result: Final = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=httpx_response, + logging_obj=logging_obj, + url_route=url_route, + result=result, + start_time=start_time, + end_time=end_time, + cache_hit=cache_hit, + request_body=request_body, + **kwargs, + ) + standard_logging_response_object = azure_speech_handler_result["result"] # rebind-ok: elif-chain + kwargs = azure_speech_handler_result["kwargs"] # rebind-ok: elif-chain contract elif self.is_vertex_ai_live_route(url_route): from .llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import ( VertexAILivePassthroughLoggingHandler, @@ -300,7 +319,7 @@ class PassThroughEndpointLogging: ): standard_logging_response_object: PassThroughEndpointLoggingResultValues | None = None logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload - if self.is_assemblyai_route(url_route): + if self.is_assemblyai_route(url_route) and not self.is_azure_speech_route(custom_llm_provider): if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True: return self.assemblyai_passthrough_logging_handler.assemblyai_passthrough_logging_handler( @@ -389,6 +408,9 @@ class PassThroughEndpointLogging: def is_comprehend_medical_route(self, custom_llm_provider: str | None) -> bool: return custom_llm_provider == "comprehendmedical" + def is_azure_speech_route(self, custom_llm_provider: str | None) -> bool: + return custom_llm_provider == AZURE_SPEECH_CUSTOM_LLM_PROVIDER + def is_langfuse_route(self, url_route: str): parsed_url: Final = urlparse(url_route) for route in self.TRACKED_LANGFUSE_ROUTES: diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index cd781abee26..e6673ec99aa 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -586,6 +586,24 @@ ], "default_model_placeholder": "azure_ai/command-r-plus" }, + { + "provider": "Azure_Speech", + "provider_display_name": "Azure AI Speech", + "litellm_provider": "azure_speech", + "credential_fields": [ + { + "key": "api_key", + "label": "Azure AI Speech Subscription Key", + "placeholder": null, + "tooltip": "The Ocp-Apim-Subscription-Key for your Azure AI Speech resource. The proxy injects it on every /azure_speech/* pass-through request. Region and API base come from AZURE_SPEECH_REGION / AZURE_SPEECH_API_BASE", + "required": true, + "field_type": "password", + "options": null, + "default_value": null + } + ], + "default_model_placeholder": "azure_speech/short-audio" + }, { "provider": "AZURE_TEXT", "provider_display_name": "Azure Text", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index aaa16fd2d44..3c2c549e89e 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -4060,6 +4060,7 @@ class LlmProviders(str, Enum): TOPAZ = "topaz" SAP_GENERATIVE_AI_HUB = "sap" ASSEMBLYAI = "assemblyai" + AZURE_SPEECH = "azure_speech" CHARITY_ENGINE = "charity_engine" GITHUB_COPILOT = "github_copilot" SNOWFLAKE = "snowflake" diff --git a/terraform/litellm/aws/locals.tf b/terraform/litellm/aws/locals.tf index bd5b97b0f50..fcf1f7b905f 100644 --- a/terraform/litellm/aws/locals.tf +++ b/terraform/litellm/aws/locals.tf @@ -86,7 +86,7 @@ locals { "/queue/chat/*", "/v1beta/*", "/interactions/*", - "/anthropic/*", "/azure/*", "/azure_ai/*", "/aws/*", "/bedrock/*", "/comprehendmedical*", + "/anthropic/*", "/azure/*", "/azure_ai/*", "/azure_speech/*", "/aws/*", "/bedrock/*", "/comprehendmedical*", "/cohere/*", "/gemini/*", "/google/*", "/vertex_ai/*", "/vertex-ai/*", "/assemblyai/*", "/eu.assemblyai/*", diff --git a/terraform/litellm/gcp/locals.tf b/terraform/litellm/gcp/locals.tf index 3861413d496..dca7b05f1c9 100644 --- a/terraform/litellm/gcp/locals.tf +++ b/terraform/litellm/gcp/locals.tf @@ -55,7 +55,7 @@ locals { "/queue/chat/*", "/v1beta/*", "/interactions/*", - "/anthropic/*", "/azure/*", "/azure_ai/*", "/aws/*", "/bedrock/*", "/comprehendmedical*", + "/anthropic/*", "/azure/*", "/azure_ai/*", "/azure_speech/*", "/aws/*", "/bedrock/*", "/comprehendmedical*", "/cohere/*", "/gemini/*", "/google/*", "/vertex_ai/*", "/vertex-ai/*", "/assemblyai/*", "/eu.assemblyai/*", diff --git a/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py b/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py index 9c61412bd6e..5ce7aa858c1 100644 --- a/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py +++ b/tests/test_litellm/proxy/middleware/test_billable_request_metrics_middleware.py @@ -116,6 +116,11 @@ def test_is_pure_asgi_not_base_http_middleware(): # Bare AWS-SDK-shaped route carries the operation in X-Amz-Target and writes SpendLogs ("/comprehendmedical", (BillableCategory.LLM, "/comprehendmedical")), ("/comprehendmedical/DetectEntitiesV2", (BillableCategory.LLM, "/comprehendmedical")), + ( + "/azure_speech/speech/recognition/conversation/cognitiveservices/v1", + (BillableCategory.LLM, "/azure_speech"), + ), + ("/azure_speech/speechtotext/v3.2/transcriptions", (BillableCategory.LLM, "/azure_speech")), ("/mcp", (BillableCategory.MCP, "/mcp")), ("/mcp/", (BillableCategory.MCP, "/mcp")), ("/mcp/tools/list", (BillableCategory.MCP, "/mcp")), diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py new file mode 100644 index 00000000000..91f7bf94281 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py @@ -0,0 +1,123 @@ +from datetime import datetime +from unittest.mock import MagicMock + +import httpx +import pytest + +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.azure_speech_passthrough_logging_handler import ( + AzureSpeechPassthroughLoggingHandler, +) +from litellm.proxy.pass_through_endpoints.success_handler import ( + PassThroughEndpointLogging, +) + +SHORT_AUDIO_URL = "https://eastus.stt.speech.microsoft.com/speech/recognition/conversation/cognitiveservices/v1?language=en-US" +BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions" +TRANSCRIPT = '{"RecognitionStatus":"Success","DisplayText":"Hello world."}' + + +def _make_response(url: str) -> httpx.Response: + request = httpx.Request("POST", url, headers={"Ocp-Apim-Subscription-Key": "server-secret"}) + return httpx.Response(200, request=request, text=TRANSCRIPT) + + +def _make_logging_obj() -> MagicMock: + logging_obj = MagicMock() + logging_obj.litellm_call_id = "test-call-id" + logging_obj.model_call_details = {} + return logging_obj + + +class TestAzureSpeechPassthroughHandler: + @pytest.mark.parametrize( + "url_route,expected_model", + [ + (SHORT_AUDIO_URL, "azure_speech/short-audio"), + (BATCH_URL, "azure_speech/batch-transcription"), + (f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription"), + ], + ) + def test_records_model_provider_and_zero_cost(self, url_route: str, expected_model: str): + logging_obj = _make_logging_obj() + + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(url_route), + logging_obj=logging_obj, + url_route=url_route, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["result"] == {"response": TRANSCRIPT} + assert handler_result["kwargs"]["model"] == expected_model + assert handler_result["kwargs"]["custom_llm_provider"] == "azure_speech" + assert handler_result["kwargs"]["response_cost"] == 0.0 + assert handler_result["kwargs"]["standard_logging_object"]["response_cost"] == 0.0 + assert handler_result["kwargs"]["standard_logging_object"]["model"] == expected_model + assert logging_obj.model_call_details["model"] == expected_model + assert logging_obj.model_call_details["custom_llm_provider"] == "azure_speech" + assert logging_obj.model_call_details["response_cost"] == 0.0 + + def test_subscription_key_never_reaches_the_logging_payload(self): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(SHORT_AUDIO_URL), + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert "server-secret" not in repr(handler_result) + + +class TestIsAzureSpeechRoute: + def test_matches_by_provider_tag(self): + assert PassThroughEndpointLogging().is_azure_speech_route("azure_speech") + + @pytest.mark.parametrize("provider", ["azure", "azure_ai", "comprehendmedical", None]) + def test_does_not_match_other_providers(self, provider: str | None): + assert not PassThroughEndpointLogging().is_azure_speech_route(provider) + + def test_config_driven_passthrough_to_azure_speech_host_is_not_claimed(self): + normalized = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=_make_response(SHORT_AUDIO_URL), + response_body={"RecognitionStatus": "Success"}, + request_body={}, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + custom_llm_provider=None, + ) + + assert normalized["kwargs"].get("model") != "azure_speech/short-audio" + assert "response_cost" not in normalized["kwargs"] + + +class TestNormalizeDispatch: + def test_normalize_routes_to_azure_speech_handler(self): + normalized = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( + httpx_response=_make_response(SHORT_AUDIO_URL), + response_body={"RecognitionStatus": "Success"}, + request_body={}, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + custom_llm_provider="azure_speech", + ) + + assert normalized["standard_logging_response_object"] == {"response": TRANSCRIPT} + assert normalized["kwargs"]["model"] == "azure_speech/short-audio" + assert normalized["kwargs"]["custom_llm_provider"] == "azure_speech" + assert normalized["kwargs"]["response_cost"] == 0.0 diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 6e82c90514d..9fe7f5b6ee9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -2,6 +2,7 @@ import asyncio import base64 import contextlib import json +import logging import os import traceback from collections.abc import Iterator, Mapping @@ -6136,3 +6137,296 @@ class TestAzureRelayDeploymentSegment: ) assert [call["model"] for call in captured] == ["gpt", "gpt"] + + +AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1" +AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions" +AZURE_SPEECH_PCM16_HEADER: Final = ( + b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" +) +AZURE_SPEECH_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + b"\x00" * 3072 +AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range(256)) * 12 +AZURE_SPEECH_TRANSCRIPT: Final = {"RecognitionStatus": "Success", "DisplayText": "The eagle has landed."} + + +@pytest.fixture +def azure_speech_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") + monkeypatch.setenv("AZURE_SPEECH_REGION", "eastus") + monkeypatch.delenv("AZURE_SPEECH_API_BASE", raising=False) + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) + yield TestClient(app) + + +class TestAzureSpeechProxyRoute: + """Drives the real FastAPI route with respx standing in for the Azure hosts only.""" + + def test_short_audio_forwards_raw_wav_bytes_with_server_key(self, azure_speech_client: TestClient) -> None: + with respx.mock(assert_all_called=True) as upstream: + route = upstream.post( + f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" + ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + params={"language": "en-US", "format": "detailed"}, + content=AZURE_SPEECH_WAV_BYTES, + headers={ + "Content-Type": "audio/wav; codecs=audio/pcm; samplerate=16000", + "Authorization": "Bearer sk-virtual", + "Ocp-Apim-Subscription-Key": "caller-supplied-key", + "x-pass-ocp-apim-subscription-key": "caller-supplied-key", + }, + ) + + assert (response.status_code, response.json()) == (200, AZURE_SPEECH_TRANSCRIPT) + sent = route.calls.last.request + assert sent.content == AZURE_SPEECH_WAV_BYTES + assert dict(sent.url.params) == {"language": "en-US", "format": "detailed"} + assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" + assert sent.headers["content-type"] == "audio/wav; codecs=audio/pcm; samplerate=16000" + assert "authorization" not in sent.headers + assert "caller-supplied-key" not in repr(sent.headers) + + def test_batch_json_goes_to_the_cognitive_services_host(self, azure_speech_client: TestClient) -> None: + body: Final = {"contentUrls": ["https://example.com/a.wav"], "locale": "en-US", "displayName": "job"} + with respx.mock(assert_all_called=True) as upstream: + route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( + return_value=httpx.Response(201, json={"self": "https://eastus.api.cognitive.microsoft.com/x"}) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", + json=body, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 201 + sent = route.calls.last.request + assert json.loads(sent.content) == body + assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" + assert "authorization" not in sent.headers + + def test_batch_multipart_upload_is_forwarded_byte_for_byte(self, azure_speech_client: TestClient) -> None: + with respx.mock(assert_all_called=True) as upstream: + route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( + return_value=httpx.Response(201, json={"status": "NotStarted"}) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", + files={"audio": ("eagle.wav", AZURE_SPEECH_NON_UTF8_WAV_BYTES, "audio/wav")}, + data={"definition": json.dumps({"locales": ["en-US"]})}, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 201 + sent = route.calls.last.request + assert sent.headers["content-type"].startswith("multipart/form-data; boundary=") + assert AZURE_SPEECH_NON_UTF8_WAV_BYTES in sent.content + assert b'name="definition"' in sent.content + assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" + assert "authorization" not in sent.headers + + def test_batch_get_is_forwarded_with_the_job_id_path(self, azure_speech_client: TestClient) -> None: + job_path: Final = f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files" + with respx.mock(assert_all_called=True) as upstream: + route = upstream.get(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( + return_value=httpx.Response(200, json={"values": []}) + ) + + response = azure_speech_client.get(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}) + + assert (response.status_code, response.json()) == (200, {"values": []}) + assert route.calls.last.request.headers["ocp-apim-subscription-key"] == "server-subscription-key" + + @pytest.mark.parametrize("method", ["GET", "POST"]) + def test_batch_requests_are_logged_as_azure_speech_not_assemblyai( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch, method: str + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[dict[str, object]] = [] # mutable-ok: test recorder accumulates callback payloads + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + self.payloads.append(kwargs["standard_logging_object"]) + + recorder: Final = _Recorder() + monkeypatch.setattr(litellm, "_async_success_callback", [*litellm._async_success_callback, recorder]) + with respx.mock(assert_all_called=True) as upstream: + upstream.request(method, f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"values": []}) + ) + + response = azure_speech_client.request( + method, + f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", + json={"locale": "en-US"} if method == "POST" else None, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 200 + assert [(p["model"], p["custom_llm_provider"], p["response_cost"]) for p in recorder.payloads] == [ + ("azure_speech/batch-transcription", "azure_speech", 0.0) + ] + + def test_api_base_wins_over_region_for_both_families( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.setenv("AZURE_SPEECH_API_BASE", "https://my-speech.cognitiveservices.azure.com") + with respx.mock(assert_all_called=True) as upstream: + short_audio = upstream.post( + f"https://my-speech.cognitiveservices.azure.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" + ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + batch = upstream.get(f"https://my-speech.cognitiveservices.azure.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"values": []}) + ) + + azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + content=AZURE_SPEECH_WAV_BYTES, + headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, + ) + azure_speech_client.get(f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", headers={"Authorization": "Bearer x"}) + + assert short_audio.called and batch.called + + @pytest.mark.parametrize("endpoint", ["openai/deployments/whisper/audio/transcriptions", "speech", "speechtotext"]) + def test_unknown_path_family_is_rejected_before_any_upstream_call( + self, azure_speech_client: TestClient, endpoint: str + ) -> None: + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(200)) + + response = azure_speech_client.post( + f"/azure_speech/{endpoint}", content=b"x", headers={"Authorization": "Bearer sk-virtual"} + ) + + assert response.status_code == 400 + assert not catch_all.called + + def test_missing_region_and_base_is_rejected_before_any_upstream_call( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("AZURE_SPEECH_REGION") + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(200)) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + content=AZURE_SPEECH_WAV_BYTES, + headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 400 + assert "AZURE_SPEECH_REGION" in response.text + assert not catch_all.called + + def test_missing_api_key_is_rejected_before_any_upstream_call( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.delenv("AZURE_SPEECH_API_KEY") + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(200)) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + content=AZURE_SPEECH_WAV_BYTES, + headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 400 + assert "AZURE_SPEECH_API_KEY" in response.text + assert not catch_all.called + + def test_azure_speech_is_a_mapped_pass_through_route(self) -> None: + from litellm.proxy._types import LiteLLMRoutes + + assert "/azure_speech" in LiteLLMRoutes.mapped_pass_through_routes.value + + +def _azure_speech_real_auth_attrs() -> dict[str, object]: + from litellm.caching.caching import DualCache + from litellm.proxy.utils import ProxyLogging + + user_api_key_cache: Final = DualCache() + return { + "prisma_client": None, + "user_api_key_cache": user_api_key_cache, + "proxy_logging_obj": ProxyLogging(user_api_key_cache=user_api_key_cache), + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "user_custom_auth": None, + "jwt_handler": None, + } + + +class TestAzureSpeechRawBodyThroughRealAuth: + """user_api_key_auth reads the body before the route runs; raw audio must not be parsed as JSON.""" + + def _post_wav( + self, monkeypatch: pytest.MonkeyPatch, path: str, api_key: str, body: bytes = AZURE_SPEECH_WAV_BYTES + ) -> httpx.Response: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") + monkeypatch.setenv("AZURE_SPEECH_REGION", "eastus") + monkeypatch.delenv("AZURE_SPEECH_API_BASE", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + with patch.multiple( # test-quality-ok: the real user_api_key_auth reads proxy_server module globals (master_key, caches) that have no injection seam + "litellm.proxy.proxy_server", **_azure_speech_real_auth_attrs() + ): + client = TestClient(app) + return client.post( + path, + params={"language": "en-US"}, + content=body, + headers={"Content-Type": "audio/wav", "Authorization": f"Bearer {api_key}"}, + ) + + @pytest.mark.parametrize("body", [AZURE_SPEECH_WAV_BYTES, AZURE_SPEECH_NON_UTF8_WAV_BYTES], ids=["ascii", "binary"]) + def test_master_key_with_raw_wav_body_reaches_azure_without_a_parse_attempt( + self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, body: bytes + ) -> None: + with respx.mock(assert_all_called=True) as upstream, caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + route = upstream.post( + f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" + ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) + + response = self._post_wav( + monkeypatch, f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", "sk-master-key", body=body + ) + + assert (response.status_code, response.json()) == (200, AZURE_SPEECH_TRANSCRIPT) + assert route.calls.last.request.content == body + assert [record.message for record in caplog.records if "request body" in record.message] == [] + + def test_wrong_litellm_key_with_raw_wav_body_is_rejected(self, monkeypatch: pytest.MonkeyPatch) -> None: + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(200)) + + response = self._post_wav(monkeypatch, f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", "sk-wrong") + + assert response.status_code in (400, 401), response.text + assert not catch_all.called + + @pytest.mark.parametrize("path", ["/v1/chat/completions", f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}"]) + def test_audio_content_type_off_the_short_audio_route_is_still_parsed_as_json( + self, monkeypatch: pytest.MonkeyPatch, path: str + ) -> None: + response = self._post_wav(monkeypatch, path, "sk-master-key", body=b'{}{"model": "gpt-4o"}') + + assert response.status_code == 400 + assert "Invalid JSON payload" in response.text diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoint_router.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoint_router.py index e3cbc2d507f..7a272a49853 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoint_router.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_endpoint_router.py @@ -159,6 +159,22 @@ def test_assemblyai_region_matching(): assert passthrough_router.get_credentials(custom_llm_provider="assemblyai", region_name=None) == "sk-us" +def test_azure_speech_dashboard_credential_resolves_through_flagged_deployment(monkeypatch): + monkeypatch.delenv("AZURE_SPEECH_API_KEY", raising=False) + CredentialAccessor.upsert_credentials([_credential("azure-speech-prod", "azure-subscription-key")]) + llm_router = litellm.Router( + model_list=[ + _flagged_deployment("azure_speech/short-audio", litellm_credential_name="azure-speech-prod"), + ] + ) + passthrough_router = _passthrough_router(llm_router) + + assert ( + passthrough_router.get_credentials(custom_llm_provider="azure_speech", region_name=None) + == "azure-subscription-key" + ) + + def test_env_fallback_when_no_router(monkeypatch): passthrough_router = _passthrough_router(None) monkeypatch.setenv("OPENAI_API_KEY", "sk-from-env") diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index de83cd00790..1a1a2fd73aa 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -80,6 +80,7 @@ export enum Providers { SageMaker = "AWS SageMaker", Azure = "Azure", Azure_AI_Studio = "Azure AI Foundry (Studio)", + Azure_Speech = "Azure AI Speech", AZURE_TEXT = "Azure Text", BASETEN = "Baseten", BYTEZ = "Bytez", @@ -193,6 +194,7 @@ export const provider_map: Record = { AUTO_ROUTER: "auto_router", Azure: "azure", Azure_AI_Studio: "azure_ai", + Azure_Speech: "azure_speech", AZURE_TEXT: "azure_text", BASETEN: "baseten", Bedrock: "bedrock", @@ -310,6 +312,7 @@ export const providerLogoMap: Partial> = { [Providers.AssemblyAI]: assemblyaiSmallLogo.src, [Providers.Azure]: microsoftAzureLogo.src, [Providers.Azure_AI_Studio]: microsoftAzureLogo.src, + [Providers.Azure_Speech]: microsoftAzureLogo.src, [Providers.AZURE_TEXT]: microsoftAzureLogo.src, [Providers.BASETEN]: basetenLogo.src, [Providers.Bedrock]: bedrockLogo.src, @@ -427,6 +430,7 @@ const providerPlaceholderMap: Partial> = { [Providers.Anthropic]: "claude-3-opus", [Providers.Azure]: "my-deployment", [Providers.Azure_AI_Studio]: "azure_ai/command-r-plus", + [Providers.Azure_Speech]: "azure_speech/short-audio", [Providers.Bedrock]: "claude-3-opus", [Providers.CHATGPT]: "chatgpt/gpt-5.4", [Providers.Cognition]: "cognition/swe-1.7", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 872875cc535..62b9921302a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1612,6 +1612,82 @@ export interface paths { patch: operations["azure_proxy_route_azure_ai__endpoint__patch"]; trace?: never; }; + "/azure_speech/{endpoint}": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Azure Speech Proxy Route + * @description Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + * `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + * with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + * + * The body is forwarded byte for byte and the proxy injects its own + * `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + * and is never forwarded. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + */ + get: operations["azure_speech_proxy_route_azure_speech__endpoint__get"]; + /** + * Azure Speech Proxy Route + * @description Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + * `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + * with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + * + * The body is forwarded byte for byte and the proxy injects its own + * `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + * and is never forwarded. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + */ + put: operations["azure_speech_proxy_route_azure_speech__endpoint__put"]; + /** + * Azure Speech Proxy Route + * @description Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + * `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + * with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + * + * The body is forwarded byte for byte and the proxy injects its own + * `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + * and is never forwarded. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + */ + post: operations["azure_speech_proxy_route_azure_speech__endpoint__post"]; + /** + * Azure Speech Proxy Route + * @description Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + * `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + * with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + * + * The body is forwarded byte for byte and the proxy injects its own + * `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + * and is never forwarded. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + */ + delete: operations["azure_speech_proxy_route_azure_speech__endpoint__delete"]; + options?: never; + head?: never; + /** + * Azure Speech Proxy Route + * @description Pass-through for the Azure AI Speech REST APIs (speech to text), e.g. + * `POST /azure_speech/speech/recognition/conversation/cognitiveservices/v1?language=en-US` + * with the raw audio as the body, or `POST /azure_speech/speechtotext/v3.2/transcriptions`. + * + * The body is forwarded byte for byte and the proxy injects its own + * `Ocp-Apim-Subscription-Key`; the caller's `Authorization` header is the LiteLLM key + * and is never forwarded. + * + * [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) + */ + patch: operations["azure_speech_proxy_route_azure_speech__endpoint__patch"]; + trace?: never; + }; "/batches": { parameters: { query?: never; @@ -43394,6 +43470,161 @@ export interface operations { }; }; }; + azure_speech_proxy_route_azure_speech__endpoint__get: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + azure_speech_proxy_route_azure_speech__endpoint__put: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + azure_speech_proxy_route_azure_speech__endpoint__post: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + azure_speech_proxy_route_azure_speech__endpoint__delete: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + azure_speech_proxy_route_azure_speech__endpoint__patch: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; list_batches_batches_get: { parameters: { query?: { From 6e1b4959d18d457f3045c97122162a37469ce975 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 03:14:13 +0000 Subject: [PATCH 083/253] feat(passthrough): deepgram streaming /v1/listen WebSocket passthrough with duration-based cost tracking Adds authenticated /deepgram/v1/listen and /deepgram/listen WebSocket routes that resolve the Deepgram credential through the pass-through router, inject Authorization: Token upstream, default the model to nova-3 when the client passes none, and relay audio and transcript frames unchanged. The shared WebSocket relay no longer assumes the first upstream frame is JSON and forwards every frame as received, keeping the Vertex AI Live setup handling on Vertex routes only. A Deepgram logging handler bills the call on Metadata.duration, falling back to the furthest Results start + duration, at the deepgram/ per-second rate from the model cost map Resolves LIT-7937 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 3 + litellm/llms/deepgram/common_utils.py | 22 ++ litellm/proxy/_lazy_features.py | 1 + litellm/proxy/_types.py | 3 + .../llm_passthrough_endpoints.py | 68 ++++- ...gram_listen_passthrough_logging_handler.py | 132 +++++++++ .../pass_through_endpoints.py | 113 ++++--- .../pass_through_endpoints/success_handler.py | 18 ++ ...gram_listen_passthrough_logging_handler.py | 254 ++++++++++++++++ .../test_deepgram_ws_passthrough_routes.py | 280 ++++++++++++++++++ .../test_pass_through_endpoints.py | 163 +++++++++- 11 files changed, 974 insertions(+), 83 deletions(-) create mode 100644 litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py diff --git a/litellm/constants.py b/litellm/constants.py index 8409a161800..c3a7a16ea39 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -317,6 +317,9 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float( # RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123 +DEEPGRAM_DEFAULT_API_BASE: Final = "https://api.deepgram.com/v1" +DEEPGRAM_LISTEN_DEFAULT_MODEL: Final = "nova-3" + BEDROCK_REALTIME_PENDING_SESSION_UPDATE_SCOPE_KEY: Final = "litellm.bedrock_realtime.pending_session_update" BEDROCK_REALTIME_SESSION_COMMITTED_SCOPE_KEY: Final = "litellm.bedrock_realtime.session_committed" BEDROCK_REALTIME_COMMITTED_FAILURE_SCOPE_KEY: Final = "litellm.bedrock_realtime.committed_failure" diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index a741b092a36..db00d048f01 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -1,5 +1,27 @@ +from types import MappingProxyType +from typing import Final + +import httpx + +from litellm.constants import DEEPGRAM_DEFAULT_API_BASE, DEEPGRAM_LISTEN_DEFAULT_MODEL from litellm.llms.base_llm.chat.transformation import BaseLLMException +_WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws", "wss": "wss", "ws": "ws"}) + class DeepgramException(BaseLLMException): pass + + +def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> str: + """ + The upstream ``/listen`` socket for a streaming transcription, keeping the client's query string as sent + and adding the default model only when the client named none + """ + listen_url: Final = httpx.URL(f"{(api_base or DEEPGRAM_DEFAULT_API_BASE).rstrip('/')}/listen") + websocket_url: Final = listen_url.copy_with(scheme=_WEBSOCKET_SCHEMES.get(listen_url.scheme, listen_url.scheme)) + params: Final = httpx.QueryParams(query_string) + query: Final = ( + query_string if params.get("model") else str(params.remove("model").add("model", DEEPGRAM_LISTEN_DEFAULT_MODEL)) + ) + return f"{websocket_url}?{query}" diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index faf95397fa5..ce48c4801d6 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -200,6 +200,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( "/cohere/", "/comprehendmedical", "/cursor/", + "/deepgram/", "/eu.assemblyai/", "/gemini/", "/gigachat/", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 680c63393e8..aa2068cb94e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -70,6 +70,7 @@ from litellm.types.utils import ( StandardLoggingVectorStoreRequest, StandardPassThroughResponseObject, TextCompletionResponse, + TranscriptionResponse, ) from litellm.types.videos.main import VideoObject @@ -487,6 +488,7 @@ class LiteLLMRoutes(enum.Enum): "/gigachat", "/watsonx", "/nvidia_nim", + "/deepgram", ] ######################################################### @@ -4694,6 +4696,7 @@ PassThroughEndpointLoggingResultValues = ( | VideoObject | StandardPassThroughResponseObject | ResponsesAPIResponse + | TranscriptionResponse ) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index b9b8cb3a22b..a5a95c34910 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -36,6 +36,7 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.deepgram.common_utils import deepgram_listen_websocket_target from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.passthrough.main import AsyncPassthroughStreamingResponse @@ -2573,7 +2574,7 @@ async def _openai_websocket_refusal( return None -class _OpenAIWebsocketRelay(Protocol): +class _WebsocketRelay(Protocol): async def __call__( self, *, @@ -2593,7 +2594,7 @@ def _proxy_general_settings() -> Mapping[str, object]: return general_settings -def _openai_websocket_relay() -> _OpenAIWebsocketRelay: +def _websocket_relay() -> _WebsocketRelay: return websocket_passthrough_request @@ -2611,6 +2612,19 @@ def _proxy_model_allowlists() -> _OpenAIWebsocketModelAllowlists: return resolve +def _negotiated_websocket_subprotocol(websocket: WebSocket) -> str | None: + """ + The first subprotocol the client offered, echoed back so browsers that carry the LiteLLM key in + ``Sec-WebSocket-Protocol`` complete the handshake + """ + requested_subprotocols: Final = tuple( + protocol.strip() + for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",") + if protocol.strip() + ) + return requested_subprotocols[0] if requested_subprotocols else None + + @router.websocket("/openai_passthrough/{endpoint:path}") @router.websocket("/openai/{endpoint:path}") async def openai_websocket_proxy_route( @@ -2618,16 +2632,11 @@ async def openai_websocket_proxy_route( endpoint: str, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], general_settings: Annotated[Mapping[str, object], Depends(_proxy_general_settings)], - relay: Annotated[_OpenAIWebsocketRelay, Depends(_openai_websocket_relay)], + relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)], model_allowlists: Annotated[_OpenAIWebsocketModelAllowlists, Depends(_proxy_model_allowlists)], ) -> None: """WebSocket passthrough for OpenAI prefixes (realtime / responses.connect).""" - requested_subprotocols: Final = tuple( - protocol.strip() - for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",") - if protocol.strip() - ) - negotiated_subprotocol: Final = requested_subprotocols[0] if requested_subprotocols else None + negotiated_subprotocol: Final = _negotiated_websocket_subprotocol(websocket) refusal: Final = await _openai_websocket_refusal(user_api_key_dict, general_settings, model_allowlists) if refusal is not None: @@ -2686,6 +2695,47 @@ async def openai_websocket_proxy_route( ) +_DEEPGRAM_WS_MISSING_KEY_REASON: Final = ( + "Required 'DEEPGRAM_API_KEY' in environment to make pass-through calls to Deepgram." +) + + +@router.websocket("/deepgram/v1/listen") +@router.websocket("/deepgram/listen") +async def deepgram_listen_websocket_route( + websocket: WebSocket, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], + relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)], +) -> None: + """ + Streaming speech to text through Deepgram's ``/v1/listen`` socket. Audio frames and transcript frames are + relayed unchanged; the call is billed on the audio duration Deepgram reports when the socket closes + """ + deepgram_api_key: Final = passthrough_endpoint_router.get_credentials( + custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value, + region_name=None, + ) + if deepgram_api_key is None: + await websocket.close(code=1011, reason=_DEEPGRAM_WS_MISSING_KEY_REASON) + return + + await websocket.accept(subprotocol=_negotiated_websocket_subprotocol(websocket)) + await relay( + websocket=websocket, + target=deepgram_listen_websocket_target( + api_base=get_secret_str("DEEPGRAM_API_BASE"), + query_string=websocket.url.query, + ), + custom_headers={ # mutable-ok: websocket_passthrough_request requires a plain dict of upstream headers + "Authorization": f"Token {deepgram_api_key}" + }, + user_api_key_dict=user_api_key_dict, + forward_headers=False, + endpoint=websocket.url.path, + accept_websocket=False, + ) + + class BaseOpenAIPassThroughHandler: @staticmethod async def _base_openai_pass_through_handler( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py new file mode 100644 index 00000000000..6a574a2e1b9 --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -0,0 +1,132 @@ +""" +Cost tracking for Deepgram's streaming ``/v1/listen`` WebSocket. Deepgram bills the audio it processed, which it +reports as ``duration`` on the closing ``Metadata`` frame; a stream that ends without one is billed on the furthest +``start + duration`` across its ``Results`` frames +""" + +import math +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import Final +from urllib.parse import parse_qs, urlparse + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.constants import DEEPGRAM_LISTEN_DEFAULT_MODEL +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy._types import PassThroughEndpointLoggingTypedDict +from litellm.types.utils import TranscriptionResponse + +DEEPGRAM_LISTEN_ROUTE_SUFFIX: Final = "/listen" + + +def _seconds(value: object) -> float | None: + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) if math.isfinite(value) and value >= 0 else None + + +def _results_frame_end(frame: Mapping[str, object]) -> float | None: + start: Final = _seconds(frame.get("start")) + duration: Final = _seconds(frame.get("duration")) + return None if start is None or duration is None else start + duration + + +def _final_transcript(frame: Mapping[str, object]) -> str | None: + if frame.get("is_final") is not True: + return None + channel: Final = frame.get("channel") + alternatives: Final = channel.get("alternatives") if isinstance(channel, Mapping) else None + first: Final = alternatives[0] if isinstance(alternatives, list) and alternatives else None + transcript: Final = first.get("transcript") if isinstance(first, Mapping) else None + return transcript if isinstance(transcript, str) and transcript else None + + +def deepgram_listen_audio_seconds(websocket_messages: Sequence[Mapping[str, object]]) -> float: + metadata_durations: Final = tuple( + duration + for frame in websocket_messages + if frame.get("type") == "Metadata" + if (duration := _seconds(frame.get("duration"))) is not None + ) + if metadata_durations: + return metadata_durations[-1] + return max( + ( + end + for frame in websocket_messages + if frame.get("type") == "Results" + if (end := _results_frame_end(frame)) is not None + ), + default=0.0, + ) + + +def deepgram_listen_transcript(websocket_messages: Sequence[Mapping[str, object]]) -> str: + return " ".join( + transcript + for frame in websocket_messages + if frame.get("type") == "Results" + if (transcript := _final_transcript(frame)) is not None + ) + + +def deepgram_listen_model(upstream_url: str) -> str: + models: Final = parse_qs(urlparse(upstream_url).query).get("model") + return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL + + +def _audio_cost(response: TranscriptionResponse, model: str) -> float | None: + try: + return litellm.completion_cost( + completion_response=response, + model=model, + custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value, + call_type="transcription", + ) + except Exception as e: # noqa: BLE001 # an unpriced model must not lose the spend row, only its cost + verbose_proxy_logger.warning("Deepgram listen passthrough: no pricing for model '%s': %s", model, e) + return None + + +class DeepgramListenPassthroughLoggingHandler: + @staticmethod + def is_deepgram_listen_route(url_route: str) -> bool: + path: Final = urlparse(url_route).path + return path.startswith("/deepgram/") and path.endswith(DEEPGRAM_LISTEN_ROUTE_SUFFIX) + + def deepgram_listen_passthrough_handler( + self, + websocket_messages: Sequence[Mapping[str, object]], + logging_obj: LiteLLMLoggingObj, + upstream_url: str, + kwargs: Mapping[str, object] = MappingProxyType({}), + ) -> PassThroughEndpointLoggingTypedDict: + model: Final = deepgram_listen_model(upstream_url) + audio_seconds: Final = deepgram_listen_audio_seconds(websocket_messages) + response: Final = TranscriptionResponse(text=deepgram_listen_transcript(websocket_messages)) + response._hidden_params["audio_transcription_duration"] = audio_seconds # pyright: ignore[reportPrivateUsage] # the cost calculator reads the billed duration off the response's hidden params + response_cost: Final = _audio_cost(response, model) + response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads a precomputed cost off the response's hidden params + + provider: Final = litellm.LlmProviders.DEEPGRAM.value + logging_obj.model = model # rebind-ok: the spend logger reads model and cost off the shared logging object + logging_obj.model_call_details["model"] = model # rebind-ok: same shared logging object + logging_obj.model_call_details["custom_llm_provider"] = provider # rebind-ok: same shared logging object + logging_obj.model_call_details["response_cost"] = response_cost # rebind-ok: same shared logging object + verbose_proxy_logger.debug( + "Deepgram listen passthrough cost tracking: model %s, audio seconds %s, cost %s", + model, + audio_seconds, + response_cost, + ) + logging_result: Final[PassThroughEndpointLoggingTypedDict] = { + "result": response, + "kwargs": { + **kwargs, + "model": model, + "custom_llm_provider": provider, + "response_cost": response_cost, + }, + } + return logging_result diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 685c19062bb..cf6985852f2 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -8,7 +8,7 @@ from base64 import b64encode from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime -from itertools import groupby +from itertools import count, groupby from typing import TYPE_CHECKING, Any, Final, TypedDict, cast from urllib.parse import urlencode, urlparse @@ -2120,6 +2120,17 @@ def _resolved_vertex_live_setup( return {**setup_data, "model": setup_model_rewriter(setup_model)} +def _json_object_frame(frame: str | bytes) -> dict[str, object] | None: + """ + The frame as a JSON object when it is one, for cost tracking; audio and non-object frames yield None + """ + try: + decoded: Final = json.loads(frame if isinstance(frame, str) else frame.decode("utf-8")) + except (json.JSONDecodeError, UnicodeDecodeError): + return None + return decoded if isinstance(decoded, dict) else None + + def _truncated_close_reason(reason: str) -> str: """ Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character @@ -2401,70 +2412,46 @@ async def websocket_passthrough_request( ) await upstream_ws.close() + def _extract_vertex_live_model_from_setup_response(setup_response: Mapping[str, object]) -> None: + extracted_model: Final = _extract_model_from_vertex_ai_setup(setup_response) + if not extracted_model: + verbose_proxy_logger.warning( + "WebSocket passthrough (%s): Failed to extract model from server setup response: %s", + endpoint, + setup_response, + ) + return + kwargs["model"] = extracted_model + kwargs["custom_llm_provider"] = "vertex_ai_language_models" + logging_obj.model = extracted_model + logging_obj.model_call_details["model"] = extracted_model + logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai_language_models" + + is_vertex_live: Final = bool(endpoint and "/vertex_ai/live" in endpoint) + json_frame_ordinal: Final = count() + + async def relay_upstream_frame(upstream_message: str | bytes) -> None: + """ + Send the frame to the client exactly as received, then keep it for cost tracking when it is a JSON + object; the Vertex AI Live setup acknowledgement only names the model, so it is read instead of kept + """ + if isinstance(upstream_message, bytes): + await websocket.send_bytes(upstream_message) + else: + await websocket.send_text(upstream_message) + message_data: Final = _json_object_frame(upstream_message) + if message_data is None: + return + if is_vertex_live and next(json_frame_ordinal) == 0: + _extract_vertex_live_model_from_setup_response(message_data) + return + websocket_messages.append(message_data) + async def forward_upstream_to_client() -> Close | None: - """Forward messages from upstream to client WebSocket, returning the upstream's close frame""" + """Relay upstream frames to the client until the upstream closes, returning its close frame""" try: - # Wait for the first response from upstream - raw_response = await upstream_ws.recv(decode=False) - # Ensure raw_response is bytes before decoding - if isinstance(raw_response, str): - raw_response = raw_response.encode("utf-8") - setup_response: Final[Mapping[str, object]] = json.loads(raw_response.decode("utf-8")) - verbose_proxy_logger.debug("Setup response: %s", setup_response) - - # Extract model and provider from setup response for Vertex AI Live - if endpoint and "/vertex_ai/live" in endpoint: - verbose_proxy_logger.debug( - "WebSocket passthrough (%s): Processing server setup response for model extraction", - endpoint, - ) - extracted_model: Final = _extract_model_from_vertex_ai_setup(setup_response) - if extracted_model: - kwargs["model"] = extracted_model - kwargs["custom_llm_provider"] = "vertex_ai_language_models" - # Update logging object with correct model - logging_obj.model = extracted_model - logging_obj.model_call_details["model"] = extracted_model - logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai_language_models" - verbose_proxy_logger.debug( - "WebSocket passthrough (%s): Successfully extracted model '%s' and set provider to 'vertex_ai' from server setup response", - endpoint, - extracted_model, - ) - else: - verbose_proxy_logger.warning( - "WebSocket passthrough (%s): Failed to extract model from server setup response: %s", - endpoint, - setup_response, - ) - else: - verbose_proxy_logger.debug( - "WebSocket passthrough (%s): Not a Vertex AI Live endpoint, skipping model extraction", - endpoint, - ) - - # Send the setup response to the client - await websocket.send_text(json.dumps(setup_response)) - - # Now continuously forward messages from upstream to client - async for upstream_message in upstream_ws: - if isinstance(upstream_message, bytes): - await websocket.send_bytes(upstream_message) - # Parse and collect for cost tracking - try: - message_data: dict[str, object] = json.loads(upstream_message.decode()) - websocket_messages.append(message_data) - except (json.JSONDecodeError, UnicodeDecodeError): - pass - else: - await websocket.send_text(upstream_message) - # Parse and collect for cost tracking - try: - message_data = json.loads(upstream_message) - websocket_messages.append(message_data) - except json.JSONDecodeError: - pass - + while True: + await relay_upstream_frame(await upstream_ws.recv()) except (ConnectionClosedOK, ConnectionClosedError) as e: verbose_proxy_logger.debug("Upstream WebSocket connection closed: %s", e) return e.rcvd diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 76a471302f4..82bb47a60ab 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -24,6 +24,9 @@ from .llm_provider_handlers.cohere_passthrough_logging_handler import ( from .llm_provider_handlers.cursor_passthrough_logging_handler import ( CursorPassthroughLoggingHandler, ) +from .llm_provider_handlers.deepgram_listen_passthrough_logging_handler import ( + DeepgramListenPassthroughLoggingHandler, +) from .llm_provider_handlers.gemini_passthrough_logging_handler import ( GeminiPassthroughLoggingHandler, ) @@ -278,6 +281,21 @@ class PassThroughEndpointLogging: standard_logging_response_object = vertex_ai_live_handler_result["result"] kwargs = vertex_ai_live_handler_result["kwargs"] + elif DeepgramListenPassthroughLoggingHandler.is_deepgram_listen_route(url_route): + deepgram_handler_result: Final = ( + DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=tuple( + message + for message in (response_body if isinstance(response_body, list) else ()) + if isinstance(message, dict) + ), + logging_obj=logging_obj, + upstream_url=str(httpx_response.request.url), + kwargs=kwargs, + ) + ) + standard_logging_response_object = deepgram_handler_result["result"] # rebind-ok: elif-chain + kwargs = deepgram_handler_result["kwargs"] # rebind-ok: elif-chain contract return_dict["standard_logging_response_object"] = standard_logging_response_object return_dict["kwargs"] = kwargs diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py new file mode 100644 index 00000000000..f00411ac10c --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -0,0 +1,254 @@ +"""Deepgram ``/v1/listen`` WebSocket passthrough: duration extraction and duration based cost tracking.""" + +import math +from collections.abc import Mapping, Sequence +from datetime import datetime +from types import SimpleNamespace +from typing import Final + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.deepgram_listen_passthrough_logging_handler import ( + DeepgramListenPassthroughLoggingHandler, + deepgram_listen_audio_seconds, + deepgram_listen_model, + deepgram_listen_transcript, +) +from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging +from litellm.types.passthrough_endpoints.pass_through_endpoints import PassthroughStandardLoggingPayload +from litellm.types.utils import StandardLoggingPayload, TranscriptionResponse + +NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000" + + +def _results(start: object, duration: object, transcript: str = "", is_final: object = True) -> dict[str, object]: + return { + "type": "Results", + "start": start, + "duration": duration, + "is_final": is_final, + "channel": {"alternatives": [{"transcript": transcript, "confidence": 0.9}]}, + } + + +def _metadata(duration: object) -> dict[str, object]: + return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": 1} + + +@pytest.mark.parametrize( + ("frames", "expected_seconds"), + [ + pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _metadata(6.25)), 6.25, id="metadata wins"), + pytest.param((_metadata(4.0), _results(0.0, 9.0), _metadata(5.5)), 5.5, id="last metadata wins"), + pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _results(1.0, 1.0)), 5.5, id="furthest results end"), + pytest.param((_results(0.0, 0.0), _metadata(0.0)), 0.0, id="zero metadata is a real zero"), + pytest.param((_results(0.0, 1.5), _metadata("6.25")), 1.5, id="string metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(True)), 1.5, id="boolean metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(-3.0)), 1.5, id="negative metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(math.nan), _metadata(math.inf)), 1.5, id="nan/inf ignored"), + pytest.param((_results("0", 2.0), _results(0.0, None), _results(0.0, 0.75)), 0.75, id="malformed results"), + pytest.param(({"type": "SpeechStarted", "timestamp": 3.0}, {"type": "UtteranceEnd"}), 0.0, id="no usage"), + pytest.param((), 0.0, id="no frames"), + ], +) +def test_deepgram_listen_audio_seconds(frames: Sequence[Mapping[str, object]], expected_seconds: float): + assert deepgram_listen_audio_seconds(frames) == expected_seconds + + +def test_deepgram_listen_transcript_joins_final_results_only(): + frames = ( + _results(0.0, 1.0, "hello wor", is_final=False), + _results(0.0, 1.5, "hello world"), + _results(1.5, 0.5, "", is_final=True), + _results(2.0, 1.0, "how are you", is_final="yes"), + {"type": "Results", "start": 3.0, "duration": 1.0, "is_final": True, "channel": {"alternatives": []}}, + _results(4.0, 1.0, "goodbye"), + _metadata(5.0), + ) + assert deepgram_listen_transcript(frames) == "hello world goodbye" + + +@pytest.mark.parametrize( + ("upstream_url", "expected_model"), + [ + (NOVA_3_URL, "nova-3"), + ("wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-2-medical", "nova-2-medical"), + ("wss://api.deepgram.com/v1/listen?model=nova-3&model=nova-2", "nova-3"), + ("wss://api.deepgram.com/v1/listen?encoding=linear16", litellm.constants.DEEPGRAM_LISTEN_DEFAULT_MODEL), + ], +) +def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, expected_model: str): + assert deepgram_listen_model(upstream_url) == expected_model + + +@pytest.mark.parametrize( + ("url_route", "expected"), + [ + ("/deepgram/v1/listen", True), + ("/deepgram/listen", True), + ("/deepgram/v1/listen?model=nova-3", True), + ("/deepgram/v1/speak", False), + ("/deepgram/v1/listen/extra", False), + ("/openai/v1/realtime", False), + ("/vertex_ai/live", False), + ("", False), + ], +) +def test_is_deepgram_listen_route(url_route: str, expected: bool): + assert DeepgramListenPassthroughLoggingHandler.is_deepgram_listen_route(url_route) is expected + + +def _logging_obj(call_id: str = "call-dg") -> LiteLLMLoggingObj: + return LiteLLMLoggingObj( + model="unknown", + messages=[{"role": "user", "content": "WebSocket connection"}], + stream=True, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id=call_id, + function_id="websocket_passthrough", + ) + + +def _registry_cost(model: str, seconds: float) -> float: + """Derives the expected charge from the live cost map rather than pinning a vendor price.""" + per_second: Final = litellm.model_cost[f"deepgram/{model}"]["input_cost_per_second"] + assert per_second > 0 + return per_second * seconds + + +def test_handler_bills_metadata_duration_at_the_registry_rate_and_names_the_model(): + frames = (_results(0.0, 5.0, "first sentence"), _results(5.0, 7.5, "second sentence"), _metadata(12.5)) + logging_obj = _logging_obj() + + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=frames, + logging_obj=logging_obj, + upstream_url=NOVA_3_URL, + kwargs={"litellm_params": {"metadata": {}}}, + ) + + result = handler_result["result"] + assert isinstance(result, TranscriptionResponse) + assert result.text == "first sentence second sentence" + assert result._hidden_params["audio_transcription_duration"] == 12.5 + assert result._hidden_params["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) + assert handler_result["kwargs"]["model"] == "nova-3" + assert handler_result["kwargs"]["custom_llm_provider"] == "deepgram" + assert handler_result["kwargs"]["litellm_params"] == {"metadata": {}} + assert logging_obj.model == "nova-3" + assert logging_obj.model_call_details["model"] == "nova-3" + assert logging_obj.model_call_details["custom_llm_provider"] == "deepgram" + assert logging_obj.model_call_details["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) + + +def test_handler_falls_back_to_results_frames_when_the_stream_ends_without_metadata(): + frames = (_results(0.0, 30.0, "a"), _results(30.0, 30.0, "b"), _results(60.0, 12.5, "c")) + + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=frames, logging_obj=_logging_obj(), upstream_url=NOVA_3_URL + ) + + assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 72.5 + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 72.5)) + + +def test_handler_charges_more_for_more_audio_on_the_same_model(): + short = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(10.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL + ) + long = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(30.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL + ) + + assert long["kwargs"]["response_cost"] == pytest.approx(3 * short["kwargs"]["response_cost"]) + assert short["kwargs"]["response_cost"] > 0 + + +def test_handler_keeps_the_spend_row_but_no_cost_for_an_unpriced_model(): + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(12.5),), + logging_obj=_logging_obj(), + upstream_url="wss://api.deepgram.com/v1/listen?model=nova-99-not-in-registry", + ) + + assert handler_result["kwargs"]["model"] == "nova-99-not-in-registry" + assert handler_result["kwargs"]["response_cost"] is None + assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 12.5 + + +class _CapturingLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[StandardLoggingPayload] = [] + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + self.payloads.append(kwargs["standard_logging_object"]) + + +@pytest.mark.asyncio +async def test_success_handler_dispatches_deepgram_listen_and_logs_duration_based_spend(monkeypatch): + """Drives the shared passthrough success handler the way the WebSocket relay does at socket close and reads + what a spend logger receives: Deepgram model and provider, the audio duration billed at the registry rate.""" + capturing_logger = _CapturingLogger() + monkeypatch.setattr(litellm, "_async_success_callback", [capturing_logger]) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "callbacks", []) + logging_obj = _logging_obj("call-dg-e2e") + frames = [_results(0.0, 5.0, "hello world", is_final=False), _results(0.0, 5.0, "hello world"), _metadata(20.0)] + user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", team_id="team-stt", user_id="user-1") + start_time = datetime.now() + passthrough_logging_payload = PassthroughStandardLoggingPayload( + url=NOVA_3_URL, request_body={}, request_method="WEBSOCKET", cost_per_request=None + ) + logging_obj.update_environment_variables( + model="unknown", + user="unknown", + optional_params={}, + litellm_params={ + "metadata": { + "user_api_key": user_api_key_dict.api_key, + "user_api_key_team_id": user_api_key_dict.team_id, + "user_api_key_user_id": user_api_key_dict.user_id, + } + }, + call_type="pass_through_endpoint", + ) + + await PassThroughEndpointLogging().pass_through_async_success_handler( + httpx_response=SimpleNamespace( + status_code=200, + text="WebSocket connection successful", + headers={}, + request=SimpleNamespace(method="WEBSOCKET", url=NOVA_3_URL), + ), + response_body=frames, + logging_obj=logging_obj, + url_route="/deepgram/v1/listen", + result="websocket_connection_successful", + start_time=start_time, + end_time=datetime.now(), + cache_hit=False, + request_body={}, + passthrough_logging_payload=passthrough_logging_payload, + litellm_params={ + "metadata": { + "user_api_key": user_api_key_dict.api_key, + "user_api_key_team_id": user_api_key_dict.team_id, + "user_api_key_user_id": user_api_key_dict.user_id, + } + }, + ) + + assert len(capturing_logger.payloads) == 1 + payload = capturing_logger.payloads[0] + assert payload["model"] == "nova-3" + assert payload["custom_llm_provider"] == "deepgram" + assert payload["response_cost"] == pytest.approx(_registry_cost("nova-3", 20.0)) + assert payload["metadata"]["user_api_key_team_id"] == "team-stt" + assert payload["id"] == "call-dg-e2e" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py new file mode 100644 index 00000000000..64804e621ba --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -0,0 +1,280 @@ +"""Deepgram ``/v1/listen`` passthrough WebSocket route: registration, auth, credential injection, target URL.""" + +from collections.abc import Mapping +from dataclasses import dataclass +from types import MappingProxyType, SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from starlette.routing import WebSocketRoute +from starlette.websockets import WebSocketDisconnect + +from litellm.proxy._lazy_features import LAZY_FEATURES +from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _websocket_relay, + deepgram_listen_websocket_route, + router, +) + +GET_CREDENTIALS: Final = ( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" +) +USER_API_KEY_AUTH: Final = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth" +LISTEN_PATHS: Final = ("/deepgram/v1/listen", "/deepgram/listen") + + +class _FakeWebSocket: + def __init__(self, path: str, query: str) -> None: + self.url = SimpleNamespace(path=path, query=query) + self.headers = {"authorization": "Bearer sk-litellm-virtual", "x-api-key": "sk-caller-secret"} + self.accepts: list[str | None] = [] + self.closed: tuple[int, str] | None = None + + async def accept(self, subprotocol: str | None = None) -> None: + self.accepts.append(subprotocol) + + async def close(self, code: int = 1000, reason: str = "") -> None: + self.closed = (code, reason) + + +@dataclass(frozen=True, slots=True) +class _RelayCall: + target: str + custom_headers: Mapping[str, str] + user_api_key_dict: UserAPIKeyAuth + forward_headers: bool + endpoint: str + accept_websocket: bool + + +class _FakeRelay: + def __init__(self) -> None: + self.calls: list[_RelayCall] = [] + + async def __call__( + self, + *, + websocket: object, + target: str, + custom_headers: dict[str, str], + user_api_key_dict: UserAPIKeyAuth, + forward_headers: bool, + endpoint: str, + accept_websocket: bool, + ) -> None: + self.calls.append( + _RelayCall( + target=target, + custom_headers=MappingProxyType(dict(custom_headers)), + user_api_key_dict=user_api_key_dict, + forward_headers=forward_headers, + endpoint=endpoint, + accept_websocket=accept_websocket, + ) + ) + + +async def _serve(websocket: _FakeWebSocket, user_api_key_dict: UserAPIKeyAuth | None = None) -> _FakeRelay: + relay = _FakeRelay() + await deepgram_listen_websocket_route( + websocket=websocket, + user_api_key_dict=user_api_key_dict or UserAPIKeyAuth(), + relay=relay, + ) + return relay + + +def test_deepgram_listen_websocket_routes_registered(): + ws_paths = {route.path for route in router.routes if isinstance(route, WebSocketRoute)} + assert set(LISTEN_PATHS) <= ws_paths + + +@pytest.mark.parametrize("path", LISTEN_PATHS) +def test_deepgram_listen_is_a_lazily_loaded_mapped_pass_through_route(path): + """The route must be reachable before the passthrough module is imported and must be authed and + billed as a mapped pass-through route like the other provider prefixes.""" + feature = next(feature for feature in LAZY_FEATURES if feature.name == "llm_passthrough") + assert feature.matches(path) + assert any(path.startswith(prefix) for prefix in LiteLLMRoutes.mapped_pass_through_routes.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", LISTEN_PATHS) +async def test_deepgram_listen_forwards_query_and_injects_only_provider_auth(path, monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket(path, "encoding=linear16&sample_rate=16000&keywords=hi%3A2&keywords=there") + caller = UserAPIKeyAuth(api_key="sk-litellm-virtual", team_id="team-stt") + + with patch(GET_CREDENTIALS, return_value="dg-provider-key") as get_credentials: + relay = await _serve(websocket, caller) + + assert get_credentials.call_args.kwargs == {"custom_llm_provider": "deepgram", "region_name": None} + assert relay.calls == [ + _RelayCall( + target=( + "wss://api.deepgram.com/v1/listen" + "?encoding=linear16&sample_rate=16000&keywords=hi%3A2&keywords=there&model=nova-3" + ), + custom_headers=MappingProxyType({"Authorization": "Token dg-provider-key"}), + user_api_key_dict=caller, + forward_headers=False, + endpoint=path, + accept_websocket=False, + ) + ] + assert websocket.accepts == [None] + assert websocket.closed is None + + +@pytest.mark.asyncio +async def test_deepgram_listen_keeps_caller_chosen_model(monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2&language=en") + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2&language=en"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("query", "expected_target"), + [ + ("", "wss://api.deepgram.com/v1/listen?model=nova-3"), + ("model=", "wss://api.deepgram.com/v1/listen?model=nova-3"), + ("model=&language=en", "wss://api.deepgram.com/v1/listen?language=en&model=nova-3"), + ], +) +async def test_deepgram_listen_defaults_to_nova_3_when_no_model_is_named(query, expected_target, monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket("/deepgram/listen", query) + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert [call.target for call in relay.calls] == [expected_target] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("api_base", "expected_target"), + [ + ("https://api.eu.deepgram.com/v1/", "wss://api.eu.deepgram.com/v1/listen?model=nova-3"), + ("http://localhost:8080/v1", "ws://localhost:8080/v1/listen?model=nova-3"), + ("wss://deepgram.internal.example/v1", "wss://deepgram.internal.example/v1/listen?model=nova-3"), + ], +) +async def test_deepgram_listen_honours_server_configured_api_base(api_base, expected_target, monkeypatch): + monkeypatch.setenv("DEEPGRAM_API_BASE", api_base) + websocket = _FakeWebSocket("/deepgram/v1/listen", "") + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert [call.target for call in relay.calls] == [expected_target] + + +@pytest.mark.asyncio +async def test_deepgram_listen_ignores_caller_supplied_api_base(monkeypatch): + """V1: the server-configured Deepgram key must only ever go to the server-configured host.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket("/deepgram/v1/listen", "api_base=wss%3A%2F%2Fattacker.example%2Fv1&model=nova-3") + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert [call.target for call in relay.calls] == [ + "wss://api.deepgram.com/v1/listen?api_base=wss%3A%2F%2Fattacker.example%2Fv1&model=nova-3" + ] + + +@pytest.mark.asyncio +async def test_deepgram_listen_closes_cleanly_when_provider_credentials_missing(): + websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-3") + + with patch(GET_CREDENTIALS, return_value=None): + relay = await _serve(websocket) + + assert websocket.closed is not None + assert websocket.closed[0] == 1011 + assert "DEEPGRAM_API_KEY" in websocket.closed[1] + assert websocket.accepts == [] + assert relay.calls == [] + + +def _app_with_relay(relay: _FakeRelay) -> FastAPI: + app = FastAPI() + app.include_router(router) + app.dependency_overrides[_websocket_relay] = lambda: relay + return app + + +def test_deepgram_listen_rejects_connections_without_a_litellm_key(): + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + + with patch(GET_CREDENTIALS, return_value="dg-provider-key") as get_credentials: + with pytest.raises(WebSocketDisconnect) as disconnect: + with client.websocket_connect("/deepgram/v1/listen?model=nova-3"): + pass + + assert disconnect.value.code == 1008 + assert relay.calls == [] + get_credentials.assert_not_called() + + +def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + caller = UserAPIKeyAuth(api_key="hashed-sk-litellm", team_id="team-stt") + + with ( + patch(GET_CREDENTIALS, return_value="dg-provider-key"), + patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=caller)) as auth, + ): + with client.websocket_connect( + "/deepgram/v1/listen?model=nova-3&punctuate=true", + headers={"Authorization": "Bearer sk-litellm-virtual"}, + ): + pass + + assert auth.await_args.kwargs["api_key"] == "Bearer sk-litellm-virtual" + assert relay.calls == [ + _RelayCall( + target="wss://api.deepgram.com/v1/listen?model=nova-3&punctuate=true", + custom_headers=MappingProxyType({"Authorization": "Token dg-provider-key"}), + user_api_key_dict=caller, + forward_headers=False, + endpoint="/deepgram/v1/listen", + accept_websocket=False, + ) + ] + + +def test_deepgram_listen_echoes_the_browser_subprotocol_that_carries_the_litellm_key(monkeypatch): + """Browsers cannot set headers, so they send the key as a subprotocol and abort the handshake unless the + server echoes that subprotocol back; the key itself must still stay off the upstream connection.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + + with ( + patch(GET_CREDENTIALS, return_value="dg-provider-key"), + patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=UserAPIKeyAuth(api_key="hashed"))), + ): + with client.websocket_connect( + "/deepgram/v1/listen?model=nova-3", + subprotocols=["openai-insecure-api-key.sk-litellm-virtual"], + ) as connection: + assert connection.accepted_subprotocol == "openai-insecure-api-key.sk-litellm-virtual" + + assert [call.custom_headers for call in relay.calls] == [ + MappingProxyType({"Authorization": "Token dg-provider-key"}) + ] + assert [call.forward_headers for call in relay.calls] == [False] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index d854ee39ff4..de13e52498f 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4844,18 +4844,21 @@ async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate(): class FakeUpstreamWebSocket: - def __init__(self, first_frame: bytes): - self._first_frame = first_frame + """Serves the given frames in order, then closes normally, the way a real websockets connection does""" + + def __init__(self, *frames: str | bytes): + self._frames = iter(frames) self.close = AsyncMock() + self.send = AsyncMock() - async def recv(self, decode: bool = True): - return self._first_frame + async def recv(self, decode: bool | None = None): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close - def __aiter__(self): - return self - - async def __anext__(self): - raise StopAsyncIteration + frame = next(self._frames, None) + if frame is None: + raise ConnectionClosedOK(rcvd=Close(1000, ""), sent=Close(1000, ""), rcvd_then_sent=True) + return frame class FakeUpstreamConnect: @@ -4876,7 +4879,7 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame(): first_frame = json.dumps( {"type": "session.created", "session": {"instructions": "Hablas español, ¿sí?"}}, ensure_ascii=False, - ).encode("utf-8") + ) upstream_ws = FakeUpstreamWebSocket(first_frame) websocket = MagicMock() @@ -4930,7 +4933,7 @@ async def test_websocket_passthrough_propagates_active_trace_context( from starlette.websockets import WebSocketState captured: dict[str, dict[str, str]] = {} - upstream_ws = FakeUpstreamWebSocket(b"{}") + upstream_ws = FakeUpstreamWebSocket("{}") def fake_connect(target, additional_headers): captured["headers"] = additional_headers @@ -5359,6 +5362,144 @@ async def test_websocket_passthrough_does_not_close_twice_when_success_logging_f websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason) +DEEPGRAM_LISTEN_TARGET = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000" +DEEPGRAM_INTERIM_FRAME = json.dumps( + { + "type": "Results", + "start": 0.0, + "duration": 1.02, + "is_final": False, + "channel": {"alternatives": [{"transcript": "hello wor", "confidence": 0.71}]}, + } +) +DEEPGRAM_FINAL_FRAME = json.dumps( + { + "type": "Results", + "start": 0.0, + "duration": 2.5, + "is_final": True, + "speech_final": True, + "channel": {"alternatives": [{"transcript": "hello world, ¿qué tal?", "confidence": 0.98}]}, + }, + ensure_ascii=False, +) +DEEPGRAM_METADATA_FRAME = json.dumps({"type": "Metadata", "request_id": "req-1", "duration": 2.5, "channels": 1}) + + +async def _relay_deepgram_listen(upstream_ws, client_receive): + """Runs the generic relay the way the Deepgram route does and returns (client websocket, success handler mock)""" + websocket = _client_websocket(client_receive) + with ( + _patched_websocket_passthrough_environment(upstream_ws), + patch( # test-quality-ok: pass_through_endpoint_logging is a module global read inside websocket_passthrough_request; there is no injection seam + "litellm.proxy.pass_through_endpoints.pass_through_endpoints." + "pass_through_endpoint_logging.pass_through_async_success_handler", + new=AsyncMock(), + ) as success_handler, + ): + await websocket_passthrough_request( + websocket=websocket, + target=DEEPGRAM_LISTEN_TARGET, + custom_headers={"Authorization": "Token dg-provider-key"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/deepgram/v1/listen", + accept_websocket=False, + ) + return websocket, success_handler + + +@pytest.mark.asyncio +async def test_websocket_passthrough_relays_deepgram_transcript_frames_verbatim_and_keeps_them_for_billing(): + """Interim, final and Metadata frames reach the client byte for byte (no JSON round trip, non-ASCII intact, + a binary frame first) and every JSON object frame is what the success handler gets to bill from.""" + upstream_ws = FakeUpstreamWebSocket( + b"\x00\x01binary-first", + DEEPGRAM_INTERIM_FRAME, + "not json at all", + DEEPGRAM_FINAL_FRAME, + DEEPGRAM_METADATA_FRAME, + ) + + websocket, success_handler = await _relay_deepgram_listen(upstream_ws, _pending_receive) + + assert [call.args[0] for call in websocket.send_bytes.await_args_list] == [b"\x00\x01binary-first"] + assert [call.args[0] for call in websocket.send_text.await_args_list] == [ + DEEPGRAM_INTERIM_FRAME, + "not json at all", + DEEPGRAM_FINAL_FRAME, + DEEPGRAM_METADATA_FRAME, + ] + success_call = success_handler.call_args.kwargs + assert success_call["url_route"] == "/deepgram/v1/listen" + assert success_call["response_body"] == [ + json.loads(DEEPGRAM_INTERIM_FRAME), + json.loads(DEEPGRAM_FINAL_FRAME), + json.loads(DEEPGRAM_METADATA_FRAME), + ] + assert success_call["httpx_response"].request.url == DEEPGRAM_LISTEN_TARGET + assert success_call["logging_obj"].model_call_details.get("custom_llm_provider") is None + websocket.close.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_websocket_passthrough_sends_deepgram_audio_bytes_and_control_text_upstream_unchanged(): + upstream_ws = RecordingUpstreamWebSocket() + audio_chunk = bytes(range(256)) * 4 + close_stream = json.dumps({"type": "CloseStream"}) + + await _relay_deepgram_listen( + upstream_ws, + AsyncMock( + side_effect=[ + {"type": "websocket.receive", "bytes": audio_chunk}, + {"type": "websocket.receive", "text": close_stream}, + {"type": "websocket.disconnect"}, + ] + ), + ) + + assert [call.args[0] for call in upstream_ws.send.await_args_list] == [audio_chunk, close_stream] + assert isinstance(upstream_ws.send.await_args_list[0].args[0], bytes) + upstream_ws.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_websocket_passthrough_vertex_live_setup_ack_names_the_model_but_is_not_billed_as_usage(): + """Vertex Live keeps its special first frame: the setup acknowledgement is forwarded verbatim, read for the + model, and left out of the frames the usage handler sees; later frames are kept as before.""" + setup_ack = json.dumps( + {"setupComplete": {}, "model": "projects/p/locations/global/publishers/google/models/gemini-live-2.5-flash"} + ) + server_content = json.dumps({"serverContent": {"turnComplete": True}, "usageMetadata": {"totalTokenCount": 12}}) + upstream_ws = FakeUpstreamWebSocket(setup_ack, server_content) + websocket = _client_websocket(_pending_receive) + + with ( + _patched_websocket_passthrough_environment(upstream_ws), + patch( # test-quality-ok: pass_through_endpoint_logging is a module global read inside websocket_passthrough_request; there is no injection seam + "litellm.proxy.pass_through_endpoints.pass_through_endpoints." + "pass_through_endpoint_logging.pass_through_async_success_handler", + new=AsyncMock(), + ) as success_handler, + ): + await websocket_passthrough_request( + websocket=websocket, + target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent", + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + ) + + assert [call.args[0] for call in websocket.send_text.await_args_list] == [setup_ack, server_content] + success_call = success_handler.call_args.kwargs + assert success_call["response_body"] == [json.loads(server_content)] + assert success_call["logging_obj"].model == "gemini-live-2.5-flash" + assert success_call["logging_obj"].model_call_details["custom_llm_provider"] == "vertex_ai_language_models" + + def _passthrough_kwargs_for_reservation( user_api_key_dict: UserAPIKeyAuth, parsed_body: Optional[dict] = None, From f2305879d06073dfa2fa9f8001659c48ae61e7d1 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 03:31:29 +0000 Subject: [PATCH 084/253] feat(proxy): price Azure Speech short audio pass-through from the recognized duration Short audio responses carry Offset and Duration in 100ns ticks; convert their sum to seconds and price it with the existing azure/speech/azure-stt entry through transcription_cost. Batch calls and responses without an integer duration stay at zero cost Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 2 + ...zure_speech_passthrough_logging_handler.py | 54 ++++++++--- .../pass_through_endpoints/success_handler.py | 1 + ...zure_speech_passthrough_logging_handler.py | 95 ++++++++++++++++--- .../test_llm_pass_through_endpoints.py | 43 +++++++++ 5 files changed, 171 insertions(+), 24 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 69de889a326..15a1d054e26 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1579,6 +1579,8 @@ AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN: Final = "api.cognitive.microsoft.com" AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER: Final = "Ocp-Apim-Subscription-Key" AZURE_SPEECH_SHORT_AUDIO_MODEL: Final = "short-audio" AZURE_SPEECH_BATCH_MODEL: Final = "batch-transcription" +AZURE_SPEECH_PRICING_MODEL: Final = "azure/speech/azure-stt" +AZURE_SPEECH_TICKS_PER_SECOND: Final = 10_000_000 BASE_MCP_ROUTE: Final = "/mcp" diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py index a7084a9545e..74587acd453 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py @@ -1,4 +1,4 @@ -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from datetime import datetime from typing import Final from urllib.parse import urlparse @@ -9,9 +9,12 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import ( AZURE_SPEECH_BATCH_MODEL, AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + AZURE_SPEECH_PRICING_MODEL, AZURE_SPEECH_SHORT_AUDIO_MODEL, AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, + AZURE_SPEECH_TICKS_PER_SECOND, ) +from litellm.cost_calculator import transcription_cost from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import ( get_standard_logging_object_payload, @@ -21,16 +24,50 @@ from litellm.types.utils import StandardPassThroughResponseObject class AzureSpeechPassthroughLoggingHandler: + @staticmethod + def _is_short_audio_route(url_route: str) -> bool: + return urlparse(url_route).path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX) + @staticmethod def _model_from_url_route(url_route: str) -> str: - path: Final = urlparse(url_route).path - if path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX): + if AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_SHORT_AUDIO_MODEL}" return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_BATCH_MODEL}" + @staticmethod + def _recognized_audio_seconds(response_body: Mapping[str, object] | Sequence[object] | None) -> float: + if not isinstance(response_body, Mapping): + return 0.0 + offset: Final = response_body.get("Offset") + duration: Final = response_body.get("Duration") + if not isinstance(offset, int) or not isinstance(duration, int): + return 0.0 + return (offset + duration) / AZURE_SPEECH_TICKS_PER_SECOND + + @staticmethod + def _response_cost(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: + if not AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): + return 0.0 + audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body) + if audio_seconds <= 0.0: + return 0.0 + try: + prompt_cost, completion_cost = transcription_cost( + model=AZURE_SPEECH_PRICING_MODEL, + custom_llm_provider="azure", + duration=audio_seconds, + ) + except Exception as e: # noqa: BLE001 # a missing price entry must not drop the spend log row + verbose_proxy_logger.warning( + "No price for %s, logging Azure Speech call at zero cost: %s", AZURE_SPEECH_PRICING_MODEL, e + ) + return 0.0 + return prompt_cost + completion_cost + @staticmethod def azure_speech_passthrough_handler( httpx_response: httpx.Response, + response_body: Mapping[str, object] | Sequence[object] | None, logging_obj: LiteLLMLoggingObj, url_route: str, result: str, @@ -40,25 +77,20 @@ class AzureSpeechPassthroughLoggingHandler: request_body: Mapping[str, object], **kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler ) -> PassThroughEndpointLoggingTypedDict: - """ - Records model and provider for an Azure AI Speech REST call. Azure bills per audio - hour after the fact and neither the short-audio response nor the batch job carries - a billable duration this path can trust, so response_cost is recorded as 0.0 rather - than estimated. - """ try: model_name: Final = AzureSpeechPassthroughLoggingHandler._model_from_url_route(url_route) + response_cost: Final = AzureSpeechPassthroughLoggingHandler._response_cost(url_route, response_body) updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict **kwargs, "model": model_name, "custom_llm_provider": AZURE_SPEECH_CUSTOM_LLM_PROVIDER, - "response_cost": 0.0, + "response_cost": response_cost, } logging_obj.model_call_details.update( model=model_name, custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER, - response_cost=0.0, + response_cost=response_cost, ) standard_logging_object: Final = get_standard_logging_object_payload( diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 919de5c1088..2ccb8ad525d 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -264,6 +264,7 @@ class PassThroughEndpointLogging: azure_speech_handler_result: Final = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( httpx_response=httpx_response, + response_body=response_body, logging_obj=logging_obj, url_route=url_route, result=result, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py index 91f7bf94281..5ffb1f7785d 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py @@ -1,9 +1,11 @@ +import json from datetime import datetime from unittest.mock import MagicMock import httpx import pytest +import litellm from litellm.proxy.pass_through_endpoints.llm_provider_handlers.azure_speech_passthrough_logging_handler import ( AzureSpeechPassthroughLoggingHandler, ) @@ -13,7 +15,24 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( SHORT_AUDIO_URL = "https://eastus.stt.speech.microsoft.com/speech/recognition/conversation/cognitiveservices/v1?language=en-US" BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions" -TRANSCRIPT = '{"RecognitionStatus":"Success","DisplayText":"Hello world."}' +TRANSCRIPT_BODY = {"RecognitionStatus": "Success", "Offset": 5000000, "Duration": 25000000, "DisplayText": "Hello world."} +TRANSCRIPT = json.dumps(TRANSCRIPT_BODY) +TRANSCRIPT_AUDIO_SECONDS = 3.0 +PRICE_PER_SECOND = 0.5 + + +@pytest.fixture(autouse=True) +def azure_stt_price(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setitem( + litellm.model_cost, + "azure/speech/azure-stt", + { + "litellm_provider": "azure", + "mode": "audio_transcription", + "input_cost_per_second": PRICE_PER_SECOND, + "output_cost_per_second": 0.0, + }, + ) def _make_response(url: str) -> httpx.Response: @@ -30,18 +49,19 @@ def _make_logging_obj() -> MagicMock: class TestAzureSpeechPassthroughHandler: @pytest.mark.parametrize( - "url_route,expected_model", + "url_route,expected_model,expected_cost", [ - (SHORT_AUDIO_URL, "azure_speech/short-audio"), - (BATCH_URL, "azure_speech/batch-transcription"), - (f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription"), + (SHORT_AUDIO_URL, "azure_speech/short-audio", TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND), + (BATCH_URL, "azure_speech/batch-transcription", 0.0), + (f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription", 0.0), ], ) - def test_records_model_provider_and_zero_cost(self, url_route: str, expected_model: str): + def test_records_model_provider_and_cost(self, url_route: str, expected_model: str, expected_cost: float): logging_obj = _make_logging_obj() handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( httpx_response=_make_response(url_route), + response_body=TRANSCRIPT_BODY, logging_obj=logging_obj, url_route=url_route, result=TRANSCRIPT, @@ -54,16 +74,65 @@ class TestAzureSpeechPassthroughHandler: assert handler_result["result"] == {"response": TRANSCRIPT} assert handler_result["kwargs"]["model"] == expected_model assert handler_result["kwargs"]["custom_llm_provider"] == "azure_speech" - assert handler_result["kwargs"]["response_cost"] == 0.0 - assert handler_result["kwargs"]["standard_logging_object"]["response_cost"] == 0.0 + assert handler_result["kwargs"]["response_cost"] == pytest.approx(expected_cost) + assert handler_result["kwargs"]["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost) assert handler_result["kwargs"]["standard_logging_object"]["model"] == expected_model assert logging_obj.model_call_details["model"] == expected_model assert logging_obj.model_call_details["custom_llm_provider"] == "azure_speech" - assert logging_obj.model_call_details["response_cost"] == 0.0 + assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost) + + @pytest.mark.parametrize( + "response_body", + [ + {"RecognitionStatus": "NoMatch", "Offset": 0, "Duration": 0}, + {"RecognitionStatus": "InitialSilenceTimeout"}, + {"Offset": "5000000", "Duration": "25000000"}, + {}, + [], + None, + ], + ) + def test_short_audio_without_recognized_duration_logs_zero_cost( + self, response_body: dict[str, object] | list[dict[str, object]] | None + ): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(SHORT_AUDIO_URL), + response_body=response_body, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["model"] == "azure_speech/short-audio" + assert handler_result["kwargs"]["response_cost"] == 0.0 + + def test_missing_price_entry_still_logs_the_row_at_zero_cost(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delitem(litellm.model_cost, "azure/speech/azure-stt") + + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(SHORT_AUDIO_URL), + response_body=TRANSCRIPT_BODY, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["model"] == "azure_speech/short-audio" + assert handler_result["kwargs"]["custom_llm_provider"] == "azure_speech" + assert handler_result["kwargs"]["response_cost"] == 0.0 def test_subscription_key_never_reaches_the_logging_payload(self): handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( httpx_response=_make_response(SHORT_AUDIO_URL), + response_body=TRANSCRIPT_BODY, logging_obj=_make_logging_obj(), url_route=SHORT_AUDIO_URL, result=TRANSCRIPT, @@ -106,18 +175,18 @@ class TestNormalizeDispatch: def test_normalize_routes_to_azure_speech_handler(self): normalized = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload( httpx_response=_make_response(SHORT_AUDIO_URL), - response_body={"RecognitionStatus": "Success"}, + response_body=TRANSCRIPT_BODY, request_body={}, logging_obj=_make_logging_obj(), url_route=SHORT_AUDIO_URL, - result=TRANSCRIPT, + result="", start_time=datetime.now(), end_time=datetime.now(), cache_hit=False, custom_llm_provider="azure_speech", ) - assert normalized["standard_logging_response_object"] == {"response": TRANSCRIPT} + assert normalized["standard_logging_response_object"] == {"response": ""} assert normalized["kwargs"]["model"] == "azure_speech/short-audio" assert normalized["kwargs"]["custom_llm_provider"] == "azure_speech" - assert normalized["kwargs"]["response_cost"] == 0.0 + assert normalized["kwargs"]["response_cost"] == pytest.approx(TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 9fe7f5b6ee9..a1102543ae8 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -6278,6 +6278,49 @@ class TestAzureSpeechProxyRoute: ("azure_speech/batch-transcription", "azure_speech", 0.0) ] + def test_short_audio_spend_is_priced_from_the_recognized_duration( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[dict[str, object]] = [] # mutable-ok: test recorder accumulates callback payloads + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + self.payloads.append(kwargs["standard_logging_object"]) + + recorder: Final = _Recorder() + monkeypatch.setattr(litellm, "_async_success_callback", [*litellm._async_success_callback, recorder]) + monkeypatch.setitem( + litellm.model_cost, + "azure/speech/azure-stt", + { + "litellm_provider": "azure", + "mode": "audio_transcription", + "input_cost_per_second": 0.25, + "output_cost_per_second": 0.0, + }, + ) + transcript: Final = {**AZURE_SPEECH_TRANSCRIPT, "Offset": 10_000_000, "Duration": 30_000_000} + with respx.mock(assert_all_called=True) as upstream: + upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json=transcript) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + content=AZURE_SPEECH_WAV_BYTES, + headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 200 + assert [(p["model"], p["custom_llm_provider"]) for p in recorder.payloads] == [ + ("azure_speech/short-audio", "azure_speech") + ] + assert recorder.payloads[0]["response_cost"] == pytest.approx(4.0 * 0.25) + def test_api_base_wins_over_region_for_both_families( self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch ) -> None: From 585c32d3f500d3788cf9a47928bc5922351d835b Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 03:53:28 +0000 Subject: [PATCH 085/253] refactor(deepgram): move listen frame parsing into llms/deepgram and drop routine docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 63 +++++++++- .../llm_passthrough_endpoints.py | 8 -- ...gram_listen_passthrough_logging_handler.py | 71 +---------- .../pass_through_endpoints.py | 8 -- .../deepgram/test_deepgram_common_utils.py | 114 ++++++++++++++++++ ...gram_listen_passthrough_logging_handler.py | 51 -------- 6 files changed, 179 insertions(+), 136 deletions(-) create mode 100644 tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index db00d048f01..f1759f94775 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -1,5 +1,8 @@ +import math +from collections.abc import Mapping, Sequence from types import MappingProxyType from typing import Final +from urllib.parse import parse_qs, urlparse import httpx @@ -14,10 +17,6 @@ class DeepgramException(BaseLLMException): def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> str: - """ - The upstream ``/listen`` socket for a streaming transcription, keeping the client's query string as sent - and adding the default model only when the client named none - """ listen_url: Final = httpx.URL(f"{(api_base or DEEPGRAM_DEFAULT_API_BASE).rstrip('/')}/listen") websocket_url: Final = listen_url.copy_with(scheme=_WEBSOCKET_SCHEMES.get(listen_url.scheme, listen_url.scheme)) params: Final = httpx.QueryParams(query_string) @@ -25,3 +24,59 @@ def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> query_string if params.get("model") else str(params.remove("model").add("model", DEEPGRAM_LISTEN_DEFAULT_MODEL)) ) return f"{websocket_url}?{query}" + + +def deepgram_listen_model(upstream_url: str) -> str: + models: Final = parse_qs(urlparse(upstream_url).query).get("model") + return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL + + +def _seconds(value: object) -> float | None: + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) if math.isfinite(value) and value >= 0 else None + + +def _results_frame_end(frame: Mapping[str, object]) -> float | None: + start: Final = _seconds(frame.get("start")) + duration: Final = _seconds(frame.get("duration")) + return None if start is None or duration is None else start + duration + + +def _final_transcript(frame: Mapping[str, object]) -> str | None: + if frame.get("is_final") is not True: + return None + channel: Final = frame.get("channel") + alternatives: Final = channel.get("alternatives") if isinstance(channel, Mapping) else None + first: Final = alternatives[0] if isinstance(alternatives, list) and alternatives else None + transcript: Final = first.get("transcript") if isinstance(first, Mapping) else None + return transcript if isinstance(transcript, str) and transcript else None + + +def deepgram_listen_audio_seconds(websocket_messages: Sequence[Mapping[str, object]]) -> float: + metadata_durations: Final = tuple( + duration + for frame in websocket_messages + if frame.get("type") == "Metadata" + if (duration := _seconds(frame.get("duration"))) is not None + ) + if metadata_durations: + return metadata_durations[-1] + return max( + ( + end + for frame in websocket_messages + if frame.get("type") == "Results" + if (end := _results_frame_end(frame)) is not None + ), + default=0.0, + ) + + +def deepgram_listen_transcript(websocket_messages: Sequence[Mapping[str, object]]) -> str: + return " ".join( + transcript + for frame in websocket_messages + if frame.get("type") == "Results" + if (transcript := _final_transcript(frame)) is not None + ) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index a5a95c34910..1abbf90cb7a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -2613,10 +2613,6 @@ def _proxy_model_allowlists() -> _OpenAIWebsocketModelAllowlists: def _negotiated_websocket_subprotocol(websocket: WebSocket) -> str | None: - """ - The first subprotocol the client offered, echoed back so browsers that carry the LiteLLM key in - ``Sec-WebSocket-Protocol`` complete the handshake - """ requested_subprotocols: Final = tuple( protocol.strip() for protocol in (websocket.headers.get("sec-websocket-protocol") or "").split(",") @@ -2707,10 +2703,6 @@ async def deepgram_listen_websocket_route( user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)], ) -> None: - """ - Streaming speech to text through Deepgram's ``/v1/listen`` socket. Audio frames and transcript frames are - relayed unchanged; the call is billed on the audio duration Deepgram reports when the socket closes - """ deepgram_api_key: Final = passthrough_endpoint_router.get_credentials( custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value, region_name=None, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py index 6a574a2e1b9..a5fea7c8020 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -1,81 +1,22 @@ -""" -Cost tracking for Deepgram's streaming ``/v1/listen`` WebSocket. Deepgram bills the audio it processed, which it -reports as ``duration`` on the closing ``Metadata`` frame; a stream that ends without one is billed on the furthest -``start + duration`` across its ``Results`` frames -""" - -import math from collections.abc import Mapping, Sequence from types import MappingProxyType from typing import Final -from urllib.parse import parse_qs, urlparse +from urllib.parse import urlparse import litellm from litellm._logging import verbose_proxy_logger -from litellm.constants import DEEPGRAM_LISTEN_DEFAULT_MODEL from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.deepgram.common_utils import ( + deepgram_listen_audio_seconds, + deepgram_listen_model, + deepgram_listen_transcript, +) from litellm.proxy._types import PassThroughEndpointLoggingTypedDict from litellm.types.utils import TranscriptionResponse DEEPGRAM_LISTEN_ROUTE_SUFFIX: Final = "/listen" -def _seconds(value: object) -> float | None: - if isinstance(value, bool) or not isinstance(value, (int, float)): - return None - return float(value) if math.isfinite(value) and value >= 0 else None - - -def _results_frame_end(frame: Mapping[str, object]) -> float | None: - start: Final = _seconds(frame.get("start")) - duration: Final = _seconds(frame.get("duration")) - return None if start is None or duration is None else start + duration - - -def _final_transcript(frame: Mapping[str, object]) -> str | None: - if frame.get("is_final") is not True: - return None - channel: Final = frame.get("channel") - alternatives: Final = channel.get("alternatives") if isinstance(channel, Mapping) else None - first: Final = alternatives[0] if isinstance(alternatives, list) and alternatives else None - transcript: Final = first.get("transcript") if isinstance(first, Mapping) else None - return transcript if isinstance(transcript, str) and transcript else None - - -def deepgram_listen_audio_seconds(websocket_messages: Sequence[Mapping[str, object]]) -> float: - metadata_durations: Final = tuple( - duration - for frame in websocket_messages - if frame.get("type") == "Metadata" - if (duration := _seconds(frame.get("duration"))) is not None - ) - if metadata_durations: - return metadata_durations[-1] - return max( - ( - end - for frame in websocket_messages - if frame.get("type") == "Results" - if (end := _results_frame_end(frame)) is not None - ), - default=0.0, - ) - - -def deepgram_listen_transcript(websocket_messages: Sequence[Mapping[str, object]]) -> str: - return " ".join( - transcript - for frame in websocket_messages - if frame.get("type") == "Results" - if (transcript := _final_transcript(frame)) is not None - ) - - -def deepgram_listen_model(upstream_url: str) -> str: - models: Final = parse_qs(urlparse(upstream_url).query).get("model") - return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL - - def _audio_cost(response: TranscriptionResponse, model: str) -> float | None: try: return litellm.completion_cost( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index cf6985852f2..449ae48b0ef 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -2121,9 +2121,6 @@ def _resolved_vertex_live_setup( def _json_object_frame(frame: str | bytes) -> dict[str, object] | None: - """ - The frame as a JSON object when it is one, for cost tracking; audio and non-object frames yield None - """ try: decoded: Final = json.loads(frame if isinstance(frame, str) else frame.decode("utf-8")) except (json.JSONDecodeError, UnicodeDecodeError): @@ -2431,10 +2428,6 @@ async def websocket_passthrough_request( json_frame_ordinal: Final = count() async def relay_upstream_frame(upstream_message: str | bytes) -> None: - """ - Send the frame to the client exactly as received, then keep it for cost tracking when it is a JSON - object; the Vertex AI Live setup acknowledgement only names the model, so it is read instead of kept - """ if isinstance(upstream_message, bytes): await websocket.send_bytes(upstream_message) else: @@ -2448,7 +2441,6 @@ async def websocket_passthrough_request( websocket_messages.append(message_data) async def forward_upstream_to_client() -> Close | None: - """Relay upstream frames to the client until the upstream closes, returning its close frame""" try: while True: await relay_upstream_frame(await upstream_ws.recv()) diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py new file mode 100644 index 00000000000..a86cb83d628 --- /dev/null +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -0,0 +1,114 @@ +import math +from collections.abc import Mapping, Sequence +from typing import Final + +import pytest + +import litellm +from litellm.llms.deepgram.common_utils import ( + deepgram_listen_audio_seconds, + deepgram_listen_model, + deepgram_listen_transcript, + deepgram_listen_websocket_target, +) + +NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000" + + +def _results(start: object, duration: object, transcript: str = "", is_final: object = True) -> dict[str, object]: + return { + "type": "Results", + "start": start, + "duration": duration, + "is_final": is_final, + "channel": {"alternatives": [{"transcript": transcript, "confidence": 0.9}]}, + } + + +def _metadata(duration: object) -> dict[str, object]: + return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": 1} + + +@pytest.mark.parametrize( + ("api_base", "query_string", "expected"), + [ + pytest.param( + None, + "model=nova-3&encoding=linear16", + "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16", + id="default", + ), + pytest.param( + None, + "encoding=linear16&sample_rate=16000", + "wss://api.deepgram.com/v1/listen?encoding=linear16&sample_rate=16000&model=nova-3", + id="model added when missing", + ), + pytest.param( + None, + "model=&encoding=linear16", + "wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-3", + id="empty model replaced", + ), + pytest.param( + "http://localhost:9000/v1/", + "model=nova-2", + "ws://localhost:9000/v1/listen?model=nova-2", + id="custom base becomes ws", + ), + pytest.param( + "wss://dg.internal/v1", + "model=nova-3&keywords=a&keywords=b", + "wss://dg.internal/v1/listen?model=nova-3&keywords=a&keywords=b", + id="repeated keys preserved", + ), + ], +) +def test_deepgram_listen_websocket_target(api_base: str | None, query_string: str, expected: str): + assert deepgram_listen_websocket_target(api_base=api_base, query_string=query_string) == expected + + +@pytest.mark.parametrize( + ("frames", "expected_seconds"), + [ + pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _metadata(6.25)), 6.25, id="metadata wins"), + pytest.param((_metadata(4.0), _results(0.0, 9.0), _metadata(5.5)), 5.5, id="last metadata wins"), + pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _results(1.0, 1.0)), 5.5, id="furthest results end"), + pytest.param((_results(0.0, 0.0), _metadata(0.0)), 0.0, id="zero metadata is a real zero"), + pytest.param((_results(0.0, 1.5), _metadata("6.25")), 1.5, id="string metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(True)), 1.5, id="boolean metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(-3.0)), 1.5, id="negative metadata is ignored"), + pytest.param((_results(0.0, 1.5), _metadata(math.nan), _metadata(math.inf)), 1.5, id="nan/inf ignored"), + pytest.param((_results("0", 2.0), _results(0.0, None), _results(0.0, 0.75)), 0.75, id="malformed results"), + pytest.param(({"type": "SpeechStarted", "timestamp": 3.0}, {"type": "UtteranceEnd"}), 0.0, id="no usage"), + pytest.param((), 0.0, id="no frames"), + ], +) +def test_deepgram_listen_audio_seconds(frames: Sequence[Mapping[str, object]], expected_seconds: float): + assert deepgram_listen_audio_seconds(frames) == expected_seconds + + +def test_deepgram_listen_transcript_joins_final_results_only(): + frames = ( + _results(0.0, 1.0, "hello wor", is_final=False), + _results(0.0, 1.5, "hello world"), + _results(1.5, 0.5, "", is_final=True), + _results(2.0, 1.0, "how are you", is_final="yes"), + {"type": "Results", "start": 3.0, "duration": 1.0, "is_final": True, "channel": {"alternatives": []}}, + _results(4.0, 1.0, "goodbye"), + _metadata(5.0), + ) + assert deepgram_listen_transcript(frames) == "hello world goodbye" + + +@pytest.mark.parametrize( + ("upstream_url", "expected_model"), + [ + (NOVA_3_URL, "nova-3"), + ("wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-2-medical", "nova-2-medical"), + ("wss://api.deepgram.com/v1/listen?model=nova-3&model=nova-2", "nova-3"), + ("wss://api.deepgram.com/v1/listen?encoding=linear16", litellm.constants.DEEPGRAM_LISTEN_DEFAULT_MODEL), + ], +) +def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, expected_model: str): + assert deepgram_listen_model(upstream_url) == expected_model diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py index f00411ac10c..40f520e344d 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -1,7 +1,5 @@ """Deepgram ``/v1/listen`` WebSocket passthrough: duration extraction and duration based cost tracking.""" -import math -from collections.abc import Mapping, Sequence from datetime import datetime from types import SimpleNamespace from typing import Final @@ -14,9 +12,6 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.llm_provider_handlers.deepgram_listen_passthrough_logging_handler import ( DeepgramListenPassthroughLoggingHandler, - deepgram_listen_audio_seconds, - deepgram_listen_model, - deepgram_listen_transcript, ) from litellm.proxy.pass_through_endpoints.success_handler import PassThroughEndpointLogging from litellm.types.passthrough_endpoints.pass_through_endpoints import PassthroughStandardLoggingPayload @@ -39,52 +34,6 @@ def _metadata(duration: object) -> dict[str, object]: return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": 1} -@pytest.mark.parametrize( - ("frames", "expected_seconds"), - [ - pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _metadata(6.25)), 6.25, id="metadata wins"), - pytest.param((_metadata(4.0), _results(0.0, 9.0), _metadata(5.5)), 5.5, id="last metadata wins"), - pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _results(1.0, 1.0)), 5.5, id="furthest results end"), - pytest.param((_results(0.0, 0.0), _metadata(0.0)), 0.0, id="zero metadata is a real zero"), - pytest.param((_results(0.0, 1.5), _metadata("6.25")), 1.5, id="string metadata is ignored"), - pytest.param((_results(0.0, 1.5), _metadata(True)), 1.5, id="boolean metadata is ignored"), - pytest.param((_results(0.0, 1.5), _metadata(-3.0)), 1.5, id="negative metadata is ignored"), - pytest.param((_results(0.0, 1.5), _metadata(math.nan), _metadata(math.inf)), 1.5, id="nan/inf ignored"), - pytest.param((_results("0", 2.0), _results(0.0, None), _results(0.0, 0.75)), 0.75, id="malformed results"), - pytest.param(({"type": "SpeechStarted", "timestamp": 3.0}, {"type": "UtteranceEnd"}), 0.0, id="no usage"), - pytest.param((), 0.0, id="no frames"), - ], -) -def test_deepgram_listen_audio_seconds(frames: Sequence[Mapping[str, object]], expected_seconds: float): - assert deepgram_listen_audio_seconds(frames) == expected_seconds - - -def test_deepgram_listen_transcript_joins_final_results_only(): - frames = ( - _results(0.0, 1.0, "hello wor", is_final=False), - _results(0.0, 1.5, "hello world"), - _results(1.5, 0.5, "", is_final=True), - _results(2.0, 1.0, "how are you", is_final="yes"), - {"type": "Results", "start": 3.0, "duration": 1.0, "is_final": True, "channel": {"alternatives": []}}, - _results(4.0, 1.0, "goodbye"), - _metadata(5.0), - ) - assert deepgram_listen_transcript(frames) == "hello world goodbye" - - -@pytest.mark.parametrize( - ("upstream_url", "expected_model"), - [ - (NOVA_3_URL, "nova-3"), - ("wss://api.deepgram.com/v1/listen?encoding=linear16&model=nova-2-medical", "nova-2-medical"), - ("wss://api.deepgram.com/v1/listen?model=nova-3&model=nova-2", "nova-3"), - ("wss://api.deepgram.com/v1/listen?encoding=linear16", litellm.constants.DEEPGRAM_LISTEN_DEFAULT_MODEL), - ], -) -def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, expected_model: str): - assert deepgram_listen_model(upstream_url) == expected_model - - @pytest.mark.parametrize( ("url_route", "expected"), [ From 6e56ba86c5ed3851f8ef4b4b309c1e85949606f9 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 04:03:22 +0000 Subject: [PATCH 086/253] test(proxy): clear leaked auth dependency override before Azure Speech real-auth tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../pass_through_endpoints/test_llm_pass_through_endpoints.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index a1102543ae8..05ef44b89b6 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -6423,6 +6423,7 @@ class TestAzureSpeechRawBodyThroughRealAuth: ) -> httpx.Response: from litellm.proxy.proxy_server import app + monkeypatch.delitem(app.dependency_overrides, user_api_key_auth, raising=False) monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") monkeypatch.setenv("AZURE_SPEECH_REGION", "eastus") monkeypatch.delenv("AZURE_SPEECH_API_BASE", raising=False) From 1d8f19e4fde7615dc1828ebf7b4bf1e0d510af85 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 04:14:26 +0000 Subject: [PATCH 087/253] fix(proxy): keep Azure Speech multipart bodies intact through auth user_api_key_auth called request.form() on multipart Azure Speech batch uploads, consuming the Starlette stream before the pass-through handler could read the raw bytes. The opaque body predicate now covers multipart on the whole /azure_speech prefix so auth caches an empty parsed body and the upload is forwarded byte for byte Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/common_utils/http_parsing_utils.py | 9 +-- .../test_llm_pass_through_endpoints.py | 63 ++++++++++++++++--- 2 files changed, 61 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 29dc36f3dba..1c17c46e5af 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -11,7 +11,6 @@ from typing_extensions import NotRequired, ReadOnly, Required from litellm._logging import verbose_proxy_logger from litellm.constants import ( AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX, - AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_REQUEST_BODY_SIZE_TO_REPAIR_MB, ) @@ -220,9 +219,11 @@ async def _read_request_body(request: Request | None) -> dict: def is_opaque_audio_pass_through_request(route: str, content_type: str) -> bool: - return route.startswith( - f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}{AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX}" - ) and _normalize_media_type(content_type).startswith("audio/") + """Azure Speech bodies (raw audio, multipart uploads) are forwarded byte for byte, so auth must not consume them.""" + media_type: Final = _normalize_media_type(content_type) + return route.startswith(f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}/") and ( + media_type.startswith("audio/") or media_type == "multipart/form-data" + ) async def read_raw_json_body(request: Request | None) -> bytes | None: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 05ef44b89b6..43fc28c34c4 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -6418,8 +6418,8 @@ def _azure_speech_real_auth_attrs() -> dict[str, object]: class TestAzureSpeechRawBodyThroughRealAuth: """user_api_key_auth reads the body before the route runs; raw audio must not be parsed as JSON.""" - def _post_wav( - self, monkeypatch: pytest.MonkeyPatch, path: str, api_key: str, body: bytes = AZURE_SPEECH_WAV_BYTES + def _post( + self, monkeypatch: pytest.MonkeyPatch, path: str, api_key: str, content_type: str, body: bytes ) -> httpx.Response: from litellm.proxy.proxy_server import app @@ -6437,9 +6437,14 @@ class TestAzureSpeechRawBodyThroughRealAuth: path, params={"language": "en-US"}, content=body, - headers={"Content-Type": "audio/wav", "Authorization": f"Bearer {api_key}"}, + headers={"Content-Type": content_type, "Authorization": f"Bearer {api_key}"}, ) + def _post_wav( + self, monkeypatch: pytest.MonkeyPatch, path: str, api_key: str, body: bytes = AZURE_SPEECH_WAV_BYTES + ) -> httpx.Response: + return self._post(monkeypatch, path, api_key, "audio/wav", body) + @pytest.mark.parametrize("body", [AZURE_SPEECH_WAV_BYTES, AZURE_SPEECH_NON_UTF8_WAV_BYTES], ids=["ascii", "binary"]) def test_master_key_with_raw_wav_body_reaches_azure_without_a_parse_attempt( self, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture, body: bytes @@ -6466,11 +6471,55 @@ class TestAzureSpeechRawBodyThroughRealAuth: assert response.status_code in (400, 401), response.text assert not catch_all.called - @pytest.mark.parametrize("path", ["/v1/chat/completions", f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}"]) - def test_audio_content_type_off_the_short_audio_route_is_still_parsed_as_json( - self, monkeypatch: pytest.MonkeyPatch, path: str + def test_master_key_with_multipart_batch_upload_is_forwarded_byte_for_byte( + self, monkeypatch: pytest.MonkeyPatch ) -> None: - response = self._post_wav(monkeypatch, path, "sk-master-key", body=b'{}{"model": "gpt-4o"}') + boundary: Final = "lit7939boundary" + multipart_body: Final = ( + f"--{boundary}\r\nContent-Disposition: form-data; name=\"definition\"\r\n\r\n".encode() + + json.dumps({"locales": ["en-US"]}).encode() + + f"\r\n--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"eagle.wav\"\r\n" + "Content-Type: audio/wav\r\n\r\n".encode() + + AZURE_SPEECH_NON_UTF8_WAV_BYTES + + f"\r\n--{boundary}--\r\n".encode() + ) + with respx.mock(assert_all_called=True) as upstream: + route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( + return_value=httpx.Response(201, json={"status": "NotStarted"}) + ) + + response = self._post( + monkeypatch, + f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", + "sk-master-key", + f"multipart/form-data; boundary={boundary}", + multipart_body, + ) + + assert (response.status_code, response.json()) == (201, {"status": "NotStarted"}) + sent = route.calls.last.request + assert sent.content == multipart_body + assert sent.headers["content-type"] == f"multipart/form-data; boundary={boundary}" + assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" + + @pytest.mark.parametrize("content_type", ["audio/wav", "multipart/form-data; boundary=x"]) + def test_wrong_litellm_key_with_multipart_batch_upload_is_rejected( + self, monkeypatch: pytest.MonkeyPatch, content_type: str + ) -> None: + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(200)) + + response = self._post( + monkeypatch, f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", "sk-wrong", content_type, b"--x--\r\n" + ) + + assert response.status_code in (400, 401), response.text + assert not catch_all.called + + def test_audio_content_type_off_the_azure_speech_route_is_still_parsed_as_json( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + response = self._post_wav(monkeypatch, "/v1/chat/completions", "sk-master-key", body=b'{}{"model": "gpt-4o"}') assert response.status_code == 400 assert "Invalid JSON payload" in response.text From 0986f404f8bc189854a9a7d88dfd4af376c84566 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 06:08:12 +0000 Subject: [PATCH 088/253] fix(proxy): match IPv4-mapped IPv6 peers against IPv4 trusted proxy ranges Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/network.py | 8 +++++++- tests/test_litellm/proxy/auth/test_network.py | 7 +++++++ 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/network.py b/litellm/proxy/auth/network.py index 32ad18d4deb..28466156100 100644 --- a/litellm/proxy/auth/network.py +++ b/litellm/proxy/auth/network.py @@ -49,11 +49,17 @@ def parse_trusted_proxy_ranges( return networks +def _unmapped(addr: ipaddress.IPv4Address | ipaddress.IPv6Address) -> ipaddress.IPv4Address | ipaddress.IPv6Address: + if isinstance(addr, ipaddress.IPv6Address) and addr.ipv4_mapped is not None: + return addr.ipv4_mapped + return addr + + def ip_in_networks(client_ip: str | None, networks: list[TrustedProxyNetwork]) -> bool: if not client_ip or not networks: return False try: - addr: Final = ipaddress.ip_address(client_ip.strip()) + addr: Final = _unmapped(ipaddress.ip_address(client_ip.strip())) except ValueError: return False return any(addr in network for network in networks) diff --git a/tests/test_litellm/proxy/auth/test_network.py b/tests/test_litellm/proxy/auth/test_network.py index b67723305e4..83233dc370d 100644 --- a/tests/test_litellm/proxy/auth/test_network.py +++ b/tests/test_litellm/proxy/auth/test_network.py @@ -57,6 +57,13 @@ def test_xff_honored_from_trusted_peer(): assert via_proxy is True +def test_ipv4_mapped_peer_and_hop_match_ipv4_trusted_ranges(): + request = make_request(headers={"x-forwarded-for": "203.0.113.9, ::ffff:10.0.0.5"}, client=("::ffff:10.0.0.1", 1)) + ip, via_proxy = resolve_client_ip(request, TRUSTED) + assert ip == "203.0.113.9" + assert via_proxy is True + + def test_spoofed_xff_from_untrusted_peer_is_ignored(): request = make_request( headers={"x-forwarded-for": "203.0.113.9"}, client=("8.8.8.8", 1) From c05095373d101f5192a4703788c940b5925bc2aa Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 06:41:23 +0000 Subject: [PATCH 089/253] fix(proxy): keep mapped-notation trusted proxy ranges matching mapped peers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/network.py | 5 +++-- tests/test_litellm/proxy/auth/test_network.py | 8 ++++++++ 2 files changed, 11 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/network.py b/litellm/proxy/auth/network.py index 28466156100..8a20e207113 100644 --- a/litellm/proxy/auth/network.py +++ b/litellm/proxy/auth/network.py @@ -59,10 +59,11 @@ def ip_in_networks(client_ip: str | None, networks: list[TrustedProxyNetwork]) - if not client_ip or not networks: return False try: - addr: Final = _unmapped(ipaddress.ip_address(client_ip.strip())) + addr: Final = ipaddress.ip_address(client_ip.strip()) except ValueError: return False - return any(addr in network for network in networks) + candidates: Final = (addr, _unmapped(addr)) + return any(candidate in network for candidate in candidates for network in networks) def _is_valid_ip(value: str) -> bool: diff --git a/tests/test_litellm/proxy/auth/test_network.py b/tests/test_litellm/proxy/auth/test_network.py index 83233dc370d..e743ce8cd23 100644 --- a/tests/test_litellm/proxy/auth/test_network.py +++ b/tests/test_litellm/proxy/auth/test_network.py @@ -64,6 +64,14 @@ def test_ipv4_mapped_peer_and_hop_match_ipv4_trusted_ranges(): assert via_proxy is True +def test_ipv4_mapped_peer_still_matches_mapped_notation_trusted_range(): + config = TrustedProxyConfig(use_forwarded_for=True, trusted_proxy_cidrs=["::ffff:10.0.0.0/104"]) + request = make_request(headers={"x-forwarded-for": "203.0.113.9"}, client=("::ffff:10.0.0.1", 1)) + ip, via_proxy = resolve_client_ip(request, config) + assert ip == "203.0.113.9" + assert via_proxy is True + + def test_spoofed_xff_from_untrusted_peer_is_ignored(): request = make_request( headers={"x-forwarded-for": "203.0.113.9"}, client=("8.8.8.8", 1) From 9fa85c5da5a6a9497a0df12cbc02866497af353d Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 08:39:01 +0000 Subject: [PATCH 090/253] fix(proxy): name the blocking guardrail in x-litellm-applied-guardrails When a guardrail hook raises, the common ProxyLogging dispatch (sequential and parallel pre_call, pipeline block, during_call and post_call metrics wrapper, streaming iterator wrapper) now records that guardrail in applied_guardrails before re-raising, and pre_call_hook folds request-declared guardrails in on its raising path. Buffered streams rebuild their response headers after the first chunk so a post_call block reached while buffering carries the blocker too Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 11 ++-- litellm/proxy/utils.py | 58 +++++++++++++++---- .../proxy/test_common_request_processing.py | 57 ++++++++++++++++++ .../proxy_logging/test_during_call_hook.py | 18 ++++++ .../proxy_logging/test_guardrail_pipeline.py | 8 ++- .../test_post_call_success_hook.py | 24 ++++++++ .../utils/proxy_logging/test_pre_call_hook.py | 40 +++++++++++++ .../proxy_logging/test_streaming_hooks.py | 38 +++++++++++- 8 files changed, 234 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 2f39e6c71bc..9a038ca79ba 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2576,11 +2576,14 @@ class ProxyBaseLLMRequestProcessing: ) async def refresh_stream_headers() -> Mapping[str, str]: - """`custom_headers` rebuilt for whichever deployment served the stream.""" - if not getattr(response, "fallback_headers_adopted", False): - return custom_headers + """`custom_headers` rebuilt once the first chunk is buffered, from `self.data` as the + guardrails left it and for whichever deployment served the stream.""" return self._stream_response_headers( - hidden_params=get_hidden_params_dict(response), + hidden_params=( + get_hidden_params_dict(response) + if getattr(response, "fallback_headers_adopted", False) + else hidden_params + ), user_api_key_dict=user_api_key_dict, logging_obj=logging_obj, version=version, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8225fef3492..fea6ce20b61 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -123,6 +123,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header from litellm.proxy.common_utils.config_sync_pubsub import publish_config_param_change from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.db.create_views import ( @@ -437,6 +438,12 @@ def _enrich_http_exception_with_guardrail_context(exc: BaseException, callback: detail.setdefault("guardrail_mode", event_hook) +def _record_raising_guardrail(request_data: Mapping[str, object], callback: object) -> None: + guardrail_name: Final[object] = getattr(callback, "guardrail_name", None) + if isinstance(request_data, dict) and isinstance(guardrail_name, str): + add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=guardrail_name) + + def _is_client_error_exception(exc: Exception) -> bool: if isinstance(exc, HTTPException): return exc.status_code < 500 @@ -1795,13 +1802,19 @@ class ProxyLogging: ) if expected_if_unmutated is not None: callback.mark_pre_call_hook_ran(expected_if_unmutated) - result: Final = await self._process_guardrail_callback( - callback=callback, - data=input_data, - user_api_key_dict=user_api_key_dict, - call_type=call_type, - event_type=GuardrailEventHooks.pre_call, - ) + try: + result: Final = await self._process_guardrail_callback( + callback=callback, + data=input_data, + user_api_key_dict=user_api_key_dict, + call_type=call_type, + event_type=GuardrailEventHooks.pre_call, + ) + except SensitiveDataRouteException: + raise + except Exception: + _record_raising_guardrail(data, callback) + raise if ( scans_raw_request and expected_if_unmutated is not None @@ -2031,6 +2044,7 @@ class ProxyLogging: callback: Final = PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name) if callback is not None: _enrich_http_exception_with_guardrail_context(original_exception, callback) + _record_raising_guardrail(data, callback) raise original_exception step_results_serializable: Final = [ @@ -2296,8 +2310,10 @@ class ProxyLogging: if data is not None: self._process_guardrail_metadata(data) return data - except Exception as e: - raise e + except Exception: + if data is not None: + self._process_guardrail_metadata(data) + raise async def _run_parallel_pre_call_guardrails( self, @@ -2355,6 +2371,8 @@ class ProxyLogging: # live kwargs. if callback.scan_raw_request and not isinstance(result, BaseException) and result is not None: callback.mark_pre_call_hook_ran(data) + if isinstance(result, BaseException) and not isinstance(result, SensitiveDataRouteException): + _record_raising_guardrail(data, callback) raised: Final = tuple(result for result in results if isinstance(result, BaseException)) blocking: Final = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None) if blocking is not None: @@ -2433,7 +2451,12 @@ class ProxyLogging: break @staticmethod - async def _run_guardrail_with_metrics(callback: object, coro: Awaitable[_T], hook_type: str) -> _T: + async def _run_guardrail_with_metrics( + callback: object, + coro: Awaitable[_T], + hook_type: str, + request_data: Mapping[str, object], + ) -> _T: """ Await `coro`, recording its latency and status to the `litellm_guardrail_latency_seconds` metric under `hook_type`, and @@ -2453,6 +2476,7 @@ class ProxyLogging: status = "error" error_type = type(e).__name__ _enrich_http_exception_with_guardrail_context(e, callback) + _record_raising_guardrail(request_data, callback) raise finally: ProxyLogging._emit_guardrail_metrics( @@ -2465,7 +2489,9 @@ class ProxyLogging: @staticmethod async def _wrap_streaming_iterator_with_enrichment( - callback: object, gen: AsyncGenerator[_T, None] + callback: object, + gen: AsyncGenerator[_T, None], + request_data: Mapping[str, object], ) -> AsyncGenerator[_T, None]: """ Yield from `gen`; if iteration raises an HTTPException with dict detail, @@ -2480,6 +2506,7 @@ class ProxyLogging: yield chunk except Exception as e: _enrich_http_exception_with_guardrail_context(e, callback) + _record_raising_guardrail(request_data, callback) raise # Cache for callback-capability detection. Keyed on a signature of @@ -2714,6 +2741,7 @@ class ProxyLogging: call_type=call_type, ), "during_call", + request_data=data, ) return await self._run_guardrail_with_metrics( @@ -2724,6 +2752,7 @@ class ProxyLogging: call_type=call_type, ), "during_call", + request_data=data, ) async def failed_tracking_alert( @@ -3242,6 +3271,7 @@ class ProxyLogging: response=response, ), "post_call", + request_data=data, ) else: guardrail_response = await self._run_guardrail_with_metrics( @@ -3252,6 +3282,7 @@ class ProxyLogging: response=response, ), "post_call", + request_data=data, ) if guardrail_response is not None: @@ -3315,6 +3346,7 @@ class ProxyLogging: response=response, ), "post_call", + request_data=data, ) else: await self._run_guardrail_with_metrics( @@ -3325,6 +3357,7 @@ class ProxyLogging: response=response, ), "post_call", + request_data=data, ) results: Final = await asyncio.gather( @@ -3388,6 +3421,7 @@ class ProxyLogging: request_data=request_data, ), "post_mcp_call", + request_data=request_data, ) return response @@ -3637,6 +3671,7 @@ class ProxyLogging: response=current_response, request_data=request_data, ), + request_data=request_data, ) else: # kind == "apply_guardrail": route through unified_guardrail @@ -3649,6 +3684,7 @@ class ProxyLogging: guardrail_to_apply=resolved_callback, buffer_until_moderated_default=(kind == "override"), ), + request_data=request_data, ) pipeline_translation: Final = ( diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 4ac687625c2..028ea29c093 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -40,9 +40,12 @@ from litellm.proxy.common_request_processing import ( _parse_event_data_for_error, _resolve_per_request_model_group_alias, _should_return_raw_model_name, + _sse_error_frames, _UpstreamClosingStreamingResponse, create_response, + sse_error_payload, ) +from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyErrorTypes, ProxyException @@ -8952,6 +8955,60 @@ class TestStreamingResponseHeadersFollowFallback: assert "llm_provider-stale-marker" not in result.headers assert result.headers["x-callback-header"] == "kept" + @pytest.mark.asyncio + async def test_streaming_block_headers_name_the_blocking_guardrail(self, monkeypatch): + processor_data: dict[str, object] = {"model": "oa", "stream": True, "metadata": {}} + + def select_data_generator(**kwargs): + async def generator(): + add_guardrail_to_applied_guardrails_header(processor_data, "stream-blocker") + _, error_obj = sse_error_payload(HTTPException(status_code=400, detail="blocked")) + for frame in _sse_error_frames(error_obj): + yield frame + + return generator() + + logging_obj = MagicMock() + logging_obj.litellm_call_id = "lit-7144-call" + logging_obj._defer_async_logging = False + logging_obj._on_deferred_stream_complete = None + logging_obj.cost_breakdown = None + processor_data["litellm_logging_obj"] = logging_obj + processor = ProxyBaseLLMRequestProcessing(data=processor_data) + + proxy_logging_obj = MagicMock(spec=ProxyLogging) + proxy_logging_obj.during_call_hook = AsyncMock(return_value=None) + proxy_logging_obj.update_request_status = AsyncMock(return_value=None) + proxy_logging_obj.post_call_success_hook = AsyncMock( + side_effect=lambda data, user_api_key_dict, response: response + ) + proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) + + async def fake_route_request(**kwargs): + async def call(): + return SimpleNamespace(_hidden_params={}, fallback_headers_adopted=False) + + return call() + + monkeypatch.setattr(litellm.proxy.common_request_processing, "route_request", fake_route_request) + + result = await processor.base_process_llm_request( + request=Request(scope={"type": "http", "headers": []}), + fastapi_response=Response(), + user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"), + route_type="acompletion", + proxy_logging_obj=proxy_logging_obj, + general_settings={}, + proxy_config=MagicMock(spec=ProxyConfig), + select_data_generator=select_data_generator, + is_streaming_request=True, + skip_pre_call_logic=True, + ) + + assert isinstance(result, JSONResponse) + assert result.status_code == 400 + assert result.headers["x-litellm-applied-guardrails"] == "stream-blocker" + class _MessagesFallbackStream: def __init__(self) -> None: diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py index 3c5d879c2dc..46f39ef6fb7 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_during_call_hook.py @@ -6,6 +6,7 @@ from typing import Any, Dict from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi import HTTPException import litellm from litellm.integrations.custom_guardrail import CustomGuardrail @@ -84,3 +85,20 @@ async def test_during_call_hook_guardrail_error_raises(proxy_logging, make_user_ user_api_key_dict=make_user_api_key_auth(), call_type="completion", ) + + +@pytest.mark.asyncio +async def test_during_call_block_names_the_blocking_guardrail_in_applied_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + g = _make_guardrail("blocker") + g.async_moderation_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked")) + monkeypatch.setattr(litellm, "callbacks", [_make_guardrail("passer"), g]) + data = {"model": "m", "metadata": {}} + with pytest.raises(HTTPException): + await proxy_logging.during_call_hook( + data=data, + user_api_key_dict=make_user_api_key_auth(), + call_type="completion", + ) + assert "blocker" in data["metadata"]["applied_guardrails"] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 077bf5a313e..cd2b7a278bb 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -527,17 +527,19 @@ def test_handle_pipeline_result_block_enriches_with_guardrail_name_and_mode(): result.step_results = [MagicMock(guardrail_name="g")] result.original_exception = original + data: dict[str, object] = {"model": "m"} saved = litellm.callbacks litellm.callbacks = [cb] try: with pytest.raises(HTTPException) as info: - ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p") + ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") finally: litellm.callbacks = saved assert info.value is original assert info.value.detail["guardrail_name"] == "g" assert info.value.detail["guardrail_mode"] == GuardrailEventHooks.pre_call + assert data["metadata"] == {"applied_guardrails": ["g"]} def test_handle_pipeline_result_block_does_not_reraise_sensitive_data_route(): @@ -617,7 +619,7 @@ async def test_run_guardrail_with_metrics_passes_result_and_records_success(monk monkeypatch.setattr(litellm, "callbacks", [prom]) out = await ProxyLogging._run_guardrail_with_metrics( - callback=MagicMock(guardrail_name="g"), coro=task(), hook_type="during_call" + callback=MagicMock(guardrail_name="g"), coro=task(), hook_type="during_call", request_data={} ) assert out == {"a": 1, "b": 2, "c": 3} @@ -643,7 +645,7 @@ async def test_run_guardrail_with_metrics_records_error_and_enriches(monkeypatch monkeypatch.setattr(litellm, "callbacks", [prom]) with pytest.raises(HTTPException): - await ProxyLogging._run_guardrail_with_metrics(callback=cb, coro=task(), hook_type="post_call") + await ProxyLogging._run_guardrail_with_metrics(callback=cb, coro=task(), hook_type="post_call", request_data={}) assert detail["guardrail_name"] == "presidio" recorded = prom._record_guardrail_metrics.call_args.kwargs diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py index 715d66db181..53d8948869f 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_post_call_success_hook.py @@ -6,6 +6,7 @@ from typing import Any from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi import HTTPException import litellm from litellm.integrations.custom_guardrail import CustomGuardrail @@ -96,3 +97,26 @@ async def test_post_call_success_hook_guardrail_returns_modified_response( data={}, response={"orig": True}, user_api_key_dict=make_user_api_key_auth() ) assert out == modified + + +@pytest.mark.asyncio +@pytest.mark.parametrize("run_in_parallel", [False, True], ids=["sequential", "parallel"]) +async def test_post_call_block_names_the_blocking_guardrail_in_applied_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch, run_in_parallel +): + def _passer_that_records(data, user_api_key_dict, response): + data["metadata"]["applied_guardrails"] = ["passer"] + + passer = _make_guardrail("passer") + passer.async_post_call_success_hook = AsyncMock(side_effect=_passer_that_records) + passer.run_in_parallel = run_in_parallel + blocker = _make_guardrail("blocker") + blocker.async_post_call_success_hook = AsyncMock(side_effect=HTTPException(status_code=400, detail="blocked")) + blocker.run_in_parallel = run_in_parallel + monkeypatch.setattr(litellm, "callbacks", [passer, blocker]) + data = {"model": "m", "metadata": {}} + with pytest.raises(HTTPException): + await proxy_logging.post_call_success_hook( + data=data, response=MagicMock(), user_api_key_dict=make_user_api_key_auth() + ) + assert data["metadata"]["applied_guardrails"] == ["passer", "blocker"] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py index 6e5cb7fcae3..dbc6fba4ab1 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_pre_call_hook.py @@ -905,3 +905,43 @@ async def test_scan_raw_request_warns_on_in_place_mutation_returning_none( ) mock_logger.warning.assert_called_once() assert "scan_raw_request" in str(mock_logger.warning.call_args) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "blocker_kwargs", + [ + pytest.param({}, id="sequential"), + pytest.param({"scan_raw_request": True}, id="scan_raw_request"), + pytest.param({"run_in_parallel": True}, id="parallel"), + ], +) +async def test_pre_call_block_names_the_blocking_guardrail_in_applied_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch, blocker_kwargs +): + monkeypatch.setattr(litellm, "callbacks", [_BlockOnSecretGuardrail(**blocker_kwargs)]) + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + data = _secret_request() + with pytest.raises(HTTPException, match="blocked"): + await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data=data, + call_type="completion", + ) + assert data["metadata"]["applied_guardrails"] == ["blocker"] + + +@pytest.mark.asyncio +async def test_pre_call_block_keeps_request_declared_guardrail_in_applied_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + monkeypatch.setattr(litellm, "callbacks", [_BlockOnSecretGuardrail(default_on=False)]) + proxy_logging.slack_alerting_instance = MagicMock(alerting=None) + data = {**_secret_request(), "metadata": {"guardrails": ["blocker", "declared-post-call"]}} + with pytest.raises(HTTPException, match="blocked"): + await proxy_logging.pre_call_hook( + user_api_key_dict=make_user_api_key_auth(), + data=data, + call_type="completion", + ) + assert data["metadata"]["applied_guardrails"] == ["blocker", "declared-post-call"] diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index ebc831b4102..52586ed2174 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -20,6 +20,7 @@ from fastapi import HTTPException import litellm from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( @@ -27,6 +28,7 @@ from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterato ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import Usage @@ -175,7 +177,7 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro yield ch cb = MagicMock(guardrail_name="g", event_hook="pre_call") - wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen()) + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen(), request_data={}) out = [ch async for ch in wrapped] snapshot = { "chunks": out, @@ -201,7 +203,7 @@ async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_r raise HTTPException(status_code=400, detail=detail) cb = MagicMock(guardrail_name="presidio", event_hook="post_call") - wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen()) + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen(), request_data={}) with pytest.raises(HTTPException): async for _ in wrapped: pass @@ -696,3 +698,35 @@ async def test_post_call_response_headers_hook_swallows_callback_error(proxy_log data={}, user_api_key_dict=make_user_api_key_auth(), response=response ) assert out == {} + + +@pytest.mark.asyncio +async def test_stream_guardrail_block_names_the_blocking_guardrail_in_applied_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _StreamBlocker(CustomGuardrail): + def __init__(self) -> None: + super().__init__(guardrail_name="stream-blocker", event_hook=GuardrailEventHooks.post_call, default_on=True) + + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object] + ) -> AsyncGenerator[object, None]: + async for _ in response: + raise HTTPException(status_code=400, detail={"error": "blocked"}) + yield # pragma: no cover + + monkeypatch.setattr(litellm, "callbacks", [_StreamBlocker()]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + + async def upstream(): + yield "chunk" + + request_data: dict[str, object] = {"metadata": {}} + with pytest.raises(HTTPException): + async for _ in proxy_logging.async_post_call_streaming_iterator_hook( + response=upstream(), + user_api_key_dict=make_user_api_key_auth(), + request_data=request_data, + ): + pass + assert request_data["metadata"]["applied_guardrails"] == ["stream-blocker"] From 9a365d2021a0c6a24be4989b4aca4bb898f60e45 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 16:10:23 +0000 Subject: [PATCH 091/253] feat(proxy): hard-block throttled Admin UI sign-ins with no credential bypass A blocked source, or source and username pair, is now refused with 429 before the database lookup and password check, in place of the soft block that held wrong guesses for 30 seconds and let a correct password through. The env admin credentials and the master key typed into the login form are refused like any other credential while blocked; recovery is the master key as an API bearer token, which never goes through the sign-in path trusted_proxy_ranges: [] now means clients connect directly, so the peer address is the source and the per-source limit stays on. Only an unset or malformed value leaves the topology unknown, warns at startup and turns the per-source limit off Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 6 +- litellm/proxy/auth/login_throttle.py | 100 +++--- litellm/proxy/auth/login_utils.py | 18 +- litellm/proxy/auth/network.py | 3 +- litellm/proxy/proxy_server.py | 4 +- .../proxy/auth/test_login_utils.py | 333 ++++++++---------- .../proxy/proxy_server/conftest.py | 7 - .../proxy_server/test_routes_login_sso.py | 59 +++- tests/test_litellm/proxy/test_proxy_server.py | 16 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 +- 10 files changed, 273 insertions(+), 279 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 08380d8b537..80147863304 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2755,7 +2755,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): max_failed_login_attempts_per_source: int | None = Field( None, ge=1, - description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set, since otherwise every client behind an ingress shares one address. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", + description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and this limit is off. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", ) max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field( None, @@ -2774,7 +2774,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): failed_login_block_seconds: int | None = Field( None, ge=1, - description="How long a blocked source address, or source address and username, stays blocked. The block is soft: a correct password still signs in, but each attempt from a blocked key takes one of 5 held slots per worker and a wrong password is held for 30 seconds before it is refused with 429. Set under `general_settings` in config.yaml. Defaults to 300", + description="How long a blocked source address, or source address and username, stays blocked. Every attempt from a blocked key, right or wrong, is refused with 429 before the password is checked; the block is not extended by refused attempts. Set under `general_settings` in config.yaml. Defaults to 300", ) allowed_routes: list | None = Field(None, description="Proxy API Endpoints you want users to be able to access") reject_clientside_metadata_tags: bool | None = Field( @@ -2885,7 +2885,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) trusted_proxy_ranges: list[str] | None = Field( None, - description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.", + description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, the per-source sign-in limit is off.", ) store_model_in_db: bool | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index a8899b9c675..0106d58e426 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -1,10 +1,10 @@ """Failed-login accounting for the Admin UI sign-in path. Wrong passwords are counted over a short window per source address and per source-and-username -pair; too many in one window blocks that key for a fixed time. Blocks are soft: a correct password -still signs in, while a wrong one from a blocked key is held open before its 429 and only a few can -be held at once, which bounds how many guesses a blocked key gets checked. A blocked pair stops +pair; too many in one window blocks that key for a fixed time. While a key is blocked every attempt +from it, right or wrong, is refused with 429 before the password is checked. A blocked pair stops counting against its source, so one script stuck on one account does not block the whole office. +Recovery is the master key over the API, which never passes through here, or waiting out the block. """ from __future__ import annotations @@ -14,8 +14,7 @@ import hashlib import ipaddress import math import time -from collections.abc import AsyncGenerator, Mapping -from contextlib import asynccontextmanager +from collections.abc import Mapping from dataclasses import dataclass from functools import cache from types import MappingProxyType @@ -37,8 +36,6 @@ DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER: Final = 5 DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60 DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300 -BLOCKED_ATTEMPT_HOLD_SECONDS: Final = 30 -MAX_HELD_ATTEMPTS_PER_KEY: Final = 5 IPV6_SOURCE_PREFIX_LENGTH: Final = 64 SOURCE_LIMIT_KEY: Final = "max_failed_login_attempts_per_source" @@ -100,11 +97,6 @@ _COUNTERS: Final = InMemoryCache( max_size_in_memory=_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS ) _BLOCKS: Final = InMemoryCache(max_size_in_memory=_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS) -_HELD_ATTEMPTS: Final[dict[str, int]] = {} # mutable-ok: in-flight hold counts rise on entry and fall on exit - - -async def _sleep(seconds: float) -> None: - await asyncio.sleep(seconds) @cache @@ -127,12 +119,25 @@ def warn_login_counters_are_per_worker(num_workers: str) -> None: def warn_source_login_limit_is_off() -> None: verbose_proxy_logger.warning( "%s is not set, so failed Admin UI sign-in attempts are limited per source address and username " - "only. Set it to the address ranges of the proxies in front of LiteLLM to also limit each " - "source address across usernames.", + "only. Set it to the address ranges of the proxies in front of LiteLLM, or to an empty list when " + "clients connect directly, to also limit each source address across usernames.", TRUSTED_PROXY_RANGES_KEY, ) +def declared_proxy_ranges(settings: Mapping[str, object]) -> tuple[str, ...] | None: + """What the operator says fronts LiteLLM: the proxy ranges, an empty tuple for none, None when unsaid. + + Only a declared topology makes the source address trustworthy enough to limit across usernames. + An unset key, or a value that is not a list of ranges, leaves it unknown and the source scope off. + """ + raw_ranges: Final = settings.get(TRUSTED_PROXY_RANGES_KEY) + if isinstance(raw_ranges, (list, tuple, set)) and not raw_ranges: + return () + cidrs: Final = tuple(normalize_cidr_ranges(raw_ranges, setting_name=TRUSTED_PROXY_RANGES_KEY)) + return cidrs or None + + def _positive_int(raw: object, key: str, default: int) -> int: if raw is None: return default @@ -221,8 +226,11 @@ class Block: @dataclass(frozen=True, slots=True) class LoginThrottle: - """Failed-login limits for one request's source address; ``source_limit`` is None when the - source scope is off because ``trusted_proxy_ranges`` is unset and the peer address is the ingress.""" + """Failed-login limits for one request's source address. + + ``source_limit`` is None when the source scope is off: ``trusted_proxy_ranges`` is unset, so the peer + address may be a shared ingress. An empty list means clients connect directly and the peer is the source. + """ client_ip: str source_limit: int | None @@ -242,15 +250,13 @@ class LoginThrottle: redis_cache: RedisCache | None, ) -> LoginThrottle: settings: Final = general_settings if general_settings is not None else _NO_SETTINGS - cidrs: Final = normalize_cidr_ranges( - settings.get(TRUSTED_PROXY_RANGES_KEY), setting_name=TRUSTED_PROXY_RANGES_KEY - ) + proxies: Final = declared_proxy_ranges(settings) resolved, _ = resolve_client_ip( - request, TrustedProxyConfig(use_forwarded_for=bool(cidrs), trusted_proxy_cidrs=cidrs) + request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ()) ) return cls( client_ip=resolved or _UNKNOWN_SOURCE, - source_limit=_source_limit(settings, resolved) if cidrs and resolved is not None else None, + source_limit=_source_limit(settings, resolved) if proxies is not None and resolved is not None else None, user_limit=_int_setting(settings, USER_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER), window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS), block_seconds=_int_setting(settings, BLOCK_KEY, DEFAULT_FAILED_LOGIN_BLOCK_SECONDS), @@ -270,18 +276,21 @@ class LoginThrottle: source_block=f"{_CACHE_KEY_PREFIX}:{{{group}}}:block:source", ) - @asynccontextmanager - async def attempt(self, username: str, *, exempt: bool = False) -> AsyncGenerator[LoginAttempt]: - if not self.enabled or exempt: - yield LoginAttempt(throttle=self, username=username, block=None) - return - keys: Final = self._keys(username) - block: Final = await self._active_block(keys) + async def attempt(self, username: str) -> LoginAttempt: + """Refuses a blocked key before any credential is looked at; otherwise hands back the attempt to settle.""" + if not self.enabled: + return LoginAttempt(throttle=self, username=username) + block: Final = await self._active_block(self._keys(username)) if block is None: - yield LoginAttempt(throttle=self, username=username, block=None) - return - slot: Final = keys.pair_block if block.scope == "user" else keys.source_block - yield LoginAttempt(throttle=self, username=username, block=block, slot=slot) + return LoginAttempt(throttle=self, username=username) + verbose_proxy_logger.warning( + "Admin UI sign-in refused: the %s is blocked for %s more seconds; username=%r source=%s", + block.scope, + block.retry_after, + username, + self.client_ip, + ) + self.refuse(block.retry_after) async def _active_block(self, keys: _Keys) -> Block | None: local: Final = self._local_block_ttls(keys) @@ -372,8 +381,6 @@ class LoginThrottle: class LoginAttempt: throttle: LoginThrottle username: str - block: Block | None - slot: str | None = None async def succeeded(self) -> None: if not self.throttle.enabled: @@ -383,8 +390,6 @@ class LoginAttempt: async def failed(self) -> None: if not self.throttle.enabled: return - if self.block is not None and self.slot is not None: - await self._hold_then_refuse(self.block, self.slot) user_block, source_block = await self.throttle.record_failure(self.username) if user_block == 0 and source_block == 0: return @@ -395,26 +400,3 @@ class LoginAttempt: self.username, self.throttle.client_ip, ) - - async def _hold_then_refuse(self, block: Block, slot: str) -> NoReturn: - held: Final = _HELD_ATTEMPTS.get(slot, 0) - if held >= MAX_HELD_ATTEMPTS_PER_KEY: - verbose_proxy_logger.warning( - "Admin UI sign-in refused at once: %s wrong attempts already held for a blocked %s; " - "username=%r source=%s", - held, - block.scope, - self.username, - self.throttle.client_ip, - ) - self.throttle.refuse(block.retry_after) - _HELD_ATTEMPTS[slot] = held + 1 - try: - await _sleep(BLOCKED_ATTEMPT_HOLD_SECONDS) - finally: - remaining: Final = _HELD_ATTEMPTS.get(slot, 1) - 1 - if remaining > 0: - _HELD_ATTEMPTS[slot] = remaining - else: - _HELD_ATTEMPTS.pop(slot, None) - self.throttle.refuse(max(block.retry_after - BLOCKED_ATTEMPT_HOLD_SECONDS, 1)) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 6f52babb255..e0d599b0017 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -186,9 +186,11 @@ async def authenticate_user( or if username/password login is disabled while SSO is configured Recovery: an admin locked out of the UI by - `disable_password_login_when_sso_enabled` can still administer the proxy over - the API with the master key (Authorization: Bearer ), which never - goes through this function. To restore UI username/password login, unset the + `disable_password_login_when_sso_enabled`, or by the failed sign-in block in + `throttle`, can still administer the proxy over the API with the master key + (Authorization: Bearer ), which never goes through this function. + No credential, the env admin credentials and the master key included, is + exempt from the block. To restore UI username/password login, unset the setting in config.yaml (or the DB-persisted general_settings) and restart the proxy; this is a deliberate, auditable config change rather than a hidden bypass. @@ -217,12 +219,8 @@ async def authenticate_user( code=500, ) - admin_credentials_match: Final = _admin_credentials_match(username, password, master_key, general_settings) - - async with throttle.attempt(username, exempt=admin_credentials_match) as attempt: - return await _sign_in( - username, password, master_key, prisma_client, attempt, general_settings, admin_credentials_match - ) + attempt: Final = await throttle.attempt(username) + return await _sign_in(username, password, master_key, prisma_client, attempt, general_settings) async def _sign_in( @@ -232,8 +230,8 @@ async def _sign_in( prisma_client: PrismaClient | None, attempt: LoginAttempt, general_settings: Mapping[str, object], - admin_credentials_match: bool, ) -> LoginResult: + admin_credentials_match: Final = _admin_credentials_match(username, password, master_key, general_settings) # Check if we can find the `username` in the db. On the UI, users can enter username=their email _user_row: LiteLLM_UserTable | None = None user_role: ( diff --git a/litellm/proxy/auth/network.py b/litellm/proxy/auth/network.py index 8a20e207113..4e8ab7512a7 100644 --- a/litellm/proxy/auth/network.py +++ b/litellm/proxy/auth/network.py @@ -1,6 +1,7 @@ from __future__ import annotations import ipaddress +from collections.abc import Sequence from typing import Any, Final from fastapi import Request @@ -19,7 +20,7 @@ class NetworkContext(BaseModel): class TrustedProxyConfig(BaseModel): use_forwarded_for: bool = False - trusted_proxy_cidrs: list[str] = Field(default_factory=list) + trusted_proxy_cidrs: Sequence[str] = Field(default_factory=tuple) def normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str = "trusted_proxy_cidrs") -> list[str]: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d00c695d968..4075ac4aff7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -328,8 +328,8 @@ from litellm.proxy.auth.fallback_model_access import router_fallback_access_chec from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import AUTO_ROUTER_LICENSE_REMEDY, LicenseCheck from litellm.proxy.auth.login_throttle import ( - TRUSTED_PROXY_RANGES_KEY, LoginThrottle, + declared_proxy_ranges, warn_login_counters_are_per_worker, warn_source_login_limit_is_off, ) @@ -5819,7 +5819,7 @@ class ProxyConfig: if os.getenv("NUM_WORKERS", "1") != "1" and redis_usage_cache is None: warn_login_counters_are_per_worker(os.getenv("NUM_WORKERS", "1")) - if not general_settings.get(TRUSTED_PROXY_RANGES_KEY): + if declared_proxy_ranges(general_settings) is None: warn_source_login_limit_is_off() _bg_hc_model_groups: Final = parse_background_health_check_model_groups(general_settings) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index d1fdb2d6a70..8670d7fef41 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -13,28 +13,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -class _RecordedSleeps: - """A sleep that records what it was asked to wait for instead of waiting.""" - - def __init__(self): - self.seconds: list[float] = [] - - async def __call__(self, seconds: float) -> None: - self.seconds.append(seconds) - - -@pytest.fixture(autouse=True) -def login_delays(monkeypatch): - """Replace the hold on a blocked wrong password, so the suite pays no wall clock and can read it back.""" - from litellm.proxy.auth import login_throttle - - recorded = _RecordedSleeps() - monkeypatch.setattr(login_throttle, "_sleep", recorded) - login_throttle._HELD_ATTEMPTS.clear() - yield recorded - login_throttle._HELD_ATTEMPTS.clear() - - def _unlimited_throttle(): """A throttle wired to real in-memory stores with limits no test can reach.""" from litellm.caching.in_memory_cache import InMemoryCache @@ -749,8 +727,8 @@ def _local_count(throttle, key: str) -> int: @pytest.mark.asyncio async def test_too_many_failures_for_one_username_block_that_pair_and_carry_retry_after(monkeypatch): - """One failure past the pair limit blocks the source for that username; the next wrong guess is held - and answered 429 with the block's remaining time, and the counter is not touched by blocked guesses.""" + """One failure past the pair limit blocks the source for that username; the next guess is answered 429 + with the block's remaining time, and the counter is not touched by blocked guesses.""" from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") @@ -764,30 +742,17 @@ async def test_too_many_failures_for_one_username_block_that_pair_and_carry_retr with pytest.raises(ProxyException) as blocked: await _guess(throttle) assert blocked.value.code == "429" - assert blocked.value.headers.get("Retry-After") == "47", "the 30s hold is taken off the remaining block" + assert blocked.value.headers.get("Retry-After") == "77" assert _local_count(throttle, keys.pair_counter) == 3, "a blocked guess is not counted again" @pytest.mark.asyncio -async def test_a_wrong_password_from_a_blocked_key_is_held_before_it_is_refused(monkeypatch, login_delays): - """The hold is the rate cap: a blocked key gets one verified guess per held slot per 30 seconds.""" - from litellm.proxy.auth.login_throttle import BLOCKED_ATTEMPT_HOLD_SECONDS +async def test_a_blocked_key_is_refused_before_the_password_is_looked_at(monkeypatch): + """The block is the rate cap: once a key is blocked, nothing from it reaches the user lookup or the + password check, so a guessing script gets no verification work out of the proxy.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.login_utils import authenticate_user - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - throttle = _throttle(user_limit=1) - - assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] - assert login_delays.seconds == [], "an unblocked wrong password is answered at once" - - assert await _fail(throttle) == "429" - assert login_delays.seconds == [BLOCKED_ATTEMPT_HOLD_SECONDS] - - -@pytest.mark.asyncio -async def test_a_correct_password_signs_in_while_its_pair_is_blocked(monkeypatch): - """The block is soft: the real user is still verified and gets in, so nobody can be locked out by - guessing at their account.""" monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") monkeypatch.setenv("DATABASE_URL", "postgresql://stub") @@ -795,13 +760,50 @@ async def test_a_correct_password_signs_in_while_its_pair_is_blocked(monkeypatch assert [await _fail(throttle, username="user@corp.com") for _ in range(3)] == ["401", "401", "429"] - result = await _db_login(throttle, "user@corp.com", "right", correct=True) - assert result.key == "sk-ui" + lookup = _known_user("user@corp.com") + verify = MagicMock(return_value=True) + with ( + patch( # test-quality-ok: the user lookup is the database boundary; a blocked attempt must not reach it + "litellm.proxy.auth.login_utils.UserRepository", lookup + ), + patch( # test-quality-ok: the password check is the expensive step; a blocked attempt must not reach it + "litellm.proxy.auth.login_utils.verify_password", verify + ), + pytest.raises(ProxyException) as refused, + ): + await authenticate_user( + username="user@corp.com", + password="right", + master_key="sk-master", + prisma_client=MagicMock(), + throttle=throttle, + ) + + assert refused.value.code == "429" + assert lookup.return_value.table.find_first.await_count == 0 + assert verify.call_count == 0 @pytest.mark.asyncio -async def test_a_correct_password_signs_in_while_its_source_is_blocked(monkeypatch): - """Same for the source-wide block: it slows guessing from that address, it does not refuse a user.""" +async def test_a_correct_password_is_refused_while_its_pair_is_blocked(monkeypatch): + """Letting the right password through would give a guesser unlimited tries, so the block is hard: the + real user waits it out, or uses the master key over the API, which never passes through here.""" + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(user_limit=1, block_seconds=90) + + assert [await _fail(throttle, username="user@corp.com") for _ in range(3)] == ["401", "401", "429"] + + with pytest.raises(ProxyException) as refused: + await _db_login(throttle, "user@corp.com", "right", correct=True) + assert refused.value.code == "429" + assert refused.value.headers.get("Retry-After") == "90" + + +@pytest.mark.asyncio +async def test_a_correct_password_is_refused_while_its_source_is_blocked(monkeypatch): + """Same for the source-wide block: every username from that address is refused until it lapses.""" monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") monkeypatch.setenv("DATABASE_URL", "postgresql://stub") @@ -811,8 +813,9 @@ async def test_a_correct_password_signs_in_while_its_source_is_blocked(monkeypat assert await _fail(throttle, username=f"other-{i}@corp.com") == "401" assert await _fail(throttle, username="other-9@corp.com") == "429", "the source is blocked for everyone" - result = await _db_login(throttle, "user@corp.com", "right", correct=True) - assert result.key == "sk-ui" + with pytest.raises(ProxyException) as refused: + await _db_login(throttle, "user@corp.com", "right", correct=True) + assert refused.value.code == "429" @pytest.mark.asyncio @@ -903,6 +906,63 @@ async def test_without_trusted_proxy_ranges_the_source_scope_is_off(monkeypatch) assert [await _fail(throttle, username=f"user-{i}@corp.com") for i in range(6)] == ["401"] * 6 +@pytest.mark.asyncio +async def test_an_empty_trusted_proxy_ranges_means_the_peer_is_the_client_and_the_source_scope_is_on(monkeypatch): + """An explicit empty list says there are no proxies: the peer address is the client, the forwarded header + is ignored, and the source-wide limit applies. Only an unset key means the topology is unknown.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.setenv("UI_PASSWORD", "right") + request = MagicMock() + request.headers = {"x-forwarded-for": "203.0.113.9"} + request.client = MagicMock() + request.client.host = "198.51.100.7" + throttle = LoginThrottle.from_request( + request, + general_settings={"trusted_proxy_ranges": [], "max_failed_login_attempts_per_source": 3}, + redis_cache=None, + ) + + assert throttle.client_ip == "198.51.100.7" + assert throttle.source_limit == 3 + assert [await _fail(throttle, username=f"user-{i}@corp.com") for i in range(4)] == ["401"] * 4 + assert await _fail(throttle, username="user-99@corp.com") == "429", "the spray is stopped by the source limit" + + +@pytest.mark.parametrize("configured", [None, 5, {"10.0.0.0/8": True}, ["", " "]]) +def test_a_trusted_proxy_ranges_value_that_names_no_ranges_leaves_the_topology_unknown(configured): + """Only a real list of ranges or an explicit empty list counts as a declaration; anything else is the same + as unset, so a typo cannot switch the source-wide block on behind a shared ingress.""" + from litellm.proxy.auth.login_throttle import LoginThrottle, declared_proxy_ranges + + settings = {"trusted_proxy_ranges": configured} if configured is not None else {} + assert declared_proxy_ranges(settings) is None + + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "198.51.100.7" + throttle = LoginThrottle.from_request(request, general_settings=settings, redis_cache=None) + assert throttle.source_limit is None + assert throttle.client_ip == "198.51.100.7" + + +def test_declared_proxy_ranges_distinguishes_none_from_empty_from_configured(): + from litellm.proxy.auth.login_throttle import declared_proxy_ranges + + assert declared_proxy_ranges({}) is None + assert declared_proxy_ranges({"trusted_proxy_ranges": []}) == () + assert declared_proxy_ranges({"trusted_proxy_ranges": ["10.0.0.0/8", " 192.168.1.1 "]}) == ( + "10.0.0.0/8", + "192.168.1.1", + ) + assert declared_proxy_ranges({"trusted_proxy_ranges": "10.0.0.0/8,172.16.0.0/12"}) == ( + "10.0.0.0/8", + "172.16.0.0/12", + ) + + @pytest.mark.asyncio async def test_with_trusted_proxy_ranges_the_source_is_the_forwarded_client(monkeypatch): """The header is walked right to left past the trusted hops, so a forged left-most entry cannot pick the bucket.""" @@ -1056,9 +1116,10 @@ async def test_the_block_time_is_fixed_and_not_refreshed_by_blocked_guesses(monk @pytest.mark.asyncio -async def test_the_configured_admin_credentials_bypass_the_throttle(monkeypatch): - """The only account that can fix a misconfiguration is exempt: no hold, no slot, even while blocked.""" - from litellm.proxy.auth import login_throttle as lt +async def test_the_configured_admin_credentials_are_not_exempt_from_the_block(monkeypatch): + """Exempting the env credentials would make them the one password worth guessing without limit, so the + right UI_PASSWORD is refused while its pair is blocked, and signs in normally once the block lapses.""" + from litellm.proxy._types import ProxyException monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") @@ -1075,9 +1136,31 @@ async def test_the_configured_admin_credentials_bypass_the_throttle(monkeypatch) "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ), ): + with pytest.raises(ProxyException) as refused: + await _guess(throttle, password="right") + assert refused.value.code == "429" + + throttle.blocks.delete_cache(throttle._keys("admin").pair_block) result = await _guess(throttle, password="right") assert result.key == "sk-ui" - assert lt._HELD_ATTEMPTS == {} + + +@pytest.mark.asyncio +async def test_the_master_key_used_as_the_ui_password_is_not_exempt_from_the_block(monkeypatch): + """Without UI_PASSWORD the master key doubles as the admin password; it gets no special treatment here + either. Lockout recovery is the master key as a bearer token over the API, which never enters this path.""" + from litellm.proxy._types import ProxyException + + monkeypatch.setenv("UI_USERNAME", "admin") + monkeypatch.delenv("UI_PASSWORD", raising=False) + monkeypatch.setenv("DATABASE_URL", "postgresql://stub") + throttle = _throttle(user_limit=1) + + assert [await _fail(throttle) for _ in range(3)] == ["401", "401", "429"] + + with pytest.raises(ProxyException) as refused: + await _guess(throttle, password="sk-master") + assert refused.value.code == "429" @pytest.mark.asyncio @@ -1180,138 +1263,29 @@ async def test_a_wrong_password_for_a_known_user_also_counts(monkeypatch): @pytest.mark.asyncio -async def test_held_attempts_from_one_blocked_key_are_capped(monkeypatch): - """Holding a wrong guess open must not let one blocked key park unlimited sockets in password checks.""" - import asyncio - +async def test_a_source_block_outranks_a_pair_block_in_the_retry_after(monkeypatch): + """When both scopes are blocked, the answer carries the source block's time, which is the one that + still applies to every other username from that address.""" from litellm.proxy._types import ProxyException - from litellm.proxy.auth import login_throttle as lt - from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - release = asyncio.Event() - - async def _park(_seconds: float) -> None: - await release.wait() - - monkeypatch.setattr(lt, "_sleep", _park) - throttle = _throttle(user_limit=1, client_ip="203.0.113.44") - slot = throttle._keys("admin").pair_block - assert [await _fail(throttle) for _ in range(2)] == ["401", "401"] - - held = [asyncio.create_task(_guess(throttle)) for _ in range(MAX_HELD_ATTEMPTS_PER_KEY)] - for _ in range(1000): - if lt._HELD_ATTEMPTS.get(slot) == MAX_HELD_ATTEMPTS_PER_KEY: - break - await asyncio.sleep(0) - assert lt._HELD_ATTEMPTS.get(slot) == MAX_HELD_ATTEMPTS_PER_KEY - - try: - with pytest.raises(ProxyException) as over_cap: - await _guess(throttle) - assert over_cap.value.code == "429" - assert over_cap.value.headers.get("Retry-After") == "300", "refused at once, for the whole block" - assert await _fail(throttle, username="someone-else@corp.com") == "401", "other keys are not affected" - finally: - release.set() - for task in held: - with pytest.raises(ProxyException): - await task - - assert lt._HELD_ATTEMPTS.get(slot) is None, "the slots are released once the held attempts answer" - - -@pytest.mark.asyncio -async def test_a_full_hold_pool_still_lets_the_right_password_in(monkeypatch): - """Five parked wrong guesses from the office must not turn the soft block into a lockout for the real user.""" - import asyncio - - from litellm.proxy.auth import login_throttle as lt - from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - monkeypatch.setenv("DATABASE_URL", "postgresql://stub") - release = asyncio.Event() - - async def _park(_seconds: float) -> None: - await release.wait() - - monkeypatch.setattr(lt, "_sleep", _park) - throttle = _throttle(user_limit=1, source_limit=3, client_ip="203.0.113.46") - assert [await _fail(throttle, username=f"spray-{i}@corp.com") for i in range(4)] == ["401"] * 4 - source_slot = throttle._keys("known@example.com").source_block - assert throttle._local_block_ttl(source_slot) > 0, "the source is blocked" - - held = [ - asyncio.create_task(_guess(throttle, username="known@example.com")) for _ in range(MAX_HELD_ATTEMPTS_PER_KEY) - ] - for _ in range(1000): - if lt._HELD_ATTEMPTS.get(source_slot) == MAX_HELD_ATTEMPTS_PER_KEY: - break - await asyncio.sleep(0) - assert lt._HELD_ATTEMPTS == {source_slot: MAX_HELD_ATTEMPTS_PER_KEY} - - try: - signed_in = await _db_login(throttle, "known@example.com", "right", correct=True) - assert signed_in.user_id == "u-1" - finally: - release.set() - for task in held: - with pytest.raises(ProxyException): - await task - - -@pytest.mark.asyncio -async def test_a_blocked_source_shares_one_held_slot_pool_across_its_blocked_usernames(monkeypatch): - """Once the source is blocked, a pair block for a username must not hand that username its own five slots.""" - import asyncio - - from litellm.proxy._types import ProxyException - from litellm.proxy.auth import login_throttle as lt - from litellm.proxy.auth.login_throttle import MAX_HELD_ATTEMPTS_PER_KEY - - monkeypatch.setenv("UI_USERNAME", "admin") - monkeypatch.setenv("UI_PASSWORD", "right") - release = asyncio.Event() - - async def _park(_seconds: float) -> None: - await release.wait() - - monkeypatch.setattr(lt, "_sleep", _park) - throttle = _throttle(user_limit=1, source_limit=3, client_ip="203.0.113.45") + throttle = _throttle(user_limit=1, source_limit=3, block_seconds=120, client_ip="203.0.113.45") assert [await _fail(throttle) for _ in range(2)] == ["401", "401"], "the admin pair is now blocked" + throttle.blocks.set_cache(throttle._keys("admin").pair_block, 1, ttl=30) assert [await _fail(throttle, username=f"spray-{i}@corp.com") for i in range(3)] == ["401"] * 3 - source_slot = throttle._keys("admin").source_block - assert throttle._local_block_ttl(source_slot) > 0, "the source is now blocked as well" + assert throttle._local_block_ttl(throttle._keys("admin").source_block) == 120, "the source is now blocked too" - usernames = ["admin", *(f"fresh-{i}@corp.com" for i in range(MAX_HELD_ATTEMPTS_PER_KEY - 1))] - held = [asyncio.create_task(_guess(throttle, username=name)) for name in usernames] - for _ in range(1000): - if lt._HELD_ATTEMPTS.get(source_slot) == MAX_HELD_ATTEMPTS_PER_KEY: - break - await asyncio.sleep(0) - assert lt._HELD_ATTEMPTS == {source_slot: MAX_HELD_ATTEMPTS_PER_KEY} - - try: - for name in ("admin", "fresh-0@corp.com", "never-seen@corp.com"): - with pytest.raises(ProxyException) as over_cap: - await _guess(throttle, username=name) - assert over_cap.value.code == "429" - assert over_cap.value.headers.get("Retry-After") == "300" - finally: - release.set() - for task in held: - with pytest.raises(ProxyException): - await task - - assert lt._HELD_ATTEMPTS == {} + for name in ("admin", "spray-0@corp.com", "never-seen@corp.com"): + with pytest.raises(ProxyException) as refused: + await _guess(throttle, username=name) + assert refused.value.code == "429" + assert refused.value.headers.get("Retry-After") == "120", name @pytest.mark.asyncio -async def test_disabling_the_control_removes_the_hold_as_well(monkeypatch, login_delays): - """The escape hatch has to turn off the whole control, not only the refusal.""" +async def test_disabling_the_control_lets_every_attempt_through(monkeypatch): + """The escape hatch has to turn off the whole control: no counting and no refusal.""" import dataclasses monkeypatch.setenv("UI_USERNAME", "admin") @@ -1319,7 +1293,7 @@ async def test_disabling_the_control_removes_the_hold_as_well(monkeypatch, login throttle = dataclasses.replace(_throttle(user_limit=1), enabled=False) assert [await _fail(throttle) for _ in range(6)] == ["401"] * 6 - assert login_delays.seconds == [] + assert _local_count(throttle, throttle._keys("admin").pair_counter) == 0 class _FakeRedis: @@ -1404,12 +1378,15 @@ async def test_redis_is_the_only_counter_while_it_answers(monkeypatch): assert await _fail(second_worker, username="user@corp.com") == "429", "the second worker sees the block" + block_keys = [k for k in redis.values if ":block:user:" in k] + assert block_keys, "the block lives in Redis, where every worker reads it" + for key in block_keys: + await redis.async_delete_cache(key) await _db_login(second_worker, "user@corp.com", "right", correct=True) assert not [k for k in redis.values if ":user:" in k and ":block:" not in k], ( "success clears the shared pair counter" ) - assert [k for k in redis.values if ":block:user:" in k], "an active block is not lifted by one success" @pytest.mark.asyncio @@ -1436,7 +1413,7 @@ async def test_a_redis_outage_falls_back_to_this_workers_own_counter(monkeypatch verbose_proxy_logger.removeHandler(handler) assert blocked.value.code == "429" - assert blocked.value.headers.get("Retry-After") == "270" + assert blocked.value.headers.get("Retry-After") == "300" assert any("Redis failed while counting Admin UI sign-in attempts" in r.getMessage() for r in records) diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index a349d378985..d1adf2a5c02 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -522,16 +522,9 @@ def reset_login_throttle(monkeypatch): Only the throttle's own keys are removed, so other cache entries remain untouched. """ from litellm.proxy import proxy_server as ps - from litellm.proxy.auth import login_throttle from litellm.proxy.auth.login_throttle import _BLOCKS, _CACHE_KEY_PREFIX, _COUNTERS - async def _no_delay(_seconds: float) -> None: - """The hold on a rejected sign-in from a blocked key, replaced so the route tests stay fast.""" - - monkeypatch.setattr(login_throttle, "_sleep", _no_delay) - def _drop_throttle_keys() -> None: - login_throttle._HELD_ATTEMPTS.clear() for store in (_COUNTERS, _BLOCKS): for key in tuple(store.cache_dict) + tuple(store.ttl_dict): if key.startswith(_CACHE_KEY_PREFIX): diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index 82e86086518..a82f078fcb9 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -553,14 +553,14 @@ def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_logi def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_throttle): - """The 429 tells the caller how long the block has left, after the 30 seconds it was already held.""" + """The 429 tells the caller how long the block has left.""" _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=77) assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] refused = client.post("/v2/login", json={"username": "admin", "password": "wrong"}) assert refused.status_code == 429 - assert refused.headers.get("retry-after") == "47" + assert refused.headers.get("retry-after") == "77" def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, reset_login_throttle): @@ -572,8 +572,8 @@ def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, res refused = client.post("/login", data={"username": "admin", "password": "wrong"}) assert refused.status_code == 429 assert refused.headers.get("content-type", "").startswith("text/html") - assert "Try again in about 47 seconds" in refused.text - assert refused.headers.get("retry-after") == "47" + assert "Try again in about 77 seconds" in refused.text + assert refused.headers.get("retry-after") == "77" def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle): @@ -608,8 +608,29 @@ def test_a_spray_across_usernames_is_not_blocked_without_trusted_proxy_ranges( assert sprayed == [401] * 8 -def test_the_configured_admin_password_still_signs_in_while_blocked(client, monkeypatch, reset_login_throttle): - """The operator must never be locked out of the console by traffic aimed at it.""" +def test_a_spray_across_usernames_is_blocked_on_the_source_with_an_empty_trusted_proxy_ranges( + client, monkeypatch, reset_login_throttle +): + """An explicit empty list says nothing fronts the proxy, so the peer address is the client and the + source scope is on. A forwarded header from an untrusted peer is ignored rather than trusted.""" + _install_real_auth(monkeypatch, trusted_proxy_ranges=[], max_failed_login_attempts_per_source=4) + + sprayed = [ + client.post( + "/v2/login", + json={"username": f"sprayed-{i}@corp.com", "password": "wrong"}, + headers={"x-forwarded-for": f"203.0.113.{i}"}, + ).status_code + for i in range(5) + ] + assert sprayed == [401] * 5 + + assert _json_login(client, "/v2/login", username="sprayed-6@corp.com") == 429 + + +def test_the_configured_admin_password_is_refused_while_blocked(client, monkeypatch, reset_login_throttle): + """The env credentials get no bypass: a bypass would make them the one password worth guessing without + limit. An operator who is blocked administers the proxy with the master key over the API meanwhile.""" from unittest.mock import AsyncMock, patch _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) @@ -625,18 +646,38 @@ def test_the_configured_admin_password_still_signs_in_while_blocked(client, monk "litellm.proxy.auth.login_utils.generate_key_helper_fn", new=AsyncMock(return_value={"token": "sk-ui"}) ), ): + assert _json_login(client, "/v2/login", password="right-password") == 429 + reset_login_throttle() assert _json_login(client, "/v2/login", password="right-password") == 200 -def test_a_database_users_correct_password_signs_in_while_blocked(client, monkeypatch, reset_login_throttle): - """The block is soft: guessing at an account slows the guesser down, it does not lock the owner out.""" +def test_the_master_key_as_a_bearer_token_still_works_while_the_ui_password_is_blocked( + client, monkeypatch, reset_login_throttle +): + """Lockout recovery: the API path with the master key never enters the sign-in throttle.""" _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + + assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] + + assert client.get("/models", headers={"Authorization": "Bearer sk-not-the-master"}).status_code >= 400 + assert client.get("/models", headers={"Authorization": "Bearer sk-test-master"}).status_code == 200 + assert _json_login(client, "/v2/login", password="right-password") == 429, "the UI block is unaffected" + + +def test_a_database_users_correct_password_is_refused_while_blocked(client, monkeypatch, reset_login_throttle): + """The block is hard: while it lasts, nothing from that source signs in as that user, right password or not, + and the block is not extended by the refused attempts.""" + _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=64) _db_user(monkeypatch, "user@corp.com") assert [_json_login(client, "/v2/login", username="user@corp.com") for _ in range(3)] == [401, 401, 429] + refused = client.post("/v2/login", json={"username": "user@corp.com", "password": "right-db-password"}) + assert refused.status_code == 429 + assert refused.headers.get("retry-after") == "64" + + reset_login_throttle() assert _json_login(client, "/v2/login", username="user@corp.com", password="right-db-password") == 200 - assert _json_login(client, "/v2/login", username="user@corp.com") == 429, "the block itself is still in force" def test_sign_in_succeeds_again_once_the_block_is_cleared(client, monkeypatch, reset_login_throttle): diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 80b1d3b1841..4c8cacf120f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3450,7 +3450,8 @@ async def test_load_config_warns_that_the_source_login_limit_is_off_without_trus tmp_path, monkeypatch, caplog ): """The per-source failed-login limit is skipped when the source cannot be attributed, and the - operator must be told so at startup; a configured range silences it.""" + operator must be told so at startup. Both a configured range and an explicit empty list (no + proxies, the peer is the source) silence it, since both keep the limit on.""" import logging from litellm.proxy.auth.login_throttle import warn_source_login_limit_is_off @@ -3465,12 +3466,13 @@ async def test_load_config_warns_that_the_source_login_limit_is_off_without_trus await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) assert "trusted_proxy_ranges is not set" in caplog.text - caplog.clear() - warn_source_login_limit_is_off.cache_clear() - config_file.write_text("model_list: []\ngeneral_settings:\n trusted_proxy_ranges: ['10.0.0.0/8']\n") - with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): - await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) - assert "trusted_proxy_ranges is not set" not in caplog.text + for configured in ("['10.0.0.0/8']", "[]"): + caplog.clear() + warn_source_login_limit_is_off.cache_clear() + config_file.write_text(f"model_list: []\ngeneral_settings:\n trusted_proxy_ranges: {configured}\n") + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file)) + assert "trusted_proxy_ranges is not set" not in caplog.text, configured @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index ee6824ca02c..871cfea2d29 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26512,7 +26512,7 @@ export interface components { enforce_fallback_model_access?: boolean | null; /** * Failed Login Block Seconds - * @description How long a blocked source address, or source address and username, stays blocked. The block is soft: a correct password still signs in, but each attempt from a blocked key takes one of 5 held slots per worker and a wrong password is held for 30 seconds before it is refused with 429. Set under `general_settings` in config.yaml. Defaults to 300 + * @description How long a blocked source address, or source address and username, stays blocked. Every attempt from a blocked key, right or wrong, is refused with 429 before the password is checked; the block is not extended by refused attempts. Set under `general_settings` in config.yaml. Defaults to 300 */ failed_login_block_seconds?: number | null; /** @@ -26566,7 +26566,7 @@ export interface components { max_batch_file_size_mb?: number | null; /** * Max Failed Login Attempts Per Source - * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set, since otherwise every client behind an ingress shares one address. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 + * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and this limit is off. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 */ max_failed_login_attempts_per_source?: number | null; /** @@ -26756,7 +26756,7 @@ export interface components { supported_db_objects?: components["schemas"]["SupportedDBObjectType"][] | null; /** * Trusted Proxy Ranges - * @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler. + * @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, the per-source sign-in limit is off. */ trusted_proxy_ranges?: string[] | null; /** From b227a8c4c99f0e24999323c6f4db5aeff2001b8e Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 16:42:04 +0000 Subject: [PATCH 092/253] refactor(proxy): move login throttle sentinels into constants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 5 +++ litellm/proxy/auth/login_throttle.py | 37 ++++++++++--------- .../proxy/auth/test_login_utils.py | 20 ++++------ .../proxy/proxy_server/conftest.py | 5 ++- 4 files changed, 36 insertions(+), 31 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 338fe0f6b85..a64932a0e4b 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1644,6 +1644,11 @@ LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_INTERVAL_SECONDS: Final = int( LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE: Final = int( os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE", 1000) ) +LOGIN_THROTTLE_CACHE_KEY_PREFIX: Final = "login_fail" +LOGIN_THROTTLE_UNKNOWN_SOURCE: Final = "unknown" +LOGIN_THROTTLE_MAX_TRACKED_COUNTERS: Final = 20_000 +LOGIN_THROTTLE_MAX_TRACKED_BLOCKS: Final = 10_000 +LOGIN_THROTTLE_NOT_BLOCKED: Final = (0, 0) LITELLM_PROXY_ADMIN_NAME: Final = "default_user_id" LITELLM_PROXY_BUDGET_NAME: Final = "litellm-proxy-budget" GLOBAL_PROXY_SPEND_CACHE_KEY: Final = f"{LITELLM_PROXY_ADMIN_NAME}:spend" diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 0106d58e426..d16cd43e36d 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -17,7 +17,6 @@ import time from collections.abc import Mapping from dataclasses import dataclass from functools import cache -from types import MappingProxyType from typing import Final, Literal, NamedTuple, NoReturn, Protocol, TypeAlias from fastapi import Request, status @@ -27,6 +26,14 @@ from redis.exceptions import RedisError from litellm._logging import verbose_proxy_logger from litellm.caching.in_memory_cache import InMemoryCache from litellm.caching.redis_cache import RedisCache, RedisCircuitBreakerOpenError +from litellm.constants import ( + EMPTY_MAPPING, + LOGIN_THROTTLE_CACHE_KEY_PREFIX, + LOGIN_THROTTLE_MAX_TRACKED_BLOCKS, + LOGIN_THROTTLE_MAX_TRACKED_COUNTERS, + LOGIN_THROTTLE_NOT_BLOCKED, + LOGIN_THROTTLE_UNKNOWN_SOURCE, +) from litellm.proxy._types import ProxyErrorTypes, ProxyException from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges, resolve_client_ip from litellm.secret_managers.main import get_secret_bool @@ -45,12 +52,6 @@ WINDOW_KEY: Final = "failed_login_window_seconds" BLOCK_KEY: Final = "failed_login_block_seconds" TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges" -_CACHE_KEY_PREFIX: Final = "login_fail" -_UNKNOWN_SOURCE: Final = "unknown" -_MAX_TRACKED_COUNTERS: Final = 20_000 -_MAX_TRACKED_BLOCKS: Final = 10_000 -_NO_SETTINGS: Final[Mapping[str, object]] = MappingProxyType({}) -_NOT_BLOCKED: Final = (0, 0) _REDIS_FAILURES: Final = (RedisError, RedisCircuitBreakerOpenError, OSError, asyncio.TimeoutError) _LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None) _SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object]) @@ -94,9 +95,11 @@ _RECORD_FAILURE_LUA: Final = ( ) _COUNTERS: Final = InMemoryCache( - max_size_in_memory=_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS + max_size_in_memory=LOGIN_THROTTLE_MAX_TRACKED_COUNTERS, default_ttl=DEFAULT_FAILED_LOGIN_WINDOW_SECONDS +) +_BLOCKS: Final = InMemoryCache( + max_size_in_memory=LOGIN_THROTTLE_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS ) -_BLOCKS: Final = InMemoryCache(max_size_in_memory=_MAX_TRACKED_BLOCKS, default_ttl=DEFAULT_FAILED_LOGIN_BLOCK_SECONDS) @cache @@ -249,13 +252,13 @@ class LoginThrottle: general_settings: Mapping[str, object] | None, redis_cache: RedisCache | None, ) -> LoginThrottle: - settings: Final = general_settings if general_settings is not None else _NO_SETTINGS + settings: Final = general_settings if general_settings is not None else EMPTY_MAPPING proxies: Final = declared_proxy_ranges(settings) resolved, _ = resolve_client_ip( request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ()) ) return cls( - client_ip=resolved or _UNKNOWN_SOURCE, + client_ip=resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE, source_limit=_source_limit(settings, resolved) if proxies is not None and resolved is not None else None, user_limit=_int_setting(settings, USER_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER), window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS), @@ -270,10 +273,10 @@ class LoginThrottle: group: Final = source_group(self.client_ip) user: Final = hashlib.sha256(username.casefold().encode("utf-8")).hexdigest() return _Keys( - pair_counter=f"{_CACHE_KEY_PREFIX}:{{{group}}}:user:{user}", - pair_block=f"{_CACHE_KEY_PREFIX}:{{{group}}}:block:user:{user}", - source_counter=f"{_CACHE_KEY_PREFIX}:{{{group}}}:source", - source_block=f"{_CACHE_KEY_PREFIX}:{{{group}}}:block:source", + pair_counter=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:user:{user}", + pair_block=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:block:user:{user}", + source_counter=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:source", + source_block=f"{LOGIN_THROTTLE_CACHE_KEY_PREFIX}:{{{group}}}:block:source", ) async def attempt(self, username: str) -> LoginAttempt: @@ -305,14 +308,14 @@ class LoginThrottle: async def _shared_block_ttls(self, keys: _Keys) -> _BlockTtls: if self.redis_cache is None: - return _NOT_BLOCKED + return LOGIN_THROTTLE_NOT_BLOCKED try: return _LUA_BLOCK_TTLS.validate_python( await self.redis_cache.async_register_script(_BLOCK_TTLS_LUA)(keys, ()) ) except _REDIS_FAILURES as err: self._warn_redis(err) - return _NOT_BLOCKED + return LOGIN_THROTTLE_NOT_BLOCKED def _local_block_ttls(self, keys: _Keys) -> _BlockTtls: return self._local_block_ttl(keys.pair_block), self._local_block_ttl(keys.source_block) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 8670d7fef41..986760c3cc5 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1361,7 +1361,7 @@ class _DownRedis(_FakeRedis): @pytest.mark.asyncio async def test_redis_is_the_only_counter_while_it_answers(monkeypatch): """Every worker must spend the same budget, see the same block, and a success must clear the pair for all.""" - from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX + from litellm.constants import LOGIN_THROTTLE_CACHE_KEY_PREFIX monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") @@ -1371,7 +1371,7 @@ async def test_redis_is_the_only_counter_while_it_answers(monkeypatch): second_worker = _throttle(user_limit=2, stores=_stores(), redis_cache=redis) assert [await _fail(first_worker, username="user@corp.com") for _ in range(3)] == ["401"] * 3 - assert not [k for k in first_worker.counters.cache_dict if str(k).startswith(_CACHE_KEY_PREFIX)], ( + assert not [k for k in first_worker.counters.cache_dict if str(k).startswith(LOGIN_THROTTLE_CACHE_KEY_PREFIX)], ( "with Redis answering, no worker may keep a counter of its own" ) assert not first_worker.blocks.cache_dict @@ -1437,7 +1437,8 @@ async def test_a_failed_redis_delete_still_clears_this_workers_counter(monkeypat async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): """Regression: throttle entries must not evict cached credentials from user_api_key_cache.""" from litellm.proxy import proxy_server as ps - from litellm.proxy.auth.login_throttle import _CACHE_KEY_PREFIX, LoginThrottle + from litellm.constants import LOGIN_THROTTLE_CACHE_KEY_PREFIX + from litellm.proxy.auth.login_throttle import LoginThrottle monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") @@ -1453,7 +1454,7 @@ async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): assert await _fail(throttle, username=f"made-up-{i}@example.com") == "401" added = set(ps.user_api_key_cache.in_memory_cache.cache_dict) - auth_cache_keys_before - assert not [k for k in added if str(k).startswith(_CACHE_KEY_PREFIX)] + assert not [k for k in added if str(k).startswith(LOGIN_THROTTLE_CACHE_KEY_PREFIX)] def test_settings_that_arrive_as_environment_strings_are_honored(): @@ -1557,17 +1558,12 @@ async def test_a_username_spray_cannot_evict_an_active_block(monkeypatch): """Counters and blocks live in separate bounded stores, so a flood of made-up pairs fills the counter store while the blocks it already earned stay in force.""" from litellm.caching.in_memory_cache import InMemoryCache - from litellm.proxy.auth.login_throttle import ( - _BLOCKS, - _COUNTERS, - _MAX_TRACKED_BLOCKS, - _MAX_TRACKED_COUNTERS, - LoginThrottle, - ) + from litellm.constants import LOGIN_THROTTLE_MAX_TRACKED_BLOCKS, LOGIN_THROTTLE_MAX_TRACKED_COUNTERS + from litellm.proxy.auth.login_throttle import _BLOCKS, _COUNTERS, LoginThrottle monkeypatch.setenv("UI_USERNAME", "admin") monkeypatch.setenv("UI_PASSWORD", "right") - assert _MAX_TRACKED_COUNTERS >= 10_000 and _MAX_TRACKED_BLOCKS >= 10_000 + assert LOGIN_THROTTLE_MAX_TRACKED_COUNTERS >= 10_000 and LOGIN_THROTTLE_MAX_TRACKED_BLOCKS >= 10_000 assert _COUNTERS is not _BLOCKS counters, blocks = InMemoryCache(max_size_in_memory=50), InMemoryCache(max_size_in_memory=50) throttle = LoginThrottle( diff --git a/tests/test_litellm/proxy/proxy_server/conftest.py b/tests/test_litellm/proxy/proxy_server/conftest.py index d1adf2a5c02..ae1b42363ef 100644 --- a/tests/test_litellm/proxy/proxy_server/conftest.py +++ b/tests/test_litellm/proxy/proxy_server/conftest.py @@ -521,13 +521,14 @@ def reset_login_throttle(monkeypatch): window, so without this a failed sign-in test could block unrelated tests later. Only the throttle's own keys are removed, so other cache entries remain untouched. """ + from litellm.constants import LOGIN_THROTTLE_CACHE_KEY_PREFIX from litellm.proxy import proxy_server as ps - from litellm.proxy.auth.login_throttle import _BLOCKS, _CACHE_KEY_PREFIX, _COUNTERS + from litellm.proxy.auth.login_throttle import _BLOCKS, _COUNTERS def _drop_throttle_keys() -> None: for store in (_COUNTERS, _BLOCKS): for key in tuple(store.cache_dict) + tuple(store.ttl_dict): - if key.startswith(_CACHE_KEY_PREFIX): + if key.startswith(LOGIN_THROTTLE_CACHE_KEY_PREFIX): store.delete_cache(key) monkeypatch.setattr(ps, "redis_usage_cache", None) From 5f64dfd8dd4deffdee76c673325590176fc48001 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 18:14:04 +0000 Subject: [PATCH 093/253] fix(proxy): price Azure Speech fast transcription and limit unpriced batch writes to admins Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 4 + .../llm_passthrough_endpoints.py | 22 +++ ...zure_speech_passthrough_logging_handler.py | 30 +++- ...zure_speech_passthrough_logging_handler.py | 39 ++++- .../test_llm_pass_through_endpoints.py | 146 ++++++++++++++++-- 5 files changed, 220 insertions(+), 21 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 15a1d054e26..d9da0cc0f64 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1574,13 +1574,17 @@ AZURE_SPEECH_CUSTOM_LLM_PROVIDER: Final = "azure_speech" AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech" AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX: Final = "/speech/" AZURE_SPEECH_BATCH_PATH_PREFIX: Final = "/speechtotext/" +AZURE_SPEECH_FAST_TRANSCRIPTION_PATH: Final = "/speechtotext/transcriptions:transcribe" +AZURE_SPEECH_UNPRICED_WRITE_METHODS: Final = frozenset({"POST", "PUT"}) AZURE_SPEECH_STT_DOMAIN: Final = "stt.speech.microsoft.com" AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN: Final = "api.cognitive.microsoft.com" AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER: Final = "Ocp-Apim-Subscription-Key" AZURE_SPEECH_SHORT_AUDIO_MODEL: Final = "short-audio" AZURE_SPEECH_BATCH_MODEL: Final = "batch-transcription" +AZURE_SPEECH_FAST_TRANSCRIPTION_MODEL: Final = "fast-transcription" AZURE_SPEECH_PRICING_MODEL: Final = "azure/speech/azure-stt" AZURE_SPEECH_TICKS_PER_SECOND: Final = 10_000_000 +AZURE_SPEECH_MILLISECONDS_PER_SECOND: Final = 1_000 BASE_MCP_ROUTE: Final = "/mcp" diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index e64eac87a7f..b8966508f81 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -33,10 +33,12 @@ from litellm.constants import ( AZURE_SPEECH_BATCH_PATH_PREFIX, AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN, AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + AZURE_SPEECH_FAST_TRANSCRIPTION_PATH, AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX, AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, AZURE_SPEECH_STT_DOMAIN, AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, + AZURE_SPEECH_UNPRICED_WRITE_METHODS, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix @@ -65,6 +67,7 @@ from litellm.proxy.common_utils.http_parsing_utils import ( get_request_body, is_json_content_type, ) +from litellm.proxy.common_utils.resource_ownership import is_proxy_admin from litellm.proxy.common_utils.sse_keepalive import ( wrap_passthrough_sse_bytes_with_keepalive_pings, ) @@ -1357,6 +1360,14 @@ def resolve_azure_speech_base_url(endpoint_path: str, api_base: str | None, regi return httpx.URL(f"https://{region}.{domain}") +def azure_speech_write_is_unpriced(method: str, endpoint_path: str) -> bool: + return ( + endpoint_path.startswith(AZURE_SPEECH_BATCH_PATH_PREFIX) + and endpoint_path != AZURE_SPEECH_FAST_TRANSCRIPTION_PATH + and method.upper() in AZURE_SPEECH_UNPRICED_WRITE_METHODS + ) + + @router.api_route( f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}/{{endpoint:path}}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: fastapi route methods must be a list @@ -1395,6 +1406,17 @@ async def azure_speech_proxy_route( "AZURE_SPEECH_REGION or AZURE_SPEECH_API_BASE in the proxy environment." ), ) + if azure_speech_write_is_unpriced( + method=request.method, endpoint_path=normalized_endpoint_path + ) and not is_proxy_admin(user_api_key_dict): + raise HTTPException( + status_code=403, + detail=( + f"{request.method} {normalized_endpoint_path} creates Azure Speech work whose cost is unknown at " + "request time, so it is limited to proxy admin keys. Use " + f"{AZURE_SPEECH_FAST_TRANSCRIPTION_PATH} for transcription that is priced per request." + ), + ) azure_speech_api_key: Final = passthrough_endpoint_router.get_credentials( custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER, region_name=None, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py index 74587acd453..588b7cc8e56 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py @@ -9,6 +9,9 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import ( AZURE_SPEECH_BATCH_MODEL, AZURE_SPEECH_CUSTOM_LLM_PROVIDER, + AZURE_SPEECH_FAST_TRANSCRIPTION_MODEL, + AZURE_SPEECH_FAST_TRANSCRIPTION_PATH, + AZURE_SPEECH_MILLISECONDS_PER_SECOND, AZURE_SPEECH_PRICING_MODEL, AZURE_SPEECH_SHORT_AUDIO_MODEL, AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, @@ -28,10 +31,16 @@ class AzureSpeechPassthroughLoggingHandler: def _is_short_audio_route(url_route: str) -> bool: return urlparse(url_route).path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX) + @staticmethod + def _is_fast_transcription_route(url_route: str) -> bool: + return urlparse(url_route).path.endswith(AZURE_SPEECH_FAST_TRANSCRIPTION_PATH) + @staticmethod def _model_from_url_route(url_route: str) -> str: if AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_SHORT_AUDIO_MODEL}" + if AzureSpeechPassthroughLoggingHandler._is_fast_transcription_route(url_route): + return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_FAST_TRANSCRIPTION_MODEL}" return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_BATCH_MODEL}" @staticmethod @@ -45,10 +54,25 @@ class AzureSpeechPassthroughLoggingHandler: return (offset + duration) / AZURE_SPEECH_TICKS_PER_SECOND @staticmethod - def _response_cost(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: - if not AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): + def _fast_transcription_audio_seconds(response_body: Mapping[str, object] | Sequence[object] | None) -> float: + if not isinstance(response_body, Mapping): return 0.0 - audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body) + duration_milliseconds: Final = response_body.get("durationMilliseconds") + if not isinstance(duration_milliseconds, int): + return 0.0 + return duration_milliseconds / AZURE_SPEECH_MILLISECONDS_PER_SECOND + + @staticmethod + def _billed_audio_seconds(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: + if AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): + return AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body) + if AzureSpeechPassthroughLoggingHandler._is_fast_transcription_route(url_route): + return AzureSpeechPassthroughLoggingHandler._fast_transcription_audio_seconds(response_body) + return 0.0 + + @staticmethod + def _response_cost(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: + audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._billed_audio_seconds(url_route, response_body) if audio_seconds <= 0.0: return 0.0 try: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py index 5ffb1f7785d..50d3de64f72 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py @@ -13,9 +13,19 @@ from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) -SHORT_AUDIO_URL = "https://eastus.stt.speech.microsoft.com/speech/recognition/conversation/cognitiveservices/v1?language=en-US" +SHORT_AUDIO_URL = ( + "https://eastus.stt.speech.microsoft.com/speech/recognition/conversation/cognitiveservices/v1?language=en-US" +) BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions" -TRANSCRIPT_BODY = {"RecognitionStatus": "Success", "Offset": 5000000, "Duration": 25000000, "DisplayText": "Hello world."} +FAST_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/transcriptions:transcribe?api-version=2024-11-15" +FAST_BODY = {"durationMilliseconds": 5061, "combinedPhrases": [{"text": "Hello world."}]} +FAST_AUDIO_SECONDS = 5.061 +TRANSCRIPT_BODY = { + "RecognitionStatus": "Success", + "Offset": 5000000, + "Duration": 25000000, + "DisplayText": "Hello world.", +} TRANSCRIPT = json.dumps(TRANSCRIPT_BODY) TRANSCRIPT_AUDIO_SECONDS = 3.0 PRICE_PER_SECOND = 0.5 @@ -52,6 +62,7 @@ class TestAzureSpeechPassthroughHandler: "url_route,expected_model,expected_cost", [ (SHORT_AUDIO_URL, "azure_speech/short-audio", TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND), + (FAST_URL, "azure_speech/fast-transcription", FAST_AUDIO_SECONDS * PRICE_PER_SECOND), (BATCH_URL, "azure_speech/batch-transcription", 0.0), (f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription", 0.0), ], @@ -61,7 +72,7 @@ class TestAzureSpeechPassthroughHandler: handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( httpx_response=_make_response(url_route), - response_body=TRANSCRIPT_BODY, + response_body={**TRANSCRIPT_BODY, **FAST_BODY}, logging_obj=logging_obj, url_route=url_route, result=TRANSCRIPT, @@ -110,6 +121,28 @@ class TestAzureSpeechPassthroughHandler: assert handler_result["kwargs"]["model"] == "azure_speech/short-audio" assert handler_result["kwargs"]["response_cost"] == 0.0 + @pytest.mark.parametrize( + "response_body", + [{"durationMilliseconds": 0}, {"durationMilliseconds": "5061"}, {"duration": 5061}, {}, [], None], + ) + def test_fast_transcription_without_duration_milliseconds_logs_zero_cost( + self, response_body: dict[str, object] | list[dict[str, object]] | None + ): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(FAST_URL), + response_body=response_body, + logging_obj=_make_logging_obj(), + url_route=FAST_URL, + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["model"] == "azure_speech/fast-transcription" + assert handler_result["kwargs"]["response_cost"] == 0.0 + def test_missing_price_entry_still_logs_the_row_at_zero_cost(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.delitem(litellm.model_cost, "azure/speech/azure-stt") diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 43fc28c34c4..eb0c607fbef 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -6141,6 +6141,7 @@ class TestAzureRelayDeploymentSegment: AZURE_SPEECH_SHORT_AUDIO_ENDPOINT: Final = "/speech/recognition/conversation/cognitiveservices/v1" AZURE_SPEECH_BATCH_ENDPOINT: Final = "/speechtotext/v3.2/transcriptions" +AZURE_SPEECH_FAST_ENDPOINT: Final = "/speechtotext/transcriptions:transcribe" AZURE_SPEECH_PCM16_HEADER: Final = ( b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" ) @@ -6149,8 +6150,7 @@ AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range AZURE_SPEECH_TRANSCRIPT: Final = {"RecognitionStatus": "Success", "DisplayText": "The eagle has landed."} -@pytest.fixture -def azure_speech_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: +def _azure_speech_test_client(monkeypatch: pytest.MonkeyPatch, caller: UserAPIKeyAuth) -> TestClient: from litellm.proxy.proxy_server import app monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") @@ -6159,8 +6159,20 @@ def azure_speech_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient] monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) litellm.in_memory_llm_clients_cache.flush_cache() - monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) - yield TestClient(app) + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: caller) + return TestClient(app) + + +@pytest.fixture +def azure_speech_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + yield _azure_speech_test_client(monkeypatch, UserAPIKeyAuth(api_key="sk-virtual")) + + +@pytest.fixture +def azure_speech_admin_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + yield _azure_speech_test_client( + monkeypatch, UserAPIKeyAuth(api_key="sk-admin", user_role=LitellmUserRoles.PROXY_ADMIN) + ) class TestAzureSpeechProxyRoute: @@ -6193,17 +6205,19 @@ class TestAzureSpeechProxyRoute: assert "authorization" not in sent.headers assert "caller-supplied-key" not in repr(sent.headers) - def test_batch_json_goes_to_the_cognitive_services_host(self, azure_speech_client: TestClient) -> None: + def test_admin_batch_job_creation_goes_to_the_cognitive_services_host( + self, azure_speech_admin_client: TestClient + ) -> None: body: Final = {"contentUrls": ["https://example.com/a.wav"], "locale": "en-US", "displayName": "job"} with respx.mock(assert_all_called=True) as upstream: route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( return_value=httpx.Response(201, json={"self": "https://eastus.api.cognitive.microsoft.com/x"}) ) - response = azure_speech_client.post( + response = azure_speech_admin_client.post( f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", json=body, - headers={"Authorization": "Bearer sk-virtual"}, + headers={"Authorization": "Bearer sk-admin"}, ) assert response.status_code == 201 @@ -6212,22 +6226,80 @@ class TestAzureSpeechProxyRoute: assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" assert "authorization" not in sent.headers - def test_batch_multipart_upload_is_forwarded_byte_for_byte(self, azure_speech_client: TestClient) -> None: + @pytest.mark.parametrize( + "method,endpoint", + [ + ("POST", AZURE_SPEECH_BATCH_ENDPOINT), + ("POST", "/speechtotext/v3.2/models"), + ("PUT", "/speechtotext/v3.2/endpoints/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab"), + ], + ) + def test_non_admin_key_cannot_create_unpriced_batch_work( + self, azure_speech_client: TestClient, method: str, endpoint: str + ) -> None: + with respx.mock(assert_all_called=False) as upstream: + catch_all = upstream.route().mock(return_value=httpx.Response(201, json={"status": "NotStarted"})) + + response = azure_speech_client.request( + method, + f"/azure_speech{endpoint}", + json={"contentUrls": ["https://example.com/a.wav"], "locale": "en-US"}, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 403, response.text + assert AZURE_SPEECH_FAST_ENDPOINT in response.text + assert not catch_all.called + + def test_non_admin_key_can_still_read_delete_and_fast_transcribe_in_the_batch_family( + self, azure_speech_client: TestClient + ) -> None: + job_path: Final = f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab" with respx.mock(assert_all_called=True) as upstream: - route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( - return_value=httpx.Response(201, json={"status": "NotStarted"}) + upstream.get(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( + return_value=httpx.Response(200, json={"status": "Succeeded"}) + ) + upstream.delete(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( + return_value=httpx.Response(204) + ) + upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"durationMilliseconds": 640, "combinedPhrases": []}) + ) + + statuses = [ + azure_speech_client.get(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}), + azure_speech_client.delete(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}), + azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", + params={"api-version": "2024-11-15"}, + files={"audio": ("eagle.wav", AZURE_SPEECH_WAV_BYTES, "audio/wav")}, + data={"definition": json.dumps({"locales": ["en-US"]})}, + headers={"Authorization": "Bearer sk-virtual"}, + ), + ] + + assert [r.status_code for r in statuses] == [200, 204, 200] + + def test_fast_transcription_multipart_upload_is_forwarded_byte_for_byte( + self, azure_speech_client: TestClient + ) -> None: + with respx.mock(assert_all_called=True) as upstream: + route = upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"durationMilliseconds": 640, "combinedPhrases": []}) ) response = azure_speech_client.post( - f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", + f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", + params={"api-version": "2024-11-15"}, files={"audio": ("eagle.wav", AZURE_SPEECH_NON_UTF8_WAV_BYTES, "audio/wav")}, data={"definition": json.dumps({"locales": ["en-US"]})}, headers={"Authorization": "Bearer sk-virtual"}, ) - assert response.status_code == 201 + assert response.status_code == 200 sent = route.calls.last.request assert sent.headers["content-type"].startswith("multipart/form-data; boundary=") + assert dict(sent.url.params) == {"api-version": "2024-11-15"} assert AZURE_SPEECH_NON_UTF8_WAV_BYTES in sent.content assert b'name="definition"' in sent.content assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" @@ -6247,7 +6319,7 @@ class TestAzureSpeechProxyRoute: @pytest.mark.parametrize("method", ["GET", "POST"]) def test_batch_requests_are_logged_as_azure_speech_not_assemblyai( - self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch, method: str + self, azure_speech_admin_client: TestClient, monkeypatch: pytest.MonkeyPatch, method: str ) -> None: from litellm.integrations.custom_logger import CustomLogger @@ -6266,11 +6338,11 @@ class TestAzureSpeechProxyRoute: return_value=httpx.Response(200, json={"values": []}) ) - response = azure_speech_client.request( + response = azure_speech_admin_client.request( method, f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", json={"locale": "en-US"} if method == "POST" else None, - headers={"Authorization": "Bearer sk-virtual"}, + headers={"Authorization": "Bearer sk-admin"}, ) assert response.status_code == 200 @@ -6278,6 +6350,50 @@ class TestAzureSpeechProxyRoute: ("azure_speech/batch-transcription", "azure_speech", 0.0) ] + def test_fast_transcription_spend_is_priced_from_duration_milliseconds( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[dict[str, object]] = [] # mutable-ok: test recorder accumulates callback payloads + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + self.payloads.append(kwargs["standard_logging_object"]) + + recorder: Final = _Recorder() + monkeypatch.setattr(litellm, "_async_success_callback", [*litellm._async_success_callback, recorder]) + monkeypatch.setitem( + litellm.model_cost, + "azure/speech/azure-stt", + { + "litellm_provider": "azure", + "mode": "audio_transcription", + "input_cost_per_second": 0.25, + "output_cost_per_second": 0.0, + }, + ) + with respx.mock(assert_all_called=True) as upstream: + upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"durationMilliseconds": 5061, "combinedPhrases": []}) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", + params={"api-version": "2024-11-15"}, + files={"audio": ("eagle.wav", AZURE_SPEECH_WAV_BYTES, "audio/wav")}, + data={"definition": json.dumps({"locales": ["en-US"]})}, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 200 + assert [(p["model"], p["custom_llm_provider"]) for p in recorder.payloads] == [ + ("azure_speech/fast-transcription", "azure_speech") + ] + assert recorder.payloads[0]["response_cost"] == pytest.approx(5.061 * 0.25) + def test_short_audio_spend_is_priced_from_the_recognized_duration( self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch ) -> None: From 849859001f5e0cce29bd973778fca71e3350e16b Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 18:18:40 +0000 Subject: [PATCH 094/253] fix(deepgram): refuse callback delivery on the /listen passthrough so sessions cannot go unbilled With callback or callback_method in the query, Deepgram sends every Results and Metadata frame to the caller's URL and only a request id down this socket, so the proxy would meter zero seconds of audio while its own Deepgram credential paid for the transcription. The route now closes such connections with 1008 before contacting Deepgram, naming the offending parameters in the close reason. Adds helper and route tests for both parameters and a nine mutation sweep, all killed Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 5 +++ .../llm_passthrough_endpoints.py | 14 +++++- .../deepgram/test_deepgram_common_utils.py | 19 ++++++++ .../test_deepgram_ws_passthrough_routes.py | 45 +++++++++++++++++++ 4 files changed, 82 insertions(+), 1 deletion(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index f1759f94775..947df37bbe7 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -10,6 +10,7 @@ from litellm.constants import DEEPGRAM_DEFAULT_API_BASE, DEEPGRAM_LISTEN_DEFAULT from litellm.llms.base_llm.chat.transformation import BaseLLMException _WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws", "wss": "wss", "ws": "ws"}) +DEEPGRAM_LISTEN_CALLBACK_PARAMS: Final = frozenset({"callback", "callback_method"}) class DeepgramException(BaseLLMException): @@ -26,6 +27,10 @@ def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> return f"{websocket_url}?{query}" +def deepgram_listen_callback_params(query_string: str) -> tuple[str, ...]: + return tuple(sorted(DEEPGRAM_LISTEN_CALLBACK_PARAMS.intersection(httpx.QueryParams(query_string).keys()))) + + def deepgram_listen_model(upstream_url: str) -> str: models: Final = parse_qs(urlparse(upstream_url).query).get("model") return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 1abbf90cb7a..7f9e0169fb2 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -36,7 +36,10 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment from litellm.llms.custom_httpx.http_handler import get_async_httpx_client -from litellm.llms.deepgram.common_utils import deepgram_listen_websocket_target +from litellm.llms.deepgram.common_utils import ( + deepgram_listen_callback_params, + deepgram_listen_websocket_target, +) from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.passthrough.main import AsyncPassthroughStreamingResponse @@ -2694,6 +2697,7 @@ async def openai_websocket_proxy_route( _DEEPGRAM_WS_MISSING_KEY_REASON: Final = ( "Required 'DEEPGRAM_API_KEY' in environment to make pass-through calls to Deepgram." ) +_DEEPGRAM_WS_CALLBACK_REASON: Final = "Deepgram callback delivery is not supported through the proxy: remove {params}" @router.websocket("/deepgram/v1/listen") @@ -2712,6 +2716,14 @@ async def deepgram_listen_websocket_route( return await websocket.accept(subprotocol=_negotiated_websocket_subprotocol(websocket)) + callback_params: Final = deepgram_listen_callback_params(websocket.url.query) + if callback_params: + await websocket.close( + code=1008, + reason=_DEEPGRAM_WS_CALLBACK_REASON.format(params=", ".join(callback_params)), + ) + return + await relay( websocket=websocket, target=deepgram_listen_websocket_target( diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index a86cb83d628..65fbf7c7870 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -7,6 +7,7 @@ import pytest import litellm from litellm.llms.deepgram.common_utils import ( deepgram_listen_audio_seconds, + deepgram_listen_callback_params, deepgram_listen_model, deepgram_listen_transcript, deepgram_listen_websocket_target, @@ -68,6 +69,24 @@ def test_deepgram_listen_websocket_target(api_base: str | None, query_string: st assert deepgram_listen_websocket_target(api_base=api_base, query_string=query_string) == expected +@pytest.mark.parametrize( + ("query_string", "expected"), + [ + pytest.param("model=nova-3&encoding=linear16", (), id="no callback"), + pytest.param("model=nova-3&callback=https%3A%2F%2Fevil.example%2Fsink", ("callback",), id="callback"), + pytest.param( + "callback_method=put&model=nova-3&callback=wss%3A%2F%2Fevil.example", + ("callback", "callback_method"), + id="callback and method", + ), + pytest.param("model=nova-3&callback_method=put", ("callback_method",), id="method alone"), + pytest.param("model=nova-3&callbacks=x&my_callback=y", (), id="only exact names match"), + ], +) +def test_deepgram_listen_callback_params(query_string: str, expected: tuple[str, ...]): + assert deepgram_listen_callback_params(query_string) == expected + + @pytest.mark.parametrize( ("frames", "expected_seconds"), [ diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 64804e621ba..5ea2b0b8ab9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -207,6 +207,30 @@ async def test_deepgram_listen_closes_cleanly_when_provider_credentials_missing( assert relay.calls == [] +@pytest.mark.asyncio +@pytest.mark.parametrize( + "query", + [ + pytest.param("model=nova-3&callback=https%3A%2F%2Fsink.example%2Fdg", id="http callback"), + pytest.param("callback=wss%3A%2F%2Fsink.example&callback_method=put&model=nova-3", id="ws callback"), + ], +) +async def test_deepgram_listen_rejects_callback_delivery_that_would_go_unbilled(query, monkeypatch): + """With ``callback`` set, Deepgram sends every Results and Metadata frame to the caller's URL and only a + request id down this socket, so the proxy would meter zero seconds of audio; refuse before contacting Deepgram.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket("/deepgram/v1/listen", query) + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert relay.calls == [] + assert websocket.closed is not None + assert websocket.closed[0] == 1008 + assert "callback" in websocket.closed[1] + assert "dg-provider-key" not in websocket.closed[1] + + def _app_with_relay(relay: _FakeRelay) -> FastAPI: app = FastAPI() app.include_router(router) @@ -228,6 +252,27 @@ def test_deepgram_listen_rejects_connections_without_a_litellm_key(): get_credentials.assert_not_called() +def test_deepgram_listen_callback_rejection_reaches_the_client_as_a_policy_close(monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + + with ( + patch(GET_CREDENTIALS, return_value="dg-provider-key"), + patch(USER_API_KEY_AUTH, new=AsyncMock(return_value=UserAPIKeyAuth(api_key="hashed"))), + ): + with pytest.raises(WebSocketDisconnect) as disconnect: + with client.websocket_connect( + "/deepgram/v1/listen?model=nova-3&callback=https%3A%2F%2Fsink.example%2Fdg", + headers={"Authorization": "Bearer sk-litellm-virtual"}, + ) as connection: + connection.receive_text() + + assert disconnect.value.code == 1008 + assert "callback" in disconnect.value.reason + assert relay.calls == [] + + def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(monkeypatch): monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) relay = _FakeRelay() From 73dea4c5676218005a42bdbe841c76ce6eb42707 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 18:58:21 +0000 Subject: [PATCH 095/253] fix(proxy): attribute only the raising layer in stream and pipeline blocks The streaming wrapper caught every exception crossing its boundary and named its own callback, so a block by an inner guardrail or a provider stream failure also named every outer guardrail. The wrapper now runs the hook over an upstream boundary that remembers the exception it raised, and skips attribution when the same exception passes through Pipeline blocks converted from SensitiveDataRouteException or ModifyResponseException into a generic guardrail_pipeline_error now still record the blocking step's guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/utils.py | 129 ++++++++++++----- .../proxy_logging/test_guardrail_pipeline.py | 26 +++- .../proxy_logging/test_streaming_hooks.py | 132 +++++++++++++++--- 3 files changed, 223 insertions(+), 64 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index fea6ce20b61..21e6ed1d24f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -12,13 +12,36 @@ import sys import threading import time import traceback -from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence +from collections.abc import ( + AsyncGenerator, + AsyncIterable, + AsyncIterator, + Awaitable, + Callable, + Coroutine, + Mapping, + Sequence, +) from dataclasses import dataclass, field from datetime import date, datetime, timedelta, timezone from email.mime.multipart import MIMEMultipart from email.mime.text import MIMEText +from functools import partial from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, Protocol, TypeVar, Union, cast, overload +from typing import ( + TYPE_CHECKING, + Any, + ClassVar, + Final, + Generic, + Literal, + Optional, + Protocol, + TypeVar, + Union, + cast, + overload, +) from typing_extensions import ReadOnly, TypedDict @@ -444,6 +467,33 @@ def _record_raising_guardrail(request_data: Mapping[str, object], callback: obje add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=guardrail_name) +class _UpstreamStreamBoundary(Generic[_T]): + """Remembers the exception the upstream iterator raised, so the wrapper around a + streaming hook can tell a pass-through failure from one the hook raised itself.""" + + __slots__ = ("_upstream", "failure") + + def __init__(self, upstream: AsyncIterable[_T]) -> None: + self._upstream: Final = upstream.__aiter__() + self.failure: BaseException | None = None + + def __aiter__(self) -> "_UpstreamStreamBoundary[_T]": + return self + + async def __anext__(self) -> _T: + try: + return await self._upstream.__anext__() + except StopAsyncIteration: + raise + except Exception as e: + self.failure = e + raise + + +class _StreamIteratorHook(Protocol[_T]): + def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ... + + def _is_client_error_exception(exc: Exception) -> bool: if isinstance(exc, HTTPException): return exc.status_code < 500 @@ -2037,14 +2087,18 @@ class ProxyLogging: _merge_pipeline_metadata_writes(data, result.modified_data) if result.terminal_action == "block": + blocking_step: Final = result.step_results[-1] if result.step_results else None + callback: Final = ( + PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name) + if blocking_step is not None + else None + ) + if callback is not None: + _record_raising_guardrail(data, callback) original_exception: Final = result.original_exception if original_exception is not None and not _exception_changes_request_flow(original_exception): - blocking_step: Final = result.step_results[-1] if result.step_results else None - if blocking_step is not None: - callback: Final = PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name) - if callback is not None: - _enrich_http_exception_with_guardrail_context(original_exception, callback) - _record_raising_guardrail(data, callback) + if callback is not None: + _enrich_http_exception_with_guardrail_context(original_exception, callback) raise original_exception step_results_serializable: Final = [ @@ -2490,23 +2544,26 @@ class ProxyLogging: @staticmethod async def _wrap_streaming_iterator_with_enrichment( callback: object, - gen: AsyncGenerator[_T, None], + response: AsyncIterable[_T], + hook: _StreamIteratorHook[_T], request_data: Mapping[str, object], ) -> AsyncGenerator[_T, None]: """ - Yield from `gen`; if iteration raises an HTTPException with dict detail, - enrich the detail with the originating callback's `guardrail_name` and - `guardrail_mode` before re-raising. Used to wrap each layer of the - async_post_call_streaming_iterator_hook chain so the enrichment is - attributed to the callback that produced the chunk pipeline at that - point in the chain. + Run `hook` over `response` and yield its chunks. If the hook itself raises, + enrich an HTTPException's dict detail with the callback's `guardrail_name` + and `guardrail_mode` and record the callback in `applied_guardrails` before + re-raising. Failures raised by `response` (the provider stream or an inner + layer of the async_post_call_streaming_iterator_hook chain) pass through + untouched, so only the layer that actually raised is attributed. """ + upstream: Final = _UpstreamStreamBoundary(response) try: - async for chunk in gen: + async for chunk in hook(response=upstream): yield chunk except Exception as e: - _enrich_http_exception_with_guardrail_context(e, callback) - _record_raising_guardrail(request_data, callback) + if e is not upstream.failure: + _enrich_http_exception_with_guardrail_context(e, callback) + _record_raising_guardrail(request_data, callback) raise # Cache for callback-capability detection. Keyed on a signature of @@ -3663,29 +3720,27 @@ class ProxyLogging: ) else kind ) - if effective_kind == "override": - current_response = self._wrap_streaming_iterator_with_enrichment( - resolved_callback, - resolved_callback.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=current_response, - request_data=request_data, - ), + hook: _StreamIteratorHook[object] = ( + partial( + resolved_callback.async_post_call_streaming_iterator_hook, + user_api_key_dict=user_api_key_dict, request_data=request_data, ) - else: - # kind == "apply_guardrail": route through unified_guardrail - current_response = self._wrap_streaming_iterator_with_enrichment( - resolved_callback, - unified_guardrail.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - request_data=request_data, - response=current_response, - guardrail_to_apply=resolved_callback, - buffer_until_moderated_default=(kind == "override"), - ), + if effective_kind == "override" + else partial( + unified_guardrail.async_post_call_streaming_iterator_hook, + user_api_key_dict=user_api_key_dict, request_data=request_data, + guardrail_to_apply=resolved_callback, + buffer_until_moderated_default=(kind == "override"), ) + ) + current_response = self._wrap_streaming_iterator_with_enrichment( + resolved_callback, + current_response, + hook, + request_data=request_data, + ) pipeline_translation: Final = ( resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py index cd2b7a278bb..5f3c09d9195 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -551,14 +551,23 @@ def test_handle_pipeline_result_block_does_not_reraise_sensitive_data_route(): session_id="sess-1", guardrail_name="pii-router", ) + cb = _make_guardrail() + cb.guardrail_name = "pii-router" result = MagicMock() result.terminal_action = "block" result.step_results = [MagicMock(guardrail_name="pii-router")] result.original_exception = original - with pytest.raises(HTTPException) as info: - ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p") + data: dict[str, object] = {"model": "m"} + saved = litellm.callbacks + litellm.callbacks = [cb] + try: + with pytest.raises(HTTPException) as info: + ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") + finally: + litellm.callbacks = saved assert info.value.status_code == 400 assert info.value.detail["error"]["type"] == "guardrail_pipeline_error" + assert data["metadata"] == {"applied_guardrails": ["pii-router"]} def test_handle_pipeline_result_block_does_not_reraise_modify_response(): @@ -571,14 +580,23 @@ def test_handle_pipeline_result_block_does_not_reraise_modify_response(): request_data={"model": "m"}, guardrail_name="masker", ) + cb = _make_guardrail() + cb.guardrail_name = "masker" result = MagicMock() result.terminal_action = "block" result.step_results = [MagicMock(guardrail_name="masker")] result.original_exception = original - with pytest.raises(HTTPException) as info: - ProxyLogging._handle_pipeline_result(result=result, data={"model": "m"}, policy_name="p") + data: dict[str, object] = {"model": "m"} + saved = litellm.callbacks + litellm.callbacks = [cb] + try: + with pytest.raises(HTTPException) as info: + ProxyLogging._handle_pipeline_result(result=result, data=data, policy_name="p") + finally: + litellm.callbacks = saved assert info.value.status_code == 400 assert info.value.detail["error"]["type"] == "guardrail_pipeline_error" + assert data["metadata"] == {"applied_guardrails": ["masker"]} def test_handle_pipeline_result_modify_response_raises_modify_exception(): diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 52586ed2174..6fb000b4fa7 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -170,6 +170,15 @@ def test_init_response_taking_too_long_task_no_slack_instance_no_error_raises(pr # --------------------------------------------------------------------------- +async def _passthrough_hook(*, response: AsyncIterator[object]) -> AsyncGenerator[object, None]: + async for chunk in response: + yield chunk + + +async def _one_chunk() -> AsyncGenerator[object, None]: + yield "chunk" + + @pytest.mark.asyncio async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(proxy_logging): async def gen(): @@ -177,7 +186,9 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro yield ch cb = MagicMock(guardrail_name="g", event_hook="pre_call") - wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=gen(), request_data={}) + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment( + callback=cb, response=gen(), hook=_passthrough_hook, request_data={} + ) out = [ch async for ch in wrapped] snapshot = { "chunks": out, @@ -197,18 +208,43 @@ async def test_wrap_streaming_iterator_with_enrichment_passes_through_chunks(pro async def test_wrap_streaming_iterator_with_enrichment_enriches_http_exception_raises(proxy_logging): detail = {"error": "blocked"} - async def boom_gen(): + async def boom_hook(*, response: AsyncIterator[object]) -> AsyncGenerator[object, None]: if False: yield # pragma: no cover raise HTTPException(status_code=400, detail=detail) cb = MagicMock(guardrail_name="presidio", event_hook="post_call") - wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment(callback=cb, gen=boom_gen(), request_data={}) + request_data: dict[str, object] = {} + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment( + callback=cb, response=_one_chunk(), hook=boom_hook, request_data=request_data + ) with pytest.raises(HTTPException): async for _ in wrapped: pass assert detail["guardrail_name"] == "presidio" assert detail["guardrail_mode"] == "post_call" + assert request_data["metadata"]["applied_guardrails"] == ["presidio"] + + +@pytest.mark.asyncio +async def test_wrap_streaming_iterator_leaves_upstream_http_exception_unattributed(proxy_logging): + detail = {"error": "upstream rejected the stream"} + + async def failing_upstream() -> AsyncGenerator[object, None]: + if False: + yield # pragma: no cover + raise HTTPException(status_code=502, detail=detail) + + cb = MagicMock(guardrail_name="presidio", event_hook="post_call") + request_data: dict[str, object] = {} + wrapped = proxy_logging._wrap_streaming_iterator_with_enrichment( + callback=cb, response=failing_upstream(), hook=_passthrough_hook, request_data=request_data + ) + with pytest.raises(HTTPException): + async for _ in wrapped: + pass + assert detail == {"error": "upstream rejected the stream"} + assert request_data == {} # --------------------------------------------------------------------------- @@ -700,33 +736,83 @@ async def test_post_call_response_headers_hook_swallows_callback_error(proxy_log assert out == {} +class _StreamBlocker(CustomGuardrail): + def __init__(self, guardrail_name: str = "stream-blocker") -> None: + super().__init__(guardrail_name=guardrail_name, event_hook=GuardrailEventHooks.post_call, default_on=True) + + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object] + ) -> AsyncGenerator[object, None]: + async for _ in response: + raise HTTPException(status_code=400, detail={"error": "blocked"}) + yield # pragma: no cover + + +class _StreamPasser(CustomGuardrail): + def __init__(self, guardrail_name: str = "stream-passer") -> None: + super().__init__(guardrail_name=guardrail_name, event_hook=GuardrailEventHooks.post_call, default_on=True) + + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object] + ) -> AsyncGenerator[object, None]: + async for chunk in response: + yield chunk + + +async def _drain_stream_chain( + proxy_logging: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, + upstream: AsyncIterator[object], + request_data: dict[str, object], +) -> None: + async for _ in proxy_logging.async_post_call_streaming_iterator_hook( + response=upstream, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ): + pass + + +async def _failing_provider_stream() -> AsyncGenerator[object, None]: + yield "chunk" + raise RuntimeError("provider connection dropped") + + @pytest.mark.asyncio async def test_stream_guardrail_block_names_the_blocking_guardrail_in_applied_guardrails( proxy_logging, make_user_api_key_auth, monkeypatch ): - class _StreamBlocker(CustomGuardrail): - def __init__(self) -> None: - super().__init__(guardrail_name="stream-blocker", event_hook=GuardrailEventHooks.post_call, default_on=True) - - async def async_post_call_streaming_iterator_hook( - self, user_api_key_dict: UserAPIKeyAuth, response: AsyncIterator[object], request_data: dict[str, object] - ) -> AsyncGenerator[object, None]: - async for _ in response: - raise HTTPException(status_code=400, detail={"error": "blocked"}) - yield # pragma: no cover - monkeypatch.setattr(litellm, "callbacks", [_StreamBlocker()]) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) - async def upstream(): - yield "chunk" - request_data: dict[str, object] = {"metadata": {}} with pytest.raises(HTTPException): - async for _ in proxy_logging.async_post_call_streaming_iterator_hook( - response=upstream(), - user_api_key_dict=make_user_api_key_auth(), - request_data=request_data, - ): - pass + await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _one_chunk(), request_data) assert request_data["metadata"]["applied_guardrails"] == ["stream-blocker"] + + +@pytest.mark.asyncio +async def test_stream_block_by_inner_guardrail_does_not_name_the_outer_layers( + proxy_logging, make_user_api_key_auth, monkeypatch +): + monkeypatch.setattr(litellm, "callbacks", [_StreamBlocker(), _StreamPasser("outer-a"), _StreamPasser("outer-b")]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + + request_data: dict[str, object] = {"metadata": {}} + with pytest.raises(HTTPException) as info: + await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _one_chunk(), request_data) + assert info.value.detail["guardrail_name"] == "stream-blocker" + assert request_data["metadata"]["applied_guardrails"] == ["stream-blocker"] + + +@pytest.mark.asyncio +async def test_stream_provider_failure_is_not_attributed_to_any_guardrail( + proxy_logging, make_user_api_key_auth, monkeypatch +): + monkeypatch.setattr(litellm, "callbacks", [_StreamPasser("outer-a"), _StreamPasser("outer-b")]) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + + request_data: dict[str, object] = {"metadata": {}} + with pytest.raises(RuntimeError, match="provider connection dropped"): + await _drain_stream_chain(proxy_logging, make_user_api_key_auth(), _failing_provider_stream(), request_data) + assert request_data["metadata"] == {} From db560ca65248d4c03c0c02a1d4cca392d0294144 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 19:02:44 +0000 Subject: [PATCH 096/253] fix(passthrough): bill every Deepgram channel, not just wall-clock duration Deepgram charges for the total processed audio across channels, so a stereo /listen session with multichannel=true costs twice its duration. The handler now multiplies the session duration by a validated channel count taken from Metadata.channels, then the widest Results channel_index, then the channels query parameter, defaulting to one. Booleans, floats, strings, zero and negative values are ignored Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 40 +++++++++++ ...gram_listen_passthrough_logging_handler.py | 8 ++- .../deepgram/test_deepgram_common_utils.py | 70 ++++++++++++++++++- ...gram_listen_passthrough_logging_handler.py | 32 ++++++++- 4 files changed, 143 insertions(+), 7 deletions(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index 947df37bbe7..6beb70b4499 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -36,6 +36,46 @@ def deepgram_listen_model(upstream_url: str) -> str: return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL +def _channel_count(value: object) -> int | None: + if isinstance(value, bool) or not isinstance(value, int): + return None + return value if value >= 1 else None + + +def _results_channel_count(frame: Mapping[str, object]) -> int | None: + channel_index: Final = frame.get("channel_index") + if not isinstance(channel_index, list) or len(channel_index) != 2: + return None + return _channel_count(channel_index[1]) + + +def _declared_channel_count(upstream_url: str) -> int | None: + declared: Final = parse_qs(urlparse(upstream_url).query).get("channels") + if not declared or not declared[0].isdigit(): + return None + return _channel_count(int(declared[0])) + + +def deepgram_listen_channel_count(websocket_messages: Sequence[Mapping[str, object]], upstream_url: str) -> int: + metadata_channels: Final = tuple( + channels + for frame in websocket_messages + if frame.get("type") == "Metadata" + if (channels := _channel_count(frame.get("channels"))) is not None + ) + if metadata_channels: + return metadata_channels[-1] + results_channels: Final = tuple( + channels + for frame in websocket_messages + if frame.get("type") == "Results" + if (channels := _results_channel_count(frame)) is not None + ) + if results_channels: + return max(results_channels) + return _declared_channel_count(upstream_url) or 1 + + def _seconds(value: object) -> float | None: if isinstance(value, bool) or not isinstance(value, (int, float)): return None diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py index a5fea7c8020..0a8d7ff3d33 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -8,6 +8,7 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.deepgram.common_utils import ( deepgram_listen_audio_seconds, + deepgram_listen_channel_count, deepgram_listen_model, deepgram_listen_transcript, ) @@ -45,8 +46,10 @@ class DeepgramListenPassthroughLoggingHandler: ) -> PassThroughEndpointLoggingTypedDict: model: Final = deepgram_listen_model(upstream_url) audio_seconds: Final = deepgram_listen_audio_seconds(websocket_messages) + channels: Final = deepgram_listen_channel_count(websocket_messages, upstream_url) + billed_seconds: Final = audio_seconds * channels response: Final = TranscriptionResponse(text=deepgram_listen_transcript(websocket_messages)) - response._hidden_params["audio_transcription_duration"] = audio_seconds # pyright: ignore[reportPrivateUsage] # the cost calculator reads the billed duration off the response's hidden params + response._hidden_params["audio_transcription_duration"] = billed_seconds # pyright: ignore[reportPrivateUsage] # the cost calculator reads the billed duration off the response's hidden params response_cost: Final = _audio_cost(response, model) response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads a precomputed cost off the response's hidden params @@ -56,9 +59,10 @@ class DeepgramListenPassthroughLoggingHandler: logging_obj.model_call_details["custom_llm_provider"] = provider # rebind-ok: same shared logging object logging_obj.model_call_details["response_cost"] = response_cost # rebind-ok: same shared logging object verbose_proxy_logger.debug( - "Deepgram listen passthrough cost tracking: model %s, audio seconds %s, cost %s", + "Deepgram listen passthrough cost tracking: model %s, audio seconds %s, channels %s, cost %s", model, audio_seconds, + channels, response_cost, ) logging_result: Final[PassThroughEndpointLoggingTypedDict] = { diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index 65fbf7c7870..d7a38f2ca5a 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -8,6 +8,7 @@ import litellm from litellm.llms.deepgram.common_utils import ( deepgram_listen_audio_seconds, deepgram_listen_callback_params, + deepgram_listen_channel_count, deepgram_listen_model, deepgram_listen_transcript, deepgram_listen_websocket_target, @@ -16,18 +17,25 @@ from litellm.llms.deepgram.common_utils import ( NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000" -def _results(start: object, duration: object, transcript: str = "", is_final: object = True) -> dict[str, object]: +def _results( + start: object, + duration: object, + transcript: str = "", + is_final: object = True, + channel_index: object = (0, 1), +) -> dict[str, object]: return { "type": "Results", "start": start, "duration": duration, "is_final": is_final, + "channel_index": list(channel_index) if isinstance(channel_index, tuple) else channel_index, "channel": {"alternatives": [{"transcript": transcript, "confidence": 0.9}]}, } -def _metadata(duration: object) -> dict[str, object]: - return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": 1} +def _metadata(duration: object, channels: object = 1) -> dict[str, object]: + return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": channels} @pytest.mark.parametrize( @@ -107,6 +115,62 @@ def test_deepgram_listen_audio_seconds(frames: Sequence[Mapping[str, object]], e assert deepgram_listen_audio_seconds(frames) == expected_seconds +@pytest.mark.parametrize( + ("frames", "upstream_url", "expected_channels"), + [ + pytest.param((_results(0.0, 2.0), _metadata(6.25)), NOVA_3_URL, 1, id="mono"), + pytest.param((_results(0.0, 2.0, channel_index=(0, 2)), _metadata(6.25, 2)), NOVA_3_URL, 2, id="stereo"), + pytest.param((_metadata(1.0, 3), _metadata(1.0, 5)), NOVA_3_URL, 5, id="last metadata wins"), + pytest.param( + (_metadata(1.0, 20), _results(0.0, 1.0, channel_index=(1, 2))), + NOVA_3_URL, + 20, + id="metadata beats channel_index", + ), + pytest.param( + (_results(0.0, 1.0, channel_index=(0, 2)), _results(0.0, 1.0, channel_index=(3, 4))), + NOVA_3_URL, + 4, + id="widest channel_index without metadata", + ), + pytest.param( + (_results(0.0, 1.0, channel_index=(0, 2)),), + f"{NOVA_3_URL}&channels=7&multichannel=true", + 2, + id="frames beat the declared query", + ), + pytest.param((), f"{NOVA_3_URL}&channels=7&multichannel=true", 7, id="declared query when no frames"), + pytest.param((), f"{NOVA_3_URL}&channels=0", 1, id="zero declared channels"), + pytest.param((), f"{NOVA_3_URL}&channels=-2", 1, id="negative declared channels"), + pytest.param((), f"{NOVA_3_URL}&channels=two", 1, id="non numeric declared channels"), + pytest.param((), NOVA_3_URL, 1, id="nothing declared"), + pytest.param((_metadata(1.0, "2"), _metadata(1.0, True), _metadata(1.0, 0)), NOVA_3_URL, 1, id="bad metadata"), + pytest.param((_metadata(1.0, 3), _metadata(1.0, True)), NOVA_3_URL, 3, id="boolean does not shadow a count"), + pytest.param((_metadata(1.0, 2.0), _metadata(1.0, -1)), NOVA_3_URL, 1, id="float and negative metadata"), + pytest.param( + (_metadata(1.0, 2), {**_results(0.0, 1.0), "channels": 9}, {"type": "UtteranceEnd", "channels": 11}), + NOVA_3_URL, + 2, + id="channels on non metadata frames ignored", + ), + pytest.param( + ( + _results(0.0, 1.0, channel_index=[0]), + _results(0.0, 1.0, channel_index=(0, "2")), + _results(0.0, 1.0, channel_index=(0, 0)), + ), + NOVA_3_URL, + 1, + id="bad channel_index", + ), + ], +) +def test_deepgram_listen_channel_count( + frames: Sequence[Mapping[str, object]], upstream_url: str, expected_channels: int +): + assert deepgram_listen_channel_count(frames, upstream_url) == expected_channels + + def test_deepgram_listen_transcript_joins_final_results_only(): frames = ( _results(0.0, 1.0, "hello wor", is_final=False), diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py index 40f520e344d..944b37f95e3 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -30,8 +30,8 @@ def _results(start: object, duration: object, transcript: str = "", is_final: ob } -def _metadata(duration: object) -> dict[str, object]: - return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": 1} +def _metadata(duration: object, channels: int = 1) -> dict[str, object]: + return {"type": "Metadata", "request_id": "req-1", "duration": duration, "channels": channels} @pytest.mark.parametrize( @@ -119,6 +119,34 @@ def test_handler_charges_more_for_more_audio_on_the_same_model(): assert short["kwargs"]["response_cost"] > 0 +def test_handler_bills_every_channel_of_a_multichannel_session(): + """Deepgram bills processed audio per channel (deepgram.com/pricing FAQ, 2026-09-17), so a stereo session must be + charged for twice its wall-clock duration or budgets can be bypassed by requesting more channels.""" + mono = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(30.0),), logging_obj=_logging_obj(), upstream_url=NOVA_3_URL + ) + stereo = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(30.0, channels=2),), + logging_obj=_logging_obj(), + upstream_url=f"{NOVA_3_URL}&multichannel=true&channels=2", + ) + + assert stereo["result"]._hidden_params["audio_transcription_duration"] == 60.0 + assert stereo["kwargs"]["response_cost"] == pytest.approx(2 * mono["kwargs"]["response_cost"]) + assert stereo["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 60.0)) + + +def test_handler_bills_the_declared_channels_when_the_stream_dies_before_any_frame_reports_them(): + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_results(0.0, 10.0, "a"),), + logging_obj=_logging_obj(), + upstream_url=f"{NOVA_3_URL}&multichannel=true&channels=3", + ) + + assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 30.0 + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 30.0)) + + def test_handler_keeps_the_spend_row_but_no_cost_for_an_unpriced_model(): handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( websocket_messages=(_metadata(12.5),), From 0f3556b1bf505864b8860b41e5d917f38b7e5468 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 17 Sep 2026 19:17:49 +0000 Subject: [PATCH 097/253] refactor(proxy): drop explanatory docstrings from the stream attribution helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 3 +-- litellm/proxy/utils.py | 11 ----------- 2 files changed, 1 insertion(+), 13 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 9a038ca79ba..849211a2dae 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2576,8 +2576,7 @@ class ProxyBaseLLMRequestProcessing: ) async def refresh_stream_headers() -> Mapping[str, str]: - """`custom_headers` rebuilt once the first chunk is buffered, from `self.data` as the - guardrails left it and for whichever deployment served the stream.""" + """`custom_headers` rebuilt for whichever deployment served the stream.""" return self._stream_response_headers( hidden_params=( get_hidden_params_dict(response) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 21e6ed1d24f..4182a864e4c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -468,9 +468,6 @@ def _record_raising_guardrail(request_data: Mapping[str, object], callback: obje class _UpstreamStreamBoundary(Generic[_T]): - """Remembers the exception the upstream iterator raised, so the wrapper around a - streaming hook can tell a pass-through failure from one the hook raised itself.""" - __slots__ = ("_upstream", "failure") def __init__(self, upstream: AsyncIterable[_T]) -> None: @@ -2548,14 +2545,6 @@ class ProxyLogging: hook: _StreamIteratorHook[_T], request_data: Mapping[str, object], ) -> AsyncGenerator[_T, None]: - """ - Run `hook` over `response` and yield its chunks. If the hook itself raises, - enrich an HTTPException's dict detail with the callback's `guardrail_name` - and `guardrail_mode` and record the callback in `applied_guardrails` before - re-raising. Failures raised by `response` (the provider stream or an inner - layer of the async_post_call_streaming_iterator_hook chain) pass through - untouched, so only the layer that actually raised is attributed. - """ upstream: Final = _UpstreamStreamBoundary(response) try: async for chunk in hook(response=upstream): From 25846d13417af204bb52154c125654dcc4da6411 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 20:29:51 +0000 Subject: [PATCH 098/253] fix(proxy): limit the whole Azure Speech batch API to proxy admin keys Ordinary keys could read, patch and delete batch transcription jobs that other keys created with the proxy's shared Azure subscription, so every /speechtotext/v3.2 method is now admin only while fast transcription stays open. Also clears SERVER_ROOT_PATH in the real-auth test helper because test_custom_proxy leaves it set at import time and the shared app then 404s pass-through routes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 1 - .../llm_passthrough_endpoints.py | 15 ++--- .../test_llm_pass_through_endpoints.py | 65 ++++++++++++------- 3 files changed, 48 insertions(+), 33 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index d9da0cc0f64..62523c5d2c3 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1575,7 +1575,6 @@ AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX: Final = "/azure_speech" AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX: Final = "/speech/" AZURE_SPEECH_BATCH_PATH_PREFIX: Final = "/speechtotext/" AZURE_SPEECH_FAST_TRANSCRIPTION_PATH: Final = "/speechtotext/transcriptions:transcribe" -AZURE_SPEECH_UNPRICED_WRITE_METHODS: Final = frozenset({"POST", "PUT"}) AZURE_SPEECH_STT_DOMAIN: Final = "stt.speech.microsoft.com" AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN: Final = "api.cognitive.microsoft.com" AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER: Final = "Ocp-Apim-Subscription-Key" diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index b8966508f81..8d18de42f45 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -38,7 +38,6 @@ from litellm.constants import ( AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX, AZURE_SPEECH_STT_DOMAIN, AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER, - AZURE_SPEECH_UNPRICED_WRITE_METHODS, BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix @@ -1360,11 +1359,10 @@ def resolve_azure_speech_base_url(endpoint_path: str, api_base: str | None, regi return httpx.URL(f"https://{region}.{domain}") -def azure_speech_write_is_unpriced(method: str, endpoint_path: str) -> bool: +def azure_speech_path_manages_shared_resources(endpoint_path: str) -> bool: return ( endpoint_path.startswith(AZURE_SPEECH_BATCH_PATH_PREFIX) and endpoint_path != AZURE_SPEECH_FAST_TRANSCRIPTION_PATH - and method.upper() in AZURE_SPEECH_UNPRICED_WRITE_METHODS ) @@ -1406,15 +1404,14 @@ async def azure_speech_proxy_route( "AZURE_SPEECH_REGION or AZURE_SPEECH_API_BASE in the proxy environment." ), ) - if azure_speech_write_is_unpriced( - method=request.method, endpoint_path=normalized_endpoint_path - ) and not is_proxy_admin(user_api_key_dict): + if azure_speech_path_manages_shared_resources(normalized_endpoint_path) and not is_proxy_admin(user_api_key_dict): raise HTTPException( status_code=403, detail=( - f"{request.method} {normalized_endpoint_path} creates Azure Speech work whose cost is unknown at " - "request time, so it is limited to proxy admin keys. Use " - f"{AZURE_SPEECH_FAST_TRANSCRIPTION_PATH} for transcription that is priced per request." + f"{request.method} {normalized_endpoint_path} manages batch transcription resources that belong to " + "the proxy's Azure Speech subscription and whose cost is unknown at request time, so it is limited " + f"to proxy admin keys. Use {AZURE_SPEECH_FAST_TRANSCRIPTION_PATH} for transcription that is priced " + "per request." ), ) azure_speech_api_key: Final = passthrough_endpoint_router.get_credentials( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index eb0c607fbef..0e8a2b9e0e9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -6232,13 +6232,17 @@ class TestAzureSpeechProxyRoute: ("POST", AZURE_SPEECH_BATCH_ENDPOINT), ("POST", "/speechtotext/v3.2/models"), ("PUT", "/speechtotext/v3.2/endpoints/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab"), + ("GET", AZURE_SPEECH_BATCH_ENDPOINT), + ("GET", f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files"), + ("PATCH", f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab"), + ("DELETE", f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab"), ], ) - def test_non_admin_key_cannot_create_unpriced_batch_work( + def test_non_admin_key_cannot_manage_shared_batch_resources( self, azure_speech_client: TestClient, method: str, endpoint: str ) -> None: with respx.mock(assert_all_called=False) as upstream: - catch_all = upstream.route().mock(return_value=httpx.Response(201, json={"status": "NotStarted"})) + catch_all = upstream.route().mock(return_value=httpx.Response(200, json={"status": "Succeeded"})) response = azure_speech_client.request( method, @@ -6251,9 +6255,23 @@ class TestAzureSpeechProxyRoute: assert AZURE_SPEECH_FAST_ENDPOINT in response.text assert not catch_all.called - def test_non_admin_key_can_still_read_delete_and_fast_transcribe_in_the_batch_family( - self, azure_speech_client: TestClient - ) -> None: + def test_non_admin_key_can_still_fast_transcribe_in_the_batch_family(self, azure_speech_client: TestClient) -> None: + with respx.mock(assert_all_called=True) as upstream: + upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"durationMilliseconds": 640, "combinedPhrases": []}) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", + params={"api-version": "2024-11-15"}, + files={"audio": ("eagle.wav", AZURE_SPEECH_WAV_BYTES, "audio/wav")}, + data={"definition": json.dumps({"locales": ["en-US"]})}, + headers={"Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 200 + + def test_admin_key_reads_and_deletes_batch_jobs(self, azure_speech_admin_client: TestClient) -> None: job_path: Final = f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab" with respx.mock(assert_all_called=True) as upstream: upstream.get(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( @@ -6262,23 +6280,15 @@ class TestAzureSpeechProxyRoute: upstream.delete(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( return_value=httpx.Response(204) ) - upstream.post(f"https://eastus.api.cognitive.microsoft.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( - return_value=httpx.Response(200, json={"durationMilliseconds": 640, "combinedPhrases": []}) - ) statuses = [ - azure_speech_client.get(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}), - azure_speech_client.delete(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}), - azure_speech_client.post( - f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", - params={"api-version": "2024-11-15"}, - files={"audio": ("eagle.wav", AZURE_SPEECH_WAV_BYTES, "audio/wav")}, - data={"definition": json.dumps({"locales": ["en-US"]})}, - headers={"Authorization": "Bearer sk-virtual"}, + azure_speech_admin_client.get(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-admin"}), + azure_speech_admin_client.delete( + f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-admin"} ), ] - assert [r.status_code for r in statuses] == [200, 204, 200] + assert [r.status_code for r in statuses] == [200, 204] def test_fast_transcription_multipart_upload_is_forwarded_byte_for_byte( self, azure_speech_client: TestClient @@ -6305,14 +6315,16 @@ class TestAzureSpeechProxyRoute: assert sent.headers["ocp-apim-subscription-key"] == "server-subscription-key" assert "authorization" not in sent.headers - def test_batch_get_is_forwarded_with_the_job_id_path(self, azure_speech_client: TestClient) -> None: + def test_batch_get_is_forwarded_with_the_job_id_path(self, azure_speech_admin_client: TestClient) -> None: job_path: Final = f"{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files" with respx.mock(assert_all_called=True) as upstream: route = upstream.get(f"https://eastus.api.cognitive.microsoft.com{job_path}").mock( return_value=httpx.Response(200, json={"values": []}) ) - response = azure_speech_client.get(f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-virtual"}) + response = azure_speech_admin_client.get( + f"/azure_speech{job_path}", headers={"Authorization": "Bearer sk-admin"} + ) assert (response.status_code, response.json()) == (200, {"values": []}) assert route.calls.last.request.headers["ocp-apim-subscription-key"] == "server-subscription-key" @@ -6445,8 +6457,8 @@ class TestAzureSpeechProxyRoute: short_audio = upstream.post( f"https://my-speech.cognitiveservices.azure.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}" ).mock(return_value=httpx.Response(200, json=AZURE_SPEECH_TRANSCRIPT)) - batch = upstream.get(f"https://my-speech.cognitiveservices.azure.com{AZURE_SPEECH_BATCH_ENDPOINT}").mock( - return_value=httpx.Response(200, json={"values": []}) + fast = upstream.post(f"https://my-speech.cognitiveservices.azure.com{AZURE_SPEECH_FAST_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"durationMilliseconds": 640, "combinedPhrases": []}) ) azure_speech_client.post( @@ -6454,9 +6466,15 @@ class TestAzureSpeechProxyRoute: content=AZURE_SPEECH_WAV_BYTES, headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, ) - azure_speech_client.get(f"/azure_speech{AZURE_SPEECH_BATCH_ENDPOINT}", headers={"Authorization": "Bearer x"}) + azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_FAST_ENDPOINT}", + params={"api-version": "2024-11-15"}, + files={"audio": ("eagle.wav", AZURE_SPEECH_WAV_BYTES, "audio/wav")}, + data={"definition": json.dumps({"locales": ["en-US"]})}, + headers={"Authorization": "Bearer sk-virtual"}, + ) - assert short_audio.called and batch.called + assert short_audio.called and fast.called @pytest.mark.parametrize("endpoint", ["openai/deployments/whisper/audio/transcriptions", "speech", "speechtotext"]) def test_unknown_path_family_is_rejected_before_any_upstream_call( @@ -6543,6 +6561,7 @@ class TestAzureSpeechRawBodyThroughRealAuth: monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") monkeypatch.setenv("AZURE_SPEECH_REGION", "eastus") monkeypatch.delenv("AZURE_SPEECH_API_BASE", raising=False) + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) litellm.in_memory_llm_clients_cache.flush_cache() with patch.multiple( # test-quality-ok: the real user_api_key_auth reads proxy_server module globals (master_key, caches) that have no injection seam From 11ec157d71b1d40100ff9d2be778477e0a145238 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 20:36:03 +0000 Subject: [PATCH 099/253] fix(deepgram): authorize the effective model and price /listen sessions at streaming rates Key auth on the Deepgram WebSocket route now sees the same model the upstream target will carry, so a key restricted to other models can no longer reach nova-3 by leaving model out of the query. user_api_key_auth_websocket keeps its signature and delegates to user_api_key_auth_websocket_for_model, which the Deepgram route calls with deepgram_listen_requested_model Sessions are priced from new deepgram/streaming/* registry rows (nova-3, nova-3-multilingual for language=multi) plus per-minute add-on rows for redact, keyterm, detect_entities and diarize, all read from Deepgram's pricing page on 2026-09-17. Models without a streaming row fall back to their pre-recorded row as before Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 49 ++++++++++ ...odel_prices_and_context_window_backup.json | 90 ++++++++++++++++++ litellm/proxy/auth/user_api_key_auth.py | 10 +- .../llm_passthrough_endpoints.py | 10 +- ...gram_listen_passthrough_logging_handler.py | 32 ++++++- model_prices_and_context_window.json | 90 ++++++++++++++++++ .../deepgram/test_deepgram_common_utils.py | 68 ++++++++++++++ ...gram_listen_passthrough_logging_handler.py | 92 +++++++++++++++++-- .../test_deepgram_ws_passthrough_routes.py | 62 ++++++++++++- 9 files changed, 481 insertions(+), 22 deletions(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index 6beb70b4499..4f18918fa71 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -11,12 +11,29 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException _WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws", "wss": "wss", "ws": "ws"}) DEEPGRAM_LISTEN_CALLBACK_PARAMS: Final = frozenset({"callback", "callback_method"}) +DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX: Final = "streaming/" +DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE: Final = "multi" +DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX: Final = "-multilingual" +DEEPGRAM_LISTEN_ADDON_PRICING_PARAMS: Final = MappingProxyType( + { + "redact": "redact", + "keyterm": "keyterm", + "detect_entities": "detect_entities", + "diarize": "diarize", + "diarize_model": "diarize", + } +) +_DISABLED_PARAM_VALUES: Final = frozenset({"", "false"}) class DeepgramException(BaseLLMException): pass +def deepgram_listen_requested_model(query_string: str) -> str: + return httpx.QueryParams(query_string).get("model") or DEEPGRAM_LISTEN_DEFAULT_MODEL + + def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> str: listen_url: Final = httpx.URL(f"{(api_base or DEEPGRAM_DEFAULT_API_BASE).rstrip('/')}/listen") websocket_url: Final = listen_url.copy_with(scheme=_WEBSOCKET_SCHEMES.get(listen_url.scheme, listen_url.scheme)) @@ -36,6 +53,38 @@ def deepgram_listen_model(upstream_url: str) -> str: return models[0] if models else DEEPGRAM_LISTEN_DEFAULT_MODEL +def _param_enabled(values: Sequence[str]) -> bool: + return any(value.strip().lower() not in _DISABLED_PARAM_VALUES for value in values) + + +def deepgram_listen_base_pricing_models(upstream_url: str) -> tuple[str, ...]: + """Registry keys to try, in order, for the per-second base rate of a streaming session: the streaming entry for + the language mode Deepgram bills (multilingual when ``language=multi``), then the plain streaming entry, then + the pre-recorded entry for models that have no streaming price of their own.""" + model: Final = deepgram_listen_model(upstream_url) + params: Final = parse_qs(urlparse(upstream_url).query) + streaming: Final = f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{model}" + multilingual: Final = params.get("language", ("",))[-1].strip().lower() == DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE + return ( + (f"{streaming}{DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX}", streaming, model) + if multilingual + else (streaming, model) + ) + + +def deepgram_listen_addon_pricing_models(upstream_url: str) -> tuple[str, ...]: + params: Final = parse_qs(urlparse(upstream_url).query) + return tuple( + sorted( + frozenset( + f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{addon}" + for param, addon in DEEPGRAM_LISTEN_ADDON_PRICING_PARAMS.items() + if _param_enabled(params.get(param, ())) + ) + ) + ) + + def _channel_count(value: object) -> int | None: if isinstance(value, bool) or not isinstance(value, int): return None diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index e87a3fec99b..7b7b3c7199c 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -20736,6 +20736,96 @@ "/v1/audio/transcriptions" ] }, + "deepgram/streaming/nova-3": { + "input_cost_per_second": 8e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0048/60 seconds = $0.00008000 per second", + "note": "Nova-3 monolingual streaming, pay as you go", + "original_pricing_per_minute": 0.0048 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/nova-3-multilingual": { + "input_cost_per_second": 9.667e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0058/60 seconds = $0.00009667 per second", + "note": "Nova-3 multilingual (language=multi) streaming, pay as you go", + "original_pricing_per_minute": 0.0058 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/redact": { + "input_cost_per_second": 3.333e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0020/60 seconds = $0.00003333 per second", + "note": "Redaction add-on (redact query param), streaming, pay as you go", + "original_pricing_per_minute": 0.002 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/keyterm": { + "input_cost_per_second": 2.167e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0013/60 seconds = $0.00002167 per second", + "note": "Keyterm Prompting add-on (keyterm query param), streaming, pay as you go", + "original_pricing_per_minute": 0.0013 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/detect_entities": { + "input_cost_per_second": 2.833e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0017/60 seconds = $0.00002833 per second", + "note": "Entity Detection add-on (detect_entities query param), streaming, pay as you go", + "original_pricing_per_minute": 0.0017 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/diarize": { + "input_cost_per_second": 3.333e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0020/60 seconds = $0.00003333 per second", + "note": "Speaker Diarization add-on (diarize / diarize_model query params), streaming, pay as you go", + "original_pricing_per_minute": 0.002 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, "deepgram/whisper": { "input_cost_per_second": 0.0001, "litellm_provider": "deepgram", diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 4cbd4213463..6842d500a82 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -629,9 +629,11 @@ def _apply_budget_limits_to_end_user_params( verbose_proxy_logger.debug("Applied budget limits to end user %s", end_user_id) -async def user_api_key_auth_websocket(websocket: WebSocket): - # Accept the WebSocket connection +async def user_api_key_auth_websocket(websocket: WebSocket) -> UserAPIKeyAuth: + return await user_api_key_auth_websocket_for_model(websocket, model=websocket.query_params.get("model")) + +async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str | None) -> UserAPIKeyAuth: ws_scope: Final = websocket.scope or {} scope_headers: Final = list(ws_scope.get("headers") or []) # ``get_request_route`` falls back to ``request.url.path`` when @@ -651,10 +653,6 @@ async def user_api_key_auth_websocket(websocket: WebSocket): request._url = websocket.url - query_params: Final = websocket.query_params - - model: Final = query_params.get("model") - async def return_body(): return _realtime_request_body(model) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 7f9e0169fb2..0bd3eb941b9 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -38,6 +38,7 @@ from litellm.llms.azure.passthrough.transformation import foreign_azure_deployme from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.deepgram.common_utils import ( deepgram_listen_callback_params, + deepgram_listen_requested_model, deepgram_listen_websocket_target, ) from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path @@ -52,6 +53,7 @@ from litellm.proxy.auth.user_api_key_auth import ( is_no_auth_dev_mode, user_api_key_auth, user_api_key_auth_websocket, + user_api_key_auth_websocket_for_model, ) from litellm.proxy.common_request_processing import open_sse_before_first_byte from litellm.proxy.common_utils.http_parsing_utils import ( @@ -2700,11 +2702,17 @@ _DEEPGRAM_WS_MISSING_KEY_REASON: Final = ( _DEEPGRAM_WS_CALLBACK_REASON: Final = "Deepgram callback delivery is not supported through the proxy: remove {params}" +async def deepgram_listen_user_api_key_auth(websocket: WebSocket) -> UserAPIKeyAuth: + return await user_api_key_auth_websocket_for_model( + websocket, model=deepgram_listen_requested_model(websocket.url.query) + ) + + @router.websocket("/deepgram/v1/listen") @router.websocket("/deepgram/listen") async def deepgram_listen_websocket_route( websocket: WebSocket, - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth_websocket)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(deepgram_listen_user_api_key_auth)], relay: Annotated[_WebsocketRelay, Depends(_websocket_relay)], ) -> None: deepgram_api_key: Final = passthrough_endpoint_router.get_credentials( diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py index 0a8d7ff3d33..fc0af400477 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -7,7 +7,9 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.deepgram.common_utils import ( + deepgram_listen_addon_pricing_models, deepgram_listen_audio_seconds, + deepgram_listen_base_pricing_models, deepgram_listen_channel_count, deepgram_listen_model, deepgram_listen_transcript, @@ -18,19 +20,39 @@ from litellm.types.utils import TranscriptionResponse DEEPGRAM_LISTEN_ROUTE_SUFFIX: Final = "/listen" -def _audio_cost(response: TranscriptionResponse, model: str) -> float | None: +def _registry_cost(response: TranscriptionResponse, pricing_model: str) -> float | None: try: return litellm.completion_cost( completion_response=response, - model=model, + model=pricing_model, custom_llm_provider=litellm.LlmProviders.DEEPGRAM.value, call_type="transcription", ) - except Exception as e: # noqa: BLE001 # an unpriced model must not lose the spend row, only its cost - verbose_proxy_logger.warning("Deepgram listen passthrough: no pricing for model '%s': %s", model, e) + except Exception as e: # noqa: BLE001 # an unpriced entry must not lose the spend row, only its cost + verbose_proxy_logger.debug("Deepgram listen passthrough: no registry price for '%s': %s", pricing_model, e) return None +def _audio_cost(response: TranscriptionResponse, upstream_url: str) -> float | None: + base_cost: Final = next( + ( + cost + for pricing_model in deepgram_listen_base_pricing_models(upstream_url) + if (cost := _registry_cost(response, pricing_model)) is not None + ), + None, + ) + if base_cost is None: + verbose_proxy_logger.warning( + "Deepgram listen passthrough: no pricing for model '%s'", deepgram_listen_model(upstream_url) + ) + return None + addon_costs: Final = tuple( + _registry_cost(response, pricing_model) for pricing_model in deepgram_listen_addon_pricing_models(upstream_url) + ) + return base_cost + sum(cost for cost in addon_costs if cost is not None) + + class DeepgramListenPassthroughLoggingHandler: @staticmethod def is_deepgram_listen_route(url_route: str) -> bool: @@ -50,7 +72,7 @@ class DeepgramListenPassthroughLoggingHandler: billed_seconds: Final = audio_seconds * channels response: Final = TranscriptionResponse(text=deepgram_listen_transcript(websocket_messages)) response._hidden_params["audio_transcription_duration"] = billed_seconds # pyright: ignore[reportPrivateUsage] # the cost calculator reads the billed duration off the response's hidden params - response_cost: Final = _audio_cost(response, model) + response_cost: Final = _audio_cost(response, upstream_url) response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads a precomputed cost off the response's hidden params provider: Final = litellm.LlmProviders.DEEPGRAM.value diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e87a3fec99b..7b7b3c7199c 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -20736,6 +20736,96 @@ "/v1/audio/transcriptions" ] }, + "deepgram/streaming/nova-3": { + "input_cost_per_second": 8e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0048/60 seconds = $0.00008000 per second", + "note": "Nova-3 monolingual streaming, pay as you go", + "original_pricing_per_minute": 0.0048 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/nova-3-multilingual": { + "input_cost_per_second": 9.667e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0058/60 seconds = $0.00009667 per second", + "note": "Nova-3 multilingual (language=multi) streaming, pay as you go", + "original_pricing_per_minute": 0.0058 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/redact": { + "input_cost_per_second": 3.333e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0020/60 seconds = $0.00003333 per second", + "note": "Redaction add-on (redact query param), streaming, pay as you go", + "original_pricing_per_minute": 0.002 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/keyterm": { + "input_cost_per_second": 2.167e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0013/60 seconds = $0.00002167 per second", + "note": "Keyterm Prompting add-on (keyterm query param), streaming, pay as you go", + "original_pricing_per_minute": 0.0013 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/detect_entities": { + "input_cost_per_second": 2.833e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0017/60 seconds = $0.00002833 per second", + "note": "Entity Detection add-on (detect_entities query param), streaming, pay as you go", + "original_pricing_per_minute": 0.0017 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, + "deepgram/streaming/diarize": { + "input_cost_per_second": 3.333e-05, + "litellm_provider": "deepgram", + "metadata": { + "calculation": "$0.0020/60 seconds = $0.00003333 per second", + "note": "Speaker Diarization add-on (diarize / diarize_model query params), streaming, pay as you go", + "original_pricing_per_minute": 0.002 + }, + "mode": "audio_transcription", + "output_cost_per_second": 0.0, + "source": "https://deepgram.com/pricing", + "supported_endpoints": [ + "/v1/listen" + ] + }, "deepgram/whisper": { "input_cost_per_second": 0.0001, "litellm_provider": "deepgram", diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index d7a38f2ca5a..802a3795ffe 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -6,10 +6,13 @@ import pytest import litellm from litellm.llms.deepgram.common_utils import ( + deepgram_listen_addon_pricing_models, deepgram_listen_audio_seconds, + deepgram_listen_base_pricing_models, deepgram_listen_callback_params, deepgram_listen_channel_count, deepgram_listen_model, + deepgram_listen_requested_model, deepgram_listen_transcript, deepgram_listen_websocket_target, ) @@ -195,3 +198,68 @@ def test_deepgram_listen_transcript_joins_final_results_only(): ) def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, expected_model: str): assert deepgram_listen_model(upstream_url) == expected_model + + +@pytest.mark.parametrize( + "query_string", + ["model=nova-2&language=en", "language=en", "model=&language=en", "", "model=nova-3-medical"], +) +def test_requested_model_is_the_model_the_upstream_target_will_carry(query_string: str): + """Authorization runs against ``deepgram_listen_requested_model``; the upstream URL is built separately, so the + two must always agree or a key could be authorized for one model and reach another.""" + target: Final = deepgram_listen_websocket_target(None, query_string) + assert deepgram_listen_requested_model(query_string) == deepgram_listen_model(target) + + +@pytest.mark.parametrize( + ("upstream_url", "expected"), + [ + pytest.param(NOVA_3_URL, ("streaming/nova-3", "nova-3"), id="monolingual"), + pytest.param(f"{NOVA_3_URL}&language=en", ("streaming/nova-3", "nova-3"), id="explicit language"), + pytest.param( + f"{NOVA_3_URL}&language=multi", + ("streaming/nova-3-multilingual", "streaming/nova-3", "nova-3"), + id="multilingual", + ), + pytest.param( + f"{NOVA_3_URL}&language=MULTI", + ("streaming/nova-3-multilingual", "streaming/nova-3", "nova-3"), + id="multilingual any case", + ), + pytest.param( + "wss://api.deepgram.com/v1/listen?model=nova-2&language=multi", + ("streaming/nova-2-multilingual", "streaming/nova-2", "nova-2"), + id="other model", + ), + ], +) +def test_deepgram_listen_base_pricing_models(upstream_url: str, expected: tuple[str, ...]): + assert deepgram_listen_base_pricing_models(upstream_url) == expected + + +@pytest.mark.parametrize( + ("upstream_url", "expected"), + [ + pytest.param(NOVA_3_URL, (), id="no add-ons"), + pytest.param(f"{NOVA_3_URL}&redact=pci", ("streaming/redact",), id="redact"), + pytest.param(f"{NOVA_3_URL}&redact=pci&redact=ssn", ("streaming/redact",), id="repeated redact once"), + pytest.param(f"{NOVA_3_URL}&keyterm=a&keyterm=b", ("streaming/keyterm",), id="keyterm"), + pytest.param(f"{NOVA_3_URL}&detect_entities=true", ("streaming/detect_entities",), id="detect_entities"), + pytest.param(f"{NOVA_3_URL}&diarize=true", ("streaming/diarize",), id="diarize"), + pytest.param(f"{NOVA_3_URL}&diarize_model=v1", ("streaming/diarize",), id="diarize_model"), + pytest.param(f"{NOVA_3_URL}&diarize=true&diarize_model=latest", ("streaming/diarize",), id="diarize both once"), + pytest.param(f"{NOVA_3_URL}&detect_entities=false&diarize=FALSE&redact=", (), id="disabled"), + pytest.param( + f"{NOVA_3_URL}&detect_entities=false&detect_entities=true", + ("streaming/detect_entities",), + id="any enabling value wins", + ), + pytest.param( + f"{NOVA_3_URL}&diarize=true&redact=pci&keyterm=x&detect_entities=true", + ("streaming/detect_entities", "streaming/diarize", "streaming/keyterm", "streaming/redact"), + id="all, sorted", + ), + ], +) +def test_deepgram_listen_addon_pricing_models(upstream_url: str, expected: tuple[str, ...]): + assert deepgram_listen_addon_pricing_models(upstream_url) == expected diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py index 944b37f95e3..a6133a7cf99 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -19,6 +19,8 @@ from litellm.types.utils import StandardLoggingPayload, TranscriptionResponse NOVA_3_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-3&encoding=linear16&sample_rate=16000" +pytestmark: Final = pytest.mark.usefixtures("local_model_cost_map") + def _results(start: object, duration: object, transcript: str = "", is_final: object = True) -> dict[str, object]: return { @@ -63,13 +65,22 @@ def _logging_obj(call_id: str = "call-dg") -> LiteLLMLoggingObj: ) -def _registry_cost(model: str, seconds: float) -> float: +def _registry_cost(pricing_model: str, seconds: float) -> float: """Derives the expected charge from the live cost map rather than pinning a vendor price.""" - per_second: Final = litellm.model_cost[f"deepgram/{model}"]["input_cost_per_second"] + per_second: Final = litellm.model_cost[f"deepgram/{pricing_model}"]["input_cost_per_second"] assert per_second > 0 return per_second * seconds +def _cost(upstream_url: str, *frames: dict[str, object]) -> float: + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=frames, logging_obj=_logging_obj(), upstream_url=upstream_url + ) + response_cost = handler_result["kwargs"]["response_cost"] + assert isinstance(response_cost, float) + return response_cost + + def test_handler_bills_metadata_duration_at_the_registry_rate_and_names_the_model(): frames = (_results(0.0, 5.0, "first sentence"), _results(5.0, 7.5, "second sentence"), _metadata(12.5)) logging_obj = _logging_obj() @@ -85,15 +96,78 @@ def test_handler_bills_metadata_duration_at_the_registry_rate_and_names_the_mode assert isinstance(result, TranscriptionResponse) assert result.text == "first sentence second sentence" assert result._hidden_params["audio_transcription_duration"] == 12.5 - assert result._hidden_params["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) - assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) + assert result._hidden_params["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5)) + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5)) assert handler_result["kwargs"]["model"] == "nova-3" assert handler_result["kwargs"]["custom_llm_provider"] == "deepgram" assert handler_result["kwargs"]["litellm_params"] == {"metadata": {}} assert logging_obj.model == "nova-3" assert logging_obj.model_call_details["model"] == "nova-3" assert logging_obj.model_call_details["custom_llm_provider"] == "deepgram" - assert logging_obj.model_call_details["response_cost"] == pytest.approx(_registry_cost("nova-3", 12.5)) + assert logging_obj.model_call_details["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 12.5)) + + +def test_handler_bills_streaming_not_prerecorded_rates(): + """Deepgram prices /v1/listen over a WebSocket separately from pre-recorded transcription, so the streaming entry + must be the one charged; the two registry rows only need to differ for this to matter, whatever their values.""" + streaming = litellm.model_cost["deepgram/streaming/nova-3"]["input_cost_per_second"] + prerecorded = litellm.model_cost["deepgram/nova-3"]["input_cost_per_second"] + assert streaming != prerecorded + + assert _cost(NOVA_3_URL, _metadata(60.0)) == pytest.approx(60.0 * streaming) + + +def test_handler_bills_multilingual_streaming_when_language_is_multi(): + monolingual = _cost(NOVA_3_URL, _metadata(60.0)) + multilingual = _cost(f"{NOVA_3_URL}&language=multi", _metadata(60.0)) + + assert multilingual == pytest.approx(_registry_cost("streaming/nova-3-multilingual", 60.0)) + assert multilingual > monolingual + + +@pytest.mark.parametrize( + ("query", "addons"), + [ + pytest.param("redact=pci", ("redact",), id="redaction"), + pytest.param("redact=pci&redact=numbers", ("redact",), id="redaction counted once"), + pytest.param("keyterm=LiteLLM&keyterm=Deepgram", ("keyterm",), id="keyterm prompting"), + pytest.param("detect_entities=true", ("detect_entities",), id="entity detection"), + pytest.param("diarize=true", ("diarize",), id="diarization"), + pytest.param("diarize_model=v1", ("diarize",), id="diarization via diarize_model"), + pytest.param("diarize=true&diarize_model=v1", ("diarize",), id="diarization counted once"), + pytest.param( + "redact=pci&keyterm=x&detect_entities=true&diarize=true", + ("redact", "keyterm", "detect_entities", "diarize"), + id="every add-on", + ), + pytest.param("detect_entities=false&diarize=False&redact=", (), id="disabled add-ons cost nothing"), + ], +) +def test_handler_adds_each_priced_add_on_once_on_top_of_the_base_rate(query: str, addons: tuple[str, ...]): + base = _cost(NOVA_3_URL, _metadata(60.0)) + expected = base + sum(_registry_cost(f"streaming/{addon}", 60.0) for addon in addons) + + assert _cost(f"{NOVA_3_URL}&{query}", _metadata(60.0)) == pytest.approx(expected) + + +def test_handler_add_ons_scale_with_channels_like_the_base_rate(): + stereo_plain = _cost(f"{NOVA_3_URL}&channels=2", _metadata(60.0, channels=2)) + stereo_redacted = _cost(f"{NOVA_3_URL}&channels=2&redact=pci", _metadata(60.0, channels=2)) + + assert stereo_redacted - stereo_plain == pytest.approx(_registry_cost("streaming/redact", 120.0)) + + +def test_handler_falls_back_to_the_prerecorded_rate_for_a_model_without_a_streaming_entry(): + assert "deepgram/streaming/nova-2" not in litellm.model_cost + + handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( + websocket_messages=(_metadata(60.0),), + logging_obj=_logging_obj(), + upstream_url="wss://api.deepgram.com/v1/listen?model=nova-2", + ) + + assert handler_result["kwargs"]["model"] == "nova-2" + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-2", 60.0)) def test_handler_falls_back_to_results_frames_when_the_stream_ends_without_metadata(): @@ -104,7 +178,7 @@ def test_handler_falls_back_to_results_frames_when_the_stream_ends_without_metad ) assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 72.5 - assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 72.5)) + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 72.5)) def test_handler_charges_more_for_more_audio_on_the_same_model(): @@ -133,7 +207,7 @@ def test_handler_bills_every_channel_of_a_multichannel_session(): assert stereo["result"]._hidden_params["audio_transcription_duration"] == 60.0 assert stereo["kwargs"]["response_cost"] == pytest.approx(2 * mono["kwargs"]["response_cost"]) - assert stereo["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 60.0)) + assert stereo["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 60.0)) def test_handler_bills_the_declared_channels_when_the_stream_dies_before_any_frame_reports_them(): @@ -144,7 +218,7 @@ def test_handler_bills_the_declared_channels_when_the_stream_dies_before_any_fra ) assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 30.0 - assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-3", 30.0)) + assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 30.0)) def test_handler_keeps_the_spend_row_but_no_cost_for_an_unpriced_model(): @@ -226,6 +300,6 @@ async def test_success_handler_dispatches_deepgram_listen_and_logs_duration_base payload = capturing_logger.payloads[0] assert payload["model"] == "nova-3" assert payload["custom_llm_provider"] == "deepgram" - assert payload["response_cost"] == pytest.approx(_registry_cost("nova-3", 20.0)) + assert payload["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 20.0)) assert payload["metadata"]["user_api_key_team_id"] == "team-stt" assert payload["id"] == "call-dg-e2e" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 5ea2b0b8ab9..1c21b75804e 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -1,10 +1,11 @@ """Deepgram ``/v1/listen`` passthrough WebSocket route: registration, auth, credential injection, target URL.""" +import asyncio from collections.abc import Mapping from dataclasses import dataclass from types import MappingProxyType, SimpleNamespace from typing import Final -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import FastAPI @@ -12,13 +13,16 @@ from fastapi.testclient import TestClient from starlette.routing import WebSocketRoute from starlette.websockets import WebSocketDisconnect +from litellm.caching.dual_cache import DualCache from litellm.proxy._lazy_features import LAZY_FEATURES from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import _cache_key_object from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( _websocket_relay, deepgram_listen_websocket_route, router, ) +from litellm.proxy.utils import hash_token GET_CREDENTIALS: Final = ( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" @@ -302,6 +306,62 @@ def test_deepgram_listen_authenticates_the_litellm_key_and_relays_to_deepgram(mo ] +async def _cache_restricted_key(virtual_key: str, models: list[str]) -> DualCache: + cache = DualCache() + await _cache_key_object( + hashed_token=hash_token(virtual_key), + user_api_key_obj=UserAPIKeyAuth(token=hash_token(virtual_key), models=models), + user_api_key_cache=cache, + proxy_logging_obj=None, + ) + return cache + + +@pytest.mark.parametrize( + ("query", "expect_relay"), + [ + pytest.param("model=nova-2", True, id="allowed model named"), + pytest.param("model=nova-3", False, id="denied model named"), + pytest.param("", False, id="model omitted, default denied"), + pytest.param("model=&language=en", False, id="model blank, default denied"), + ], +) +def test_deepgram_listen_authorizes_the_model_it_will_actually_send_upstream(query, expect_relay, monkeypatch): + """A key allowed only ``nova-2`` must not reach ``nova-3`` by leaving ``model`` out and letting the proxy fill + in its default: the real key auth path must see the same model the upstream target will carry.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"])) + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + + with ( + patch(GET_CREDENTIALS, return_value="dg-provider-key"), + patch.multiple( # test-quality-ok: the real key auth path reads these proxy_server globals and has no injection seam + "litellm.proxy.proxy_server", + master_key="sk-master", + prisma_client=MagicMock(), + user_api_key_cache=cache, + llm_model_list=None, + llm_router=None, + ), + ): + if expect_relay: + with client.websocket_connect( + f"/deepgram/v1/listen?{query}", headers={"Authorization": "Bearer sk-only-nova-2"} + ): + pass + assert [call.target for call in relay.calls] == [f"wss://api.deepgram.com/v1/listen?{query}"] + return + with pytest.raises(WebSocketDisconnect) as disconnect: + with client.websocket_connect( + f"/deepgram/v1/listen?{query}", headers={"Authorization": "Bearer sk-only-nova-2"} + ): + pass + + assert disconnect.value.code == 1008 + assert relay.calls == [] + + def test_deepgram_listen_echoes_the_browser_subprotocol_that_carries_the_litellm_key(monkeypatch): """Browsers cannot set headers, so they send the key as a subprotocol and abort the handshake unless the server echoes that subprotocol back; the key itself must still stay off the upstream connection.""" From 3084d2af31794f552d1584558fc1efadfdcc7dc9 Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 20:51:50 +0000 Subject: [PATCH 100/253] test(utils): allow /v1/listen in the registry supported_endpoints schema The deepgram/streaming/* rows added for the Deepgram WebSocket passthrough declare /v1/listen as their endpoint, so the registry validation test needs it in the enum, the same way /vertex_ai/live was added for that passthrough Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/test_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index f219f26b353..77336836e04 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -945,6 +945,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "/v1/audio/speech", "/v1/ocr", "/vertex_ai/live", + "/v1/listen", "/v1beta/interactions", ], }, From 6e7c3f68a1f3cebe75c2e019490db5cfb143f56f Mon Sep 17 00:00:00 2001 From: yassin Date: Thu, 17 Sep 2026 21:06:07 +0000 Subject: [PATCH 101/253] test(deepgram): pin litellm.max_budget to zero in the model authorization route test Under xdist the per-test litellm reload is skipped, so a leaked max_budget from another proxy test sent the real key auth path into the global spend lookup, which the MagicMock prisma client cannot await Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_deepgram_ws_passthrough_routes.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 1c21b75804e..85baaacf1e9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -13,6 +13,7 @@ from fastapi.testclient import TestClient from starlette.routing import WebSocketRoute from starlette.websockets import WebSocketDisconnect +import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy._lazy_features import LAZY_FEATURES from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth @@ -330,6 +331,7 @@ def test_deepgram_listen_authorizes_the_model_it_will_actually_send_upstream(que """A key allowed only ``nova-2`` must not reach ``nova-3`` by leaving ``model`` out and letting the proxy fill in its default: the real key auth path must see the same model the upstream target will carry.""" monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + monkeypatch.setattr(litellm, "max_budget", 0.0) cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"])) relay = _FakeRelay() client = TestClient(_app_with_relay(relay)) From 7e5b3b49d46ae1773e3f0450adbda334c198f108 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 22:33:54 +0000 Subject: [PATCH 102/253] fix(proxy): treat non-positive project max_budget as unbudgeted 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 | 23 +++++++++++++++++++ 2 files changed, 24 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e80a4a677ce..e6313790380 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -5612,7 +5612,7 @@ async def _project_max_budget_check( if project_object.litellm_budget_table is not None: max_budget = project_object.litellm_budget_table.max_budget - if max_budget is None or not math.isfinite(max_budget): + if max_budget is None or max_budget <= 0 or not math.isfinite(max_budget): return from litellm.proxy.proxy_server import get_current_spend diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 56ee0aac249..08ed1575fdd 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7378,6 +7378,29 @@ async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "project_budget" +@pytest.mark.asyncio +@pytest.mark.parametrize("max_budget", [0.0, -1.0]) +async def test_project_max_budget_check_treats_non_positive_budget_as_unbudgeted(max_budget): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.auth.auth_checks import _project_max_budget_check + + real_spend_counter_cache = DualCache() + real_spend_counter_cache.in_memory_cache.set_cache(key="spend:project:p-budget", value=12.5) + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + + with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock + "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache + ): + await _project_max_budget_check( + project_object=_project_with_budget(spend=12.5, max_budget=max_budget), + valid_token=UserAPIKeyAuth(api_key="hashed-key", project_id="p-budget"), + proxy_logging_obj=proxy_logging_obj, + ) + + proxy_logging_obj.budget_alerts.assert_not_awaited() + + def test_is_user_proxy_admin_rejects_view_only_admin(): """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an Admin Viewer answering True here would gain every write route. Read parity for From 1c0342332beb96ba3c13ede34563c7e1ccd3c8a2 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 22:40:07 +0000 Subject: [PATCH 103/253] refactor(proxy): keep project spend enqueue within lint ceilings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 67 ++++++++-------------- 1 file changed, 25 insertions(+), 42 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 4f04a1b3a1a..e0985d5b5ae 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -706,17 +706,11 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) - try: - await self._update_project_db( - response_cost=response_cost, - project_id=project_id, - prisma_client=prisma_client, - ) - except Exception: - verbose_proxy_logger.debug( - "_batch_database_updates: _update_project_db failed: %s", - traceback.format_exc(), - ) + await self._update_project_db( + response_cost=response_cost, + project_id=project_id, + prisma_client=prisma_client, + ) try: await self._update_tag_db( @@ -985,27 +979,16 @@ class DBSpendUpdateWriter: response_cost: float | None, project_id: str | None, prisma_client: PrismaClient | None, - ): - try: - if project_id is None or prisma_client is None: - return - - await self.spend_update_queue.add_update( - update=SpendUpdateQueueItem( - entity_type=Litellm_EntityType.PROJECT, - entity_id=project_id, - response_cost=response_cost, - ) + ) -> None: + if project_id is None or prisma_client is None: + return + await self.spend_update_queue.add_update( + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.PROJECT, + entity_id=project_id, + response_cost=response_cost, ) - except Exception as e: - spend_log_error( - "Spend tracking - failed to enqueue project spend update. project_id=%s, response_cost=%s - %s", - project_id, - response_cost, - str(e), - exc=e, - ) - raise e + ) async def _update_agent_db( self, @@ -1246,17 +1229,17 @@ class DBSpendUpdateWriter: "Spend tracking - committing spend updates from Redis to DB: " "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, " "projects=%d, tags=%d, agents=%d, model_access_groups=%d", - len(db_spend_update_transactions.get("key_list_transactions") or {}), - len(db_spend_update_transactions.get("user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_list_transactions") or {}), - len(db_spend_update_transactions.get("org_list_transactions") or {}), - len(db_spend_update_transactions.get("end_user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_member_list_transactions") or {}), - len(db_spend_update_transactions.get("org_member_list_transactions") or {}), - len(db_spend_update_transactions.get("project_list_transactions") or {}), - len(db_spend_update_transactions.get("tag_list_transactions") or {}), - len(db_spend_update_transactions.get("agent_list_transactions") or {}), - len(db_spend_update_transactions.get("model_access_group_list_transactions") or {}), + len(db_spend_update_transactions.get("key_list_transactions") or ()), + len(db_spend_update_transactions.get("user_list_transactions") or ()), + len(db_spend_update_transactions.get("team_list_transactions") or ()), + len(db_spend_update_transactions.get("org_list_transactions") or ()), + len(db_spend_update_transactions.get("end_user_list_transactions") or ()), + len(db_spend_update_transactions.get("team_member_list_transactions") or ()), + len(db_spend_update_transactions.get("org_member_list_transactions") or ()), + len(db_spend_update_transactions.get("project_list_transactions") or ()), + len(db_spend_update_transactions.get("tag_list_transactions") or ()), + len(db_spend_update_transactions.get("agent_list_transactions") or ()), + len(db_spend_update_transactions.get("model_access_group_list_transactions") or ()), ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, From 70164dabd2795613a814f3e6d76c344e5f5ce618 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 23:20:10 +0000 Subject: [PATCH 104/253] fix(proxy): isolate project spend enqueue failures from sibling spend writes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 16 +++++++---- .../proxy/db/test_db_spend_update_writer.py | 27 +++++++++++++++++++ 2 files changed, 38 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e0985d5b5ae..2d6c7a0a74c 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -706,11 +706,17 @@ class DBSpendUpdateWriter: traceback.format_exc(), ) - await self._update_project_db( - response_cost=response_cost, - project_id=project_id, - prisma_client=prisma_client, - ) + try: + await self._update_project_db( + response_cost=response_cost, + project_id=project_id, + prisma_client=prisma_client, + ) + except Exception: # noqa: BLE001 # a project enqueue failure must not skip the sibling spend writes + verbose_proxy_logger.debug( + "_batch_database_updates: _update_project_db failed: %s", + traceback.format_exc(), + ) try: await self._update_tag_db( diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index f9494cd6332..9a2b15d8a30 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1175,6 +1175,33 @@ async def test_batch_database_updates_without_project_id_touches_no_project_row( assert transactions["project_list_transactions"] == {} +@pytest.mark.asyncio +async def test_project_enqueue_failure_does_not_stop_sibling_spend_updates(): + db_writer: Final = DBSpendUpdateWriter() + db_writer._update_project_db = AsyncMock(side_effect=RuntimeError("project queue boom")) + db_writer._update_tag_db = AsyncMock() + db_writer._update_agent_db = AsyncMock() + db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock() + + await db_writer._batch_database_updates( + response_cost=0.1, + user_id="u1", + hashed_token="t1", + team_id="team-1", + org_id=None, + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.1, "request_tags": ["t"]}, + project_id="proj-1", + ) + + db_writer._update_project_db.assert_awaited_once() + db_writer._update_tag_db.assert_awaited_once() + db_writer._update_agent_db.assert_awaited_once() + db_writer.add_spend_log_transaction_to_daily_user_transaction.assert_awaited_once() + + @pytest.mark.asyncio async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_id(): """ From e4778cd2d1d14a5efed40546031d474fe53eab63 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 17 Sep 2026 16:24:51 -0700 Subject: [PATCH 105/253] test(proxy): drive the real spend queue in the project enqueue isolation test Substitutes a queue that rejects project items instead of replacing writer methods with mocks, so the test asserts the sibling key, team and tag spend actually landed in the queue rather than that a mock was awaited. --- .../proxy/db/test_db_spend_update_writer.py | 31 +++++++++++-------- 1 file changed, 18 insertions(+), 13 deletions(-) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 9a2b15d8a30..3d61b1dba12 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -15,8 +15,9 @@ import pytest from redis.exceptions import DataError import litellm -from litellm.proxy._types import Litellm_EntityType +from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) @@ -1176,30 +1177,34 @@ async def test_batch_database_updates_without_project_id_touches_no_project_row( @pytest.mark.asyncio -async def test_project_enqueue_failure_does_not_stop_sibling_spend_updates(): +async def test_failed_project_enqueue_does_not_drop_the_rest_of_the_batch(): + class _ProjectRejectingQueue(SpendUpdateQueue): + async def add_update(self, update: SpendUpdateQueueItem): + if update.get("entity_type") is Litellm_EntityType.PROJECT: + raise RuntimeError("project enqueue boom") + await super().add_update(update) + db_writer: Final = DBSpendUpdateWriter() - db_writer._update_project_db = AsyncMock(side_effect=RuntimeError("project queue boom")) - db_writer._update_tag_db = AsyncMock() - db_writer._update_agent_db = AsyncMock() - db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock() + db_writer.spend_update_queue = _ProjectRejectingQueue() await db_writer._batch_database_updates( - response_cost=0.1, + response_cost=0.25, user_id="u1", hashed_token="t1", team_id="team-1", - org_id=None, + org_id="org-1", end_user_id=None, prisma_client=MagicMock(), litellm_proxy_budget_name=None, - payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.1, "request_tags": ["t"]}, + payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25, "request_tags": ["tag-1"]}, project_id="proj-1", ) - db_writer._update_project_db.assert_awaited_once() - db_writer._update_tag_db.assert_awaited_once() - db_writer._update_agent_db.assert_awaited_once() - db_writer.add_spend_log_transaction_to_daily_user_transaction.assert_awaited_once() + transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() + assert transactions["project_list_transactions"] == {} + assert transactions["tag_list_transactions"] == {"tag-1": 0.25} + assert transactions["key_list_transactions"] == {"t1": 0.25} + assert transactions["team_list_transactions"] == {"team-1": 0.25} @pytest.mark.asyncio From 89289d4f7753d3d780ce0cf3b679c409220c2b88 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 17 Sep 2026 16:26:26 -0700 Subject: [PATCH 106/253] fix(proxy): surface a failed project spend enqueue at error level The call-site guard kept a project enqueue failure from skipping the sibling spend writes, but logged it at debug only, so a dropped project charge was invisible on a default log level. Log it through spend_log_error inside the helper and re-raise, the way the org and agent helpers already do. --- litellm/proxy/db/db_spend_update_writer.py | 22 +++++++++---- .../proxy/db/test_db_spend_update_writer.py | 33 +++++++++++-------- 2 files changed, 36 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 2d6c7a0a74c..f715897b314 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -988,13 +988,23 @@ class DBSpendUpdateWriter: ) -> None: if project_id is None or prisma_client is None: return - await self.spend_update_queue.add_update( - update=SpendUpdateQueueItem( - entity_type=Litellm_EntityType.PROJECT, - entity_id=project_id, - response_cost=response_cost, + try: + await self.spend_update_queue.add_update( + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.PROJECT, + entity_id=project_id, + response_cost=response_cost, + ) ) - ) + except Exception as e: + spend_log_error( + "Spend tracking - failed to enqueue project spend update. project_id=%s, response_cost=%s - %s", + project_id, + response_cost, + str(e), + exc=e, + ) + raise e async def _update_agent_db( self, diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 3d61b1dba12..2efc0d99b8a 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1,6 +1,7 @@ import asyncio import copy import json +import logging import re @@ -15,6 +16,7 @@ import pytest from redis.exceptions import DataError import litellm +from litellm._logging import verbose_proxy_logger from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue @@ -1177,7 +1179,9 @@ async def test_batch_database_updates_without_project_id_touches_no_project_row( @pytest.mark.asyncio -async def test_failed_project_enqueue_does_not_drop_the_rest_of_the_batch(): +async def test_failed_project_enqueue_is_reported_and_does_not_drop_the_rest_of_the_batch( + caplog: pytest.LogCaptureFixture, +): class _ProjectRejectingQueue(SpendUpdateQueue): async def add_update(self, update: SpendUpdateQueueItem): if update.get("entity_type") is Litellm_EntityType.PROJECT: @@ -1187,18 +1191,21 @@ async def test_failed_project_enqueue_does_not_drop_the_rest_of_the_batch(): db_writer: Final = DBSpendUpdateWriter() db_writer.spend_update_queue = _ProjectRejectingQueue() - await db_writer._batch_database_updates( - response_cost=0.25, - user_id="u1", - hashed_token="t1", - team_id="team-1", - org_id="org-1", - end_user_id=None, - prisma_client=MagicMock(), - litellm_proxy_budget_name=None, - payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25, "request_tags": ["tag-1"]}, - project_id="proj-1", - ) + with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name): + await db_writer._batch_database_updates( + response_cost=0.25, + user_id="u1", + hashed_token="t1", + team_id="team-1", + org_id="org-1", + end_user_id=None, + prisma_client=MagicMock(), + litellm_proxy_budget_name=None, + payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25, "request_tags": ["tag-1"]}, + project_id="proj-1", + ) + + assert any("proj-1" in record.getMessage() for record in caplog.records if record.levelno >= logging.ERROR) transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() assert transactions["project_list_transactions"] == {} From ea37596b884056b2c818f8945d71e394b5126c7d Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 01:01:26 +0000 Subject: [PATCH 107/253] feat(mcp): let proxy admins force-close live MCP sessions and revoke stored user credentials Adds an admin-only DELETE /v1/mcp/sessions that terminates stateful MCP gateway sessions on the current worker by session id prefix and/or by the LiteLLM user that opened them, tombstones the terminated ids so a client reusing one gets 404 instead of a silently recreated stateless session, and lets PROXY_ADMIN name a user_id on the BYOK and OAuth credential delete routes. Full and view-only admins can list every user's stored credential metadata for a server (never the secret). The dashboard gains Disconnect controls on the Live Connections tab and a User Credentials tab with Revoke controls, both hidden from read-only admins. Resolves LIT-8001 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 2 + litellm/proxy/_experimental/mcp_server/db.py | 32 ++ .../proxy/_experimental/mcp_server/server.py | 72 +++- litellm/proxy/_lazy_openapi_snapshot.json | 237 ++++++++++- litellm/proxy/_types.py | 10 + .../mcp_management_endpoints.py | 129 +++++- litellm/types/mcp.py | 8 + .../mcp_server/test_db_credentials.py | 47 +++ .../mcp_server/test_mcp_server.py | 177 ++++++++ .../test_mcp_management_endpoints.py | 398 ++++++++++++++++++ ...MCPGatewaySessionsTab.integration.test.tsx | 90 +++- .../_components/MCPGatewaySessionsTab.tsx | 146 ++++++- ...rUserCredentialsPanel.integration.test.tsx | 102 +++++ .../MCPServerUserCredentialsPanel.tsx | 212 ++++++++++ .../_components/mcp_server_view.tsx | 21 + .../mcp-servers/_components/mcp_servers.tsx | 4 +- .../src/components/mcp_tools/types.tsx | 20 + .../src/components/networking.tsx | 36 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 132 +++++- 19 files changed, 1826 insertions(+), 49 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.tsx diff --git a/litellm/constants.py b/litellm/constants.py index d4827bb7483..2ebc9beb632 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -183,6 +183,8 @@ MCP_TOOL_LISTING_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_TOOL_LISTING_TIME MCP_METADATA_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "10.0")) MCP_HEALTH_CHECK_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0")) MCP_TOOL_LISTING_MAX_PAGES: Final = 1000 +MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH: Final = 8 +MCP_ADMIN_TERMINATED_SESSION_IDS_MAX: Final = 1024 # Allowlist of commands permitted for MCP stdio transport. # Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation. diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 789b2ffaef4..a04e2f5c9b8 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -24,6 +24,7 @@ from litellm.proxy._types import ( MCPApprovalStatus, MCPEnvVar, MCPEnvVarScope, + MCPServerUserCredentialListItem, MCPSubmissionsSummary, NewMCPServerRequest, SpecialMCPServerName, @@ -1504,6 +1505,37 @@ async def get_user_oauth_credential( return _parse_oauth_payload(decoded) +def _server_user_credential_item( + row: "prisma_db_models.LiteLLM_MCPUserCredentials", +) -> MCPServerUserCredentialListItem: + oauth_payload: Final = _decode_oauth_payload(row.credential_b64) + if oauth_payload is None: + return MCPServerUserCredentialListItem( + user_id=row.user_id, + credential_type="byok", + updated_at=row.updated_at.isoformat(), + ) + return MCPServerUserCredentialListItem( + user_id=row.user_id, + credential_type="oauth2", + expires_at=oauth_payload.get("expires_at"), + connected_at=oauth_payload.get("connected_at"), + updated_at=row.updated_at.isoformat(), + ) + + +async def list_server_user_credentials( + prisma_client: PrismaClient, + server_id: str, +) -> tuple[MCPServerUserCredentialListItem, ...]: + """Every user's stored credential for one server, typed but without the secret, for admins.""" + rows: Final = await _db_find_user_credential_rows( + prisma_client, + {"server_id": server_id}, # mutable-ok: prisma where-inputs must be plain dicts + ) + return tuple(_server_user_credential_item(row) for row in rows) + + async def list_user_oauth_credentials( prisma_client: PrismaClient, user_id: str, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 524bac747ad..80d9274859d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -14,7 +14,7 @@ import time import traceback import types import uuid -from collections import Counter +from collections import Counter, deque from collections.abc import AsyncIterator, Callable, Mapping, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol @@ -28,7 +28,11 @@ from starlette.types import Message, Receive, Scope, Send from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_logger -from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG +from litellm.constants import ( + MAXIMUM_TRACEBACK_LINES_TO_LOG, + MCP_ADMIN_TERMINATED_SESSION_IDS_MAX, + MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH, +) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, @@ -91,6 +95,7 @@ from litellm.types.mcp import ( MCPGatewaySession, MCPGatewaySessionGroupCount, MCPGatewaySessionsResponse, + MCPGatewaySessionsTerminateResponse, MCPSpecVersion, ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer @@ -618,6 +623,9 @@ if MCP_AVAILABLE: _stateful_session_locks: Final[dict[str, asyncio.Lock]] = {} _stateful_session_active_request_counts: Final[dict[str, int]] = {} _stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown + _admin_terminated_session_ids: Final[deque[str]] = deque( # mutable-ok: bounded ring, appended on admin termination + maxlen=MCP_ADMIN_TERMINATED_SESSION_IDS_MAX + ) class _TerminableTransport(Protocol): async def terminate(self) -> None: ... @@ -3850,7 +3858,7 @@ if MCP_AVAILABLE: client_info: Final = _stateful_session_client_info.get(session_id) key_auth: Final = auth_user.user_api_key_auth return MCPGatewaySession( - session_id_prefix=session_id[:8], + session_id_prefix=session_id[:MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH], client_name=client_info.name if client_info is not None else None, client_version=client_info.version if client_info is not None else None, user_id=key_auth.user_id if key_auth is not None else None, @@ -3885,6 +3893,53 @@ if MCP_AVAILABLE: sessions=sessions, ) + def _session_matches_admin_selector( + session_id: str, + auth_user: MCPAuthenticatedUser, + session_id_prefix: str | None, + user_id: str | None, + ) -> bool: + if session_id_prefix is not None and not session_id.startswith(session_id_prefix): + return False + if user_id is None: + return True + key_auth: Final = auth_user.user_api_key_auth + return key_auth is not None and key_auth.user_id == user_id + + async def terminate_mcp_gateway_sessions( + *, + session_id_prefix: str | None = None, + user_id: str | None = None, + ) -> MCPGatewaySessionsTerminateResponse: + """Force-close every live stateful session on this worker matching the selector. + + The transport is terminated (open streams close), all per-session + tracking is dropped, and the id is remembered so a client that keeps + sending it receives 404 and has to ``initialize`` again, which re-runs + admission. Only sessions held by this worker process are affected. + """ + now: Final = time.monotonic() + server_instances: Final = _stateful_server_instances() + targets: Final = tuple( + (session_id, auth_user) + for session_id, auth_user in tuple(_stateful_session_auth_contexts.items()) + if session_id in server_instances + and _session_matches_admin_selector(session_id, auth_user, session_id_prefix, user_id) + ) + terminated: Final = tuple(_gateway_session_for(session_id, auth_user, now) for session_id, auth_user in targets) + for session_id, _ in targets: + _admin_terminated_session_ids.append(session_id) + transport = server_instances.pop(session_id, None) + _remove_stateful_session_tracking(session_id) + if transport is not None: + await transport.terminate() + verbose_logger.warning("MCP session '%s' terminated by an administrator.", session_id) + return MCPGatewaySessionsTerminateResponse( + worker_pid=os.getpid(), + terminated_sessions=len(terminated), + sessions=terminated, + ) + async def _read_request_body_for_routing( receive: Receive, ) -> tuple[list[Message], bytes]: @@ -4009,6 +4064,17 @@ if MCP_AVAILABLE: await success_response(scope, receive, send) return True + if _session_id in _admin_terminated_session_ids: + terminated_response: Final = JSONResponse( + status_code=404, + content={ # mutable-ok: JSONResponse content must be a plain dict + "error": "Not Found", + "details": "mcp-session-id was terminated by an administrator. Send initialize to start a new session.", + }, + ) + await terminated_response(scope, receive, send) + return True + # Non-DELETE: strip stale session ID to allow new session creation verbose_logger.warning( "MCP session ID '%s' not found in this worker's memory. " diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 82b709feb92..1e9b1f5b743 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -27648,6 +27648,32 @@ "title": "MCPGatewaySessionsResponse", "type": "object" }, + "MCPGatewaySessionsTerminateResponse": { + "description": "Stateful sessions an administrator force-closed on this proxy worker.", + "properties": { + "sessions": { + "items": { + "$ref": "#/components/schemas/MCPGatewaySession" + }, + "title": "Sessions", + "type": "array" + }, + "terminated_sessions": { + "title": "Terminated Sessions", + "type": "integer" + }, + "worker_pid": { + "title": "Worker Pid", + "type": "integer" + } + }, + "required": [ + "worker_pid", + "terminated_sessions" + ], + "title": "MCPGatewaySessionsTerminateResponse", + "type": "object" + }, "MCPOAuthUserCredentialRequest": { "description": "Stores a user's OAuth2 token for an OpenAPI MCP server.", "properties": { @@ -27744,6 +27770,56 @@ "title": "MCPOAuthUserCredentialStatus", "type": "object" }, + "MCPServerUserCredentialListItem": { + "description": "One user's stored credential for an MCP server, as an admin sees it. Never carries the secret.", + "properties": { + "connected_at": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Connected At" + }, + "credential_type": { + "enum": [ + "oauth2", + "byok" + ], + "title": "Credential Type", + "type": "string" + }, + "expires_at": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Expires At" + }, + "updated_at": { + "title": "Updated At", + "type": "string" + }, + "user_id": { + "title": "User Id", + "type": "string" + } + }, + "required": [ + "user_id", + "credential_type", + "updated_at" + ], + "title": "MCPServerUserCredentialListItem", + "type": "object" + }, "MCPSubmissionsSummary": { "properties": { "active": { @@ -29920,7 +29996,7 @@ }, "/v1/mcp/server/{server_id}/oauth-user-credential": { "delete": { - "description": "Revoke the calling user's stored OAuth2 token for an MCP server", + "description": "Revoke the calling user's stored OAuth2 token for an MCP server. A proxy admin may pass user_id to revoke another user's stored token.", "operationId": "delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete", "parameters": [ { @@ -29931,6 +30007,23 @@ "title": "Server Id", "type": "string" } + }, + { + "in": "query", + "name": "user_id", + "required": false, + "schema": { + "anyOf": [ + { + "minLength": 1, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "User Id" + } } ], "responses": { @@ -30130,7 +30223,7 @@ }, "/v1/mcp/server/{server_id}/user-credential": { "delete": { - "description": "Delete the calling user's stored API key for a BYOK MCP server", + "description": "Delete the calling user's stored API key for a BYOK MCP server. A proxy admin may pass user_id to revoke another user's stored key.", "operationId": "delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete", "parameters": [ { @@ -30141,6 +30234,23 @@ "title": "Server Id", "type": "string" } + }, + { + "in": "query", + "name": "user_id", + "required": false, + "schema": { + "anyOf": [ + { + "minLength": 1, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "User Id" + } } ], "responses": { @@ -30232,6 +30342,58 @@ ] } }, + "/v1/mcp/server/{server_id}/user-credentials": { + "get": { + "description": "List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets)", + "operationId": "list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get", + "parameters": [ + { + "in": "path", + "name": "server_id", + "required": true, + "schema": { + "title": "Server Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "items": { + "$ref": "#/components/schemas/MCPServerUserCredentialListItem" + }, + "title": "Response List Mcp Server User Credentials V1 Mcp Server Server Id User Credentials Get", + "type": "array" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "List Mcp Server User Credentials", + "tags": [ + "mcp_management" + ] + } + }, "/v1/mcp/server/{server_id}/user-env-vars": { "delete": { "description": "Clear the calling user's per-user MCP env var values for this server.", @@ -30383,6 +30545,77 @@ } }, "/v1/mcp/sessions": { + "delete": { + "description": "Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix and/or by the LiteLLM user that opened them (proxy admin only).", + "operationId": "delete_mcp_gateway_sessions_v1_mcp_sessions_delete", + "parameters": [ + { + "in": "query", + "name": "session_id_prefix", + "required": false, + "schema": { + "anyOf": [ + { + "minLength": 8, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Session Id Prefix" + } + }, + { + "in": "query", + "name": "user_id", + "required": false, + "schema": { + "anyOf": [ + { + "minLength": 1, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "User Id" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/MCPGatewaySessionsTerminateResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Delete Mcp Gateway Sessions", + "tags": [ + "mcp_management" + ] + }, "get": { "description": "Live stateful MCP gateway sessions on this proxy worker, grouped by AI client and by user.", "operationId": "get_mcp_gateway_sessions_v1_mcp_sessions_get", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index ff32d4784df..ee92e64b046 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1716,6 +1716,16 @@ class MCPUserCredentialListItem(LiteLLMPydanticObjectBase): connected_at: str | None = None # ISO-8601 +class MCPServerUserCredentialListItem(LiteLLMPydanticObjectBase): + """One user's stored credential for an MCP server, as an admin sees it. Never carries the secret.""" + + user_id: str + credential_type: Literal["oauth2", "byok"] + expires_at: str | None = None + connected_at: str | None = None + updated_at: str + + class MCPUserEnvVarsRequest(LiteLLMPydanticObjectBase): """Payload for storing the calling user's per-user env var values.""" diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 82a1cdcdd00..728ba9568e8 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -47,7 +47,7 @@ except ImportError: import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._uuid import uuid -from litellm.constants import LITELLM_PROXY_ADMIN_NAME +from litellm.constants import LITELLM_PROXY_ADMIN_NAME, MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, @@ -145,6 +145,7 @@ if MCP_AVAILABLE: get_user_env_vars, get_user_env_vars_bulk, get_user_oauth_credential, + list_server_user_credentials, list_user_oauth_credentials, mcp_oauth_token_identity, merge_user_env_vars, @@ -180,6 +181,7 @@ if MCP_AVAILABLE: MCPApprovalStatus, MCPOAuthUserCredentialRequest, MCPOAuthUserCredentialStatus, + MCPServerUserCredentialListItem, MCPSubmissionsSummary, MCPTransport, MCPUserCredentialListItem, @@ -221,6 +223,7 @@ if MCP_AVAILABLE: MCPAuth, MCPCredentials, MCPGatewaySessionsResponse, + MCPGatewaySessionsTerminateResponse, normalize_upstream_header_name, ) from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -662,6 +665,31 @@ if MCP_AVAILABLE: """ return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + def _resolve_credential_target_user_id(user_api_key_dict: UserAPIKeyAuth, requested_user_id: str | None) -> str: + """The user whose stored MCP credential a request acts on. + + Defaults to the caller. Naming another user is a revocation and needs + ``PROXY_ADMIN``; a read-only admin or a regular user gets 403. + """ + caller_user_id: Final = user_api_key_dict.user_id or "" + if requested_user_id is not None and requested_user_id != caller_user_id: + if not _user_is_full_admin(user_api_key_dict): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict + "error": "Proxy admin access required to revoke another user's MCP credential.", + }, + ) + return requested_user_id + if not caller_user_id: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "User ID not found in token" + }, # mutable-ok: FastAPI HTTPException detail requires a plain dict + ) + return caller_user_id + def _is_restricted_virtual_key_request(user_api_key_dict: UserAPIKeyAuth) -> bool: """Best-effort detection for route-restricted virtual keys. @@ -1373,6 +1401,41 @@ if MCP_AVAILABLE: return get_mcp_gateway_sessions_report() + @router.delete( + "/sessions", + description=( + "Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix " + "and/or by the LiteLLM user that opened them (proxy admin only)." + ), + dependencies=(Depends(user_api_key_auth),), + response_model=MCPGatewaySessionsTerminateResponse, + ) + @management_endpoint_wrapper + async def delete_mcp_gateway_sessions( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + session_id_prefix: Annotated[str | None, Query(min_length=MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH)] = None, + user_id: Annotated[str | None, Query(min_length=1)] = None, + ) -> MCPGatewaySessionsTerminateResponse: + if not _user_is_full_admin(user_api_key_dict): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict + "error": "Proxy admin access required to terminate MCP gateway sessions.", + }, + ) + if session_id_prefix is None and user_id is None: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict + "error": "Provide session_id_prefix and/or user_id to select the sessions to terminate.", + }, + ) + from litellm.proxy._experimental.mcp_server.server import ( + terminate_mcp_gateway_sessions, + ) + + return await terminate_mcp_gateway_sessions(session_id_prefix=session_id_prefix, user_id=user_id) + @router.get( "/server/submissions", description="Returns all MCP servers submitted by non-admin users (admin review queue). Mirrors GET /guardrails/submissions.", @@ -2261,7 +2324,10 @@ if MCP_AVAILABLE: @router.delete( "/server/{server_id}/user-credential", - description="Delete the calling user's stored API key for a BYOK MCP server", + description=( + "Delete the calling user's stored API key for a BYOK MCP server. " + "A proxy admin may pass user_id to revoke another user's stored key." + ), dependencies=[Depends(user_api_key_auth)], response_model=MCPUserCredentialResponse, ) @@ -2269,24 +2335,20 @@ if MCP_AVAILABLE: async def delete_mcp_user_credential( server_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + user_id: Annotated[str | None, Query(min_length=1)] = None, ): - """Remove the calling user's BYOK credential.""" + """Remove the target user's BYOK credential (the caller unless an admin names another user).""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - user_id: Final = user_api_key_dict.user_id or "" - if not user_id: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "User ID not found in token"}, - ) + target_user_id: Final = _resolve_credential_target_user_id(user_api_key_dict, user_id) try: - await delete_user_credential(prisma_client, user_id, server_id) + await delete_user_credential(prisma_client, target_user_id, server_id) except RecordNotFoundError: pass # Already deleted or didn't exist from litellm.proxy._experimental.mcp_server.server import ( _invalidate_byok_cred_cache, ) - _invalidate_byok_cred_cache(user_id, server_id) + _invalidate_byok_cred_cache(target_user_id, server_id) return MCPUserCredentialResponse(server_id=server_id, has_credential=False) # ── OAuth2 user-credential endpoints ────────────────────────────────────── @@ -2362,7 +2424,10 @@ if MCP_AVAILABLE: @router.delete( "/server/{server_id}/oauth-user-credential", - description="Revoke the calling user's stored OAuth2 token for an MCP server", + description=( + "Revoke the calling user's stored OAuth2 token for an MCP server. " + "A proxy admin may pass user_id to revoke another user's stored token." + ), dependencies=[Depends(user_api_key_auth)], response_model=MCPOAuthUserCredentialStatus, ) @@ -2370,29 +2435,25 @@ if MCP_AVAILABLE: async def delete_mcp_oauth_user_credential( server_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + user_id: Annotated[str | None, Query(min_length=1)] = None, ): - """Revoke/delete the user's OAuth2 credential.""" + """Revoke the target user's OAuth2 credential (the caller unless an admin names another user).""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - user_id: Final = user_api_key_dict.user_id or "" - if not user_id: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "User ID not found in token"}, - ) + target_user_id: Final = _resolve_credential_target_user_id(user_api_key_dict, user_id) # Only delete if the stored credential is actually an OAuth2 token. # This prevents accidentally deleting a BYOK credential if one exists # for the same (user_id, server_id) pair. - cred_to_delete: Final = await get_user_oauth_credential(prisma_client, user_id, server_id) + cred_to_delete: Final = await get_user_oauth_credential(prisma_client, target_user_id, server_id) if cred_to_delete is not None: try: - await delete_user_credential(prisma_client, user_id, server_id) + await delete_user_credential(prisma_client, target_user_id, server_id) except RecordNotFoundError: pass # Already gone — treat as a successful delete from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 global_mcp_server_manager, ) - await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id) + await global_mcp_server_manager.invalidate_user_oauth_token_cache(target_user_id, server_id) return MCPOAuthUserCredentialStatus( server_id=server_id, has_credential=False, @@ -2481,6 +2542,30 @@ if MCP_AVAILABLE: ) return items + @router.get( + "/server/{server_id}/user-credentials", + description="List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets)", + dependencies=(Depends(user_api_key_auth),), + response_model=list[MCPServerUserCredentialListItem], + ) + @management_endpoint_wrapper + async def list_mcp_server_user_credentials( + server_id: str, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + ) -> tuple[MCPServerUserCredentialListItem, ...]: + if user_api_key_dict.user_role not in ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ # mutable-ok: FastAPI HTTPException detail requires a plain dict + "error": "Admin access required to view MCP server user credentials.", + }, + ) + prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") + return await list_server_user_credentials(prisma_client, server_id) + # ── Per-user MCP env var endpoints ──────────────────────────────────────── async def _authorize_and_fetch_mcp_server( diff --git a/litellm/types/mcp.py b/litellm/types/mcp.py index 2d06bb9a009..c5a26c997b7 100644 --- a/litellm/types/mcp.py +++ b/litellm/types/mcp.py @@ -464,3 +464,11 @@ class MCPGatewaySessionsResponse(BaseModel): by_client: list[MCPGatewaySessionGroupCount] = Field(default_factory=list) by_user: list[MCPGatewaySessionGroupCount] = Field(default_factory=list) sessions: list[MCPGatewaySession] = Field(default_factory=list) + + +class MCPGatewaySessionsTerminateResponse(BaseModel): + """Stateful sessions an administrator force-closed on this proxy worker.""" + + worker_pid: int + terminated_sessions: int + sessions: list[MCPGatewaySession] = Field(default_factory=list) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 60a5e1a22bb..cfcff73b857 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -212,6 +212,53 @@ async def test_purge_user_oauth_credentials_for_server_invalidates_each_user(): assert set(invalidations) == {("alice", "srv-1"), ("bob", "srv-1")} +@pytest.mark.asyncio +async def test_list_server_user_credentials_types_each_row_without_leaking_the_secret(): + """The admin view of one server's stored credentials names the user and the kind of + credential (OAuth2 vs BYOK) and echoes OAuth expiry, but never the token or key itself.""" + from litellm.proxy._experimental.mcp_server.db import list_server_user_credentials + + oauth_row = _legacy_row( + json.dumps( + { + "type": "oauth2", + "access_token": "tok-alice", + "expires_at": "2026-12-31T00:00:00+00:00", + "connected_at": "2026-01-01T00:00:00+00:00", + } + ) + ) + oauth_row.user_id = "alice" + oauth_row.updated_at = datetime(2026, 1, 1, tzinfo=timezone.utc) + byok_row = _byok_row("carol") + byok_row.updated_at = datetime(2026, 2, 1, tzinfo=timezone.utc) + prisma = MagicMock() + prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[oauth_row, byok_row]) + + items = await list_server_user_credentials(prisma, "srv-1") + + prisma.db.litellm_mcpusercredentials.find_many.assert_awaited_once_with(where={"server_id": "srv-1"}) + assert [item.model_dump() for item in items] == [ + { + "user_id": "alice", + "credential_type": "oauth2", + "expires_at": "2026-12-31T00:00:00+00:00", + "connected_at": "2026-01-01T00:00:00+00:00", + "updated_at": "2026-01-01T00:00:00+00:00", + }, + { + "user_id": "carol", + "credential_type": "byok", + "expires_at": None, + "connected_at": None, + "updated_at": "2026-02-01T00:00:00+00:00", + }, + ] + serialized = "".join(item.model_dump_json() for item in items) + assert "tok-alice" not in serialized + assert "sk-byok-carol" not in serialized + + @pytest.mark.asyncio async def test_purge_user_oauth_credentials_for_server_spares_byok_rows(): """Regression: the purge used to delete_many on server_id alone, wiping BYOK API keys that share diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ef424255f04..4311a69d465 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2871,6 +2871,183 @@ def test_remove_stateful_session_tracking_drops_client_info(): assert session_id not in mcp_server._stateful_session_client_info +def _admin_terminate_fixture(mcp_server): + def auth_user(user_id: str): + return mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key=f"key-{user_id}", user_id=user_id), + ) + + contexts = { + "alice-session-1": auth_user("alice"), + "alice-session-2": auth_user("alice"), + "bob-session-1": auth_user("bob"), + "anon-session-1": mcp_server.MCPAuthenticatedUser(user_api_key_auth=None), + "gone-session-1": auth_user("alice"), + } + transports = { + session_id: MagicMock(terminate=AsyncMock()) + for session_id in ("alice-session-1", "alice-session-2", "bob-session-1", "anon-session-1") + } + return contexts, transports + + +@pytest.mark.asyncio +async def test_terminate_mcp_gateway_sessions_by_user_closes_every_live_session_of_that_user(): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import session_manager_stateful + except ImportError: + pytest.skip("MCP server not available") + + contexts, transports = _admin_terminate_fixture(mcp_server) + live_transports = dict(transports) + last_seen = {session_id: 100.0 for session_id in contexts} + locks = {session_id: asyncio.Lock() for session_id in contexts} + + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + session_manager_stateful, "_server_instances", live_transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_context_last_seen, last_seen, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_locks, locks, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_owners, {session_id: "owner" for session_id in contexts}, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_active_request_counts, {}, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_client_info, {}, clear=True + ), + ): + result = await mcp_server.terminate_mcp_gateway_sessions(user_id="alice") + + assert set(live_transports) == {"bob-session-1", "anon-session-1"} + assert set(mcp_server._stateful_session_auth_contexts) == {"bob-session-1", "anon-session-1", "gone-session-1"} + assert set(mcp_server._stateful_session_locks) == {"bob-session-1", "anon-session-1", "gone-session-1"} + assert set(mcp_server._stateful_session_owners) == {"bob-session-1", "anon-session-1", "gone-session-1"} + assert set(mcp_server._stateful_session_auth_context_last_seen) == { + "bob-session-1", + "anon-session-1", + "gone-session-1", + } + + transports["alice-session-1"].terminate.assert_awaited_once() + transports["alice-session-2"].terminate.assert_awaited_once() + transports["bob-session-1"].terminate.assert_not_awaited() + transports["anon-session-1"].terminate.assert_not_awaited() + assert result.terminated_sessions == 2 + assert sorted(session.session_id_prefix for session in result.sessions) == ["alice-se", "alice-se"] + assert {session.user_id for session in result.sessions} == {"alice"} + assert "key-alice" not in result.model_dump_json() + assert "alice-session-1" not in result.model_dump_json() + + +@pytest.mark.asyncio +async def test_terminate_mcp_gateway_sessions_prefix_and_user_must_both_match(): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import session_manager_stateful + except ImportError: + pytest.skip("MCP server not available") + + contexts, transports = _admin_terminate_fixture(mcp_server) + live_transports = dict(transports) + + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + session_manager_stateful, "_server_instances", live_transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_client_info, {}, clear=True + ), + ): + mismatch = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="alice-session-1", user_id="bob") + assert mismatch.terminated_sessions == 0 + assert set(live_transports) == set(transports) + + stale = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="gone-session-1") + assert stale.terminated_sessions == 0 + + exact = await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix="alice-session-1", user_id="alice") + assert exact.terminated_sessions == 1 + assert set(live_transports) == {"alice-session-2", "bob-session-1", "anon-session-1"} + + +@pytest.mark.asyncio +async def test_admin_terminated_session_id_gets_404_instead_of_a_fresh_stateless_session(): + """Once an admin closes a session, a client replaying its id must not be silently upgraded to a + new stateless session by the stale-header path; it gets 404 and has to initialize again.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import session_manager_stateful + from starlette.types import Scope + except ImportError: + pytest.skip("MCP server not available") + + session_id = "admin-closed-session-1" + live_transports = {session_id: MagicMock(terminate=AsyncMock())} + contexts = { + session_id: mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key="key-alice", user_id="alice"), + ) + } + + def scope_with_session_header() -> Scope: + return { + "type": "http", + "method": "POST", + "headers": [(b"content-type", b"application/json"), (b"mcp-session-id", session_id.encode())], + } + + try: + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + session_manager_stateful, "_server_instances", live_transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_client_info, {}, clear=True + ), + ): + await mcp_server.terminate_mcp_gateway_sessions(session_id_prefix=session_id) + + terminated_scope = scope_with_session_header() + send = AsyncMock() + handled = await mcp_server._handle_stale_mcp_session( + terminated_scope, AsyncMock(), send, session_manager_stateful + ) + + assert handled is True + statuses = [m["status"] for (m,), _ in send.await_args_list if m["type"] == "http.response.start"] + assert statuses == [404] + assert [k for k, _ in terminated_scope["headers"]] == [b"content-type", b"mcp-session-id"] + + unknown_scope = scope_with_session_header() + unknown_scope["headers"][1] = (b"mcp-session-id", b"never-seen-session") + assert ( + await mcp_server._handle_stale_mcp_session( + unknown_scope, AsyncMock(), AsyncMock(), session_manager_stateful + ) + is False + ) + assert [k for k, _ in unknown_scope["headers"]] == [b"content-type"] + finally: + mcp_server._admin_terminated_session_ids.clear() + + @pytest.mark.asyncio async def test_initialize_request_with_existing_session_tracks_new_session(): try: diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 7c874aff3df..2f33018599f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -5136,6 +5136,261 @@ async def test_delete_mcp_oauth_user_credential_invalidates_when_record_already_ assert result.has_credential is False +def _make_admin_auth(role: LitellmUserRoles = LitellmUserRoles.PROXY_ADMIN) -> "UserAPIKeyAuth": + return UserAPIKeyAuth(api_key="sk-admin", user_id="admin-user", user_role=role) + + +@pytest.mark.asyncio +async def test_admin_revokes_another_users_byok_credential(): + """A proxy admin naming user_id deletes and cache-invalidates that user's stored key, not their own.""" + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_user_credential, + ) + + delete_mock = AsyncMock(return_value=None) + invalidate_mock = MagicMock() + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the credential row delete + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=delete_mock, + ), + patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam + mcp_server, "_invalidate_byok_cred_cache", new=invalidate_mock + ), + ): + result = await delete_mcp_user_credential( + server_id="srv-byok-admin", + user_api_key_dict=_make_admin_auth(), + user_id="mallory", + ) + + delete_mock.assert_awaited_once() + assert delete_mock.await_args.args[1:] == ("mallory", "srv-byok-admin") + invalidate_mock.assert_called_once_with("mallory", "srv-byok-admin") + assert result.has_credential is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +async def test_non_full_admin_cannot_revoke_another_users_byok_credential(role): + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_user_credential, + ) + + delete_mock = AsyncMock(return_value=None) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the credential row delete + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=delete_mock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await delete_mcp_user_credential( + server_id="srv-byok-forbidden", + user_api_key_dict=_make_admin_auth(role), + user_id="mallory", + ) + + assert exc_info.value.status_code == 403 + delete_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_user_naming_themselves_still_deletes_own_byok_credential(): + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_user_credential, + ) + + delete_mock = AsyncMock(return_value=None) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the credential row delete + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=delete_mock, + ), + patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam + mcp_server, "_invalidate_byok_cred_cache", new=MagicMock() + ), + ): + await delete_mcp_user_credential( + server_id="srv-byok-self", + user_api_key_dict=_make_user_auth("user-self"), + user_id="user-self", + ) + + assert delete_mock.await_args.args[1:] == ("user-self", "srv-byok-self") + + +@pytest.mark.asyncio +async def test_admin_revokes_another_users_oauth_credential(): + """A proxy admin naming user_id reads, deletes, and cache-invalidates that user's OAuth token.""" + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_oauth_user_credential, + ) + + get_mock = AsyncMock(return_value={"type": "oauth2", "access_token": "mallory-tok"}) + delete_mock = AsyncMock(return_value=None) + invalidate_mock = AsyncMock(return_value=None) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the stored OAuth token read + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential", + new=get_mock, + ), + patch( # test-quality-ok: endpoint test stubs the credential row delete + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=delete_mock, + ), + patch.object( # test-quality-ok: the OAuth cache lives on the global manager; the suite's only seam + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + result = await delete_mcp_oauth_user_credential( + server_id="srv-oauth-admin", + user_api_key_dict=_make_admin_auth(), + user_id="mallory", + ) + + assert get_mock.await_args.args[1:] == ("mallory", "srv-oauth-admin") + assert delete_mock.await_args.args[1:] == ("mallory", "srv-oauth-admin") + invalidate_mock.assert_awaited_once_with("mallory", "srv-oauth-admin") + assert result.has_credential is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +async def test_non_full_admin_cannot_revoke_another_users_oauth_credential(role): + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_oauth_user_credential, + ) + + get_mock = AsyncMock(return_value={"type": "oauth2", "access_token": "mallory-tok"}) + delete_mock = AsyncMock(return_value=None) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the stored OAuth token read + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential", + new=get_mock, + ), + patch( # test-quality-ok: endpoint test stubs the credential row delete + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=delete_mock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await delete_mcp_oauth_user_credential( + server_id="srv-oauth-forbidden", + user_api_key_dict=_make_admin_auth(role), + user_id="mallory", + ) + + assert exc_info.value.status_code == 403 + get_mock.assert_not_awaited() + delete_mock.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +async def test_admin_lists_every_users_credential_for_a_server(role): + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._types import MCPServerUserCredentialListItem + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + list_mcp_server_user_credentials, + ) + + items = ( + MCPServerUserCredentialListItem(user_id="alice", credential_type="byok", updated_at="2026-01-01T00:00:00"), + MCPServerUserCredentialListItem(user_id="bob", credential_type="oauth2", updated_at="2026-01-02T00:00:00"), + ) + list_mock = AsyncMock(return_value=items) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the credential row listing + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_server_user_credentials", + new=list_mock, + ), + ): + result = await list_mcp_server_user_credentials( + server_id="srv-list-admin", + user_api_key_dict=_make_admin_auth(role), + ) + + assert list_mock.await_args.args[1:] == ("srv-list-admin",) + assert [(item.user_id, item.credential_type) for item in result] == [("alice", "byok"), ("bob", "oauth2")] + + +@pytest.mark.asyncio +async def test_non_admin_cannot_list_a_servers_user_credentials(): + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + list_mcp_server_user_credentials, + ) + + list_mock = AsyncMock(return_value=()) + with ( + patch( # test-quality-ok: endpoint test stubs the Prisma client lookup + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( # test-quality-ok: endpoint test stubs the credential row listing + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_server_user_credentials", + new=list_mock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await list_mcp_server_user_credentials( + server_id="srv-list-forbidden", + user_api_key_dict=_make_user_auth("user-plain"), + ) + + assert exc_info.value.status_code == 403 + list_mock.assert_not_awaited() + + @pytest.mark.asyncio async def test_list_mcp_user_credentials_batch_server_fetch(): """list_mcp_user_credentials uses a single batch DB call, not N+1 queries.""" @@ -7321,3 +7576,146 @@ class TestGetMCPGatewaySessions: assert [(group.label, group.count) for group in result.by_client] == [("cursor", 1)] assert [(group.label, group.count) for group in result.by_user] == [("alice", 1)] assert "sk-live-secret" not in result.model_dump_json() + + +class TestDeleteMCPGatewaySessions: + @pytest.fixture(autouse=True) + def _forget_admin_terminated_ids(self): + from litellm.proxy._experimental.mcp_server import server as mcp_server + + yield + mcp_server._admin_terminated_session_ids.clear() + + @pytest.mark.asyncio + @pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) + async def test_non_full_admin_forbidden_before_any_session_is_touched(self, role): + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_gateway_sessions, + ) + + session_id = "gateway-terminate-forbidden-1" + transport = MagicMock(terminate=AsyncMock()) + auth_user = mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-live", user_id="alice"), + ) + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + mcp_server.session_manager_stateful, "_server_instances", {session_id: transport} + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, {session_id: auth_user}, clear=True + ), + ): + with pytest.raises(HTTPException) as exc_info: + await delete_mcp_gateway_sessions( + user_api_key_dict=generate_mock_user_api_key_auth(user_role=role), + session_id_prefix=session_id, + user_id=None, + ) + assert exc_info.value.status_code == 403 + transport.terminate.assert_not_awaited() + assert session_id in mcp_server._stateful_session_auth_contexts + + @pytest.mark.asyncio + async def test_requires_a_selector(self): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_gateway_sessions, + ) + + with pytest.raises(HTTPException) as exc_info: + await delete_mcp_gateway_sessions( + user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN), + session_id_prefix=None, + user_id=None, + ) + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_admin_terminates_only_the_selected_session(self): + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_gateway_sessions, + ) + from litellm.types.mcp import MCPGatewaySessionsTerminateResponse + + target_id = "11111111-target-session" + other_id = "22222222-other-session" + target_transport = MagicMock(terminate=AsyncMock()) + other_transport = MagicMock(terminate=AsyncMock()) + transports = {target_id: target_transport, other_id: other_transport} + contexts = { + target_id: mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-live-target", user_id="alice"), + ), + other_id: mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key="sk-live-other", user_id="bob"), + ), + } + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + mcp_server.session_manager_stateful, "_server_instances", transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + ): + result = await delete_mcp_gateway_sessions( + user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN), + session_id_prefix=target_id[:8], + user_id=None, + ) + assert target_id not in transports + assert other_id in transports + assert target_id not in mcp_server._stateful_session_auth_contexts + assert other_id in mcp_server._stateful_session_auth_contexts + + target_transport.terminate.assert_awaited_once() + other_transport.terminate.assert_not_awaited() + assert isinstance(result, MCPGatewaySessionsTerminateResponse) + assert result.terminated_sessions == 1 + assert [(s.session_id_prefix, s.user_id) for s in result.sessions] == [(target_id[:8], "alice")] + assert target_id not in result.model_dump_json() + assert "sk-live-target" not in result.model_dump_json() + + @pytest.mark.asyncio + async def test_admin_terminates_every_session_of_the_selected_user(self): + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_gateway_sessions, + ) + + def auth_user(user_id: str): + return mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key=f"sk-live-{user_id}", user_id=user_id), + ) + + transports = { + "bob-session-1": MagicMock(terminate=AsyncMock()), + "bob-session-2": MagicMock(terminate=AsyncMock()), + "alice-session-1": MagicMock(terminate=AsyncMock()), + } + contexts = { + "bob-session-1": auth_user("bob"), + "bob-session-2": auth_user("bob"), + "alice-session-1": auth_user("alice"), + } + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + mcp_server.session_manager_stateful, "_server_instances", transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + ): + result = await delete_mcp_gateway_sessions( + user_api_key_dict=generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN), + session_id_prefix=None, + user_id="bob", + ) + assert set(transports) == {"alice-session-1"} + assert set(mcp_server._stateful_session_auth_contexts) == {"alice-session-1"} + + assert result.terminated_sessions == 2 + assert {s.user_id for s in result.sessions} == {"bob"} + assert "sk-live-bob" not in result.model_dump_json() diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.integration.test.tsx index 11328ff1a3d..f1ddf709038 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.integration.test.tsx @@ -1,13 +1,15 @@ import React from "react"; import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { MCPGatewaySessionsTab, formatIdleSeconds } from "./MCPGatewaySessionsTab"; +import { MCPGatewaySessionsTab, describeTerminateResult, formatIdleSeconds } from "./MCPGatewaySessionsTab"; import * as networking from "@/components/networking"; -import type { MCPGatewaySessionsResponse } from "@/components/mcp_tools/types"; +import type { MCPGatewaySessionsResponse, MCPGatewaySessionsTerminateResponse } from "@/components/mcp_tools/types"; vi.mock("@/components/networking", () => ({ fetchMCPGatewaySessions: vi.fn(), + terminateMCPGatewaySessions: vi.fn(), })); const REPORT: MCPGatewaySessionsResponse = { @@ -64,11 +66,11 @@ const REPORT: MCPGatewaySessionsResponse = { ], }; -const renderTab = () => { +const renderTab = ({ canTerminate = false }: { canTerminate?: boolean } = {}) => { const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } }); return render( - + , ); }; @@ -83,6 +85,17 @@ describe("formatIdleSeconds", () => { }); }); +describe("describeTerminateResult", () => { + it("pluralizes the session count and names the worker", () => { + expect(describeTerminateResult({ worker_pid: 9, terminated_sessions: 1, sessions: [] })).toBe( + "Disconnected 1 session on worker pid 9.", + ); + expect(describeTerminateResult({ worker_pid: 9, terminated_sessions: 0, sessions: [] })).toBe( + "Disconnected 0 sessions on worker pid 9.", + ); + }); +}); + describe("MCPGatewaySessionsTab", () => { beforeEach(() => { vi.clearAllMocks(); @@ -135,4 +148,73 @@ describe("MCPGatewaySessionsTab", () => { expect(alert).toHaveTextContent("Could not load live connections"); expect(alert).toHaveTextContent("Admin access required"); }); + + it("hides every disconnect control from a read-only admin", async () => { + vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT); + renderTab({ canTerminate: false }); + + await screen.findByRole("region", { name: "Live sessions" }); + expect(screen.queryByRole("button", { name: /^Disconnect/ })).not.toBeInTheDocument(); + }); + + it("disconnects one session by its displayed prefix after confirmation and refetches", async () => { + const user = userEvent.setup(); + const terminated: MCPGatewaySessionsTerminateResponse = { + worker_pid: 4242, + terminated_sessions: 1, + sessions: [REPORT.sessions[1]], + }; + vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT); + vi.mocked(networking.terminateMCPGatewaySessions).mockResolvedValue(terminated); + renderTab({ canTerminate: true }); + + await user.click(await screen.findByRole("button", { name: "Disconnect session bbbb2222" })); + expect(networking.terminateMCPGatewaySessions).not.toHaveBeenCalled(); + const dialog = await screen.findByRole("alertdialog"); + expect(dialog).toHaveTextContent("session bbbb2222"); + await user.click(within(dialog).getByRole("button", { name: "Disconnect" })); + + const status = await screen.findByText("Disconnected 1 session on worker pid 4242.", { exact: false }); + expect(status).toBeInTheDocument(); + expect(networking.terminateMCPGatewaySessions).toHaveBeenCalledWith("token", { session_id_prefix: "bbbb2222" }); + expect(networking.fetchMCPGatewaySessions).toHaveBeenCalledTimes(2); + }); + + it("disconnects every session of a user from the by-user table", async () => { + const user = userEvent.setup(); + vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT); + vi.mocked(networking.terminateMCPGatewaySessions).mockResolvedValue({ + worker_pid: 4242, + terminated_sessions: 2, + sessions: [REPORT.sessions[0], REPORT.sessions[1]], + }); + renderTab({ canTerminate: true }); + + const byUser = await screen.findByRole("region", { name: "Sessions by user" }); + expect(within(byUser).queryByRole("button", { name: /\(unknown\)/ })).not.toBeInTheDocument(); + await user.click(within(byUser).getByRole("button", { name: "Disconnect all sessions for user alice" })); + const dialog = await screen.findByRole("alertdialog"); + expect(dialog).toHaveTextContent("every live session opened by user alice"); + await user.click(within(dialog).getByRole("button", { name: "Disconnect" })); + + expect(await screen.findByText(/Disconnected 2 sessions on worker pid 4242\./)).toBeInTheDocument(); + expect(networking.terminateMCPGatewaySessions).toHaveBeenCalledWith("token", { user_id: "alice" }); + }); + + it("shows the API error when a disconnect is refused", async () => { + const user = userEvent.setup(); + vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(REPORT); + vi.mocked(networking.terminateMCPGatewaySessions).mockRejectedValue( + new Error("Proxy admin access required to terminate MCP gateway sessions."), + ); + renderTab({ canTerminate: true }); + + await user.click(await screen.findByRole("button", { name: "Disconnect session aaaa1111" })); + await user.click(within(await screen.findByRole("alertdialog")).getByRole("button", { name: "Disconnect" })); + + const alert = await screen.findByRole("alert"); + expect(alert).toHaveTextContent("Could not disconnect"); + expect(alert).toHaveTextContent("Proxy admin access required to terminate MCP gateway sessions."); + expect(screen.getByRole("region", { name: "Live sessions" })).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.tsx index 18f44d5d090..04a095f8792 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPGatewaySessionsTab.tsx @@ -1,14 +1,27 @@ "use client"; -import React from "react"; -import { useQuery } from "@tanstack/react-query"; -import { RefreshCw } from "lucide-react"; +import React, { useState } from "react"; +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; +import { RefreshCw, Unplug } from "lucide-react"; import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; +import { + AlertDialog, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/components/ui/alert-dialog"; import { Button } from "@/components/ui/button"; import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; -import { fetchMCPGatewaySessions } from "@/components/networking"; -import type { MCPGatewaySessionGroupCount, MCPGatewaySessionsResponse } from "@/components/mcp_tools/types"; +import { fetchMCPGatewaySessions, terminateMCPGatewaySessions } from "@/components/networking"; +import type { + MCPGatewaySessionGroupCount, + MCPGatewaySessionSelector, + MCPGatewaySessionsResponse, + MCPGatewaySessionsTerminateResponse, +} from "@/components/mcp_tools/types"; import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory"; const mcpGatewaySessionKeys = createQueryKeys("mcpGatewaySessions"); @@ -28,6 +41,16 @@ function groupLabel(label: string | null): string { return label === "" ? '""' : label; } +export function describeSelector(selector: MCPGatewaySessionSelector): string { + if (selector.user_id !== undefined) return `every live session opened by user ${groupLabel(selector.user_id)}`; + return `session ${selector.session_id_prefix}`; +} + +export function describeTerminateResult(result: MCPGatewaySessionsTerminateResponse): string { + const noun = result.terminated_sessions === 1 ? "session" : "sessions"; + return `Disconnected ${result.terminated_sessions} ${noun} on worker pid ${result.worker_pid}.`; +} + function StatCard({ label, value }: { label: string; value: number }) { return (

@@ -37,14 +60,37 @@ function StatCard({ label, value }: { label: string; value: number }) { ); } +function DisconnectUserButton({ + userId, + onDisconnectUser, +}: { + userId: string | null; + onDisconnectUser: (userId: string) => void; +}) { + if (userId === null || userId === "") return null; + return ( + + ); +} + function GroupCountTable({ title, groups, labelHeader, + onDisconnectUser, }: { title: string; groups: MCPGatewaySessionGroupCount[]; labelHeader: string; + onDisconnectUser?: (userId: string) => void; }) { return (
@@ -54,6 +100,7 @@ function GroupCountTable({ {labelHeader} Sessions + {onDisconnectUser ? Actions : null} @@ -61,6 +108,11 @@ function GroupCountTable({ {groupLabel(group.label)} {group.count} + {onDisconnectUser ? ( + + + + ) : null} ))} @@ -73,10 +125,12 @@ function SessionsBody({ data, error, isLoading, + onDisconnect, }: { data: MCPGatewaySessionsResponse | undefined; error: Error | null; isLoading: boolean; + onDisconnect: ((selector: MCPGatewaySessionSelector) => void) | null; }) { if (isLoading) { return ( @@ -117,7 +171,12 @@ function SessionsBody({
- + onDisconnect({ user_id: userId }) : undefined} + />

@@ -134,11 +193,12 @@ function SessionsBody({ Client IP Idle In flight + {onDisconnect ? Actions : null} - {data.sessions.map((session) => ( - + {data.sessions.map((session, index) => ( + {session.session_id_prefix} {session.client_name === null ? ( @@ -169,6 +229,19 @@ function SessionsBody({ {session.client_ip || "-"} {formatIdleSeconds(session.idle_seconds)} {session.in_flight_requests} + {onDisconnect ? ( + + + + ) : null} ))} @@ -180,9 +253,12 @@ function SessionsBody({ interface MCPGatewaySessionsTabProps { accessToken: string | null; + canTerminate: boolean; } -export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProps) { +export function MCPGatewaySessionsTab({ accessToken, canTerminate }: MCPGatewaySessionsTabProps) { + const queryClient = useQueryClient(); + const [pendingSelector, setPendingSelector] = useState(null); const queryOptions = { queryKey: mcpGatewaySessionKeys.lists(), queryFn: () => fetchMCPGatewaySessions(accessToken!), @@ -190,6 +266,15 @@ export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProp refetchInterval: REFETCH_INTERVAL_MS, }; const { data, error, isLoading, isFetching, refetch } = useQuery(queryOptions); + const terminate = useMutation({ + mutationFn: (selector) => terminateMCPGatewaySessions(accessToken!, selector), + onSettled: () => queryClient.invalidateQueries({ queryKey: mcpGatewaySessionKeys.lists() }), + }); + const confirmDisconnect = () => { + if (pendingSelector === null) return; + terminate.mutate(pendingSelector); + setPendingSelector(null); + }; return (
@@ -214,7 +299,48 @@ export function MCPGatewaySessionsTab({ accessToken }: MCPGatewaySessionsTabProp
- + {terminate.isError ? ( + + Could not disconnect + {terminate.error.message} + + ) : null} + {terminate.isSuccess ? ( + + Disconnected + + {describeTerminateResult(terminate.data)} Clients holding those sessions must send a new initialize request, + which re-runs authentication. Sessions on other proxy workers are not affected. + + + ) : null} + + + + !open && setPendingSelector(null)}> + + + Disconnect MCP session + + {pendingSelector ? `This force-closes ${describeSelector(pendingSelector)} on this proxy worker. ` : ""} + In-flight requests fail and the client must initialize again before it can call tools. + + + + + + + +

); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.integration.test.tsx new file mode 100644 index 00000000000..a20dd032d33 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.integration.test.tsx @@ -0,0 +1,102 @@ +import React from "react"; +import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel"; +import * as networking from "@/components/networking"; +import type { MCPServerUserCredentialListItem } from "@/components/mcp_tools/types"; + +vi.mock("@/components/networking", () => ({ + fetchMCPServerUserCredentials: vi.fn(), + revokeMCPServerUserCredential: vi.fn(), +})); + +const ITEMS: MCPServerUserCredentialListItem[] = [ + { + user_id: "alice", + credential_type: "oauth2", + expires_at: "2026-12-31T00:00:00+00:00", + connected_at: "2026-01-01T00:00:00+00:00", + updated_at: "2026-01-01T00:00:00+00:00", + }, + { + user_id: "carol", + credential_type: "byok", + expires_at: null, + connected_at: null, + updated_at: "2026-02-01T00:00:00+00:00", + }, +]; + +const renderPanel = ({ canRevoke = false }: { canRevoke?: boolean } = {}) => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } }); + return render( + + + , + ); +}; + +describe("MCPServerUserCredentialsPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("lists each user's credential type without a revoke control for a read-only admin", async () => { + vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue(ITEMS); + renderPanel({ canRevoke: false }); + + const table = await screen.findByRole("region", { name: "Stored user credentials" }); + expect(within(table).getByRole("row", { name: /alice/ })).toHaveTextContent("OAuth2"); + expect(within(table).getByRole("row", { name: /carol/ })).toHaveTextContent("BYOK API key"); + expect(screen.queryByRole("button", { name: /^Revoke credential/ })).not.toBeInTheDocument(); + expect(networking.fetchMCPServerUserCredentials).toHaveBeenCalledWith("token", "srv-1"); + }); + + it("revokes the selected user's credential through the route for its type and refetches", async () => { + const user = userEvent.setup(); + vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValueOnce(ITEMS).mockResolvedValueOnce([ITEMS[1]]); + vi.mocked(networking.revokeMCPServerUserCredential).mockResolvedValue(undefined); + renderPanel({ canRevoke: true }); + + await user.click(await screen.findByRole("button", { name: "Revoke credential for user alice" })); + expect(networking.revokeMCPServerUserCredential).not.toHaveBeenCalled(); + const dialog = await screen.findByRole("alertdialog"); + expect(dialog).toHaveTextContent("OAuth2 credential stored for user alice"); + await user.click(within(dialog).getByRole("button", { name: "Revoke" })); + + expect(await screen.findByText(/OAuth2 credential for user alice was deleted/)).toBeInTheDocument(); + expect(networking.revokeMCPServerUserCredential).toHaveBeenCalledWith("token", "srv-1", "alice", "oauth2"); + const table = await screen.findByRole("region", { name: "Stored user credentials" }); + expect(within(table).queryByRole("row", { name: /alice/ })).not.toBeInTheDocument(); + expect(within(table).getByRole("row", { name: /carol/ })).toBeInTheDocument(); + }); + + it("shows the API error when a revoke is refused and keeps the list", async () => { + const user = userEvent.setup(); + vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue(ITEMS); + vi.mocked(networking.revokeMCPServerUserCredential).mockRejectedValue( + new Error("Proxy admin access required to revoke another user's MCP credential."), + ); + renderPanel({ canRevoke: true }); + + await user.click(await screen.findByRole("button", { name: "Revoke credential for user carol" })); + await user.click(within(await screen.findByRole("alertdialog")).getByRole("button", { name: "Revoke" })); + + const alert = await screen.findByRole("alert"); + expect(alert).toHaveTextContent("Could not revoke credential"); + expect(alert).toHaveTextContent("Proxy admin access required to revoke another user's MCP credential."); + expect(networking.revokeMCPServerUserCredential).toHaveBeenCalledWith("token", "srv-1", "carol", "byok"); + expect(screen.getByRole("region", { name: "Stored user credentials" })).toBeInTheDocument(); + }); + + it("shows the API error when the list cannot be loaded", async () => { + vi.mocked(networking.fetchMCPServerUserCredentials).mockRejectedValue(new Error("Admin access required")); + renderPanel(); + + const alert = await screen.findByRole("alert"); + expect(alert).toHaveTextContent("Could not load user credentials"); + expect(alert).toHaveTextContent("Admin access required"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.tsx new file mode 100644 index 00000000000..20b679450c3 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerUserCredentialsPanel.tsx @@ -0,0 +1,212 @@ +"use client"; + +import React, { useState } from "react"; +import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; +import { RefreshCw, ShieldOff } from "lucide-react"; +import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; +import { + AlertDialog, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/components/ui/alert-dialog"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { fetchMCPServerUserCredentials, revokeMCPServerUserCredential } from "@/components/networking"; +import type { MCPServerUserCredentialListItem } from "@/components/mcp_tools/types"; +import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory"; + +const mcpServerUserCredentialKeys = createQueryKeys("mcpServerUserCredentials"); + +export function credentialTypeLabel(credentialType: MCPServerUserCredentialListItem["credential_type"]): string { + return credentialType === "oauth2" ? "OAuth2" : "BYOK API key"; +} + +export function formatTimestamp(value: string | null): string { + if (value === null) return "-"; + const parsed = new Date(value); + return Number.isNaN(parsed.getTime()) ? value : parsed.toLocaleString(); +} + +function CredentialsBody({ + items, + error, + isLoading, + onRevoke, +}: { + items: MCPServerUserCredentialListItem[] | undefined; + error: Error | null; + isLoading: boolean; + onRevoke: ((item: MCPServerUserCredentialListItem) => void) | null; +}) { + if (isLoading) { + return ( +
+ +

Loading user credentials...

+
+ ); + } + if (error) { + return ( + + Could not load user credentials + {error.message} + + ); + } + if (!items) return null; + if (items.length === 0) { + return ( +
+

No user has a stored credential for this server.

+
+ ); + } + return ( +
+ + + + User + Type + Connected + Expires + Updated + {onRevoke ? Actions : null} + + + + {items.map((item) => ( + + {item.user_id} + + {credentialTypeLabel(item.credential_type)} + + {formatTimestamp(item.connected_at)} + {formatTimestamp(item.expires_at)} + {formatTimestamp(item.updated_at)} + {onRevoke ? ( + + + + ) : null} + + ))} + +
+
+ ); +} + +interface MCPServerUserCredentialsPanelProps { + serverId: string; + accessToken: string | null; + canRevoke: boolean; +} + +export function MCPServerUserCredentialsPanel({ + serverId, + accessToken, + canRevoke, +}: MCPServerUserCredentialsPanelProps) { + const queryClient = useQueryClient(); + const [pendingItem, setPendingItem] = useState(null); + const queryKey = mcpServerUserCredentialKeys.detail(serverId); + const { data, error, isLoading, isFetching, refetch } = useQuery({ + queryKey, + queryFn: () => fetchMCPServerUserCredentials(accessToken!, serverId), + enabled: !!accessToken, + }); + const revoke = useMutation({ + mutationFn: (item) => revokeMCPServerUserCredential(accessToken!, serverId, item.user_id, item.credential_type), + onSettled: () => queryClient.invalidateQueries({ queryKey }), + }); + const confirmRevoke = () => { + if (pendingItem === null) return; + revoke.mutate(pendingItem); + setPendingItem(null); + }; + + return ( +
+
+
+

User Credentials

+

+ Per-user OAuth2 tokens and BYOK API keys stored for this server. Revoking one deletes it from the database + and clears the cached copy, so the user must connect again before the gateway will call this server for + them. +

+
+ +
+ + {revoke.isError ? ( + + Could not revoke credential + {revoke.error.message} + + ) : null} + {revoke.isSuccess ? ( + + Credential revoked + + The stored {credentialTypeLabel(revoke.variables.credential_type)} credential for user{" "} + {revoke.variables.user_id} was deleted. + + + ) : null} + + + + !open && setPendingItem(null)}> + + + Revoke stored credential + + {pendingItem + ? `This deletes the ${credentialTypeLabel(pendingItem.credential_type)} credential stored for user ${pendingItem.user_id}. ` + : ""} + Their next MCP request to this server fails until they connect again. + + + + + + + + +
+ ); +} + +export default MCPServerUserCredentialsPanel; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index 475392620d4..c23b48ef672 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -9,7 +9,9 @@ import { MCPServer, handleTransport, handleAuth } from "@/components/mcp_tools/t // TODO: Move Tools viewer from index file import { MCPToolsViewer } from "."; import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; +import { MCPServerUserCredentialsPanel } from "./MCPServerUserCredentialsPanel"; import { getSecureItem } from "@/utils/secureStorage"; +import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import MCPServerCostDisplay from "./mcp_server_cost_display"; import { getMaskedAndFullUrl } from "./utils"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; @@ -63,6 +65,8 @@ export const MCPServerView: React.FC = ({ const [showFullUrl, setShowFullUrl] = useState(false); const [copiedStates, setCopiedStates] = useState>({}); const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex); + const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole); + const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole); const handleSuccess = (updated: MCPServer) => { setEditing(false); @@ -142,6 +146,11 @@ export const MCPServerView: React.FC = ({ Settings )} + {canViewUserCredentials && ( + + User Credentials + + )} {/* Overview Panel */} @@ -387,6 +396,18 @@ export const MCPServerView: React.FC = ({ )} + + {canViewUserCredentials && ( + + + + + + )}
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 00d79022103..3a197786774 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -1,4 +1,4 @@ -import { isAdminRole, isProxyAdminTierRole } from "@/utils/roles"; +import { isAdminRole, isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import { CircleHelp, Search } from "lucide-react"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; @@ -755,7 +755,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) )} {isProxyAdminTierRole(userRole) && ( - + )} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index fe04eb4969b..2f6a3f1d0d9 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -587,3 +587,23 @@ export interface MCPGatewaySessionsResponse { by_user: MCPGatewaySessionGroupCount[]; sessions: MCPGatewaySession[]; } + +export interface MCPGatewaySessionsTerminateResponse { + worker_pid: number; + terminated_sessions: number; + sessions: MCPGatewaySession[]; +} + +export type MCPGatewaySessionSelector = + | { session_id_prefix: string; user_id?: undefined } + | { user_id: string; session_id_prefix?: undefined }; + +export type MCPServerUserCredentialType = "oauth2" | "byok"; + +export interface MCPServerUserCredentialListItem { + user_id: string; + credential_type: MCPServerUserCredentialType; + expires_at: string | null; + connected_at: string | null; + updated_at: string; +} diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a8ea1b66488..9381e2c07c9 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -97,7 +97,14 @@ import type { ModelBudgetUsage, ModelMaxBudget } from "./key_team_helpers/ModelM import type { ObjectPermission } from "./object_permission_types"; import type { components } from "@/lib/http/schema"; import { jsonFields } from "./common_components/check_openapi_schema"; -import type { MCPGatewaySessionsResponse, MCPUserEnvVarsStatus } from "./mcp_tools/types"; +import type { + MCPGatewaySessionSelector, + MCPGatewaySessionsResponse, + MCPGatewaySessionsTerminateResponse, + MCPServerUserCredentialListItem, + MCPServerUserCredentialType, + MCPUserEnvVarsStatus, +} from "./mcp_tools/types"; import type { CoordinationRedisSettings, CoordinationRedisSettingsResponse, @@ -5112,6 +5119,33 @@ export const fetchMCPSubmissions = async (accessToken: string) => { export const fetchMCPGatewaySessions = async (accessToken: string): Promise => apiClient.get(`/v1/mcp/sessions`, { accessToken }); +export const terminateMCPGatewaySessions = async ( + accessToken: string, + selector: MCPGatewaySessionSelector, +): Promise => + apiClient.delete(`/v1/mcp/sessions`, { accessToken, query: { ...selector } }); + +export const fetchMCPServerUserCredentials = async ( + accessToken: string, + serverId: string, +): Promise => + apiClient.get(`/v1/mcp/server/${encodeURIComponent(serverId)}/user-credentials`, { + accessToken, + }); + +export const revokeMCPServerUserCredential = async ( + accessToken: string, + serverId: string, + userId: string, + credentialType: MCPServerUserCredentialType, +): Promise => { + const route = credentialType === "oauth2" ? "oauth-user-credential" : "user-credential"; + await apiClient.delete(`/v1/mcp/server/${encodeURIComponent(serverId)}/${route}`, { + accessToken, + query: { user_id: userId }, + }); +}; + export const approveMCPServer = async (accessToken: string, serverId: string) => { try { const url = (proxyBaseUrl ? `${proxyBaseUrl}` : "") + `/v1/mcp/server/${encodeURIComponent(serverId)}/approve`; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f1779c7cb5a..43e1e16125c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -18902,7 +18902,7 @@ export interface paths { post: operations["store_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_post"]; /** * Delete Mcp Oauth User Credential - * @description Revoke the calling user's stored OAuth2 token for an MCP server + * @description Revoke the calling user's stored OAuth2 token for an MCP server. A proxy admin may pass user_id to revoke another user's stored token. */ delete: operations["delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete"]; options?: never; @@ -18966,7 +18966,7 @@ export interface paths { post: operations["store_mcp_user_credential_v1_mcp_server__server_id__user_credential_post"]; /** * Delete Mcp User Credential - * @description Delete the calling user's stored API key for a BYOK MCP server + * @description Delete the calling user's stored API key for a BYOK MCP server. A proxy admin may pass user_id to revoke another user's stored key. */ delete: operations["delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete"]; options?: never; @@ -18974,6 +18974,26 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/mcp/server/{server_id}/user-credentials": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * List Mcp Server User Credentials + * @description List every user's stored BYOK or OAuth2 credential for an MCP server (admin only, no secrets) + */ + get: operations["list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/mcp/server/{server_id}/user-env-vars": { parameters: { query?: never; @@ -19016,7 +19036,11 @@ export interface paths { get: operations["get_mcp_gateway_sessions_v1_mcp_sessions_get"]; put?: never; post?: never; - delete?: never; + /** + * Delete Mcp Gateway Sessions + * @description Force-close live stateful MCP gateway sessions on this proxy worker, selected by session id prefix and/or by the LiteLLM user that opened them (proxy admin only). + */ + delete: operations["delete_mcp_gateway_sessions_v1_mcp_sessions_delete"]; options?: never; head?: never; patch?: never; @@ -32383,6 +32407,18 @@ export interface components { /** Worker Pid */ worker_pid: number; }; + /** + * MCPGatewaySessionsTerminateResponse + * @description Stateful sessions an administrator force-closed on this proxy worker. + */ + MCPGatewaySessionsTerminateResponse: { + /** Sessions */ + sessions?: components["schemas"]["MCPGatewaySession"][]; + /** Terminated Sessions */ + terminated_sessions: number; + /** Worker Pid */ + worker_pid: number; + }; /** * MCPOAuthUserCredentialRequest * @description Stores a user's OAuth2 token for an OpenAPI MCP server. @@ -32487,6 +32523,25 @@ export interface components { [key: string]: unknown; }; }; + /** + * MCPServerUserCredentialListItem + * @description One user's stored credential for an MCP server, as an admin sees it. Never carries the secret. + */ + MCPServerUserCredentialListItem: { + /** Connected At */ + connected_at?: string | null; + /** + * Credential Type + * @enum {string} + */ + credential_type: "oauth2" | "byok"; + /** Expires At */ + expires_at?: string | null; + /** Updated At */ + updated_at: string; + /** User Id */ + user_id: string; + }; /** MCPSubmissionsSummary */ MCPSubmissionsSummary: { /** Active */ @@ -65109,7 +65164,9 @@ export interface operations { }; delete_mcp_oauth_user_credential_v1_mcp_server__server_id__oauth_user_credential_delete: { parameters: { - query?: never; + query?: { + user_id?: string | null; + }; header?: never; path: { server_id: string; @@ -65241,7 +65298,9 @@ export interface operations { }; delete_mcp_user_credential_v1_mcp_server__server_id__user_credential_delete: { parameters: { - query?: never; + query?: { + user_id?: string | null; + }; header?: never; path: { server_id: string; @@ -65270,6 +65329,37 @@ export interface operations { }; }; }; + list_mcp_server_user_credentials_v1_mcp_server__server_id__user_credentials_get: { + parameters: { + query?: never; + header?: never; + path: { + server_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["MCPServerUserCredentialListItem"][]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; get_mcp_user_env_vars_v1_mcp_server__server_id__user_env_vars_get: { parameters: { query?: never; @@ -65387,6 +65477,38 @@ export interface operations { }; }; }; + delete_mcp_gateway_sessions_v1_mcp_sessions_delete: { + parameters: { + query?: { + session_id_prefix?: string | null; + user_id?: string | null; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["MCPGatewaySessionsTerminateResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; get_mcp_tools_v1_mcp_tools_get: { parameters: { query?: never; From aea33b4b507d139bc8b5a1d93aa055205602a67f Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 01:37:12 +0000 Subject: [PATCH 108/253] fix(ui): hide MCP disconnect and revoke controls from view-only admin sessions The auth hook normalizes proxy_admin_viewer to Admin for page access, so the role check alone let a view-only admin see Disconnect and Revoke buttons that the backend refuses with 403. Thread isViewOnly from useAuthorized into the MCP servers page and gate both mutation controls on it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_components/mcp_server_view.test.tsx | 62 +++++++++++++++---- .../_components/mcp_server_view.tsx | 4 +- .../_components/mcp_servers.test.tsx | 48 ++++++++++++++ .../mcp-servers/_components/mcp_servers.tsx | 8 ++- .../src/app/(dashboard)/mcp-servers/page.tsx | 4 +- .../src/components/mcp_tools/types.tsx | 1 + 6 files changed, 110 insertions(+), 17 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx index 02d168bf7f4..da564f23de5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.test.tsx @@ -1,7 +1,9 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi, beforeEach } from "vitest"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { MCPServerView } from "./mcp_server_view"; +import * as networking from "@/components/networking"; import type { MCPServer } from "@/components/mcp_tools/types"; vi.mock(".", () => ({ @@ -13,6 +15,12 @@ vi.mock("./mcp_server_edit", () => ({ EDIT_OAUTH_UI_STATE_KEY: "litellm-mcp-oauth-edit-state", })); +vi.mock("@/components/networking", async (importOriginal) => ({ + ...(await importOriginal()), + fetchMCPServerUserCredentials: vi.fn(), + revokeMCPServerUserCredential: vi.fn(), +})); + const baseServer = { server_id: "srv-1", server_name: "demo server", @@ -25,19 +33,38 @@ const baseServer = { const renderView = (overrides: Partial = {}, props: Record = {}) => render( - , + + + , ); +const openUserCredentials = async (props: Record) => { + vi.mocked(networking.fetchMCPServerUserCredentials).mockResolvedValue([ + { + user_id: "alice", + credential_type: "byok", + expires_at: null, + connected_at: null, + updated_at: "2026-01-01T00:00:00+00:00", + }, + ]); + renderView({}, props); + await userEvent.click(screen.getByRole("tab", { name: "User Credentials" })); + return within(await screen.findByRole("region", { name: "Stored user credentials" })).getByRole("row", { + name: /alice/, + }); +}; + describe("MCPServerView", () => { beforeEach(() => { vi.clearAllMocks(); @@ -149,4 +176,15 @@ describe("MCPServerView", () => { expect(await screen.findByText("All tools enabled")).toBeInTheDocument(); }); + + it("lets a full admin revoke a stored user credential", async () => { + const row = await openUserCredentials({}); + expect(within(row).getByRole("button", { name: "Revoke credential for user alice" })).toBeInTheDocument(); + }); + + it("shows stored credentials to a view-only admin session without a revoke control", async () => { + const row = await openUserCredentials({ isViewOnly: true }); + expect(row).toHaveTextContent("BYOK API key"); + expect(within(row).queryByRole("button", { name: /^Revoke credential/ })).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx index c23b48ef672..a346c7d986b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx @@ -25,6 +25,7 @@ interface MCPServerViewProps { accessToken: string | null; userRole: string | null; userID: string | null; + isViewOnly?: boolean; availableAccessGroups: string[]; initialTabIndex?: number; } @@ -55,6 +56,7 @@ export const MCPServerView: React.FC = ({ accessToken, userRole, userID, + isViewOnly = false, availableAccessGroups, initialTabIndex = 0, }) => { @@ -66,7 +68,7 @@ export const MCPServerView: React.FC = ({ const [copiedStates, setCopiedStates] = useState>({}); const [selectedTabIndex, setSelectedTabIndex] = useState(returningFromEditOAuth ? 2 : initialTabIndex); const canViewUserCredentials = userRole !== null && isProxyAdminTierRole(userRole); - const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole); + const canRevokeUserCredentials = userRole !== null && isProxyAdminRole(userRole) && !isViewOnly; const handleSuccess = (updated: MCPServer) => { setEditing(false); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 8abb8855e3d..1394a923174 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -17,6 +17,8 @@ vi.mock("@/components/networking", () => ({ updateConfigFieldSetting: vi.fn().mockResolvedValue(undefined), deleteConfigFieldSetting: vi.fn().mockResolvedValue(undefined), listMCPUserEnvVarStatus: vi.fn().mockResolvedValue([]), + fetchMCPGatewaySessions: vi.fn(), + terminateMCPGatewaySessions: vi.fn(), })); const createQueryClient = () => @@ -400,4 +402,50 @@ describe("MCPServers", () => { // The server list refresh must NOT trigger a second health check expect(networking.fetchMCPServerHealth).toHaveBeenCalledTimes(1); }); + + const liveSessionsReport = { + worker_pid: 4242, + total_sessions: 1, + by_client: [{ label: "claude-code", count: 1 }], + by_user: [{ label: "alice", count: 1 }], + sessions: [ + { + session_id_prefix: "aaaa1111", + client_name: "claude-code", + client_version: "1.0.0", + user_id: "alice", + user_email: "alice@example.com", + key_alias: "alice-key", + team_id: null, + team_alias: null, + client_ip: "10.0.0.1", + idle_seconds: 5, + in_flight_requests: 0, + }, + ], + }; + + const openLiveConnections = async (props: { isViewOnly?: boolean }) => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPGatewaySessions).mockResolvedValue(liveSessionsReport); + render( + + + , + ); + await userEvent.click(await screen.findByRole("tab", { name: "Live Connections" })); + return within(await screen.findByRole("region", { name: "Live sessions" })).getByRole("row", { name: /aaaa1111/ }); + }; + + it("lets a full admin disconnect a live session", async () => { + const row = await openLiveConnections({ isViewOnly: false }); + expect(within(row).getByRole("button", { name: "Disconnect session aaaa1111" })).toBeInTheDocument(); + }); + + it("shows live sessions to a view-only admin session without any disconnect control", async () => { + const row = await openLiveConnections({ isViewOnly: true }); + expect(row).toHaveTextContent("alice@example.com"); + expect(within(row).queryByRole("button", { name: /^Disconnect/ })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /^Disconnect all/ })).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 3a197786774..aa0a031c55c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -109,7 +109,7 @@ const readToolsOAuthServerId = (): string | null => { } }; -const MCPServers: React.FC = ({ accessToken, userRole, userID }) => { +const MCPServers: React.FC = ({ accessToken, userRole, userID, isViewOnly = false }) => { const { data: mcpServers, isLoading: isLoadingServers, refetch } = useMCPServers(); // Fetch health status for all servers @@ -578,6 +578,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) accessToken={accessToken} userID={userID} userRole={userRole} + isViewOnly={isViewOnly} availableAccessGroups={uniqueMcpAccessGroups} initialTabIndex={selectedServerId === toolsTabServerId ? 1 : 0} /> @@ -755,7 +756,10 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) )} {isProxyAdminTierRole(userRole) && ( - + )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx index 462c48360cd..c297dcb8d33 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/page.tsx @@ -4,6 +4,6 @@ import { MCPServers } from "./_components"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export default function McpServers() { - const { accessToken, userRole, userId } = useAuthorized(); - return ; + const { accessToken, userRole, userId, isViewOnly } = useAuthorized(); + return ; } diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 2f6a3f1d0d9..f009429693b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -517,6 +517,7 @@ export interface MCPServerProps { accessToken: string | null; userRole: string | null; userID: string | null; + isViewOnly?: boolean; } export interface MCPToolsetTool { From 9f6e5242853a601670774871073d9175d15edfd0 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 01:43:03 +0000 Subject: [PATCH 109/253] fix(deepgram): ignore zero-duration Metadata frames when billing streamed audio Deepgram sends a Metadata frame with duration 0 on connect. When the closing Metadata frame is not collected before the socket closes, that handshake frame used to become the billed duration and the session logged zero spend. Only a positive Metadata duration is treated as authoritative now; otherwise the furthest Results end time is billed Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 2 +- .../llms/deepgram/test_deepgram_common_utils.py | 6 ++++++ 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index 4f18918fa71..ede5b4f157d 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -152,7 +152,7 @@ def deepgram_listen_audio_seconds(websocket_messages: Sequence[Mapping[str, obje duration for frame in websocket_messages if frame.get("type") == "Metadata" - if (duration := _seconds(frame.get("duration"))) is not None + if (duration := _seconds(frame.get("duration"))) is not None and duration > 0 ) if metadata_durations: return metadata_durations[-1] diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index 802a3795ffe..8b61d5bffa0 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -105,6 +105,12 @@ def test_deepgram_listen_callback_params(query_string: str, expected: tuple[str, pytest.param((_metadata(4.0), _results(0.0, 9.0), _metadata(5.5)), 5.5, id="last metadata wins"), pytest.param((_results(0.0, 2.0), _results(2.0, 3.5), _results(1.0, 1.0)), 5.5, id="furthest results end"), pytest.param((_results(0.0, 0.0), _metadata(0.0)), 0.0, id="zero metadata is a real zero"), + pytest.param( + (_metadata(0.0), _results(0.0, 2.0), _results(2.0, 3.5)), + 5.5, + id="handshake metadata zero does not hide streamed results", + ), + pytest.param((_metadata(0.0), _results(0.0, 2.0), _metadata(0.0)), 2.0, id="only zero metadata frames"), pytest.param((_results(0.0, 1.5), _metadata("6.25")), 1.5, id="string metadata is ignored"), pytest.param((_results(0.0, 1.5), _metadata(True)), 1.5, id="boolean metadata is ignored"), pytest.param((_results(0.0, 1.5), _metadata(-3.0)), 1.5, id="negative metadata is ignored"), From fd834f6f8b18c5b2a5d642627631f1d4bf566ab4 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 02:10:50 +0000 Subject: [PATCH 110/253] fix(mcp): broadcast BYOK and OAuth credential eviction to peer workers and expire admin session tombstones Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 3 +- .../mcp_server/byok_credential_cache.py | 38 ++++++ .../mcp_server/byok_oauth_endpoints.py | 2 +- .../mcp_server/oauth2_token_cache.py | 7 +- .../proxy/_experimental/mcp_server/server.py | 121 +++++++++--------- .../mcp_management_endpoints.py | 4 +- litellm/proxy/proxy_server.py | 3 +- .../mcp_server/test_byok_credential_cache.py | 57 +++++++++ .../mcp_server/test_byok_oauth_endpoints.py | 32 ++++- .../mcp_server/test_mcp_server.py | 70 +++++++++- .../mcp_server/test_oauth2_token_cache.py | 26 ++++ .../test_mcp_management_endpoints.py | 6 +- tests/test_litellm/proxy/test_proxy_server.py | 71 ++++++++++ 13 files changed, 366 insertions(+), 74 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/byok_credential_cache.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py diff --git a/litellm/constants.py b/litellm/constants.py index 2ebc9beb632..72fc6ed5160 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -184,7 +184,8 @@ MCP_METADATA_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_METADATA_TIMEOUT", "1 MCP_HEALTH_CHECK_TIMEOUT: Final = float(os.getenv("LITELLM_MCP_HEALTH_CHECK_TIMEOUT", "10.0")) MCP_TOOL_LISTING_MAX_PAGES: Final = 1000 MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH: Final = 8 -MCP_ADMIN_TERMINATED_SESSION_IDS_MAX: Final = 1024 +MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS: Final = 60 +MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE: Final = 4096 # Allowlist of commands permitted for MCP stdio transport. # Prevents arbitrary command execution via /mcp-rest/test/* endpoints or server creation. diff --git a/litellm/proxy/_experimental/mcp_server/byok_credential_cache.py b/litellm/proxy/_experimental/mcp_server/byok_credential_cache.py new file mode 100644 index 00000000000..3892015c405 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/byok_credential_cache.py @@ -0,0 +1,38 @@ +"""Per-worker cache of stored BYOK credentials, keyed so peer workers can evict it over the auth cache pub/sub.""" + +from dataclasses import dataclass +from typing import Final + +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.constants import MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE, MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS + +_CACHE_KEY_PREFIX: Final = "mcp_byok_credential" + + +@dataclass(frozen=True, slots=True) +class CachedByokCredential: + credential: str | None + + +byok_credential_cache: Final = InMemoryCache( + max_size_in_memory=MCP_BYOK_CREDENTIAL_CACHE_MAX_SIZE, + default_ttl=MCP_BYOK_CREDENTIAL_CACHE_TTL_SECONDS, +) + + +def byok_credential_cache_key(user_id: str, server_id: str) -> str: + return f"{_CACHE_KEY_PREFIX}:{user_id}:{server_id}" + + +def get_cached_byok_credential(user_id: str, server_id: str) -> CachedByokCredential | None: + cached: Final = byok_credential_cache.get_cache( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # InMemoryCache is untyped + byok_credential_cache_key(user_id, server_id) + ) + return cached if isinstance(cached, CachedByokCredential) else None + + +def cache_byok_credential(user_id: str, server_id: str, credential: str | None) -> None: + byok_credential_cache.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped + byok_credential_cache_key(user_id, server_id), + CachedByokCredential(credential=credential), + ) diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index 0ab76588b1f..2c63e0a96d8 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -865,7 +865,7 @@ async def byok_token( _invalidate_byok_cred_cache, ) - _invalidate_byok_cred_cache(user_id, server_id) + await _invalidate_byok_cred_cache(user_id, server_id) except Exception as exc: verbose_proxy_logger.error( "byok_token: failed to store user credential for user=%s server=%s: %s", diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 42edc2999ab..3742d7b4ccc 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -295,12 +295,15 @@ class MCPPerUserTokenCache: ) async def delete(self, user_id: str, server_id: str) -> None: - """Invalidate the cached token (removes from both in-memory and Redis layers).""" + """Invalidate the cached token in Redis, here, and in every peer worker's in-memory layer.""" try: + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( # noqa: PLC0415 # proxy import cycle + evict_and_broadcast, + ) from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 key: Final = self._cache_key(user_id, server_id) - await user_api_key_cache.async_delete_cache(key) + await evict_and_broadcast((key,), user_api_key_cache) except Exception as exc: verbose_logger.debug( "MCPPerUserTokenCache.delete failed for user=%s server=%s: %s", diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 80d9274859d..a7ca2775fb1 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -14,7 +14,7 @@ import time import traceback import types import uuid -from collections import Counter, deque +from collections import Counter from collections.abc import AsyncIterator, Callable, Mapping, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol @@ -30,7 +30,6 @@ from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.constants import ( MAXIMUM_TRACEBACK_LINES_TO_LOG, - MCP_ADMIN_TERMINATED_SESSION_IDS_MAX, MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -42,6 +41,12 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, _is_mcp_admitted_user_subject, ) +from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache, + byok_credential_cache_key, + cache_byok_credential, + get_cached_byok_credential, +) from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( get_request_base_url, ) @@ -86,6 +91,9 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( + publish_auth_cache_invalidation, +) from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, get_chain_id_from_headers, @@ -107,13 +115,6 @@ if TYPE_CHECKING: from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload -# Short-lived in-memory cache for BYOK credentials. -# Keyed by (user_id, server_id); value is (credential_or_None, monotonic_timestamp). -# Storing the credential value (not just a bool) means _get_byok_credential and -# _check_byok_credential share a single DB round-trip per TTL window. -_byok_cred_cache: Final[dict[tuple[str, str], tuple[str | None, float]]] = {} -_BYOK_CRED_CACHE_TTL: Final = 60 # seconds -_BYOK_CRED_CACHE_MAX_SIZE: Final = 4096 # cap to prevent unbounded growth _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each # `initialize` creates a session that survives until the idle timeout, so @@ -132,20 +133,11 @@ _MCP_TRANSPORT_SPAN_SCOPE_KEY: Final = "litellm_otel_transport_span" _MCP_DESTINATIONS_SCOPE_KEY: Final = "litellm_otel_request_destinations" -def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: - """Remove a (user_id, server_id) entry from the BYOK credential cache. - - Call this after storing or deleting a credential so subsequent calls - see the fresh value rather than a stale cached result. - """ - _byok_cred_cache.pop((user_id, server_id), None) - - -def _write_byok_cred_cache(user_id: str, server_id: str, credential: str | None) -> None: - """Write a credential value to the cache, evicting all entries if at capacity.""" - if len(_byok_cred_cache) >= _BYOK_CRED_CACHE_MAX_SIZE: - _byok_cred_cache.clear() - _byok_cred_cache[(user_id, server_id)] = (credential, time.monotonic()) +async def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: + """Drop a stored-or-deleted BYOK credential from this worker's cache and from every peer worker's.""" + cache_key: Final = byok_credential_cache_key(user_id, server_id) + byok_credential_cache.delete_cache(cache_key) + await publish_auth_cache_invalidation(cache_key=cache_key) # Check if MCP is available @@ -623,9 +615,7 @@ if MCP_AVAILABLE: _stateful_session_locks: Final[dict[str, asyncio.Lock]] = {} _stateful_session_active_request_counts: Final[dict[str, int]] = {} _stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown - _admin_terminated_session_ids: Final[deque[str]] = deque( # mutable-ok: bounded ring, appended on admin termination - maxlen=MCP_ADMIN_TERMINATED_SESSION_IDS_MAX - ) + _admin_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: admin-closed id -> last replay class _TerminableTransport(Protocol): async def terminate(self) -> None: ... @@ -697,6 +687,7 @@ if MCP_AVAILABLE: for session_id in list(_stateful_session_auth_context_last_seen): if session_id not in _stateful_session_auth_contexts: _remove_stateful_session_tracking(session_id) + _forget_expired_admin_terminated_session_ids(now) async def _enforce_stateful_session_cap_for_owner(owner: str) -> bool: """ @@ -2819,35 +2810,28 @@ if MCP_AVAILABLE: mcp_server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None, ) -> str | None: - """Retrieve the stored BYOK credential for a user+server pair. - - Uses the shared _byok_cred_cache to avoid a DB round-trip on every - tool call within the TTL window. - """ + """Retrieve the stored BYOK credential for a user+server pair, served from the worker cache within its TTL.""" if not mcp_server.is_byok: return None user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or "" if not user_id: return None - cache_key: Final = (user_id, mcp_server.server_id) - cached: Final = _byok_cred_cache.get(cache_key) + cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) if cached is not None: - credential, ts = cached - if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL: - return credential + return cached.credential from litellm.proxy._experimental.mcp_server.db import get_user_credential from litellm.proxy.proxy_server import prisma_client if prisma_client is None: return None - credential = await get_user_credential( + credential: Final = await get_user_credential( prisma_client=prisma_client, user_id=user_id, server_id=mcp_server.server_id, ) - _write_byok_cred_cache(user_id, mcp_server.server_id, credential) + cache_byok_credential(user_id, mcp_server.server_id, credential) return credential async def _check_byok_credential( @@ -2876,27 +2860,23 @@ if MCP_AVAILABLE: headers={"WWW-Authenticate": get_byok_www_authenticate()}, ) - # Check shared credential cache before hitting the DB. - cache_key: Final = (user_id, mcp_server.server_id) - cached: Final = _byok_cred_cache.get(cache_key) + cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) if cached is not None: - cached_cred, ts = cached - if time.monotonic() - ts < _BYOK_CRED_CACHE_TTL: - if cached_cred is None: - raise HTTPException( - status_code=401, - detail={ - "error": "byok_auth_required", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": ( - "No stored credential found for this BYOK server. " - "Complete the OAuth authorization flow to provide your API key." - ), - }, - headers={"WWW-Authenticate": get_byok_www_authenticate()}, - ) - return + if cached.credential is None: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={"WWW-Authenticate": get_byok_www_authenticate()}, + ) + return from litellm.proxy._experimental.mcp_server.db import get_user_credential from litellm.proxy.proxy_server import prisma_client @@ -2920,7 +2900,7 @@ if MCP_AVAILABLE: user_id=user_id, server_id=mcp_server.server_id, ) - _write_byok_cred_cache(user_id, mcp_server.server_id, credential) + cache_byok_credential(user_id, mcp_server.server_id, credential) if credential is None: raise HTTPException( status_code=401, @@ -3906,6 +3886,24 @@ if MCP_AVAILABLE: key_auth: Final = auth_user.user_api_key_auth return key_auth is not None and key_auth.user_id == user_id + def _forget_expired_admin_terminated_session_ids(now: float) -> None: + for session_id in [ + session_id + for session_id, last_replayed in _admin_terminated_session_ids.items() + if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + ]: + del _admin_terminated_session_ids[session_id] + + def _is_admin_terminated_session_id(session_id: str, now: float) -> bool: + last_replayed: Final = _admin_terminated_session_ids.get(session_id) + if last_replayed is None: + return False + if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: + del _admin_terminated_session_ids[session_id] + return False + _admin_terminated_session_ids[session_id] = now + return True + async def terminate_mcp_gateway_sessions( *, session_id_prefix: str | None = None, @@ -3919,6 +3917,7 @@ if MCP_AVAILABLE: admission. Only sessions held by this worker process are affected. """ now: Final = time.monotonic() + _forget_expired_admin_terminated_session_ids(now) server_instances: Final = _stateful_server_instances() targets: Final = tuple( (session_id, auth_user) @@ -3928,7 +3927,7 @@ if MCP_AVAILABLE: ) terminated: Final = tuple(_gateway_session_for(session_id, auth_user, now) for session_id, auth_user in targets) for session_id, _ in targets: - _admin_terminated_session_ids.append(session_id) + _admin_terminated_session_ids[session_id] = now transport = server_instances.pop(session_id, None) _remove_stateful_session_tracking(session_id) if transport is not None: @@ -4064,7 +4063,7 @@ if MCP_AVAILABLE: await success_response(scope, receive, send) return True - if _session_id in _admin_terminated_session_ids: + if _is_admin_terminated_session_id(_session_id, time.monotonic()): terminated_response: Final = JSONResponse( status_code=404, content={ # mutable-ok: JSONResponse content must be a plain dict diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 728ba9568e8..5326cf3415f 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2317,7 +2317,7 @@ if MCP_AVAILABLE: _invalidate_byok_cred_cache, ) - _invalidate_byok_cred_cache(user_id, server_id) + await _invalidate_byok_cred_cache(user_id, server_id) return MCPUserCredentialResponse(server_id=server_id, has_credential=True) # save=False: credential not persisted return MCPUserCredentialResponse(server_id=server_id, has_credential=False) @@ -2348,7 +2348,7 @@ if MCP_AVAILABLE: _invalidate_byok_cred_cache, ) - _invalidate_byok_cred_cache(target_user_id, server_id) + await _invalidate_byok_cred_cache(target_user_id, server_id) return MCPUserCredentialResponse(server_id=server_id, has_credential=False) # ── OAuth2 user-credential endpoints ────────────────────────────────────── diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index d7d8413d2ce..48122549aea 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -307,6 +307,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import ( from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.openai_like.model_info import MODEL_INFO_REFRESH_SECONDS from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot from litellm.proxy._types import * from litellm.proxy.analytics_endpoints.analytics_endpoints import ( @@ -7535,7 +7536,7 @@ class ProxyConfig: subscriber: Final = AuthCacheInvalidationSubscriber( redis_cache=redis_cache, user_api_key_cache=user_api_key_cache, - additional_in_memory_caches=(spend_counter_cache.in_memory_cache,), + additional_in_memory_caches=(spend_counter_cache.in_memory_cache, byok_credential_cache), ) self.auth_cache_invalidation_subscriber = subscriber subscriber.start() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py new file mode 100644 index 00000000000..8ec5b8642bc --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_credential_cache.py @@ -0,0 +1,57 @@ +import json + +import pytest + +from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + CachedByokCredential, + byok_credential_cache, + byok_credential_cache_key, + cache_byok_credential, + get_cached_byok_credential, +) +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + +class _FakeRedisCache: + namespace = None + + def init_async_client(self) -> object: + return object() + + +@pytest.fixture(autouse=True) +def _empty_cache(): + byok_credential_cache.flush_cache() + yield + byok_credential_cache.flush_cache() + + +def test_a_cached_negative_lookup_is_distinguishable_from_a_miss(): + assert get_cached_byok_credential("u-1", "srv-1") is None + cache_byok_credential("u-1", "srv-1", None) + assert get_cached_byok_credential("u-1", "srv-1") == CachedByokCredential(credential=None) + cache_byok_credential("u-1", "srv-1", "sk-stored") + assert get_cached_byok_credential("u-1", "srv-1") == CachedByokCredential(credential="sk-stored") + assert get_cached_byok_credential("u-1", "srv-2") is None + + +def test_peer_worker_invalidation_message_evicts_the_cached_credential(): + """The key a mutating worker broadcasts must be the key every other worker caches under.""" + cache_byok_credential("mallory", "srv-byok", "sk-revoked") + cache_byok_credential("alice", "srv-byok", "sk-kept") + subscriber = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(), # pyright: ignore[reportArgumentType] # subscriber is never started; only its message handler runs + user_api_key_cache=UserApiKeyCache(), + additional_in_memory_caches=(byok_credential_cache,), + ) + + subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler + { + "type": "message", + "data": json.dumps({"cache_key": byok_credential_cache_key("mallory", "srv-byok")}).encode(), + } + ) + + assert get_cached_byok_credential("mallory", "srv-byok") is None + assert get_cached_byok_credential("alice", "srv-byok") == CachedByokCredential(credential="sk-kept") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 55accfb169d..df0d7916baf 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -592,7 +592,7 @@ async def test_check_byok_credential_missing_credential(monkeypatch): monkeypatch.delenv("PROXY_BASE_URL", raising=False) monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) - monkeypatch.setattr(server_module, "_byok_cred_cache", {}) + server_module.byok_credential_cache.flush_cache() mock_prisma = MagicMock() with ( @@ -628,7 +628,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk from litellm.types.mcp_server.mcp_server_manager import MCPServer monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy") - monkeypatch.setattr(mcp_module, "_byok_cred_cache", {}) + mcp_module.byok_credential_cache.flush_cache() server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True) prisma = MagicMock() prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None) @@ -677,6 +677,34 @@ async def test_check_byok_credential_has_credential(): await _check_byok_credential(server, user_auth) +@pytest.mark.asyncio +async def test_invalidate_byok_cred_cache_evicts_locally_and_broadcasts_the_same_key(): + """A revoked credential must stop being served here and on every peer worker within the TTL.""" + from litellm.proxy._experimental.mcp_server import server as server_module + from litellm.proxy._experimental.mcp_server.byok_credential_cache import byok_credential_cache_key + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer(server_id="byok-revoke", name="byok-server", transport=MCPTransport.http, is_byok=True) + user_auth = UserAPIKeyAuth(user_id="mallory", api_key="sk-test") + server_module.byok_credential_cache.flush_cache() + db_lookup = AsyncMock(side_effect=["sk-before-revoke", None]) + publish = AsyncMock() + + with ( + patch("litellm.proxy._experimental.mcp_server.db.get_user_credential", new=db_lookup), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch.object(server_module, "publish_auth_cache_invalidation", new=publish), + ): + assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" + assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" + await server_module._invalidate_byok_cred_cache("mallory", "byok-revoke") + assert await server_module._get_byok_credential(server, user_auth) is None + + assert db_lookup.await_count == 2 + publish.assert_awaited_once_with(cache_key=byok_credential_cache_key("mallory", "byok-revoke")) + + @pytest.mark.asyncio async def test_check_byok_credential_db_unavailable_fails_closed(): """BYOK server with no prisma_client → 503, not silent pass. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 4311a69d465..f496f22d5d2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2989,9 +2989,10 @@ async def test_admin_terminated_session_id_gets_404_instead_of_a_fresh_stateless """Once an admin closes a session, a client replaying its id must not be silently upgraded to a new stateless session by the stale-header path; it gets 404 and has to initialize again.""" try: + from starlette.types import Scope + from litellm.proxy._experimental.mcp_server import server as mcp_server from litellm.proxy._experimental.mcp_server.server import session_manager_stateful - from starlette.types import Scope except ImportError: pytest.skip("MCP server not available") @@ -3048,6 +3049,73 @@ async def test_admin_terminated_session_id_gets_404_instead_of_a_fresh_stateless mcp_server._admin_terminated_session_ids.clear() +@pytest.mark.asyncio +async def test_admin_terminated_session_id_stays_refused_while_replayed_and_is_forgotten_like_an_idle_session(): + """The refusal window slides on every replay, so a client that keeps retrying is never silently + upgraded to a stateless session no matter how many other sessions an admin closes later; an id + nobody has replayed for a full idle timeout is dropped from the table by the idle sweep.""" + try: + from starlette.types import Scope + + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import session_manager_stateful + except ImportError: + pytest.skip("MCP server not available") + + idle_timeout = mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + retrying_id, silent_id = "admin-closed-retrying", "admin-closed-silent" + contexts = { + session_id: mcp_server.MCPAuthenticatedUser( + user_api_key_auth=UserAPIKeyAuth(api_key="key-alice", user_id="alice"), + ) + for session_id in (retrying_id, silent_id) + } + live_transports = {session_id: MagicMock(terminate=AsyncMock()) for session_id in contexts} + + async def replay(session_id: str, now: float) -> tuple[bool, list[bytes]]: + scope: Scope = { + "type": "http", + "method": "POST", + "headers": [(b"content-type", b"application/json"), (b"mcp-session-id", session_id.encode())], + } + with patch.object(mcp_server.time, "monotonic", return_value=now): + handled = await mcp_server._handle_stale_mcp_session( + scope, AsyncMock(), AsyncMock(), session_manager_stateful + ) + return handled, [k for k, _ in scope["headers"]] + + try: + with ( + patch.object( # test-quality-ok: the transport registry is a module-level singleton; the suite's only seam + session_manager_stateful, "_server_instances", live_transports + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_contexts, contexts, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_client_info, {}, clear=True + ), + patch.dict( # test-quality-ok: the session tables are module-level singletons; the suite's only seam + mcp_server._stateful_session_auth_context_last_seen, {}, clear=True + ), + ): + with patch.object(mcp_server.time, "monotonic", return_value=1000.0): + closed = await mcp_server.terminate_mcp_gateway_sessions(user_id="alice") + assert closed.terminated_sessions == 2 + + for elapsed in (idle_timeout - 1, 2 * idle_timeout - 2, 3 * idle_timeout - 3): + assert await replay(retrying_id, 1000.0 + elapsed) == (True, [b"content-type", b"mcp-session-id"]) + + await mcp_server._purge_expired_stateful_session_auth_contexts(now=1000.0 + idle_timeout) + assert set(mcp_server._admin_terminated_session_ids) == {retrying_id} + + assert await replay(silent_id, 1000.0 + idle_timeout) == (False, [b"content-type"]) + assert await replay(retrying_id, 1000.0 + 4 * idle_timeout) == (False, [b"content-type"]) + assert mcp_server._admin_terminated_session_ids == {} + finally: + mcp_server._admin_terminated_session_ids.clear() + + @pytest.mark.asyncio async def test_initialize_request_with_existing_session_tracks_new_session(): try: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py index f7567efcabc..301b3887922 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -395,6 +395,32 @@ async def test_invalidate_clears_every_identity_for_a_server(): assert mock_client.post.call_count == 3 +@pytest.mark.asyncio +async def test_per_user_token_delete_evicts_locally_and_broadcasts_to_peer_workers(): + """Revoking a user's OAuth token must not leave peer workers serving it from their in-memory layer.""" + from litellm.proxy import proxy_server + from litellm.proxy._experimental.mcp_server.oauth2_token_cache import MCPPerUserTokenCache + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + local_cache = UserApiKeyCache() + publish = AsyncMock() + token_cache = MCPPerUserTokenCache() + key = token_cache._cache_key("mallory", "srv-oauth") # pyright: ignore[reportPrivateUsage] # asserting the broadcast names the stored key + local_cache.in_memory_cache.set_cache(key, "encrypted-token") + + with ( + patch.object(proxy_server, "user_api_key_cache", local_cache), + patch( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", + new=publish, + ), + ): + await token_cache.delete("mallory", "srv-oauth") + + assert local_cache.in_memory_cache.get_cache(key) is None + publish.assert_awaited_once_with(cache_key=key) + + @pytest.mark.asyncio async def test_m2m_mint_uses_admin_entered_token_url_when_issuer_yield_empties_resolved(): """A pinned issuer empties the resolved token_url while configured_token_url keeps the diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 2f33018599f..a69cfe0ac7d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -5152,7 +5152,7 @@ async def test_admin_revokes_another_users_byok_credential(): ) delete_mock = AsyncMock(return_value=None) - invalidate_mock = MagicMock() + invalidate_mock = AsyncMock() with ( patch( # test-quality-ok: endpoint test stubs the Prisma client lookup "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -5174,7 +5174,7 @@ async def test_admin_revokes_another_users_byok_credential(): delete_mock.assert_awaited_once() assert delete_mock.await_args.args[1:] == ("mallory", "srv-byok-admin") - invalidate_mock.assert_called_once_with("mallory", "srv-byok-admin") + invalidate_mock.assert_awaited_once_with("mallory", "srv-byok-admin") assert result.has_credential is False @@ -5231,7 +5231,7 @@ async def test_user_naming_themselves_still_deletes_own_byok_credential(): new=delete_mock, ), patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam - mcp_server, "_invalidate_byok_cred_cache", new=MagicMock() + mcp_server, "_invalidate_byok_cred_cache", new=AsyncMock() ), ): await delete_mcp_user_credential( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 41c4956dba6..acd13fef9a6 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -13919,3 +13919,74 @@ async def test_token_counter_loads_a_custom_tokenizer_once_per_identifier_revisi ] finally: litellm.utils._select_custom_tokenizer_helper.cache_clear() + + +@pytest.mark.asyncio +async def test_auth_cache_invalidation_subscriber_evicts_byok_credentials_cached_by_this_worker(): + """A peer worker's BYOK revocation broadcast must reach this worker's BYOK credential cache.""" + from redis.asyncio import Redis + + from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache, + byok_credential_cache_key, + cache_byok_credential, + get_cached_byok_credential, + ) + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + class _QueuePubSub: + def __init__(self, messages: list[object]) -> None: + self.queue: asyncio.Queue[object] = asyncio.Queue() + for message in messages: + self.queue.put_nowait(message) + + async def subscribe(self, *channels: str) -> None: + return None + + async def get_message(self, *, ignore_subscribe_messages: bool, timeout: float) -> object | None: + try: + return await asyncio.wait_for(self.queue.get(), timeout) + except asyncio.TimeoutError: + return None + + async def aclose(self) -> None: + return None + + class _PubSubRedisClient(Redis): + def __init__(self, pubsub: _QueuePubSub) -> None: + self._scripted_pubsub = pubsub + + def pubsub(self) -> _QueuePubSub: + return self._scripted_pubsub + + class _FakeRedisCache: + namespace = None + + def __init__(self, client: object) -> None: + self._client = client + + def init_async_client(self) -> object: + return self._client + + byok_credential_cache.flush_cache() + cache_byok_credential("mallory", "srv-byok", "sk-revoked-elsewhere") + message: Final = { + "type": "message", + "data": json.dumps({"cache_key": byok_credential_cache_key("mallory", "srv-byok")}).encode(), + } + proxy_config: Final = proxy_server_module.ProxyConfig() + proxy_config.start_auth_cache_invalidation_subscriber( + redis_cache=_FakeRedisCache(_PubSubRedisClient(_QueuePubSub([message]))), # pyright: ignore[reportArgumentType] # fake pub/sub capable redis; no live redis in this unit test + user_api_key_cache=UserApiKeyCache(), + ) + try: + for _ in range(200): + if get_cached_byok_credential("mallory", "srv-byok") is None: + break + await asyncio.sleep(0.01) + evicted: Final = get_cached_byok_credential("mallory", "srv-byok") is None + finally: + await proxy_config.stop_auth_cache_invalidation_subscriber() + byok_credential_cache.flush_cache() + + assert evicted, "the subscriber does not evict the BYOK credential cache on a peer worker's broadcast" From ce48a3fbcce70cc17f877de30b9c2a64ed4bc278 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 02:49:41 +0000 Subject: [PATCH 111/253] test(mcp): give the new cache and tombstone patches TQ008 reasons and match the keyword eviction call Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/mcp_tests/test_per_user_oauth_cache.py | 4 +--- .../mcp_server/test_byok_oauth_endpoints.py | 12 +++++++++--- .../_experimental/mcp_server/test_mcp_server.py | 8 ++++++-- .../mcp_server/test_oauth2_token_cache.py | 6 ++++-- 4 files changed, 20 insertions(+), 10 deletions(-) diff --git a/tests/mcp_tests/test_per_user_oauth_cache.py b/tests/mcp_tests/test_per_user_oauth_cache.py index 141b906fce9..ac453df8fa5 100644 --- a/tests/mcp_tests/test_per_user_oauth_cache.py +++ b/tests/mcp_tests/test_per_user_oauth_cache.py @@ -331,9 +331,7 @@ class TestMCPPerUserTokenCache: with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_dual_cache): await cache.delete("alice", "slack-test") - mock_dual_cache.async_delete_cache.assert_called_once_with( - "mcp:per_user_token:alice:slack-test" - ) + mock_dual_cache.async_delete_cache.assert_called_once_with(key="mcp:per_user_token:alice:slack-test") mock_dual_cache.async_set_cache.assert_not_called() @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index df0d7916baf..87e23893616 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -692,9 +692,15 @@ async def test_invalidate_byok_cred_cache_evicts_locally_and_broadcasts_the_same publish = AsyncMock() with ( - patch("litellm.proxy._experimental.mcp_server.db.get_user_credential", new=db_lookup), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - patch.object(server_module, "publish_auth_cache_invalidation", new=publish), + patch( # test-quality-ok: the DB row lookup is the only seam below the credential resolver; no Prisma fake exists + "litellm.proxy._experimental.mcp_server.db.get_user_credential", new=db_lookup + ), + patch( # test-quality-ok: the resolver reads the module-level prisma_client singleton; the suite's only seam + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch.object( # test-quality-ok: the redis publisher is module-level; asserting the broadcast without a redis + server_module, "publish_auth_cache_invalidation", new=publish + ), ): assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index f496f22d5d2..0e2ceb98987 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -3078,7 +3078,9 @@ async def test_admin_terminated_session_id_stays_refused_while_replayed_and_is_f "method": "POST", "headers": [(b"content-type", b"application/json"), (b"mcp-session-id", session_id.encode())], } - with patch.object(mcp_server.time, "monotonic", return_value=now): + with patch.object( # test-quality-ok: the stale-session handler reads the clock directly; no injectable now + mcp_server.time, "monotonic", return_value=now + ): handled = await mcp_server._handle_stale_mcp_session( scope, AsyncMock(), AsyncMock(), session_manager_stateful ) @@ -3099,7 +3101,9 @@ async def test_admin_terminated_session_id_stays_refused_while_replayed_and_is_f mcp_server._stateful_session_auth_context_last_seen, {}, clear=True ), ): - with patch.object(mcp_server.time, "monotonic", return_value=1000.0): + with patch.object( # test-quality-ok: termination stamps the tombstone from the clock directly; no injectable now + mcp_server.time, "monotonic", return_value=1000.0 + ): closed = await mcp_server.terminate_mcp_gateway_sessions(user_id="alice") assert closed.terminated_sessions == 2 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py index 301b3887922..30d0f17a099 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -409,8 +409,10 @@ async def test_per_user_token_delete_evicts_locally_and_broadcasts_to_peer_worke local_cache.in_memory_cache.set_cache(key, "encrypted-token") with ( - patch.object(proxy_server, "user_api_key_cache", local_cache), - patch( + patch.object( # test-quality-ok: the token cache reads the module-level user_api_key_cache singleton; the suite's only seam + proxy_server, "user_api_key_cache", local_cache + ), + patch( # test-quality-ok: the redis publisher is module-level; asserting the broadcast without a redis "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", new=publish, ), From 85d338da8bc9bc1f9ea8fa7cc4ccb0f1940e4407 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 04:25:44 +0000 Subject: [PATCH 112/253] feat(dashscope): add qwen3.8-omni-flash and qwen3.8-flash to the model cost map Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 74 +++++++++++++++++++ model_prices_and_context_window.json | 74 +++++++++++++++++++ 2 files changed, 148 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 355e2b90e96..3409d9b4e46 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -16690,6 +16690,43 @@ "supports_tool_choice": true, "supports_vision": true }, + "dashscope/qwen3.8-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_web_search": true + }, + "dashscope/qwen3.8-omni-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", @@ -18594,6 +18631,43 @@ "supports_tool_choice": true, "supports_vision": true }, + "qwen_ai_platform/qwen3.8-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "qwen_ai_platform", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_web_search": true + }, + "qwen_ai_platform/qwen3.8-omni-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "qwen_ai_platform", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "qwen_ai_platform/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "qwen_ai_platform", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 355e2b90e96..3409d9b4e46 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16690,6 +16690,43 @@ "supports_tool_choice": true, "supports_vision": true }, + "dashscope/qwen3.8-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_web_search": true + }, + "dashscope/qwen3.8-omni-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", @@ -18594,6 +18631,43 @@ "supports_tool_choice": true, "supports_vision": true }, + "qwen_ai_platform/qwen3.8-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "qwen_ai_platform", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_web_search": true + }, + "qwen_ai_platform/qwen3.8-omni-flash": { + "cache_read_input_token_cost": 1.6e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "qwen_ai_platform", + "max_input_tokens": 991808, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.7e-07, + "source": "https://docs.modelstudio.console.alibabacloud.com/en/model-studio/model-pricing", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true + }, "qwen_ai_platform/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "qwen_ai_platform", From 4863e1775ba4630df7d6596b164b509d836357a2 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 04:33:16 +0000 Subject: [PATCH 113/253] feat(guardrails): add TypeSafe Jev relevance-based compaction guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- .../guardrails/auto_router_compression.py | 2 +- .../guardrail_hooks/typesafe/__init__.py | 80 ++++ .../guardrail_hooks/typesafe/typesafe.py | 390 ++++++++++++++++++ litellm/types/guardrails.py | 7 +- .../guardrails/guardrail_hooks/typesafe.py | 58 +++ .../guardrail_hooks/test_typesafe.py | 274 ++++++++++++ .../add_model/buildAutoRouterCompression.ts | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 9 files changed, 812 insertions(+), 5 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index b244678e201..a215cc573d7 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -10006,7 +10006,7 @@ }, "unreachable_fallback": { "default": "fail_closed", - "description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.", + "description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.", "enum": [ "fail_closed", "fail_open" diff --git a/litellm/proxy/guardrails/auto_router_compression.py b/litellm/proxy/guardrails/auto_router_compression.py index c37b9fff1f0..335419c6372 100644 --- a/litellm/proxy/guardrails/auto_router_compression.py +++ b/litellm/proxy/guardrails/auto_router_compression.py @@ -21,7 +21,7 @@ if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.router import Router -COMPRESSION_GUARDRAIL_PROVIDERS: Final = frozenset({"headroom", "compresr"}) +COMPRESSION_GUARDRAIL_PROVIDERS: Final = frozenset({"headroom", "compresr", "typesafe"}) _NO_COMPRESSION: Final = "none" # A ContextVar, not metadata: metadata reaches spend logs the caller can read, and a diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py new file mode 100644 index 00000000000..c347c863a1b --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -0,0 +1,80 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Final, cast + +from litellm.types.guardrails import ( + GuardrailEventHooks, + Mode, + SupportedGuardrailIntegrations, +) + +from .typesafe import TypeSafeGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def _coerce_event_hook( + mode: str | list[str] | Mode, +) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode: + if isinstance(mode, Mode): + return mode + if isinstance(mode, list): + return [GuardrailEventHooks(item) for item in mode] + return GuardrailEventHooks(mode) + + +def _get_optional_value(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> object: + if optional_params is not None: + value: Final = getattr(optional_params, attribute_name, None) + if value is not None: + return cast(object, value) + return cast(object, getattr(litellm_params, attribute_name, None)) + + +def _optional_float(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> float | None: + value: Final = _get_optional_value(litellm_params, optional_params, attribute_name) + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) + + +def _optional_int(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> int | None: + value: Final = _get_optional_value(litellm_params, optional_params, attribute_name) + if isinstance(value, bool) or not isinstance(value, int): + return None + return value + + +def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> TypeSafeGuardrail: + import litellm + + optional_params: Final = getattr(litellm_params, "optional_params", None) + + _callback: Final = TypeSafeGuardrail( + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + model=litellm_params.model, + relevance_threshold=_optional_float(litellm_params, optional_params, "relevance_threshold"), + min_chars_to_evaluate=_optional_int(litellm_params, optional_params, "min_chars_to_evaluate"), + max_result_chars_in_state=_optional_int(litellm_params, optional_params, "max_result_chars_in_state"), + guardrail_name=guardrail["guardrail_name"], + event_hook=_coerce_event_hook(litellm_params.mode), + default_on=litellm_params.default_on or False, + unreachable_fallback=( + litellm_params.unreachable_fallback if "unreachable_fallback" in litellm_params.model_fields_set else None + ), + ) + litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback manager is untyped + _callback + ) + return _callback + + +guardrail_initializer_registry: Final = { + SupportedGuardrailIntegrations.TYPESAFE.value: initialize_guardrail, +} + +guardrail_class_registry: Final = { + SupportedGuardrailIntegrations.TYPESAFE.value: TypeSafeGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py new file mode 100644 index 00000000000..ef192030e95 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -0,0 +1,390 @@ +"""TypeSafe (Jev) relevance-based compaction guardrail. + +Instead of summarizing tool output, the guardrail asks TypeSafe's Jev model +one yes/no question per completed tool exchange ("is this result still needed +for the current task?") over ``POST {api_base}/v1/systemone`` and blanks the +tool results Jev judges no longer relevant. The assistant tool-call rows stay +intact, so the conversation remains well-formed while the dead context stops +consuming input tokens. + +Exchanges follow litellm's own compression protection policy: system rows, the +last user row, and the last assistant row (which, expanded over its tool +exchange, covers the most recent exchange) are never evaluated or rewritten. +""" + +from __future__ import annotations + +import asyncio +import time +from collections.abc import Mapping, Sequence +from typing import TYPE_CHECKING, Annotated, Final, Literal, TypeGuard, cast + +import httpx +from fastapi import HTTPException +from httpx import Response as HttpxResponse +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.compression.compress import get_protected_indices +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, # pyright: ignore[reportUnknownVariableType] # decorator is untyped in custom_guardrail +) +from litellm.litellm_core_utils.prompt_templates.factory import group_tool_exchanges +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler + httpxSpecialProvider, +) +from litellm.proxy.guardrails.guardrail_hooks.content_text import content_to_text +from litellm.secret_managers.main import get_secret_str +from litellm.types.guardrails import GuardrailEventHooks, Mode +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import ( + Logging as LiteLLMLoggingObj, + ) + from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrailConfigModel, + ) + +DEFAULT_API_BASE: Final = "https://api.typesafe.ai" +DEFAULT_MODEL: Final = "jev-latest" +DEFAULT_RELEVANCE_THRESHOLD: Final = 0.2 +DEFAULT_MIN_CHARS_TO_EVALUATE: Final = 200 +DEFAULT_MAX_RESULT_CHARS_IN_STATE: Final = 4000 +_MAX_EXCHANGES_EVALUATED: Final = 200 +# The shared GuardrailCallback client carries no per-call bound; an on-request +# guardrail must not hold the caller's request for the client's pooled timeout. +_JEV_TIMEOUT_SECONDS: Final = 30.0 +DROPPED_RESULT_TEXT: Final = ( + "[Tool result removed by TypeSafe compaction: judged no longer relevant to the current task]" +) + + +def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip + return isinstance(value, dict) + + +def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip + return isinstance(value, list) + + +def _safe_response_text(response: object, limit: int = 500) -> str: + try: + text: Final = getattr(response, "text", "") + except httpx.DecodingError: + return "" + return (text or "")[:limit] + + +class _JevNoulAnswer(BaseModel): + model_config = ConfigDict(frozen=True, allow_inf_nan=False) + + type: Literal["noul"] + noul: Annotated[float, Field(ge=0.0, le=1.0)] + + +class _JevSystemOneResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + answers: Mapping[str, _JevNoulAnswer] + + +_JEV_RESPONSE_ADAPTER: Final = TypeAdapter(_JevSystemOneResponse) + + +def _question_instructions(question_id: str) -> str: + return ( + f"Is tool exchange `{question_id}` in `tool_exchanges` still needed by the assistant to " + "complete `task`? Answer yes if its result contains information the assistant has not yet " + "fully used or will need again; answer no if it is off-topic, superseded, or already " + "incorporated into later messages." + ) + + +def _tool_call_entries(assistant_message: Mapping[str, object]) -> list[dict[str, object]]: + tool_calls: Final = assistant_message.get("tool_calls") + if not _is_object_list(tool_calls): + return [] + entries: Final[list[dict[str, object]]] = [] + for tool_call in tool_calls: + if not _is_str_object_dict(tool_call): + continue + function = tool_call.get("function") + fn = function if _is_str_object_dict(function) else tool_call + entries.append({"name": fn.get("name"), "arguments": fn.get("arguments")}) + return entries + + +def _protected_indices(messages: Sequence[Mapping[str, object]]) -> frozenset[int]: + """Rows typesafe must not rewrite, expanded over whole tool exchanges. + + ``get_protected_indices`` covers system rows, the last user row, the last + assistant row, and cache_control prefixes. Expanding over exchanges keeps an + exchange atomic: the last assistant row protects its own tool results too, + so the most recent exchange is never evaluated. + """ + protected: Final = frozenset(get_protected_indices(messages)) + return protected | frozenset( + index + for group in group_tool_exchanges(messages) + if any(member in protected for member in group) + for index in group + ) + + +class TypeSafeGuardrail(CustomGuardrail): + def __init__( + self, + api_base: str | None = None, + api_key: str | None = None, + model: str | None = None, + relevance_threshold: float | None = None, + min_chars_to_evaluate: int | None = None, + max_result_chars_in_state: int | None = None, + unreachable_fallback: str | None = None, + guardrail_name: str | None = None, + event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, + default_on: bool = False, + async_handler: AsyncHTTPHandler | None = None, + ): + raw_api_base: Final = (api_base or get_secret_str("TYPESAFE_API_BASE") or DEFAULT_API_BASE).rstrip("/") + self.typesafe_api_base = raw_api_base + self.typesafe_api_key = api_key or get_secret_str("TYPESAFE_API_KEY") + if not self.typesafe_api_key: + raise ValueError( + "TypeSafe guardrail requires an API key. Set `api_key` in the " + "guardrail config or the TYPESAFE_API_KEY env var." + ) + self.jev_model = model or DEFAULT_MODEL + self.relevance_threshold = DEFAULT_RELEVANCE_THRESHOLD if relevance_threshold is None else relevance_threshold + self.min_chars_to_evaluate = ( + DEFAULT_MIN_CHARS_TO_EVALUATE if min_chars_to_evaluate is None else min_chars_to_evaluate + ) + self.max_result_chars_in_state = ( + DEFAULT_MAX_RESULT_CHARS_IN_STATE if max_result_chars_in_state is None else max_result_chars_in_state + ) + self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( + "fail_closed" if unreachable_fallback == "fail_closed" else "fail_open" + ) + self.async_handler: AsyncHTTPHandler = async_handler or get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback, + ) + super().__init__( # pyright: ignore[reportUnknownMemberType] # CustomGuardrail.__init__ is untyped + guardrail_name=guardrail_name, + event_hook=event_hook, + default_on=default_on, + ) + + def _handle_failure(self, error: str, log_detail: dict[str, object]) -> None: + """fail_open logs and the caller forwards uncompacted; fail_closed raises. + Upstream bodies go to server logs only; the raised HTTPException is generic.""" + if self.unreachable_fallback == "fail_open": + verbose_proxy_logger.warning( + "TypeSafe: %s; fail_open configured, forwarding request uncompacted. detail=%s", + error, + log_detail, + ) + return + verbose_proxy_logger.error("TypeSafe: %s. detail=%s", error, log_detail) + raise HTTPException(status_code=500, detail={"error": error}) + + def _candidate_exchanges(self, messages: list[dict[str, object]]) -> list[tuple[int, ...]]: + """Message-index groups eligible for relevance evaluation, oldest first. + + A candidate is a completed tool exchange: an assistant row that made + tool calls plus at least one ``tool``/``function`` row answering it, + with no member protected, and enough combined tool-result text to be + worth an evaluation call. + """ + protected: Final = _protected_indices(messages) + candidates: Final[list[tuple[int, ...]]] = [] + for group in group_tool_exchanges(messages): + if len(group) < 2: + continue + if messages[group[0]].get("role") != "assistant": + continue + if any(member in protected for member in group): + continue + tool_text = "".join( + content_to_text(messages[index].get("content")) + for index in group[1:] + if messages[index].get("role") in ("tool", "function") + ) + if not tool_text or len(tool_text) < self.min_chars_to_evaluate: + continue + candidates.append(group) + return candidates[-_MAX_EXCHANGES_EVALUATED:] + + def _build_state(self, messages: list[dict[str, object]], candidates: list[tuple[int, ...]]) -> dict[str, object]: + task: Final = next( + ( + content_to_text(messages[index].get("content")) + for index in range(len(messages) - 1, -1, -1) + if messages[index].get("role") == "user" + ), + "", + ) + system: Final = "\n\n".join( + content_to_text(message.get("content")) for message in messages if message.get("role") == "system" + ) + tool_exchanges: Final[dict[str, object]] = {} + for ordinal, group in enumerate(candidates): + result_text = "".join( + content_to_text(messages[index].get("content")) + for index in group[1:] + if messages[index].get("role") in ("tool", "function") + ) + tool_exchanges[f"e{ordinal}"] = { + "tool_calls": _tool_call_entries(messages[group[0]]), + "result": result_text[: self.max_result_chars_in_state], + } + return {"task": task, "system": system, "tool_exchanges": tool_exchanges} + + async def _call_systemone(self, state: dict[str, object], question_ids: list[str]) -> _JevSystemOneResponse | None: + """Evaluate each exchange. Returns the response, or None when the service + failed and fail_open applies.""" + payload: Final[dict[str, object]] = { + "model": self.jev_model, + "state": state, + "questions": { + question_id: {"type": "noul", "instructions": _question_instructions(question_id)} + for question_id in question_ids + }, + } + try: + raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped + url=f"{self.typesafe_api_base}/v1/systemone", + json=payload, + headers={ + "Authorization": f"Bearer {self.typesafe_api_key}", + "Content-Type": "application/json", + }, + timeout=_JEV_TIMEOUT_SECONDS, + ) + except asyncio.CancelledError: + raise + except httpx.HTTPStatusError as e: + resp: Final = getattr(e, "response", None) + self._handle_failure( + "TypeSafe evaluation service returned an error", + {"status_code": getattr(resp, "status_code", None), "body": _safe_response_text(resp)}, + ) + return None + except (httpx.RequestError, litellm.Timeout) as e: + self._handle_failure("TypeSafe evaluation service request failed", {"detail": str(e)}) + return None + except Exception as e: + self._handle_failure("TypeSafe evaluation service request failed", {"detail": str(e)}) + return None + if not 200 <= raw_response.status_code < 300: + self._handle_failure( + "TypeSafe evaluation service returned an error", + {"status_code": raw_response.status_code, "body": _safe_response_text(raw_response)}, + ) + return None + try: + body: Final = cast(object, raw_response.json()) + except (ValueError, httpx.DecodingError, RecursionError): + self._handle_failure( + "TypeSafe evaluation service returned an unreadable response", + {"body": _safe_response_text(raw_response)}, + ) + return None + try: + return _JEV_RESPONSE_ADAPTER.validate_python(body) + except ValidationError: + self._handle_failure( + "TypeSafe evaluation service returned unexpected response shape", + {"body": _safe_response_text(raw_response)}, + ) + return None + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + logging_obj: LiteLLMLoggingObj | None = None, + ) -> GenericGuardrailAPIInputs: + if input_type != "request": + return inputs + + structured_messages: Final = inputs.get("structured_messages") + if not _is_object_list(structured_messages) or not structured_messages: + return inputs + messages: Final = [m for m in structured_messages if _is_str_object_dict(m)] + if len(messages) != len(structured_messages): + return inputs + + candidates: Final = self._candidate_exchanges(messages) + if not candidates: + verbose_proxy_logger.debug("TypeSafe: no completed tool exchanges eligible for evaluation") + return inputs + + question_ids: Final = [f"e{ordinal}" for ordinal in range(len(candidates))] + state: Final = self._build_state(messages, candidates) + + start_time: Final = time.monotonic() + response: Final = await self._call_systemone(state, question_ids) + end_time: Final = time.monotonic() + if response is None: + return inputs + + dropped_ordinals: Final = frozenset( + ordinal + for ordinal in range(len(candidates)) + if (answer := response.answers.get(f"e{ordinal}")) is not None and answer.noul < self.relevance_threshold + ) + dropped_tool_indices: Final[frozenset[int]] = frozenset( + index + for ordinal in dropped_ordinals + for index in candidates[ordinal][1:] + if messages[index].get("role") in ("tool", "function") + ) + if not dropped_tool_indices: + verbose_proxy_logger.debug("TypeSafe: all evaluated exchanges still relevant; request unchanged") + return inputs + + compacted_messages: Final = [ + {**message, "content": DROPPED_RESULT_TEXT} if index in dropped_tool_indices else message + for index, message in enumerate(messages) + ] + chars_removed: Final = sum( + len(content_to_text(messages[index].get("content"))) - len(DROPPED_RESULT_TEXT) + for index in dropped_tool_indices + ) + exchanges_dropped: Final = len(dropped_ordinals) + verbose_proxy_logger.info( + "TypeSafe: evaluated %s tool exchange(s), dropped %s, ~%s chars removed", + len(candidates), + exchanges_dropped, + chars_removed, + ) + self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper + guardrail_json_response={ + "exchanges_evaluated": len(candidates), + "exchanges_dropped": exchanges_dropped, + "chars_removed": chars_removed, + "model": self.jev_model, + }, + request_data=request_data, + guardrail_status="success", + guardrail_provider="typesafe", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + return {**inputs, "structured_messages": compacted_messages} # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime + + @staticmethod + def get_config_model() -> type[TypeSafeGuardrailConfigModel] | None: + from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrailConfigModel, + ) + + return TypeSafeGuardrailConfigModel diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 6346e13f3ba..400cadd69e7 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -62,6 +62,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.singulr import ( from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( ToolPermissionGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import ( VigilGuardGuardrailConfigModel, ) @@ -138,6 +141,7 @@ class SupportedGuardrailIntegrations(Enum): SINGULR = "singulr" HEADROOM = "headroom" COMPRESR = "compresr" + TYPESAFE = "typesafe" STRAIKER = "straiker" ALICE = "alice" AGENT_365 = "agent_365" @@ -1055,7 +1059,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up default="fail_closed", description=( "Behavior when a guardrail endpoint is unreachable due to network errors. " - "Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. " + "Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. " "'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed." ), ) @@ -1171,6 +1175,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o LakeraV2GuardrailConfigModel, HeadroomGuardrailConfigModel, CompresrGuardrailConfigModel, + TypeSafeGuardrailConfigModel, RepelloAIGuardrailConfigModel, LassoGuardrailConfigModel, DeepKeepGuardrailConfigModel, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py b/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py new file mode 100644 index 00000000000..4e742bfc7be --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py @@ -0,0 +1,58 @@ +from typing import Literal + +from pydantic import BaseModel, Field + +from .base import GuardrailConfigModel + + +class TypeSafeGuardrailOptionalParams(BaseModel): + """Optional tuning knobs for the TypeSafe (Jev) compaction guardrail.""" + + relevance_threshold: float | None = Field( + default=None, + description=( + "Relevance cutoff in [0, 1]. A completed tool exchange is dropped when Jev " + "scores the probability that it is still needed below this value. Defaults to 0.2." + ), + ) + min_chars_to_evaluate: int | None = Field( + default=None, + description=( + "Skip tool exchanges whose combined tool-result text is shorter than this many characters. Defaults to 200." + ), + ) + max_result_chars_in_state: int | None = Field( + default=None, + description=( + "Tool result text is truncated to this many characters when sent to the Jev evaluator. Defaults to 4000." + ), + ) + + +class TypeSafeGuardrailConfigModel(GuardrailConfigModel[TypeSafeGuardrailOptionalParams]): + api_key: str | None = Field( + default=None, + description="TypeSafe API key, sent as a Bearer token. Falls back to the TYPESAFE_API_KEY env var.", + ) + api_base: str | None = Field( + default=None, + description=( + "Base URL of the TypeSafe API. Falls back to the TYPESAFE_API_BASE env var, then https://api.typesafe.ai." + ), + ) + model: str | None = Field( + default=None, + description="TypeSafe evaluation model (not the LLM). Defaults to 'jev-latest'.", + ) + unreachable_fallback: Literal["fail_closed", "fail_open"] = Field( + default="fail_open", + description=( + "Behavior when the TypeSafe evaluation service is unreachable or errors. " + "'fail_open' (default) forwards the request uncompacted. 'fail_closed' " + "raises an error instead." + ), + ) + + @staticmethod + def ui_friendly_name() -> str: + return "TypeSafe (Jev) Compaction" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py new file mode 100644 index 00000000000..293e4f8c344 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -0,0 +1,274 @@ +""" +Unit tests for the TypeSafe (Jev) compaction guardrail. + +Tests cover: +- exchanges scored below relevance_threshold have their tool rows blanked while + assistant tool-call rows and kept exchanges pass through verbatim, without + mutating the caller's message list +- protected rows (system, last user, and the last tool exchange via the + last-assistant rule) are never sent to Jev even when long +- exchanges under min_chars_to_evaluate are skipped +- request shape: POST {api_base}/v1/systemone with Bearer auth, one noul + question per candidate keyed e, task = last user text, results truncated + to max_result_chars_in_state +- identity return when there are no candidates or nothing is dropped +- fail_open forwards uncompacted on service failure; fail_closed raises +- response input_type passthrough and initialize_guardrail wiring +""" + +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrail, + guardrail_class_registry, + guardrail_initializer_registry, + initialize_guardrail, +) +from litellm.proxy.guardrails.guardrail_hooks.typesafe.typesafe import DROPPED_RESULT_TEXT +from litellm.types.guardrails import SupportedGuardrailIntegrations +from litellm.types.utils import GenericGuardrailAPIInputs + +FAKE_API_BASE = "https://typesafe.example.com" +FAKE_API_KEY = "ts_test-key" + +SYSTEM_TEXT = "You are a research assistant." +USER_TEXT = "Which 2026 EV has the longest range?" +TOOL_OUTPUT_LONG = "Result: EV range comparison. " * 40 # > 200 chars +TOOL_OUTPUT_SHORT = "short" + + +def _exchange(call_id: str, tool_text: str, name: str = "web_search") -> list[dict]: + return [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": call_id, + "type": "function", + "function": {"name": name, "arguments": '{"query": "ev"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": call_id, "name": name, "content": tool_text}, + ] + + +def _messages(*, tail: list | None = None) -> list[dict]: + base = [ + {"role": "system", "content": SYSTEM_TEXT}, + {"role": "user", "content": USER_TEXT}, + ] + return base + (tail or []) + + +def _make_guardrail(handler: MagicMock | None = None, **kwargs) -> TypeSafeGuardrail: + defaults = dict( + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + guardrail_name="typesafe", + default_on=True, + async_handler=handler or _make_handler({"e0": 0.9}), + ) + defaults.update(kwargs) + return TypeSafeGuardrail(**defaults) + + +def _make_handler(answers: dict[str, float], status: int = 200) -> MagicMock: + response = MagicMock() + response.status_code = status + response.json.return_value = { + "model": "jev-1.13.0", + "answers": {qid: {"type": "noul", "noul": score} for qid, score in answers.items()}, + "usage": {"input_tokens": 10, "output_tokens": 1}, + } + response.text = "" + handler = MagicMock() + handler.post = AsyncMock(return_value=response) + return handler + + +def _inputs(messages: list) -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs(structured_messages=messages) + + +async def _apply(guardrail: TypeSafeGuardrail, messages: list, input_type: str = "request"): + return await guardrail.apply_guardrail( + inputs=_inputs(messages), + request_data={}, + input_type=input_type, # pyright: ignore[reportArgumentType] # test uses the same literal domain + logging_obj=None, + ) + + +@pytest.mark.asyncio +async def test_low_noul_exchange_blanked_high_kept_and_input_not_mutated(): + handler = _make_handler({"e0": 0.1, "e1": 0.95}) + guardrail = _make_guardrail(handler) + messages = _messages( + tail=[ + *_exchange("call_1", TOOL_OUTPUT_LONG), + *_exchange("call_2", TOOL_OUTPUT_LONG), + {"role": "assistant", "content": "still thinking"}, + ] + ) + snapshot = [dict(m) for m in messages] + + result = await _apply(guardrail, messages) + out = result["structured_messages"] + + assert out[3]["content"] == DROPPED_RESULT_TEXT + assert out[3]["tool_call_id"] == "call_1" + assert out[3]["role"] == "tool" + assert out[5]["content"] == TOOL_OUTPUT_LONG + assert out[2] == messages[2] + assert out[4] == messages[4] + assert out[6]["content"] == "still thinking" + assert messages == snapshot + + +@pytest.mark.asyncio +async def test_last_exchange_and_protected_rows_never_evaluated(): + handler = _make_handler({"e0": 0.05}) + guardrail = _make_guardrail(handler) + # Ends on a tool result: the last assistant row is protected, so the whole + # last exchange is out of scope even though its text is long. + messages = _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), *_exchange("call_2", TOOL_OUTPUT_LONG)]) + + result = await _apply(guardrail, messages) + + payload = handler.post.call_args.kwargs["json"] + assert list(payload["questions"]) == ["e0"] + assert list(payload["state"]["tool_exchanges"]) == ["e0"] + assert payload["state"]["task"] == USER_TEXT + assert payload["state"]["system"] == SYSTEM_TEXT + out = result["structured_messages"] + assert out[3]["content"] == DROPPED_RESULT_TEXT + assert out[5]["content"] == TOOL_OUTPUT_LONG + + +@pytest.mark.asyncio +async def test_short_exchange_not_sent(): + handler = _make_handler({"e0": 0.9}) + guardrail = _make_guardrail(handler) + messages = _messages( + tail=[ + *_exchange("call_1", TOOL_OUTPUT_SHORT), + *_exchange("call_2", TOOL_OUTPUT_LONG), + {"role": "assistant", "content": "done"}, + ] + ) + result = await _apply(guardrail, messages) + payload = handler.post.call_args.kwargs["json"] + assert list(payload["questions"]) == ["e0"] + exchange = payload["state"]["tool_exchanges"]["e0"] + assert exchange["result"] == TOOL_OUTPUT_LONG + assert result is not None + + +@pytest.mark.asyncio +async def test_request_body_shape_and_truncation(): + handler = _make_handler({"e0": 0.9}) + guardrail = _make_guardrail(handler, max_result_chars_in_state=50) + messages = _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "done"}]) + await _apply(guardrail, messages) + + kwargs = handler.post.call_args.kwargs + assert kwargs["url"].endswith("/v1/systemone") + assert kwargs["url"].startswith(FAKE_API_BASE) + assert kwargs["headers"]["Authorization"] == f"Bearer {FAKE_API_KEY}" + assert kwargs["headers"]["Content-Type"] == "application/json" + payload = kwargs["json"] + assert payload["model"] == "jev-latest" + assert list(payload["questions"]) == ["e0"] + assert payload["questions"]["e0"]["type"] == "noul" + assert "e0" in payload["questions"]["e0"]["instructions"] + assert payload["state"]["task"] == USER_TEXT + exchange = payload["state"]["tool_exchanges"]["e0"] + assert exchange["result"] == TOOL_OUTPUT_LONG[:50] + assert exchange["tool_calls"] == [{"name": "web_search", "arguments": '{"query": "ev"}'}] + + +@pytest.mark.asyncio +async def test_no_candidates_returns_identity_and_skips_http(): + handler = _make_handler({}) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[{"role": "assistant", "content": "plain answer"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + handler.post.assert_not_called() + + +@pytest.mark.asyncio +async def test_all_above_threshold_returns_identity(): + handler = _make_handler({"e0": 0.9}) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_fail_open_returns_inputs_on_exception(): + handler = MagicMock() + handler.post = AsyncMock(side_effect=Exception("connection refused")) + guardrail = _make_guardrail(handler, unreachable_fallback="fail_open") + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_fail_closed_raises_http_exception(): + handler = MagicMock() + handler.post = AsyncMock(side_effect=Exception("connection refused")) + guardrail = _make_guardrail(handler, unreachable_fallback="fail_closed") + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + with pytest.raises(HTTPException): + await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + + +@pytest.mark.asyncio +async def test_fail_open_on_non_2xx(): + handler = _make_handler({"e0": 0.9}, status=500) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_response_input_type_passthrough(): + handler = _make_handler({"e0": 0.05}) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG)])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="response", logging_obj=None) + assert result is inputs + handler.post.assert_not_called() + + +def test_initialize_guardrail_applies_optional_params_and_registry_keys(): + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="typesafe", + mode="pre_call", + api_key=FAKE_API_KEY, + api_base=FAKE_API_BASE, + optional_params={ + "relevance_threshold": 0.5, + "min_chars_to_evaluate": 10, + "max_result_chars_in_state": 100, + }, + ) + callback = initialize_guardrail(litellm_params, {"guardrail_name": "jev-compaction"}) + assert isinstance(callback, TypeSafeGuardrail) + assert callback.relevance_threshold == 0.5 + assert callback.min_chars_to_evaluate == 10 + assert callback.max_result_chars_in_state == 100 + assert callback.unreachable_fallback == "fail_open" + assert guardrail_initializer_registry[SupportedGuardrailIntegrations.TYPESAFE.value] is initialize_guardrail + assert guardrail_class_registry[SupportedGuardrailIntegrations.TYPESAFE.value] is TypeSafeGuardrail diff --git a/ui/litellm-dashboard/src/components/add_model/buildAutoRouterCompression.ts b/ui/litellm-dashboard/src/components/add_model/buildAutoRouterCompression.ts index c3503f78afb..292b497f6df 100644 --- a/ui/litellm-dashboard/src/components/add_model/buildAutoRouterCompression.ts +++ b/ui/litellm-dashboard/src/components/add_model/buildAutoRouterCompression.ts @@ -14,7 +14,7 @@ export const NO_COMPRESSION = "none"; /** Guardrail providers that compress prompts, mirroring COMPRESSION_GUARDRAIL_PROVIDERS in * litellm/proxy/guardrails/auto_router_compression.py. Both are selectable per hop. */ -export const COMPRESSION_GUARDRAIL_PROVIDERS: readonly string[] = ["headroom", "compresr"]; +export const COMPRESSION_GUARDRAIL_PROVIDERS: readonly string[] = ["headroom", "compresr", "typesafe"]; export const isCompressionGuardrailProvider = (provider: unknown): boolean => typeof provider === "string" && COMPRESSION_GUARDRAIL_PROVIDERS.includes(provider.toLowerCase()); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fd882937e79..16d61bd9b0f 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -24272,7 +24272,7 @@ export interface components { timeout?: number | null; /** * Unreachable Fallback - * @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed. + * @description Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', 'compresr', and 'typesafe'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed. * @default fail_closed * @enum {string} */ From dffb6a38d95f5e2fa2c259f06f50243c46ca2ab2 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 04:35:58 +0000 Subject: [PATCH 114/253] refactor(guardrails): tighten typesafe guardrail typing and error handling Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/__init__.py | 42 +++++-------- .../guardrail_hooks/typesafe/typesafe.py | 63 +++++++------------ .../guardrail_hooks/test_typesafe.py | 3 +- 3 files changed, 40 insertions(+), 68 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py index c347c863a1b..40f040eb676 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -1,12 +1,17 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Final, cast +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel from litellm.types.guardrails import ( GuardrailEventHooks, Mode, SupportedGuardrailIntegrations, ) +from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrailOptionalParams, +) from .typesafe import TypeSafeGuardrail @@ -24,40 +29,27 @@ def _coerce_event_hook( return GuardrailEventHooks(mode) -def _get_optional_value(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> object: - if optional_params is not None: - value: Final = getattr(optional_params, attribute_name, None) - if value is not None: - return cast(object, value) - return cast(object, getattr(litellm_params, attribute_name, None)) - - -def _optional_float(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> float | None: - value: Final = _get_optional_value(litellm_params, optional_params, attribute_name) - if isinstance(value, bool) or not isinstance(value, (int, float)): - return None - return float(value) - - -def _optional_int(litellm_params: LitellmParams, optional_params: object | None, attribute_name: str) -> int | None: - value: Final = _get_optional_value(litellm_params, optional_params, attribute_name) - if isinstance(value, bool) or not isinstance(value, int): - return None - return value +def _optional_params(litellm_params: LitellmParams) -> TypeSafeGuardrailOptionalParams: + value: Final = litellm_params.optional_params + if isinstance(value, TypeSafeGuardrailOptionalParams): + return value + if isinstance(value, BaseModel): + return TypeSafeGuardrailOptionalParams.model_validate(value.model_dump()) + return TypeSafeGuardrailOptionalParams() def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> TypeSafeGuardrail: import litellm - optional_params: Final = getattr(litellm_params, "optional_params", None) + optional_params: Final = _optional_params(litellm_params) _callback: Final = TypeSafeGuardrail( api_base=litellm_params.api_base, api_key=litellm_params.api_key, model=litellm_params.model, - relevance_threshold=_optional_float(litellm_params, optional_params, "relevance_threshold"), - min_chars_to_evaluate=_optional_int(litellm_params, optional_params, "min_chars_to_evaluate"), - max_result_chars_in_state=_optional_int(litellm_params, optional_params, "max_result_chars_in_state"), + relevance_threshold=optional_params.relevance_threshold, + min_chars_to_evaluate=optional_params.min_chars_to_evaluate, + max_result_chars_in_state=optional_params.max_result_chars_in_state, guardrail_name=guardrail["guardrail_name"], event_hook=_coerce_event_hook(litellm_params.mode), default_on=litellm_params.default_on or False, diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index ef192030e95..360219e0ad5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -3,13 +3,7 @@ Instead of summarizing tool output, the guardrail asks TypeSafe's Jev model one yes/no question per completed tool exchange ("is this result still needed for the current task?") over ``POST {api_base}/v1/systemone`` and blanks the -tool results Jev judges no longer relevant. The assistant tool-call rows stay -intact, so the conversation remains well-formed while the dead context stops -consuming input tokens. - -Exchanges follow litellm's own compression protection policy: system rows, the -last user row, and the last assistant row (which, expanded over its tool -exchange, covers the most recent exchange) are never evaluated or rewritten. +tool results Jev judges no longer relevant. """ from __future__ import annotations @@ -24,7 +18,6 @@ from fastapi import HTTPException from httpx import Response as HttpxResponse from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError -import litellm from litellm._logging import verbose_proxy_logger from litellm.compression.compress import get_protected_indices from litellm.integrations.custom_guardrail import ( @@ -56,8 +49,6 @@ DEFAULT_RELEVANCE_THRESHOLD: Final = 0.2 DEFAULT_MIN_CHARS_TO_EVALUATE: Final = 200 DEFAULT_MAX_RESULT_CHARS_IN_STATE: Final = 4000 _MAX_EXCHANGES_EVALUATED: Final = 200 -# The shared GuardrailCallback client carries no per-call bound; an on-request -# guardrail must not hold the caller's request for the client's pooled timeout. _JEV_TIMEOUT_SECONDS: Final = 30.0 DROPPED_RESULT_TEXT: Final = ( "[Tool result removed by TypeSafe compaction: judged no longer relevant to the current task]" @@ -72,9 +63,11 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin return isinstance(value, list) -def _safe_response_text(response: object, limit: int = 500) -> str: +def _safe_response_text(response: HttpxResponse | None, limit: int = 500) -> str: + if response is None: + return "" try: - text: Final = getattr(response, "text", "") + text: Final = response.text except httpx.DecodingError: return "" return (text or "")[:limit] @@ -120,13 +113,7 @@ def _tool_call_entries(assistant_message: Mapping[str, object]) -> list[dict[str def _protected_indices(messages: Sequence[Mapping[str, object]]) -> frozenset[int]: - """Rows typesafe must not rewrite, expanded over whole tool exchanges. - - ``get_protected_indices`` covers system rows, the last user row, the last - assistant row, and cache_control prefixes. Expanding over exchanges keeps an - exchange atomic: the last assistant row protects its own tool results too, - so the most recent exchange is never evaluated. - """ + """``get_protected_indices`` expanded over whole tool exchanges, so the most recent exchange is never evaluated.""" protected: Final = frozenset(get_protected_indices(messages)) return protected | frozenset( index @@ -180,8 +167,7 @@ class TypeSafeGuardrail(CustomGuardrail): ) def _handle_failure(self, error: str, log_detail: dict[str, object]) -> None: - """fail_open logs and the caller forwards uncompacted; fail_closed raises. - Upstream bodies go to server logs only; the raised HTTPException is generic.""" + """fail_open logs and returns; fail_closed raises a generic 502 (upstream bodies stay in server logs).""" if self.unreachable_fallback == "fail_open": verbose_proxy_logger.warning( "TypeSafe: %s; fail_open configured, forwarding request uncompacted. detail=%s", @@ -190,16 +176,10 @@ class TypeSafeGuardrail(CustomGuardrail): ) return verbose_proxy_logger.error("TypeSafe: %s. detail=%s", error, log_detail) - raise HTTPException(status_code=500, detail={"error": error}) + raise HTTPException(status_code=502, detail={"error": error}) def _candidate_exchanges(self, messages: list[dict[str, object]]) -> list[tuple[int, ...]]: - """Message-index groups eligible for relevance evaluation, oldest first. - - A candidate is a completed tool exchange: an assistant row that made - tool calls plus at least one ``tool``/``function`` row answering it, - with no member protected, and enough combined tool-result text to be - worth an evaluation call. - """ + """Completed tool exchanges eligible for evaluation: unprotected, and long enough to be worth a call.""" protected: Final = _protected_indices(messages) candidates: Final[list[tuple[int, ...]]] = [] for group in group_tool_exchanges(messages): @@ -245,8 +225,7 @@ class TypeSafeGuardrail(CustomGuardrail): return {"task": task, "system": system, "tool_exchanges": tool_exchanges} async def _call_systemone(self, state: dict[str, object], question_ids: list[str]) -> _JevSystemOneResponse | None: - """Evaluate each exchange. Returns the response, or None when the service - failed and fail_open applies.""" + """Returns the response, or None when the service failed and fail_open applies.""" payload: Final[dict[str, object]] = { "model": self.jev_model, "state": state, @@ -267,18 +246,18 @@ class TypeSafeGuardrail(CustomGuardrail): ) except asyncio.CancelledError: raise - except httpx.HTTPStatusError as e: - resp: Final = getattr(e, "response", None) - self._handle_failure( - "TypeSafe evaluation service returned an error", - {"status_code": getattr(resp, "status_code", None), "body": _safe_response_text(resp)}, - ) - return None - except (httpx.RequestError, litellm.Timeout) as e: - self._handle_failure("TypeSafe evaluation service request failed", {"detail": str(e)}) - return None except Exception as e: - self._handle_failure("TypeSafe evaluation service request failed", {"detail": str(e)}) + detail: Final[dict[str, object]] = ( + { + "error_type": type(e).__name__, + "detail": str(e), + "status_code": e.response.status_code, + "body": _safe_response_text(e.response), + } + if isinstance(e, httpx.HTTPStatusError) + else {"error_type": type(e).__name__, "detail": str(e)} + ) + self._handle_failure("TypeSafe evaluation service request failed", detail) return None if not 200 <= raw_response.status_code < 300: self._handle_failure( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py index 293e4f8c344..8b732091d7c 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -227,8 +227,9 @@ async def test_fail_closed_raises_http_exception(): handler.post = AsyncMock(side_effect=Exception("connection refused")) guardrail = _make_guardrail(handler, unreachable_fallback="fail_closed") inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) - with pytest.raises(HTTPException): + with pytest.raises(HTTPException) as exc_info: await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert exc_info.value.status_code == 502 @pytest.mark.asyncio From 27b566590123e494107d1b39bab5a411a8d045b0 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 04:45:28 +0000 Subject: [PATCH 115/253] fix(dashscope): mark qwen3.8-flash as vision and video capable with explicit cache creation cost Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 6 ++++++ model_prices_and_context_window.json | 6 ++++++ 2 files changed, 12 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3409d9b4e46..e903b239881 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -16691,6 +16691,7 @@ "supports_vision": true }, "dashscope/qwen3.8-flash": { + "cache_creation_input_token_cost": 2e-07, "cache_read_input_token_cost": 1.6e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "dashscope", @@ -16705,6 +16706,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, "supports_web_search": true }, "dashscope/qwen3.8-omni-flash": { @@ -18632,6 +18635,7 @@ "supports_vision": true }, "qwen_ai_platform/qwen3.8-flash": { + "cache_creation_input_token_cost": 2e-07, "cache_read_input_token_cost": 1.6e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "qwen_ai_platform", @@ -18646,6 +18650,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, "supports_web_search": true }, "qwen_ai_platform/qwen3.8-omni-flash": { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3409d9b4e46..e903b239881 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16691,6 +16691,7 @@ "supports_vision": true }, "dashscope/qwen3.8-flash": { + "cache_creation_input_token_cost": 2e-07, "cache_read_input_token_cost": 1.6e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "dashscope", @@ -16705,6 +16706,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, "supports_web_search": true }, "dashscope/qwen3.8-omni-flash": { @@ -18632,6 +18635,7 @@ "supports_vision": true }, "qwen_ai_platform/qwen3.8-flash": { + "cache_creation_input_token_cost": 2e-07, "cache_read_input_token_cost": 1.6e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "qwen_ai_platform", @@ -18646,6 +18650,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true, "supports_web_search": true }, "qwen_ai_platform/qwen3.8-omni-flash": { From ca8c9062d80f1ee5f2e34765d0ba3dbc4fd7af72 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 04:48:11 +0000 Subject: [PATCH 116/253] fix(guardrails): satisfy strict lint gates in typesafe guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/typesafe.py | 44 ++++++++++++------- 1 file changed, 28 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index 360219e0ad5..9fca6eaed20 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -11,7 +11,7 @@ from __future__ import annotations import asyncio import time from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Annotated, Final, Literal, TypeGuard, cast +from typing import TYPE_CHECKING, Annotated, Final, Literal import httpx from fastapi import HTTPException @@ -55,12 +55,22 @@ DROPPED_RESULT_TEXT: Final = ( ) -def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip - return isinstance(value, dict) +_STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object]) -def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip - return isinstance(value, list) +def _as_str_object_dict(value: object) -> dict[str, object] | None: + try: + return _STR_OBJECT_DICT_ADAPTER.validate_python(value) + except ValidationError: + return None + + +def _as_object_list(value: object) -> list[object] | None: + try: + return _OBJECT_LIST_ADAPTER.validate_python(value) + except ValidationError: + return None def _safe_response_text(response: HttpxResponse | None, limit: int = 500) -> str: @@ -99,15 +109,16 @@ def _question_instructions(question_id: str) -> str: def _tool_call_entries(assistant_message: Mapping[str, object]) -> list[dict[str, object]]: - tool_calls: Final = assistant_message.get("tool_calls") - if not _is_object_list(tool_calls): + tool_calls: Final = _as_object_list(assistant_message.get("tool_calls")) + if tool_calls is None: return [] entries: Final[list[dict[str, object]]] = [] for tool_call in tool_calls: - if not _is_str_object_dict(tool_call): + parsed_call = _as_str_object_dict(tool_call) + if parsed_call is None: continue - function = tool_call.get("function") - fn = function if _is_str_object_dict(function) else tool_call + function = _as_str_object_dict(parsed_call.get("function")) + fn = function if function is not None else parsed_call entries.append({"name": fn.get("name"), "arguments": fn.get("arguments")}) return entries @@ -137,7 +148,7 @@ class TypeSafeGuardrail(CustomGuardrail): event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, default_on: bool = False, async_handler: AsyncHTTPHandler | None = None, - ): + ) -> None: raw_api_base: Final = (api_base or get_secret_str("TYPESAFE_API_BASE") or DEFAULT_API_BASE).rstrip("/") self.typesafe_api_base = raw_api_base self.typesafe_api_key = api_key or get_secret_str("TYPESAFE_API_KEY") @@ -266,7 +277,7 @@ class TypeSafeGuardrail(CustomGuardrail): ) return None try: - body: Final = cast(object, raw_response.json()) + body: Final[object] = raw_response.json() # pyright: ignore[reportAny] # httpx Response.json() is untyped except (ValueError, httpx.DecodingError, RecursionError): self._handle_failure( "TypeSafe evaluation service returned an unreadable response", @@ -293,12 +304,13 @@ class TypeSafeGuardrail(CustomGuardrail): if input_type != "request": return inputs - structured_messages: Final = inputs.get("structured_messages") - if not _is_object_list(structured_messages) or not structured_messages: + structured_messages: Final = _as_object_list(inputs.get("structured_messages")) + if not structured_messages: return inputs - messages: Final = [m for m in structured_messages if _is_str_object_dict(m)] - if len(messages) != len(structured_messages): + parsed_messages: Final = [_as_str_object_dict(m) for m in structured_messages] + if any(m is None for m in parsed_messages): return inputs + messages: Final = [m for m in parsed_messages if m is not None] candidates: Final = self._candidate_exchanges(messages) if not candidates: From fd4476b1308feba1da8e86cd162401ce1bb2da4c Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 04:49:50 +0000 Subject: [PATCH 117/253] fix(guardrails): bound typesafe tuning params, preserve result tail, log fail-open status Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/typesafe.py | 26 +++++++++++++++++- .../guardrails/guardrail_hooks/typesafe.py | 7 ++++- .../guardrail_hooks/test_typesafe.py | 27 ++++++++++++------- 3 files changed, 49 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index 9fca6eaed20..b3cf5bcf800 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -53,6 +53,7 @@ _JEV_TIMEOUT_SECONDS: Final = 30.0 DROPPED_RESULT_TEXT: Final = ( "[Tool result removed by TypeSafe compaction: judged no longer relevant to the current task]" ) +_ELISION_MARKER: Final = "\n... [middle truncated] ...\n" _STR_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -99,6 +100,17 @@ class _JevSystemOneResponse(BaseModel): _JEV_RESPONSE_ADAPTER: Final = TypeAdapter(_JevSystemOneResponse) +def _truncate_for_state(text: str, max_chars: int) -> str: + """Keeps the head and tail within ``max_chars`` so Jev sees both ends of a long result.""" + if len(text) <= max_chars: + return text + if max_chars <= len(_ELISION_MARKER): + return text[:max_chars] + budget: Final = max_chars - len(_ELISION_MARKER) + head: Final = budget // 2 + return text[:head] + _ELISION_MARKER + text[len(text) - (budget - head) :] + + def _question_instructions(question_id: str) -> str: return ( f"Is tool exchange `{question_id}` in `tool_exchanges` still needed by the assistant to " @@ -231,7 +243,7 @@ class TypeSafeGuardrail(CustomGuardrail): ) tool_exchanges[f"e{ordinal}"] = { "tool_calls": _tool_call_entries(messages[group[0]]), - "result": result_text[: self.max_result_chars_in_state], + "result": _truncate_for_state(result_text, self.max_result_chars_in_state), } return {"task": task, "system": system, "tool_exchanges": tool_exchanges} @@ -324,6 +336,18 @@ class TypeSafeGuardrail(CustomGuardrail): response: Final = await self._call_systemone(state, question_ids) end_time: Final = time.monotonic() if response is None: + self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper + guardrail_json_response={ + "error": "TypeSafe evaluation unavailable; request forwarded uncompacted", + "model": self.jev_model, + }, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + guardrail_provider="typesafe", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) return inputs dropped_ordinals: Final = frozenset( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py b/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py index 4e742bfc7be..59482d2e190 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/typesafe.py @@ -10,6 +10,8 @@ class TypeSafeGuardrailOptionalParams(BaseModel): relevance_threshold: float | None = Field( default=None, + ge=0.0, + le=1.0, description=( "Relevance cutoff in [0, 1]. A completed tool exchange is dropped when Jev " "scores the probability that it is still needed below this value. Defaults to 0.2." @@ -17,14 +19,17 @@ class TypeSafeGuardrailOptionalParams(BaseModel): ) min_chars_to_evaluate: int | None = Field( default=None, + ge=0, description=( "Skip tool exchanges whose combined tool-result text is shorter than this many characters. Defaults to 200." ), ) max_result_chars_in_state: int | None = Field( default=None, + ge=1, description=( - "Tool result text is truncated to this many characters when sent to the Jev evaluator. Defaults to 4000." + "Tool result text is truncated to this many characters when sent to the Jev evaluator, " + "keeping the head and tail. Defaults to 4000." ), ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py index 8b732091d7c..e2e5fc2fefc 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -40,7 +40,7 @@ TOOL_OUTPUT_LONG = "Result: EV range comparison. " * 40 # > 200 chars TOOL_OUTPUT_SHORT = "short" -def _exchange(call_id: str, tool_text: str, name: str = "web_search") -> list[dict]: +def _exchange(call_id: str, tool_text: str, name: str = "web_search") -> list[dict[str, object]]: return [ { "role": "assistant", @@ -57,7 +57,7 @@ def _exchange(call_id: str, tool_text: str, name: str = "web_search") -> list[di ] -def _messages(*, tail: list | None = None) -> list[dict]: +def _messages(*, tail: list[dict[str, object]] | None = None) -> list[dict[str, object]]: base = [ {"role": "system", "content": SYSTEM_TEXT}, {"role": "user", "content": USER_TEXT}, @@ -65,16 +65,21 @@ def _messages(*, tail: list | None = None) -> list[dict]: return base + (tail or []) -def _make_guardrail(handler: MagicMock | None = None, **kwargs) -> TypeSafeGuardrail: - defaults = dict( +def _make_guardrail( + handler: MagicMock | None = None, + *, + max_result_chars_in_state: int | None = None, + unreachable_fallback: str | None = None, +) -> TypeSafeGuardrail: + return TypeSafeGuardrail( api_base=FAKE_API_BASE, api_key=FAKE_API_KEY, guardrail_name="typesafe", default_on=True, async_handler=handler or _make_handler({"e0": 0.9}), + max_result_chars_in_state=max_result_chars_in_state, + unreachable_fallback=unreachable_fallback, ) - defaults.update(kwargs) - return TypeSafeGuardrail(**defaults) def _make_handler(answers: dict[str, float], status: int = 200) -> MagicMock: @@ -91,11 +96,13 @@ def _make_handler(answers: dict[str, float], status: int = 200) -> MagicMock: return handler -def _inputs(messages: list) -> GenericGuardrailAPIInputs: +def _inputs(messages: list[dict[str, object]]) -> GenericGuardrailAPIInputs: return GenericGuardrailAPIInputs(structured_messages=messages) -async def _apply(guardrail: TypeSafeGuardrail, messages: list, input_type: str = "request"): +async def _apply( + guardrail: TypeSafeGuardrail, messages: list[dict[str, object]], input_type: str = "request" +) -> GenericGuardrailAPIInputs: return await guardrail.apply_guardrail( inputs=_inputs(messages), request_data={}, @@ -188,7 +195,9 @@ async def test_request_body_shape_and_truncation(): assert "e0" in payload["questions"]["e0"]["instructions"] assert payload["state"]["task"] == USER_TEXT exchange = payload["state"]["tool_exchanges"]["e0"] - assert exchange["result"] == TOOL_OUTPUT_LONG[:50] + assert len(exchange["result"]) == 50 + assert exchange["result"].startswith(TOOL_OUTPUT_LONG[:10]) + assert exchange["result"].endswith(TOOL_OUTPUT_LONG[-11:]) assert exchange["tool_calls"] == [{"name": "web_search", "arguments": '{"query": "ev"}'}] From 91d4c579e43bd5429741254ea0ff94956edc96fc Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 05:06:48 +0000 Subject: [PATCH 118/253] refactor(guardrails): freeze or suppress mutable constructions in typesafe guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/__init__.py | 15 +- .../guardrail_hooks/typesafe/typesafe.py | 148 ++++++++++-------- .../guardrail_hooks/test_typesafe.py | 2 +- 3 files changed, 91 insertions(+), 74 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py index 40f040eb676..c1b05597a6c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -1,5 +1,6 @@ from __future__ import annotations +from types import MappingProxyType from typing import TYPE_CHECKING, Final from pydantic import BaseModel @@ -25,7 +26,9 @@ def _coerce_event_hook( if isinstance(mode, Mode): return mode if isinstance(mode, list): - return [GuardrailEventHooks(item) for item in mode] + return [ + GuardrailEventHooks(item) for item in mode + ] # mutable-ok: CustomGuardrail event_hook contract wants a list return GuardrailEventHooks(mode) @@ -63,10 +66,8 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> return _callback -guardrail_initializer_registry: Final = { - SupportedGuardrailIntegrations.TYPESAFE.value: initialize_guardrail, -} +guardrail_initializer_registry: Final = MappingProxyType( + {SupportedGuardrailIntegrations.TYPESAFE.value: initialize_guardrail} +) -guardrail_class_registry: Final = { - SupportedGuardrailIntegrations.TYPESAFE.value: TypeSafeGuardrail, -} +guardrail_class_registry: Final = MappingProxyType({SupportedGuardrailIntegrations.TYPESAFE.value: TypeSafeGuardrail}) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index b3cf5bcf800..384bfd2c54b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -11,6 +11,7 @@ from __future__ import annotations import asyncio import time from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Literal import httpx @@ -120,19 +121,20 @@ def _question_instructions(question_id: str) -> str: ) -def _tool_call_entries(assistant_message: Mapping[str, object]) -> list[dict[str, object]]: +def _tool_call_entry(tool_call: object) -> dict[str, object] | None: + parsed_call = _as_str_object_dict(tool_call) + if parsed_call is None: + return None + function = _as_str_object_dict(parsed_call.get("function")) + fn = function if function is not None else parsed_call + return {"name": fn.get("name"), "arguments": fn.get("arguments")} # mutable-ok: serialized to JSON + + +def _tool_call_entries(assistant_message: Mapping[str, object]) -> tuple[dict[str, object], ...]: tool_calls: Final = _as_object_list(assistant_message.get("tool_calls")) if tool_calls is None: - return [] - entries: Final[list[dict[str, object]]] = [] - for tool_call in tool_calls: - parsed_call = _as_str_object_dict(tool_call) - if parsed_call is None: - continue - function = _as_str_object_dict(parsed_call.get("function")) - fn = function if function is not None else parsed_call - entries.append({"name": fn.get("name"), "arguments": fn.get("arguments")}) - return entries + return () + return tuple(entry for tool_call in tool_calls if (entry := _tool_call_entry(tool_call)) is not None) def _protected_indices(messages: Sequence[Mapping[str, object]]) -> frozenset[int]: @@ -199,30 +201,32 @@ class TypeSafeGuardrail(CustomGuardrail): ) return verbose_proxy_logger.error("TypeSafe: %s. detail=%s", error, log_detail) - raise HTTPException(status_code=502, detail={"error": error}) + raise HTTPException(status_code=502, detail={"error": error}) # mutable-ok: FastAPI wants a dict detail - def _candidate_exchanges(self, messages: list[dict[str, object]]) -> list[tuple[int, ...]]: + def _candidate_exchanges(self, messages: Sequence[dict[str, object]]) -> tuple[tuple[int, ...], ...]: """Completed tool exchanges eligible for evaluation: unprotected, and long enough to be worth a call.""" protected: Final = _protected_indices(messages) - candidates: Final[list[tuple[int, ...]]] = [] - for group in group_tool_exchanges(messages): - if len(group) < 2: - continue - if messages[group[0]].get("role") != "assistant": - continue - if any(member in protected for member in group): - continue - tool_text = "".join( - content_to_text(messages[index].get("content")) - for index in group[1:] - if messages[index].get("role") in ("tool", "function") - ) - if not tool_text or len(tool_text) < self.min_chars_to_evaluate: - continue - candidates.append(group) + candidates: Final = tuple( + group + for group in group_tool_exchanges(messages) + if len(group) >= 2 + and messages[group[0]].get("role") == "assistant" + and not any(member in protected for member in group) + and len(self._exchange_tool_text(messages, group)) >= self.min_chars_to_evaluate + ) return candidates[-_MAX_EXCHANGES_EVALUATED:] - def _build_state(self, messages: list[dict[str, object]], candidates: list[tuple[int, ...]]) -> dict[str, object]: + @staticmethod + def _exchange_tool_text(messages: Sequence[dict[str, object]], group: tuple[int, ...]) -> str: + return "".join( + content_to_text(messages[index].get("content")) + for index in group[1:] + if messages[index].get("role") in ("tool", "function") + ) + + def _build_state( + self, messages: Sequence[dict[str, object]], candidates: tuple[tuple[int, ...], ...] + ) -> dict[str, object]: task: Final = next( ( content_to_text(messages[index].get("content")) @@ -234,26 +238,29 @@ class TypeSafeGuardrail(CustomGuardrail): system: Final = "\n\n".join( content_to_text(message.get("content")) for message in messages if message.get("role") == "system" ) - tool_exchanges: Final[dict[str, object]] = {} - for ordinal, group in enumerate(candidates): - result_text = "".join( - content_to_text(messages[index].get("content")) - for index in group[1:] - if messages[index].get("role") in ("tool", "function") - ) - tool_exchanges[f"e{ordinal}"] = { + tool_exchanges: Final = { # mutable-ok: accumulated once, serialized to JSON + f"e{ordinal}": { # mutable-ok: serialized to JSON "tool_calls": _tool_call_entries(messages[group[0]]), - "result": _truncate_for_state(result_text, self.max_result_chars_in_state), + "result": _truncate_for_state( + self._exchange_tool_text(messages, group), self.max_result_chars_in_state + ), } - return {"task": task, "system": system, "tool_exchanges": tool_exchanges} + for ordinal, group in enumerate(candidates) + } + return {"task": task, "system": system, "tool_exchanges": tool_exchanges} # mutable-ok: serialized to JSON - async def _call_systemone(self, state: dict[str, object], question_ids: list[str]) -> _JevSystemOneResponse | None: + async def _call_systemone( + self, state: dict[str, object], question_ids: Sequence[str] + ) -> _JevSystemOneResponse | None: """Returns the response, or None when the service failed and fail_open applies.""" - payload: Final[dict[str, object]] = { + payload: Final[dict[str, object]] = { # mutable-ok: serialized to JSON by httpx "model": self.jev_model, "state": state, - "questions": { - question_id: {"type": "noul", "instructions": _question_instructions(question_id)} + "questions": { # mutable-ok: serialized to JSON + question_id: { + "type": "noul", + "instructions": _question_instructions(question_id), + } # mutable-ok: serialized to JSON for question_id in question_ids }, } @@ -261,7 +268,7 @@ class TypeSafeGuardrail(CustomGuardrail): raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped url=f"{self.typesafe_api_base}/v1/systemone", json=payload, - headers={ + headers={ # mutable-ok: httpx header contract is a dict "Authorization": f"Bearer {self.typesafe_api_key}", "Content-Type": "application/json", }, @@ -271,21 +278,24 @@ class TypeSafeGuardrail(CustomGuardrail): raise except Exception as e: detail: Final[dict[str, object]] = ( - { + { # mutable-ok: log detail record "error_type": type(e).__name__, "detail": str(e), "status_code": e.response.status_code, "body": _safe_response_text(e.response), } if isinstance(e, httpx.HTTPStatusError) - else {"error_type": type(e).__name__, "detail": str(e)} + else {"error_type": type(e).__name__, "detail": str(e)} # mutable-ok: log detail record ) self._handle_failure("TypeSafe evaluation service request failed", detail) return None if not 200 <= raw_response.status_code < 300: self._handle_failure( "TypeSafe evaluation service returned an error", - {"status_code": raw_response.status_code, "body": _safe_response_text(raw_response)}, + { + "status_code": raw_response.status_code, + "body": _safe_response_text(raw_response), + }, # mutable-ok: log detail record ) return None try: @@ -293,7 +303,7 @@ class TypeSafeGuardrail(CustomGuardrail): except (ValueError, httpx.DecodingError, RecursionError): self._handle_failure( "TypeSafe evaluation service returned an unreadable response", - {"body": _safe_response_text(raw_response)}, + {"body": _safe_response_text(raw_response)}, # mutable-ok: log detail record ) return None try: @@ -301,7 +311,7 @@ class TypeSafeGuardrail(CustomGuardrail): except ValidationError: self._handle_failure( "TypeSafe evaluation service returned unexpected response shape", - {"body": _safe_response_text(raw_response)}, + {"body": _safe_response_text(raw_response)}, # mutable-ok: log detail record ) return None @@ -319,17 +329,17 @@ class TypeSafeGuardrail(CustomGuardrail): structured_messages: Final = _as_object_list(inputs.get("structured_messages")) if not structured_messages: return inputs - parsed_messages: Final = [_as_str_object_dict(m) for m in structured_messages] + parsed_messages: Final = tuple(_as_str_object_dict(m) for m in structured_messages) if any(m is None for m in parsed_messages): return inputs - messages: Final = [m for m in parsed_messages if m is not None] + messages: Final = tuple(m for m in parsed_messages if m is not None) candidates: Final = self._candidate_exchanges(messages) if not candidates: verbose_proxy_logger.debug("TypeSafe: no completed tool exchanges eligible for evaluation") return inputs - question_ids: Final = [f"e{ordinal}" for ordinal in range(len(candidates))] + question_ids: Final = tuple(f"e{ordinal}" for ordinal in range(len(candidates))) state: Final = self._build_state(messages, candidates) start_time: Final = time.monotonic() @@ -337,10 +347,12 @@ class TypeSafeGuardrail(CustomGuardrail): end_time: Final = time.monotonic() if response is None: self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper - guardrail_json_response={ - "error": "TypeSafe evaluation unavailable; request forwarded uncompacted", - "model": self.jev_model, - }, + guardrail_json_response=MappingProxyType( + { + "error": "TypeSafe evaluation unavailable; request forwarded uncompacted", + "model": self.jev_model, + } + ), request_data=request_data, guardrail_status="guardrail_failed_to_respond", guardrail_provider="typesafe", @@ -365,8 +377,10 @@ class TypeSafeGuardrail(CustomGuardrail): verbose_proxy_logger.debug("TypeSafe: all evaluated exchanges still relevant; request unchanged") return inputs - compacted_messages: Final = [ - {**message, "content": DROPPED_RESULT_TEXT} if index in dropped_tool_indices else message + compacted_messages: Final = [ # mutable-ok: structured_messages contract is a list of dicts + {**message, "content": DROPPED_RESULT_TEXT} + if index in dropped_tool_indices + else message # mutable-ok: JSON message row for index, message in enumerate(messages) ] chars_removed: Final = sum( @@ -381,12 +395,14 @@ class TypeSafeGuardrail(CustomGuardrail): chars_removed, ) self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper - guardrail_json_response={ - "exchanges_evaluated": len(candidates), - "exchanges_dropped": exchanges_dropped, - "chars_removed": chars_removed, - "model": self.jev_model, - }, + guardrail_json_response=MappingProxyType( + { + "exchanges_evaluated": len(candidates), + "exchanges_dropped": exchanges_dropped, + "chars_removed": chars_removed, + "model": self.jev_model, + } + ), request_data=request_data, guardrail_status="success", guardrail_provider="typesafe", @@ -394,7 +410,7 @@ class TypeSafeGuardrail(CustomGuardrail): end_time=end_time, duration=end_time - start_time, ) - return {**inputs, "structured_messages": compacted_messages} # pyright: ignore[reportReturnType] # plain dicts satisfy AllMessageValues at runtime + return {**inputs, "structured_messages": compacted_messages} # pyright: ignore[reportReturnType] # mutable-ok: inputs protocol is a plain dict # plain dicts satisfy AllMessageValues at runtime @staticmethod def get_config_model() -> type[TypeSafeGuardrailConfigModel] | None: diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py index e2e5fc2fefc..92840c489c3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -198,7 +198,7 @@ async def test_request_body_shape_and_truncation(): assert len(exchange["result"]) == 50 assert exchange["result"].startswith(TOOL_OUTPUT_LONG[:10]) assert exchange["result"].endswith(TOOL_OUTPUT_LONG[-11:]) - assert exchange["tool_calls"] == [{"name": "web_search", "arguments": '{"query": "ev"}'}] + assert list(exchange["tool_calls"]) == [{"name": "web_search", "arguments": '{"query": "ev"}'}] @pytest.mark.asyncio From 351afc85196a26b6180e82183c2794b5be3b8958 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 05:10:48 +0000 Subject: [PATCH 119/253] test(guardrails): cover typesafe failure paths and edge shapes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/test_typesafe.py | 116 +++++++++++++++++- 1 file changed, 115 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py index 92840c489c3..a936d725bd3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -16,7 +16,7 @@ Tests cover: - response input_type passthrough and initialize_guardrail wiring """ -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, PropertyMock import pytest from fastapi import HTTPException @@ -282,3 +282,117 @@ def test_initialize_guardrail_applies_optional_params_and_registry_keys(): assert callback.unreachable_fallback == "fail_open" assert guardrail_initializer_registry[SupportedGuardrailIntegrations.TYPESAFE.value] is initialize_guardrail assert guardrail_class_registry[SupportedGuardrailIntegrations.TYPESAFE.value] is TypeSafeGuardrail + + +def test_missing_api_key_raises(monkeypatch): + monkeypatch.delenv("TYPESAFE_API_KEY", raising=False) + with pytest.raises(ValueError, match="requires an API key"): + TypeSafeGuardrail(api_key=None) + + +def test_get_config_model_and_ui_name(): + from litellm.types.proxy.guardrails.guardrail_hooks.typesafe import ( + TypeSafeGuardrailConfigModel, + ) + + assert TypeSafeGuardrail.get_config_model() is TypeSafeGuardrailConfigModel + assert TypeSafeGuardrailConfigModel.ui_friendly_name() == "TypeSafe (Jev) Compaction" + + +@pytest.mark.asyncio +async def test_non_list_and_non_dict_messages_return_identity(): + guardrail = _make_guardrail() + not_a_list = GenericGuardrailAPIInputs(structured_messages={"role": "user"}) + assert await guardrail.apply_guardrail(inputs=not_a_list, request_data={}, input_type="request", logging_obj=None) is not_a_list + with_bad_row = _inputs(_messages(tail=[["not", "a", "dict"]])) + assert await guardrail.apply_guardrail(inputs=with_bad_row, request_data={}, input_type="request", logging_obj=None) is with_bad_row + + +def test_odd_tool_call_shapes_yield_no_entries(): + from litellm.proxy.guardrails.guardrail_hooks.typesafe.typesafe import _tool_call_entries + + assert _tool_call_entries({"tool_calls": "not-a-list"}) == () + assert _tool_call_entries({"tool_calls": None}) == () + assert list(_tool_call_entries({"tool_calls": [42]})) == [] + entries = _tool_call_entries({"tool_calls": [{"function": {"name": "web_search", "arguments": "{}"}}]}) + assert list(entries) == [{"name": "web_search", "arguments": "{}"}] + + +@pytest.mark.asyncio +async def test_short_max_chars_uses_prefix_slice(): + handler = _make_handler({"e0": 0.9}) + guardrail = _make_guardrail(handler, max_result_chars_in_state=5) + await _apply(guardrail, _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = handler.post.call_args.kwargs["json"]["state"]["tool_exchanges"]["e0"]["result"] + assert result == TOOL_OUTPUT_LONG[:5] + + +@pytest.mark.asyncio +async def test_unreadable_json_body_fails_open(): + handler = MagicMock() + response = MagicMock() + response.status_code = 200 + response.text = "not json" + response.json.side_effect = ValueError("no json") + handler.post = AsyncMock(return_value=response) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_malformed_answers_shape_fails_open(): + handler = MagicMock() + response = MagicMock() + response.status_code = 200 + response.text = '{"answers": "oops"}' + response.json.return_value = {"answers": "oops"} + handler.post = AsyncMock(return_value=response) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_http_status_error_includes_status_and_undecodable_body(): + import httpx + + response = MagicMock() + response.status_code = 503 + type(response).text = PropertyMock(side_effect=httpx.DecodingError("bad codec")) + handler = MagicMock() + handler.post = AsyncMock( + side_effect=httpx.HTTPStatusError("unavailable", request=MagicMock(), response=response) + ) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + assert result is inputs + + +@pytest.mark.asyncio +async def test_cancelled_jev_call_propagates(): + import asyncio + + handler = MagicMock() + handler.post = AsyncMock(side_effect=asyncio.CancelledError()) + guardrail = _make_guardrail(handler) + inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + with pytest.raises(asyncio.CancelledError): + await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + + +def test_optional_params_defaults_and_event_hook_coercion(): + from litellm.proxy.guardrails.guardrail_hooks.typesafe import _coerce_event_hook, _optional_params + from litellm.types.guardrails import GuardrailEventHooks, LitellmParams + + assert _coerce_event_hook("pre_call") is GuardrailEventHooks.pre_call + assert _coerce_event_hook(["pre_call", "post_call"]) == [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ] + litellm_params = LitellmParams(guardrail="typesafe", mode="pre_call", api_key=FAKE_API_KEY) + params = _optional_params(litellm_params) + assert params.relevance_threshold is None From 87e1c6b3ba19b184735601323eef6d3002f5a0c2 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 05:25:20 +0000 Subject: [PATCH 120/253] fix(guardrails): pin mutable-ok suppressions to constructed literals Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrails/guardrail_hooks/typesafe/__init__.py | 4 ++-- .../guardrails/guardrail_hooks/typesafe/typesafe.py | 12 ++++++------ 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py index c1b05597a6c..abfc60c669a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -26,9 +26,9 @@ def _coerce_event_hook( if isinstance(mode, Mode): return mode if isinstance(mode, list): - return [ + return [ # mutable-ok: CustomGuardrail event_hook contract wants a list GuardrailEventHooks(item) for item in mode - ] # mutable-ok: CustomGuardrail event_hook contract wants a list + ] return GuardrailEventHooks(mode) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index 384bfd2c54b..cb4f21b3261 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -257,10 +257,10 @@ class TypeSafeGuardrail(CustomGuardrail): "model": self.jev_model, "state": state, "questions": { # mutable-ok: serialized to JSON - question_id: { + question_id: { # mutable-ok: serialized to JSON "type": "noul", "instructions": _question_instructions(question_id), - } # mutable-ok: serialized to JSON + } for question_id in question_ids }, } @@ -292,10 +292,10 @@ class TypeSafeGuardrail(CustomGuardrail): if not 200 <= raw_response.status_code < 300: self._handle_failure( "TypeSafe evaluation service returned an error", - { + { # mutable-ok: log detail record "status_code": raw_response.status_code, "body": _safe_response_text(raw_response), - }, # mutable-ok: log detail record + }, ) return None try: @@ -378,9 +378,9 @@ class TypeSafeGuardrail(CustomGuardrail): return inputs compacted_messages: Final = [ # mutable-ok: structured_messages contract is a list of dicts - {**message, "content": DROPPED_RESULT_TEXT} + {**message, "content": DROPPED_RESULT_TEXT} # mutable-ok: JSON message row if index in dropped_tool_indices - else message # mutable-ok: JSON message row + else message for index, message in enumerate(messages) ] chars_removed: Final = sum( From 22e6947fef68c2e3af137cdca78b3d8811fbb8a4 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 05:44:25 +0000 Subject: [PATCH 121/253] fix(guardrails): keep typesafe guardrail log payloads JSON-serializable Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/typesafe.py | 25 ++++++++----------- 1 file changed, 10 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py index cb4f21b3261..9df5c204a77 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/typesafe.py @@ -11,7 +11,6 @@ from __future__ import annotations import asyncio import time from collections.abc import Mapping, Sequence -from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Literal import httpx @@ -347,12 +346,10 @@ class TypeSafeGuardrail(CustomGuardrail): end_time: Final = time.monotonic() if response is None: self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper - guardrail_json_response=MappingProxyType( - { - "error": "TypeSafe evaluation unavailable; request forwarded uncompacted", - "model": self.jev_model, - } - ), + guardrail_json_response={ # mutable-ok: must stay JSON-serializable for shared logging + "error": "TypeSafe evaluation unavailable; request forwarded uncompacted", + "model": self.jev_model, + }, request_data=request_data, guardrail_status="guardrail_failed_to_respond", guardrail_provider="typesafe", @@ -395,14 +392,12 @@ class TypeSafeGuardrail(CustomGuardrail): chars_removed, ) self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] # untyped base helper - guardrail_json_response=MappingProxyType( - { - "exchanges_evaluated": len(candidates), - "exchanges_dropped": exchanges_dropped, - "chars_removed": chars_removed, - "model": self.jev_model, - } - ), + guardrail_json_response={ # mutable-ok: must stay JSON-serializable for shared logging + "exchanges_evaluated": len(candidates), + "exchanges_dropped": exchanges_dropped, + "chars_removed": chars_removed, + "model": self.jev_model, + }, request_data=request_data, guardrail_status="success", guardrail_provider="typesafe", From 0765f6d571b3a27de6bc9a7456da977bf62705e4 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 06:05:43 +0000 Subject: [PATCH 122/253] fix(guardrails): keep typesafe registries as dicts so guardrail discovery finds them Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/typesafe/__init__.py | 11 +++---- .../guardrail_hooks/test_typesafe.py | 29 +++++++++++++------ 2 files changed, 26 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py index abfc60c669a..dcea75d3a98 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/typesafe/__init__.py @@ -1,6 +1,5 @@ from __future__ import annotations -from types import MappingProxyType from typing import TYPE_CHECKING, Final from pydantic import BaseModel @@ -66,8 +65,10 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) -> return _callback -guardrail_initializer_registry: Final = MappingProxyType( - {SupportedGuardrailIntegrations.TYPESAFE.value: initialize_guardrail} -) +guardrail_initializer_registry: Final = { # mutable-ok: guardrail_registry discovery checks isinstance(registry, dict) + SupportedGuardrailIntegrations.TYPESAFE.value: initialize_guardrail, +} -guardrail_class_registry: Final = MappingProxyType({SupportedGuardrailIntegrations.TYPESAFE.value: TypeSafeGuardrail}) +guardrail_class_registry: Final = { # mutable-ok: guardrail_registry discovery checks isinstance(registry, dict) + SupportedGuardrailIntegrations.TYPESAFE.value: TypeSafeGuardrail, +} diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py index a936d725bd3..2d1db07a1a0 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_typesafe.py @@ -36,7 +36,7 @@ FAKE_API_KEY = "ts_test-key" SYSTEM_TEXT = "You are a research assistant." USER_TEXT = "Which 2026 EV has the longest range?" -TOOL_OUTPUT_LONG = "Result: EV range comparison. " * 40 # > 200 chars +TOOL_OUTPUT_LONG = "Result: EV range comparison. " * 40 TOOL_OUTPUT_SHORT = "short" @@ -141,8 +141,6 @@ async def test_low_noul_exchange_blanked_high_kept_and_input_not_mutated(): async def test_last_exchange_and_protected_rows_never_evaluated(): handler = _make_handler({"e0": 0.05}) guardrail = _make_guardrail(handler) - # Ends on a tool result: the last assistant row is protected, so the whole - # last exchange is out of scope even though its text is long. messages = _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), *_exchange("call_2", TOOL_OUTPUT_LONG)]) result = await _apply(guardrail, messages) @@ -303,9 +301,15 @@ def test_get_config_model_and_ui_name(): async def test_non_list_and_non_dict_messages_return_identity(): guardrail = _make_guardrail() not_a_list = GenericGuardrailAPIInputs(structured_messages={"role": "user"}) - assert await guardrail.apply_guardrail(inputs=not_a_list, request_data={}, input_type="request", logging_obj=None) is not_a_list + assert ( + await guardrail.apply_guardrail(inputs=not_a_list, request_data={}, input_type="request", logging_obj=None) + is not_a_list + ) with_bad_row = _inputs(_messages(tail=[["not", "a", "dict"]])) - assert await guardrail.apply_guardrail(inputs=with_bad_row, request_data={}, input_type="request", logging_obj=None) is with_bad_row + assert ( + await guardrail.apply_guardrail(inputs=with_bad_row, request_data={}, input_type="request", logging_obj=None) + is with_bad_row + ) def test_odd_tool_call_shapes_yield_no_entries(): @@ -322,7 +326,9 @@ def test_odd_tool_call_shapes_yield_no_entries(): async def test_short_max_chars_uses_prefix_slice(): handler = _make_handler({"e0": 0.9}) guardrail = _make_guardrail(handler, max_result_chars_in_state=5) - await _apply(guardrail, _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) + await _apply( + guardrail, _messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}]) + ) result = handler.post.call_args.kwargs["json"]["state"]["tool_exchanges"]["e0"]["result"] assert result == TOOL_OUTPUT_LONG[:5] @@ -363,9 +369,7 @@ async def test_http_status_error_includes_status_and_undecodable_body(): response.status_code = 503 type(response).text = PropertyMock(side_effect=httpx.DecodingError("bad codec")) handler = MagicMock() - handler.post = AsyncMock( - side_effect=httpx.HTTPStatusError("unavailable", request=MagicMock(), response=response) - ) + handler.post = AsyncMock(side_effect=httpx.HTTPStatusError("unavailable", request=MagicMock(), response=response)) guardrail = _make_guardrail(handler) inputs = _inputs(_messages(tail=[*_exchange("call_1", TOOL_OUTPUT_LONG), {"role": "assistant", "content": "x"}])) result = await guardrail.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) @@ -396,3 +400,10 @@ def test_optional_params_defaults_and_event_hook_coercion(): litellm_params = LitellmParams(guardrail="typesafe", mode="pre_call", api_key=FAKE_API_KEY) params = _optional_params(litellm_params) assert params.relevance_threshold is None + + +def test_typesafe_initializer_discoverable_via_hook_registries(): + from litellm.proxy.guardrails.guardrail_registry import get_guardrail_initializer_from_hooks + + initializers = get_guardrail_initializer_from_hooks() + assert initializers["typesafe"] is initialize_guardrail From a187a5bfc6e64de75bd59aca3103bf2059deeb45 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 07:54:01 +0000 Subject: [PATCH 123/253] fix(proxy): keep requested model guardrails and key disable_fallbacks on rate-limit fallback The local rate-limit fallback path introduced in #40596 restores the request from a snapshot taken before the pre-call pass, so the guardrails resolved for the requested model were dropped when another deployment was selected, and the raw request-body disable_fallbacks field gated the fallback before the key-level disable_fallbacks override had been applied Carry the requested model's merged guardrail list onto every fallback pass and read the effective disable_fallbacks value after the pre-call pass Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 23 +++++- .../proxy/test_common_request_processing.py | 70 ++++++++++++++++++- 2 files changed, 87 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f650b6d0b28..607ce3441e9 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -94,7 +94,11 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression from litellm.proxy.route_llm_request import route_request -from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails +from litellm.proxy.utils import ( + ProxyLogging, + _check_and_merge_model_level_guardrails, + _merge_guardrails_with_existing, +) from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -1934,6 +1938,7 @@ class ProxyBaseLLMRequestProcessing: user_api_base: str | None = None, model: str | None = None, llm_router: Router | None = None, + requested_model_guardrails: list | None = None, ) -> tuple[dict, LiteLLMLoggingObj]: start_time: Final = datetime.now() # start before calling guardrail hooks @@ -2098,7 +2103,11 @@ class ProxyBaseLLMRequestProcessing: # could otherwise spoof an unguarded model_info.id while requesting # a guarded alias and bypass guardrails (veria-ai HIGH on #29654). self.data = _check_and_merge_model_level_guardrails( - data=self.data, + data=( + self.data + if requested_model_guardrails is None + else _merge_guardrails_with_existing(data=self.data, model_level_guardrails=requested_model_guardrails) + ), llm_router=llm_router, trust_client_model_info=False, ) @@ -2163,7 +2172,7 @@ class ProxyBaseLLMRequestProcessing: configured_fallbacks: Final = ( self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict) - if llm_router is not None and not self.data.get("disable_fallbacks") + if llm_router is not None else None ) pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None @@ -2208,6 +2217,7 @@ class ProxyBaseLLMRequestProcessing: original_model, fallback_models, ) + requested_model_guardrails: Final = self._request_guardrails(rate_limited_data) try: for fallback_model in fallback_models: @@ -2231,6 +2241,7 @@ class ProxyBaseLLMRequestProcessing: model=fallback_model, route_type=route_type, llm_router=llm_router, + requested_model_guardrails=requested_model_guardrails, ) except ProxyRateLimitError: continue @@ -2241,6 +2252,12 @@ class ProxyBaseLLMRequestProcessing: self.data = rate_limited_data raise original_exc + @staticmethod + def _request_guardrails(data: dict) -> list | None: + metadata: Final = data.get("metadata") + guardrails: Final = metadata.get("guardrails") if isinstance(metadata, dict) else None + return guardrails if isinstance(guardrails, list) else None + @staticmethod def _configured_fallbacks(llm_router: Router, user_api_key_dict: UserAPIKeyAuth) -> list | None: key_router_settings: Final = user_api_key_dict.router_settings diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index d465deace15..f440711ac29 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -6711,6 +6711,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: monkeypatch: pytest.MonkeyPatch, user_api_key_dict: ProxyUserAPIKeyAuth, fallbacks: list[dict[str, list[str]]], + model_guardrails: dict[str, list[str]] | None = None, ) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]: """Real v3 limiter (the default ``parallel_request_limiter``) wired in through the ``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real: @@ -6738,9 +6739,17 @@ class TestPreCallWithFallbacksOnLocalRateLimit: proxy_logging_obj = MagicMock(spec=ProxyLogging) proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter) + guardrails_by_group = model_guardrails or {} router = litellm.Router( model_list=[ - {"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}} + { + "model_name": group, + "litellm_params": { + "model": "openai/gpt-4.1-nano", + "api_key": "fake", + **({"guardrails": guardrails_by_group[group]} if group in guardrails_by_group else {}), + }, + } for chain in fallbacks for group in (*chain.keys(), *(m for models in chain.values() for m in models)) ], @@ -6752,7 +6761,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: def _otel_key( rpm_limit: int | None = None, model_rpm_limit: dict[str, int] | None = None, - disable_fallbacks: bool = False, + disable_fallbacks: bool | None = None, ) -> ProxyUserAPIKeyAuth: from opentelemetry.sdk.trace import TracerProvider @@ -6763,7 +6772,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit: rpm_limit=rpm_limit, metadata={ **({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}), - **({"disable_fallbacks": True} if disable_fallbacks else {}), + **({"disable_fallbacks": disable_fallbacks} if disable_fallbacks is not None else {}), }, ) @@ -6906,6 +6915,61 @@ class TestPreCallWithFallbacksOnLocalRateLimit: assert exc_info.value.status_code == 429 assert rig[3] == [primary_model, primary_model] + @pytest.mark.asyncio + async def test_key_metadata_disable_fallbacks_false_overrides_request_body(self, monkeypatch: pytest.MonkeyPatch): + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + key = self._otel_key(model_rpm_limit={primary_model: 1}, disable_fallbacks=False) + rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}]) + request = { + "model": primary_model, + "messages": [{"role": "user", "content": "hi"}], + "disable_fallbacks": True, + } + + await self._pre_call(dict(request), key, rig) + _, (data, _) = await self._pre_call(dict(request), key, rig) + + assert data["model"] == fallback_model + assert data["disable_fallbacks"] is False + assert rig[3] == [primary_model, primary_model, fallback_model] + + @pytest.mark.asyncio + async def test_fallback_keeps_requested_model_guardrails(self, monkeypatch: pytest.MonkeyPatch): + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + guardrail = "pii-guard-for-primary" + key = self._otel_key(model_rpm_limit={primary_model: 1}) + rig = self._v3_limiter_rig( + monkeypatch, key, [{primary_model: [fallback_model]}], model_guardrails={primary_model: [guardrail]} + ) + run_limiter = rig[0].pre_call_hook + + async def limiter_then_guardrail( + user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str + ) -> dict[str, object]: + limited = await run_limiter(user_api_key_dict=user_api_key_dict, data=data, call_type=call_type) + if guardrail not in (limited["metadata"].get("guardrails") or []): + return limited + return { + **limited, + "messages": [ + {**m, "content": str(m["content"]).replace("123-45-6789", "[REDACTED-SSN]")} + for m in limited["messages"] + ], + } + + rig[0].pre_call_hook = AsyncMock(side_effect=limiter_then_guardrail) + request = {"model": primary_model, "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}]} + + await self._pre_call(dict(request), key, rig) + _, (data, _) = await self._pre_call(dict(request), key, rig) + + assert data["model"] == fallback_model + assert guardrail in data["metadata"]["guardrails"] + assert data["messages"] == [{"role": "user", "content": "my ssn is [REDACTED-SSN]"}] + assert rig[3] == [primary_model, primary_model, fallback_model] + class _RecordingSuccessLogger(CustomLogger): def __init__(self): From 318b027782b0b6ab1aec7103a619d308661894e6 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 08:26:39 +0000 Subject: [PATCH 124/253] fix(otel): keep caller traceparent and tracestate on pass-through relays Pass-through requests inject the proxy span into upstream headers since #40669, which replaced an explicit x-pass-traceparent with an unrelated trace and dropped its x-pass-tracestate. Keep the caller's context when the carrier already names a different trace, and keep the proxy child span for same-trace or missing headers. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/otel/plumbing/context.py | 19 +++++- .../otel/test_otel_v2_components.py | 63 ++++++++++++++----- .../test_pass_through_endpoints.py | 59 +++++++++++++++++ 3 files changed, 125 insertions(+), 16 deletions(-) diff --git a/litellm/integrations/otel/plumbing/context.py b/litellm/integrations/otel/plumbing/context.py index 1282e654365..19243d64c64 100644 --- a/litellm/integrations/otel/plumbing/context.py +++ b/litellm/integrations/otel/plumbing/context.py @@ -325,12 +325,27 @@ def _outgoing_trace_context(parent_span: object) -> Context | None: return None +def _propagated_context(headers: Mapping[str, str], request_context: Context) -> Context: + """``request_context`` when it continues the trace ``headers`` already name, else the + caller's own context, so an explicit upstream ``traceparent`` (``x-pass-traceparent``) + is never swapped for an unrelated trace and its ``tracestate`` survives.""" + caller: Final = extract_traceparent(headers) + if caller is None: + return request_context + caller_span: Final = get_current_span(caller).get_span_context() + request_span: Final = get_current_span(request_context).get_span_context() + if not caller_span.is_valid or caller_span.trace_id == request_span.trace_id: + return request_context + return caller + + def inject_trace_context(headers: Mapping[str, str], parent_span: object = None) -> dict[str, str]: """``headers`` plus W3C ``traceparent``/``tracestate`` for this request's span. Parent preference: ``parent_span`` (the request span auth stashed on the key), then the anchored request root span, then the ambient active span. Only trace context is - injected, never Baggage. Unchanged when no valid span exists anywhere. + injected, never Baggage. Unchanged when no valid span exists anywhere. A ``traceparent`` + already in ``headers`` from a different trace is forwarded as-is instead of replaced. """ context: Final = _outgoing_trace_context(parent_span) if context is None: @@ -338,7 +353,7 @@ def inject_trace_context(headers: Mapping[str, str], parent_span: object = None) carrier: Final = { # mutable-ok: OpenTelemetry propagator requires a mutable carrier key: value for key, value in headers.items() if key.lower() not in _W3C_TRACE_HEADERS } - _PROPAGATOR.inject(carrier, context=context) + _PROPAGATOR.inject(carrier, context=_propagated_context(headers, context)) return carrier diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/test_litellm/integrations/otel/test_otel_v2_components.py index 72f6213e880..0c95049ce05 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -477,19 +477,21 @@ def _test_tracer(): return provider.get_tracer("test") +_CALLER_TRACEPARENT = "00-11111111111111111111111111111111-2222222222222222-01" + + def test_inject_trace_context_prefers_request_root_span(): def run(): tracer = _test_tracer() - with tracer.start_as_current_span("root") as root: + inbound = TraceContextTextMapPropagator().extract({"traceparent": _CALLER_TRACEPARENT}) + with tracer.start_as_current_span("root", context=inbound) as root: ctx_mod.set_request_root_span(root) - result = ctx_mod.inject_trace_context( - {"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"} - ) + result = ctx_mod.inject_trace_context({"traceparent": _CALLER_TRACEPARENT}) propagated = get_current_span(TraceContextTextMapPropagator().extract(result)) return result, root, propagated result, root, propagated = ContextVarContext().run(run) - assert result["traceparent"] != "00-11111111111111111111111111111111-2222222222222222-01" + assert result["traceparent"] != _CALLER_TRACEPARENT assert propagated.get_span_context().trace_id == root.get_span_context().trace_id assert propagated.get_span_context().span_id == root.get_span_context().span_id @@ -507,24 +509,57 @@ def test_inject_trace_context_uses_ambient_span_without_request_root(): assert propagated.get_span_context().span_id == ambient.get_span_context().span_id -def test_inject_trace_context_replaces_stale_trace_headers(): +def test_inject_trace_context_replaces_same_trace_headers_with_request_span(): def run(): tracer = _test_tracer() - with tracer.start_as_current_span("ambient") as ambient: - headers = { - "Traceparent": "00-" + "a" * 32 + "-" + "b" * 16 + "-01", - "Tracestate": "vendor=old", - "x-keep": "1", - } + headers = {"Traceparent": _CALLER_TRACEPARENT, "Tracestate": "vendor=caller", "x-keep": "1"} + inbound = TraceContextTextMapPropagator().extract({key.lower(): value for key, value in headers.items()}) + with tracer.start_as_current_span("ambient", context=inbound) as ambient: result = ctx_mod.inject_trace_context(headers) propagated = get_current_span(TraceContextTextMapPropagator().extract(result)) return result, ambient, propagated result, ambient, propagated = ContextVarContext().run(run) assert sum(key.lower() == "traceparent" for key in result) == 1 - assert not any(key.lower() == "tracestate" for key in result) + assert sum(key.lower() == "tracestate" for key in result) == 1 assert result["x-keep"] == "1" - assert propagated.get_span_context().trace_id == ambient.get_span_context().trace_id + assert result["tracestate"] == "vendor=caller" + assert propagated.get_span_context().span_id == ambient.get_span_context().span_id + + +def test_inject_trace_context_keeps_caller_traceparent_from_another_trace(): + def run(): + tracer = _test_tracer() + parent = tracer.start_span("litellm_request") + with tracer.start_as_current_span("ambient") as ambient: + ctx_mod.set_request_root_span(ambient) + headers = {"Traceparent": _CALLER_TRACEPARENT, "Tracestate": "vendor=caller", "x-keep": "1"} + result = ctx_mod.inject_trace_context(headers, parent_span=parent) + propagated = get_current_span(TraceContextTextMapPropagator().extract(result)) + return result, parent, propagated + + result, parent, propagated = ContextVarContext().run(run) + assert result["traceparent"] == _CALLER_TRACEPARENT + assert result["tracestate"] == "vendor=caller" + assert result["x-keep"] == "1" + assert sum(key.lower() == "traceparent" for key in result) == 1 + assert sum(key.lower() == "tracestate" for key in result) == 1 + assert propagated.get_span_context().trace_id != parent.get_span_context().trace_id + + +def test_inject_trace_context_replaces_malformed_caller_traceparent(): + def run(): + tracer = _test_tracer() + parent = tracer.start_span("litellm_request") + with tracer.start_as_current_span("ambient"): + headers = {"traceparent": "not-a-traceparent", "tracestate": "vendor=caller"} + result = ctx_mod.inject_trace_context(headers, parent_span=parent) + propagated = get_current_span(TraceContextTextMapPropagator().extract(result)) + return result, parent, propagated + + result, parent, propagated = ContextVarContext().run(run) + assert propagated.get_span_context().span_id == parent.get_span_context().span_id + assert "tracestate" not in result def test_inject_trace_context_prefers_explicit_parent_span_over_root_and_ambient(): diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index e3e7ad618e0..df4842a7887 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4328,6 +4328,65 @@ async def test_pass_through_request_propagates_active_trace_context(span_source: assert propagated.get_span_context().span_id == span.get_span_context().span_id +async def _relay_with_trace_headers(inbound_headers: dict[str, str], forward_headers: bool): + from opentelemetry.sdk.trace import TracerProvider + + captured: dict[str, httpx.Headers] = {} + + def transport_handler(upstream_request: httpx.Request) -> httpx.Response: + captured["headers"] = upstream_request.headers + return httpx.Response(200, json={"ok": True}, request=upstream_request) + + fake_client, cleanup = _inject_fake_passthrough_client(httpx.MockTransport(transport_handler), timeout=None) + tracer = TracerProvider().get_tracer("test") + try: + with ExitStack() as stack: + _enter_relay_logging_mocks(stack, {}) + span = tracer.start_span("litellm_request") + stack.callback(span.end) + request = _relay_client_request(method="POST") + request.headers = Headers(inbound_headers) + response = await pass_through_request( + request=request, + target="http://internal-api.test/v1/generate", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", parent_otel_span=span), + forward_headers=forward_headers, + ) + finally: + cleanup() + await fake_client.aclose() + assert response.status_code == 200 + return captured["headers"], span + + +@pytest.mark.asyncio +@pytest.mark.parametrize("forward_headers", [False, True]) +async def test_pass_through_request_keeps_x_pass_trace_headers_when_otel_span_is_active(forward_headers: bool): + caller_traceparent = "00-11111111111111111111111111111111-2222222222222222-01" + + upstream_headers, span = await _relay_with_trace_headers( + {"x-pass-traceparent": caller_traceparent, "x-pass-tracestate": "vendor=caller"}, + forward_headers=forward_headers, + ) + + assert upstream_headers["traceparent"] == caller_traceparent + assert upstream_headers["tracestate"] == "vendor=caller" + assert format(span.get_span_context().trace_id, "032x") not in upstream_headers["traceparent"] + + +@pytest.mark.asyncio +async def test_pass_through_request_without_caller_trace_headers_still_propagates_proxy_span(): + from opentelemetry.trace import get_current_span + from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator + + upstream_headers, span = await _relay_with_trace_headers({"x-pass-anthropic-beta": "beta-1"}, forward_headers=False) + + propagated = get_current_span(TraceContextTextMapPropagator().extract(upstream_headers)) + assert propagated.get_span_context().span_id == span.get_span_context().span_id + assert upstream_headers["anthropic-beta"] == "beta-1" + + @pytest.mark.asyncio async def test_pass_through_request_relays_non_json_body_without_buffering(): """ From b346eefd1f230b38f32e37ecb7a4ee8f73effb8b Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 08:33:46 +0000 Subject: [PATCH 125/253] fix(proxy): rerun requested model guardrail merge on fallback instead of carrying a raw list Structured guardrail entries (dicts) are unhashable, so merging a carried list through _merge_guardrails_with_existing raised TypeError on the fallback path. Resolve the requested model alias through _check_and_merge_model_level_guardrails instead and add a regression test for structured request guardrails Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 31 +++++++------------ litellm/proxy/utils.py | 8 +++-- .../proxy/test_common_request_processing.py | 20 ++++++++++++ 3 files changed, 36 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 607ce3441e9..174b9ceff93 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -94,11 +94,7 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression from litellm.proxy.route_llm_request import route_request -from litellm.proxy.utils import ( - ProxyLogging, - _check_and_merge_model_level_guardrails, - _merge_guardrails_with_existing, -) +from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -1938,7 +1934,7 @@ class ProxyBaseLLMRequestProcessing: user_api_base: str | None = None, model: str | None = None, llm_router: Router | None = None, - requested_model_guardrails: list | None = None, + rate_limited_model: str | None = None, ) -> tuple[dict, LiteLLMLoggingObj]: start_time: Final = datetime.now() # start before calling guardrail hooks @@ -2102,12 +2098,15 @@ class ProxyBaseLLMRequestProcessing: # model_info when allow_client_pricing_override is set, so a caller # could otherwise spoof an unguarded model_info.id while requesting # a guarded alias and bypass guardrails (veria-ai HIGH on #29654). + merged_for_requested: Final = ( + self.data + if rate_limited_model is None + else _check_and_merge_model_level_guardrails( + data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model + ) + ) self.data = _check_and_merge_model_level_guardrails( - data=( - self.data - if requested_model_guardrails is None - else _merge_guardrails_with_existing(data=self.data, model_level_guardrails=requested_model_guardrails) - ), + data=merged_for_requested, llm_router=llm_router, trust_client_model_info=False, ) @@ -2217,8 +2216,6 @@ class ProxyBaseLLMRequestProcessing: original_model, fallback_models, ) - requested_model_guardrails: Final = self._request_guardrails(rate_limited_data) - try: for fallback_model in fallback_models: if fallback_model == original_model: @@ -2241,7 +2238,7 @@ class ProxyBaseLLMRequestProcessing: model=fallback_model, route_type=route_type, llm_router=llm_router, - requested_model_guardrails=requested_model_guardrails, + rate_limited_model=original_model, ) except ProxyRateLimitError: continue @@ -2252,12 +2249,6 @@ class ProxyBaseLLMRequestProcessing: self.data = rate_limited_data raise original_exc - @staticmethod - def _request_guardrails(data: dict) -> list | None: - metadata: Final = data.get("metadata") - guardrails: Final = metadata.get("guardrails") if isinstance(metadata, dict) else None - return guardrails if isinstance(guardrails, list) else None - @staticmethod def _configured_fallbacks(llm_router: Router, user_api_key_dict: UserAPIKeyAuth) -> list | None: key_router_settings: Final = user_api_key_dict.router_settings diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 950ac5e9906..ad0c0a82d49 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -7497,6 +7497,7 @@ def _check_and_merge_model_level_guardrails( data: dict, llm_router: Router | None, trust_client_model_info: bool = True, + model_alias: str | None = None, ) -> dict: """ Check if the model has guardrails defined and merge them with existing guardrails in the request data. @@ -7504,6 +7505,7 @@ def _check_and_merge_model_level_guardrails( Args: data: The request data dict llm_router: The LLM router instance to get deployment info from + model_alias: Resolve guardrails for this model group instead of data["model"] trust_client_model_info: If False, ignore metadata.model_info.id and resolve guardrails by alias-union only. Set to False on the pre_call path because add_litellm_data_to_request preserves @@ -7548,13 +7550,13 @@ def _check_and_merge_model_level_guardrails( # set on ANY eligible deployment still fires (#29652; addresses # veria-ai HIGH on the single-deployment fallback that would skip # non-first deployments). - model_alias: Final = data.get("model") - if not isinstance(model_alias, str) or not model_alias: + alias: Final = model_alias if model_alias is not None else data.get("model") + if not isinstance(alias, str) or not alias: return data # Pass team_id so team-scoped public model names resolve the same way # route_request resolves them; otherwise team-scoped deployments are # invisible to this lookup and their guardrails are silently dropped. - deployments: Final = llm_router.get_model_list(model_name=model_alias, team_id=team_id) or [] + deployments: Final = llm_router.get_model_list(model_name=alias, team_id=team_id) or [] seen: Final[set] = set() union: Final[list] = [] for dep in deployments: diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index f440711ac29..135ce019813 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -6970,6 +6970,26 @@ class TestPreCallWithFallbacksOnLocalRateLimit: assert data["messages"] == [{"role": "user", "content": "my ssn is [REDACTED-SSN]"}] assert rig[3] == [primary_model, primary_model, fallback_model] + @pytest.mark.asyncio + async def test_fallback_keeps_structured_request_guardrails(self, monkeypatch: pytest.MonkeyPatch): + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + structured_guardrail = {"pii-guard": {"extra_body": {"threshold": 0.5}}} + key = self._otel_key(model_rpm_limit={primary_model: 1}) + rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}]) + request = { + "model": primary_model, + "messages": [{"role": "user", "content": "hi"}], + "guardrails": [structured_guardrail], + } + + await self._pre_call(dict(request), key, rig) + _, (data, _) = await self._pre_call(dict(request), key, rig) + + assert data["model"] == fallback_model + assert data["metadata"]["guardrails"] == [structured_guardrail] + assert rig[3] == [primary_model, primary_model, fallback_model] + class _RecordingSuccessLogger(CustomLogger): def __init__(self): From a052da69744e6d09ba312e07a181b451552edffc Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 11 Sep 2026 23:54:35 +0000 Subject: [PATCH 126/253] feat(keys): let team service account keys use key management endpoints for their own team Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 9 + litellm/proxy/auth/route_checks.py | 5 +- .../key_management_endpoints.py | 41 ++++- .../team_member_permission_checks.py | 36 ++-- .../proxy/auth/test_route_checks.py | 69 +++++++ .../test_key_management_endpoints.py | 130 ++++++++++++- .../test_team_member_permission_checks.py | 172 ++++++++++++++++-- 7 files changed, 415 insertions(+), 47 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b0d31df92ce..cc44a86c691 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3270,6 +3270,15 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob user_role=LitellmUserRoles.PROXY_ADMIN, ) + @property + def is_team_service_account(self) -> bool: + return ( + self.user_id is None + and self.team_id is not None + and bool(self.metadata) + and self.metadata.get("service_account_id") is not None + ) + def user_api_key_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool: """Return True if the caller's role grants unscoped read access to all diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 0a6b618805d..001ab9a30b0 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -326,7 +326,10 @@ class RouteChecks: pass elif route.startswith("/v1/mcp/") or route.startswith("/mcp-rest/"): pass # authN/authZ handled by api itself - elif RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token): + elif RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token) or ( + valid_token.is_team_service_account + and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.key_management_routes.value) + ): pass elif valid_token.allowed_routes is not None: # check if route is in allowed_routes (exact match or prefix match) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 802a7c3e469..42fbe95487f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -509,6 +509,16 @@ def _get_user_in_team(team_table: LiteLLM_TeamTableCachedObj, user_id: str | Non return None +def _get_caller_team_role( + team_table: LiteLLM_TeamTableCachedObj, + user_api_key_dict: UserAPIKeyAuth, +) -> Literal["admin", "user"] | None: + if user_api_key_dict.is_team_service_account and user_api_key_dict.team_id == team_table.team_id: + return "user" + member: Final = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) + return None if member is None else member.role + + def _calculate_key_rotation_time(rotation_interval: str) -> datetime: """ Helper function to calculate the next rotation time for a key based on the rotation interval. @@ -603,7 +613,7 @@ def _team_key_operation_team_member_check( detail=f"User={assigned_user_id} not assigned to team={team_table.team_id}", ) - team_member_object: Final = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) + caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) is_admin: Final = ( user_api_key_dict.user_role is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -611,22 +621,22 @@ def _team_key_operation_team_member_check( if is_admin: return True - elif team_member_object is None: + elif caller_team_role is None: raise HTTPException( status_code=400, detail=f"User={user_api_key_dict.user_id} not assigned to team={team_table.team_id}", ) elif ( "allowed_team_member_roles" in team_key_generation - and team_member_object.role not in team_key_generation["allowed_team_member_roles"] + and caller_team_role not in team_key_generation["allowed_team_member_roles"] ): raise HTTPException( status_code=400, - detail=f"Team member role {team_member_object.role} not in allowed_team_member_roles={team_key_generation['allowed_team_member_roles']}", + detail=f"Team member role {caller_team_role} not in allowed_team_member_roles={team_key_generation['allowed_team_member_roles']}", ) TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( - team_member_object=team_member_object, + team_member_role=caller_team_role, team_table=team_table, route=route, ) @@ -747,6 +757,12 @@ def key_generation_check( Check if admin has restricted key creation to certain roles for teams or individuals """ + if user_api_key_dict.is_team_service_account and data.team_id != user_api_key_dict.team_id: + raise HTTPException( + status_code=403, + detail=f"Service account keys can only create keys for their own team. team_id={user_api_key_dict.team_id}", + ) + ## check if key is for team or individual is_team_key: Final = _is_team_key(data=data) _is_admin: Final = ( @@ -2223,6 +2239,10 @@ async def generate_service_account_key_fn( prisma_client=prisma_client, ) + if data.metadata is None or data.metadata.get("service_account_id") is None: + service_account_id: Final = (data.metadata or {}).get("service_account_id") or data.key_alias or str(uuid.uuid4()) + data.metadata = {**(data.metadata or {}), "service_account_id": service_account_id} # rebind-ok: stamping the generated service_account_id onto the request model so it persists on the key + verbose_proxy_logger.debug("entered /key/generate") custom_key_generate_hook: Final[Callable[..., Awaitable[Mapping[str, object]]] | None] = _custom_key_generate_hook( @@ -3837,8 +3857,10 @@ async def validate_key_team_change( detail=f"Key={key.token} has a rpm_limit={key.rpm_limit} which is greater than the team's rpm_limit={team.rpm_limit}.", ) + team_table: Final = cast(LiteLLM_TeamTableCachedObj, team) + # Check if the key's user_id is a member of the team - member_object: Final = _get_user_in_team(team_table=cast(LiteLLM_TeamTableCachedObj, team), user_id=key.user_id) + member_object: Final = _get_user_in_team(team_table=team_table, user_id=key.user_id) if key.user_id is not None: if not member_object: raise HTTPException( @@ -3854,8 +3876,11 @@ async def validate_key_team_change( team_obj=team, ) or TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( - team_member_object=member_object, - team_table=cast(LiteLLM_TeamTableCachedObj, team), + team_member_role=_get_caller_team_role( + team_table=team_table, + user_api_key_dict=change_initiated_by, + ), + team_table=team_table, route=KeyManagementRoutes.KEY_UPDATE.value, ) ): diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index 86c7a7bd947..a7e75caeb69 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -1,4 +1,4 @@ -from typing import Final +from typing import Final, Literal from litellm.proxy._types import ( KeyManagementRoutes, @@ -6,7 +6,6 @@ from litellm.proxy._types import ( LiteLLM_VerificationToken, LiteLLMRoutes, LitellmUserRoles, - Member, ProxyErrorTypes, ProxyException, UserAPIKeyAuth, @@ -27,7 +26,6 @@ DEFAULT_TEAM_MEMBER_PERMISSIONS: Final = BASELINE_TEAM_MEMBER_PERMISSIONS class TeamMemberPermissionChecks: @staticmethod def get_permissions_for_team_member( - team_member_object: Member, team_table: LiteLLM_TeamTableCachedObj, ) -> list[KeyManagementRoutes]: """ @@ -67,7 +65,7 @@ class TeamMemberPermissionChecks: Main handler for checking if a team member can update a key """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_user_in_team, + _get_caller_team_role, ) # 1. Don't execute these checks if the user role is proxy admin @@ -87,12 +85,12 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - # 4. Extract `Member` object from `team_table` - key_assigned_user_in_team: Final = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) + # 4. Resolve the caller's role in the key's team (service accounts act as "user") + caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # 5. Check if the team member has permissions for the endpoint has_permission: Final = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( - team_member_object=key_assigned_user_in_team, + team_member_role=caller_team_role, team_table=team_table, route=route, ) @@ -106,7 +104,7 @@ class TeamMemberPermissionChecks: @staticmethod def does_team_member_have_permissions_for_endpoint( - team_member_object: Member | None, + team_member_role: Literal["admin", "user"] | None, team_table: LiteLLM_TeamTableCachedObj, route: str, ) -> bool | None: @@ -116,13 +114,12 @@ class TeamMemberPermissionChecks: # permission checks only run for non-admin users # Non-Admin user trying to access information about a team's key - if team_member_object is None: + if team_member_role is None: return False - if team_member_object.role == "admin": + if team_member_role == "admin": return True _team_member_permissions: Final = TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=team_member_object, team_table=team_table, ) team_member_permissions = TeamMemberPermissionChecks._get_list_of_route_enum_as_str(_team_member_permissions) @@ -156,7 +153,7 @@ class TeamMemberPermissionChecks: from fastapi import HTTPException from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_user_in_team, + _get_caller_team_role, ) # No-op when the request does not assign any access groups. @@ -177,20 +174,19 @@ class TeamMemberPermissionChecks: ), ) - team_member_object: Final = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) + caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) # Team admins always bypass (consistent with other member-permission checks). - if team_member_object is not None and team_member_object.role == "admin": + if caller_team_role == "admin": return permissions: Final = ( TeamMemberPermissionChecks._get_list_of_route_enum_as_str( TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=team_member_object, team_table=team_table, ) ) - if team_member_object is not None + if caller_team_role is not None else [] ) @@ -214,7 +210,7 @@ class TeamMemberPermissionChecks: Returns True if the user belongs to the team that the key is assigned to """ from litellm.proxy.management_endpoints.key_management_endpoints import ( - _get_user_in_team, + _get_caller_team_role, ) from litellm.proxy.proxy_server import prisma_client, user_api_key_cache @@ -228,9 +224,9 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - # 4. Extract `Member` object from `team_table` - team_member_object: Final = _get_user_in_team(team_table=team_table, user_id=user_api_key_dict.user_id) - return team_member_object is not None + # 4. Resolve the caller's role in the key's team (service accounts act as "user") + caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) + return caller_team_role is not None @staticmethod def get_all_available_team_member_permissions() -> list[str]: diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 72c7011e6f5..b1405fc35cb 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3967,3 +3967,72 @@ def test_auto_router_session_read_grant_rejects_other_methods_paths_and_scopes( RouteChecks.should_call_route(route, valid_token, request) assert error.value.status_code == 403 + + +@pytest.mark.parametrize("route", ["/key/generate", "/key/update"]) +def test_team_service_account_key_allowed_key_management_routes(route): + """A service account key (user_id=None, team_id set, metadata.service_account_id) + can reach key-management routes; team scoping is enforced in the handlers.""" + valid_token = UserAPIKeyAuth( + api_key="sk", + team_id="t1", + user_id=None, + metadata={"service_account_id": "ci"}, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + result = RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=None, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + assert result is None + + +def test_team_service_account_key_rejected_for_non_key_management_route(): + """The service account carve-out does not extend past key-management routes.""" + valid_token = UserAPIKeyAuth( + api_key="sk", + team_id="t1", + user_id=None, + metadata={"service_account_id": "ci"}, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + with pytest.raises(Exception, match="Only proxy admin can be used to generate, delete, update"): + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=None, + route="/team/new", + request=request, + valid_token=valid_token, + request_data={}, + ) + + +def test_team_key_without_service_account_marker_still_rejected(): + """A team key without metadata.service_account_id is not a service account + and still cannot reach key-management routes.""" + valid_token = UserAPIKeyAuth( + api_key="sk", + team_id="t1", + user_id=None, + metadata={}, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + with pytest.raises(Exception, match="Only proxy admin can be used to generate, delete, update"): + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=None, + route="/key/generate", + request=request, + valid_token=valid_token, + request_data={}, + ) 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 cc0a7631b59..f865e959a78 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 @@ -18,6 +18,7 @@ import inspect from litellm.proxy._types import ( GenerateKeyRequest, + KeyManagementRoutes, NewUserRequest, LiteLLM_BudgetTable, LiteLLM_ObjectPermissionBase, @@ -3099,7 +3100,7 @@ async def test_validate_key_team_change_with_member_permissions(): # Verify the permission check was called with correct parameters mock_has_perms.assert_called_once_with( - team_member_object=mock_member_object, + team_member_role=mock_member_object.role, team_table=mock_team, route=KeyManagementRoutes.KEY_UPDATE.value, ) @@ -19733,3 +19734,130 @@ async def test_bulk_update_team_keys_runs_custom_key_policy_per_key(monkeypatch) assert [policy_request.operation for policy_request in received] == ["update", "update"] assert [policy_request.effective_key.max_budget for policy_request in received] == [50.0, 50.0] assert [policy_request.effective_key.team_id for policy_request in received] == ["team-abc", "team-abc"] + + +class TestServiceAccountKeyGenerationCheck: + """Service account keys (user_id=None, team_id set, metadata.service_account_id) + may only create keys for their own team.""" + + def _service_account_token(self, team_id: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-sa", + user_id=None, + team_id=team_id, + metadata={"service_account_id": "sa-1"}, + ) + + def test_other_team_denied(self): + data = GenerateKeyRequest(team_id="team-b") + with pytest.raises(HTTPException) as exc_info: + key_generation_check( + team_table=None, + user_api_key_dict=self._service_account_token(team_id="team-a"), + data=data, + route=KeyManagementRoutes.KEY_GENERATE, + ) + assert exc_info.value.status_code == 403 + + def test_personal_key_denied(self): + """team_id=None would mint a personal key; service accounts may only + create keys for their own team.""" + data = GenerateKeyRequest() + with pytest.raises(HTTPException) as exc_info: + key_generation_check( + team_table=None, + user_api_key_dict=self._service_account_token(team_id="team-a"), + data=data, + route=KeyManagementRoutes.KEY_GENERATE, + ) + assert exc_info.value.status_code == 403 + + def test_own_team_with_permission_allowed(self): + team_table = LiteLLM_TeamTableCachedObj( + team_id="team-a", + members_with_roles=[], + team_member_permissions=["/key/generate"], + ) + data = GenerateKeyRequest(team_id="team-a") + assert ( + key_generation_check( + team_table=team_table, + user_api_key_dict=self._service_account_token(team_id="team-a"), + data=data, + route=KeyManagementRoutes.KEY_GENERATE, + ) + is True + ) + + def test_own_team_without_permission_denied(self): + team_table = LiteLLM_TeamTableCachedObj( + team_id="team-a", + members_with_roles=[], + team_member_permissions=["/key/info"], + ) + data = GenerateKeyRequest(team_id="team-a") + with pytest.raises(ProxyException) as exc_info: + key_generation_check( + team_table=team_table, + user_api_key_dict=self._service_account_token(team_id="team-a"), + data=data, + route=KeyManagementRoutes.KEY_GENERATE, + ) + assert str(exc_info.value.code) == "401" + + +def _stub_service_account_generation(monkeypatch): + """Stub the DB lookups generate_service_account_key_fn needs so the test + exercises only the service_account_id stamping and user_id clearing.""" + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints import key_management_endpoints as kme + + mock_helper = AsyncMock(return_value=MagicMock()) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(kme, "validate_team_id_used_in_service_account_request", AsyncMock()) + monkeypatch.setattr(kme, "_common_key_generation_helper", mock_helper) + return mock_helper + + +@pytest.mark.asyncio +async def test_generate_service_account_key_stamps_service_account_id(monkeypatch): + """generate_service_account_key_fn must stamp metadata.service_account_id + (key_alias fallback) so the key is identifiable as a service account by + is_team_service_account and check_if_token_is_service_account.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_service_account_key_fn, + ) + + mock_helper = _stub_service_account_generation(monkeypatch) + data = GenerateKeyRequest(team_id="team-a", key_alias="sa-alias") + + await generate_service_account_key_fn( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"), + litellm_changed_by=None, + ) + + assert data.metadata is not None + assert data.metadata["service_account_id"] == "sa-alias" + assert data.user_id is None + mock_helper.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_generate_service_account_key_generates_uuid_when_no_alias(monkeypatch): + """Without key_alias, service_account_id falls back to a generated uuid.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_service_account_key_fn, + ) + + _stub_service_account_generation(monkeypatch) + data = GenerateKeyRequest(team_id="team-a") + + await generate_service_account_key_fn( + data=data, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1"), + litellm_changed_by=None, + ) + + assert data.metadata is not None + assert data.metadata["service_account_id"] diff --git a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py index 36c61eddbb2..55d29724e38 100644 --- a/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py +++ b/tests/test_litellm/proxy/management_helpers/test_team_member_permission_checks.py @@ -3,7 +3,12 @@ from unittest.mock import MagicMock import pytest -from litellm.proxy._types import KeyManagementRoutes, Member, ProxyException +from litellm.proxy._types import ( + KeyManagementRoutes, + Member, + ProxyException, + UserAPIKeyAuth, +) from litellm.proxy.management_helpers.team_member_permission_checks import ( BASELINE_TEAM_MEMBER_PERMISSIONS, TeamMemberPermissionChecks, @@ -21,22 +26,16 @@ class TestGetPermissionsForTeamMember: def test_none_permissions_returns_defaults(self): """When team_member_permissions is None, return DEFAULT_TEAM_MEMBER_PERMISSIONS.""" team = _make_team_table(None) - member = MagicMock(spec=Member) - result = TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=member, team_table=team - ) + result = TeamMemberPermissionChecks.get_permissions_for_team_member(team_table=team) assert set(result) == set(BASELINE_TEAM_MEMBER_PERMISSIONS) def test_empty_list_includes_baseline(self): """When team_member_permissions is [], baseline permissions are still included.""" team = _make_team_table([]) - member = MagicMock(spec=Member) - result = TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=member, team_table=team - ) + result = TeamMemberPermissionChecks.get_permissions_for_team_member(team_table=team) assert KeyManagementRoutes.KEY_INFO in result assert KeyManagementRoutes.KEY_HEALTH in result @@ -44,11 +43,8 @@ class TestGetPermissionsForTeamMember: def test_explicit_permissions_include_baseline(self): """When explicit permissions are set, baseline is always included.""" team = _make_team_table(["/key/generate", "/key/delete"]) - member = MagicMock(spec=Member) - result = TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=member, team_table=team - ) + result = TeamMemberPermissionChecks.get_permissions_for_team_member(team_table=team) assert KeyManagementRoutes.KEY_GENERATE in result assert KeyManagementRoutes.KEY_DELETE in result @@ -58,11 +54,8 @@ class TestGetPermissionsForTeamMember: def test_explicit_permissions_with_baseline_no_duplicates(self): """When explicit permissions already include baseline, no duplicates.""" team = _make_team_table(["/key/info", "/key/generate"]) - member = MagicMock(spec=Member) - result = TeamMemberPermissionChecks.get_permissions_for_team_member( - team_member_object=member, team_table=team - ) + result = TeamMemberPermissionChecks.get_permissions_for_team_member(team_table=team) # Using set ensures no duplicates from the implementation assert KeyManagementRoutes.KEY_INFO in result @@ -402,3 +395,148 @@ class TestEnforceMemberCanAssignAccessGroups: team_table=self._team(["/key/generate", self.AG_PERMISSION]), access_group_ids=["ag-1"], ) + + +class TestDoesTeamMemberHavePermissionsForEndpoint: + def _team(self, team_member_permissions, team_id="team-a"): + team = MagicMock() + team.team_id = team_id + team.team_member_permissions = team_member_permissions + return team + + def test_none_role_returns_false(self): + """A caller with no team membership is denied.""" + result = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( + team_member_role=None, + team_table=self._team(["/key/update"]), + route=KeyManagementRoutes.KEY_UPDATE.value, + ) + assert result is False + + def test_admin_role_always_allowed(self): + """Team admins bypass the member permission list.""" + result = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( + team_member_role="admin", + team_table=self._team([]), + route=KeyManagementRoutes.KEY_UPDATE.value, + ) + assert result is True + + def test_user_role_with_permission_allowed(self): + result = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( + team_member_role="user", + team_table=self._team(["/key/update"]), + route=KeyManagementRoutes.KEY_UPDATE.value, + ) + assert result is True + + def test_user_role_without_permission_raises(self): + with pytest.raises(ProxyException) as exc: + TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( + team_member_role="user", + team_table=self._team(["/key/generate"]), + route=KeyManagementRoutes.KEY_UPDATE.value, + ) + assert str(exc.value.code) == "401" + assert exc.value.type == "team_member_permission_error" + + +class TestCanTeamMemberExecuteKeyManagementEndpointServiceAccount: + def _service_account_token(self, team_id: str) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + user_id=None, + team_id=team_id, + metadata={"service_account_id": "sa-1"}, + ) + + @pytest.mark.asyncio + async def test_service_account_same_team_with_permission(self, monkeypatch): + """A service account key can manage keys in its own team when the + team grants the route via team_member_permissions.""" + from litellm.proxy.management_helpers import ( + team_member_permission_checks as module, + ) + + async def _mock_get_team_object(**kwargs): + team = MagicMock() + team.team_id = "team-a" + team.members_with_roles = [] + team.team_member_permissions = ["/key/update"] + return team + + monkeypatch.setattr(module, "get_team_object", _mock_get_team_object) + + existing_key_row = MagicMock() + existing_key_row.team_id = "team-a" + + result = await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=self._service_account_token(team_id="team-a"), + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + existing_key_row=existing_key_row, + ) + assert result is None + + @pytest.mark.asyncio + async def test_service_account_same_team_without_permission(self, monkeypatch): + """A service account key is denied when the team's + team_member_permissions does not include the route.""" + from litellm.proxy.management_helpers import ( + team_member_permission_checks as module, + ) + + async def _mock_get_team_object(**kwargs): + team = MagicMock() + team.team_id = "team-a" + team.members_with_roles = [] + team.team_member_permissions = ["/key/generate"] + return team + + monkeypatch.setattr(module, "get_team_object", _mock_get_team_object) + + existing_key_row = MagicMock() + existing_key_row.team_id = "team-a" + + with pytest.raises(ProxyException) as exc: + await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=self._service_account_token(team_id="team-a"), + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + existing_key_row=existing_key_row, + ) + assert str(exc.value.code) == "401" + assert exc.value.type == "team_member_permission_error" + + @pytest.mark.asyncio + async def test_service_account_different_team_denied(self, monkeypatch): + """A service account key cannot manage keys in another team, even if + that team grants the route to its members.""" + from litellm.proxy.management_helpers import ( + team_member_permission_checks as module, + ) + + async def _mock_get_team_object(**kwargs): + team = MagicMock() + team.team_id = "team-b" + team.members_with_roles = [] + team.team_member_permissions = ["/key/update"] + return team + + monkeypatch.setattr(module, "get_team_object", _mock_get_team_object) + + existing_key_row = MagicMock() + existing_key_row.team_id = "team-b" + + with pytest.raises(ProxyException) as exc: + await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=self._service_account_token(team_id="team-a"), + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + existing_key_row=existing_key_row, + ) + assert str(exc.value.code) == "401" + assert exc.value.type == "team_member_permission_error" From 062e17a7ac7fb1b78373e5198c40962c6252c699 Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 11 Sep 2026 23:55:49 +0000 Subject: [PATCH 127/253] refactor(keys): keep validate_key_team_change permission check on the key's assigned member Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/management_endpoints/key_management_endpoints.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 42fbe95487f..f0a689d6c50 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -3876,10 +3876,7 @@ async def validate_key_team_change( team_obj=team, ) or TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( - team_member_role=_get_caller_team_role( - team_table=team_table, - user_api_key_dict=change_initiated_by, - ), + team_member_role=None if member_object is None else member_object.role, team_table=team_table, route=KeyManagementRoutes.KEY_UPDATE.value, ) From 2542b9cfe8f6cf84f27943f5926c218fd027fa27 Mon Sep 17 00:00:00 2001 From: shivam Date: Sat, 12 Sep 2026 00:04:53 +0000 Subject: [PATCH 128/253] style(keys): fix ruff format and drop comment that restates the call Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/key_management_endpoints.py | 9 +++++++-- .../management_helpers/team_member_permission_checks.py | 4 +--- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index f0a689d6c50..705ea01f690 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2240,8 +2240,13 @@ async def generate_service_account_key_fn( ) if data.metadata is None or data.metadata.get("service_account_id") is None: - service_account_id: Final = (data.metadata or {}).get("service_account_id") or data.key_alias or str(uuid.uuid4()) - data.metadata = {**(data.metadata or {}), "service_account_id": service_account_id} # rebind-ok: stamping the generated service_account_id onto the request model so it persists on the key + service_account_id: Final = ( + (data.metadata or {}).get("service_account_id") or data.key_alias or str(uuid.uuid4()) + ) + data.metadata = { # rebind-ok: stamping the generated service_account_id onto the request model so it persists on the key + **(data.metadata or {}), + "service_account_id": service_account_id, + } verbose_proxy_logger.debug("entered /key/generate") diff --git a/litellm/proxy/management_helpers/team_member_permission_checks.py b/litellm/proxy/management_helpers/team_member_permission_checks.py index a7e75caeb69..a076d8240c6 100644 --- a/litellm/proxy/management_helpers/team_member_permission_checks.py +++ b/litellm/proxy/management_helpers/team_member_permission_checks.py @@ -85,10 +85,9 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - # 4. Resolve the caller's role in the key's team (service accounts act as "user") caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) - # 5. Check if the team member has permissions for the endpoint + # 4. Check if the team member has permissions for the endpoint has_permission: Final = TeamMemberPermissionChecks.does_team_member_have_permissions_for_endpoint( team_member_role=caller_team_role, team_table=team_table, @@ -224,7 +223,6 @@ class TeamMemberPermissionChecks: check_db_only=True, ) - # 4. Resolve the caller's role in the key's team (service accounts act as "user") caller_team_role: Final = _get_caller_team_role(team_table=team_table, user_api_key_dict=user_api_key_dict) return caller_team_role is not None From b78793b153116c24ced32ef88e207160d35da606 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 09:16:57 +0000 Subject: [PATCH 129/253] fix(auth): limit team service account route carve-out to /key/generate and /key/update The key_management_routes group also contains /spend/logs, /team/daily/activity and other routes whose handlers scope non-admin callers by user_id. A userless service account key would have reached them unscoped, so the route check now uses a dedicated two-route allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 5 +++++ litellm/proxy/auth/route_checks.py | 4 +++- tests/test_litellm/proxy/auth/test_route_checks.py | 8 +++++--- 3 files changed, 13 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index cc44a86c691..a31b6f834b3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -660,6 +660,11 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.AUTO_ROUTER_MANAGE.value, ] + team_service_account_key_routes = [ + KeyManagementRoutes.KEY_GENERATE.value, + KeyManagementRoutes.KEY_UPDATE.value, + ] + management_routes = ( [ # user diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 001ab9a30b0..1b9fd7c42bf 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -328,7 +328,9 @@ class RouteChecks: pass # authN/authZ handled by api itself elif RouteChecks.check_passthrough_route_access(route=route, user_api_key_dict=valid_token) or ( valid_token.is_team_service_account - and RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.key_management_routes.value) + and RouteChecks.check_route_access( + route=route, allowed_routes=LiteLLMRoutes.team_service_account_key_routes.value + ) ): pass elif valid_token.allowed_routes is not None: diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index b1405fc35cb..72c59223549 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3993,8 +3993,10 @@ def test_team_service_account_key_allowed_key_management_routes(route): assert result is None -def test_team_service_account_key_rejected_for_non_key_management_route(): - """The service account carve-out does not extend past key-management routes.""" +@pytest.mark.parametrize("route", ["/team/new", "/spend/logs", "/key/delete", "/key/regenerate"]) +def test_team_service_account_key_rejected_outside_generate_and_update(route): + """The service account carve-out covers only /key/generate and /key/update; other + key-management routes lack team scoping for a userless caller and stay denied.""" valid_token = UserAPIKeyAuth( api_key="sk", team_id="t1", @@ -4008,7 +4010,7 @@ def test_team_service_account_key_rejected_for_non_key_management_route(): RouteChecks.non_proxy_admin_allowed_routes_check( user_obj=None, _user_role=None, - route="/team/new", + route=route, request=request, valid_token=valid_token, request_data={}, From fade26b969880633ce4f87d0a81d01a2372f76ac Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 09:25:34 +0000 Subject: [PATCH 130/253] feat(proxy): derive the per-username sign-in allowance from the address limit The per-address-and-username allowance is now half the effective address allowance, rounded up, instead of a separate max_failed_login_attempts_per_user setting. A per-address override therefore raises or effectively removes both limits for that address, and no second override table is needed Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 9 +-- litellm/proxy/auth/login_throttle.py | 13 +++-- .../proxy/auth/test_login_utils.py | 56 ++++++++++++++++++- .../proxy_server/test_routes_login_sso.py | 18 +++--- tests/test_litellm/proxy/test_proxy_server.py | 2 - ui/litellm-dashboard/src/lib/http/schema.d.ts | 9 +-- 6 files changed, 76 insertions(+), 31 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 80147863304..801babbbb2d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2755,16 +2755,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): max_failed_login_attempts_per_source: int | None = Field( None, ge=1, - description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and this limit is off. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", + description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded up, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", ) max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field( None, - description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins. Set under `general_settings` in config.yaml", - ) - max_failed_login_attempts_per_user: int | None = Field( - None, - ge=1, - description="Failed Admin UI sign-in attempts allowed from one source address for one username within `failed_login_window_seconds`. One more blocks that address for that username for `failed_login_block_seconds`, and its further failures stop counting against `max_failed_login_attempts_per_source`, so a script stuck on one account does not block everyone behind the same address. Set under `general_settings` in config.yaml. Defaults to 5", + description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A very large value opts the address out of both limits. Set under `general_settings` in config.yaml", ) failed_login_window_seconds: int | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index d16cd43e36d..012d557f334 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -39,7 +39,6 @@ from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges from litellm.secret_managers.main import get_secret_bool DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 10 -DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER: Final = 5 DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60 DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300 @@ -47,7 +46,6 @@ IPV6_SOURCE_PREFIX_LENGTH: Final = 64 SOURCE_LIMIT_KEY: Final = "max_failed_login_attempts_per_source" SOURCE_LIMIT_OVERRIDES_KEY: Final = "max_failed_login_attempts_per_source_overrides" -USER_LIMIT_KEY: Final = "max_failed_login_attempts_per_user" WINDOW_KEY: Final = "failed_login_window_seconds" BLOCK_KEY: Final = "failed_login_block_seconds" TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges" @@ -204,6 +202,11 @@ def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: return matches[-1][1] if matches else default +def user_limit_for(source_limit: int) -> int: + """Failures allowed for one username from one address: half the address allowance, rounded up.""" + return (source_limit + 1) // 2 + + def source_group(client_ip: str) -> str: """The bucket an address is counted in: IPv4 as is, IPv6 by its /64, so one prefix holder cannot rotate.""" address: Final = _parse_address(client_ip) @@ -233,6 +236,7 @@ class LoginThrottle: ``source_limit`` is None when the source scope is off: ``trusted_proxy_ranges`` is unset, so the peer address may be a shared ingress. An empty list means clients connect directly and the peer is the source. + ``user_limit`` is derived from the address allowance either way, see ``user_limit_for``. """ client_ip: str @@ -257,10 +261,11 @@ class LoginThrottle: resolved, _ = resolve_client_ip( request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ()) ) + source_limit: Final = _source_limit(settings, resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE) return cls( client_ip=resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE, - source_limit=_source_limit(settings, resolved) if proxies is not None and resolved is not None else None, - user_limit=_int_setting(settings, USER_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_USER), + source_limit=source_limit if proxies is not None and resolved is not None else None, + user_limit=user_limit_for(source_limit), window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS), block_seconds=_int_setting(settings, BLOCK_KEY, DEFAULT_FAILED_LOGIN_BLOCK_SECONDS), counters=_COUNTERS, diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 986760c3cc5..d6a094d56cb 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1471,7 +1471,6 @@ def test_settings_that_arrive_as_environment_strings_are_honored(): general_settings={ "trusted_proxy_ranges": "10.0.0.0/8", "max_failed_login_attempts_per_source": " 70 ", - "max_failed_login_attempts_per_user": "7", "failed_login_window_seconds": "not-a-number", "failed_login_block_seconds": "-5", }, @@ -1479,7 +1478,7 @@ def test_settings_that_arrive_as_environment_strings_are_honored(): ) assert throttle.source_limit == 70 - assert throttle.user_limit == 7 + assert throttle.user_limit == 35, "the per-username allowance is half the address allowance" assert throttle.window_seconds == 60, "garbage falls back to the default" assert throttle.block_seconds == 300, "a value below one would block nothing or forever" @@ -1503,6 +1502,59 @@ def test_the_defaults_are_the_agreed_ones(): ) +@pytest.mark.parametrize( + ("source_limit", "expected_user_limit"), + [(1, 1), (2, 1), (3, 2), (10, 5), (1_000_000, 500_000)], + ids=["one-stays-one", "two-halves-to-one", "odd-rounds-up", "default", "opt-out"], +) +def test_the_per_username_allowance_is_half_the_address_allowance_rounded_up(source_limit, expected_user_limit): + from litellm.proxy.auth.login_throttle import user_limit_for + + assert user_limit_for(source_limit) == expected_user_limit + + +def test_a_per_address_override_also_raises_that_address_per_username_allowance(): + """One override opts an address out of both limits, so operators need no second override table.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + settings = { + "trusted_proxy_ranges": ["10.0.0.0/8"], + "max_failed_login_attempts_per_source": 10, + "max_failed_login_attempts_per_source_overrides": {"203.0.113.0/24": 1_000_000}, + } + + def _from(client_ip: str) -> LoginThrottle: + request = MagicMock() + request.headers = {"x-forwarded-for": client_ip} + request.client = MagicMock() + request.client.host = "10.0.0.1" + return LoginThrottle.from_request(request, general_settings=settings, redis_cache=None) + + exempt = _from("203.0.113.9") + assert (exempt.source_limit, exempt.user_limit) == (1_000_000, 500_000) + + ordinary = _from("198.51.100.4") + assert (ordinary.source_limit, ordinary.user_limit) == (10, 5) + + +def test_the_per_username_allowance_follows_the_peer_override_when_the_source_scope_is_off(): + """Without trusted_proxy_ranges the address is not blocked, but its override still sizes the pair limit.""" + from litellm.proxy.auth.login_throttle import LoginThrottle + + request = MagicMock() + request.headers = {} + request.client = MagicMock() + request.client.host = "192.0.2.8" + throttle = LoginThrottle.from_request( + request, + general_settings={"max_failed_login_attempts_per_source_overrides": {"192.0.2.8": 40}}, + redis_cache=None, + ) + + assert throttle.source_limit is None + assert throttle.user_limit == 20 + + def test_the_disable_flag_is_read_once_not_per_login_attempt(monkeypatch): """Regression: the kill switch was read through the secret manager on every unauthenticated request.""" from litellm.proxy.auth import login_throttle diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py index a82f078fcb9..88f8be4e49a 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_login_sso.py @@ -531,7 +531,7 @@ def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset """ _install_real_auth( monkeypatch, - max_failed_login_attempts_per_user=10, + max_failed_login_attempts_per_source=20, control_plane_url="https://cp.example.com", ) @@ -544,7 +544,7 @@ def test_budget_is_shared_across_every_login_endpoint(client, monkeypatch, reset def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_login_throttle): """The database lookup is case-insensitive, so casing must not partition the counter.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=3) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=6) assert [_json_login(client, "/v2/login", username="admin@corp.com") for _ in range(2)] == [401] * 2 assert [_json_login(client, "/v2/login", username="ADMIN@corp.com") for _ in range(2)] == [401] * 2 @@ -554,7 +554,7 @@ def test_budget_is_shared_across_username_casing(client, monkeypatch, reset_logi def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_throttle): """The 429 tells the caller how long the block has left.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=77) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=77) assert [_json_login(client, "/v2/login") for _ in range(2)] == [401, 401] @@ -565,7 +565,7 @@ def test_a_refused_attempt_carries_retry_after(client, monkeypatch, reset_login_ def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, reset_login_throttle): """The no-JavaScript form must render a wait page when its POST is throttled.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=77) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=77) assert [_form_login(client) for _ in range(2)] == [401, 401] @@ -578,7 +578,7 @@ def test_the_form_returns_a_human_readable_lockout_page(client, monkeypatch, res def test_a_second_username_from_the_same_source_still_gets_through(client, monkeypatch, reset_login_throttle): """The pair block is per username, so one account's block cannot take the office down with it.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2) assert [_json_login(client, "/v2/login", username="admin") for _ in range(3)] == [401, 401, 429] @@ -633,7 +633,7 @@ def test_the_configured_admin_password_is_refused_while_blocked(client, monkeypa limit. An operator who is blocked administers the proxy with the master key over the API meanwhile.""" from unittest.mock import AsyncMock, patch - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2) monkeypatch.setenv("DATABASE_URL", "postgresql://stub") assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] @@ -655,7 +655,7 @@ def test_the_master_key_as_a_bearer_token_still_works_while_the_ui_password_is_b client, monkeypatch, reset_login_throttle ): """Lockout recovery: the API path with the master key never enters the sign-in throttle.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2) assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] @@ -667,7 +667,7 @@ def test_the_master_key_as_a_bearer_token_still_works_while_the_ui_password_is_b def test_a_database_users_correct_password_is_refused_while_blocked(client, monkeypatch, reset_login_throttle): """The block is hard: while it lasts, nothing from that source signs in as that user, right password or not, and the block is not extended by the refused attempts.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1, failed_login_block_seconds=64) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2, failed_login_block_seconds=64) _db_user(monkeypatch, "user@corp.com") assert [_json_login(client, "/v2/login", username="user@corp.com") for _ in range(3)] == [401, 401, 429] @@ -682,7 +682,7 @@ def test_a_database_users_correct_password_is_refused_while_blocked(client, monk def test_sign_in_succeeds_again_once_the_block_is_cleared(client, monkeypatch, reset_login_throttle): """A cleared store lets the same username straight back to a plain credential check.""" - _install_real_auth(monkeypatch, max_failed_login_attempts_per_user=1) + _install_real_auth(monkeypatch, max_failed_login_attempts_per_source=2) assert [_json_login(client, "/v2/login") for _ in range(3)] == [401, 401, 429] diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 4c8cacf120f..9e2b85e765d 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -13519,13 +13519,11 @@ async def test_login_throttle_settings_are_not_hot_applied_from_the_database(): ps.general_settings.clear() await ProxyConfig()._update_general_settings( db_general_settings={ - "max_failed_login_attempts_per_user": 999, "max_failed_login_attempts_per_source": 999, "failed_login_window_seconds": 1, "failed_login_block_seconds": 1, } ) - assert "max_failed_login_attempts_per_user" not in ps.general_settings assert "max_failed_login_attempts_per_source" not in ps.general_settings assert "failed_login_window_seconds" not in ps.general_settings assert "failed_login_block_seconds" not in ps.general_settings diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 871cfea2d29..98bb3182916 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26566,21 +26566,16 @@ export interface components { max_batch_file_size_mb?: number | null; /** * Max Failed Login Attempts Per Source - * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and this limit is off. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 + * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded up, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 */ max_failed_login_attempts_per_source?: number | null; /** * Max Failed Login Attempts Per Source Overrides - * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins. Set under `general_settings` in config.yaml + * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A very large value opts the address out of both limits. Set under `general_settings` in config.yaml */ max_failed_login_attempts_per_source_overrides?: { [key: string]: number; } | null; - /** - * Max Failed Login Attempts Per User - * @description Failed Admin UI sign-in attempts allowed from one source address for one username within `failed_login_window_seconds`. One more blocks that address for that username for `failed_login_block_seconds`, and its further failures stop counting against `max_failed_login_attempts_per_source`, so a script stuck on one account does not block everyone behind the same address. Set under `general_settings` in config.yaml. Defaults to 5 - */ - max_failed_login_attempts_per_user?: number | null; /** * Max File Size Mb * @description max file size in MB for /v1/files uploads, for any purpose, if a file is larger than this size it will be rejected before being forwarded to the provider From e68ea077f169d179e13610edbca937415563937c Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 09:26:29 +0000 Subject: [PATCH 131/253] style(keys): keep rebind-ok reason within the line limit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/management_endpoints/key_management_endpoints.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 705ea01f690..6d4b7f05532 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2243,7 +2243,7 @@ async def generate_service_account_key_fn( service_account_id: Final = ( (data.metadata or {}).get("service_account_id") or data.key_alias or str(uuid.uuid4()) ) - data.metadata = { # rebind-ok: stamping the generated service_account_id onto the request model so it persists on the key + data.metadata = { # rebind-ok: stamp the service_account_id onto the request so it persists on the key **(data.metadata or {}), "service_account_id": service_account_id, } From b50265d75f719bfd2dd27b46c237a0981e6f1cbe Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 09:34:44 +0000 Subject: [PATCH 132/253] refactor(keys): freeze the service account route tuple and stamp metadata without seeding mutable literals Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 4 ++-- .../management_endpoints/key_management_endpoints.py | 9 ++++----- 2 files changed, 6 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a31b6f834b3..2ca6abf7873 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -660,10 +660,10 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.AUTO_ROUTER_MANAGE.value, ] - team_service_account_key_routes = [ + team_service_account_key_routes = ( KeyManagementRoutes.KEY_GENERATE.value, KeyManagementRoutes.KEY_UPDATE.value, - ] + ) management_routes = ( [ diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 6d4b7f05532..f6f47a8ffc9 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2240,13 +2240,12 @@ async def generate_service_account_key_fn( ) if data.metadata is None or data.metadata.get("service_account_id") is None: - service_account_id: Final = ( - (data.metadata or {}).get("service_account_id") or data.key_alias or str(uuid.uuid4()) - ) - data.metadata = { # rebind-ok: stamp the service_account_id onto the request so it persists on the key - **(data.metadata or {}), + service_account_id: Final = data.key_alias or str(uuid.uuid4()) + stamped_metadata: Final = { # mutable-ok: GenerateKeyRequest.metadata is a plain dict field + **(data.metadata or MappingProxyType({})), "service_account_id": service_account_id, } + data.metadata = stamped_metadata # rebind-ok: the request carries the stamp so it persists on the key verbose_proxy_logger.debug("entered /key/generate") From b6bb212248ff27676e63917ac3050a4809d827ea Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 09:40:05 +0000 Subject: [PATCH 133/253] refactor(proxy): raise the sign-in block explicitly and type the empty settings mapping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 10 +++++----- tests/test_litellm/proxy/auth/test_login_utils.py | 2 +- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 012d557f334..2ad75463629 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -17,7 +17,7 @@ import time from collections.abc import Mapping from dataclasses import dataclass from functools import cache -from typing import Final, Literal, NamedTuple, NoReturn, Protocol, TypeAlias +from typing import Final, Literal, NamedTuple, Protocol, TypeAlias from fastapi import Request, status from pydantic import TypeAdapter, ValidationError @@ -256,7 +256,7 @@ class LoginThrottle: general_settings: Mapping[str, object] | None, redis_cache: RedisCache | None, ) -> LoginThrottle: - settings: Final = general_settings if general_settings is not None else EMPTY_MAPPING + settings: Final[Mapping[str, object]] = general_settings if general_settings is not None else EMPTY_MAPPING proxies: Final = declared_proxy_ranges(settings) resolved, _ = resolve_client_ip( request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ()) @@ -298,7 +298,7 @@ class LoginThrottle: username, self.client_ip, ) - self.refuse(block.retry_after) + raise self.refused(block.retry_after) async def _active_block(self, keys: _Keys) -> Block | None: local: Final = self._local_block_ttls(keys) @@ -375,8 +375,8 @@ class LoginThrottle: ) @staticmethod - def refuse(retry_after: int) -> NoReturn: - raise ProxyException( + def refused(retry_after: int) -> ProxyException: + return ProxyException( message="Too many failed sign-in attempts. Try again later.", type=ProxyErrorTypes.auth_error, param="username", diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index d6a094d56cb..4d4a10a9fa4 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1436,8 +1436,8 @@ async def test_a_failed_redis_delete_still_clears_this_workers_counter(monkeypat @pytest.mark.asyncio async def test_counters_do_not_share_the_key_authentication_cache(monkeypatch): """Regression: throttle entries must not evict cached credentials from user_api_key_cache.""" - from litellm.proxy import proxy_server as ps from litellm.constants import LOGIN_THROTTLE_CACHE_KEY_PREFIX + from litellm.proxy import proxy_server as ps from litellm.proxy.auth.login_throttle import LoginThrottle monkeypatch.setenv("UI_USERNAME", "admin") From 0da2f5b96cb88031fc359a65506b7cc137677564 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 09:54:54 +0000 Subject: [PATCH 134/253] feat(proxy): round the per-username sign-in allowance down and exempt an address with an override of 0 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 4 +- litellm/proxy/auth/login_throttle.py | 27 ++++++--- .../proxy/auth/test_login_utils.py | 58 ++++++++++++++----- ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 +- 4 files changed, 66 insertions(+), 27 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 801babbbb2d..04ea7712545 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2755,11 +2755,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): max_failed_login_attempts_per_source: int | None = Field( None, ge=1, - description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded up, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", + description="Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded down but at least 1, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10", ) max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field( None, - description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A very large value opts the address out of both limits. Set under `general_settings` in config.yaml", + description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml", ) failed_login_window_seconds: int | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 2ad75463629..273cf8b9f93 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -43,6 +43,7 @@ DEFAULT_FAILED_LOGIN_WINDOW_SECONDS: Final = 60 DEFAULT_FAILED_LOGIN_BLOCK_SECONDS: Final = 300 IPV6_SOURCE_PREFIX_LENGTH: Final = 64 +EXEMPT: Final = 0 SOURCE_LIMIT_KEY: Final = "max_failed_login_attempts_per_source" SOURCE_LIMIT_OVERRIDES_KEY: Final = "max_failed_login_attempts_per_source_overrides" @@ -157,6 +158,13 @@ def _int_setting(settings: Mapping[str, object], key: str, default: int) -> int: return _positive_int(settings.get(key), key, default) +def _override_limit(raw: object, default: int) -> int: + """A per-address override: a limit of 1 or more, or ``EXEMPT`` (0) to leave that address unlimited.""" + if str(raw).strip() == str(EXEMPT): + return EXEMPT + return _positive_int(raw, SOURCE_LIMIT_OVERRIDES_KEY, default) + + def _parse_address(client_ip: str) -> ipaddress.IPv4Address | ipaddress.IPv6Address | None: """The address as it is limited and counted: an IPv4-mapped IPv6 address is its IPv4 address.""" try: @@ -179,7 +187,10 @@ def _parse_network(raw_range: str) -> _Network | None: def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: - """Failure allowance for this address: the most specific configured range containing it, else the default.""" + """Failure allowance for this address: the most specific configured range containing it, else the default. + + ``EXEMPT`` (0) means the operator opted this address out of both limits. + """ default: Final = _int_setting(settings, SOURCE_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE) raw_overrides: Final = settings.get(SOURCE_LIMIT_OVERRIDES_KEY) if raw_overrides is None: @@ -195,7 +206,7 @@ def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: if address is None: return default matches: Final = sorted( - (network.prefixlen, _positive_int(raw_limit, SOURCE_LIMIT_OVERRIDES_KEY, default)) + (network.prefixlen, _override_limit(raw_limit, default)) for raw_range, raw_limit in overrides.items() if (network := _parse_network(raw_range)) is not None and address in network ) @@ -203,8 +214,8 @@ def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: def user_limit_for(source_limit: int) -> int: - """Failures allowed for one username from one address: half the address allowance, rounded up.""" - return (source_limit + 1) // 2 + """Failures allowed for one username from one address: half the address allowance, rounded down, at least 1.""" + return max(source_limit // 2, 1) def source_group(client_ip: str) -> str: @@ -236,7 +247,8 @@ class LoginThrottle: ``source_limit`` is None when the source scope is off: ``trusted_proxy_ranges`` is unset, so the peer address may be a shared ingress. An empty list means clients connect directly and the peer is the source. - ``user_limit`` is derived from the address allowance either way, see ``user_limit_for``. + ``user_limit`` is derived from the address allowance either way, see ``user_limit_for``. An address whose + override is ``EXEMPT`` gets a disabled throttle: nothing is counted or blocked for it. """ client_ip: str @@ -262,16 +274,17 @@ class LoginThrottle: request, TrustedProxyConfig(use_forwarded_for=bool(proxies), trusted_proxy_cidrs=proxies or ()) ) source_limit: Final = _source_limit(settings, resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE) + exempt: Final = source_limit == EXEMPT return cls( client_ip=resolved or LOGIN_THROTTLE_UNKNOWN_SOURCE, - source_limit=source_limit if proxies is not None and resolved is not None else None, + source_limit=source_limit if proxies is not None and resolved is not None and not exempt else None, user_limit=user_limit_for(source_limit), window_seconds=_int_setting(settings, WINDOW_KEY, DEFAULT_FAILED_LOGIN_WINDOW_SECONDS), block_seconds=_int_setting(settings, BLOCK_KEY, DEFAULT_FAILED_LOGIN_BLOCK_SECONDS), counters=_COUNTERS, blocks=_BLOCKS, redis_cache=redis_cache, - enabled=not _rate_limit_disabled(), + enabled=not exempt and not _rate_limit_disabled(), ) def _keys(self, username: str) -> _Keys: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 4d4a10a9fa4..d6e0a489939 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -6,6 +6,7 @@ to login_utils.py for better reusability. """ import os +from collections.abc import Mapping from contextlib import ExitStack from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -1504,39 +1505,64 @@ def test_the_defaults_are_the_agreed_ones(): @pytest.mark.parametrize( ("source_limit", "expected_user_limit"), - [(1, 1), (2, 1), (3, 2), (10, 5), (1_000_000, 500_000)], - ids=["one-stays-one", "two-halves-to-one", "odd-rounds-up", "default", "opt-out"], + [(1, 1), (2, 1), (3, 1), (10, 5), (11, 5), (70, 35)], + ids=["one-stays-one", "two-halves-to-one", "odd-rounds-down", "default", "eleven-rounds-down", "even"], ) -def test_the_per_username_allowance_is_half_the_address_allowance_rounded_up(source_limit, expected_user_limit): +def test_the_per_username_allowance_is_half_the_address_allowance_rounded_down_at_least_one( + source_limit, expected_user_limit +): from litellm.proxy.auth.login_throttle import user_limit_for assert user_limit_for(source_limit) == expected_user_limit -def test_a_per_address_override_also_raises_that_address_per_username_allowance(): - """One override opts an address out of both limits, so operators need no second override table.""" +def _throttle_behind_trusted_proxy(client_ip: str, settings: Mapping[str, object]) -> "LoginThrottle": from litellm.proxy.auth.login_throttle import LoginThrottle + request = MagicMock() + request.headers = {"x-forwarded-for": client_ip} + request.client = MagicMock() + request.client.host = "10.0.0.1" + return LoginThrottle.from_request(request, general_settings=settings, redis_cache=None) + + +def test_a_per_address_override_also_raises_that_address_per_username_allowance(): + """One override sizes both limits for an address, so operators need no second override table.""" settings = { "trusted_proxy_ranges": ["10.0.0.0/8"], "max_failed_login_attempts_per_source": 10, - "max_failed_login_attempts_per_source_overrides": {"203.0.113.0/24": 1_000_000}, + "max_failed_login_attempts_per_source_overrides": {"203.0.113.0/24": 50}, } - def _from(client_ip: str) -> LoginThrottle: - request = MagicMock() - request.headers = {"x-forwarded-for": client_ip} - request.client = MagicMock() - request.client.host = "10.0.0.1" - return LoginThrottle.from_request(request, general_settings=settings, redis_cache=None) + raised = _throttle_behind_trusted_proxy("203.0.113.9", settings) + assert (raised.source_limit, raised.user_limit) == (50, 25) - exempt = _from("203.0.113.9") - assert (exempt.source_limit, exempt.user_limit) == (1_000_000, 500_000) - - ordinary = _from("198.51.100.4") + ordinary = _throttle_behind_trusted_proxy("198.51.100.4", settings) assert (ordinary.source_limit, ordinary.user_limit) == (10, 5) +@pytest.mark.asyncio +async def test_an_override_of_zero_exempts_that_address_from_both_limits(): + """Regression: opting an address out used to mean guessing a large enough number.""" + settings = { + "trusted_proxy_ranges": ["10.0.0.0/8"], + "max_failed_login_attempts_per_source": 1, + "max_failed_login_attempts_per_source_overrides": {"203.0.113.7": 0, "203.0.113.0/24": 3}, + } + + exempt = _throttle_behind_trusted_proxy("203.0.113.7", settings) + assert exempt.enabled is False + assert exempt.source_limit is None + attempt = await exempt.attempt("scanner@example.com") + for _ in range(5): + await attempt.failed() + await exempt.attempt("scanner@example.com") + + sibling = _throttle_behind_trusted_proxy("203.0.113.8", settings) + assert sibling.enabled is True + assert (sibling.source_limit, sibling.user_limit) == (3, 1) + + def test_the_per_username_allowance_follows_the_peer_override_when_the_source_scope_is_off(): """Without trusted_proxy_ranges the address is not blocked, but its override still sizes the pair limit.""" from litellm.proxy.auth.login_throttle import LoginThrottle diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 98bb3182916..e6e1627c4ca 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26566,12 +26566,12 @@ export interface components { max_batch_file_size_mb?: number | null; /** * Max Failed Login Attempts Per Source - * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded up, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 + * @description Failed Admin UI sign-in attempts allowed from one source address, across every username, within `failed_login_window_seconds`. One more blocks that address for `failed_login_block_seconds`. Half this value, rounded down but at least 1, is the allowance for one username from that address; one more blocks that address for that username only, and its further failures stop counting toward the address limit, so a script stuck on one account does not block everyone behind a shared address. The per-address limit is only enforced when `trusted_proxy_ranges` is set: to the proxies in front of LiteLLM, or to an empty list when clients connect directly. Left unset, the peer address may be a shared ingress and only the per-username half runs. IPv6 addresses are grouped by /64. Set under `general_settings` in config.yaml. Defaults to 10 */ max_failed_login_attempts_per_source?: number | null; /** * Max Failed Login Attempts Per Source Overrides - * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A very large value opts the address out of both limits. Set under `general_settings` in config.yaml + * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml */ max_failed_login_attempts_per_source_overrides?: { [key: string]: number; From 65765d6550fea8c7d1ac9fc700ea5d8171da8b54 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 09:57:35 +0000 Subject: [PATCH 135/253] test(proxy): import LoginThrottle under TYPE_CHECKING for the throttle helper annotation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/auth/test_login_utils.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index d6e0a489939..52ae3ca092b 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -8,11 +8,14 @@ to login_utils.py for better reusability. import os from collections.abc import Mapping from contextlib import ExitStack -from typing import Final +from typing import TYPE_CHECKING, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest +if TYPE_CHECKING: + from litellm.proxy.auth.login_throttle import LoginThrottle + def _unlimited_throttle(): """A throttle wired to real in-memory stores with limits no test can reach.""" From cfdf4fa5dc043f378d38c09a5b6ea87874076041 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 10:20:39 +0000 Subject: [PATCH 136/253] fix(proxy): break ties between equivalent login limit overrides deterministically Two spellings of one network share a prefix length, so the exemption wins the tie, then the higher limit, regardless of mapping order Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 2 +- litellm/proxy/auth/login_throttle.py | 12 +++++++++--- .../test_litellm/proxy/auth/test_login_utils.py | 17 +++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 4 files changed, 28 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 04ea7712545..851b2e7b28e 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2759,7 +2759,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) max_failed_login_attempts_per_source_overrides: dict[str, int] | None = Field( None, - description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml", + description="Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins (between equivalent keys such as '1.2.3.4' and '1.2.3.4/32', an exemption wins, then the higher limit), and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml", ) failed_login_window_seconds: int | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 273cf8b9f93..e4d9ffd4c5b 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -186,10 +186,16 @@ def _parse_network(raw_range: str) -> _Network | None: return None +def _precedence(network: _Network, limit: int) -> tuple[int, bool, int]: + """Sort key for competing overrides: the longest prefix wins, then an exemption, then the higher limit.""" + return (network.prefixlen, limit == EXEMPT, limit) + + def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: """Failure allowance for this address: the most specific configured range containing it, else the default. - ``EXEMPT`` (0) means the operator opted this address out of both limits. + ``EXEMPT`` (0) means the operator opted this address out of both limits. Between equivalent keys such as + ``1.2.3.4`` and ``1.2.3.4/32`` an exemption wins, then the higher limit. """ default: Final = _int_setting(settings, SOURCE_LIMIT_KEY, DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE) raw_overrides: Final = settings.get(SOURCE_LIMIT_OVERRIDES_KEY) @@ -206,11 +212,11 @@ def _source_limit(settings: Mapping[str, object], client_ip: str) -> int: if address is None: return default matches: Final = sorted( - (network.prefixlen, _override_limit(raw_limit, default)) + _precedence(network, _override_limit(raw_limit, default)) for raw_range, raw_limit in overrides.items() if (network := _parse_network(raw_range)) is not None and address in network ) - return matches[-1][1] if matches else default + return matches[-1][-1] if matches else default def user_limit_for(source_limit: int) -> int: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 52ae3ca092b..0cac0364219 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -1023,6 +1023,23 @@ def test_source_overrides_pick_the_most_specific_matching_range(): assert _limit("::ffff:203.0.113.10") == 200 +@pytest.mark.parametrize( + ("overrides", "expected"), + [ + ({"203.0.113.7": 0, "203.0.113.7/32": 5}, None), + ({"203.0.113.7/32": 5, "203.0.113.7": 0}, None), + ({"203.0.113.0/24": 3, "203.0.113.9/24": 8}, 8), + ({"203.0.113.9/24": 8, "203.0.113.0/24": 3}, 8), + ], + ids=["exact-then-slash32", "slash32-then-exact", "low-then-high", "high-then-low"], +) +def test_equivalent_override_keys_resolve_to_the_exemption_then_the_higher_limit(overrides, expected): + """Two spellings of the same network are a config mistake, so precedence must not depend on dict order.""" + settings = {"trusted_proxy_ranges": ["10.0.0.0/8"], "max_failed_login_attempts_per_source_overrides": overrides} + + assert _throttle_behind_trusted_proxy("203.0.113.7", settings).source_limit == expected + + def test_ipv6_sources_are_grouped_by_their_64_bit_prefix(): """A /64 holder has 2^64 addresses; counting each one separately would hand them unlimited fresh buckets.""" from litellm.proxy.auth.login_throttle import source_group diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index e6e1627c4ca..5b1532c3629 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26571,7 +26571,7 @@ export interface components { max_failed_login_attempts_per_source?: number | null; /** * Max Failed Login Attempts Per Source Overrides - * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins, and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml + * @description Per-address overrides of `max_failed_login_attempts_per_source`, keyed by IP address or CIDR range, e.g. {'1.2.3.4': 200, '5.6.0.0/24': 500}. The most specific matching range wins (between equivalent keys such as '1.2.3.4' and '1.2.3.4/32', an exemption wins, then the higher limit), and the per-username allowance for that address follows as half the override. A value of 0 exempts the address from both limits. Set under `general_settings` in config.yaml */ max_failed_login_attempts_per_source_overrides?: { [key: string]: number; From cacd12b87dc778f335ad6ac41dad02ba00312000 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 11:28:35 +0000 Subject: [PATCH 137/253] fix(proxy): treat a malformed trusted_proxy_ranges entry as an undeclared topology A list with an entry that is not an address or CIDR range no longer switches the per-source Admin UI sign-in limit on against the direct peer address, so a typo cannot make a shared ingress address the bucket for every user behind it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 2 +- litellm/proxy/auth/login_throttle.py | 19 ++++++++++--------- .../proxy/auth/test_login_utils.py | 19 ++++++++++++++++--- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 4 files changed, 28 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 851b2e7b28e..eee37016347 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2880,7 +2880,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): ) trusted_proxy_ranges: list[str] | None = Field( None, - description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, the per-source sign-in limit is off.", + description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off.", ) store_model_in_db: bool | None = Field( None, diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index e4d9ffd4c5b..49eb7a5f20d 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -120,9 +120,9 @@ def warn_login_counters_are_per_worker(num_workers: str) -> None: @cache def warn_source_login_limit_is_off() -> None: verbose_proxy_logger.warning( - "%s is not set, so failed Admin UI sign-in attempts are limited per source address and username " - "only. Set it to the address ranges of the proxies in front of LiteLLM, or to an empty list when " - "clients connect directly, to also limit each source address across usernames.", + "%s is not set or not a valid list of ranges, so failed Admin UI sign-in attempts are limited per " + "source address and username only. Set it to the address ranges of the proxies in front of LiteLLM, " + "or to an empty list when clients connect directly, to also limit each source address across usernames.", TRUSTED_PROXY_RANGES_KEY, ) @@ -131,13 +131,16 @@ def declared_proxy_ranges(settings: Mapping[str, object]) -> tuple[str, ...] | N """What the operator says fronts LiteLLM: the proxy ranges, an empty tuple for none, None when unsaid. Only a declared topology makes the source address trustworthy enough to limit across usernames. - An unset key, or a value that is not a list of ranges, leaves it unknown and the source scope off. + An unset key, a value that is not a list of ranges, or a list with an entry that is not an address + or range leaves it unknown and the source scope off. """ raw_ranges: Final = settings.get(TRUSTED_PROXY_RANGES_KEY) if isinstance(raw_ranges, (list, tuple, set)) and not raw_ranges: return () cidrs: Final = tuple(normalize_cidr_ranges(raw_ranges, setting_name=TRUSTED_PROXY_RANGES_KEY)) - return cidrs or None + if not cidrs or any(_parse_network(cidr, TRUSTED_PROXY_RANGES_KEY) is None for cidr in cidrs): + return None + return cidrs def _positive_int(raw: object, key: str, default: int) -> int: @@ -176,13 +179,11 @@ def _parse_address(client_ip: str) -> ipaddress.IPv4Address | ipaddress.IPv6Addr return address -def _parse_network(raw_range: str) -> _Network | None: +def _parse_network(raw_range: str, setting_name: str = SOURCE_LIMIT_OVERRIDES_KEY) -> _Network | None: try: return ipaddress.ip_network(raw_range.strip(), strict=False) except ValueError: - verbose_proxy_logger.warning( - "Invalid address or range %r in %s; skipping", raw_range, SOURCE_LIMIT_OVERRIDES_KEY - ) + verbose_proxy_logger.warning("Invalid address or range %r in %s; skipping", raw_range, setting_name) return None diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 0cac0364219..90186f4b8c0 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -934,10 +934,23 @@ async def test_an_empty_trusted_proxy_ranges_means_the_peer_is_the_client_and_th assert await _fail(throttle, username="user-99@corp.com") == "429", "the spray is stopped by the source limit" -@pytest.mark.parametrize("configured", [None, 5, {"10.0.0.0/8": True}, ["", " "]]) +@pytest.mark.parametrize( + "configured", + [ + None, + 5, + {"10.0.0.0/8": True}, + ["", " "], + ["not-a-range"], + ["10.0.0.0/8, 172.16.0.0/12"], + ["10.0.0.0/8", "10.0.0.0/33"], + "10.0.0.0/8;172.16.0.0/12", + ], +) def test_a_trusted_proxy_ranges_value_that_names_no_ranges_leaves_the_topology_unknown(configured): - """Only a real list of ranges or an explicit empty list counts as a declaration; anything else is the same - as unset, so a typo cannot switch the source-wide block on behind a shared ingress.""" + """Only a list of valid ranges or an explicit empty list counts as a declaration; anything else, including a + list with one bad entry, is the same as unset, so a typo cannot switch the source-wide block on against + the shared ingress address and lock out everyone behind it.""" from litellm.proxy.auth.login_throttle import LoginThrottle, declared_proxy_ranges settings = {"trusted_proxy_ranges": configured} if configured is not None else {} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 5b1532c3629..9cd578ea0a4 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26751,7 +26751,7 @@ export interface components { supported_db_objects?: components["schemas"]["SupportedDBObjectType"][] | null; /** * Trusted Proxy Ranges - * @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, the per-source sign-in limit is off. + * @description CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler, and whose X-Forwarded-For is used to attribute Admin UI sign-in attempts to a source address. Set it to an empty list when clients connect directly, so the peer address is the source. Left unset, or containing an entry that is not an address or CIDR range, the per-source sign-in limit is off. */ trusted_proxy_ranges?: string[] | null; /** From 8fc74a78bc0137b07af8a5c494b62cf77acc44ef Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 11:50:30 +0000 Subject: [PATCH 138/253] fix(proxy): reject blank trusted_proxy_ranges entries before they are dropped Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/login_throttle.py | 29 ++++++++++++++----- .../proxy/auth/test_login_utils.py | 5 ++++ 2 files changed, 27 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/auth/login_throttle.py b/litellm/proxy/auth/login_throttle.py index 49eb7a5f20d..b7f7eaceff4 100644 --- a/litellm/proxy/auth/login_throttle.py +++ b/litellm/proxy/auth/login_throttle.py @@ -35,7 +35,7 @@ from litellm.constants import ( LOGIN_THROTTLE_UNKNOWN_SOURCE, ) from litellm.proxy._types import ProxyErrorTypes, ProxyException -from litellm.proxy.auth.network import TrustedProxyConfig, normalize_cidr_ranges, resolve_client_ip +from litellm.proxy.auth.network import TrustedProxyConfig, resolve_client_ip from litellm.secret_managers.main import get_secret_bool DEFAULT_MAX_FAILED_LOGIN_ATTEMPTS_PER_SOURCE: Final = 10 @@ -54,6 +54,7 @@ TRUSTED_PROXY_RANGES_KEY: Final = "trusted_proxy_ranges" _REDIS_FAILURES: Final = (RedisError, RedisCircuitBreakerOpenError, OSError, asyncio.TimeoutError) _LOCAL_BLOCK_EXPIRY: Final = TypeAdapter[float | None](float | None) _SOURCE_LIMIT_OVERRIDES: Final = TypeAdapter[Mapping[str, object]](Mapping[str, object]) +_RANGE_ENTRIES: Final = TypeAdapter[tuple[object, ...]](tuple[object, ...]) Scope: TypeAlias = Literal["user", "source"] @@ -134,13 +135,27 @@ def declared_proxy_ranges(settings: Mapping[str, object]) -> tuple[str, ...] | N An unset key, a value that is not a list of ranges, or a list with an entry that is not an address or range leaves it unknown and the source scope off. """ - raw_ranges: Final = settings.get(TRUSTED_PROXY_RANGES_KEY) - if isinstance(raw_ranges, (list, tuple, set)) and not raw_ranges: - return () - cidrs: Final = tuple(normalize_cidr_ranges(raw_ranges, setting_name=TRUSTED_PROXY_RANGES_KEY)) - if not cidrs or any(_parse_network(cidr, TRUSTED_PROXY_RANGES_KEY) is None for cidr in cidrs): + entries: Final = _configured_range_entries(settings.get(TRUSTED_PROXY_RANGES_KEY)) + if entries is None or any(_parse_network(entry, TRUSTED_PROXY_RANGES_KEY) is None for entry in entries): + return None + return entries + + +def _configured_range_entries(raw_ranges: object) -> tuple[str, ...] | None: + """Every configured entry, blanks included, so a stray empty string fails validation like any other typo.""" + if raw_ranges is None: + return None + if isinstance(raw_ranges, str): + return tuple(part.strip() for part in raw_ranges.split(",")) + try: + return tuple(str(entry).strip() for entry in _RANGE_ENTRIES.validate_python(raw_ranges)) + except ValidationError: + verbose_proxy_logger.warning( + "Invalid %s value: expected a list of address ranges, got %s", + TRUSTED_PROXY_RANGES_KEY, + type(raw_ranges).__name__, + ) return None - return cidrs def _positive_int(raw: object, key: str, default: int) -> int: diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index 90186f4b8c0..55ece36252d 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -944,7 +944,12 @@ async def test_an_empty_trusted_proxy_ranges_means_the_peer_is_the_client_and_th ["not-a-range"], ["10.0.0.0/8, 172.16.0.0/12"], ["10.0.0.0/8", "10.0.0.0/33"], + ["10.0.0.0/8", " "], + ["10.0.0.0/8", ""], + ["10.0.0.0/8", None], "10.0.0.0/8;172.16.0.0/12", + "10.0.0.0/8,", + "", ], ) def test_a_trusted_proxy_ranges_value_that_names_no_ranges_leaves_the_topology_unknown(configured): From ca405879a2ba9915b4fec8a35efc48360a719d1e Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:22:57 +0000 Subject: [PATCH 139/253] fix(cohere): set embed v3 context length to 512 tokens Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 16 ++++++++-------- model_prices_and_context_window.json | 16 ++++++++-------- 2 files changed, 16 insertions(+), 16 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 92df776adf4..60f48ec9679 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -22170,8 +22170,8 @@ "embed-english-light-v3.0": { "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0 }, @@ -22188,8 +22188,8 @@ "input_cost_per_image": 0.0001, "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "metadata": { "notes": "'supports_image_input' is a deprecated field. Use 'supports_embedding_image_input' instead." }, @@ -22210,8 +22210,8 @@ "embed-multilingual-v3.0": { "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0, "supports_embedding_image_input": true @@ -22219,8 +22219,8 @@ "embed-multilingual-light-v3.0": { "input_cost_per_token": 0.0001, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0, "supports_embedding_image_input": true diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 92df776adf4..60f48ec9679 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -22170,8 +22170,8 @@ "embed-english-light-v3.0": { "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0 }, @@ -22188,8 +22188,8 @@ "input_cost_per_image": 0.0001, "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "metadata": { "notes": "'supports_image_input' is a deprecated field. Use 'supports_embedding_image_input' instead." }, @@ -22210,8 +22210,8 @@ "embed-multilingual-v3.0": { "input_cost_per_token": 1e-07, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0, "supports_embedding_image_input": true @@ -22219,8 +22219,8 @@ "embed-multilingual-light-v3.0": { "input_cost_per_token": 0.0001, "litellm_provider": "cohere", - "max_input_tokens": 1024, - "max_tokens": 1024, + "max_input_tokens": 512, + "max_tokens": 512, "mode": "embedding", "output_cost_per_token": 0.0, "supports_embedding_image_input": true From b6410d563bd36caf13a5a629aa2d5a8ceb858aa8 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 15:46:33 +0000 Subject: [PATCH 140/253] fix(scim): accept entitlements and roles entries without a value on SCIM user PUT SCIMMultiValuedAttribute required value, so a PUT /scim/v2/Users/{id} that carried an IdP-specific entitlements entry such as {"groups": [...]} failed body validation with 422 and the suspend (active: false) never reached update_user. value is now optional and unknown members are kept, so the suspend is applied, keys are blocked, and the entries are stored as sent Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_lazy_openapi_snapshot.json | 15 ++-- .../management_endpoints/scim/scim_v2.py | 2 +- .../proxy/management_endpoints/scim_v2.py | 4 +- .../scim/test_scim_patch_user.py | 18 ++++- .../scim/test_scim_v2_endpoints.py | 72 +++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 +- 6 files changed, 106 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index b244678e201..e889daf6c67 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -38245,6 +38245,7 @@ "type": "object" }, "SCIMMultiValuedAttribute": { + "additionalProperties": true, "properties": { "display": { "anyOf": [ @@ -38280,13 +38281,17 @@ "title": "Type" }, "value": { - "title": "Value", - "type": "string" + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Value" } }, - "required": [ - "value" - ], "title": "SCIMMultiValuedAttribute", "type": "object" }, diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 34c1ad42435..2b74dc1e838 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -2107,7 +2107,7 @@ def _handle_multi_valued_attribute_update(path: str, op_type: str, value: object except ValidationError: raise HTTPException( status_code=400, - detail={"error": f"Invalid value for {base}: expected a list of objects with a 'value' sub-attribute"}, + detail={"error": f"Invalid value for {base}: expected a list of objects or strings"}, ) dumped: Final = [attr.model_dump(exclude_none=True) for attr in attrs] diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 61fd5c36b16..6f2c48ab283 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -61,7 +61,9 @@ class SCIMUserGroup(BaseModel): class SCIMMultiValuedAttribute(BaseModel): - value: str + model_config = ConfigDict(extra="allow") + + value: str | None = None display: str | None = None type: str | None = None primary: bool | None = None diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py index c3af4208d37..edcce16ab41 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py @@ -422,7 +422,7 @@ def test_apply_patch_ops_invalid_entitlements_value_raises_400(): patch_ops = SCIMPatchOp( Operations=[ SCIMPatchOperation( - op="replace", path="entitlements", value=[{"display": "no value"}] + op="replace", path="entitlements", value=[42] ) ] ) @@ -433,6 +433,22 @@ def test_apply_patch_ops_invalid_entitlements_value_raises_400(): assert exc_info.value.status_code == 400 +def test_apply_patch_ops_replace_entitlements_without_value_member_is_stored_as_sent(): + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation( + op="replace", path="entitlements", value=[{"groups": ["S0506MKA55L"]}] + ) + ] + ) + + update_data, _ = _apply_patch_ops( + existing_user=_user_with_metadata({}), patch_ops=patch_ops + ) + + assert update_data["metadata"]["scim_entitlements"] == [{"groups": ["S0506MKA55L"]}] + + def test_apply_patch_ops_add_without_value_raises_400_naming_value_member(): patch_ops = SCIMPatchOp( Operations=[SCIMPatchOperation(op="add", path="entitlements")] 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 364ec4aad61..6d85691d8ac 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 @@ -1,3 +1,4 @@ +import json import logging import time from collections.abc import Callable, Mapping, Sequence @@ -1303,6 +1304,77 @@ async def test_update_user_success(mocker): assert call_args[1]["data"]["teams"] == ["new-team"] +@pytest.mark.asyncio +async def test_update_user_put_with_valueless_entitlements_deactivates_user(scim_test_client, mocker): + """A suspend PUT whose entitlements entries carry no `value` member (an IdP-specific shape) + must not be rejected by body validation: the user is deactivated and the entries are stored as sent""" + existing_user = mocker.MagicMock() + existing_user.teams = [] + existing_user.metadata = {"scim_active": True} + + updated_user = { + "user_id": "suspend-me", + "user_email": "suspend@example.com", + "user_alias": None, + "teams": [], + "metadata": "{}", + } + response_scim_user = SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="suspend-me", + userName="suspend-me", + active=False, + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) + + mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(return_value=existing_user), + ) + mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) + set_keys_blocked_mock = mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable + "litellm.proxy.management_endpoints.scim.scim_v2._set_user_keys_blocked", + AsyncMock(return_value=1), + ) + mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=response_scim_user), + ) + + async with scim_test_client as client: + response = await client.put( + "/scim/v2/Users/suspend-me", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": "suspend-me", + "emails": [{"value": "suspend@example.com", "primary": True}], + "entitlements": [{"groups": ["S0506MKA55L", "S0506MKA56M"]}], + "roles": [{"display": "Viewer"}], + "active": False, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["active"] is False + + written_metadata = json.loads(mock_prisma_client.db.litellm_usertable.update.call_args.kwargs["data"]["metadata"]) + assert written_metadata["scim_active"] is False + assert written_metadata["scim_entitlements"] == [{"groups": ["S0506MKA55L", "S0506MKA56M"]}] + assert written_metadata["scim_roles"] == [{"display": "Viewer"}] + set_keys_blocked_mock.assert_awaited_once_with(user_id="suspend-me", blocked=True) + + @pytest.mark.asyncio @pytest.mark.parametrize("groups", [None, []], ids=["groups-omitted", "groups-empty"]) async def test_update_user_without_groups_preserves_memberships_and_role(mocker, monkeypatch, groups): diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fd882937e79..4883611c3b1 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -36612,7 +36612,9 @@ export interface components { /** Type */ type?: string | null; /** Value */ - value: string; + value?: string | null; + } & { + [key: string]: unknown; }; /** SCIMPatchOp */ SCIMPatchOp: { From 817cdcef6422df4eee08b6022909c941b4750264 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 16:00:54 +0000 Subject: [PATCH 141/253] test(scim): drop redundant docstring from the valueless entitlements PUT test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/management_endpoints/scim/test_scim_v2_endpoints.py | 2 -- 1 file changed, 2 deletions(-) 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 6d85691d8ac..dbcf622bbb1 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 @@ -1306,8 +1306,6 @@ async def test_update_user_success(mocker): @pytest.mark.asyncio async def test_update_user_put_with_valueless_entitlements_deactivates_user(scim_test_client, mocker): - """A suspend PUT whose entitlements entries carry no `value` member (an IdP-specific shape) - must not be rejected by body validation: the user is deactivated and the entries are stored as sent""" existing_user = mocker.MagicMock() existing_user.teams = [] existing_user.metadata = {"scim_active": True} From dba9ff801f6b3474e7861b7c066396ebe19e3cf3 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 17:18:49 +0000 Subject: [PATCH 142/253] fix(proxy): reset sibling tpm/rpm counters when shared window rolls over Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../hooks/parallel_request_limiter_v3.py | 39 +++++++++++ .../hooks/test_parallel_request_limiter_v3.py | 69 +++++++++++++++++++ 2 files changed, 108 insertions(+) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 89826694bb6..57b7d4f249d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -91,10 +91,16 @@ else: _REQUEST_RATE_LIMIT_DATA: Final = TypeAdapter(Mapping[str, object]) +def _sibling_counter_keys(window_key: str) -> tuple[str, str]: + prefix: Final = window_key.removesuffix(":window") + return f"{prefix}:requests", f"{prefix}:tokens" + + BATCH_RATE_LIMITER_SCRIPT: Final = """ local results = {} local now = tonumber(ARGV[1]) local window_size = tonumber(ARGV[2]) +local reset_windows = {} -- Process each window/counter pair for i = 1, #KEYS, 2 do @@ -106,6 +112,11 @@ for i = 1, #KEYS, 2 do local window_start = redis.call('GET', window_key) if not window_start or (now - tonumber(window_start)) >= window_size then -- Reset window and counter + if not reset_windows[window_key] then + local prefix = string.sub(window_key, 1, -(#':window') - 1) + redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') + reset_windows[window_key] = true + end redis.call('SET', window_key, tostring(now)) redis.call('SET', counter_key, increment_value) redis.call('EXPIRE', window_key, window_size) @@ -151,6 +162,7 @@ CHECK_AND_INCREMENT_BY_N_SCRIPT: Final = """ local time_reply = redis.call('TIME') local now = tonumber(time_reply[1]) local descriptor_count = #KEYS / 2 +local reset_windows = {} -- Pass 1: read state, validate. Abort without writing if any over limit. local descriptor_state = {} @@ -201,6 +213,11 @@ for i = 1, descriptor_count do if window_expired then active_window_start = now + if not reset_windows[window_key] then + local prefix = string.sub(window_key, 1, -(#':window') - 1) + redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') + reset_windows[window_key] = true + end redis.call('SET', window_key, tostring(now)) redis.call('SET', counter_key, increment) redis.call('EXPIRE', window_key, window_size) @@ -1019,6 +1036,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): This follows the same logic as the Redis Lua script but uses async cache operations. """ results: Final[list[CacheCounterValue | None]] = [] + reset_windows: Final[set[str]] = set() # mutable-ok: tracks windows reset during this call # Process each window/counter pair for i in range(0, len(keys), 2): @@ -1036,6 +1054,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Check if window exists and is valid if window_start is None or (now_int - int(window_start)) >= window_size: # Reset window and counter + if window_key not in reset_windows: + for sibling_counter_key in _sibling_counter_keys(window_key): + await self.internal_usage_cache.async_set_cache( + key=sibling_counter_key, + value=0, + ttl=window_size, + litellm_parent_otel_span=None, + local_only=True, + ) + reset_windows.add(window_key) await self.internal_usage_cache.async_set_cache( key=window_key, value=str(now_int), @@ -2049,9 +2077,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Pass 2: apply increments. statuses: Final[list[RateLimitStatus]] = [] + reset_windows: Final[set[str]] = set() # mutable-ok: tracks windows reset during this call for meta, state in zip(per_counter_meta, descriptor_state): new_counter = meta["increment"] if state["window_expired"] else state["current"] + meta["increment"] if state["window_expired"]: + if meta["window_key"] not in reset_windows: + for sibling_counter_key in _sibling_counter_keys(meta["window_key"]): + await self.internal_usage_cache.async_set_cache( + key=sibling_counter_key, + value=0, + ttl=meta["window_size"], + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) + reset_windows.add(meta["window_key"]) await self.internal_usage_cache.async_set_cache( key=meta["window_key"], value=str(now_int), diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 85023b94207..0b7710eda70 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -5602,6 +5602,75 @@ async def _reserved_tokens_for( return int(await local_cache.async_get_cache(key=tokens_key) or 0) +@pytest.mark.asyncio +async def test_tpm_reservation_resets_sibling_tokens_with_request_window(monkeypatch): + monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "true") + time_controller = TimeController() + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + time_provider=time_controller.now, + ) + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-window-reset-siblings"), + tpm_limit=1000, + rpm_limit=1000, + ) + + async def request(call_id): + data = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 200, + "litellm_call_id": call_id, + "metadata": { + "user_api_key": user_api_key_dict.api_key, + "user_api_key_user_id": user_api_key_dict.user_id, + }, + } + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=data, + call_type="completion", + ) + await handler.async_log_success_event( + kwargs={ + "litellm_call_id": call_id, + "litellm_params": { + "metadata": { + "user_api_key": user_api_key_dict.api_key, + "user_api_key_user_id": user_api_key_dict.user_id, + "model_group": "gpt-4o", + } + }, + "standard_logging_object": { + "metadata": { + "user_api_key_hash": user_api_key_dict.api_key, + "user_api_key_user_id": user_api_key_dict.user_id, + } + }, + }, + response_obj=ModelResponse( + model="gpt-4o", + usage=Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300), + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + + tokens_key = handler.create_rate_limit_keys( + key="api_key", value=user_api_key_dict.api_key, rate_limit_type="tokens" + ) + for index in range(3): + await request(f"call-{index}") + assert await local_cache.async_get_cache(key=tokens_key) == (index + 1) * 300 + + time_controller.advance(61) + await request("call-after-window-reset") + assert await local_cache.async_get_cache(key=tokens_key) == 300 + + @pytest.mark.asyncio @pytest.mark.parametrize( "key_metadata, team_metadata, expected_output_estimate, tier", From 02c0ee4b5d477aa0c40aa67a97d55bf51296f34c Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 17:20:05 +0000 Subject: [PATCH 143/253] fix(proxy): drop _add_general_settings_from_db_config re-added by merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_server.py | 72 ----------------------------------- 1 file changed, 72 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6f191a42dcc..c3870a3e899 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7115,78 +7115,6 @@ class ProxyConfig: invalid_groups, ) - def _add_general_settings_from_db_config( - self, config_data: dict, general_settings: dict, proxy_logging_obj: ProxyLogging - ) -> None: - """ - Adds general settings from DB config to litellm proxy - - Args: - config_data: dict - general_settings: dict - global general_settings currently in use - proxy_logging_obj: ProxyLogging - """ - _general_settings: Final = config_data.get("general_settings", {}) - - if _general_settings is not None and "alerting" in _general_settings: - if ( - general_settings is not None - and general_settings.get("alerting", None) is not None - and isinstance(general_settings["alerting"], list) - and _general_settings.get("alerting", None) is not None - and isinstance(_general_settings["alerting"], list) - ): - # Merge DB and YAML/config alerting values instead of overriding - _yaml_alerting: Final = set(general_settings["alerting"]) - _db_alerting: Final = set(_general_settings["alerting"]) - _merged_alerting = list(_yaml_alerting.union(_db_alerting)) - # Preserve order: YAML values first, then DB values - _merged_alerting = list(general_settings["alerting"]) + [ - item for item in _general_settings["alerting"] if item not in general_settings["alerting"] - ] - verbose_proxy_logger.debug( - "Merging alerting values: YAML=%s, DB=%s, Merged=%s", - general_settings["alerting"], - _general_settings["alerting"], - _merged_alerting, - ) - general_settings["alerting"] = _merged_alerting - # Use update_values to properly set alerting for both slack and email - proxy_logging_obj.update_values( - alerting=general_settings["alerting"], - ) - elif general_settings is None: - general_settings = {} - general_settings["alerting"] = _general_settings["alerting"] - # Use update_values to properly set alerting for both slack and email - proxy_logging_obj.update_values( - alerting=general_settings["alerting"], - ) - elif isinstance(general_settings, dict): - general_settings["alerting"] = _general_settings["alerting"] - # Use update_values to properly set alerting for both slack and email - proxy_logging_obj.update_values( - alerting=general_settings["alerting"], - ) - - if _general_settings is not None and "alert_types" in _general_settings: - general_settings["alert_types"] = _general_settings["alert_types"] - proxy_logging_obj.alert_types = general_settings["alert_types"] - proxy_logging_obj.slack_alerting_instance.update_values( - alert_types=general_settings["alert_types"], llm_router=llm_router - ) - - if _general_settings is not None and "alert_to_webhook_url" in _general_settings: - general_settings["alert_to_webhook_url"] = _general_settings["alert_to_webhook_url"] - proxy_logging_obj.slack_alerting_instance.update_values( - alert_to_webhook_url=general_settings["alert_to_webhook_url"], - llm_router=llm_router, - ) - - if _general_settings is not None and "plugins" in _general_settings: - general_settings["plugins"] = _general_settings["plugins"] - register_plugins_from_config(general_settings) - async def _reschedule_spend_log_cleanup_job(self): """ Reschedule the spend log cleanup job based on current general_settings. From 8dfda9312370f2a1f4ebb531b5ba69c924ef8c10 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 17:32:31 +0000 Subject: [PATCH 144/253] test(proxy): drop explanatory docstrings from routing_groups regression tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/test_proxy_server.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 82565f000bf..ce809eda847 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4958,8 +4958,6 @@ def _routing_groups_router(): @pytest.mark.asyncio async def test_invalid_db_routing_groups_do_not_abort_other_router_settings(): - """Regression: an overlapping routing_groups value persisted in DB used to raise out of - _add_router_settings_from_db_config, which skipped SSO / guardrail loading downstream.""" from unittest.mock import AsyncMock, MagicMock from litellm.proxy.proxy_server import ProxyConfig @@ -9341,8 +9339,6 @@ def test_update_config_writes_only_sent_section(_update_config_setup): def test_update_config_rejects_overlapping_routing_groups_before_writing(_update_config_setup): - """Regression: overlapping groups were persisted and only failed at router reload, where the - failure took SSO and the other DB-backed settings down with it.""" existing_groups = [{"group_name": "g1", "models": ["m1"], "routing_strategy": "least-busy"}] client, prisma, restore = _update_config_setup( initial_rows={"router_settings": {"num_retries": 2, "routing_groups": existing_groups}} From 48f5dc61760abfce73ed2e99c8991891b0822edc Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 10:48:07 -0700 Subject: [PATCH 145/253] chore(proxy): keep the lazy OpenAPI snapshot as CI's Python 3.12 generates it --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fa046ef0a72..b244678e201 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19346,7 +19346,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " }, "500": { "content": { From 639c71e5bb03b7c4307fece027db7d87e4835ab5 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 18:09:20 +0000 Subject: [PATCH 146/253] chore(deps): bump anyio from 4.13.0 to 4.14.2 Bumps [anyio](https://github.com/agronholm/anyio) from 4.13.0 to 4.14.2. - [Release notes](https://github.com/agronholm/anyio/releases) - [Commits](https://github.com/agronholm/anyio/compare/4.13.0...4.14.2) --- updated-dependencies: - dependency-name: anyio dependency-version: 4.14.2 dependency-type: indirect ... Signed-off-by: dependabot[bot] --- uv.lock | 388 ++++++++++++++++++++++++++++---------------------------- 1 file changed, 194 insertions(+), 194 deletions(-) diff --git a/uv.lock b/uv.lock index a5e60c68515..f04fa5a17c1 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-09-14T23:55:55.024292355Z" +exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values. exclude-newer-span = "P3D" [manifest] @@ -225,9 +225,9 @@ name = "aiologic" version = "0.17.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "sniffio", marker = "python_full_version < '3.13'" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, - { name = "wrapt", marker = "python_full_version < '3.13'" }, + { name = "sniffio" }, + { name = "typing-extensions" }, + { name = "wrapt" }, ] sdist = { url = "https://files.pythonhosted.org/packages/53/a7/809482759f40079f4c4328c7318bf569ae25d457f5017aad30a1b9aafedc/aiologic-0.17.0.tar.gz", hash = "sha256:65aa058e858c94cd208badb188e7f00b54dcabb3ba85b34f794db98074d108b9", size = 251625, upload-time = "2026-06-14T12:24:35.367Z" } wheels = [ @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] @@ -519,14 +519,14 @@ name = "aurelio-sdk" version = "0.0.19" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiofiles", marker = "python_full_version < '3.14'" }, - { name = "aiohttp", marker = "python_full_version < '3.14'" }, - { name = "colorlog", marker = "python_full_version < '3.14'" }, - { name = "pydantic", marker = "python_full_version < '3.14'" }, - { name = "python-dotenv", marker = "python_full_version < '3.14'" }, - { name = "requests", marker = "python_full_version < '3.14'" }, - { name = "requests-toolbelt", marker = "python_full_version < '3.14'" }, - { name = "tornado", marker = "python_full_version < '3.14'" }, + { name = "aiofiles" }, + { name = "aiohttp" }, + { name = "colorlog" }, + { name = "pydantic" }, + { name = "python-dotenv" }, + { name = "requests" }, + { name = "requests-toolbelt" }, + { name = "tornado" }, ] sdist = { url = "https://files.pythonhosted.org/packages/27/0e/c2e369ad173fb3d76448e46d10beb3dcc53388318933ddf8169a3f21a810/aurelio_sdk-0.0.19.tar.gz", hash = "sha256:14107e7440ff2efd0b4a08c52fb595e7680bd4bc973a0ddfb3b64157c6666b91", size = 15258, upload-time = "2025-03-24T14:37:32.203Z" } wheels = [ @@ -538,9 +538,9 @@ name = "aws-sdk-bedrock-runtime" version = "0.11.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "smithy-aws-core", extra = ["eventstream", "json"], marker = "python_full_version >= '3.12'" }, - { name = "smithy-core", marker = "python_full_version >= '3.12'" }, - { name = "smithy-http", extra = ["aiohttp"], marker = "python_full_version >= '3.12'" }, + { name = "smithy-aws-core", extra = ["eventstream", "json"] }, + { name = "smithy-core" }, + { name = "smithy-http", extra = ["aiohttp"] }, ] sdist = { url = "https://files.pythonhosted.org/packages/8e/b3/9c225cbfe9f17ea2e3d75a0fdd0b325ef79839b9c09a376bda63a7bf3bb3/aws_sdk_bedrock_runtime-0.11.0.tar.gz", hash = "sha256:f2c45d34625bf6a7b56375e29a53a16b376880bda771e4bbf7d84491622eb193", size = 173854, upload-time = "2026-08-24T21:17:16.304Z" } wheels = [ @@ -549,7 +549,7 @@ wheels = [ [package.optional-dependencies] awscrt = [ - { name = "smithy-http", extra = ["awscrt"], marker = "python_full_version >= '3.12'" }, + { name = "smithy-http", extra = ["awscrt"] }, ] [[package]] @@ -1207,7 +1207,7 @@ name = "colorlog" version = "6.10.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "colorama", marker = "python_full_version < '3.14' and sys_platform == 'win32'" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/a2/61/f083b5ac52e505dfc1c624eafbf8c7589a0d7f32daa398d2e7590efa5fda/colorlog-6.10.1.tar.gz", hash = "sha256:eb4ae5cb65fe7fec7773c2306061a8e63e02efc2c72eba9d27b0fa23c94f1321", size = 17162, upload-time = "2025-10-16T16:14:11.978Z" } wheels = [ @@ -1231,7 +1231,7 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/66/54/eb9bfc647b19f2009dd5c7f5ec51c4e6ca831725f1aea7a993034f483147/contourpy-1.3.2.tar.gz", hash = "sha256:b6945942715a034c671b7fc54f9588126b0b8bf23db2696e3ca8328f3ff0ab54", size = 13466130, upload-time = "2025-04-15T17:47:53.79Z" } wheels = [ @@ -1304,7 +1304,7 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/58/01/1253e6698a07380cd31a736d248a3f2a50a7c88779a1813da27503cadc2a/contourpy-1.3.3.tar.gz", hash = "sha256:083e12155b210502d0bca491432bb04d56dc3432f95a979b429f2848c3dbe880", size = 13466174, upload-time = "2025-07-26T12:03:12.549Z" } @@ -1574,8 +1574,8 @@ name = "culsans" version = "0.11.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiologic", marker = "python_full_version < '3.13'" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "aiologic" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/d9/e3/49afa1bc180e0d28008ec6bcdf82a4072d1c7a41032b5b759b60814ca4b0/culsans-0.11.0.tar.gz", hash = "sha256:0b43d0d05dce6106293d114c86e3fb4bfc63088cfe8ff08ed3fe36891447fe33", size = 107546, upload-time = "2025-12-31T23:15:38.196Z" } wheels = [ @@ -1829,7 +1829,7 @@ name = "exceptiongroup" version = "1.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "typing-extensions", marker = "python_full_version < '3.11'" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" } wheels = [ @@ -2412,11 +2412,11 @@ resolution-markers = [ "python_full_version >= '3.14'", ] dependencies = [ - { name = "google-auth", marker = "python_full_version >= '3.14'" }, - { name = "googleapis-common-protos", marker = "python_full_version >= '3.14'" }, - { name = "proto-plus", marker = "python_full_version >= '3.14'" }, - { name = "protobuf", marker = "python_full_version >= '3.14'" }, - { name = "requests", marker = "python_full_version >= '3.14'" }, + { name = "google-auth" }, + { name = "googleapis-common-protos" }, + { name = "proto-plus" }, + { name = "protobuf" }, + { name = "requests" }, ] sdist = { url = "https://files.pythonhosted.org/packages/09/cd/63f1557235c2440fe0577acdbc32577c5c002684c58c7f4d770a92366a24/google_api_core-2.25.2.tar.gz", hash = "sha256:1c63aa6af0d0d5e37966f157a77f9396d820fba59f9e43e9415bc3dc5baff300", size = 166266, upload-time = "2025-10-03T00:07:34.778Z" } wheels = [ @@ -2425,8 +2425,8 @@ wheels = [ [package.optional-dependencies] grpc = [ - { name = "grpcio", marker = "python_full_version >= '3.14'" }, - { name = "grpcio-status", marker = "python_full_version >= '3.14'" }, + { name = "grpcio" }, + { name = "grpcio-status" }, ] [[package]] @@ -2440,11 +2440,11 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "google-auth", marker = "python_full_version < '3.14'" }, - { name = "googleapis-common-protos", marker = "python_full_version < '3.14'" }, - { name = "proto-plus", marker = "python_full_version < '3.14'" }, - { name = "protobuf", marker = "python_full_version < '3.14'" }, - { name = "requests", marker = "python_full_version < '3.14'" }, + { name = "google-auth" }, + { name = "googleapis-common-protos" }, + { name = "proto-plus" }, + { name = "protobuf" }, + { name = "requests" }, ] sdist = { url = "https://files.pythonhosted.org/packages/16/ce/502a57fb0ec752026d24df1280b162294b22a0afb98a326084f9a979138b/google_api_core-2.30.3.tar.gz", hash = "sha256:e601a37f148585319b26db36e219df68c5d07b6382cff2d580e83404e44d641b", size = 177001, upload-time = "2026-04-10T00:41:28.035Z" } wheels = [ @@ -2453,8 +2453,8 @@ wheels = [ [package.optional-dependencies] grpc = [ - { name = "grpcio", marker = "python_full_version < '3.14'" }, - { name = "grpcio-status", marker = "python_full_version < '3.14'" }, + { name = "grpcio" }, + { name = "grpcio-status" }, ] [[package]] @@ -2623,12 +2623,12 @@ resolution-markers = [ "python_full_version >= '3.14'", ] dependencies = [ - { name = "google-api-core", version = "2.25.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.14'" }, - { name = "google-auth", marker = "python_full_version >= '3.14'" }, - { name = "google-cloud-core", marker = "python_full_version >= '3.14'" }, - { name = "google-crc32c", marker = "python_full_version >= '3.14'" }, - { name = "google-resumable-media", marker = "python_full_version >= '3.14'" }, - { name = "requests", marker = "python_full_version >= '3.14'" }, + { name = "google-api-core", version = "2.25.2", source = { registry = "https://pypi.org/simple" } }, + { name = "google-auth" }, + { name = "google-cloud-core" }, + { name = "google-crc32c" }, + { name = "google-resumable-media" }, + { name = "requests" }, ] sdist = { url = "https://files.pythonhosted.org/packages/bd/ef/7cefdca67a6c8b3af0ec38612f9e78e5a9f6179dd91352772ae1a9849246/google_cloud_storage-3.4.1.tar.gz", hash = "sha256:6f041a297e23a4b485fad8c305a7a6e6831855c208bcbe74d00332a909f82268", size = 17238203, upload-time = "2025-10-08T18:43:39.665Z" } wheels = [ @@ -2646,12 +2646,12 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "google-api-core", version = "2.30.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14'" }, - { name = "google-auth", marker = "python_full_version < '3.14'" }, - { name = "google-cloud-core", marker = "python_full_version < '3.14'" }, - { name = "google-crc32c", marker = "python_full_version < '3.14'" }, - { name = "google-resumable-media", marker = "python_full_version < '3.14'" }, - { name = "requests", marker = "python_full_version < '3.14'" }, + { name = "google-api-core", version = "2.30.3", source = { registry = "https://pypi.org/simple" } }, + { name = "google-auth" }, + { name = "google-cloud-core" }, + { name = "google-crc32c" }, + { name = "google-resumable-media" }, + { name = "requests" }, ] sdist = { url = "https://files.pythonhosted.org/packages/4c/47/205eb8e9a1739b5345843e5a425775cbdc472cc38e7eda082ba5b8d02450/google_cloud_storage-3.10.1.tar.gz", hash = "sha256:97db9aa4460727982040edd2bd13ff3d5e2260b5331ad22895802da1fc2a5286", size = 17309950, upload-time = "2026-03-23T09:35:23.409Z" } wheels = [ @@ -4081,13 +4081,13 @@ name = "langchain-classic" version = "1.0.7" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "langchain-core", marker = "python_full_version >= '3.11'" }, - { name = "langchain-text-splitters", marker = "python_full_version >= '3.11'" }, - { name = "langsmith", marker = "python_full_version >= '3.11'" }, - { name = "pydantic", marker = "python_full_version >= '3.11'" }, - { name = "pyyaml", marker = "python_full_version >= '3.11'" }, - { name = "requests", marker = "python_full_version >= '3.11'" }, - { name = "sqlalchemy", marker = "python_full_version >= '3.11'" }, + { name = "langchain-core" }, + { name = "langchain-text-splitters" }, + { name = "langsmith" }, + { name = "pydantic" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "sqlalchemy" }, ] sdist = { url = "https://files.pythonhosted.org/packages/9b/78/84b5065816f348c39fefa4316f209f0135e8410216340a953bec17d9e4e4/langchain_classic-1.0.7.tar.gz", hash = "sha256:debbec8065e69b95108d2652e8d5c44f4516e19aa8d716c02ed2211c3aee099d", size = 10554118, upload-time = "2026-05-07T15:46:56.8Z" } wheels = [ @@ -4102,18 +4102,18 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "aiohttp", marker = "python_full_version < '3.11'" }, - { name = "dataclasses-json", marker = "python_full_version < '3.11'" }, - { name = "httpx-sse", marker = "python_full_version < '3.11'" }, - { name = "langchain", marker = "python_full_version < '3.11'" }, - { name = "langchain-core", marker = "python_full_version < '3.11'" }, - { name = "langsmith", marker = "python_full_version < '3.11'" }, - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "pydantic-settings", marker = "python_full_version < '3.11'" }, - { name = "pyyaml", marker = "python_full_version < '3.11'" }, - { name = "requests", marker = "python_full_version < '3.11'" }, - { name = "sqlalchemy", marker = "python_full_version < '3.11'" }, - { name = "tenacity", marker = "python_full_version < '3.11'" }, + { name = "aiohttp" }, + { name = "dataclasses-json" }, + { name = "httpx-sse" }, + { name = "langchain" }, + { name = "langchain-core" }, + { name = "langsmith" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } }, + { name = "pydantic-settings" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "sqlalchemy" }, + { name = "tenacity" }, ] sdist = { url = "https://files.pythonhosted.org/packages/83/49/2ff5354273809e9811392bc24bcffda545a196070666aef27bc6aacf1c21/langchain_community-0.3.31.tar.gz", hash = "sha256:250e4c1041539130f6d6ac6f9386cb018354eafccd917b01a4cff1950b80fd81", size = 33241237, upload-time = "2025-10-07T20:17:57.857Z" } wheels = [ @@ -4131,19 +4131,19 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "aiohttp", marker = "python_full_version >= '3.11'" }, - { name = "dataclasses-json", marker = "python_full_version >= '3.11'" }, - { name = "httpx-sse", marker = "python_full_version >= '3.11'" }, - { name = "langchain-classic", marker = "python_full_version >= '3.11'" }, - { name = "langchain-core", marker = "python_full_version >= '3.11'" }, - { name = "langsmith", marker = "python_full_version >= '3.11'" }, - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "aiohttp" }, + { name = "dataclasses-json" }, + { name = "httpx-sse" }, + { name = "langchain-classic" }, + { name = "langchain-core" }, + { name = "langsmith" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "pydantic-settings", marker = "python_full_version >= '3.11'" }, - { name = "pyyaml", marker = "python_full_version >= '3.11'" }, - { name = "requests", marker = "python_full_version >= '3.11'" }, - { name = "sqlalchemy", marker = "python_full_version >= '3.11'" }, - { name = "tenacity", marker = "python_full_version >= '3.11'" }, + { name = "pydantic-settings" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "sqlalchemy" }, + { name = "tenacity" }, ] sdist = { url = "https://files.pythonhosted.org/packages/53/97/a03585d42b9bdb6fbd935282d6e3348b10322a24e6ce12d0c99eb461d9af/langchain_community-0.4.1.tar.gz", hash = "sha256:f3b211832728ee89f169ddce8579b80a085222ddb4f4ed445a46e977d17b1e85", size = 33241144, upload-time = "2025-10-27T15:20:32.504Z" } wheels = [ @@ -4215,7 +4215,7 @@ name = "langchain-text-splitters" version = "1.1.2" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "langchain-core", marker = "python_full_version >= '3.11'" }, + { name = "langchain-core" }, ] sdist = { url = "https://files.pythonhosted.org/packages/26/9f/6c545900fefb7b00ddfa3f16b80d61338a0ec68c31c5451eeeab99082760/langchain_text_splitters-1.1.2.tar.gz", hash = "sha256:782a723db0a4746ac91e251c7c1d57fd23636e4f38ed733074e28d7a86f41627", size = 293580, upload-time = "2026-04-16T14:20:39.162Z" } wheels = [ @@ -4959,16 +4959,16 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "aiohttp", marker = "python_full_version < '3.11'" }, - { name = "chevron", marker = "python_full_version < '3.11'" }, - { name = "jsonpickle", marker = "python_full_version < '3.11'" }, - { name = "langchain-community", version = "0.3.31", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "packaging", marker = "python_full_version < '3.11'" }, - { name = "pydantic", marker = "python_full_version < '3.11'" }, - { name = "pyhumps", marker = "python_full_version < '3.11'" }, - { name = "requests", marker = "python_full_version < '3.11'" }, - { name = "setuptools", marker = "python_full_version < '3.11'" }, - { name = "tenacity", marker = "python_full_version < '3.11'" }, + { name = "aiohttp" }, + { name = "chevron" }, + { name = "jsonpickle" }, + { name = "langchain-community", version = "0.3.31", source = { registry = "https://pypi.org/simple" } }, + { name = "packaging" }, + { name = "pydantic" }, + { name = "pyhumps" }, + { name = "requests" }, + { name = "setuptools" }, + { name = "tenacity" }, ] sdist = { url = "https://files.pythonhosted.org/packages/4a/6f/9ca1acf766848aaf5f0ac4140c34c91ad0dbfad2654359699644be3352c9/lunary-1.4.36.tar.gz", hash = "sha256:53f002f385c83d9c0e6368e7999923acffbde987f53c5205c2c249c38ee2d75c", size = 20253, upload-time = "2026-02-09T20:49:30.56Z" } wheels = [ @@ -4986,16 +4986,16 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "aiohttp", marker = "python_full_version >= '3.11'" }, - { name = "chevron", marker = "python_full_version >= '3.11'" }, - { name = "jsonpickle", marker = "python_full_version >= '3.11'" }, - { name = "langchain-community", version = "0.4.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "packaging", marker = "python_full_version >= '3.11'" }, - { name = "pydantic", marker = "python_full_version >= '3.11'" }, - { name = "pyhumps", marker = "python_full_version >= '3.11'" }, - { name = "requests", marker = "python_full_version >= '3.11'" }, - { name = "setuptools", marker = "python_full_version >= '3.11'" }, - { name = "tenacity", marker = "python_full_version >= '3.11'" }, + { name = "aiohttp" }, + { name = "chevron" }, + { name = "jsonpickle" }, + { name = "langchain-community", version = "0.4.1", source = { registry = "https://pypi.org/simple" } }, + { name = "packaging" }, + { name = "pydantic" }, + { name = "pyhumps" }, + { name = "requests" }, + { name = "setuptools" }, + { name = "tenacity" }, ] sdist = { url = "https://files.pythonhosted.org/packages/37/ef/1acbc6957585cc0110e648d787663871717ced3df27fcd3cb5e18fa418f3/lunary-1.4.37.tar.gz", hash = "sha256:1781091e9dceffcc28ebc4be7e085c9fec4102d98d7ca945ed0021e9ce03c36f", size = 20248, upload-time = "2026-02-12T08:15:02.091Z" } wheels = [ @@ -8787,10 +8787,10 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "joblib", marker = "python_full_version < '3.11'" }, - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "threadpoolctl", marker = "python_full_version < '3.11'" }, + { name = "joblib" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } }, + { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" } }, + { name = "threadpoolctl" }, ] sdist = { url = "https://files.pythonhosted.org/packages/98/c2/a7855e41c9d285dfe86dc50b250978105dce513d6e459ea66a6aeb0e1e0c/scikit_learn-1.7.2.tar.gz", hash = "sha256:20e9e49ecd130598f1ca38a1d85090e1a600147b9c02fa6f15d69cb53d968fda", size = 7193136, upload-time = "2025-09-09T08:21:29.075Z" } wheels = [ @@ -8837,11 +8837,11 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "joblib", marker = "python_full_version >= '3.11'" }, - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "joblib" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, - { name = "threadpoolctl", marker = "python_full_version >= '3.11'" }, + { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" } }, + { name = "threadpoolctl" }, ] sdist = { url = "https://files.pythonhosted.org/packages/0e/d4/40988bf3b8e34feec1d0e6a051446b1f66225f8529b9309becaeef62b6c4/scikit_learn-1.8.0.tar.gz", hash = "sha256:9bccbb3b40e3de10351f8f5068e105d0f4083b1a65fa07b6634fbc401a6287fd", size = 7335585, upload-time = "2025-12-10T07:08:53.618Z" } wheels = [ @@ -8891,7 +8891,7 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" } }, ] sdist = { url = "https://files.pythonhosted.org/packages/0f/37/6964b830433e654ec7485e45a00fc9a27cf868d622838f6b6d9c5ec0d532/scipy-1.15.3.tar.gz", hash = "sha256:eae3cf522bc7df64b42cad3925c876e1b0b6c35c1337c93e12c0f366f55b0eaf", size = 59419214, upload-time = "2025-05-08T16:13:05.955Z" } wheels = [ @@ -8953,7 +8953,7 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/7a/97/5a3609c4f8d58b039179648e62dd220f89864f56f7357f5d4f45c29eb2cc/scipy-1.17.1.tar.gz", hash = "sha256:95d8e012d8cb8816c226aef832200b1d45109ed4464303e997c5b13122b297c0", size = 30573822, upload-time = "2026-02-23T00:26:24.851Z" } @@ -9038,20 +9038,20 @@ name = "semantic-router" version = "0.1.15" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aiohttp", marker = "python_full_version < '3.14'" }, - { name = "aurelio-sdk", marker = "python_full_version < '3.14'" }, - { name = "colorama", marker = "python_full_version < '3.14'" }, - { name = "colorlog", marker = "python_full_version < '3.14'" }, - { name = "litellm", marker = "python_full_version < '3.14'" }, - { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12'" }, + { name = "aiohttp" }, + { name = "aurelio-sdk" }, + { name = "colorama" }, + { name = "colorlog" }, + { name = "litellm" }, + { name = "numpy", version = "1.26.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.12' or python_full_version >= '3.14'" }, { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12' and python_full_version < '3.14'" }, - { name = "openai", marker = "python_full_version < '3.14'" }, - { name = "pydantic", marker = "python_full_version < '3.14'" }, - { name = "pyyaml", marker = "python_full_version < '3.14'" }, - { name = "regex", marker = "python_full_version < '3.14'" }, - { name = "tiktoken", marker = "python_full_version < '3.14'" }, - { name = "tornado", marker = "python_full_version < '3.14'" }, - { name = "urllib3", marker = "python_full_version < '3.14'" }, + { name = "openai" }, + { name = "pydantic" }, + { name = "pyyaml" }, + { name = "regex" }, + { name = "tiktoken" }, + { name = "tornado" }, + { name = "urllib3" }, ] sdist = { url = "https://files.pythonhosted.org/packages/dc/a9/1a689e916e8b280f1fd8fb335cc059be626a22fe4533baa045d32fcd6de5/semantic_router-0.1.15.tar.gz", hash = "sha256:328256ddc3c2b713101ec69561d6585aecbf1198ea3461e1486289d8c3a35288", size = 95605, upload-time = "2026-05-23T12:58:15.444Z" } wheels = [ @@ -9134,9 +9134,9 @@ name = "smithy-aws-core" version = "0.11.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "aws-sdk-signers", marker = "python_full_version >= '3.12'" }, - { name = "smithy-core", marker = "python_full_version >= '3.12'" }, - { name = "smithy-http", marker = "python_full_version >= '3.12'" }, + { name = "aws-sdk-signers" }, + { name = "smithy-core" }, + { name = "smithy-http" }, ] sdist = { url = "https://files.pythonhosted.org/packages/7d/d3/501c0023548173416109ac42298ca33b708469dc922005770811a597949f/smithy_aws_core-0.11.0.tar.gz", hash = "sha256:29ee89976a520a87e3db557e03e115fdc21a0a60b81161e95174395a1b064da1", size = 38791, upload-time = "2026-08-24T21:16:59.631Z" } wheels = [ @@ -9145,10 +9145,10 @@ wheels = [ [package.optional-dependencies] eventstream = [ - { name = "smithy-aws-event-stream", marker = "python_full_version >= '3.12'" }, + { name = "smithy-aws-event-stream" }, ] json = [ - { name = "smithy-json", marker = "python_full_version >= '3.12'" }, + { name = "smithy-json" }, ] [[package]] @@ -9156,7 +9156,7 @@ name = "smithy-aws-event-stream" version = "0.3.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "smithy-core", marker = "python_full_version >= '3.12'" }, + { name = "smithy-core" }, ] sdist = { url = "https://files.pythonhosted.org/packages/38/0e/6efb3a4ed92c0f1ada6de060ac92e7115a1e34d0ab1fb99a6056734a88ea/smithy_aws_event_stream-0.3.0.tar.gz", hash = "sha256:a0e227367a973144e205a075d0a424f95c92f26656a1018d08900da2ae547c49", size = 12818, upload-time = "2026-05-05T18:04:14.317Z" } wheels = [ @@ -9177,7 +9177,7 @@ name = "smithy-http" version = "0.5.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "smithy-core", marker = "python_full_version >= '3.12'" }, + { name = "smithy-core" }, ] sdist = { url = "https://files.pythonhosted.org/packages/98/78/b5f3113d6c8f0bc1f9777a7f5ca84b892d29efac05850e14f7d4f7e645b5/smithy_http-0.5.0.tar.gz", hash = "sha256:bb4a19672f7c7eeb872a308f777eb505281a5bafb1ee3d1ea9c760c06c352510", size = 31122, upload-time = "2026-08-24T21:16:56.488Z" } wheels = [ @@ -9186,11 +9186,11 @@ wheels = [ [package.optional-dependencies] aiohttp = [ - { name = "aiohttp", marker = "python_full_version >= '3.12'" }, - { name = "yarl", marker = "python_full_version >= '3.12'" }, + { name = "aiohttp" }, + { name = "yarl" }, ] awscrt = [ - { name = "awscrt", marker = "python_full_version >= '3.12'" }, + { name = "awscrt" }, ] [[package]] @@ -9198,8 +9198,8 @@ name = "smithy-json" version = "0.3.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "ijson", marker = "python_full_version >= '3.12'" }, - { name = "smithy-core", marker = "python_full_version >= '3.12'" }, + { name = "ijson" }, + { name = "smithy-core" }, ] sdist = { url = "https://files.pythonhosted.org/packages/c7/ac/04164eefb3da7479f52f6535b4b39cc8384c292cb2bb74279f2acc4f4b4d/smithy_json-0.3.0.tar.gz", hash = "sha256:c81c7034587e01bc64767cbbecb05a7d65ca9070612fd94e8a03e80540290a22", size = 7956, upload-time = "2026-08-20T17:55:32.177Z" } wheels = [ @@ -9277,23 +9277,23 @@ resolution-markers = [ "python_full_version < '3.11'", ] dependencies = [ - { name = "alabaster", marker = "python_full_version < '3.11'" }, - { name = "babel", marker = "python_full_version < '3.11'" }, - { name = "colorama", marker = "python_full_version < '3.11' and sys_platform == 'win32'" }, - { name = "docutils", version = "0.21.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, - { name = "imagesize", marker = "python_full_version < '3.11'" }, - { name = "jinja2", marker = "python_full_version < '3.11'" }, - { name = "packaging", marker = "python_full_version < '3.11'" }, - { name = "pygments", marker = "python_full_version < '3.11'" }, - { name = "requests", marker = "python_full_version < '3.11'" }, - { name = "snowballstemmer", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-applehelp", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-devhelp", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-htmlhelp", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-jsmath", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-qthelp", marker = "python_full_version < '3.11'" }, - { name = "sphinxcontrib-serializinghtml", marker = "python_full_version < '3.11'" }, - { name = "tomli", marker = "python_full_version < '3.11'" }, + { name = "alabaster" }, + { name = "babel" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "docutils", version = "0.21.2", source = { registry = "https://pypi.org/simple" } }, + { name = "imagesize" }, + { name = "jinja2" }, + { name = "packaging" }, + { name = "pygments" }, + { name = "requests" }, + { name = "snowballstemmer" }, + { name = "sphinxcontrib-applehelp" }, + { name = "sphinxcontrib-devhelp" }, + { name = "sphinxcontrib-htmlhelp" }, + { name = "sphinxcontrib-jsmath" }, + { name = "sphinxcontrib-qthelp" }, + { name = "sphinxcontrib-serializinghtml" }, + { name = "tomli" }, ] sdist = { url = "https://files.pythonhosted.org/packages/6f/6d/be0b61178fe2cdcb67e2a92fc9ebb488e3c51c4f74a36a7824c0adf23425/sphinx-8.1.3.tar.gz", hash = "sha256:43c1911eecb0d3e161ad78611bc905d1ad0e523e4ddc202a58a821773dc4c927", size = 8184611, upload-time = "2024-10-13T20:27:13.93Z" } wheels = [ @@ -9308,23 +9308,23 @@ resolution-markers = [ "python_full_version == '3.11.*'", ] dependencies = [ - { name = "alabaster", marker = "python_full_version == '3.11.*'" }, - { name = "babel", marker = "python_full_version == '3.11.*'" }, - { name = "colorama", marker = "python_full_version == '3.11.*' and sys_platform == 'win32'" }, - { name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, - { name = "imagesize", marker = "python_full_version == '3.11.*'" }, - { name = "jinja2", marker = "python_full_version == '3.11.*'" }, - { name = "packaging", marker = "python_full_version == '3.11.*'" }, - { name = "pygments", marker = "python_full_version == '3.11.*'" }, - { name = "requests", marker = "python_full_version == '3.11.*'" }, - { name = "roman-numerals", marker = "python_full_version == '3.11.*'" }, - { name = "snowballstemmer", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-applehelp", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-devhelp", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-htmlhelp", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-jsmath", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-qthelp", marker = "python_full_version == '3.11.*'" }, - { name = "sphinxcontrib-serializinghtml", marker = "python_full_version == '3.11.*'" }, + { name = "alabaster" }, + { name = "babel" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" } }, + { name = "imagesize" }, + { name = "jinja2" }, + { name = "packaging" }, + { name = "pygments" }, + { name = "requests" }, + { name = "roman-numerals" }, + { name = "snowballstemmer" }, + { name = "sphinxcontrib-applehelp" }, + { name = "sphinxcontrib-devhelp" }, + { name = "sphinxcontrib-htmlhelp" }, + { name = "sphinxcontrib-jsmath" }, + { name = "sphinxcontrib-qthelp" }, + { name = "sphinxcontrib-serializinghtml" }, ] sdist = { url = "https://files.pythonhosted.org/packages/42/50/a8c6ccc36d5eacdfd7913ddccd15a9cee03ecafc5ee2bc40e1f168d85022/sphinx-9.0.4.tar.gz", hash = "sha256:594ef59d042972abbc581d8baa577404abe4e6c3b04ef61bd7fc2acbd51f3fa3", size = 8710502, upload-time = "2025-12-04T07:45:27.343Z" } wheels = [ @@ -9341,23 +9341,23 @@ resolution-markers = [ "python_full_version == '3.12.*'", ] dependencies = [ - { name = "alabaster", marker = "python_full_version >= '3.12'" }, - { name = "babel", marker = "python_full_version >= '3.12'" }, - { name = "colorama", marker = "python_full_version >= '3.12' and sys_platform == 'win32'" }, - { name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, - { name = "imagesize", marker = "python_full_version >= '3.12'" }, - { name = "jinja2", marker = "python_full_version >= '3.12'" }, - { name = "packaging", marker = "python_full_version >= '3.12'" }, - { name = "pygments", marker = "python_full_version >= '3.12'" }, - { name = "requests", marker = "python_full_version >= '3.12'" }, - { name = "roman-numerals", marker = "python_full_version >= '3.12'" }, - { name = "snowballstemmer", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-applehelp", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-devhelp", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-htmlhelp", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-jsmath", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-qthelp", marker = "python_full_version >= '3.12'" }, - { name = "sphinxcontrib-serializinghtml", marker = "python_full_version >= '3.12'" }, + { name = "alabaster" }, + { name = "babel" }, + { name = "colorama", marker = "sys_platform == 'win32'" }, + { name = "docutils", version = "0.22.4", source = { registry = "https://pypi.org/simple" } }, + { name = "imagesize" }, + { name = "jinja2" }, + { name = "packaging" }, + { name = "pygments" }, + { name = "requests" }, + { name = "roman-numerals" }, + { name = "snowballstemmer" }, + { name = "sphinxcontrib-applehelp" }, + { name = "sphinxcontrib-devhelp" }, + { name = "sphinxcontrib-htmlhelp" }, + { name = "sphinxcontrib-jsmath" }, + { name = "sphinxcontrib-qthelp" }, + { name = "sphinxcontrib-serializinghtml" }, ] sdist = { url = "https://files.pythonhosted.org/packages/cd/bd/f08eb0f4eed5c83f1ba2a3bd18f7745a2b1525fad70660a1c00224ec468a/sphinx-9.1.0.tar.gz", hash = "sha256:7741722357dd75f8190766926071fed3bdc211c74dd2d7d4df5404da95930ddb", size = 8718324, upload-time = "2025-12-31T15:09:27.646Z" } wheels = [ @@ -9505,8 +9505,8 @@ name = "standard-aifc" version = "3.13.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "audioop-lts", marker = "python_full_version >= '3.13'" }, - { name = "standard-chunk", marker = "python_full_version >= '3.13'" }, + { name = "audioop-lts" }, + { name = "standard-chunk" }, ] sdist = { url = "https://files.pythonhosted.org/packages/c4/53/6050dc3dde1671eb3db592c13b55a8005e5040131f7509cef0215212cb84/standard_aifc-3.13.0.tar.gz", hash = "sha256:64e249c7cb4b3daf2fdba4e95721f811bde8bdfc43ad9f936589b7bb2fae2e43", size = 15240, upload-time = "2024-10-30T16:01:31.772Z" } wheels = [ @@ -9527,7 +9527,7 @@ name = "standard-sunau" version = "3.13.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "audioop-lts", marker = "python_full_version >= '3.13'" }, + { name = "audioop-lts" }, ] sdist = { url = "https://files.pythonhosted.org/packages/66/e3/ce8d38cb2d70e05ffeddc28bb09bad77cfef979eb0a299c9117f7ed4e6a9/standard_sunau-3.13.0.tar.gz", hash = "sha256:b319a1ac95a09a2378a8442f403c66f4fd4b36616d6df6ae82b8e536ee790908", size = 9368, upload-time = "2024-10-30T16:01:41.626Z" } wheels = [ @@ -9561,8 +9561,8 @@ name = "taskgroup" version = "0.2.2" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, - { name = "typing-extensions", marker = "python_full_version < '3.11'" }, + { name = "exceptiongroup" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/f0/8d/e218e0160cc1b692e6e0e5ba34e8865dbb171efeb5fc9a704544b3020605/taskgroup-0.2.2.tar.gz", hash = "sha256:078483ac3e78f2e3f973e2edbf6941374fbea81b9c5d0a96f51d297717f4752d", size = 11504, upload-time = "2025-01-03T09:24:13.761Z" } wheels = [ From 860b0203c0d97249607632cda73f9e6726e3428e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:17:34 -0700 Subject: [PATCH 147/253] test(bedrock): expect canonical session tags in the dynamic auth params propagation test --- ..._bedrock_dynamic_auth_params_unit_tests.py | 24 +++++++++---------- 1 file changed, 11 insertions(+), 13 deletions(-) diff --git a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py index 6c059423f74..2d1d2815026 100644 --- a/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py +++ b/tests/llm_translation/test_bedrock_dynamic_auth_params_unit_tests.py @@ -185,19 +185,19 @@ class DummyCredentials: ], ) @pytest.mark.parametrize( - "param_name, param_value", + "param_name, param_value, expected_credentials_value", [ - ("aws_session_token", "dummy_session_token"), - ("aws_session_name", "dummy_session_name"), - ("aws_profile_name", "dummy_profile_name"), - ("aws_role_name", "dummy_role_name"), - ("aws_web_identity_token", "dummy_web_identity_token"), - ("aws_sts_endpoint", "dummy_sts_endpoint"), - ("aws_external_id", "dummy_external_id"), - ("aws_session_tags", [{"Key": "team", "Value": "genai"}]), + ("aws_session_token", "dummy_session_token", "dummy_session_token"), + ("aws_session_name", "dummy_session_name", "dummy_session_name"), + ("aws_profile_name", "dummy_profile_name", "dummy_profile_name"), + ("aws_role_name", "dummy_role_name", "dummy_role_name"), + ("aws_web_identity_token", "dummy_web_identity_token", "dummy_web_identity_token"), + ("aws_sts_endpoint", "dummy_sts_endpoint", "dummy_sts_endpoint"), + ("aws_external_id", "dummy_external_id", "dummy_external_id"), + ("aws_session_tags", [{"Key": "team", "Value": "genai"}], ({"Key": "team", "Value": "genai"},)), ], ) -def test_dynamic_aws_params_propagation(model, param_name, param_value): +def test_dynamic_aws_params_propagation(model, param_name, param_value, expected_credentials_value): """ When passed to litellm.completion, each dynamic AWS authentication parameter should propagate down to the get_credentials() call in BaseAWSLLM. @@ -282,6 +282,4 @@ def test_dynamic_aws_params_propagation(model, param_name, param_value): ) # We now assert that get_credentials() was called with the dynamic param. - assert ( - dummy_get_credentials.called_kwargs.get(param_name) == param_value - ) + assert dummy_get_credentials.called_kwargs.get(param_name) == expected_credentials_value From 3219a875d0dc350ba129b881fa1d9dc003ba5b17 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 18:20:18 +0000 Subject: [PATCH 148/253] fix(proxy): serialize in-memory window rollover and reset siblings once per expired window Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../hooks/parallel_request_limiter_v3.py | 61 +++++++------- .../hooks/test_parallel_request_limiter_v3.py | 82 +++++++++++++++++++ 2 files changed, 115 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 57b7d4f249d..cdec5922fff 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -100,7 +100,6 @@ BATCH_RATE_LIMITER_SCRIPT: Final = """ local results = {} local now = tonumber(ARGV[1]) local window_size = tonumber(ARGV[2]) -local reset_windows = {} -- Process each window/counter pair for i = 1, #KEYS, 2 do @@ -112,11 +111,8 @@ for i = 1, #KEYS, 2 do local window_start = redis.call('GET', window_key) if not window_start or (now - tonumber(window_start)) >= window_size then -- Reset window and counter - if not reset_windows[window_key] then - local prefix = string.sub(window_key, 1, -(#':window') - 1) - redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') - reset_windows[window_key] = true - end + local prefix = string.sub(window_key, 1, -(#':window') - 1) + redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') redis.call('SET', window_key, tostring(now)) redis.call('SET', counter_key, increment_value) redis.call('EXPIRE', window_key, window_size) @@ -1035,8 +1031,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): Implement sliding window rate limiting logic using in-memory cache operations. This follows the same logic as the Redis Lua script but uses async cache operations. """ + async with self._check_and_increment_lock: + return await self._in_memory_cache_sliding_window(keys=keys, now_int=now_int, window_size=window_size) + + async def _in_memory_cache_sliding_window( + self, + keys: list[str], + now_int: int, + window_size: int, + ) -> CacheCounterValues: results: Final[list[CacheCounterValue | None]] = [] - reset_windows: Final[set[str]] = set() # mutable-ok: tracks windows reset during this call # Process each window/counter pair for i in range(0, len(keys), 2): @@ -1054,16 +1058,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Check if window exists and is valid if window_start is None or (now_int - int(window_start)) >= window_size: # Reset window and counter - if window_key not in reset_windows: - for sibling_counter_key in _sibling_counter_keys(window_key): - await self.internal_usage_cache.async_set_cache( - key=sibling_counter_key, - value=0, - ttl=window_size, - litellm_parent_otel_span=None, - local_only=True, - ) - reset_windows.add(window_key) + for sibling_counter_key in _sibling_counter_keys(window_key): + await self.internal_usage_cache.async_set_cache( + key=sibling_counter_key, + value=0, + ttl=window_size, + litellm_parent_otel_span=None, + local_only=True, + ) await self.internal_usage_cache.async_set_cache( key=window_key, value=str(now_int), @@ -2076,21 +2078,24 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) # Pass 2: apply increments. + expired_windows: Final[Mapping[str, int]] = { + meta["window_key"]: meta["window_size"] + for meta, state in zip(per_counter_meta, descriptor_state) + if state["window_expired"] + } + for window_key, window_size in expired_windows.items(): + for sibling_counter_key in _sibling_counter_keys(window_key): + await self.internal_usage_cache.async_set_cache( + key=sibling_counter_key, + value=0, + ttl=window_size, + litellm_parent_otel_span=parent_otel_span, + local_only=True, + ) statuses: Final[list[RateLimitStatus]] = [] - reset_windows: Final[set[str]] = set() # mutable-ok: tracks windows reset during this call for meta, state in zip(per_counter_meta, descriptor_state): new_counter = meta["increment"] if state["window_expired"] else state["current"] + meta["increment"] if state["window_expired"]: - if meta["window_key"] not in reset_windows: - for sibling_counter_key in _sibling_counter_keys(meta["window_key"]): - await self.internal_usage_cache.async_set_cache( - key=sibling_counter_key, - value=0, - ttl=meta["window_size"], - litellm_parent_otel_span=parent_otel_span, - local_only=True, - ) - reset_windows.add(meta["window_key"]) await self.internal_usage_cache.async_set_cache( key=meta["window_key"], value=str(now_int), diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 0b7710eda70..5889b1b513f 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -18,11 +18,13 @@ from fastapi import HTTPException import litellm from litellm import Router from litellm.caching.caching import DualCache +from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, + RateLimitDescriptor, RequestRateLimiterStash, _request_stash, get_or_create_request_stash, @@ -5671,6 +5673,86 @@ async def test_tpm_reservation_resets_sibling_tokens_with_request_window(monkeyp assert await local_cache.async_get_cache(key=tokens_key) == 300 +@pytest.mark.asyncio +async def test_atomic_tpm_reservation_rollover_resets_sibling_requests_counter(): + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + window_size = 60 + now_int = int(time.time()) + window_key = "{api_key:atomic-rollover}:window" + requests_key = handler.create_rate_limit_keys("api_key", "atomic-rollover", "requests") + tokens_key = handler.create_rate_limit_keys("api_key", "atomic-rollover", "tokens") + for key, value in ((window_key, str(now_int - window_size - 1)), (requests_key, 3), (tokens_key, 900)): + await local_cache.async_set_cache(key=key, value=value, ttl=window_size) + + tpm_pass = await handler.atomic_check_and_increment_by_n( + descriptors=[ + RateLimitDescriptor( + key="api_key", + value="atomic-rollover", + rate_limit={"tokens_per_unit": 1000, "window_size": window_size}, + ) + ], + increments=[{"tokens": 200}], + ) + assert tpm_pass["overall_code"] == "OK" + assert await local_cache.async_get_cache(key=tokens_key) == 200 + + rpm_pass = await handler.should_rate_limit( + descriptors=[ + RateLimitDescriptor( + key="api_key", + value="atomic-rollover", + rate_limit={"requests_per_unit": 5, "window_size": window_size}, + ) + ], + skip_tpm_check=True, + ) + assert rpm_pass["overall_code"] == "OK" + assert [status["limit_remaining"] for status in rpm_pass["statuses"]] == [4] + assert await local_cache.async_get_cache(key=requests_key) == 1 + + +class _YieldingInMemoryCache(InMemoryCache): + async def async_get_cache(self, key: str, **kwargs: object) -> object: + value = await super().async_get_cache(key, **kwargs) + await asyncio.sleep(0) + return value + + +@pytest.mark.asyncio +async def test_window_rollover_reset_does_not_erase_concurrent_sibling_increment(): + local_cache = DualCache(in_memory_cache=_YieldingInMemoryCache()) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + window_size = 60 + now_int = int(time.time()) + window_key = "{api_key:concurrent-rollover}:window" + requests_key = handler.create_rate_limit_keys("api_key", "concurrent-rollover", "requests") + tokens_key = handler.create_rate_limit_keys("api_key", "concurrent-rollover", "tokens") + for key, value in ((window_key, str(now_int - window_size - 1)), (requests_key, 3), (tokens_key, 900)): + await local_cache.async_set_cache(key=key, value=value, ttl=window_size) + + tpm_descriptor = RateLimitDescriptor( + key="api_key", + value="concurrent-rollover", + rate_limit={"tokens_per_unit": 1000, "window_size": window_size}, + ) + rpm_pass, tpm_pass = await asyncio.gather( + handler.in_memory_cache_sliding_window( + keys=[window_key, requests_key], now_int=now_int, window_size=window_size + ), + handler.atomic_check_and_increment_by_n( + descriptors=[tpm_descriptor], + increments=[{"requests": 0, "tokens": 200}], + ), + ) + + assert rpm_pass == [str(now_int), 1] + assert tpm_pass["overall_code"] == "OK" + assert await local_cache.async_get_cache(key=requests_key) == 1 + assert await local_cache.async_get_cache(key=tokens_key) == 200 + + @pytest.mark.asyncio @pytest.mark.parametrize( "key_metadata, team_metadata, expected_output_estimate, tier", From a4da989aa9eac99ce0b8fb12cd1f1fdf1109e90a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:42:52 -0700 Subject: [PATCH 149/253] test(bedrock): type the STS recording helper --- .../llms/bedrock/test_base_aws_llm.py | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index c91bc604b06..6b9450afed4 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -10,6 +10,7 @@ from fastapi.testclient import TestClient +from collections.abc import Callable from datetime import datetime, timedelta, timezone from typing import Any, Dict, Optional from unittest.mock import MagicMock, patch @@ -3557,14 +3558,14 @@ def test_run_aws_signing_leaves_the_default_executor_free_for_other_providers(): assert signing_thread.startswith("aws-signing") -def _recording_boto3_client(recorded: Dict[str, Any]): +def _recording_boto3_client(recorded: dict[str, dict[str, object]]) -> Callable[..., MagicMock]: """boto3.client replacement that records the STS client kwargs and the assume-role params.""" - def _client(service_name, **client_kwargs): + def _client(service_name: str, **client_kwargs: object) -> MagicMock: recorded["client_kwargs"] = client_kwargs sts = MagicMock() - def _assume(**params): + def _assume(**params: object) -> dict[str, object]: recorded["assume_role"] = params return { "Credentials": { @@ -3575,7 +3576,7 @@ def _recording_boto3_client(recorded: Dict[str, Any]): } } - def _assume_web_identity(**params): + def _assume_web_identity(**params: object) -> dict[str, object]: recorded["assume_role_with_web_identity"] = params return { "Credentials": { @@ -3608,7 +3609,7 @@ def test_resolve_credentials_forwards_static_keys_role_session_and_external_id() aws_sts_endpoint="https://custom-sts.example", aws_session_tags=[{"Key": "team", "Value": "genai"}, {"Key": "cost-center", "Value": "42"}], ) - recorded: Dict[str, Any] = {} + recorded: dict[str, dict[str, object]] = {} with ( patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), @@ -3648,7 +3649,7 @@ def test_resolve_credentials_rejects_malformed_session_tags(malformed_tags): aws_session_name="litellm-session", aws_session_tags=malformed_tags, ) - recorded: Dict[str, Any] = {} + recorded: dict[str, dict[str, object]] = {} with ( patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), @@ -3669,7 +3670,7 @@ def test_resolve_credentials_forwards_web_identity_token(): aws_role_name="arn:aws:iam::123456789012:role/litellm-wif", aws_session_name="litellm-wif-session", ) - recorded: Dict[str, Any] = {} + recorded: dict[str, dict[str, object]] = {} with ( patch.dict(os.environ, _os_environ_without_aws_keys(), clear=True), From cf08440557bd0a377f1ce082497d037294c6aca0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:42:52 -0700 Subject: [PATCH 150/253] fix(azure-ai): price FLUX.2 Flex by its resolved cost map key --- .../azure_ai/image_generation/cost_calculator.py | 5 +---- .../test_azure_ai_flux2_image_generation.py | 13 +++++++++++++ 2 files changed, 14 insertions(+), 4 deletions(-) diff --git a/litellm/llms/azure_ai/image_generation/cost_calculator.py b/litellm/llms/azure_ai/image_generation/cost_calculator.py index 086293f26c0..35d0f4fb6c3 100644 --- a/litellm/llms/azure_ai/image_generation/cost_calculator.py +++ b/litellm/llms/azure_ai/image_generation/cost_calculator.py @@ -42,9 +42,6 @@ def cost_calculator( if input_cost_per_pixel: from litellm.cost_calculator import default_image_cost_calculator - cost_model: Final = ( - model if model.startswith(f"{litellm.LlmProviders.AZURE_AI.value}/") else f"azure_ai/{model}" - ) width: Final = optional_params.get("width") if optional_params else None height: Final = optional_params.get("height") if optional_params else None pixel_size: Final = ( @@ -53,7 +50,7 @@ def cost_calculator( else size or image_response.size ) return default_image_cost_calculator( - model=cost_model, + model=_model_info["key"], custom_llm_provider=litellm.LlmProviders.AZURE_AI.value, size=pixel_size, n=num_images, diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py index c1d46b7e919..388cca9f054 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py @@ -172,6 +172,19 @@ def test_flux2_cost_uses_mapped_dimensions_after_response_transformation(dimensi ) == pytest.approx(5e-08 * 2048 * 1024 * 2) +def test_flux2_flex_cost_accepts_lowercase_model_spelling(): + response: Final = ImageResponse(data=[ImageObject(b64_json="aW1n"), ImageObject(b64_json="aW1n")]) + + cost: Final = litellm.completion_cost( + model="azure_ai/flux.2-flex", + completion_response=response, + optional_params={"width": 1536, "height": 1024, "num_images": 2}, + call_type="image_generation", + ) + + assert cost == pytest.approx(5e-08 * 1536 * 1024 * 2) + + def test_flux2_response_preserves_mapped_dimensions(): config = AzureFoundryFluxImageGenerationConfig() params = config.map_openai_params( From 307df09792c1088a2d7b63e424c8040fa7746595 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 18:51:43 +0000 Subject: [PATCH 151/253] test(mcp): assert the self-revoke response instead of echoing the delete mock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_mcp_management_endpoints.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index a69cfe0ac7d..afadd6f3d19 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -24,6 +24,7 @@ from litellm.proxy._types import ( LiteLLM_MCPServerTable, LitellmUserRoles, MCPTransport, + MCPUserCredentialResponse, NewMCPServerRequest, UpdateMCPServerRequest, UserAPIKeyAuth, @@ -5220,7 +5221,11 @@ async def test_user_naming_themselves_still_deletes_own_byok_credential(): delete_mcp_user_credential, ) - delete_mock = AsyncMock(return_value=None) + deleted_rows: list[tuple[str, str]] = [] # mutable-ok: test-local recorder for the fake delete boundary + + async def _fake_delete_user_credential(_prisma_client: object, user_id: str, server_id: str) -> None: + deleted_rows.append((user_id, server_id)) + with ( patch( # test-quality-ok: endpoint test stubs the Prisma client lookup "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", @@ -5228,19 +5233,20 @@ async def test_user_naming_themselves_still_deletes_own_byok_credential(): ), patch( # test-quality-ok: endpoint test stubs the credential row delete "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", - new=delete_mock, + new=_fake_delete_user_credential, ), patch.object( # test-quality-ok: the cache invalidator is module scoped; the suite's only seam mcp_server, "_invalidate_byok_cred_cache", new=AsyncMock() ), ): - await delete_mcp_user_credential( + result = await delete_mcp_user_credential( server_id="srv-byok-self", user_api_key_dict=_make_user_auth("user-self"), user_id="user-self", ) - assert delete_mock.await_args.args[1:] == ("user-self", "srv-byok-self") + assert deleted_rows == [("user-self", "srv-byok-self")] + assert result == MCPUserCredentialResponse(server_id="srv-byok-self", has_credential=False) @pytest.mark.asyncio From 73fddb999e3849ffa31897426a82353ac3705521 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:53:26 -0700 Subject: [PATCH 152/253] fix(router): resolve stream_timeout before generic timeouts on the passthrough route --- litellm/passthrough/timeout_utils.py | 56 +++++++++-------- litellm/router.py | 4 +- .../test_pass_through_endpoints.py | 38 ++++++------ tests/test_litellm/test_router.py | 61 +++++++++---------- 4 files changed, 80 insertions(+), 79 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index fb649a9eeaf..0170d93c156 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -1,9 +1,20 @@ import sys from typing import Final +from pydantic import BaseModel, ConfigDict + DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0 +class _TimeoutFields(BaseModel): + model_config = ConfigDict(frozen=True) + + stream: bool = False + stream_timeout: float | None = None + timeout: float | None = None + request_timeout: float | None = None + + def resolve_pass_through_request_timeout( endpoint_timeout: float | None = None, ) -> float: @@ -33,8 +44,8 @@ def resolve_pass_through_request_timeout( def resolve_llm_passthrough_timeout( kwargs: dict | None = None, litellm_params: dict | None = None, - router_timeout: float | None = None, - router_stream_timeout: float | None = None, + router_timeout: float | str | None = None, + router_stream_timeout: float | str | None = None, ) -> float: """ Resolve upstream httpx timeout for SDK native passthrough (e.g. Bedrock /converse, @@ -44,28 +55,23 @@ def resolve_llm_passthrough_timeout( timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout -> 600s default. - Streaming (``kwargs["stream"]`` truthy) additionally consults ``stream_timeout`` at each - level before the non-streaming key, matching ``Router._get_stream_timeout`` on the - completion route: kwargs stream_timeout -> kwargs timeout/request_timeout -> - litellm_params stream_timeout -> litellm_params timeout/request_timeout -> - router_stream_timeout -> router_timeout -> pass_through_request_timeout -> 600s. + Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before + any generic timeout, matching ``Router._get_stream_timeout`` on the completion route: + kwargs stream_timeout -> litellm_params stream_timeout -> router_stream_timeout, then the + non-streaming chain above. """ - kwargs = kwargs or {} - litellm_params = litellm_params or {} - is_stream: Final[bool] = bool(kwargs.get("stream", False)) - - keys: Final[tuple[str, ...]] = ( - ("stream_timeout", "timeout", "request_timeout") if is_stream else ("timeout", "request_timeout") + request: Final = _TimeoutFields.model_validate(kwargs or {}) + deployment: Final = _TimeoutFields.model_validate(litellm_params or {}) + stream_candidates: Final = ( + (request.stream_timeout, deployment.stream_timeout, router_stream_timeout) if request.stream else () ) - for source in (kwargs, litellm_params): - for key in keys: - val = source.get(key) - if val is not None: - return float(val) - - if is_stream and router_stream_timeout is not None: - return float(router_stream_timeout) - if router_timeout is not None: - return float(router_timeout) - - return resolve_pass_through_request_timeout() + candidates: Final = ( + *stream_candidates, + request.timeout, + request.request_timeout, + deployment.timeout, + deployment.request_timeout, + router_timeout, + ) + resolved: Final = next((float(val) for val in candidates if val is not None), None) + return resolved if resolved is not None else resolve_pass_through_request_timeout() diff --git a/litellm/router.py b/litellm/router.py index 7ce7ba30502..a5523d7af79 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3880,7 +3880,9 @@ class Router: float(self._explicit_timeout) if isinstance(self._explicit_timeout, (int, float)) else None ) _router_stream_timeout: Final = ( - float(self.stream_timeout) if isinstance(self.stream_timeout, (int, float)) else None + self.stream_timeout + if self.stream_timeout is not None + else self.default_litellm_params.get("stream_timeout") ) kwargs["timeout"] = resolve_llm_passthrough_timeout( kwargs=kwargs, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 7f8663ea860..54336800db4 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1120,7 +1120,6 @@ def test_resolve_llm_passthrough_timeout_precedence(): def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): - # streaming: stream_timeout wins at each level, then falls through to the non-stream keys assert ( resolve_llm_passthrough_timeout( kwargs={"stream": True, "stream_timeout": 1800, "timeout": 45}, @@ -1129,36 +1128,35 @@ def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): ) assert ( resolve_llm_passthrough_timeout( - kwargs={"stream": True}, + kwargs={"stream": True, "timeout": 45}, litellm_params={"stream_timeout": 1800, "timeout": 90}, ) == 1800.0 ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True, "timeout": 45}, + litellm_params={"timeout": 90}, + router_timeout=120, + router_stream_timeout=1800, + ) + == 1800.0 + ) + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": True}, + router_stream_timeout="1800", + ) + == 1800.0 + ) assert ( resolve_llm_passthrough_timeout( kwargs={"stream": True}, litellm_params={"timeout": 90}, - router_stream_timeout=1800, + router_timeout=120, ) == 90.0 ) - assert ( - resolve_llm_passthrough_timeout( - kwargs={"stream": True}, - router_timeout=120, - router_stream_timeout=1800, - ) - == 1800.0 - ) - assert ( - resolve_llm_passthrough_timeout( - kwargs={"stream": True}, - router_timeout=120, - ) - == 120.0 - ) - - # non-streaming: stream_timeout is ignored everywhere assert ( resolve_llm_passthrough_timeout( kwargs={"stream": False, "stream_timeout": 1800}, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 4c39f8ba4e4..e92e5de23dc 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -5480,13 +5480,17 @@ def test_update_kwargs_with_deployment_uses_pass_through_request_timeout(): assert kwargs["timeout"] == 6.0 +def _passthrough_timeout(router: litellm.Router, deployment: dict, stream: bool) -> float: + kwargs: Final[dict] = {"stream": stream} + router._update_kwargs_with_deployment( + deployment=deployment, + kwargs=kwargs, + function_name="_ageneric_api_call_with_fallbacks", + ) + return kwargs["timeout"] + + def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): - """ - The SDK-native passthrough route (anthropic /v1/messages, bedrock /converse) resolves - its upstream timeout separately from the completion route. A streaming call must get - stream_timeout (deployment litellm_params first, then router_settings), while a - non-streaming call on the same deployment keeps the non-stream resolution. - """ router = litellm.Router( model_list=[ { @@ -5503,6 +5507,7 @@ def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): "litellm_params": { "model": "anthropic/claude-sonnet-4-5", "api_key": "fake-key", + "timeout": 60, }, }, ], @@ -5511,37 +5516,27 @@ def test_update_kwargs_with_deployment_passthrough_honors_stream_timeout(): ) per_deployment, router_default = router.model_list - kwargs: dict = {"stream": True} - router._update_kwargs_with_deployment( - deployment=per_deployment, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 1800.0 + assert _passthrough_timeout(router, per_deployment, stream=True) == 1800.0 + assert _passthrough_timeout(router, router_default, stream=True) == 900.0 + assert _passthrough_timeout(router, per_deployment, stream=False) == 60.0 + assert _passthrough_timeout(router, router_default, stream=False) == 60.0 - kwargs = {"stream": True} - router._update_kwargs_with_deployment( - deployment=router_default, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 900.0 - kwargs = {"stream": False} - router._update_kwargs_with_deployment( - deployment=per_deployment, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", +def test_update_kwargs_with_deployment_passthrough_router_stream_timeout_sources(): + deployment: Final[dict] = { + "model_name": "anthropic-router-default", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5", "api_key": "fake-key"}, + } + string_router = litellm.Router(model_list=[deployment], timeout=120, stream_timeout="900") + default_router = litellm.Router( + model_list=[deployment], + timeout=120, + default_litellm_params={"stream_timeout": 700}, ) - assert kwargs["timeout"] == 60.0 - kwargs = {"stream": False} - router._update_kwargs_with_deployment( - deployment=router_default, - kwargs=kwargs, - function_name="_ageneric_api_call_with_fallbacks", - ) - assert kwargs["timeout"] == 120.0 + assert _passthrough_timeout(string_router, string_router.model_list[0], stream=True) == 900.0 + assert _passthrough_timeout(default_router, default_router.model_list[0], stream=True) == 700.0 + assert _passthrough_timeout(default_router, default_router.model_list[0], stream=False) == 120.0 @pytest.mark.asyncio From c1e39810ecbcc3def67502932273b5e7a0bc94a8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 11:58:04 -0700 Subject: [PATCH 153/253] fix(mistral): send reasoning_effort as a level the model accepts Declare the live-verified reasoning_effort_levels on the Mistral cost-map entries and round an undeclared request to the nearest declared level (up to the weakest level at least as strong, down to the strongest when the request exceeds the ceiling). Codex's default medium no longer 400s on mistral-medium-latest, mistral-small-latest, or the vibe-cli family; an entry that declares nothing keeps forwarding the value verbatim --- litellm/llms/mistral/chat/transformation.py | 19 ++++- ...odel_prices_and_context_window_backup.json | 77 +++++++++++++++++++ .../reasoning_effort_capability.py | 17 ++++ model_prices_and_context_window.json | 77 +++++++++++++++++++ .../test_mistral_chat_transformation.py | 44 +++++++++-- .../test_reasoning_effort_capability.py | 20 +++++ 6 files changed, 246 insertions(+), 8 deletions(-) diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 6316128e6fa..8cbe4291054 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, cast, get_type_hints, ove import httpx +from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.prompt_templates.common_utils import ( handle_messages_with_content_list_to_str_conversion, @@ -20,6 +21,10 @@ from litellm.llms.openai.chat.gpt_transformation import ( OpenAIChatCompletionStreamingHandler, OpenAIGPTConfig, ) +from litellm.router_utils.reasoning_effort_capability import ( + declared_reasoning_efforts_for_model, + nearest_declared_reasoning_effort, +) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.mistral import MistralThinkingBlock, MistralToolCallMessage from litellm.types.llms.openai import AllMessageValues @@ -30,6 +35,18 @@ if TYPE_CHECKING: import tiktoken +def _accepted_reasoning_effort(model: str, requested: str) -> str: + declared: Final = declared_reasoning_efforts_for_model(model, "mistral") + if declared is None: + return requested + accepted: Final = nearest_declared_reasoning_effort(requested, declared) + if accepted != requested: + verbose_logger.debug( + "mistral: %s takes reasoning_effort %s, sending %s in place of %s", model, declared, accepted, requested + ) + return accepted + + class MistralConfig(OpenAIGPTConfig): """ Reference: https://docs.mistral.ai/api/ @@ -170,7 +187,7 @@ class MistralConfig(OpenAIGPTConfig): if param == "response_format": optional_params["response_format"] = value if param == "reasoning_effort" and "magistral" not in model.lower(): - optional_params["reasoning_effort"] = value + optional_params["reasoning_effort"] = _accepted_reasoning_effort(model, value) if param in ("reasoning_effort", "thinking") and "magistral" in model.lower(): # Flag that we need to add reasoning system prompt optional_params["_add_reasoning_prompt"] = True diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 0826f05cb11..f8707f9e403 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -36993,6 +36993,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37072,6 +37076,15 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "none", + "minimal", + "low", + "medium", + "high", + "xhigh", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-2", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37089,6 +37102,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37106,6 +37124,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37123,6 +37146,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37140,6 +37168,15 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "none", + "minimal", + "low", + "medium", + "high", + "xhigh", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-2", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37435,6 +37472,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37495,6 +37536,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37512,6 +37557,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37545,6 +37594,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37575,6 +37628,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -59375,6 +59432,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62700,6 +62761,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62717,6 +62782,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62734,6 +62803,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62751,6 +62824,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index 7b145c15a07..f880c846500 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -103,6 +103,23 @@ def declared_reasoning_efforts_for_model(model: str, custom_llm_provider: str) - return declared_reasoning_efforts(entry) +REASONING_EFFORT_STRENGTH_ORDER: Final = ("none", "minimal", "low", "medium", "high", "xhigh", "max") +_STRENGTH_RANK: Final = MappingProxyType({effort: rank for rank, effort in enumerate(REASONING_EFFORT_STRENGTH_ORDER)}) + + +def nearest_declared_reasoning_effort(requested: str, declared: Sequence[str]) -> str: + """Rounds a request up to the weakest declared level at least as strong as it, and down to the + strongest declared level when it asks for more than the model has, so the caller gets no less + reasoning than it asked for instead of a rejected call. A level outside the strength order is + returned as is for upstream to judge.""" + ranked: Final = sorted( + (effort for effort in declared if effort in _STRENGTH_RANK), key=lambda effort: _STRENGTH_RANK[effort] + ) + if requested in ranked or requested not in _STRENGTH_RANK or not ranked: + return requested + return next((effort for effort in ranked if _STRENGTH_RANK[effort] >= _STRENGTH_RANK[requested]), ranked[-1]) + + def _supports_none_reasoning_effort(model_info: Mapping[str, object], flag: object) -> bool: """Opt-in only where a request path refuses the level. AzureOpenAIGPT5Config raises UnsupportedParamsError on reasoning_effort='none' without an explicit true, and it is selected diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 0826f05cb11..f8707f9e403 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -36993,6 +36993,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37072,6 +37076,15 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "none", + "minimal", + "low", + "medium", + "high", + "xhigh", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-2", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37089,6 +37102,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37106,6 +37124,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37123,6 +37146,11 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "low", + "high", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-3", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37140,6 +37168,15 @@ "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, + "reasoning_effort_levels": [ + "none", + "minimal", + "low", + "medium", + "high", + "xhigh", + "max" + ], "source": "https://docs.mistral.ai/models/zai-glm-5-2", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37435,6 +37472,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37495,6 +37536,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37512,6 +37557,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37545,6 +37594,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -37575,6 +37628,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -59375,6 +59432,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62700,6 +62761,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62717,6 +62782,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62734,6 +62803,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 7.5e-06, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", "supports_assistant_prefill": true, "supports_function_calling": true, @@ -62751,6 +62824,10 @@ "max_tokens": 262144, "mode": "chat", "output_cost_per_token": 6e-07, + "reasoning_effort_levels": [ + "none", + "high" + ], "source": "https://docs.mistral.ai/models/model-cards/mistral-small-4-0-26-03", "supports_assistant_prefill": true, "supports_function_calling": true, diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index 38639f23050..a68ba9f570c 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -63,7 +63,7 @@ class TestMistralReasoningSupport: assert "reasoning_effort" not in supported_params_normal assert "thinking" not in supported_params_normal - def test_map_openai_params_reasoning_effort(self): + def test_map_openai_params_reasoning_effort(self, local_model_cost_map): """Test that reasoning_effort parameter is properly mapped for magistral models.""" mistral_config = MistralConfig() @@ -87,21 +87,51 @@ class TestMistralReasoningSupport: ) assert "_add_reasoning_prompt" not in result_normal - assert result_normal["reasoning_effort"] == "low" + assert result_normal["reasoning_effort"] == "high" @pytest.mark.parametrize( - ("model", "reasoning_effort"), - [("mistral-medium-latest", "high"), ("zai-glm-5-2", "xhigh")], + ("model", "requested", "sent"), + [ + ("mistral-medium-latest", "high", "high"), + ("mistral-medium-latest", "none", "none"), + ("mistral-medium-latest", "low", "high"), + ("mistral-medium-latest", "medium", "high"), + ("mistral-medium-latest", "xhigh", "high"), + ("mistral-small-latest", "medium", "high"), + ("mistral-vibe-cli-latest", "medium", "high"), + ("zai-glm-5", "none", "low"), + ("zai-glm-5", "medium", "high"), + ("zai-glm-5", "xhigh", "max"), + ("zai-glm-5-2", "medium", "medium"), + ("zai-glm-5-2", "xhigh", "xhigh"), + ], ) - def test_reasoning_effort_forwarded_verbatim_for_reasoning_models(self, model, reasoning_effort): + def test_reasoning_effort_is_sent_as_a_level_the_model_accepts(self, local_model_cost_map, model, requested, sent): import litellm optional_params = litellm.get_optional_params( model=model, custom_llm_provider="mistral", - reasoning_effort=reasoning_effort, + reasoning_effort=requested, ) - assert optional_params["reasoning_effort"] == reasoning_effort + assert optional_params["reasoning_effort"] == sent + + def test_reasoning_effort_is_forwarded_verbatim_when_the_map_declares_no_levels( + self, local_model_cost_map, monkeypatch + ): + import litellm + + monkeypatch.setitem( + litellm.model_cost, + "mistral/undeclared-reasoner", + {"litellm_provider": "mistral", "mode": "chat", "supports_reasoning": True}, + ) + optional_params = litellm.get_optional_params( + model="undeclared-reasoner", + custom_llm_provider="mistral", + reasoning_effort="medium", + ) + assert optional_params["reasoning_effort"] == "medium" def test_reasoning_effort_stays_unsupported_for_non_reasoning_models(self): import litellm diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index ccd6766b13a..ca4a3517bfc 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -4,6 +4,7 @@ import litellm from litellm.router_utils.reasoning_effort_capability import ( deployment_is_catalog_mapped, intersect_supported_reasoning_efforts, + nearest_declared_reasoning_effort, resolve_supported_reasoning_efforts, ) @@ -415,3 +416,22 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels: "high", "xhigh", ) + + +class TestNearestDeclaredReasoningEffort: + def test_a_declared_level_is_kept(self): + assert nearest_declared_reasoning_effort("high", ("none", "high")) == "high" + assert nearest_declared_reasoning_effort("none", ("none", "high")) == "none" + + def test_an_undeclared_level_rounds_up_to_the_next_declared_one(self): + assert nearest_declared_reasoning_effort("medium", ("none", "high")) == "high" + assert nearest_declared_reasoning_effort("none", ("low", "high", "max")) == "low" + assert nearest_declared_reasoning_effort("xhigh", ("low", "high", "max")) == "max" + + def test_a_level_above_the_ceiling_takes_the_strongest_declared_one(self): + assert nearest_declared_reasoning_effort("max", ("none", "high")) == "high" + assert nearest_declared_reasoning_effort("xhigh", ("none", "low", "medium", "high")) == "high" + + def test_a_level_outside_the_strength_order_is_left_for_upstream(self): + assert nearest_declared_reasoning_effort("turbo", ("none", "high")) == "turbo" + assert nearest_declared_reasoning_effort("medium", ()) == "medium" From 84f7adec2e4d07ec00cea03b75d0509cde817418 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:06:24 +0000 Subject: [PATCH 154/253] fix(passthrough): match deepgram listen routes served under a path prefix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../deepgram_listen_passthrough_logging_handler.py | 2 +- .../test_deepgram_listen_passthrough_logging_handler.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py index fc0af400477..c0953ae7b85 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -57,7 +57,7 @@ class DeepgramListenPassthroughLoggingHandler: @staticmethod def is_deepgram_listen_route(url_route: str) -> bool: path: Final = urlparse(url_route).path - return path.startswith("/deepgram/") and path.endswith(DEEPGRAM_LISTEN_ROUTE_SUFFIX) + return "/deepgram/" in path and path.endswith(DEEPGRAM_LISTEN_ROUTE_SUFFIX) def deepgram_listen_passthrough_handler( self, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py index a6133a7cf99..2ae56709b11 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -42,6 +42,7 @@ def _metadata(duration: object, channels: int = 1) -> dict[str, object]: ("/deepgram/v1/listen", True), ("/deepgram/listen", True), ("/deepgram/v1/listen?model=nova-3", True), + ("/litellm/deepgram/v1/listen", True), ("/deepgram/v1/speak", False), ("/deepgram/v1/listen/extra", False), ("/openai/v1/realtime", False), From afa8b8a9037c17dc28c2d06649464582dffe1f67 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:11:39 +0000 Subject: [PATCH 155/253] test(docs): read only the first column of the router_settings reference table Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/documentation_tests/test_router_settings.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/documentation_tests/test_router_settings.py b/tests/documentation_tests/test_router_settings.py index 75032f80dfa..7e3d0c07459 100644 --- a/tests/documentation_tests/test_router_settings.py +++ b/tests/documentation_tests/test_router_settings.py @@ -51,9 +51,7 @@ try: if general_settings_section: # Extract the table rows, which contain the documented keys table_content = general_settings_section.group(1) - doc_key_pattern = re.compile( - r"\|\s*([^\|]+?)\s*\|" - ) # Capture the key from each row of the table + doc_key_pattern = re.compile(r"^\|\s*([^\|]+?)\s*\|", re.MULTILINE) documented_keys.update(doc_key_pattern.findall(table_content)) except Exception as e: raise Exception( From aad89de4b4cb64aa8c921dd5315c236d46c1d7f0 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:13:56 +0000 Subject: [PATCH 156/253] build(deps): bump anyio to 4.14.2 in uv.lock to clear GHSA-5p39-cfhj-2xmp and GHSA-82r6-8w77-94w6 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- uv.lock | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/uv.lock b/uv.lock index a5e60c68515..04c31bea7d3 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-09-14T23:55:55.024292355Z" +exclude-newer = "2026-09-15T19:13:10.278189214Z" exclude-newer-span = "P3D" [manifest] @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] From b5ad17e8482f924e1fafb71563d2e4819f105ddb Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:11:39 +0000 Subject: [PATCH 157/253] test(docs): read only the first column of the router_settings reference table Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/documentation_tests/test_router_settings.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/documentation_tests/test_router_settings.py b/tests/documentation_tests/test_router_settings.py index 75032f80dfa..7e3d0c07459 100644 --- a/tests/documentation_tests/test_router_settings.py +++ b/tests/documentation_tests/test_router_settings.py @@ -51,9 +51,7 @@ try: if general_settings_section: # Extract the table rows, which contain the documented keys table_content = general_settings_section.group(1) - doc_key_pattern = re.compile( - r"\|\s*([^\|]+?)\s*\|" - ) # Capture the key from each row of the table + doc_key_pattern = re.compile(r"^\|\s*([^\|]+?)\s*\|", re.MULTILINE) documented_keys.update(doc_key_pattern.findall(table_content)) except Exception as e: raise Exception( From 2299846b48e4c29868bdbd7a5e145493e4de067a Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 19:18:02 +0000 Subject: [PATCH 158/253] fix(bedrock_mantle): align gpt-5.6-sol pricing with the AWS in-region model card Co-authored-by: kusumakarb Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 16 ++++++++-------- model_prices_and_context_window.json | 16 ++++++++-------- ...st_bedrock_mantle_responses_transformation.py | 16 ++++++++++++++++ 3 files changed, 32 insertions(+), 16 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 94e20fa0213..b91c7274a1d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -57990,14 +57990,14 @@ "supports_tool_choice": true }, "bedrock_mantle/openai.gpt-5.6-sol": { - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, "search_context_cost_per_query": { "search_context_size_high": 0.012, "search_context_size_low": 0.012, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 94e20fa0213..b91c7274a1d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -57990,14 +57990,14 @@ "supports_tool_choice": true }, "bedrock_mantle/openai.gpt-5.6-sol": { - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, "search_context_cost_per_query": { "search_context_size_high": 0.012, "search_context_size_low": 0.012, diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 901c005f5a3..81c39f736a5 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -1839,6 +1839,22 @@ class TestBedrockMantleResponsesSigV4: class TestBedrockMantleResponsesPricing: + @pytest.mark.parametrize( + "model", + ["openai.gpt-5.6-sol", "openai.gpt-5.6-terra", "openai.gpt-5.6-luna"], + ) + def test_mantle_matches_in_region_converse_pricing(self, local_cost_map, model): + mantle = litellm.model_cost[f"bedrock_mantle/{model}"] + converse = litellm.model_cost[f"us.{model}"] + + cost_fields = [k for k in converse if "cost" in k and k != "search_context_cost_per_query"] + assert cost_fields, "expected cost fields on the converse entry" + for field in cost_fields: + assert mantle.get(field) == pytest.approx(converse[field]), ( + f"{model}: {field} is {mantle.get(field)} on bedrock_mantle " + f"but {converse[field]} on us. (bedrock_converse)" + ) + def test_models_registered(self, local_cost_map): assert "bedrock_mantle/openai.gpt-5.5" in litellm.bedrock_mantle_models assert "bedrock_mantle/openai.gpt-5.4" in litellm.bedrock_mantle_models From 31f1c5d33107d6cd5179c3ec8c6074a8d5742f87 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 19:18:14 +0000 Subject: [PATCH 159/253] fix(openrouter): refresh glm-5.3, glm-latest, deepseek-flash-latest and hy3 pricing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 22 +++++++++---------- model_prices_and_context_window.json | 22 +++++++++---------- 2 files changed, 22 insertions(+), 22 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b91c7274a1d..f8a537fb922 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -65896,9 +65896,9 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 1.4e-06, - "output_cost_per_token": 4.4e-06, - "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 9.1e-07, + "output_cost_per_token": 2.86e-06, + "cache_read_input_token_cost": 1.69e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943717, @@ -70619,14 +70619,14 @@ "supports_web_search": false }, "openrouter/~deepseek/deepseek-flash-latest": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 1.5e-07, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6e-07, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -70911,14 +70911,14 @@ "supports_web_search": false }, "openrouter/~z-ai/glm-latest": { - "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost": 1.46625e-07, "input_cost_per_token": 9e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 2.805e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -74062,14 +74062,14 @@ "supports_web_search": false }, "openrouter/tencent/hy3": { - "cache_read_input_token_cost": 3.3e-08, - "input_cost_per_token": 1.32e-07, + "cache_read_input_token_cost": 2.0625e-08, + "input_cost_per_token": 8.25e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 5.28e-07, + "output_cost_per_token": 3.3e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b91c7274a1d..f8a537fb922 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -65896,9 +65896,9 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 1.4e-06, - "output_cost_per_token": 4.4e-06, - "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 9.1e-07, + "output_cost_per_token": 2.86e-06, + "cache_read_input_token_cost": 1.69e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 943717, @@ -70619,14 +70619,14 @@ "supports_web_search": false }, "openrouter/~deepseek/deepseek-flash-latest": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 1.5e-07, + "cache_read_input_token_cost": 4.2e-09, + "input_cost_per_token": 1.4e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 6e-07, + "output_cost_per_token": 4.2e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -70911,14 +70911,14 @@ "supports_web_search": false }, "openrouter/~z-ai/glm-latest": { - "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost": 1.46625e-07, "input_cost_per_token": 9e-07, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 235929, "max_tokens": 235929, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 2.805e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -74062,14 +74062,14 @@ "supports_web_search": false }, "openrouter/tencent/hy3": { - "cache_read_input_token_cost": 3.3e-08, - "input_cost_per_token": 1.32e-07, + "cache_read_input_token_cost": 2.0625e-08, + "input_cost_per_token": 8.25e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 5.28e-07, + "output_cost_per_token": 3.3e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, From a774e71175177641c2914dce33ac5d8d106ca131 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 19:18:45 +0000 Subject: [PATCH 160/253] test(bedrock_mantle): document mantle in-region pricing invariant Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../test_bedrock_mantle_responses_transformation.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index 81c39f736a5..951911ac066 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -1844,6 +1844,11 @@ class TestBedrockMantleResponsesPricing: ["openai.gpt-5.6-sol", "openai.gpt-5.6-terra", "openai.gpt-5.6-luna"], ) def test_mantle_matches_in_region_converse_pricing(self, local_cost_map, model): + """bedrock-mantle serves these models In-Region only, and the AWS model + cards price In-Region and Geo CRIS identically -- so every cost field on + the mantle key must equal the `us.` converse key. A price change applied + to one namespace but not the other shows up here. + """ mantle = litellm.model_cost[f"bedrock_mantle/{model}"] converse = litellm.model_cost[f"us.{model}"] From b4dc081c275da7b631faa9cf59782a3cb5452141 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:22:34 -0700 Subject: [PATCH 161/253] fix(mistral): never round reasoning_effort none onto the strength ladder --- litellm/router_utils/reasoning_effort_capability.py | 8 +++++--- .../llms/mistral/test_mistral_chat_transformation.py | 3 ++- .../router_utils/test_reasoning_effort_capability.py | 6 +++++- 3 files changed, 12 insertions(+), 5 deletions(-) diff --git a/litellm/router_utils/reasoning_effort_capability.py b/litellm/router_utils/reasoning_effort_capability.py index f880c846500..1d7656f253e 100644 --- a/litellm/router_utils/reasoning_effort_capability.py +++ b/litellm/router_utils/reasoning_effort_capability.py @@ -103,15 +103,17 @@ def declared_reasoning_efforts_for_model(model: str, custom_llm_provider: str) - return declared_reasoning_efforts(entry) -REASONING_EFFORT_STRENGTH_ORDER: Final = ("none", "minimal", "low", "medium", "high", "xhigh", "max") +REASONING_EFFORT_STRENGTH_ORDER: Final = ("minimal", "low", "medium", "high", "xhigh", "max") _STRENGTH_RANK: Final = MappingProxyType({effort: rank for rank, effort in enumerate(REASONING_EFFORT_STRENGTH_ORDER)}) def nearest_declared_reasoning_effort(requested: str, declared: Sequence[str]) -> str: """Rounds a request up to the weakest declared level at least as strong as it, and down to the strongest declared level when it asks for more than the model has, so the caller gets no less - reasoning than it asked for instead of a rejected call. A level outside the strength order is - returned as is for upstream to judge.""" + reasoning than it asked for instead of a rejected call. none is the off switch rather than a + strength, so it is never rounded onto the ladder and no level is rounded down to it: a caller + who turned reasoning off must not be billed for it, and a model that cannot turn it off says so + itself. A level outside the strength order is likewise returned as is for upstream to judge.""" ranked: Final = sorted( (effort for effort in declared if effort in _STRENGTH_RANK), key=lambda effort: _STRENGTH_RANK[effort] ) diff --git a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py index a68ba9f570c..8fb3b3c43df 100644 --- a/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/test_litellm/llms/mistral/test_mistral_chat_transformation.py @@ -99,7 +99,8 @@ class TestMistralReasoningSupport: ("mistral-medium-latest", "xhigh", "high"), ("mistral-small-latest", "medium", "high"), ("mistral-vibe-cli-latest", "medium", "high"), - ("zai-glm-5", "none", "low"), + ("zai-glm-5", "none", "none"), + ("zai-glm-5", "minimal", "low"), ("zai-glm-5", "medium", "high"), ("zai-glm-5", "xhigh", "max"), ("zai-glm-5-2", "medium", "medium"), diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index ca4a3517bfc..adee44aa8a3 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -425,9 +425,13 @@ class TestNearestDeclaredReasoningEffort: def test_an_undeclared_level_rounds_up_to_the_next_declared_one(self): assert nearest_declared_reasoning_effort("medium", ("none", "high")) == "high" - assert nearest_declared_reasoning_effort("none", ("low", "high", "max")) == "low" + assert nearest_declared_reasoning_effort("minimal", ("low", "high", "max")) == "low" assert nearest_declared_reasoning_effort("xhigh", ("low", "high", "max")) == "max" + def test_none_is_a_switch_that_is_never_rounded_in_either_direction(self): + assert nearest_declared_reasoning_effort("none", ("low", "high", "max")) == "none" + assert nearest_declared_reasoning_effort("medium", ("none",)) == "medium" + def test_a_level_above_the_ceiling_takes_the_strongest_declared_one(self): assert nearest_declared_reasoning_effort("max", ("none", "high")) == "high" assert nearest_declared_reasoning_effort("xhigh", ("none", "low", "medium", "high")) == "high" From 1e7c5400fd792c2b7f5d9915ac832c4147f256dc Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:27:25 +0000 Subject: [PATCH 162/253] fix(proxy): canonicalize azure speech paths and bill uploaded short audio Resolve dot segments in the /azure_speech endpoint path before the endpoint family and the admin-only batch guard are decided, so the guard and the forwarded upstream path agree. Bill short-audio requests for the longer of the uploaded audio duration and the recognized duration, so a NoMatch or silence response still charges for the audio Azure processed Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llm_passthrough_endpoints.py | 16 ++- ...zure_speech_passthrough_logging_handler.py | 40 +++++++- ...zure_speech_passthrough_logging_handler.py | 99 ++++++++++++++++--- .../test_llm_pass_through_endpoints.py | 75 ++++++++++++++ 4 files changed, 207 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ab4b636cb60..ed70369506d 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -12,6 +12,7 @@ import hmac import inspect import json import os +import posixpath import re from collections.abc import AsyncGenerator, Callable, Mapping, Sequence from dataclasses import dataclass @@ -1408,6 +1409,18 @@ def azure_speech_path_manages_shared_resources(endpoint_path: str) -> bool: ) +def canonical_azure_speech_endpoint_path(endpoint: str) -> str: + """ + The path Azure will actually serve, with ``.`` and ``..`` segments resolved, so the + endpoint family and the admin guard are decided on the same path the upstream request uses. + """ + raw_path: Final = httpx.URL(endpoint).path + resolved_path: Final = posixpath.normpath(f"/{raw_path.lstrip('/')}") + if raw_path.endswith("/") and resolved_path != "/": + return f"{resolved_path}/" + return resolved_path + + @router.api_route( f"{AZURE_SPEECH_PASS_THROUGH_ROUTE_PREFIX}/{{endpoint:path}}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: fastapi route methods must be a list @@ -1430,8 +1443,7 @@ async def azure_speech_proxy_route( [Docs](https://docs.litellm.ai/docs/pass_through/azure_speech) """ - endpoint_path: Final = httpx.URL(endpoint).path - normalized_endpoint_path: Final = endpoint_path if endpoint_path.startswith("/") else f"/{endpoint_path}" + normalized_endpoint_path: Final = canonical_azure_speech_endpoint_path(endpoint) base_url: Final = resolve_azure_speech_base_url( endpoint_path=normalized_endpoint_path, api_base=get_secret_str(secret_name="AZURE_SPEECH_API_BASE"), diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py index 588b7cc8e56..8cd1b137de8 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py @@ -18,6 +18,7 @@ from litellm.constants import ( AZURE_SPEECH_TICKS_PER_SECOND, ) from litellm.cost_calculator import transcription_cost +from litellm.litellm_core_utils.audio_utils.utils import calculate_request_duration from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import ( get_standard_logging_object_payload, @@ -53,6 +54,23 @@ class AzureSpeechPassthroughLoggingHandler: return 0.0 return (offset + duration) / AZURE_SPEECH_TICKS_PER_SECOND + @staticmethod + def _uploaded_audio_seconds(httpx_response: httpx.Response) -> float: + try: + uploaded_audio: Final = httpx_response.request.content + except RuntimeError: + return 0.0 + return calculate_request_duration(uploaded_audio) or 0.0 + + @staticmethod + def _short_audio_seconds( + httpx_response: httpx.Response, response_body: Mapping[str, object] | Sequence[object] | None + ) -> float: + return max( + AzureSpeechPassthroughLoggingHandler._uploaded_audio_seconds(httpx_response), + AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body), + ) + @staticmethod def _fast_transcription_audio_seconds(response_body: Mapping[str, object] | Sequence[object] | None) -> float: if not isinstance(response_body, Mapping): @@ -63,16 +81,26 @@ class AzureSpeechPassthroughLoggingHandler: return duration_milliseconds / AZURE_SPEECH_MILLISECONDS_PER_SECOND @staticmethod - def _billed_audio_seconds(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: + def _billed_audio_seconds( + url_route: str, + httpx_response: httpx.Response, + response_body: Mapping[str, object] | Sequence[object] | None, + ) -> float: if AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route): - return AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body) + return AzureSpeechPassthroughLoggingHandler._short_audio_seconds(httpx_response, response_body) if AzureSpeechPassthroughLoggingHandler._is_fast_transcription_route(url_route): return AzureSpeechPassthroughLoggingHandler._fast_transcription_audio_seconds(response_body) return 0.0 @staticmethod - def _response_cost(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float: - audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._billed_audio_seconds(url_route, response_body) + def _response_cost( + url_route: str, + httpx_response: httpx.Response, + response_body: Mapping[str, object] | Sequence[object] | None, + ) -> float: + audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._billed_audio_seconds( + url_route, httpx_response, response_body + ) if audio_seconds <= 0.0: return 0.0 try: @@ -103,7 +131,9 @@ class AzureSpeechPassthroughLoggingHandler: ) -> PassThroughEndpointLoggingTypedDict: try: model_name: Final = AzureSpeechPassthroughLoggingHandler._model_from_url_route(url_route) - response_cost: Final = AzureSpeechPassthroughLoggingHandler._response_cost(url_route, response_body) + response_cost: Final = AzureSpeechPassthroughLoggingHandler._response_cost( + url_route, httpx_response, response_body + ) updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict **kwargs, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py index 50d3de64f72..670f8e65823 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py @@ -1,5 +1,8 @@ +import io import json +import wave from datetime import datetime +from typing import Final from unittest.mock import MagicMock import httpx @@ -29,6 +32,25 @@ TRANSCRIPT_BODY = { TRANSCRIPT = json.dumps(TRANSCRIPT_BODY) TRANSCRIPT_AUDIO_SECONDS = 3.0 PRICE_PER_SECOND = 0.5 +WAV_SAMPLE_RATE: Final = 16000 +UNRECOGNIZED_BODIES: Final = ( + {"RecognitionStatus": "NoMatch", "Offset": 0, "Duration": 0}, + {"RecognitionStatus": "InitialSilenceTimeout"}, + {"Offset": "5000000", "Duration": "25000000"}, + {}, + [], + None, +) + + +def _pcm16_wav(seconds: float) -> bytes: + buffer: Final = io.BytesIO() + with wave.open(buffer, "wb") as wav: + wav.setnchannels(1) + wav.setsampwidth(2) + wav.setframerate(WAV_SAMPLE_RATE) + wav.writeframes(b"\x00\x00" * int(seconds * WAV_SAMPLE_RATE)) + return buffer.getvalue() @pytest.fixture(autouse=True) @@ -45,8 +67,8 @@ def azure_stt_price(monkeypatch: pytest.MonkeyPatch): ) -def _make_response(url: str) -> httpx.Response: - request = httpx.Request("POST", url, headers={"Ocp-Apim-Subscription-Key": "server-secret"}) +def _make_response(url: str, uploaded: bytes = b"") -> httpx.Response: + request = httpx.Request("POST", url, headers={"Ocp-Apim-Subscription-Key": "server-secret"}, content=uploaded) return httpx.Response(200, request=request, text=TRANSCRIPT) @@ -92,22 +114,13 @@ class TestAzureSpeechPassthroughHandler: assert logging_obj.model_call_details["custom_llm_provider"] == "azure_speech" assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost) - @pytest.mark.parametrize( - "response_body", - [ - {"RecognitionStatus": "NoMatch", "Offset": 0, "Duration": 0}, - {"RecognitionStatus": "InitialSilenceTimeout"}, - {"Offset": "5000000", "Duration": "25000000"}, - {}, - [], - None, - ], - ) - def test_short_audio_without_recognized_duration_logs_zero_cost( - self, response_body: dict[str, object] | list[dict[str, object]] | None + @pytest.mark.parametrize("response_body", UNRECOGNIZED_BODIES) + @pytest.mark.parametrize("uploaded", [b"", b"not audio at all"]) + def test_short_audio_with_neither_recognized_nor_decodable_audio_logs_zero_cost( + self, response_body: dict[str, object] | list[dict[str, object]] | None, uploaded: bytes ): handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( - httpx_response=_make_response(SHORT_AUDIO_URL), + httpx_response=_make_response(SHORT_AUDIO_URL, uploaded), response_body=response_body, logging_obj=_make_logging_obj(), url_route=SHORT_AUDIO_URL, @@ -121,6 +134,60 @@ class TestAzureSpeechPassthroughHandler: assert handler_result["kwargs"]["model"] == "azure_speech/short-audio" assert handler_result["kwargs"]["response_cost"] == 0.0 + @pytest.mark.parametrize("response_body", UNRECOGNIZED_BODIES) + def test_short_audio_bills_the_uploaded_audio_when_nothing_was_recognized( + self, response_body: dict[str, object] | list[dict[str, object]] | None + ): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(SHORT_AUDIO_URL, _pcm16_wav(seconds=2.0)), + response_body=response_body, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["response_cost"] == pytest.approx(2.0 * PRICE_PER_SECOND) + + @pytest.mark.parametrize( + "uploaded_seconds,expected_seconds", + [(1.0, TRANSCRIPT_AUDIO_SECONDS), (TRANSCRIPT_AUDIO_SECONDS + 2.0, TRANSCRIPT_AUDIO_SECONDS + 2.0)], + ) + def test_short_audio_bills_the_longer_of_uploaded_and_recognized_audio( + self, uploaded_seconds: float, expected_seconds: float + ): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(SHORT_AUDIO_URL, _pcm16_wav(seconds=uploaded_seconds)), + response_body=TRANSCRIPT_BODY, + logging_obj=_make_logging_obj(), + url_route=SHORT_AUDIO_URL, + result=TRANSCRIPT, + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["response_cost"] == pytest.approx(expected_seconds * PRICE_PER_SECOND) + + def test_fast_transcription_ignores_the_uploaded_multipart_body(self): + handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler( + httpx_response=_make_response(FAST_URL, _pcm16_wav(seconds=30.0)), + response_body=FAST_BODY, + logging_obj=_make_logging_obj(), + url_route=FAST_URL, + result=json.dumps(FAST_BODY), + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + request_body={}, + ) + + assert handler_result["kwargs"]["response_cost"] == pytest.approx(FAST_AUDIO_SECONDS * PRICE_PER_SECOND) + @pytest.mark.parametrize( "response_body", [{"durationMilliseconds": 0}, {"durationMilliseconds": "5061"}, {"duration": 5061}, {}, [], None], diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 9d210ac2f4f..636980eb6e3 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -33,6 +33,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( _proxy_general_settings, anthropic_proxy_route, azure_proxy_route, + azure_speech_proxy_route, bedrock_llm_proxy_route, bedrock_proxy_route, create_pass_through_route, @@ -6443,6 +6444,7 @@ AZURE_SPEECH_PCM16_HEADER: Final = ( b"RIFF\x24\x0c\x00\x00WAVEfmt \x10\x00\x00\x00\x01\x00\x01\x00\x80\x3e\x00\x00\x00\x7d\x00\x00\x02\x00\x10\x00data\x00\x0c\x00\x00" ) AZURE_SPEECH_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + b"\x00" * 3072 +AZURE_SPEECH_WAV_SECONDS: Final = 3072 / (16000 * 2) AZURE_SPEECH_NON_UTF8_WAV_BYTES: Final = AZURE_SPEECH_PCM16_HEADER + bytes(range(256)) * 12 AZURE_SPEECH_TRANSCRIPT: Final = {"RecognitionStatus": "Success", "DisplayText": "The eagle has landed."} @@ -6826,6 +6828,79 @@ class TestAzureSpeechProxyRoute: assert "/azure_speech" in LiteLLMRoutes.mapped_pass_through_routes.value + def test_short_audio_with_no_recognized_speech_is_billed_for_the_uploaded_audio( + self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch + ) -> None: + from litellm.integrations.custom_logger import CustomLogger + + class _Recorder(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.payloads: list[dict[str, object]] = [] # mutable-ok: test recorder accumulates callback payloads + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + self.payloads.append(kwargs["standard_logging_object"]) + + recorder: Final = _Recorder() + monkeypatch.setattr(litellm, "_async_success_callback", [*litellm._async_success_callback, recorder]) + monkeypatch.setitem( + litellm.model_cost, + "azure/speech/azure-stt", + { + "litellm_provider": "azure", + "mode": "audio_transcription", + "input_cost_per_second": 0.25, + "output_cost_per_second": 0.0, + }, + ) + with respx.mock(assert_all_called=True) as upstream: + upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock( + return_value=httpx.Response(200, json={"RecognitionStatus": "NoMatch", "Offset": 0, "Duration": 0}) + ) + + response = azure_speech_client.post( + f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}", + content=AZURE_SPEECH_WAV_BYTES, + headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"}, + ) + + assert response.status_code == 200 + assert [p["model"] for p in recorder.payloads] == ["azure_speech/short-audio"] + assert recorder.payloads[0]["response_cost"] == pytest.approx(AZURE_SPEECH_WAV_SECONDS * 0.25) + + +class TestAzureSpeechProxyRoutePathTraversal: + """Calls the route function directly because httpx clients resolve dot segments before sending.""" + + @pytest.mark.parametrize( + "endpoint", + [ + f"speech/..{AZURE_SPEECH_BATCH_ENDPOINT}", + f"speech/recognition/../..{AZURE_SPEECH_BATCH_ENDPOINT}/", + f"speech/./..{AZURE_SPEECH_BATCH_ENDPOINT}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab", + ], + ) + @pytest.mark.asyncio + async def test_dot_segments_cannot_reach_shared_batch_resources_with_a_non_admin_key( + self, monkeypatch: pytest.MonkeyPatch, endpoint: str + ) -> None: + monkeypatch.setenv("AZURE_SPEECH_API_KEY", "server-subscription-key") + monkeypatch.setenv("AZURE_SPEECH_REGION", "eastus") + monkeypatch.delenv("AZURE_SPEECH_API_BASE", raising=False) + request: Final = MagicMock(spec=Request) + request.method = "GET" + + with pytest.raises(HTTPException) as denied: + await azure_speech_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-virtual"), + ) + + assert denied.value.status_code == 403 + assert AZURE_SPEECH_FAST_ENDPOINT in str(denied.value.detail) + def _azure_speech_real_auth_attrs() -> dict[str, object]: from litellm.caching.caching import DualCache From d4ee62eb8c3083a48fb21f72da49c14a578dc993 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:29:32 -0700 Subject: [PATCH 163/253] refactor(passthrough): type the timeout resolver mapping parameters --- litellm/passthrough/timeout_utils.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index 0170d93c156..829277105e3 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -1,4 +1,5 @@ import sys +from collections.abc import Mapping from typing import Final from pydantic import BaseModel, ConfigDict @@ -42,8 +43,8 @@ def resolve_pass_through_request_timeout( def resolve_llm_passthrough_timeout( - kwargs: dict | None = None, - litellm_params: dict | None = None, + kwargs: Mapping[str, object] | None = None, + litellm_params: Mapping[str, object] | None = None, router_timeout: float | str | None = None, router_stream_timeout: float | str | None = None, ) -> float: From 12b7268ae986bde9dc08c19fff8b188eec084fa7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:29:55 -0700 Subject: [PATCH 164/253] feat(vscode): add LiteLLM language model provider extension --- .github/workflows/test-vscode-extension.yml | 65 + vscode-extension/.gitignore | 3 + vscode-extension/.vscodeignore | 10 + vscode-extension/LICENSE | 21 + vscode-extension/README.md | 31 + vscode-extension/package-lock.json | 3570 +++++++++++++++++++ vscode-extension/package.json | 87 + vscode-extension/src/extension.ts | 17 + vscode-extension/src/gateway.ts | 114 + vscode-extension/src/messages.ts | 188 + vscode-extension/src/models.ts | 170 + vscode-extension/src/provider.ts | 136 + vscode-extension/src/stream.ts | 90 + vscode-extension/src/vscode.proposed.d.ts | 15 + vscode-extension/test/gateway.test.ts | 205 ++ vscode-extension/test/messages.test.ts | 168 + vscode-extension/test/models.test.ts | 185 + vscode-extension/test/provider.test.ts | 258 ++ vscode-extension/test/stream.test.ts | 93 + vscode-extension/test/vscode-mock.ts | 82 + vscode-extension/tsconfig.json | 19 + vscode-extension/vitest.config.mts | 13 + 22 files changed, 5540 insertions(+) create mode 100644 .github/workflows/test-vscode-extension.yml create mode 100644 vscode-extension/.gitignore create mode 100644 vscode-extension/.vscodeignore create mode 100644 vscode-extension/LICENSE create mode 100644 vscode-extension/README.md create mode 100644 vscode-extension/package-lock.json create mode 100644 vscode-extension/package.json create mode 100644 vscode-extension/src/extension.ts create mode 100644 vscode-extension/src/gateway.ts create mode 100644 vscode-extension/src/messages.ts create mode 100644 vscode-extension/src/models.ts create mode 100644 vscode-extension/src/provider.ts create mode 100644 vscode-extension/src/stream.ts create mode 100644 vscode-extension/src/vscode.proposed.d.ts create mode 100644 vscode-extension/test/gateway.test.ts create mode 100644 vscode-extension/test/messages.test.ts create mode 100644 vscode-extension/test/models.test.ts create mode 100644 vscode-extension/test/provider.test.ts create mode 100644 vscode-extension/test/stream.test.ts create mode 100644 vscode-extension/test/vscode-mock.ts create mode 100644 vscode-extension/tsconfig.json create mode 100644 vscode-extension/vitest.config.mts diff --git a/.github/workflows/test-vscode-extension.yml b/.github/workflows/test-vscode-extension.yml new file mode 100644 index 00000000000..886268d9e2c --- /dev/null +++ b/.github/workflows/test-vscode-extension.yml @@ -0,0 +1,65 @@ +name: VS Code Extension +permissions: + contents: read + +on: + pull_request: + branches: + - main + - litellm_internal_staging + - litellm_oss_staging + - "litellm_**" + paths: + - "vscode-extension/**" + - ".github/workflows/test-vscode-extension.yml" + push: + branches: + - main + paths: + - "vscode-extension/**" + - ".github/workflows/test-vscode-extension.yml" + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }} + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + +jobs: + vscode-extension: + runs-on: ubuntu-latest + timeout-minutes: 10 + defaults: + run: + working-directory: vscode-extension + + steps: + - name: Checkout repository + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + fetch-depth: 1 + persist-credentials: false + + - name: Set up Node.js + uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0 + with: + node-version: "24" + cache: npm + cache-dependency-path: vscode-extension/package-lock.json + + - name: Install dependencies + run: npm ci + + - name: Typecheck + run: npm run typecheck + + - name: Unit tests + run: npm test + + - name: Package extension + run: npm run package + + - name: Upload VSIX + uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1 + with: + name: litellm-vscode + path: vscode-extension/*.vsix + if-no-files-found: error diff --git a/vscode-extension/.gitignore b/vscode-extension/.gitignore new file mode 100644 index 00000000000..a08e1da2de7 --- /dev/null +++ b/vscode-extension/.gitignore @@ -0,0 +1,3 @@ +node_modules/ +dist/ +*.vsix diff --git a/vscode-extension/.vscodeignore b/vscode-extension/.vscodeignore new file mode 100644 index 00000000000..dc7c667ce84 --- /dev/null +++ b/vscode-extension/.vscodeignore @@ -0,0 +1,10 @@ +.gitignore +.vscodeignore +node_modules/** +src/** +test/** +tsconfig.json +package-lock.json +**/*.map +**/*.vsix +vitest.config.mts diff --git a/vscode-extension/LICENSE b/vscode-extension/LICENSE new file mode 100644 index 00000000000..dd11dc52350 --- /dev/null +++ b/vscode-extension/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2023 Berri AI + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/vscode-extension/README.md b/vscode-extension/README.md new file mode 100644 index 00000000000..47c1a3afd27 --- /dev/null +++ b/vscode-extension/README.md @@ -0,0 +1,31 @@ +# LiteLLM for VS Code + +Chat in VS Code with every model your [LiteLLM AI Gateway](https://docs.litellm.ai) exposes. The extension registers LiteLLM as a language model provider, so the gateway's models show up in the chat model picker next to the built-in ones, with the price and reasoning effort controls the gateway reports for each of them + +## What you get + +The model list comes from the gateway's `GET /model_group/info` endpoint, scoped to the virtual key you configure, so the picker shows exactly the chat models that key can use. Each model carries its input and output price per 1M tokens in the picker and in the Language Models editor, and its context limits come from the gateway too, so VS Code sizes prompts correctly. A model whose gateway entry lists `supported_reasoning_efforts` gets a Reasoning Effort submenu in the picker's Configure Model menu, and the chosen effort is sent as `reasoning_effort` on every request to that model. Requests go to `POST /v1/chat/completions` on the gateway as streaming chat completions with tools and images passed through, so routing, fallbacks, guardrails, and spend tracking all apply as usual + +## Setup + +1. Install the extension +2. Run `Chat: Manage Language Models` from the Command Palette and pick `LiteLLM` +3. Enter a name for the connection, the gateway URL (for example `https://litellm.example.com`), and a LiteLLM virtual key. The key is stored in VS Code's secret storage +4. Open the chat model picker. The gateway's chat models are listed under the name you chose, each with its price + +Add the same provider again with another name to reach a second gateway or a second key. Run `LiteLLM: Refresh Models` after the gateway's model list changes. To change the key of an existing connection or to drop it, use the gear on its row in the Language Models editor (`Update API Key`, `Delete`); to change the URL, open its entry with `Open in Language Models (JSON)` from the same menu. If the stored key is ever lost the editor shows a `missing its API key` row for that connection until you update the key + +## Requirements + +VS Code 1.109 or newer and a LiteLLM AI Gateway the key can reach. The key needs access to at least one model group whose mode is `chat` + +## Development + +``` +npm ci +npm run typecheck +npm test +npm run package +``` + +`npm run package` writes a `.vsix` you can install with `code --install-extension litellm-vscode-.vsix` diff --git a/vscode-extension/package-lock.json b/vscode-extension/package-lock.json new file mode 100644 index 00000000000..378579dc024 --- /dev/null +++ b/vscode-extension/package-lock.json @@ -0,0 +1,3570 @@ +{ + "name": "litellm-vscode", + "version": "0.1.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "litellm-vscode", + "version": "0.1.0", + "license": "MIT", + "dependencies": { + "openai": "^7.18.0" + }, + "devDependencies": { + "@types/node": "^22.20.3", + "@types/vscode": "^1.109.0", + "@vscode/vsce": "^4.0.0", + "esbuild": "^0.28.2", + "typescript": "^5.9.3", + "vitest": "^4.1.11" + }, + "engines": { + "vscode": "^1.109.0" + } + }, + "node_modules/@azure/abort-controller": { + "version": "2.2.0", + "resolved": "https://registry.npmjs.org/@azure/abort-controller/-/abort-controller-2.2.0.tgz", + "integrity": "sha512-fNAjWnA/nZ2jz31kxR/AqRaUT8ewHBw/WuBIosK0moMy1C9e5ValbDfFdIxJzVOOYaYkV/b2F1S4H/aHiqfVQg==", + "dev": true, + "license": "MIT", + "dependencies": { + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-auth": { + "version": "1.11.0", + "resolved": "https://registry.npmjs.org/@azure/core-auth/-/core-auth-1.11.0.tgz", + "integrity": "sha512-IUZydyTUkDnYdstOW9pFOOUQlBjAepK5teihDE3x6yxsPJs/hsAaaYpeGxdxrgtOiJbBKSjKW7MDk7AEhb4LRg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/abort-controller": "^2.1.2", + "@azure/core-util": "^1.13.0", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-client": { + "version": "1.11.1", + "resolved": "https://registry.npmjs.org/@azure/core-client/-/core-client-1.11.1.tgz", + "integrity": "sha512-2QygG2F76ZpMP2eMztiJvAiFMu71M9rDeU7vO/QKg5Css7MgM4frUOslFjhVjRhbGaCNPtz/S8M6y46/fFKVuQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/abort-controller": "^2.1.2", + "@azure/core-auth": "^1.10.0", + "@azure/core-rest-pipeline": "^1.22.0", + "@azure/core-tracing": "^1.3.0", + "@azure/core-util": "^1.13.0", + "@azure/logger": "^1.3.0", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-process": { + "version": "1.0.0", + "resolved": "https://registry.npmjs.org/@azure/core-process/-/core-process-1.0.0.tgz", + "integrity": "sha512-/shnJ+ooO8WPxDhPEeI/2oRQuubn16gZ6CvlbpWbEswZfzwI9tI/sMAHmF3x1LuQ9yZYXfLW3TjzGMLEC5blKg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-rest-pipeline": { + "version": "1.25.0", + "resolved": "https://registry.npmjs.org/@azure/core-rest-pipeline/-/core-rest-pipeline-1.25.0.tgz", + "integrity": "sha512-bMs8ekJLjX8wPV+9IPBges1SLPyuDtE9g5gLDWOpxzKcoOFQnpLGkbcT1tdw3FaAmDS1gnPmMmJ6y/T5B96kIA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/abort-controller": "^2.1.2", + "@azure/core-auth": "^1.10.0", + "@azure/core-tracing": "^1.3.0", + "@azure/core-util": "^1.13.0", + "@azure/logger": "^1.3.0", + "@typespec/ts-http-runtime": "^0.3.4", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-tracing": { + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/@azure/core-tracing/-/core-tracing-1.4.0.tgz", + "integrity": "sha512-eGwxD0AtncrxeBM4tG8R55Pc3rdX1hNW2WibJAgYpCVA6E93mvvVH+LcssoVjOBrSKWS55yEIHsk0X8ctHmfOQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/core-util": { + "version": "1.14.0", + "resolved": "https://registry.npmjs.org/@azure/core-util/-/core-util-1.14.0.tgz", + "integrity": "sha512-9n2pWK61veAuN0V20t9lOuoV4CFMdyAZ1ygZzvBGk/pBBJRib/PjL9PLXa/aI2CcPpyHfqVsxxqLCYl6uZlfDw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/abort-controller": "^2.1.2", + "@typespec/ts-http-runtime": "^0.3.0", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/identity": { + "version": "4.13.3", + "resolved": "https://registry.npmjs.org/@azure/identity/-/identity-4.13.3.tgz", + "integrity": "sha512-zGQPtqvXPgSA8yfV2CkIQ1qirqk0p9AIVpC5uEkdXQYcKl07QHvyaGYRnZOk0AsQUmxNb4wfkcwY5di8Z5xa9A==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/abort-controller": "^2.0.0", + "@azure/core-auth": "^1.9.0", + "@azure/core-client": "^1.9.2", + "@azure/core-process": "^1.0.0", + "@azure/core-rest-pipeline": "^1.17.0", + "@azure/core-tracing": "^1.0.0", + "@azure/core-util": "^1.11.0", + "@azure/logger": "^1.0.0", + "@azure/msal-browser": "^5.5.0", + "@azure/msal-node": "^6.0.0", + "open": "^10.1.0", + "tslib": "^2.2.0" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/logger": { + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/@azure/logger/-/logger-1.4.0.tgz", + "integrity": "sha512-rbAE25KUfjU/s3XHUdJgceoCP5dEOpMx85J04kF+QMdta73XkuG9JGHHinch+XIoKpBdqljin+KqURpJriSzLA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@typespec/ts-http-runtime": "^0.3.0", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@azure/msal-browser": { + "version": "5.22.0", + "resolved": "https://registry.npmjs.org/@azure/msal-browser/-/msal-browser-5.22.0.tgz", + "integrity": "sha512-5kgu9xeEKgGc2JeidxAtU15NJTqiH/CMCRRQAJ4Rac56kB7KVg91vbNmn+z3RO1vNomPN69UvjKG9h1Pghx6dQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/msal-common": "16.14.1" + }, + "engines": { + "node": ">=0.8.0" + } + }, + "node_modules/@azure/msal-common": { + "version": "16.14.1", + "resolved": "https://registry.npmjs.org/@azure/msal-common/-/msal-common-16.14.1.tgz", + "integrity": "sha512-Or6xhPNyi4zHW25158yxBoyxuCqNSPa5YBVqfF1J5Ks4MJWBo/USXdp05DQIPu1Zli00YZu6t0+h6KvHJealxQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=0.8.0" + } + }, + "node_modules/@azure/msal-node": { + "version": "6.0.1", + "resolved": "https://registry.npmjs.org/@azure/msal-node/-/msal-node-6.0.1.tgz", + "integrity": "sha512-ixSO1Y/kCVRthRs+hSx/5qkwaunX1/RAePhlMN0wIpIQ4WEZ6AREGGnGd1AP0qspHVsAwzWQKGubpjh50JyJIQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/msal-common": "16.14.1", + "jsonwebtoken": "^9.0.0" + }, + "engines": { + "node": ">=20" + } + }, + "node_modules/@esbuild/aix-ppc64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/aix-ppc64/-/aix-ppc64-0.28.2.tgz", + "integrity": "sha512-XExcO+dvLKvVtNTibSTBej1NCAbaGhWn9Ww1ZPx80qsahhPFe/8jgWP0IchNe0F3HwkU7n8ejhH8bjonqht8mQ==", + "cpu": [ + "ppc64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "aix" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-arm": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-arm/-/android-arm-0.28.2.tgz", + "integrity": "sha512-kXXoiPVVGQcnIYGOeaovwOURpniDBpSq4A03qkQ+BMQqtGG6HYap3xne9C1O1yo4TR3qxlCX5IqqmX6fFo2Lqg==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-arm64/-/android-arm64-0.28.2.tgz", + "integrity": "sha512-5YfKeeI8qWfBZIX+u2xZC3Zlb3Os/gLS2sbEKM+I4ZOcsWmHS2WLysCcQZDAFRslDUU5Oiq44gf6PYN1vGwG5A==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/android-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/android-x64/-/android-x64-0.28.2.tgz", + "integrity": "sha512-O387ite7SzUyCcy3JQX4P4bLtEA7bLLkx+esve5JHnyYfNTxcVpXZo9jhdB0lTKN44gztELTdU7nS8Nr16Fs1Q==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/darwin-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.28.2.tgz", + "integrity": "sha512-n4KqkOQrraxHJcgjM1RvwbigfQKIKJVpM7xp+KsxiyUSrRdIXnt73VhrPAx0fV44hgfmIVKjxMN9J1t5jySVkw==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/darwin-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/darwin-x64/-/darwin-x64-0.28.2.tgz", + "integrity": "sha512-uq6suIWYP37qzGddBKPw5QEQPi6HiLGsO7UmkpfyaYNQ3D+rN6w6WfwH+nuqcGXWvawGwxOEroO4YGnFh95azw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/freebsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/freebsd-arm64/-/freebsd-arm64-0.28.2.tgz", + "integrity": "sha512-n+I0BTSRIoy+d6RPKnEVwql5UwBJolytvY4mAOIEJorKlqgPII8ix6slVVrfZ5Tnj7glIZvloylbB/EJPMWEXw==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/freebsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/freebsd-x64/-/freebsd-x64-0.28.2.tgz", + "integrity": "sha512-78XJTJkvPs0kz2w61301PJjXl4g7q3JqiYMZ/M/yVI73EHBrCRTgkhu9oqG7vPqq+a/yadEW8aD+agKlk5xrmg==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-arm": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-arm/-/linux-arm-0.28.2.tgz", + "integrity": "sha512-XlDnu2q5yoqems+xay6wSAcg9DDD7K9RLKZEBOMZm3ckNpJBvOX20tSfby8KfrrhINDyv9V2YVZKY/SpoGJI8w==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-arm64/-/linux-arm64-0.28.2.tgz", + "integrity": "sha512-pW4AC0P3it8c7do9MVM4p51FzHzdM/TZrerurgRcHJ2WTa1VQ1CIq18xncfpBJw4ojkiZZrKW2yIBWBP92j6Ug==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-ia32": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-ia32/-/linux-ia32-0.28.2.tgz", + "integrity": "sha512-CYbnj78HsIeA+DhgUKgFCfvNsTHFhMMrinUrMZpDXJXKN8T3XViTZ/+wtHeVxEWY8ewSzTFN+nRmSwO2tZaLUQ==", + "cpu": [ + "ia32" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-loong64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-loong64/-/linux-loong64-0.28.2.tgz", + "integrity": "sha512-buwkd8nsph4R+ajRvw0qM5Hja/TXQow3ptzWO2EbG/cqcIkHloRrdlBtQlshyYGTNFvfkfJ5tpPLVkY4DtsPfQ==", + "cpu": [ + "loong64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-mips64el": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-mips64el/-/linux-mips64el-0.28.2.tgz", + "integrity": "sha512-ZVykbDyk7519VwiNb9Lcj9m8XM6v5V9uKPvrEMkkEedVewf+0itkhahp4HDpgERXhwLRpWFypsGbG/J8s0QjJA==", + "cpu": [ + "mips64el" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-ppc64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-ppc64/-/linux-ppc64-0.28.2.tgz", + "integrity": "sha512-CAXl+Dtd9UUuJd8pKKdwh6MLm3MUMiqMPmhZ3tTSXPqfyQ3vDl6R5hZdZ/kYojK4ofXtdfSv1tFq8XzWx3heNQ==", + "cpu": [ + "ppc64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-riscv64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-riscv64/-/linux-riscv64-0.28.2.tgz", + "integrity": "sha512-GeXCej4IQtU1B+QlDV8W/RRvbzI3O/Stss+/bCXv4lZls5WGRtu2a+3JkA3i4qIUlMXpcHebWpF8AkJhATowuA==", + "cpu": [ + "riscv64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-s390x": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-s390x/-/linux-s390x-0.28.2.tgz", + "integrity": "sha512-3H1weTYZPxt/WOhByszQZybS9w5lKzUn1FDMsgEChbHWQwHYQQRfBxgCcZvPhjHfKyJjIievvMmEUawJrdY9Dg==", + "cpu": [ + "s390x" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/linux-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/linux-x64/-/linux-x64-0.28.2.tgz", + "integrity": "sha512-4xTZr1FUmSoQW4XIWmit3tzQrUTZM+N3P0XV8xROKYF50XfI7xeO90+1bZvNwxIufQ9hDQVRJH5YhgPVF8A/HQ==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/netbsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/netbsd-arm64/-/netbsd-arm64-0.28.2.tgz", + "integrity": "sha512-sSATRjPeDBg3pdgHoQfoYBob11Kk1FGa9lui5RIHZCoCkJa9QKlvl3/vKz2usCmYYjs7ymJR/2Nnsqe+Hjt5nw==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "netbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/netbsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/netbsd-x64/-/netbsd-x64-0.28.2.tgz", + "integrity": "sha512-lqnzCV+mM0gIADaKihiCg6ifgfU2L3h5E33rNQBN1Y4MaVGnzryzmvvf7UHxprpQdE8hpqLolJ9Rl+SkIRDpyw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "netbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openbsd-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openbsd-arm64/-/openbsd-arm64-0.28.2.tgz", + "integrity": "sha512-AL2qJILH7lNjrDmCQDvdxMfAUIv8KMNZOvrwAQ8i8//ntL9FflhOyMJ8OZSMBb8/AWXe3/5v5S20y3zCoZWKoQ==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "openbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openbsd-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openbsd-x64/-/openbsd-x64-0.28.2.tgz", + "integrity": "sha512-QtiuPytchRyC4rwUKhexJdQKvDuZ6hWloi3igqPQNUJCS1/v9EiO3UTOXR6A3FoMo4fnAKbWJdqaIwhOzh8qEw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "openbsd" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/openharmony-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/openharmony-arm64/-/openharmony-arm64-0.28.2.tgz", + "integrity": "sha512-WkhYDmpTjLvGlScA1rwjRUmhl4k8oXR3cIbtqWmELgU/dFeHHlEllxDvdWcNJV9rbzCexB5vz8gtNewWLgCT7Q==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "openharmony" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/sunos-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/sunos-x64/-/sunos-x64-0.28.2.tgz", + "integrity": "sha512-GPMSkTOtMnv2U2F8gxe4Io6qmVs+YKyp832Etqqxr0hFngmXQ3rzwytelm3GIn7T4VviRUlf3sOgBOiTdvaf7g==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "sunos" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-arm64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-arm64/-/win32-arm64-0.28.2.tgz", + "integrity": "sha512-PIhhEkE9uPBleRBrQEJpUn7MBnibZzbGzYWPmY3x+YoVg/95zbjB4CxPPOQ8l5tYYM4mMaCthF8/1DIfBQQyWQ==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-ia32": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-ia32/-/win32-ia32-0.28.2.tgz", + "integrity": "sha512-YmJbfTlvU7Sdn9BB+4PRES4oB6pxgS37MAONj+hBr/cpXS1aBPKXxNnDbu+QCWPj0o9dgyxeq79g6c5P8KeuYA==", + "cpu": [ + "ia32" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@esbuild/win32-x64": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.28.2.tgz", + "integrity": "sha512-5ebpxr3nWMzrL/rnUI755Jkuee0bHL/Gq0WTF9lvcpv73wAp5eu8MfBUgWK9bhWvZjj7yX8etf/8tI8Ney695g==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">=18" + } + }, + "node_modules/@jridgewell/sourcemap-codec": { + "version": "1.6.0", + "resolved": "https://registry.npmjs.org/@jridgewell/sourcemap-codec/-/sourcemap-codec-1.6.0.tgz", + "integrity": "sha512-T7jf+5zgsZHwNJ4lvQ7/aezbyk0nNX+zJVWpmHA7VYsEx7a7qr5Rg5IbtJFqkgze5Y2sruq1RUY8Q837Od7iFw==", + "dev": true, + "license": "MIT" + }, + "node_modules/@napi-rs/keyring": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring/-/keyring-1.3.0.tgz", + "integrity": "sha512-WrOw/bcXm0f9qHkumlT1QlArXSTWqaY9sunsDpOk+yCCorCKMxvWT/a3xko4EYHVdeZoh00yI2TydXn6eyICDA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 10" + }, + "funding": { + "type": "github", + "url": "https://github.com/sponsors/Brooooooklyn" + }, + "optionalDependencies": { + "@napi-rs/keyring-darwin-arm64": "1.3.0", + "@napi-rs/keyring-darwin-x64": "1.3.0", + "@napi-rs/keyring-freebsd-x64": "1.3.0", + "@napi-rs/keyring-linux-arm-gnueabihf": "1.3.0", + "@napi-rs/keyring-linux-arm64-gnu": "1.3.0", + "@napi-rs/keyring-linux-arm64-musl": "1.3.0", + "@napi-rs/keyring-linux-riscv64-gnu": "1.3.0", + "@napi-rs/keyring-linux-x64-gnu": "1.3.0", + "@napi-rs/keyring-linux-x64-musl": "1.3.0", + "@napi-rs/keyring-win32-arm64-msvc": "1.3.0", + "@napi-rs/keyring-win32-ia32-msvc": "1.3.0", + "@napi-rs/keyring-win32-x64-msvc": "1.3.0" + } + }, + "node_modules/@napi-rs/keyring-darwin-arm64": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-darwin-arm64/-/keyring-darwin-arm64-1.3.0.tgz", + "integrity": "sha512-pl76hJvdYUBn6I24bXiOBMA9nbDapo3I5B+f3OorjDU4dUMSypXeKbOVehJe8fhgTiH24flMyTS3aAIy43xegQ==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-darwin-x64": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-darwin-x64/-/keyring-darwin-x64-1.3.0.tgz", + "integrity": "sha512-YcJtEV5LA3cvA4z3BurgxH5IhTsW1JfIvcAAcqcecwk06Si9F9NqkxbZVIfDwQ8oRHgaBmT3zZJnLAotCrVahw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-freebsd-x64": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-freebsd-x64/-/keyring-freebsd-x64-1.3.0.tgz", + "integrity": "sha512-vlLf31TGhfRAaxLDBhg8b89ss0HHD/lyNmL5F3UjSaz5CUXElsJmKYq9fqA/B+cZKUEUcLHHGhF0I/CqcFdaVw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-arm-gnueabihf": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-arm-gnueabihf/-/keyring-linux-arm-gnueabihf-1.3.0.tgz", + "integrity": "sha512-KiWdMMu/Inz/bHHIAGrnF7r54FZDYXuHO6UFF/rhIrshUsxbMG1Rl9lEymNtqqsVo927G0VYcb02FzWQ3iBQRQ==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-arm64-gnu": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-arm64-gnu/-/keyring-linux-arm64-gnu-1.3.0.tgz", + "integrity": "sha512-eyKGpY40lm9Jvs1aD294XRH4y7+TlJM0YVAryZeXA6TX0mb4gMkxVXwSQv7MCwgah7raeUd0dKUb4BPAYIgcMg==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-arm64-musl": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-arm64-musl/-/keyring-linux-arm64-musl-1.3.0.tgz", + "integrity": "sha512-iIK6JWHXAJqDrEyLY3TmswwloVyt2vj+04TZnew+uSJ9gnDO8EwRbp3/iw3LpWaXiDO7VomGO6y8I0Id8uBZSw==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-riscv64-gnu": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-riscv64-gnu/-/keyring-linux-riscv64-gnu-1.3.0.tgz", + "integrity": "sha512-/PGqrwn6EwgtK6vccASSXJRfOSP4vN1F4ASsIQ+7MdrK6hNvAJ1FZPrIuD5gGGdxezo3F++To2Wq7DbuGIeuNQ==", + "cpu": [ + "riscv64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-x64-gnu": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-x64-gnu/-/keyring-linux-x64-gnu-1.3.0.tgz", + "integrity": "sha512-2PDK1WKWTu9lBGq9VvNEkSlQD3O7YwVpmnyN2M3cy4v7NJ/8gDMd9GXv3G+FVXN13uhp4gnnPBS+ScefmEeD2A==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-linux-x64-musl": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-linux-x64-musl/-/keyring-linux-x64-musl-1.3.0.tgz", + "integrity": "sha512-oJ2HkX8YUo46QBkn0pG+HuIKQNqr523q6vBobCn+P95s4C4K6/kLBqHY/1bg5J4ap31DzsznhnFKcfBNBsjCnw==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-win32-arm64-msvc": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-win32-arm64-msvc/-/keyring-win32-arm64-msvc-1.3.0.tgz", + "integrity": "sha512-tOd3c/uAaeoE4ycVlmAdSvygz0Zt3zdca6Y7gokBeIbaRDWpjDIUOpU3MvML59XAaqyuKGsVVu0F/DZb1lHPmw==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-win32-ia32-msvc": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-win32-ia32-msvc/-/keyring-win32-ia32-msvc-1.3.0.tgz", + "integrity": "sha512-sPSqeAFZMGqP1R++M2JTza7GQJJ/TpCo6JU6Vcd4jnebvOaEDs9b7eipakU1PJdSvhpC2yXMCNRk9gXfrhuwHQ==", + "cpu": [ + "ia32" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@napi-rs/keyring-win32-x64-msvc": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/@napi-rs/keyring-win32-x64-msvc/-/keyring-win32-x64-msvc-1.3.0.tgz", + "integrity": "sha512-4DnCWXwDc0HRKwyRlG5y0VhKZW2tNRQfKKfyj6IX/KWfDNyq9hn4n+GL1auyDcOO/v8PwnhmYo2+rOOqCkvvOg==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 10" + } + }, + "node_modules/@oxc-project/types": { + "version": "0.150.0", + "resolved": "https://registry.npmjs.org/@oxc-project/types/-/types-0.150.0.tgz", + "integrity": "sha512-rDS5/31E9HfPl/CIzGrn0DOlvBbXFseQ5URJ9sYMfstbKLD/c6Gm9vmRzRGDdAXyOIL4zmO37lc9RIwYqVruZw==", + "dev": true, + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/oxc-project" + } + }, + "node_modules/@rolldown/binding-android-arm-eabi": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm-eabi/-/binding-android-arm-eabi-1.2.9.tgz", + "integrity": "sha512-tNISae1QEf/vkb3xkRcjV5SEdzPE97We5IVaa2Z8jSszQPZ8U60B/YCYpw4QI7VidYsBtKavczXf+DyDs9WGxw==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-android-arm64": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-android-arm64/-/binding-android-arm64-1.2.9.tgz", + "integrity": "sha512-YC8YsI30o606GTZi0VyzYlsDKFP8W61i/QzayHDkLbNEz/IShqAmTa+hsJRj13xTHA0H+6fk4b2UmGn+Q/cMlg==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-darwin-arm64": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-arm64/-/binding-darwin-arm64-1.2.9.tgz", + "integrity": "sha512-IwhlH3qK5urrY8hZiEgGkHKEFN901p/p2bjxCxJlr4GyNnF7wYpUvK+Y43uaRYuC4hpfjzbR3SJC3arX1jGvmw==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-darwin-x64": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-darwin-x64/-/binding-darwin-x64-1.2.9.tgz", + "integrity": "sha512-XxpJfVzFh+jilRxIXUqcfYAYcunIc/XEzIizsOL1fcJee5Sf7H3mH8WlLmfHfluz5amqR88QQo9izKtmMlavAw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-freebsd-x64": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-freebsd-x64/-/binding-freebsd-x64-1.2.9.tgz", + "integrity": "sha512-kSfvhmgeWyfkbT3p/1s5vSgboogoah2zkm9fX2zjg2hHxSV7T4KhMWRUUaRk4OXNqoD3QAUeRqLcs1aZOK4U1g==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-arm-gnueabihf": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm-gnueabihf/-/binding-linux-arm-gnueabihf-1.2.9.tgz", + "integrity": "sha512-1RVzG17pxqbTfYLC352JlLt6kKLG+6Hr30n8DlIJqsnV5luUDd2Qdx9Ayw1Cabfyb1K9k0jXEZ7evxkRoT+uiw==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-arm64-gnu": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-gnu/-/binding-linux-arm64-gnu-1.2.9.tgz", + "integrity": "sha512-BXqPvZ2drqVD+/Z8UpKwcs4Mp7grM+eGFku4CAEKrEtcbAsUpzREphK1sogCRZGreVPiMkiiBtw0n3TPteuqvw==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-arm64-musl": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-arm64-musl/-/binding-linux-arm64-musl-1.2.9.tgz", + "integrity": "sha512-11vWvo8YDwLzukt27J3aYDWU+gg2P7J+ZOmiJ0hkF5BXZDW7pVya7r40MXDy6ya0i9KamoENSVKIugvJNgFXIA==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-ppc64-gnu": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-ppc64-gnu/-/binding-linux-ppc64-gnu-1.2.9.tgz", + "integrity": "sha512-a1tijMkdwsIARtc0F39ApURROkf3NwqinI6TOiSSWCTR7dT96dffNvMUtDHnq64wKNTIZOIlzKrFvvFUznJiyw==", + "cpu": [ + "ppc64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-s390x-gnu": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-s390x-gnu/-/binding-linux-s390x-gnu-1.2.9.tgz", + "integrity": "sha512-x6SQNdAvv4c3hWqTMaWuawzMX9myaCs/yEmlGsxJzkdClnHW7FbrjQuSiRDhuSYzEYoEMhsaJy9qHG/XNemJPQ==", + "cpu": [ + "s390x" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-x64-gnu": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-gnu/-/binding-linux-x64-gnu-1.2.9.tgz", + "integrity": "sha512-9s0AZ8BFK5/n7B/TBoa2yJE3gI3KURrbXcPBlsAsvjU4VeJKgE90y1YtNxyEUIcHPQkg6/yfF3qihUrcM/Kf0Q==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-linux-x64-musl": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-linux-x64-musl/-/binding-linux-x64-musl-1.2.9.tgz", + "integrity": "sha512-P7VWAmV+WdJluH7ovnRGoiv2i8To7GAZ+kGzfGup635cyL7SyYl3lSUaA3Gp5THf0n/Co5EyEqb2zbqq+nMOHQ==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MIT", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-openharmony-arm64": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-openharmony-arm64/-/binding-openharmony-arm64-1.2.9.tgz", + "integrity": "sha512-1qixtsE4BK8h+yS3BfmZ09UhA7O/N4IACva6YBr7EBvCJraByTuRcgOTaiA62Tm0vey3UcKXLOaoGHtYmNGEVg==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "openharmony" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-win32-arm64-msvc": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-win32-arm64-msvc/-/binding-win32-arm64-msvc-1.2.9.tgz", + "integrity": "sha512-ok8IQjcEPs1AKZfuEUznVBrJw+gK4soq+bx8b1X2XoMqVClarc1q5JDmVtWXY1xfr6ZuHTAsPXHTgTrqKTZeww==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/binding-win32-x64-msvc": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/@rolldown/binding-win32-x64-msvc/-/binding-win32-x64-msvc-1.2.9.tgz", + "integrity": "sha512-Ip2mXoU0hM0boq3Rf+ekuT653OROSo6aSYcPT1VHE4q52KvyxgFkQgrgb/IEsxOuvQ2fZZbs8khJAyCEPM24/g==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MIT", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": "^20.19.0 || >=22.12.0" + } + }, + "node_modules/@rolldown/pluginutils": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.1.tgz", + "integrity": "sha512-2j9bGt5Jh8hj+vPtgzPtl72j0yRxHAyumoo6TNfAjsLB04UtpSvPbPcDcBMxz7n+9CYB0c1GxQFxYRg2jimqGw==", + "dev": true, + "license": "MIT" + }, + "node_modules/@secretlint/core": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/core/-/core-10.2.2.tgz", + "integrity": "sha512-6rdwBwLP9+TO3rRjMVW1tX+lQeo5gBbxl1I5F8nh8bgGtKwdlCMhMKsBWzWg1ostxx/tIG7OjZI0/BxsP8bUgw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@secretlint/profiler": "^10.2.2", + "@secretlint/types": "^10.2.2", + "debug": "^4.4.1", + "structured-source": "^4.0.0" + }, + "engines": { + "node": ">=20.0.0" + } + }, + "node_modules/@secretlint/profiler": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/profiler/-/profiler-10.2.2.tgz", + "integrity": "sha512-qm9rWfkh/o8OvzMIfY8a5bCmgIniSpltbVlUVl983zDG1bUuQNd1/5lUEeWx5o/WJ99bXxS7yNI4/KIXfHexig==", + "dev": true, + "license": "MIT" + }, + "node_modules/@secretlint/secretlint-rule-no-dotenv": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/secretlint-rule-no-dotenv/-/secretlint-rule-no-dotenv-10.2.2.tgz", + "integrity": "sha512-KJRbIShA9DVc5Va3yArtJ6QDzGjg3PRa1uYp9As4RsyKtKSSZjI64jVca57FZ8gbuk4em0/0Jq+uy6485wxIdg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@secretlint/types": "^10.2.2" + }, + "engines": { + "node": ">=20.0.0" + } + }, + "node_modules/@secretlint/secretlint-rule-preset-recommend": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/secretlint-rule-preset-recommend/-/secretlint-rule-preset-recommend-10.2.2.tgz", + "integrity": "sha512-K3jPqjva8bQndDKJqctnGfwuAxU2n9XNCPtbXVI5JvC7FnQiNg/yWlQPbMUlBXtBoBGFYp08A94m6fvtc9v+zA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=20.0.0" + } + }, + "node_modules/@secretlint/source-creator": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/source-creator/-/source-creator-10.2.2.tgz", + "integrity": "sha512-h6I87xJfwfUTgQ7irWq7UTdq/Bm1RuQ/fYhA3dtTIAop5BwSFmZyrchph4WcoEvbN460BWKmk4RYSvPElIIvxw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@secretlint/types": "^10.2.2", + "istextorbinary": "^9.5.0" + }, + "engines": { + "node": ">=20.0.0" + } + }, + "node_modules/@secretlint/types": { + "version": "10.2.2", + "resolved": "https://registry.npmjs.org/@secretlint/types/-/types-10.2.2.tgz", + "integrity": "sha512-Nqc90v4lWCXyakD6xNyNACBJNJ0tNCwj2WNk/7ivyacYHxiITVgmLUFXTBOeCdy79iz6HtN9Y31uw/jbLrdOAg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=20.0.0" + } + }, + "node_modules/@standard-schema/spec": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@standard-schema/spec/-/spec-1.1.0.tgz", + "integrity": "sha512-l2aFy5jALhniG5HgqrD6jXLi/rUWrKvqN/qJx6yoJsgKhblVd+iqqU4RCXavm/jPityDo5TCvKMnpjKnOriy0w==", + "dev": true, + "license": "MIT" + }, + "node_modules/@types/chai": { + "version": "5.2.3", + "resolved": "https://registry.npmjs.org/@types/chai/-/chai-5.2.3.tgz", + "integrity": "sha512-Mw558oeA9fFbv65/y4mHtXDs9bPnFMZAL/jxdPFUpOHHIXX91mcgEHbS5Lahr+pwZFR8A7GQleRWeI6cGFC2UA==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/deep-eql": "*", + "assertion-error": "^2.0.1" + } + }, + "node_modules/@types/deep-eql": { + "version": "4.0.2", + "resolved": "https://registry.npmjs.org/@types/deep-eql/-/deep-eql-4.0.2.tgz", + "integrity": "sha512-c9h9dVVMigMPc4bwTvC5dxqtqJZwQPePsWjPlpSOnojbor6pGqdk541lfA7AqFQr5pB1BRdq0juY9db81BwyFw==", + "dev": true, + "license": "MIT" + }, + "node_modules/@types/estree": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/@types/estree/-/estree-1.0.9.tgz", + "integrity": "sha512-GhdPgy1el4/ImP05X05Uw4cw2/M93BCUmnEvWZNStlCzEKME4Fkk+YpoA5OiHNQmoS7Cafb8Xa3Pya8m1Qrzeg==", + "dev": true, + "license": "MIT" + }, + "node_modules/@types/node": { + "version": "22.20.3", + "resolved": "https://registry.npmjs.org/@types/node/-/node-22.20.3.tgz", + "integrity": "sha512-DZmzkmwHzXrLPAXPyKNDzlIwMMUZCVacoD25ywdy5YTKGbOx/2ld+Q38Im2zJ0vBuZP5Prd3VZutKZyXwkOS8A==", + "dev": true, + "license": "MIT", + "dependencies": { + "undici-types": "~6.21.0" + } + }, + "node_modules/@types/vscode": { + "version": "1.138.0", + "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.138.0.tgz", + "integrity": "sha512-DhlvrucJa8n6EHst7qQcXOcjLKs/VRTQy55plLO+RN74jonYsr/t9dd66zBnAyz1tpDig0okQAxuoPHpVNje2g==", + "dev": true, + "license": "MIT" + }, + "node_modules/@typespec/ts-http-runtime": { + "version": "0.3.9", + "resolved": "https://registry.npmjs.org/@typespec/ts-http-runtime/-/ts-http-runtime-0.3.9.tgz", + "integrity": "sha512-edSdeAqkdxBVzA1yL1LrLCml1YjyCVvPMtMqJpbF+6K609tHe8V6sQUzFQSGcYNhcuhOceZtjvN32+mpIth30A==", + "dev": true, + "license": "MIT", + "dependencies": { + "http-proxy-agent": "^7.0.0", + "https-proxy-agent": "^7.0.0", + "tslib": "^2.6.2" + }, + "engines": { + "node": ">=22.0.0" + } + }, + "node_modules/@vitest/expect": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-4.1.11.tgz", + "integrity": "sha512-VX2x5vNJXET47KAFzwERI+KRMtTTCSWTfSMKsW7JsUsXV4psq++e3DvZpuTDOpHcxytiDs6p2nhVb2tVDiiUYw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@standard-schema/spec": "^1.1.0", + "@types/chai": "^5.2.2", + "@vitest/spy": "4.1.11", + "@vitest/utils": "4.1.11", + "chai": "^6.2.2", + "tinyrainbow": "^3.1.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/mocker": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-4.1.11.tgz", + "integrity": "sha512-2XJVD55d1o5AZous5CCGKS74g/riOj9odEt2bQpCVZeblHyHdnMeFl4jl0XjU21stf4mbjUkew2eXQZt65g5CQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/spy": "4.1.11", + "estree-walker": "^3.0.3", + "magic-string": "^0.30.21" + }, + "funding": { + "url": "https://opencollective.com/vitest" + }, + "peerDependencies": { + "msw": "^2.4.9", + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0" + }, + "peerDependenciesMeta": { + "msw": { + "optional": true + }, + "vite": { + "optional": true + } + } + }, + "node_modules/@vitest/pretty-format": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-4.1.11.tgz", + "integrity": "sha512-yiZzPbGTS9Sr/JpFl8zHrcIkAofNbFV6k21vIgQN/cY/oxZeXhJv5sc/MBJ5jFKWmWs+oJHw0UXLZjmf931+Vw==", + "dev": true, + "license": "MIT", + "dependencies": { + "tinyrainbow": "^3.1.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/runner": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-4.1.11.tgz", + "integrity": "sha512-LztvUgdwMNJMIkj3hQnnxiC2Xy1zNxq928W/xhjCLaNCzqTZOudjwbQf6v9IntZGPw132i2Lq2rgTRZHD3JHNw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/utils": "4.1.11", + "pathe": "^2.0.3" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/snapshot": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-4.1.11.tgz", + "integrity": "sha512-pN7ikn1ON7h8ee4gIAp4AzyK+zBtJPzVbqOgu5LCEh4VaJVbPQcgYQYJIMGQPXVeJJq1fnfazis7a5pFNPahog==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/pretty-format": "4.1.11", + "@vitest/utils": "4.1.11", + "magic-string": "^0.30.21", + "pathe": "^2.0.3" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/spy": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-4.1.11.tgz", + "integrity": "sha512-apNa/prQy2qCeywhnixOHPRCgGNhvg7T4Dapfl1GahLp/R+uhBm5cPyFoNVyqsNd2h1nJxL6BqqdIjiABL60YA==", + "dev": true, + "license": "MIT", + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vitest/utils": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-4.1.11.tgz", + "integrity": "sha512-zTCVGpyFsGWBhllOyKlTw/vnr6D9qxsfSDyfbyZmTyjHw5N/VuvzHpHoQjm2ZJzn4RJgx5w4r7V0er69CmLgPQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/pretty-format": "4.1.11", + "convert-source-map": "^2.0.0", + "tinyrainbow": "^3.1.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + } + }, + "node_modules/@vscode/vsce": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/@vscode/vsce/-/vsce-4.0.0.tgz", + "integrity": "sha512-NImwuLaenMmb5D5Jer9/lzi/F9ZQUBOp8Azhj/BVYcTFgixv8KehFXqEUDjQlD2tAiw2E6dDGyjTuAB//di60A==", + "dev": true, + "license": "MIT", + "dependencies": { + "@azure/identity": "^4.13.2", + "@napi-rs/keyring": "^1.3.0", + "@secretlint/core": "^10.2.2", + "@secretlint/secretlint-rule-no-dotenv": "^10.2.2", + "@secretlint/secretlint-rule-preset-recommend": "^10.2.2", + "@secretlint/source-creator": "^10.2.2", + "@secretlint/types": "^10.2.2", + "@vscode/vsce-sign": "^2.1.0", + "azure-devops-node-api": "^12.5.0", + "cockatiel": "^3.2.1", + "commander": "^12.1.0", + "hosted-git-info": "^4.1.0", + "jsonc-parser": "^3.3.1", + "marked": "^18.0.11", + "mime": "^1.6.0", + "minimatch": "^10.2.6", + "parse5": "^8.0.1", + "proper-lockfile": "^4.1.2", + "read": "^1.0.7", + "semver": "^7.8.5", + "tinyglobby": "^0.2.17", + "typed-rest-client": "^1.8.11", + "url-join": "^4.0.1", + "xml2js": "^0.5.0", + "yauzl": "^3.4.0", + "yazl": "^2.5.1" + }, + "bin": { + "vsce": "vsce" + }, + "engines": { + "node": ">= 22" + } + }, + "node_modules/@vscode/vsce-sign": { + "version": "2.1.0", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign/-/vsce-sign-2.1.0.tgz", + "integrity": "sha512-9AQrqazrBgTgRSuwleLVXUrIUphY02/SFCh2TKYoLV/xifJAdblhdmEmw5gUrYSPQ3sRwNs9iyCMD14sATEE6g==", + "dev": true, + "hasInstallScript": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optionalDependencies": { + "@vscode/vsce-sign-alpine-arm64": "2.0.6", + "@vscode/vsce-sign-alpine-x64": "2.0.6", + "@vscode/vsce-sign-darwin-arm64": "2.0.6", + "@vscode/vsce-sign-darwin-x64": "2.0.6", + "@vscode/vsce-sign-linux-arm": "2.0.6", + "@vscode/vsce-sign-linux-arm64": "2.0.6", + "@vscode/vsce-sign-linux-x64": "2.0.6", + "@vscode/vsce-sign-win32-arm64": "2.0.6", + "@vscode/vsce-sign-win32-x64": "2.0.6" + } + }, + "node_modules/@vscode/vsce-sign-alpine-arm64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-alpine-arm64/-/vsce-sign-alpine-arm64-2.0.6.tgz", + "integrity": "sha512-wKkJBsvKF+f0GfsUuGT0tSW0kZL87QggEiqNqK6/8hvqsXvpx8OsTEc3mnE1kejkh5r+qUyQ7PtF8jZYN0mo8Q==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "alpine" + ] + }, + "node_modules/@vscode/vsce-sign-alpine-x64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-alpine-x64/-/vsce-sign-alpine-x64-2.0.6.tgz", + "integrity": "sha512-YoAGlmdK39vKi9jA18i4ufBbd95OqGJxRvF3n6ZbCyziwy3O+JgOpIUPxv5tjeO6gQfx29qBivQ8ZZTUF2Ba0w==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "alpine" + ] + }, + "node_modules/@vscode/vsce-sign-darwin-arm64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-darwin-arm64/-/vsce-sign-darwin-arm64-2.0.6.tgz", + "integrity": "sha512-5HMHaJRIQuozm/XQIiJiA0W9uhdblwwl2ZNDSSAeXGO9YhB9MH5C4KIHOmvyjUnKy4UCuiP43VKpIxW1VWP4tQ==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@vscode/vsce-sign-darwin-x64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-darwin-x64/-/vsce-sign-darwin-x64-2.0.6.tgz", + "integrity": "sha512-25GsUbTAiNfHSuRItoQafXOIpxlYj+IXb4/qarrXu7kmbH94jlm5sdWSCKrrREs8+GsXF1b+l3OB7VJy5jsykw==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@vscode/vsce-sign-linux-arm": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-linux-arm/-/vsce-sign-linux-arm-2.0.6.tgz", + "integrity": "sha512-UndEc2Xlq4HsuMPnwu7420uqceXjs4yb5W8E2/UkaHBB9OWCwMd3/bRe/1eLe3D8kPpxzcaeTyXiK3RdzS/1CA==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@vscode/vsce-sign-linux-arm64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-linux-arm64/-/vsce-sign-linux-arm64-2.0.6.tgz", + "integrity": "sha512-cfb1qK7lygtMa4NUl2582nP7aliLYuDEVpAbXJMkDq1qE+olIw/es+C8j1LJwvcRq1I2yWGtSn3EkDp9Dq5FdA==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@vscode/vsce-sign-linux-x64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-linux-x64/-/vsce-sign-linux-x64-2.0.6.tgz", + "integrity": "sha512-/olerl1A4sOqdP+hjvJ1sbQjKN07Y3DVnxO4gnbn/ahtQvFrdhUi0G1VsZXDNjfqmXw57DmPi5ASnj/8PGZhAA==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@vscode/vsce-sign-win32-arm64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-win32-arm64/-/vsce-sign-win32-arm64-2.0.6.tgz", + "integrity": "sha512-ivM/MiGIY0PJNZBoGtlRBM/xDpwbdlCWomUWuLmIxbi1Cxe/1nooYrEQoaHD8ojVRgzdQEUzMsRbyF5cJJgYOg==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/@vscode/vsce-sign-win32-x64": { + "version": "2.0.6", + "resolved": "https://registry.npmjs.org/@vscode/vsce-sign-win32-x64/-/vsce-sign-win32-x64-2.0.6.tgz", + "integrity": "sha512-mgth9Kvze+u8CruYMmhHw6Zgy3GRX2S+Ed5oSokDEK5vPEwGGKnmuXua9tmFhomeAnhgJnL4DCna3TiNuGrBTQ==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "SEE LICENSE IN LICENSE.txt", + "optional": true, + "os": [ + "win32" + ] + }, + "node_modules/agent-base": { + "version": "7.1.4", + "resolved": "https://registry.npmjs.org/agent-base/-/agent-base-7.1.4.tgz", + "integrity": "sha512-MnA+YT8fwfJPgBx3m60MNqakm30XOkyIoH1y6huTQvC0PwZG7ki8NacLBcrPbNoo8vEZy7Jpuk7+jMO+CUovTQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 14" + } + }, + "node_modules/assertion-error": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/assertion-error/-/assertion-error-2.0.1.tgz", + "integrity": "sha512-Izi8RQcffqCeNVgFigKli1ssklIbpHnCYc6AknXGYoB6grJqyeby7jv12JUQgmTAnIDnbck1uxksT4dzN3PWBA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=12" + } + }, + "node_modules/azure-devops-node-api": { + "version": "12.5.0", + "resolved": "https://registry.npmjs.org/azure-devops-node-api/-/azure-devops-node-api-12.5.0.tgz", + "integrity": "sha512-R5eFskGvOm3U/GzeAuxRkUsAl0hrAwGgWn6zAd2KrZmrEhWZVqLew4OOupbQlXUuojUzpGtq62SmdhJ06N88og==", + "dev": true, + "license": "MIT", + "dependencies": { + "tunnel": "0.0.6", + "typed-rest-client": "^1.8.4" + } + }, + "node_modules/balanced-match": { + "version": "4.0.4", + "resolved": "https://registry.npmjs.org/balanced-match/-/balanced-match-4.0.4.tgz", + "integrity": "sha512-BLrgEcRTwX2o6gGxGOCNyMvGSp35YofuYzw9h1IMTRmKqttAZZVU67bdb9Pr2vUHA8+j3i2tJfjO6C6+4myGTA==", + "dev": true, + "license": "MIT", + "engines": { + "node": "18 || 20 || >=22" + } + }, + "node_modules/binaryextensions": { + "version": "6.11.0", + "resolved": "https://registry.npmjs.org/binaryextensions/-/binaryextensions-6.11.0.tgz", + "integrity": "sha512-sXnYK/Ij80TO3lcqZVV2YgfKN5QjUWIRk/XSm2J/4bd/lPko3lvk0O4ZppH6m+6hB2/GTu+ptNwVFe1xh+QLQw==", + "dev": true, + "license": "Artistic-2.0", + "dependencies": { + "editions": "^6.21.0" + }, + "engines": { + "node": ">=4" + }, + "funding": { + "url": "https://bevry.me/fund" + } + }, + "node_modules/boundary": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/boundary/-/boundary-2.0.0.tgz", + "integrity": "sha512-rJKn5ooC9u8q13IMCrW0RSp31pxBCHE3y9V/tp3TdWSLf8Em3p6Di4NBpfzbJge9YjjFEsD0RtFEjtvHL5VyEA==", + "dev": true, + "license": "BSD-2-Clause" + }, + "node_modules/brace-expansion": { + "version": "5.0.12", + "resolved": "https://registry.npmjs.org/brace-expansion/-/brace-expansion-5.0.12.tgz", + "integrity": "sha512-YovQ3rzhaLMIrDjNDMkNS01tea93qhEhG5xy8f6+R0l+dw3Ki+5sCoIoI942iuLZTHWogWktgwVDhU09iNEimQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "balanced-match": "^4.0.2" + }, + "engines": { + "node": "20 || >=22" + } + }, + "node_modules/buffer-crc32": { + "version": "0.2.13", + "resolved": "https://registry.npmjs.org/buffer-crc32/-/buffer-crc32-0.2.13.tgz", + "integrity": "sha512-VO9Ht/+p3SN7SKWqcrgEzjGbRSJYTx+Q1pTQC0wrWqHx0vpJraQ6GtHx8tvcg1rlK1byhU5gccxgOgj7B0TDkQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": "*" + } + }, + "node_modules/buffer-equal-constant-time": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/buffer-equal-constant-time/-/buffer-equal-constant-time-1.0.1.tgz", + "integrity": "sha512-zRpUiDwd/xk6ADqPMATG8vc9VPrkck7T07OIx0gnjmJAnHnTVXNQG3vfvWNuiZIkwu9KrKdA1iJKfsfTVxE6NA==", + "dev": true, + "license": "BSD-3-Clause" + }, + "node_modules/bundle-name": { + "version": "4.1.0", + "resolved": "https://registry.npmjs.org/bundle-name/-/bundle-name-4.1.0.tgz", + "integrity": "sha512-tjwM5exMg6BGRI+kNmTntNsvdZS1X8BFYS6tnJ2hdH0kVxM6/eVZ2xy+FqStSWvYmtfFMDLIxurorHwDKfDz5Q==", + "dev": true, + "license": "MIT", + "dependencies": { + "run-applescript": "^7.0.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/call-bind-apply-helpers": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/call-bind-apply-helpers/-/call-bind-apply-helpers-1.0.2.tgz", + "integrity": "sha512-Sp1ablJ0ivDkSzjcaJdxEunN5/XvksFJ2sMBFfq6x0ryhQV/2b/KwFe21cMpmHtPOSij8K99/wSfoEuTObmuMQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "function-bind": "^1.1.2" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/call-bound": { + "version": "1.0.4", + "resolved": "https://registry.npmjs.org/call-bound/-/call-bound-1.0.4.tgz", + "integrity": "sha512-+ys997U96po4Kx/ABpBCqhA9EuxJaQWDQg7295H4hBphv3IZg0boBKuwYpt4YXp6MZ5AmZQnU/tyMTlRpaSejg==", + "dev": true, + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.2", + "get-intrinsic": "^1.3.0" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/chai": { + "version": "6.2.2", + "resolved": "https://registry.npmjs.org/chai/-/chai-6.2.2.tgz", + "integrity": "sha512-NUPRluOfOiTKBKvWPtSD4PhFvWCqOi0BGStNWs57X9js7XGTprSmFoz5F0tWhR4WPjNeR9jXqdC7/UpSJTnlRg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + } + }, + "node_modules/cockatiel": { + "version": "3.2.1", + "resolved": "https://registry.npmjs.org/cockatiel/-/cockatiel-3.2.1.tgz", + "integrity": "sha512-gfrHV6ZPkquExvMh9IOkKsBzNDk6sDuZ6DdBGUBkvFnTCqCxzpuq48RySgP0AnaqQkw2zynOFj9yly6T1Q2G5Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=16" + } + }, + "node_modules/commander": { + "version": "12.1.0", + "resolved": "https://registry.npmjs.org/commander/-/commander-12.1.0.tgz", + "integrity": "sha512-Vw8qHK3bZM9y/P10u3Vib8o/DdkvA2OtPtZvD871QKjy74Wj1WSKFILMPRPSdUSx5RFK1arlJzEtA4PkFgnbuA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + } + }, + "node_modules/convert-source-map": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/convert-source-map/-/convert-source-map-2.0.0.tgz", + "integrity": "sha512-Kvp459HrV2FEJ1CAsi1Ku+MY3kasH19TFykTz2xWmMeq6bk2NU3XXvfJ+Q61m0xktWwt+1HSYf3JZsTms3aRJg==", + "dev": true, + "license": "MIT" + }, + "node_modules/debug": { + "version": "4.4.3", + "resolved": "https://registry.npmjs.org/debug/-/debug-4.4.3.tgz", + "integrity": "sha512-RGwwWnwQvkVfavKVt22FGLw+xYSdzARwm0ru6DhTVA3umU5hZc28V3kO4stgYryrTlLpuvgI9GiijltAjNbcqA==", + "dev": true, + "license": "MIT", + "dependencies": { + "ms": "^2.1.3" + }, + "engines": { + "node": ">=6.0" + }, + "peerDependenciesMeta": { + "supports-color": { + "optional": true + } + } + }, + "node_modules/default-browser": { + "version": "5.5.1", + "resolved": "https://registry.npmjs.org/default-browser/-/default-browser-5.5.1.tgz", + "integrity": "sha512-m1pAzaJgZ/gssEqlOhJkPJp8Xly7QyW6xcrkUa2KKcDeDSEMP7X8xipU3snUcfisTQx0w1AGae+9UtJSfVnXGw==", + "dev": true, + "license": "MIT", + "dependencies": { + "bundle-name": "^4.1.0", + "default-browser-id": "^5.0.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/default-browser-id": { + "version": "5.0.1", + "resolved": "https://registry.npmjs.org/default-browser-id/-/default-browser-id-5.0.1.tgz", + "integrity": "sha512-x1VCxdX4t+8wVfd1so/9w+vQ4vx7lKd2Qp5tDRutErwmR85OgmfX7RlLRMWafRMY7hbEiXIbudNrjOAPa/hL8Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/define-lazy-prop": { + "version": "3.0.0", + "resolved": "https://registry.npmjs.org/define-lazy-prop/-/define-lazy-prop-3.0.0.tgz", + "integrity": "sha512-N+MeXYoqr3pOgn8xfyRPREN7gHakLYjhsHhWGT3fWAiL4IkAt0iDw14QiiEm2bE30c5XX5q0FtAA3CK5f9/BUg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=12" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/detect-libc": { + "version": "2.1.2", + "resolved": "https://registry.npmjs.org/detect-libc/-/detect-libc-2.1.2.tgz", + "integrity": "sha512-Btj2BOOO83o3WyH59e8MgXsxEQVcarkUOpEYrubB0urwnN10yQ364rsiByU11nZlqWYZm05i/of7io4mzihBtQ==", + "dev": true, + "license": "Apache-2.0", + "engines": { + "node": ">=8" + } + }, + "node_modules/dunder-proto": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/dunder-proto/-/dunder-proto-1.0.1.tgz", + "integrity": "sha512-KIN/nDJBQRcXw0MLVhZE9iQHmG68qAVIBg9CqmUYjmQIhgij9U5MFvrqkUL5FbtyyzZuOeOt0zdeRe4UY7ct+A==", + "dev": true, + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.1", + "es-errors": "^1.3.0", + "gopd": "^1.2.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/ecdsa-sig-formatter": { + "version": "1.0.11", + "resolved": "https://registry.npmjs.org/ecdsa-sig-formatter/-/ecdsa-sig-formatter-1.0.11.tgz", + "integrity": "sha512-nagl3RYrbNv6kQkeJIpt6NJZy8twLB/2vtz6yN9Z4vRKHN4/QZJIEbqohALSgwKdnksuY3k5Addp5lg8sVoVcQ==", + "dev": true, + "license": "Apache-2.0", + "dependencies": { + "safe-buffer": "^5.0.1" + } + }, + "node_modules/editions": { + "version": "6.22.0", + "resolved": "https://registry.npmjs.org/editions/-/editions-6.22.0.tgz", + "integrity": "sha512-UgGlf8IW75je7HZjNDpJdCv4cGJWIi6yumFdZ0R7A8/CIhQiWUjyGLCxdHpd8bmyD1gnkfUNK0oeOXqUS2cpfQ==", + "dev": true, + "license": "Artistic-2.0", + "dependencies": { + "version-range": "^4.15.0" + }, + "engines": { + "ecmascript": ">= es5", + "node": ">=4" + }, + "funding": { + "url": "https://bevry.me/fund" + } + }, + "node_modules/entities": { + "version": "8.1.0", + "resolved": "https://registry.npmjs.org/entities/-/entities-8.1.0.tgz", + "integrity": "sha512-kxL7msIffSuh9aaFAMD7rxAIuTRMAHMeBtgHW2yUdWw732ZNh4MehkF2gdjvtdmikkaIP9bFDDJOPlsvm7avrA==", + "dev": true, + "license": "BSD-2-Clause", + "engines": { + "node": ">=20.19.0" + }, + "funding": { + "url": "https://github.com/fb55/entities?sponsor=1" + } + }, + "node_modules/es-define-property": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/es-define-property/-/es-define-property-1.0.1.tgz", + "integrity": "sha512-e3nRfgfUZ4rNGL232gUgX06QNyyez04KdjFrF+LTRoOXmrOgFKDg4BCdsjW8EnT69eqdYGmRpJwiPVYNrCaW3g==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/es-errors": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/es-errors/-/es-errors-1.3.0.tgz", + "integrity": "sha512-Zf5H2Kxt2xjTvbJvP2ZWLEICxA6j+hAmMzIlypy4xcBg1vKVnx89Wy0GbS+kf5cwCVFFzdCFh2XSCFNULS6csw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/es-module-lexer": { + "version": "2.3.2", + "resolved": "https://registry.npmjs.org/es-module-lexer/-/es-module-lexer-2.3.2.tgz", + "integrity": "sha512-poHGpORABojJJucnV9KbOavETW8lBVnphkW77ER5/BQ5Fz7oXSoCNek7IH3vR5nRjdsEz926ibFYX8KtLQmdyw==", + "dev": true, + "license": "MIT" + }, + "node_modules/es-object-atoms": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/es-object-atoms/-/es-object-atoms-1.1.2.tgz", + "integrity": "sha512-HWcBoN6NileqtSydK2FqHbS/LoDd2pqrnQHLyJzBj4kOp/ky2MWMN694xOfkK8/SnUsW2DH7EfyVlydKCsm1Zw==", + "dev": true, + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/esbuild": { + "version": "0.28.2", + "resolved": "https://registry.npmjs.org/esbuild/-/esbuild-0.28.2.tgz", + "integrity": "sha512-HKVLS8dvII+xoKW9kmqxbRKrnWEXfJJr/FZhhJmiqIB0e053QNYFqOBouTMO/k5sID4MvCiUCvv8b9M4h32wIA==", + "dev": true, + "hasInstallScript": true, + "license": "MIT", + "bin": { + "esbuild": "bin/esbuild" + }, + "engines": { + "node": ">=18" + }, + "optionalDependencies": { + "@esbuild/aix-ppc64": "0.28.2", + "@esbuild/android-arm": "0.28.2", + "@esbuild/android-arm64": "0.28.2", + "@esbuild/android-x64": "0.28.2", + "@esbuild/darwin-arm64": "0.28.2", + "@esbuild/darwin-x64": "0.28.2", + "@esbuild/freebsd-arm64": "0.28.2", + "@esbuild/freebsd-x64": "0.28.2", + "@esbuild/linux-arm": "0.28.2", + "@esbuild/linux-arm64": "0.28.2", + "@esbuild/linux-ia32": "0.28.2", + "@esbuild/linux-loong64": "0.28.2", + "@esbuild/linux-mips64el": "0.28.2", + "@esbuild/linux-ppc64": "0.28.2", + "@esbuild/linux-riscv64": "0.28.2", + "@esbuild/linux-s390x": "0.28.2", + "@esbuild/linux-x64": "0.28.2", + "@esbuild/netbsd-arm64": "0.28.2", + "@esbuild/netbsd-x64": "0.28.2", + "@esbuild/openbsd-arm64": "0.28.2", + "@esbuild/openbsd-x64": "0.28.2", + "@esbuild/openharmony-arm64": "0.28.2", + "@esbuild/sunos-x64": "0.28.2", + "@esbuild/win32-arm64": "0.28.2", + "@esbuild/win32-ia32": "0.28.2", + "@esbuild/win32-x64": "0.28.2" + } + }, + "node_modules/estree-walker": { + "version": "3.0.3", + "resolved": "https://registry.npmjs.org/estree-walker/-/estree-walker-3.0.3.tgz", + "integrity": "sha512-7RUKfXgSMMkzt6ZuXmqapOurLGPPfgj6l9uRZ7lRGolvk0y2yocc35LdcxKC5PQZdn2DMqioAQ2NoWcrTKmm6g==", + "dev": true, + "license": "MIT", + "dependencies": { + "@types/estree": "^1.0.0" + } + }, + "node_modules/expect-type": { + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.4.0.tgz", + "integrity": "sha512-KfYbmpRm0VbLjEvVa9yGwCi9GI34xvi7A/HXYWQO65CSD2u3MczUJSuwXKFIxlGsgBQizV9q5J9NHj4VG0n+pA==", + "dev": true, + "license": "Apache-2.0", + "engines": { + "node": ">=12.0.0" + } + }, + "node_modules/fdir": { + "version": "6.5.0", + "resolved": "https://registry.npmjs.org/fdir/-/fdir-6.5.0.tgz", + "integrity": "sha512-tIbYtZbucOs0BRGqPJkshJUYdL+SDH7dVM8gjy+ERp3WAUjLEFJE+02kanyHtwjWOnwrKYBiwAmM0p4kLJAnXg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=12.0.0" + }, + "peerDependencies": { + "picomatch": "^3 || ^4" + }, + "peerDependenciesMeta": { + "picomatch": { + "optional": true + } + } + }, + "node_modules/fsevents": { + "version": "2.3.3", + "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.3.tgz", + "integrity": "sha512-5xoDfX+fL7faATnagmWPpbFtwh/R77WmMMqqHGS65C3vvB0YHrgF+B1YmZ3441tMj5n63k0212XNoJwzlhffQw==", + "dev": true, + "hasInstallScript": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": "^8.16.0 || ^10.6.0 || >=11.0.0" + } + }, + "node_modules/function-bind": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/function-bind/-/function-bind-1.1.2.tgz", + "integrity": "sha512-7XHNxH7qX9xG5mIwxkhumTox/MIRNcOgDrxWsMt2pAr23WHp6MrRlN7FBSFpCpr+oVO0F744iUgR82nJMfG2SA==", + "dev": true, + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/get-intrinsic": { + "version": "1.3.0", + "resolved": "https://registry.npmjs.org/get-intrinsic/-/get-intrinsic-1.3.0.tgz", + "integrity": "sha512-9fSjSaos/fRIVIp+xSJlE6lfwhES7LNtKaCBIamHsjr2na1BiABJPo0mOjjz8GJDURarmCPGqaiVg5mfjb98CQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "call-bind-apply-helpers": "^1.0.2", + "es-define-property": "^1.0.1", + "es-errors": "^1.3.0", + "es-object-atoms": "^1.1.1", + "function-bind": "^1.1.2", + "get-proto": "^1.0.1", + "gopd": "^1.2.0", + "has-symbols": "^1.1.0", + "hasown": "^2.0.2", + "math-intrinsics": "^1.1.0" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/get-proto": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/get-proto/-/get-proto-1.0.1.tgz", + "integrity": "sha512-sTSfBjoXBp89JvIKIefqw7U2CCebsc74kiY6awiGogKtoSGbgjYE/G/+l9sF3MWFPNc9IcoOC4ODfKHfxFmp0g==", + "dev": true, + "license": "MIT", + "dependencies": { + "dunder-proto": "^1.0.1", + "es-object-atoms": "^1.0.0" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/gopd": { + "version": "1.2.0", + "resolved": "https://registry.npmjs.org/gopd/-/gopd-1.2.0.tgz", + "integrity": "sha512-ZUKRh6/kUFoAiTAtTYPZJ3hw9wNxx+BIBOijnlG9PnrJsCcSjs1wyyD6vJpaYtgnzDrKYRSqf3OO6Rfa93xsRg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/graceful-fs": { + "version": "4.2.11", + "resolved": "https://registry.npmjs.org/graceful-fs/-/graceful-fs-4.2.11.tgz", + "integrity": "sha512-RbJ5/jmFcNNCcDV5o9eTnBLJ/HszWV0P73bc+Ff4nS/rJj+YaS6IGyiOL0VoBYX+l1Wrl3k63h/KrH+nhJ0XvQ==", + "dev": true, + "license": "ISC" + }, + "node_modules/has-symbols": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/has-symbols/-/has-symbols-1.1.0.tgz", + "integrity": "sha512-1cDNdwJ2Jaohmb3sg4OmKaMBwuC48sYni5HUw2DvsC8LjGTLK9h+eb1X6RyuOHe4hT0ULCW68iomhjUoKUqlPQ==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/hasown": { + "version": "2.0.4", + "resolved": "https://registry.npmjs.org/hasown/-/hasown-2.0.4.tgz", + "integrity": "sha512-T2UbfbBEF32wiepXIsMlTW9+dDYC6wMh/t/vYA4tuOMKqWz/n3vr1NFSxQiyP+zk2mXsoMA/i/7qV6LKut1t1A==", + "dev": true, + "license": "MIT", + "dependencies": { + "function-bind": "^1.1.2" + }, + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/hosted-git-info": { + "version": "4.1.0", + "resolved": "https://registry.npmjs.org/hosted-git-info/-/hosted-git-info-4.1.0.tgz", + "integrity": "sha512-kyCuEOWjJqZuDbRHzL8V93NzQhwIB71oFWSyzVo+KPZI+pnQPPxucdkrOZvkLRnrf5URsQM+IJ09Dw29cRALIA==", + "dev": true, + "license": "ISC", + "dependencies": { + "lru-cache": "^6.0.0" + }, + "engines": { + "node": ">=10" + } + }, + "node_modules/http-proxy-agent": { + "version": "7.0.2", + "resolved": "https://registry.npmjs.org/http-proxy-agent/-/http-proxy-agent-7.0.2.tgz", + "integrity": "sha512-T1gkAiYYDWYx3V5Bmyu7HcfcvL7mUrTWiM6yOfa3PIphViJ/gFPbvidQ+veqSOHci/PxBcDabeUNCzpOODJZig==", + "dev": true, + "license": "MIT", + "dependencies": { + "agent-base": "^7.1.0", + "debug": "^4.3.4" + }, + "engines": { + "node": ">= 14" + } + }, + "node_modules/https-proxy-agent": { + "version": "7.0.6", + "resolved": "https://registry.npmjs.org/https-proxy-agent/-/https-proxy-agent-7.0.6.tgz", + "integrity": "sha512-vK9P5/iUfdl95AI+JVyUuIcVtd4ofvtrOr3HNtM2yxC9bnMbEdp3x01OhQNnjb8IJYi38VlTE3mBXwcfvywuSw==", + "dev": true, + "license": "MIT", + "dependencies": { + "agent-base": "^7.1.2", + "debug": "4" + }, + "engines": { + "node": ">= 14" + } + }, + "node_modules/is-docker": { + "version": "3.0.0", + "resolved": "https://registry.npmjs.org/is-docker/-/is-docker-3.0.0.tgz", + "integrity": "sha512-eljcgEDlEns/7AXFosB5K/2nCM4P7FQPkGc/DWLy5rmFEWvZayGrik1d9/QIY5nJ4f9YsVvBkA6kJpHn9rISdQ==", + "dev": true, + "license": "MIT", + "bin": { + "is-docker": "cli.js" + }, + "engines": { + "node": "^12.20.0 || ^14.13.1 || >=16.0.0" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/is-inside-container": { + "version": "1.0.0", + "resolved": "https://registry.npmjs.org/is-inside-container/-/is-inside-container-1.0.0.tgz", + "integrity": "sha512-KIYLCCJghfHZxqjYBE7rEy0OBuTd5xCHS7tHVgvCLkx7StIoaxwNW3hCALgEUjFfeRk+MG/Qxmp/vtETEF3tRA==", + "dev": true, + "license": "MIT", + "dependencies": { + "is-docker": "^3.0.0" + }, + "bin": { + "is-inside-container": "cli.js" + }, + "engines": { + "node": ">=14.16" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/is-wsl": { + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/is-wsl/-/is-wsl-3.1.1.tgz", + "integrity": "sha512-e6rvdUCiQCAuumZslxRJWR/Doq4VpPR82kqclvcS0efgt430SlGIk05vdCN58+VrzgtIcfNODjozVielycD4Sw==", + "dev": true, + "license": "MIT", + "dependencies": { + "is-inside-container": "^1.0.0" + }, + "engines": { + "node": ">=16" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/istextorbinary": { + "version": "9.5.0", + "resolved": "https://registry.npmjs.org/istextorbinary/-/istextorbinary-9.5.0.tgz", + "integrity": "sha512-5mbUj3SiZXCuRf9fT3ibzbSSEWiy63gFfksmGfdOzujPjW3k+z8WvIBxcJHBoQNlaZaiyB25deviif2+osLmLw==", + "dev": true, + "license": "Artistic-2.0", + "dependencies": { + "binaryextensions": "^6.11.0", + "editions": "^6.21.0", + "textextensions": "^6.11.0" + }, + "engines": { + "node": ">=4" + }, + "funding": { + "url": "https://bevry.me/fund" + } + }, + "node_modules/jsonc-parser": { + "version": "3.3.1", + "resolved": "https://registry.npmjs.org/jsonc-parser/-/jsonc-parser-3.3.1.tgz", + "integrity": "sha512-HUgH65KyejrUFPvHFPbqOY0rsFip3Bo5wb4ngvdi1EpCYWUQDC5V+Y7mZws+DLkr4M//zQJoanu1SP+87Dv1oQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/jsonwebtoken": { + "version": "9.0.3", + "resolved": "https://registry.npmjs.org/jsonwebtoken/-/jsonwebtoken-9.0.3.tgz", + "integrity": "sha512-MT/xP0CrubFRNLNKvxJ2BYfy53Zkm++5bX9dtuPbqAeQpTVe0MQTFhao8+Cp//EmJp244xt6Drw/GVEGCUj40g==", + "dev": true, + "license": "MIT", + "dependencies": { + "jws": "^4.0.1", + "lodash.includes": "^4.3.0", + "lodash.isboolean": "^3.0.3", + "lodash.isinteger": "^4.0.4", + "lodash.isnumber": "^3.0.3", + "lodash.isplainobject": "^4.0.6", + "lodash.isstring": "^4.0.1", + "lodash.once": "^4.0.0", + "ms": "^2.1.1", + "semver": "^7.5.4" + }, + "engines": { + "node": ">=12", + "npm": ">=6" + } + }, + "node_modules/jwa": { + "version": "2.0.1", + "resolved": "https://registry.npmjs.org/jwa/-/jwa-2.0.1.tgz", + "integrity": "sha512-hRF04fqJIP8Abbkq5NKGN0Bbr3JxlQ+qhZufXVr0DvujKy93ZCbXZMHDL4EOtodSbCWxOqR8MS1tXA5hwqCXDg==", + "dev": true, + "license": "MIT", + "dependencies": { + "buffer-equal-constant-time": "^1.0.1", + "ecdsa-sig-formatter": "1.0.11", + "safe-buffer": "^5.0.1" + } + }, + "node_modules/jws": { + "version": "4.0.1", + "resolved": "https://registry.npmjs.org/jws/-/jws-4.0.1.tgz", + "integrity": "sha512-EKI/M/yqPncGUUh44xz0PxSidXFr/+r0pA70+gIYhjv+et7yxM+s29Y+VGDkovRofQem0fs7Uvf4+YmAdyRduA==", + "dev": true, + "license": "MIT", + "dependencies": { + "jwa": "^2.0.1", + "safe-buffer": "^5.0.1" + } + }, + "node_modules/lightningcss": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss/-/lightningcss-1.33.0.tgz", + "integrity": "sha512-WkUDrojuJs0xkgGf2udWxa3yGBRxPtxUkB79i6aCZLRgc7PM8fZe9TosfPDcvEpQZbuFASnHYmRLBLUbmLOIIA==", + "dev": true, + "license": "MPL-2.0", + "dependencies": { + "detect-libc": "^2.0.3" + }, + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + }, + "optionalDependencies": { + "lightningcss-android-arm64": "1.33.0", + "lightningcss-darwin-arm64": "1.33.0", + "lightningcss-darwin-x64": "1.33.0", + "lightningcss-freebsd-x64": "1.33.0", + "lightningcss-linux-arm-gnueabihf": "1.33.0", + "lightningcss-linux-arm64-gnu": "1.33.0", + "lightningcss-linux-arm64-musl": "1.33.0", + "lightningcss-linux-x64-gnu": "1.33.0", + "lightningcss-linux-x64-musl": "1.33.0", + "lightningcss-win32-arm64-msvc": "1.33.0", + "lightningcss-win32-x64-msvc": "1.33.0" + } + }, + "node_modules/lightningcss-android-arm64": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-android-arm64/-/lightningcss-android-arm64-1.33.0.tgz", + "integrity": "sha512-gEpRTalKdosp4Bb8qWtc2iOgE5SeIHlpS1up9bFq2wAyYhl1UdTObYiHe98zEM9SQvSoqQZ1IQD0JNpg3Ml5pg==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "android" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-darwin-arm64": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-darwin-arm64/-/lightningcss-darwin-arm64-1.33.0.tgz", + "integrity": "sha512-Sciaz8eenNTKn9b3t7+xr0ipTp9YxKQY4npwQ3mrRuL0BAVHBLyZxofhaKBAVtzmtRZ/zTyo0/to4B1uWG/Djg==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-darwin-x64": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-darwin-x64/-/lightningcss-darwin-x64-1.33.0.tgz", + "integrity": "sha512-Z5UPAxzrjlWNNyGy6i65cJzzvgJ5D3T6wMvs+gWpY9d7qRhANrxqAp6LhxIgZhWEw18RfJTGcRxjuLIBr+m8XQ==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-freebsd-x64": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-freebsd-x64/-/lightningcss-freebsd-x64-1.33.0.tgz", + "integrity": "sha512-QQM/Ti/hQajJwCY+RiWuCZ9sdtI/XQk7nDK5vC8kkdwixezOlDgvDx7+RT+QjK6FcFT4MpsuoBnHIo/O3StRRg==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "freebsd" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-linux-arm-gnueabihf": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-linux-arm-gnueabihf/-/lightningcss-linux-arm-gnueabihf-1.33.0.tgz", + "integrity": "sha512-N7FVBe6iS24MlM6R/4RBTxGhQheZGs7tiQ9U32UtF75NzP5Q7xWPRqLBCKxlRQRk3rY1jCIPLzx7WzOhuUIRLQ==", + "cpu": [ + "arm" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-linux-arm64-gnu": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-linux-arm64-gnu/-/lightningcss-linux-arm64-gnu-1.33.0.tgz", + "integrity": "sha512-j2v/itmy4HlNxlc6voKXYgBqNi0Ng2LShg4z7GufpEgs05P+2suBVyi9I6YHq5uoVFx9ETin3eCEhLVyXGQnKg==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MPL-2.0", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-linux-arm64-musl": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-linux-arm64-musl/-/lightningcss-linux-arm64-musl-1.33.0.tgz", + "integrity": "sha512-yiO5ROMuYQgXbC60yjZU5CYSFZGKXL0HFATXt9mHJn1+zW55oCtMI9NfcVhYLMFDL7gV7oBPon/EmMMGg2OvtQ==", + "cpu": [ + "arm64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MPL-2.0", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-linux-x64-gnu": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-linux-x64-gnu/-/lightningcss-linux-x64-gnu-1.33.0.tgz", + "integrity": "sha512-ar+Ju7LmcN0Jo4FpL4hpFybwNG9/3A/Br5KW2n2jyODg3MEZXaDYADdemoNS+BDNfMgKvylJLj4S5tyRActuAg==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "glibc" + ], + "license": "MPL-2.0", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-linux-x64-musl": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-linux-x64-musl/-/lightningcss-linux-x64-musl-1.33.0.tgz", + "integrity": "sha512-RYiYbkokw0trfKqqzfF55lginwEPrD3OJDfTuJzFs1MK6iFnDenaz1fqLLtX4ITG3OktJQXOeTaw1awrBAlZPw==", + "cpu": [ + "x64" + ], + "dev": true, + "libc": [ + "musl" + ], + "license": "MPL-2.0", + "optional": true, + "os": [ + "linux" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-win32-arm64-msvc": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-win32-arm64-msvc/-/lightningcss-win32-arm64-msvc-1.33.0.tgz", + "integrity": "sha512-1K+MPfLSFVpphzpdbfkhlWk6wBrTObBzS2T6db10PNOZgR9GoVsAWzwNyuhUYYbTp23j+4RrncfujZ4uAzXvwA==", + "cpu": [ + "arm64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lightningcss-win32-x64-msvc": { + "version": "1.33.0", + "resolved": "https://registry.npmjs.org/lightningcss-win32-x64-msvc/-/lightningcss-win32-x64-msvc-1.33.0.tgz", + "integrity": "sha512-OlEICDx/Xl0FqSp4bry8zFnCvGpig3Gl4gCquvYwHuqJKEC1+n9NgDniFvqHGmMv1ZkqDJrDqKKSykTDX+ehuA==", + "cpu": [ + "x64" + ], + "dev": true, + "license": "MPL-2.0", + "optional": true, + "os": [ + "win32" + ], + "engines": { + "node": ">= 12.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/parcel" + } + }, + "node_modules/lodash.includes": { + "version": "4.3.0", + "resolved": "https://registry.npmjs.org/lodash.includes/-/lodash.includes-4.3.0.tgz", + "integrity": "sha512-W3Bx6mdkRTGtlJISOvVD/lbqjTlPPUDTMnlXZFnVwi9NKJ6tiAk6LVdlhZMm17VZisqhKcgzpO5Wz91PCt5b0w==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.isboolean": { + "version": "3.0.3", + "resolved": "https://registry.npmjs.org/lodash.isboolean/-/lodash.isboolean-3.0.3.tgz", + "integrity": "sha512-Bz5mupy2SVbPHURB98VAcw+aHh4vRV5IPNhILUCsOzRmsTmSQ17jIuqopAentWoehktxGd9e/hbIXq980/1QJg==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.isinteger": { + "version": "4.0.4", + "resolved": "https://registry.npmjs.org/lodash.isinteger/-/lodash.isinteger-4.0.4.tgz", + "integrity": "sha512-DBwtEWN2caHQ9/imiNeEA5ys1JoRtRfY3d7V9wkqtbycnAmTvRRmbHKDV4a0EYc678/dia0jrte4tjYwVBaZUA==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.isnumber": { + "version": "3.0.3", + "resolved": "https://registry.npmjs.org/lodash.isnumber/-/lodash.isnumber-3.0.3.tgz", + "integrity": "sha512-QYqzpfwO3/CWf3XP+Z+tkQsfaLL/EnUlXWVkIk5FUPc4sBdTehEqZONuyRt2P67PXAk+NXmTBcc97zw9t1FQrw==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.isplainobject": { + "version": "4.0.6", + "resolved": "https://registry.npmjs.org/lodash.isplainobject/-/lodash.isplainobject-4.0.6.tgz", + "integrity": "sha512-oSXzaWypCMHkPC3NvBEaPHf0KsA5mvPrOPgQWDsbg8n7orZ290M0BmC/jgRZ4vcJ6DTAhjrsSYgdsW/F+MFOBA==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.isstring": { + "version": "4.0.1", + "resolved": "https://registry.npmjs.org/lodash.isstring/-/lodash.isstring-4.0.1.tgz", + "integrity": "sha512-0wJxfxH1wgO3GrbuP+dTTk7op+6L41QCXbGINEmD+ny/G/eCqGzxyCsh7159S+mgDDcoarnBw6PC1PS5+wUGgw==", + "dev": true, + "license": "MIT" + }, + "node_modules/lodash.once": { + "version": "4.1.1", + "resolved": "https://registry.npmjs.org/lodash.once/-/lodash.once-4.1.1.tgz", + "integrity": "sha512-Sb487aTOCr9drQVL8pIxOzVhafOjZN9UU54hiN8PU3uAiSV7lx1yYNpbNmex2PK6dSJoNTSJUUswT651yww3Mg==", + "dev": true, + "license": "MIT" + }, + "node_modules/lru-cache": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-6.0.0.tgz", + "integrity": "sha512-Jo6dJ04CmSjuznwJSS3pUeWmd/H0ffTlkXXgwZi+eq1UCmqQwCh+eLsYOYCwY991i2Fah4h1BEMCx4qThGbsiA==", + "dev": true, + "license": "ISC", + "dependencies": { + "yallist": "^4.0.0" + }, + "engines": { + "node": ">=10" + } + }, + "node_modules/magic-string": { + "version": "0.30.21", + "resolved": "https://registry.npmjs.org/magic-string/-/magic-string-0.30.21.tgz", + "integrity": "sha512-vd2F4YUyEXKGcLHoq+TEyCjxueSeHnFxyyjNp80yg0XV4vUhnDer/lvvlqM/arB5bXQN5K2/3oinyCRyx8T2CQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@jridgewell/sourcemap-codec": "^1.5.5" + } + }, + "node_modules/marked": { + "version": "18.0.13", + "resolved": "https://registry.npmjs.org/marked/-/marked-18.0.13.tgz", + "integrity": "sha512-xTxVzZsBFwunP6HDmtBkabUQEYArnP7/rMDGmPj9SlrKlQ4i8MdYVow+nJL0eOqwpUqhzBoTBRADGN6uYwPyOw==", + "dev": true, + "license": "MIT", + "bin": { + "marked": "bin/marked.js" + }, + "engines": { + "node": ">= 20" + } + }, + "node_modules/math-intrinsics": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz", + "integrity": "sha512-/IXtbwEk5HTPyEwyKX6hGkYXxM9nbj64B+ilVJnC/R6B0pH5G4V3b0pVbL7DBj4tkhBAppbQUlf6F6Xl9LHu1g==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + } + }, + "node_modules/mime": { + "version": "1.6.0", + "resolved": "https://registry.npmjs.org/mime/-/mime-1.6.0.tgz", + "integrity": "sha512-x0Vn8spI+wuJ1O6S7gnbaQg8Pxh4NNHb7KSINmEWKiPE4RKOplvijn+NkmYmmRgP68mc70j2EbeTFRsrswaQeg==", + "dev": true, + "license": "MIT", + "bin": { + "mime": "cli.js" + }, + "engines": { + "node": ">=4" + } + }, + "node_modules/minimatch": { + "version": "10.2.6", + "resolved": "https://registry.npmjs.org/minimatch/-/minimatch-10.2.6.tgz", + "integrity": "sha512-vpLQEs+VLCr1nU0BXS07maYoFwlDAH0gngQuuttxIwutDFEMHq2blX+8vpgxDdK3J1PwjCJiep77OitTZ4Ll1A==", + "dev": true, + "license": "BlueOak-1.0.0", + "dependencies": { + "brace-expansion": "^5.0.8" + }, + "engines": { + "node": "18 || 20 || >=22" + }, + "funding": { + "url": "https://github.com/sponsors/isaacs" + } + }, + "node_modules/ms": { + "version": "2.1.3", + "resolved": "https://registry.npmjs.org/ms/-/ms-2.1.3.tgz", + "integrity": "sha512-6FlzubTLZG3J2a/NVCAleEhjzq5oxgHyaCU9yYXvcLsvoVaHJq/s5xXI6/XXP6tz7R9xAOtHnSO/tXtF3WRTlA==", + "dev": true, + "license": "MIT" + }, + "node_modules/mute-stream": { + "version": "0.0.8", + "resolved": "https://registry.npmjs.org/mute-stream/-/mute-stream-0.0.8.tgz", + "integrity": "sha512-nnbWWOkoWyUsTjKrhgD0dcz22mdkSnpYqbEjIm2nhwhuxlSkpywJmBo8h0ZqJdkp73mb90SssHkN4rsRaBAfAA==", + "dev": true, + "license": "ISC" + }, + "node_modules/nanoid": { + "version": "3.3.19", + "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.19.tgz", + "integrity": "sha512-Y2tUNy4ouw6tq5oDSKeQYGOyhkUBhNOcGV/02KC+6kd9eDGqdZd++mjMiIDilrBYvjEnCYvVtsuHCuP+okSfug==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/ai" + } + ], + "license": "MIT", + "bin": { + "nanoid": "bin/nanoid.cjs" + }, + "engines": { + "node": "^10 || ^12 || ^13.7 || ^14 || >=15.0.1" + } + }, + "node_modules/object-inspect": { + "version": "1.13.4", + "resolved": "https://registry.npmjs.org/object-inspect/-/object-inspect-1.13.4.tgz", + "integrity": "sha512-W67iLl4J2EXEGTbfeHCffrjDfitvLANg0UlX3wFUUSTx92KXRFegMHUVgSqE+wvhAbi4WqjGg9czysTV2Epbew==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/obug": { + "version": "2.2.1", + "resolved": "https://registry.npmjs.org/obug/-/obug-2.2.1.tgz", + "integrity": "sha512-XrsrhT5sybtKI6wakr2SPOlGZWWYbUXZ7a0jT8/QOeAPau+1X/bSegNe5YR75oJmEZQbKningirmGOEJCIk61Q==", + "dev": true, + "funding": [ + "https://github.com/sponsors/sxzz", + "https://opencollective.com/debug" + ], + "license": "MIT", + "engines": { + "node": ">=12.20.0" + } + }, + "node_modules/open": { + "version": "10.2.0", + "resolved": "https://registry.npmjs.org/open/-/open-10.2.0.tgz", + "integrity": "sha512-YgBpdJHPyQ2UE5x+hlSXcnejzAvD0b22U2OuAP+8OnlJT+PjWPxtgmGqKKc+RgTM63U9gN0YzrYc71R2WT/hTA==", + "dev": true, + "license": "MIT", + "dependencies": { + "default-browser": "^5.2.1", + "define-lazy-prop": "^3.0.0", + "is-inside-container": "^1.0.0", + "wsl-utils": "^0.1.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/openai": { + "version": "7.18.0", + "resolved": "https://registry.npmjs.org/openai/-/openai-7.18.0.tgz", + "integrity": "sha512-S+xxaUf9VzIHEPDUpUFnRgvR2Ho0K1yZEaSJ8gd7xQlCNm4sdOD7/QaKWLRTYK3MK6KhwpD3cOEC0F4OPG698A==", + "license": "Apache-2.0", + "engines": { + "node": ">=22.0.0" + }, + "peerDependencies": { + "@aws-sdk/credential-provider-node": ">=3.972.0 <4", + "@smithy/hash-node": ">=4.3.0 <5", + "@smithy/signature-v4": ">=5.4.0 <6", + "undici": ">=5 <9", + "ws": "^8.21.0", + "zod": "^3.25 || ^4.0" + }, + "peerDependenciesMeta": { + "@aws-sdk/credential-provider-node": { + "optional": true + }, + "@smithy/hash-node": { + "optional": true + }, + "@smithy/signature-v4": { + "optional": true + }, + "undici": { + "optional": true + }, + "ws": { + "optional": true + }, + "zod": { + "optional": true + } + } + }, + "node_modules/parse5": { + "version": "8.0.1", + "resolved": "https://registry.npmjs.org/parse5/-/parse5-8.0.1.tgz", + "integrity": "sha512-z1e/HMG90obSGeidlli3hj7cbocou0/wa5HacvI3ASx34PecNjNQeaHNo5WIZpWofN9kgkqV1q5YvXe3F0FoPw==", + "dev": true, + "license": "MIT", + "dependencies": { + "entities": "^8.0.0" + }, + "funding": { + "url": "https://github.com/inikulin/parse5?sponsor=1" + } + }, + "node_modules/pathe": { + "version": "2.0.3", + "resolved": "https://registry.npmjs.org/pathe/-/pathe-2.0.3.tgz", + "integrity": "sha512-WUjGcAqP1gQacoQe+OBJsFA7Ld4DyXuUIjZ5cc75cLHvJ7dtNsTugphxIADwspS+AraAUePCKrSVtPLFj/F88w==", + "dev": true, + "license": "MIT" + }, + "node_modules/pend": { + "version": "1.2.0", + "resolved": "https://registry.npmjs.org/pend/-/pend-1.2.0.tgz", + "integrity": "sha512-F3asv42UuXchdzt+xXqfW1OGlVBe+mxa2mqI0pg5yAHZPvFmY3Y6drSf/GQ1A86WgWEN9Kzh/WrgKa6iGcHXLg==", + "dev": true, + "license": "MIT" + }, + "node_modules/picocolors": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/picocolors/-/picocolors-1.1.1.tgz", + "integrity": "sha512-xceH2snhtb5M9liqDsmEw56le376mTZkEX/jEb/RxNFyegNul7eNslCXP9FDj/Lcu0X8KEyMceP2ntpaHrDEVA==", + "dev": true, + "license": "ISC" + }, + "node_modules/picomatch": { + "version": "4.0.7", + "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.7.tgz", + "integrity": "sha512-qcJu88Q2IWqJsDD529JKMdwGm/dvInW4HvQnRwiH9JtihJvzGOscDtHE3x1pBKeUOTysQ8kVmLnJ2kJu7yhcGA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=12" + }, + "funding": { + "url": "https://github.com/sponsors/jonschlinkert" + } + }, + "node_modules/postcss": { + "version": "8.5.28", + "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.28.tgz", + "integrity": "sha512-RRuzqDtt5Y9h3quz5hWhK+TPnsmVs6WwSU6LkJMeY4HstUEDuYTG8UJSdawMRzmzAtV+KEoG8N3Qg2qLy5vM/A==", + "dev": true, + "funding": [ + { + "type": "opencollective", + "url": "https://opencollective.com/postcss/" + }, + { + "type": "tidelift", + "url": "https://tidelift.com/funding/github/npm/postcss" + }, + { + "type": "github", + "url": "https://github.com/sponsors/ai" + } + ], + "license": "MIT", + "dependencies": { + "nanoid": "^3.3.18", + "picocolors": "^1.1.1", + "source-map-js": "^1.2.1" + }, + "engines": { + "node": "^10 || ^12 || >=14" + } + }, + "node_modules/proper-lockfile": { + "version": "4.1.2", + "resolved": "https://registry.npmjs.org/proper-lockfile/-/proper-lockfile-4.1.2.tgz", + "integrity": "sha512-TjNPblN4BwAWMXU8s9AEz4JmQxnD1NNL7bNOY/AKUzyamc379FWASUhc/K1pL2noVb+XmZKLL68cjzLsiOAMaA==", + "dev": true, + "license": "MIT", + "dependencies": { + "graceful-fs": "^4.2.4", + "retry": "^0.12.0", + "signal-exit": "^3.0.2" + } + }, + "node_modules/qs": { + "version": "6.16.0", + "resolved": "https://registry.npmjs.org/qs/-/qs-6.16.0.tgz", + "integrity": "sha512-h6fhOIaRrID2CbEY2fqs+7t+UXZo+MLAnU5gRIq85uFtdiUPCdsApMlHhXogKVM4HM2DVbIjGNTTYH2OcmP1vA==", + "dev": true, + "license": "BSD-3-Clause", + "dependencies": { + "es-define-property": "^1.0.1", + "side-channel": "^1.1.1" + }, + "engines": { + "node": ">=0.6" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/read": { + "version": "1.0.7", + "resolved": "https://registry.npmjs.org/read/-/read-1.0.7.tgz", + "integrity": "sha512-rSOKNYUmaxy0om1BNjMN4ezNT6VKK+2xF4GBhc81mkH7L60i6dp8qPYrkndNLT3QPphoII3maL9PVC9XmhHwVQ==", + "dev": true, + "license": "ISC", + "dependencies": { + "mute-stream": "~0.0.4" + }, + "engines": { + "node": ">=0.8" + } + }, + "node_modules/retry": { + "version": "0.12.0", + "resolved": "https://registry.npmjs.org/retry/-/retry-0.12.0.tgz", + "integrity": "sha512-9LkiTwjUh6rT555DtE9rTX+BKByPfrMzEAtnlEtdEwr3Nkffwiihqe2bWADg+OQRjt9gl6ICdmB/ZFDCGAtSow==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">= 4" + } + }, + "node_modules/rolldown": { + "version": "1.2.9", + "resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.2.9.tgz", + "integrity": "sha512-hx/Pv0N1haXRb11qkfnK5MXB/iqr7i0yjWQqmO9uHqZpBgQSqzc8UsSnEpalsh+j1I8qQ2CkXAkJC8Br3dKSlg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@oxc-project/types": "=0.150.0", + "@rolldown/pluginutils": "^1.0.0" + }, + "bin": { + "rolldown": "bin/cli.mjs" + }, + "engines": { + "node": "^20.19.0 || >=22.12.0" + }, + "optionalDependencies": { + "@rolldown/binding-android-arm-eabi": "1.2.9", + "@rolldown/binding-android-arm64": "1.2.9", + "@rolldown/binding-darwin-arm64": "1.2.9", + "@rolldown/binding-darwin-x64": "1.2.9", + "@rolldown/binding-freebsd-x64": "1.2.9", + "@rolldown/binding-linux-arm-gnueabihf": "1.2.9", + "@rolldown/binding-linux-arm64-gnu": "1.2.9", + "@rolldown/binding-linux-arm64-musl": "1.2.9", + "@rolldown/binding-linux-ppc64-gnu": "1.2.9", + "@rolldown/binding-linux-s390x-gnu": "1.2.9", + "@rolldown/binding-linux-x64-gnu": "1.2.9", + "@rolldown/binding-linux-x64-musl": "1.2.9", + "@rolldown/binding-openharmony-arm64": "1.2.9", + "@rolldown/binding-win32-arm64-msvc": "1.2.9", + "@rolldown/binding-win32-x64-msvc": "1.2.9" + } + }, + "node_modules/run-applescript": { + "version": "7.1.0", + "resolved": "https://registry.npmjs.org/run-applescript/-/run-applescript-7.1.0.tgz", + "integrity": "sha512-DPe5pVFaAsinSaV6QjQ6gdiedWDcRCbUuiQfQa2wmWV7+xC9bGulGI8+TdRmoFkAPaBXk8CrAbnlY2ISniJ47Q==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/safe-buffer": { + "version": "5.2.1", + "resolved": "https://registry.npmjs.org/safe-buffer/-/safe-buffer-5.2.1.tgz", + "integrity": "sha512-rp3So07KcdmmKbGvgaNxQSJr7bGVSVk5S9Eq1F+ppbRo70+YeaDxkw5Dd8NPN+GD6bjnYm2VuPuCXmpuYvmCXQ==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/feross" + }, + { + "type": "patreon", + "url": "https://www.patreon.com/feross" + }, + { + "type": "consulting", + "url": "https://feross.org/support" + } + ], + "license": "MIT" + }, + "node_modules/sax": { + "version": "1.6.1", + "resolved": "https://registry.npmjs.org/sax/-/sax-1.6.1.tgz", + "integrity": "sha512-42tBVwLWnaQvW5zc4HbZrTuWccECCZfBi92FDuwtqxasH+JbPB3/FOKb1m222K42R4WxuxzzMsTswfzgtSu64Q==", + "dev": true, + "license": "BlueOak-1.0.0", + "engines": { + "node": ">=11.0.0" + } + }, + "node_modules/semver": { + "version": "7.8.5", + "resolved": "https://registry.npmjs.org/semver/-/semver-7.8.5.tgz", + "integrity": "sha512-Y7/KDsb8LjooZpwaqGyulO6DQlksgCncchHGk+sZIY4SBvUocMBEFH5Ur1fI4dV+Jvl0w6cjvucaIi40puRioA==", + "dev": true, + "license": "ISC", + "bin": { + "semver": "bin/semver.js" + }, + "engines": { + "node": ">=10" + } + }, + "node_modules/side-channel": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/side-channel/-/side-channel-1.1.1.tgz", + "integrity": "sha512-6x6dK6zJdpTzF4sQeNYxwtvBzf6Eg4GtlesS94HOvTudUeyK2WXAaIfmDgsyslYrRBeFIlsi54AYsFGUuhmvrQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "object-inspect": "^1.13.4", + "side-channel-list": "^1.0.1", + "side-channel-map": "^1.0.1", + "side-channel-weakmap": "^1.0.2" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-list": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/side-channel-list/-/side-channel-list-1.0.1.tgz", + "integrity": "sha512-mjn/0bi/oUURjc5Xl7IaWi/OJJJumuoJFQJfDDyO46+hBWsfaVM65TBHq2eoZBhzl9EchxOijpkbRC8SVBQU0w==", + "dev": true, + "license": "MIT", + "dependencies": { + "es-errors": "^1.3.0", + "object-inspect": "^1.13.4" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-map": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/side-channel-map/-/side-channel-map-1.0.1.tgz", + "integrity": "sha512-VCjCNfgMsby3tTdo02nbjtM/ewra6jPHmpThenkTYh8pG9ucZ/1P8So4u4FGBek/BjpOVsDCMoLA/iuBKIFXRA==", + "dev": true, + "license": "MIT", + "dependencies": { + "call-bound": "^1.0.2", + "es-errors": "^1.3.0", + "get-intrinsic": "^1.2.5", + "object-inspect": "^1.13.3" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/side-channel-weakmap": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/side-channel-weakmap/-/side-channel-weakmap-1.0.2.tgz", + "integrity": "sha512-WPS/HvHQTYnHisLo9McqBHOJk2FkHO/tlpvldyrnem4aeQp4hai3gythswg6p01oSoTl58rcpiFAjF2br2Ak2A==", + "dev": true, + "license": "MIT", + "dependencies": { + "call-bound": "^1.0.2", + "es-errors": "^1.3.0", + "get-intrinsic": "^1.2.5", + "object-inspect": "^1.13.3", + "side-channel-map": "^1.0.1" + }, + "engines": { + "node": ">= 0.4" + }, + "funding": { + "url": "https://github.com/sponsors/ljharb" + } + }, + "node_modules/siginfo": { + "version": "2.0.0", + "resolved": "https://registry.npmjs.org/siginfo/-/siginfo-2.0.0.tgz", + "integrity": "sha512-ybx0WO1/8bSBLEWXZvEd7gMW3Sn3JFlW3TvX1nREbDLRNQNaeNN8WK0meBwPdAaOI7TtRRRJn/Es1zhrrCHu7g==", + "dev": true, + "license": "ISC" + }, + "node_modules/signal-exit": { + "version": "3.0.7", + "resolved": "https://registry.npmjs.org/signal-exit/-/signal-exit-3.0.7.tgz", + "integrity": "sha512-wnD2ZE+l+SPC/uoS0vXeE9L1+0wuaMqKlfz9AMUo38JsyLSBWSFcHR1Rri62LZc12vLr1gb3jl7iwQhgwpAbGQ==", + "dev": true, + "license": "ISC" + }, + "node_modules/source-map-js": { + "version": "1.2.1", + "resolved": "https://registry.npmjs.org/source-map-js/-/source-map-js-1.2.1.tgz", + "integrity": "sha512-UXWMKhLOwVKb728IUtQPXxfYU+usdybtUrK/8uGE8CQMvrhOpwvzDBwj0QhSL7MQc7vIsISBG8VQ8+IDQxpfQA==", + "dev": true, + "license": "BSD-3-Clause", + "engines": { + "node": ">=0.10.0" + } + }, + "node_modules/stackback": { + "version": "0.0.2", + "resolved": "https://registry.npmjs.org/stackback/-/stackback-0.0.2.tgz", + "integrity": "sha512-1XMJE5fQo1jGH6Y/7ebnwPOBEkIEnT4QF32d5R1+VXdXveM0IBMJt8zfaxX1P3QhVwrYe+576+jkANtSS2mBbw==", + "dev": true, + "license": "MIT" + }, + "node_modules/std-env": { + "version": "4.2.0", + "resolved": "https://registry.npmjs.org/std-env/-/std-env-4.2.0.tgz", + "integrity": "sha512-oCUKSupKTHX53EyjDtuZQ64pjLJ6yYCtpmEw0goYxtjG9KpbRe8KAsl2tBUGU9DyMcJ0RwJ8GqJAFzMXcXW1Rw==", + "dev": true, + "license": "MIT" + }, + "node_modules/structured-source": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/structured-source/-/structured-source-4.0.0.tgz", + "integrity": "sha512-qGzRFNJDjFieQkl/sVOI2dUjHKRyL9dAJi2gCPGJLbJHBIkyOHxjuocpIEfbLioX+qSJpvbYdT49/YCdMznKxA==", + "dev": true, + "license": "BSD-2-Clause", + "dependencies": { + "boundary": "^2.0.0" + } + }, + "node_modules/textextensions": { + "version": "6.11.0", + "resolved": "https://registry.npmjs.org/textextensions/-/textextensions-6.11.0.tgz", + "integrity": "sha512-tXJwSr9355kFJI3lbCkPpUH5cP8/M0GGy2xLO34aZCjMXBaK3SoPnZwr/oWmo1FdCnELcs4npdCIOFtq9W3ruQ==", + "dev": true, + "license": "Artistic-2.0", + "dependencies": { + "editions": "^6.21.0" + }, + "engines": { + "node": ">=4" + }, + "funding": { + "url": "https://bevry.me/fund" + } + }, + "node_modules/tinybench": { + "version": "2.9.0", + "resolved": "https://registry.npmjs.org/tinybench/-/tinybench-2.9.0.tgz", + "integrity": "sha512-0+DUvqWMValLmha6lr4kD8iAMK1HzV0/aKnCtWb9v9641TnP/MFb7Pc2bxoxQjTXAErryXVgUOfv2YqNllqGeg==", + "dev": true, + "license": "MIT" + }, + "node_modules/tinyexec": { + "version": "1.3.1", + "resolved": "https://registry.npmjs.org/tinyexec/-/tinyexec-1.3.1.tgz", + "integrity": "sha512-GCvB3aoys96IuDFBMcTB46JOR6mdMtAToqwiW8JlWhsoh1mhHi/xn9ss/Dg7N555GiJyEt2qzoG/NHCwM6h1EA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=18" + } + }, + "node_modules/tinyglobby": { + "version": "0.2.17", + "resolved": "https://registry.npmjs.org/tinyglobby/-/tinyglobby-0.2.17.tgz", + "integrity": "sha512-wXR/dYpcqKmfWpEdZjiKJOwCNFndD0DMnrW/cYjVGttEkBfVgcLFHoNrlj47mjOVic9yyNu65alsgF4NQyTa2g==", + "dev": true, + "license": "MIT", + "dependencies": { + "fdir": "^6.5.0", + "picomatch": "^4.0.4" + }, + "engines": { + "node": ">=12.0.0" + }, + "funding": { + "url": "https://github.com/sponsors/SuperchupuDev" + } + }, + "node_modules/tinyrainbow": { + "version": "3.1.1", + "resolved": "https://registry.npmjs.org/tinyrainbow/-/tinyrainbow-3.1.1.tgz", + "integrity": "sha512-yau8yJdTt989Mm0Bd/236QnzEiPf2xLLTqUZRUJOo/3CB078LSwzei343DgtJVmfJKJE3TMINY1u42SQsP6mXw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=14.0.0" + } + }, + "node_modules/tslib": { + "version": "2.8.1", + "resolved": "https://registry.npmjs.org/tslib/-/tslib-2.8.1.tgz", + "integrity": "sha512-oJFu94HQb+KVduSUQL7wnpmqnfmLsOA/nAh6b6EH0wCEoK0/mPeXU6c3wKDV83MkOuHPRHtSXKKU99IBazS/2w==", + "dev": true, + "license": "0BSD" + }, + "node_modules/tunnel": { + "version": "0.0.6", + "resolved": "https://registry.npmjs.org/tunnel/-/tunnel-0.0.6.tgz", + "integrity": "sha512-1h/Lnq9yajKY2PEbBadPXj3VxsDDu844OnaAo52UVmIzIvwwtBPIuNvkjuzBlTWpfJyUbG3ez0KSBibQkj4ojg==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=0.6.11 <=0.7.0 || >=0.7.3" + } + }, + "node_modules/typed-rest-client": { + "version": "1.8.11", + "resolved": "https://registry.npmjs.org/typed-rest-client/-/typed-rest-client-1.8.11.tgz", + "integrity": "sha512-5UvfMpd1oelmUPRbbaVnq+rHP7ng2cE4qoQkQeAqxRL6PklkxsM0g32/HL0yfvruK6ojQ5x8EE+HF4YV6DtuCA==", + "dev": true, + "license": "MIT", + "dependencies": { + "qs": "^6.9.1", + "tunnel": "0.0.6", + "underscore": "^1.12.1" + } + }, + "node_modules/typescript": { + "version": "5.9.3", + "resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz", + "integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==", + "dev": true, + "license": "Apache-2.0", + "bin": { + "tsc": "bin/tsc", + "tsserver": "bin/tsserver" + }, + "engines": { + "node": ">=14.17" + } + }, + "node_modules/underscore": { + "version": "1.13.8", + "resolved": "https://registry.npmjs.org/underscore/-/underscore-1.13.8.tgz", + "integrity": "sha512-DXtD3ZtEQzc7M8m4cXotyHR+FAS18C64asBYY5vqZexfYryNNnDc02W4hKg3rdQuqOYas1jkseX0+nZXjTXnvQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/undici-types": { + "version": "6.21.0", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", + "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", + "dev": true, + "license": "MIT" + }, + "node_modules/url-join": { + "version": "4.0.1", + "resolved": "https://registry.npmjs.org/url-join/-/url-join-4.0.1.tgz", + "integrity": "sha512-jk1+QP6ZJqyOiuEI9AEWQfju/nB2Pw466kbA0LEZljHwKeMgd9WrAEgEGxjPDD2+TNbbb37rTyhEfrCXfuKXnA==", + "dev": true, + "license": "MIT" + }, + "node_modules/version-range": { + "version": "4.15.0", + "resolved": "https://registry.npmjs.org/version-range/-/version-range-4.15.0.tgz", + "integrity": "sha512-Ck0EJbAGxHwprkzFO966t4/5QkRuzh+/I1RxhLgUKKwEn+Cd8NwM60mE3AqBZg5gYODoXW0EFsQvbZjRlvdqbg==", + "dev": true, + "license": "Artistic-2.0", + "engines": { + "node": ">=4" + }, + "funding": { + "url": "https://bevry.me/fund" + } + }, + "node_modules/vite": { + "version": "8.3.0", + "resolved": "https://registry.npmjs.org/vite/-/vite-8.3.0.tgz", + "integrity": "sha512-lhZBVvEHefgE+HQZC9O7EBJgCU/nVzFNl7vkS4RE0APtWLP02/8QVIkQtzBxPquh7lq5/78NHipTj7ODQ6XuyQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "lightningcss": "^1.33.0", + "picomatch": "^4.0.7", + "postcss": "^8.5.28", + "rolldown": "~1.2.6", + "tinyglobby": "^0.2.17" + }, + "bin": { + "vite": "bin/vite.js" + }, + "engines": { + "node": "^20.19.0 || >=22.12.0" + }, + "funding": { + "url": "https://github.com/vitejs/vite?sponsor=1" + }, + "optionalDependencies": { + "fsevents": "~2.3.3" + }, + "peerDependencies": { + "@types/node": "^20.19.0 || >=22.12.0", + "@vitejs/devtools": "^0.7.1", + "esbuild": "^0.27.0 || ^0.28.0", + "jiti": ">=1.21.0", + "less": "^4.0.0", + "sass": "^1.70.0", + "sass-embedded": "^1.70.0", + "stylus": ">=0.54.8", + "sugarss": "^5.0.0", + "terser": "^5.16.0", + "tsx": "^4.8.1", + "yaml": "^2.4.2" + }, + "peerDependenciesMeta": { + "@types/node": { + "optional": true + }, + "@vitejs/devtools": { + "optional": true + }, + "esbuild": { + "optional": true + }, + "jiti": { + "optional": true + }, + "less": { + "optional": true + }, + "sass": { + "optional": true + }, + "sass-embedded": { + "optional": true + }, + "stylus": { + "optional": true + }, + "sugarss": { + "optional": true + }, + "terser": { + "optional": true + }, + "tsx": { + "optional": true + }, + "yaml": { + "optional": true + } + } + }, + "node_modules/vitest": { + "version": "4.1.11", + "resolved": "https://registry.npmjs.org/vitest/-/vitest-4.1.11.tgz", + "integrity": "sha512-fhACrNXUidIbGSBr5FlbuBkO7VWC1ZyLl0DO4CU2DrQoAPxX84Ysxs+HeGQpii5lZWV1Q4gBZTTu49mF+A6Edw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@vitest/expect": "4.1.11", + "@vitest/mocker": "4.1.11", + "@vitest/pretty-format": "4.1.11", + "@vitest/runner": "4.1.11", + "@vitest/snapshot": "4.1.11", + "@vitest/spy": "4.1.11", + "@vitest/utils": "4.1.11", + "es-module-lexer": "^2.0.0", + "expect-type": "^1.3.0", + "magic-string": "^0.30.21", + "obug": "^2.1.1", + "pathe": "^2.0.3", + "picomatch": "^4.0.3", + "std-env": "^4.0.0-rc.1", + "tinybench": "^2.9.0", + "tinyexec": "^1.0.2", + "tinyglobby": "^0.2.15", + "tinyrainbow": "^3.1.0", + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0", + "why-is-node-running": "^2.3.0" + }, + "bin": { + "vitest": "vitest.mjs" + }, + "engines": { + "node": "^20.0.0 || ^22.0.0 || >=24.0.0" + }, + "funding": { + "url": "https://opencollective.com/vitest" + }, + "peerDependencies": { + "@edge-runtime/vm": "*", + "@opentelemetry/api": "^1.9.0", + "@types/node": "^20.0.0 || ^22.0.0 || >=24.0.0", + "@vitest/browser-playwright": "4.1.11", + "@vitest/browser-preview": "4.1.11", + "@vitest/browser-webdriverio": "4.1.11", + "@vitest/coverage-istanbul": "4.1.11", + "@vitest/coverage-v8": "4.1.11", + "@vitest/ui": "4.1.11", + "happy-dom": "*", + "jsdom": "*", + "vite": "^6.0.0 || ^7.0.0 || ^8.0.0" + }, + "peerDependenciesMeta": { + "@edge-runtime/vm": { + "optional": true + }, + "@opentelemetry/api": { + "optional": true + }, + "@types/node": { + "optional": true + }, + "@vitest/browser-playwright": { + "optional": true + }, + "@vitest/browser-preview": { + "optional": true + }, + "@vitest/browser-webdriverio": { + "optional": true + }, + "@vitest/coverage-istanbul": { + "optional": true + }, + "@vitest/coverage-v8": { + "optional": true + }, + "@vitest/ui": { + "optional": true + }, + "happy-dom": { + "optional": true + }, + "jsdom": { + "optional": true + }, + "vite": { + "optional": false + } + } + }, + "node_modules/why-is-node-running": { + "version": "2.3.0", + "resolved": "https://registry.npmjs.org/why-is-node-running/-/why-is-node-running-2.3.0.tgz", + "integrity": "sha512-hUrmaWBdVDcxvYqnyh09zunKzROWjbZTiNy8dBEjkS7ehEDQibXJ7XvlmtbwuTclUiIyN+CyXQD4Vmko8fNm8w==", + "dev": true, + "license": "MIT", + "dependencies": { + "siginfo": "^2.0.0", + "stackback": "0.0.2" + }, + "bin": { + "why-is-node-running": "cli.js" + }, + "engines": { + "node": ">=8" + } + }, + "node_modules/wsl-utils": { + "version": "0.1.0", + "resolved": "https://registry.npmjs.org/wsl-utils/-/wsl-utils-0.1.0.tgz", + "integrity": "sha512-h3Fbisa2nKGPxCpm89Hk33lBLsnaGBvctQopaBSOW/uIs6FTe1ATyAnKFJrzVs9vpGdsTe73WF3V4lIsk4Gacw==", + "dev": true, + "license": "MIT", + "dependencies": { + "is-wsl": "^3.1.0" + }, + "engines": { + "node": ">=18" + }, + "funding": { + "url": "https://github.com/sponsors/sindresorhus" + } + }, + "node_modules/xml2js": { + "version": "0.5.0", + "resolved": "https://registry.npmjs.org/xml2js/-/xml2js-0.5.0.tgz", + "integrity": "sha512-drPFnkQJik/O+uPKpqSgr22mpuFHqKdbS835iAQrUC73L2F5WkboIRd63ai/2Yg6I1jzifPFKH2NTK+cfglkIA==", + "dev": true, + "license": "MIT", + "dependencies": { + "sax": ">=0.6.0", + "xmlbuilder": "~11.0.0" + }, + "engines": { + "node": ">=4.0.0" + } + }, + "node_modules/xmlbuilder": { + "version": "11.0.1", + "resolved": "https://registry.npmjs.org/xmlbuilder/-/xmlbuilder-11.0.1.tgz", + "integrity": "sha512-fDlsI/kFEx7gLvbecc0/ohLG50fugQp8ryHzMTuW9vSa1GJ0XYWKnhsUx7oie3G98+r56aTQIUB4kht42R3JvA==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=4.0" + } + }, + "node_modules/yallist": { + "version": "4.0.0", + "resolved": "https://registry.npmjs.org/yallist/-/yallist-4.0.0.tgz", + "integrity": "sha512-3wdGidZyq5PB084XLES5TpOSRA3wjXAlIWMhum2kRcv/41Sn2emQ0dycQW4uZXLejwKvg6EsvbdlVL+FYEct7A==", + "dev": true, + "license": "ISC" + }, + "node_modules/yauzl": { + "version": "3.4.0", + "resolved": "https://registry.npmjs.org/yauzl/-/yauzl-3.4.0.tgz", + "integrity": "sha512-jIH9yLR9wqr0wOS0TpBvo/g/2UgZH5qePVbjgRliiF0BYvOZyaBknKsF+x9Iht0O6sqgnB93rCICdOZFecJuDw==", + "dev": true, + "license": "MIT", + "dependencies": { + "pend": "~1.2.0" + }, + "engines": { + "node": ">=12" + } + }, + "node_modules/yazl": { + "version": "2.5.1", + "resolved": "https://registry.npmjs.org/yazl/-/yazl-2.5.1.tgz", + "integrity": "sha512-phENi2PLiHnHb6QBVot+dJnaAZ0xosj7p3fWl+znIjBDlnMI2PsZCJZ306BPTFOaHf5qdDEI8x5qFrSOBN5vrw==", + "dev": true, + "license": "MIT", + "dependencies": { + "buffer-crc32": "~0.2.3" + } + } + } +} diff --git a/vscode-extension/package.json b/vscode-extension/package.json new file mode 100644 index 00000000000..a8e52b14eda --- /dev/null +++ b/vscode-extension/package.json @@ -0,0 +1,87 @@ +{ + "name": "litellm-vscode", + "displayName": "LiteLLM", + "description": "Chat with every model behind your LiteLLM AI Gateway in VS Code, with live pricing and reasoning effort controls in the model picker", + "version": "0.1.0", + "publisher": "litellm", + "license": "MIT", + "repository": { + "type": "git", + "url": "https://github.com/BerriAI/litellm.git", + "directory": "vscode-extension" + }, + "homepage": "https://docs.litellm.ai", + "bugs": { + "url": "https://github.com/BerriAI/litellm/issues" + }, + "engines": { + "vscode": "^1.109.0" + }, + "categories": [ + "AI", + "Chat" + ], + "keywords": [ + "litellm", + "ai gateway", + "llm", + "chat", + "copilot" + ], + "main": "./dist/extension.js", + "activationEvents": [], + "contributes": { + "languageModelChatProviders": [ + { + "vendor": "litellm", + "displayName": "LiteLLM", + "configuration": { + "type": "object", + "properties": { + "baseUrl": { + "type": "string", + "title": "Gateway URL", + "description": "Base URL of your LiteLLM AI Gateway, for example https://litellm.example.com", + "default": "http://localhost:4000" + }, + "apiKey": { + "type": "string", + "title": "API key", + "description": "A LiteLLM virtual key. The models offered are the ones this key can access", + "secret": true + } + }, + "required": [ + "baseUrl", + "apiKey" + ] + } + } + ], + "commands": [ + { + "command": "litellm.refreshModels", + "title": "Refresh Models", + "category": "LiteLLM" + } + ] + }, + "scripts": { + "build": "esbuild src/extension.ts --bundle --outfile=dist/extension.js --external:vscode --format=cjs --platform=node --target=node22", + "typecheck": "tsc --noEmit", + "test": "vitest run", + "vscode:prepublish": "npm run typecheck && npm run build", + "package": "vsce package --no-dependencies" + }, + "dependencies": { + "openai": "^7.18.0" + }, + "devDependencies": { + "@types/node": "^22.20.3", + "@types/vscode": "^1.109.0", + "@vscode/vsce": "^4.0.0", + "esbuild": "^0.28.2", + "typescript": "^5.9.3", + "vitest": "^4.1.11" + } +} diff --git a/vscode-extension/src/extension.ts b/vscode-extension/src/extension.ts new file mode 100644 index 00000000000..fe0cafb2fbc --- /dev/null +++ b/vscode-extension/src/extension.ts @@ -0,0 +1,17 @@ +import * as vscode from "vscode"; +import { createGatewayClient } from "./gateway"; +import { LiteLLMChatProvider } from "./provider"; + +export const VENDOR = "litellm"; +export const REFRESH_COMMAND = "litellm.refreshModels"; + +export function activate(context: vscode.ExtensionContext): void { + const provider = new LiteLLMChatProvider(createGatewayClient()); + context.subscriptions.push( + provider, + vscode.lm.registerLanguageModelChatProvider(VENDOR, provider), + vscode.commands.registerCommand(REFRESH_COMMAND, () => provider.refresh()), + ); +} + +export function deactivate(): void {} diff --git a/vscode-extension/src/gateway.ts b/vscode-extension/src/gateway.ts new file mode 100644 index 00000000000..0ffe5e0f3f3 --- /dev/null +++ b/vscode-extension/src/gateway.ts @@ -0,0 +1,114 @@ +import OpenAI from "openai"; +import type { ChatCompletionChunk, ChatCompletionCreateParamsStreaming } from "openai/resources/chat/completions"; +import packageJson from "../package.json"; +import { parseModelGroups, type ConfigurationValues, type ModelGroupInfo } from "./models"; + +export interface GatewayConfig { + readonly baseUrl: string; + readonly apiKey: string; +} + +export type GatewayConfigResult = + | { readonly kind: "ok"; readonly config: GatewayConfig } + | { readonly kind: "unconfigured" } + | { readonly kind: "missing_fields"; readonly fields: readonly string[] } + | { readonly kind: "invalid_url"; readonly baseUrl: string }; + +export type ModelGroupsResult = + | { readonly kind: "ok"; readonly groups: readonly ModelGroupInfo[] } + | { readonly kind: "http_error"; readonly status: number; readonly body: string } + | { readonly kind: "invalid_response"; readonly reason: string }; + +export interface GatewayClient { + listModelGroups(config: GatewayConfig, signal: AbortSignal): Promise; + streamChatCompletion( + config: GatewayConfig, + params: ChatCompletionCreateParamsStreaming, + signal: AbortSignal, + ): Promise>; +} + +export const USER_AGENT = `litellm-vscode/${packageJson.version}`; +export const ERROR_SUMMARY_LIMIT = 200; + +const GATEWAY_PROTOCOLS: ReadonlySet = new Set(["http:", "https:"]); + +const parsesAsHttpUrl = (value: string): boolean => { + try { + return GATEWAY_PROTOCOLS.has(new URL(value).protocol); + } catch { + return false; + } +}; + +export const gatewayRoot = (baseUrl: string): string | undefined => { + const root = baseUrl.trim().replace(/\/+$/, "").replace(/\/v1$/, ""); + return parsesAsHttpUrl(root) ? root : undefined; +}; + +const nonEmptyString = (value: unknown): string | undefined => + typeof value === "string" && value.trim() !== "" ? value.trim() : undefined; + +export const gatewayConfigFrom = (configuration: ConfigurationValues | undefined): GatewayConfigResult => { + if (configuration === undefined) { + return { kind: "unconfigured" }; + } + const baseUrl = nonEmptyString(configuration.baseUrl); + const apiKey = nonEmptyString(configuration.apiKey); + if (baseUrl === undefined || apiKey === undefined) { + const fields = [...(baseUrl === undefined ? ["Gateway URL"] : []), ...(apiKey === undefined ? ["API key"] : [])]; + return { kind: "missing_fields", fields }; + } + const root = gatewayRoot(baseUrl); + return root === undefined ? { kind: "invalid_url", baseUrl } : { kind: "ok", config: { baseUrl: root, apiKey } }; +}; + +export const modelGroupInfoUrl = (root: string): string => `${root}/model_group/info`; + +export const openAiBaseUrl = (root: string): string => `${root}/v1`; + +const isRecord = (value: unknown): value is Record => typeof value === "object" && value !== null; + +const errorMessageIn = (body: string): string | undefined => { + try { + const parsed: unknown = JSON.parse(body); + if (!isRecord(parsed)) { + return undefined; + } + if (isRecord(parsed.error) && typeof parsed.error.message === "string") { + return parsed.error.message; + } + return typeof parsed.detail === "string" ? parsed.detail : undefined; + } catch { + return undefined; + } +}; + +export const summarizeErrorBody = (body: string): string => { + const message = (errorMessageIn(body) ?? body).replace(/\s+/g, " ").trim(); + return message.length > ERROR_SUMMARY_LIMIT ? `${message.slice(0, ERROR_SUMMARY_LIMIT)}...` : message; +}; + +export const createGatewayClient = (fetchImpl: typeof fetch = fetch): GatewayClient => ({ + async listModelGroups(config, signal) { + const response = await fetchImpl(modelGroupInfoUrl(config.baseUrl), { + headers: { Authorization: `Bearer ${config.apiKey}`, "User-Agent": USER_AGENT }, + signal, + }); + if (!response.ok) { + return { kind: "http_error", status: response.status, body: await response.text() }; + } + const parsed = parseModelGroups(await response.json()); + return parsed.kind === "ok" ? parsed : { kind: "invalid_response", reason: parsed.reason }; + }, + streamChatCompletion(config, params, signal) { + const client = new OpenAI({ + apiKey: config.apiKey, + baseURL: openAiBaseUrl(config.baseUrl), + defaultHeaders: { "User-Agent": USER_AGENT }, + fetch: fetchImpl, + maxRetries: 0, + }); + return client.chat.completions.create(params, { signal }); + }, +}); diff --git a/vscode-extension/src/messages.ts b/vscode-extension/src/messages.ts new file mode 100644 index 00000000000..855f85435ad --- /dev/null +++ b/vscode-extension/src/messages.ts @@ -0,0 +1,188 @@ +import type * as vscode from "vscode"; +import type { + ChatCompletionAssistantMessageParam, + ChatCompletionContentPart, + ChatCompletionCreateParamsStreaming, + ChatCompletionMessageParam, + ChatCompletionMessageToolCall, + ChatCompletionTool, + ChatCompletionToolMessageParam, +} from "openai/resources/chat/completions"; +import { estimateTokens } from "./models"; + +export interface ChatRequestInput { + readonly model: string; + readonly messages: readonly vscode.LanguageModelChatRequestMessage[]; + readonly tools: readonly vscode.LanguageModelChatTool[]; + readonly requireToolCall: boolean; + readonly reasoningEffort: string | undefined; + readonly modelOptions: { readonly [key: string]: unknown }; +} + +interface TextPart { + readonly value: string; +} + +interface ToolCallPart { + readonly callId: string; + readonly name: string; + readonly input: object; +} + +interface ToolResultPart { + readonly callId: string; + readonly content: ReadonlyArray; +} + +interface DataPart { + readonly mimeType: string; + readonly data: Uint8Array; +} + +const USER_ROLE = 1; +const ASSISTANT_ROLE = 2; +const SYSTEM_ROLE = 3; + +export const ESTIMATED_TOKENS_PER_IMAGE = 1000; + +const isRecord = (value: unknown): value is Record => typeof value === "object" && value !== null; + +const isTextPart = (part: unknown): part is TextPart => isRecord(part) && typeof part.value === "string"; + +const isToolCallPart = (part: unknown): part is ToolCallPart => + isRecord(part) && typeof part.callId === "string" && typeof part.name === "string" && isRecord(part.input); + +const isToolResultPart = (part: unknown): part is ToolResultPart => + isRecord(part) && typeof part.callId === "string" && Array.isArray(part.content); + +const isDataPart = (part: unknown): part is DataPart => + isRecord(part) && typeof part.mimeType === "string" && part.data instanceof Uint8Array; + +const isImagePart = (part: unknown): part is DataPart => isDataPart(part) && part.mimeType.startsWith("image/"); + +const dataUrl = (part: DataPart): string => `data:${part.mimeType};base64,${Buffer.from(part.data).toString("base64")}`; + +const textOf = (part: unknown): string => { + if (isTextPart(part)) { + return part.value; + } + if (isDataPart(part) && part.mimeType.startsWith("text/")) { + return Buffer.from(part.data).toString("utf8"); + } + if (isRecord(part) && "value" in part) { + return JSON.stringify(part.value); + } + return ""; +}; + +const contentParts = (parts: readonly unknown[]): readonly ChatCompletionContentPart[] => + parts.flatMap((part): readonly ChatCompletionContentPart[] => { + if (isImagePart(part)) { + return [{ type: "image_url", image_url: { url: dataUrl(part) } }]; + } + const text = textOf(part); + return text === "" ? [] : [{ type: "text", text }]; + }); + +const toolMessage = (part: ToolResultPart): ChatCompletionToolMessageParam => ({ + role: "tool", + tool_call_id: part.callId, + content: part.content.filter((item) => !isImagePart(item)).map(textOf).join(""), +}); + +const userMessages = (parts: readonly unknown[]): readonly ChatCompletionMessageParam[] => { + const toolResults = parts.filter(isToolResultPart); + const toolResultImages = toolResults.flatMap((result) => result.content.filter(isImagePart)); + const remaining = parts.filter((part) => !isToolResultPart(part)); + const userContent = contentParts([...remaining, ...toolResultImages]); + const userMessage: readonly ChatCompletionMessageParam[] = + userContent.length === 0 ? [] : [{ role: "user", content: [...userContent] }]; + return [...toolResults.map(toolMessage), ...userMessage]; +}; + +const toolCall = (part: ToolCallPart): ChatCompletionMessageToolCall => ({ + id: part.callId, + type: "function", + function: { name: part.name, arguments: JSON.stringify(part.input) }, +}); + +const assistantMessages = (parts: readonly unknown[]): readonly ChatCompletionAssistantMessageParam[] => { + const text = parts.filter(isTextPart).map((part) => part.value).join(""); + const toolCalls = parts.filter(isToolCallPart).map(toolCall); + if (text === "" && toolCalls.length === 0) { + return []; + } + return [ + { + role: "assistant", + content: text === "" ? null : text, + ...(toolCalls.length === 0 ? {} : { tool_calls: toolCalls }), + }, + ]; +}; + +const convertMessage = (message: vscode.LanguageModelChatRequestMessage): readonly ChatCompletionMessageParam[] => { + const role: number = message.role; + switch (role) { + case USER_ROLE: + return userMessages(message.content); + case ASSISTANT_ROLE: + return assistantMessages(message.content); + case SYSTEM_ROLE: + return [{ role: "system", content: message.content.map(textOf).join("") }]; + default: + return []; + } +}; + +export const toChatCompletionMessages = ( + messages: readonly vscode.LanguageModelChatRequestMessage[], +): readonly ChatCompletionMessageParam[] => messages.flatMap(convertMessage); + +const imagePartsIn = (parts: readonly unknown[]): readonly DataPart[] => [ + ...parts.filter(isImagePart), + ...parts.filter(isToolResultPart).flatMap((result) => result.content.filter(isImagePart)), +]; + +const withoutImageData = (key: string, value: unknown): unknown => (key === "image_url" ? undefined : value); + +export const estimateMessageTokens = (message: vscode.LanguageModelChatRequestMessage): number => { + const converted = convertMessage(message); + if (converted.length === 0) { + return 0; + } + const images = imagePartsIn(message.content).length; + return estimateTokens(JSON.stringify(converted, withoutImageData)) + images * ESTIMATED_TOKENS_PER_IMAGE; +}; + +const toTool = (tool: vscode.LanguageModelChatTool): ChatCompletionTool => ({ + type: "function", + function: { + name: tool.name, + description: tool.description, + ...(tool.inputSchema === undefined ? {} : { parameters: tool.inputSchema as Record }), + }, +}); + +const NUMERIC_OPTIONS = ["temperature", "top_p", "max_tokens", "presence_penalty", "frequency_penalty", "seed"] as const; + +const forwardedModelOptions = (modelOptions: { readonly [key: string]: unknown }): Record => + Object.fromEntries( + NUMERIC_OPTIONS.flatMap((key) => { + const value = modelOptions[key]; + return typeof value === "number" ? [[key, value] as const] : []; + }), + ); + +export const buildChatCompletionParams = (input: ChatRequestInput): ChatCompletionCreateParamsStreaming => ({ + model: input.model, + messages: [...toChatCompletionMessages(input.messages)], + stream: true, + stream_options: { include_usage: true }, + ...forwardedModelOptions(input.modelOptions), + ...(input.tools.length === 0 ? {} : { tools: input.tools.map(toTool) }), + ...(input.tools.length === 0 || !input.requireToolCall ? {} : { tool_choice: "required" }), + ...(input.reasoningEffort === undefined + ? {} + : { reasoning_effort: input.reasoningEffort as ChatCompletionCreateParamsStreaming["reasoning_effort"] }), +}); diff --git a/vscode-extension/src/models.ts b/vscode-extension/src/models.ts new file mode 100644 index 00000000000..d76f6557b8b --- /dev/null +++ b/vscode-extension/src/models.ts @@ -0,0 +1,170 @@ +export interface ModelGroupInfo { + readonly modelGroup: string; + readonly providers: readonly string[]; + readonly mode: string | undefined; + readonly maxInputTokens: number | undefined; + readonly maxOutputTokens: number | undefined; + readonly inputCostPerToken: number | undefined; + readonly outputCostPerToken: number | undefined; + readonly supportsVision: boolean; + readonly supportsFunctionCalling: boolean; + readonly supportedReasoningEfforts: readonly string[]; +} + +export type ConfigurationValues = { readonly [key: string]: unknown }; + +export interface ConfigurationSchemaProperty { + readonly type: "string"; + readonly title: string; + readonly enum: readonly string[]; + readonly enumItemLabels: readonly string[]; + readonly default: string; + readonly group: "navigation"; +} + +export interface ConfigurationSchema { + readonly properties: { readonly [key: string]: ConfigurationSchemaProperty }; +} + +export interface ModelDescriptor { + readonly id: string; + readonly name: string; + readonly family: string; + readonly version: string; + readonly detail: string; + readonly tooltip: string; + readonly maxInputTokens: number; + readonly maxOutputTokens: number; + readonly imageInput: boolean; + readonly toolCalling: boolean; + readonly configurationSchema: ConfigurationSchema | undefined; +} + +export type ModelGroupsParseResult = + | { readonly kind: "ok"; readonly groups: readonly ModelGroupInfo[] } + | { readonly kind: "invalid"; readonly reason: string }; + +export const REASONING_EFFORT_KEY = "reasoningEffort"; +export const GATEWAY_DEFAULT_EFFORT = "default"; +export const ASSUMED_MAX_INPUT_TOKENS = 128000; +export const ASSUMED_MAX_OUTPUT_TOKENS = 4096; +export const MARKDOWN_LINE_BREAK = " \n"; + +const isRecord = (value: unknown): value is Record => typeof value === "object" && value !== null; + +const optionalNumber = (value: unknown): number | undefined => + typeof value === "number" && Number.isFinite(value) ? value : undefined; + +const optionalString = (value: unknown): string | undefined => (typeof value === "string" ? value : undefined); + +const stringList = (value: unknown): readonly string[] => + Array.isArray(value) ? value.filter((item): item is string => typeof item === "string") : []; + +const parseGroup = (value: unknown): ModelGroupInfo | undefined => { + if (!isRecord(value) || typeof value.model_group !== "string") { + return undefined; + } + return { + modelGroup: value.model_group, + providers: stringList(value.providers), + mode: optionalString(value.mode), + maxInputTokens: optionalNumber(value.max_input_tokens), + maxOutputTokens: optionalNumber(value.max_output_tokens), + inputCostPerToken: optionalNumber(value.input_cost_per_token), + outputCostPerToken: optionalNumber(value.output_cost_per_token), + supportsVision: value.supports_vision === true, + supportsFunctionCalling: value.supports_function_calling === true, + supportedReasoningEfforts: stringList(value.supported_reasoning_efforts), + }; +}; + +export const parseModelGroups = (body: unknown): ModelGroupsParseResult => { + if (!isRecord(body) || !Array.isArray(body.data)) { + return { kind: "invalid", reason: "response has no data array" }; + } + const groups = body.data.map(parseGroup).filter((group): group is ModelGroupInfo => group !== undefined); + return { kind: "ok", groups }; +}; + +const isChatGroup = (group: ModelGroupInfo): boolean => group.mode === undefined || group.mode === "chat"; + +export const formatUsdPerMillionTokens = (costPerToken: number): string => { + const perMillion = costPerToken * 1_000_000; + const digits = perMillion === 0 || perMillion >= 0.01 ? perMillion.toFixed(2) : perMillion.toPrecision(2); + return `$${digits}`; +}; + +const priceLine = (label: string, costPerToken: number | undefined): string => + costPerToken === undefined ? `${label}: no price configured` : `${label}: ${formatUsdPerMillionTokens(costPerToken)} per 1M tokens`; + +const pricingDetail = (group: ModelGroupInfo): string => { + if (group.inputCostPerToken === undefined && group.outputCostPerToken === undefined) { + return "No pricing configured"; + } + const input = group.inputCostPerToken === undefined ? "n/a" : formatUsdPerMillionTokens(group.inputCostPerToken); + const output = group.outputCostPerToken === undefined ? "n/a" : formatUsdPerMillionTokens(group.outputCostPerToken); + return `${input} in / ${output} out per 1M tokens`; +}; + +const capitalize = (value: string): string => value.charAt(0).toUpperCase() + value.slice(1); + +const effortSchema = (efforts: readonly string[]): ConfigurationSchema | undefined => { + if (efforts.length === 0) { + return undefined; + } + return { + properties: { + [REASONING_EFFORT_KEY]: { + type: "string", + title: "Reasoning Effort", + enum: [GATEWAY_DEFAULT_EFFORT, ...efforts], + enumItemLabels: ["Gateway default", ...efforts.map(capitalize)], + default: GATEWAY_DEFAULT_EFFORT, + group: "navigation", + }, + }, + }; +}; + +const tooltipFor = (group: ModelGroupInfo): string => { + const providers = group.providers.length === 0 ? "" : ` via ${group.providers.join(", ")}`; + const context = + group.maxInputTokens === undefined || group.maxOutputTokens === undefined + ? `Context: unknown, assuming ${ASSUMED_MAX_INPUT_TOKENS} in / ${ASSUMED_MAX_OUTPUT_TOKENS} out tokens` + : `Context: ${group.maxInputTokens} in / ${group.maxOutputTokens} out tokens`; + const efforts = + group.supportedReasoningEfforts.length === 0 + ? "Reasoning effort: not configurable" + : `Reasoning effort: ${group.supportedReasoningEfforts.join(", ")}`; + return [ + `LiteLLM model group ${group.modelGroup}${providers}`, + priceLine("Input", group.inputCostPerToken), + priceLine("Output", group.outputCostPerToken), + context, + efforts, + ].join(MARKDOWN_LINE_BREAK); +}; + +const describeGroup = (group: ModelGroupInfo): ModelDescriptor => ({ + id: group.modelGroup, + name: group.modelGroup, + family: group.modelGroup, + version: "1.0", + detail: pricingDetail(group), + tooltip: tooltipFor(group), + maxInputTokens: group.maxInputTokens ?? ASSUMED_MAX_INPUT_TOKENS, + maxOutputTokens: group.maxOutputTokens ?? ASSUMED_MAX_OUTPUT_TOKENS, + imageInput: group.supportsVision, + toolCalling: group.supportsFunctionCalling, + configurationSchema: effortSchema(group.supportedReasoningEfforts), +}); + +export const describeModels = (groups: readonly ModelGroupInfo[]): readonly ModelDescriptor[] => + groups.filter(isChatGroup).map(describeGroup); + +export const reasoningEffortFrom = (configuration: ConfigurationValues | undefined): string | undefined => { + const effort = configuration?.[REASONING_EFFORT_KEY]; + return typeof effort === "string" && effort !== GATEWAY_DEFAULT_EFFORT ? effort : undefined; +}; + +export const estimateTokens = (text: string): number => Math.ceil(text.length / 4); diff --git a/vscode-extension/src/provider.ts b/vscode-extension/src/provider.ts new file mode 100644 index 00000000000..3cafd449147 --- /dev/null +++ b/vscode-extension/src/provider.ts @@ -0,0 +1,136 @@ +import * as vscode from "vscode"; +import { gatewayConfigFrom, summarizeErrorBody, type GatewayClient, type GatewayConfig, type GatewayConfigResult, type ModelGroupsResult } from "./gateway"; +import { buildChatCompletionParams, estimateMessageTokens } from "./messages"; +import { describeModels, estimateTokens, reasoningEffortFrom, type ModelDescriptor } from "./models"; +import { responseParts, type ResponsePart } from "./stream"; + +export interface LiteLLMModel extends vscode.LanguageModelChatInformation { + readonly gateway: GatewayConfig; +} + +export const TRUNCATED_MESSAGE = "The model stopped at its output token limit before finishing the response"; + +const RECONFIGURE_HINT = + 'Fix it from the gear on its row in Manage Language Models: "Update API Key" for the key, "Open in Language Models (JSON)" for the URL'; + +const configurationProblem = (result: Exclude): string => { + switch (result.kind) { + case "missing_fields": + return `LiteLLM provider is missing its ${result.fields.join(" and ")}. ${RECONFIGURE_HINT}`; + case "invalid_url": + return `LiteLLM gateway URL "${result.baseUrl}" is not an http or https URL. ${RECONFIGURE_HINT}`; + } +}; + +const discoveryFailure = (result: Exclude, baseUrl: string): string => { + switch (result.kind) { + case "http_error": + return `LiteLLM gateway at ${baseUrl} answered ${result.status} for /model_group/info: ${summarizeErrorBody(result.body)}`; + case "invalid_response": + return `LiteLLM gateway at ${baseUrl} returned an unexpected /model_group/info payload: ${result.reason}`; + } +}; + +const toModel = (descriptor: ModelDescriptor, gateway: GatewayConfig): LiteLLMModel => ({ + id: descriptor.id, + name: descriptor.name, + family: descriptor.family, + version: descriptor.version, + detail: descriptor.detail, + tooltip: descriptor.tooltip, + maxInputTokens: descriptor.maxInputTokens, + maxOutputTokens: descriptor.maxOutputTokens, + capabilities: { imageInput: descriptor.imageInput, toolCalling: descriptor.toolCalling }, + ...(descriptor.configurationSchema === undefined ? {} : { configurationSchema: descriptor.configurationSchema }), + gateway, +}); + +const toVscodePart = (part: ResponsePart): vscode.LanguageModelResponsePart => { + switch (part.kind) { + case "text": + return new vscode.LanguageModelTextPart(part.value); + case "tool_call": + return new vscode.LanguageModelToolCallPart(part.callId, part.name, part.input); + case "invalid_tool_call": + throw new Error(`Model returned invalid JSON arguments for tool ${part.name}: ${part.arguments}`); + case "truncated": + throw new Error(TRUNCATED_MESSAGE); + } +}; + +const withAbortSignal = async (token: vscode.CancellationToken, run: (signal: AbortSignal) => Promise): Promise => { + const controller = new AbortController(); + const subscription = token.onCancellationRequested(() => controller.abort()); + try { + return await run(controller.signal); + } finally { + subscription.dispose(); + } +}; + +export class LiteLLMChatProvider implements vscode.LanguageModelChatProvider, vscode.Disposable { + private readonly changeEmitter = new vscode.EventEmitter(); + readonly onDidChangeLanguageModelChatInformation = this.changeEmitter.event; + + constructor(private readonly gateway: GatewayClient) {} + + refresh(): void { + this.changeEmitter.fire(); + } + + dispose(): void { + this.changeEmitter.dispose(); + } + + async provideLanguageModelChatInformation( + options: vscode.PrepareLanguageModelChatModelOptions, + token: vscode.CancellationToken, + ): Promise { + const configured = gatewayConfigFrom(options.configuration); + if (configured.kind === "unconfigured") { + return []; + } + if (configured.kind !== "ok") { + throw new Error(configurationProblem(configured)); + } + const result = await withAbortSignal(token, (signal) => this.gateway.listModelGroups(configured.config, signal)); + if (result.kind !== "ok") { + throw new Error(discoveryFailure(result, configured.config.baseUrl)); + } + return describeModels(result.groups).map((descriptor) => toModel(descriptor, configured.config)); + } + + async provideLanguageModelChatResponse( + model: LiteLLMModel, + messages: readonly vscode.LanguageModelChatRequestMessage[], + options: vscode.ProvideLanguageModelChatResponseOptions, + progress: vscode.Progress, + token: vscode.CancellationToken, + ): Promise { + const params = buildChatCompletionParams({ + model: model.id, + messages, + tools: options.tools ?? [], + requireToolCall: options.toolMode === vscode.LanguageModelChatToolMode.Required, + reasoningEffort: reasoningEffortFrom(options.modelConfiguration), + modelOptions: options.modelOptions ?? {}, + }); + try { + await withAbortSignal(token, async (signal) => { + const chunks = await this.gateway.streamChatCompletion(model.gateway, params, signal); + for await (const part of responseParts(chunks)) { + progress.report(toVscodePart(part)); + } + }); + } catch (error) { + if (token.isCancellationRequested) { + return; + } + throw error; + } + } + + async provideTokenCount(_model: LiteLLMModel, text: string | vscode.LanguageModelChatRequestMessage): Promise { + return typeof text === "string" ? estimateTokens(text) : estimateMessageTokens(text); + } +} diff --git a/vscode-extension/src/stream.ts b/vscode-extension/src/stream.ts new file mode 100644 index 00000000000..f1677262bbf --- /dev/null +++ b/vscode-extension/src/stream.ts @@ -0,0 +1,90 @@ +import type { ChatCompletionChunk } from "openai/resources/chat/completions"; + +export type ResponsePart = + | { readonly kind: "text"; readonly value: string } + | { readonly kind: "tool_call"; readonly callId: string; readonly name: string; readonly input: object } + | { readonly kind: "invalid_tool_call"; readonly callId: string; readonly name: string; readonly arguments: string } + | { readonly kind: "truncated" }; + +interface PendingToolCall { + readonly index: number; + readonly callId: string; + readonly name: string; + readonly arguments: string; +} + +export type PendingToolCalls = readonly PendingToolCall[]; + +export interface ChunkOutcome { + readonly pending: PendingToolCalls; + readonly parts: readonly ResponsePart[]; +} + +export const NO_PENDING_TOOL_CALLS: PendingToolCalls = []; + +type ToolCallDelta = NonNullable[number]; + +const nonEmpty = (value: string | undefined): string | undefined => (value === undefined || value === "" ? undefined : value); + +const targetOf = (pending: PendingToolCalls, delta: ToolCallDelta): PendingToolCall | undefined => { + const id = nonEmpty(delta.id); + if (id !== undefined) { + return pending.find((call) => call.callId === id); + } + const sameIndex = pending.filter((call) => call.index === delta.index); + return sameIndex.at(-1) ?? (delta.index === undefined ? pending.at(-1) : undefined); +}; + +const mergeToolCallDelta = (pending: PendingToolCalls, delta: ToolCallDelta): PendingToolCalls => { + const target = targetOf(pending, delta); + const base: PendingToolCall = target ?? { index: delta.index ?? pending.length, callId: nonEmpty(delta.id) ?? "", name: "", arguments: "" }; + const merged: PendingToolCall = { + ...base, + name: nonEmpty(delta.function?.name) ?? base.name, + arguments: base.arguments + (delta.function?.arguments ?? ""), + }; + return target === undefined ? [...pending, merged] : pending.map((call) => (call === target ? merged : call)); +}; + +export const applyChunk = (pending: PendingToolCalls, chunk: ChatCompletionChunk): ChunkOutcome => { + const choice = chunk.choices[0]; + if (choice === undefined) { + return { pending, parts: [] }; + } + const text = typeof choice.delta.content === "string" && choice.delta.content !== "" ? [{ kind: "text", value: choice.delta.content } as const] : []; + const truncated = choice.finish_reason === "length" ? [{ kind: "truncated" } as const] : []; + const nextPending = (choice.delta.tool_calls ?? []).reduce(mergeToolCallDelta, pending); + return { pending: nextPending, parts: [...text, ...truncated] }; +}; + +const parseArguments = (raw: string): object | undefined => { + if (raw.trim() === "") { + return {}; + } + try { + const parsed: unknown = JSON.parse(raw); + return typeof parsed === "object" && parsed !== null ? parsed : undefined; + } catch { + return undefined; + } +}; + +const finishToolCall = (call: PendingToolCall): ResponsePart => { + const input = parseArguments(call.arguments); + return input === undefined + ? { kind: "invalid_tool_call", callId: call.callId, name: call.name, arguments: call.arguments } + : { kind: "tool_call", callId: call.callId, name: call.name, input }; +}; + +export const flushToolCalls = (pending: PendingToolCalls): readonly ResponsePart[] => + [...pending].sort((left, right) => left.index - right.index).map(finishToolCall); + +export async function* responseParts(chunks: AsyncIterable): AsyncGenerator { + let pending: PendingToolCalls = NO_PENDING_TOOL_CALLS; + for await (const chunk of chunks) { + const outcome = applyChunk(pending, chunk); + pending = outcome.pending; + yield* outcome.parts; + } + yield* flushToolCalls(pending); +} diff --git a/vscode-extension/src/vscode.proposed.d.ts b/vscode-extension/src/vscode.proposed.d.ts new file mode 100644 index 00000000000..6665aaa38ef --- /dev/null +++ b/vscode-extension/src/vscode.proposed.d.ts @@ -0,0 +1,15 @@ +import type { ConfigurationSchema, ConfigurationValues } from "./models"; + +declare module "vscode" { + interface LanguageModelChatInformation { + readonly configurationSchema?: ConfigurationSchema; + } + + interface PrepareLanguageModelChatModelOptions { + readonly configuration?: ConfigurationValues; + } + + interface ProvideLanguageModelChatResponseOptions { + readonly modelConfiguration?: ConfigurationValues; + } +} diff --git a/vscode-extension/test/gateway.test.ts b/vscode-extension/test/gateway.test.ts new file mode 100644 index 00000000000..3a70a68b377 --- /dev/null +++ b/vscode-extension/test/gateway.test.ts @@ -0,0 +1,205 @@ +import { createServer, type IncomingMessage, type Server, type ServerResponse } from "node:http"; +import type { AddressInfo } from "node:net"; +import { afterEach, describe, expect, it } from "vitest"; +import { + ERROR_SUMMARY_LIMIT, + createGatewayClient, + gatewayConfigFrom, + gatewayRoot, + modelGroupInfoUrl, + openAiBaseUrl, + summarizeErrorBody, + USER_AGENT, + type GatewayConfig, +} from "../src/gateway"; +import { buildChatCompletionParams } from "../src/messages"; + +interface RecordedRequest { + readonly method: string | undefined; + readonly url: string | undefined; + readonly authorization: string | undefined; + readonly userAgent: string | undefined; + readonly body: string; +} + +type Handler = (request: RecordedRequest, response: ServerResponse) => void; + +const readBody = (request: IncomingMessage): Promise => + new Promise((resolve) => { + const chunks: Buffer[] = []; + request.on("data", (chunk: Buffer) => chunks.push(chunk)); + request.on("end", () => resolve(Buffer.concat(chunks).toString("utf8"))); + }); + +const servers: Server[] = []; + +const startGateway = (handler: Handler): Promise<{ readonly url: string; readonly requests: readonly RecordedRequest[] }> => + new Promise((resolve) => { + const requests: RecordedRequest[] = []; + const server = createServer(async (request, response) => { + const recorded: RecordedRequest = { + method: request.method, + url: request.url, + authorization: request.headers.authorization, + userAgent: request.headers["user-agent"], + body: await readBody(request), + }; + requests.push(recorded); + handler(recorded, response); + }); + servers.push(server); + server.listen(0, "127.0.0.1", () => { + const { port } = server.address() as AddressInfo; + resolve({ url: `http://127.0.0.1:${port}`, requests }); + }); + }); + +afterEach(() => { + servers.splice(0).forEach((server) => server.close()); +}); + +const configFor = (baseUrl: string, apiKey: string): GatewayConfig => { + const result = gatewayConfigFrom({ baseUrl, apiKey }); + if (result.kind !== "ok") { + throw new Error(result.kind); + } + return result.config; +}; + +const sse = (response: ServerResponse, events: readonly object[]): void => { + response.writeHead(200, { "content-type": "text/event-stream" }); + events.forEach((event) => response.write(`data: ${JSON.stringify(event)}\n\n`)); + response.end("data: [DONE]\n\n"); +}; + +describe("gateway URLs", () => { + it("accepts the gateway root with or without a trailing slash or /v1", () => { + expect(gatewayRoot("https://litellm.example.com/")).toBe("https://litellm.example.com"); + expect(gatewayRoot("https://litellm.example.com/v1")).toBe("https://litellm.example.com"); + expect(gatewayRoot(" http://localhost:4000 ")).toBe("http://localhost:4000"); + expect(modelGroupInfoUrl("https://litellm.example.com")).toBe("https://litellm.example.com/model_group/info"); + expect(openAiBaseUrl("https://litellm.example.com")).toBe("https://litellm.example.com/v1"); + }); + + it("rejects anything that is not an http or https URL", () => { + expect(gatewayRoot("litellm.example.com")).toBeUndefined(); + expect(gatewayRoot("ftp://litellm.example.com")).toBeUndefined(); + expect(gatewayRoot("")).toBeUndefined(); + }); +}); + +describe("gatewayConfigFrom", () => { + it("distinguishes the unconfigured probe, a lost secret, and a bad URL from a usable configuration", () => { + expect(gatewayConfigFrom(undefined)).toEqual({ kind: "unconfigured" }); + expect(gatewayConfigFrom({ baseUrl: "http://localhost:4000", apiKey: undefined })).toEqual({ kind: "missing_fields", fields: ["API key"] }); + expect(gatewayConfigFrom({ baseUrl: " ", apiKey: "" })).toEqual({ kind: "missing_fields", fields: ["Gateway URL", "API key"] }); + expect(gatewayConfigFrom({ baseUrl: "localhost:4000", apiKey: "sk" })).toEqual({ kind: "invalid_url", baseUrl: "localhost:4000" }); + expect(gatewayConfigFrom({ baseUrl: " http://localhost:4000/v1/ ", apiKey: " sk-test " })).toEqual({ + kind: "ok", + config: { baseUrl: "http://localhost:4000", apiKey: "sk-test" }, + }); + }); +}); + +describe("summarizeErrorBody", () => { + it("prefers the gateway's error message and caps the length", () => { + expect(summarizeErrorBody('{"error":{"message":"invalid key","type":"auth_error","param":"sk-...abcd"}}')).toBe("invalid key"); + expect(summarizeErrorBody('{"detail":"Not Found"}')).toBe("Not Found"); + expect(summarizeErrorBody("\n 502 Bad Gateway\n")).toBe(" 502 Bad Gateway "); + const long = summarizeErrorBody("x".repeat(ERROR_SUMMARY_LIMIT + 50)); + expect(long).toBe(`${"x".repeat(ERROR_SUMMARY_LIMIT)}...`); + }); +}); + +describe("listModelGroups", () => { + it("calls /model_group/info with the virtual key and this extension's user agent", async () => { + const gateway = await startGateway((_request, response) => { + response.writeHead(200, { "content-type": "application/json" }); + response.end(JSON.stringify({ data: [{ model_group: "gpt-5.6", mode: "chat", input_cost_per_token: 4e-6 }] })); + }); + const result = await createGatewayClient().listModelGroups(configFor(`${gateway.url}/v1`, "sk-test"), new AbortController().signal); + expect(result).toEqual({ + kind: "ok", + groups: [expect.objectContaining({ modelGroup: "gpt-5.6", inputCostPerToken: 4e-6 })], + }); + expect(gateway.requests).toEqual([ + expect.objectContaining({ method: "GET", url: "/model_group/info", authorization: "Bearer sk-test", userAgent: USER_AGENT }), + ]); + }); + + it("reports the gateway's status and body when the key is rejected", async () => { + const gateway = await startGateway((_request, response) => { + response.writeHead(401, { "content-type": "application/json" }); + response.end('{"error":{"message":"invalid key"}}'); + }); + expect(await createGatewayClient().listModelGroups({ baseUrl: gateway.url, apiKey: "sk-bad" }, new AbortController().signal)).toEqual({ + kind: "http_error", + status: 401, + body: '{"error":{"message":"invalid key"}}', + }); + }); + + it("reports a payload that is not a model group listing", async () => { + const gateway = await startGateway((_request, response) => { + response.writeHead(200, { "content-type": "application/json" }); + response.end('{"object":"list","models":[]}'); + }); + expect(await createGatewayClient().listModelGroups({ baseUrl: gateway.url, apiKey: "sk" }, new AbortController().signal)).toEqual({ + kind: "invalid_response", + reason: "response has no data array", + }); + }); +}); + +describe("streamChatCompletion", () => { + it("streams /v1/chat/completions through the gateway with the chosen reasoning effort", async () => { + const gateway = await startGateway((_request, response) => + sse(response, [ + { id: "c", object: "chat.completion.chunk", created: 0, model: "gpt-5.6", choices: [{ index: 0, delta: { content: "Hi" }, finish_reason: null }] }, + { id: "c", object: "chat.completion.chunk", created: 0, model: "gpt-5.6", choices: [{ index: 0, delta: {}, finish_reason: "stop" }] }, + ]), + ); + const params = buildChatCompletionParams({ + model: "gpt-5.6", + messages: [{ role: 1, content: [{ value: "hello" }], name: undefined }], + tools: [], + requireToolCall: false, + reasoningEffort: "high", + modelOptions: {}, + }); + const chunks = await createGatewayClient().streamChatCompletion(configFor(`${gateway.url}/`, "sk-test"), params, new AbortController().signal); + const contents: string[] = []; + for await (const chunk of chunks) { + contents.push(chunk.choices[0]?.delta.content ?? ""); + } + expect(contents.join("")).toBe("Hi"); + const [request] = gateway.requests; + expect(request).toMatchObject({ method: "POST", url: "/v1/chat/completions", authorization: "Bearer sk-test", userAgent: USER_AGENT }); + expect(JSON.parse(request?.body ?? "{}")).toMatchObject({ + model: "gpt-5.6", + stream: true, + stream_options: { include_usage: true }, + reasoning_effort: "high", + messages: [{ role: "user", content: [{ type: "text", text: "hello" }] }], + }); + }); + + it("leaves retries to the gateway instead of resending a failed request", async () => { + const gateway = await startGateway((_request, response) => { + response.writeHead(502, { "content-type": "application/json" }); + response.end('{"error":{"message":"upstream unavailable"}}'); + }); + const params = buildChatCompletionParams({ + model: "gpt-5.6", + messages: [{ role: 1, content: [{ value: "hello" }], name: undefined }], + tools: [], + requireToolCall: false, + reasoningEffort: undefined, + modelOptions: {}, + }); + await expect( + createGatewayClient().streamChatCompletion({ baseUrl: gateway.url, apiKey: "sk-test" }, params, new AbortController().signal), + ).rejects.toThrow(/upstream unavailable/); + expect(gateway.requests).toHaveLength(1); + }); +}); diff --git a/vscode-extension/test/messages.test.ts b/vscode-extension/test/messages.test.ts new file mode 100644 index 00000000000..ba11e479ad2 --- /dev/null +++ b/vscode-extension/test/messages.test.ts @@ -0,0 +1,168 @@ +import { describe, expect, it } from "vitest"; +import type * as vscode from "vscode"; +import { ESTIMATED_TOKENS_PER_IMAGE, buildChatCompletionParams, estimateMessageTokens, toChatCompletionMessages, type ChatRequestInput } from "../src/messages"; + +const USER = 1 as vscode.LanguageModelChatMessageRole; +const ASSISTANT = 2 as vscode.LanguageModelChatMessageRole; +const SYSTEM = 3 as vscode.LanguageModelChatMessageRole; + +const message = (role: vscode.LanguageModelChatMessageRole, content: readonly unknown[]): vscode.LanguageModelChatRequestMessage => ({ + role, + content, + name: undefined, +}); + +const text = (value: string): unknown => ({ value }); +const image = (bytes: readonly number[], mimeType = "image/png"): unknown => ({ mimeType, data: Uint8Array.from(bytes) }); +const toolCall = (callId: string, name: string, input: object): unknown => ({ callId, name, input }); +const toolResult = (callId: string, content: readonly unknown[]): unknown => ({ callId, content }); + +const request = (overrides: Partial = {}): ChatRequestInput => ({ + model: "gpt-5.6", + messages: [message(USER, [text("hi")])], + tools: [], + requireToolCall: false, + reasoningEffort: undefined, + modelOptions: {}, + ...overrides, +}); + +describe("toChatCompletionMessages", () => { + it("maps system, user, and assistant text", () => { + expect( + toChatCompletionMessages([ + message(SYSTEM, [text("be terse")]), + message(USER, [text("hello "), text("there")]), + message(ASSISTANT, [text("hi")]), + ]), + ).toEqual([ + { role: "system", content: "be terse" }, + { role: "user", content: [{ type: "text", text: "hello " }, { type: "text", text: "there" }] }, + { role: "assistant", content: "hi" }, + ]); + }); + + it("sends user images as data URLs", () => { + expect(toChatCompletionMessages([message(USER, [text("what is this"), image([1, 2, 3])])])).toEqual([ + { + role: "user", + content: [ + { type: "text", text: "what is this" }, + { type: "image_url", image_url: { url: "data:image/png;base64,AQID" } }, + ], + }, + ]); + }); + + it("round-trips tool calls and puts tool results before the user's follow-up text", () => { + expect( + toChatCompletionMessages([ + message(ASSISTANT, [text("checking"), toolCall("call_1", "read_file", { path: "a.ts" })]), + message(USER, [toolResult("call_1", [text("export const a = 1;")]), text("thanks")]), + ]), + ).toEqual([ + { + role: "assistant", + content: "checking", + tool_calls: [{ id: "call_1", type: "function", function: { name: "read_file", arguments: '{"path":"a.ts"}' } }], + }, + { role: "tool", tool_call_id: "call_1", content: "export const a = 1;" }, + { role: "user", content: [{ type: "text", text: "thanks" }] }, + ]); + }); + + it("emits a content-less assistant turn that only called tools", () => { + expect(toChatCompletionMessages([message(ASSISTANT, [toolCall("c", "t", {})])])).toEqual([ + { role: "assistant", content: null, tool_calls: [{ id: "c", type: "function", function: { name: "t", arguments: "{}" } }] }, + ]); + }); + + it("drops an assistant turn with neither text nor tool calls", () => { + expect(toChatCompletionMessages([message(USER, [text("hi")]), message(ASSISTANT, [text("")]), message(USER, [text("again")])])).toEqual([ + { role: "user", content: [{ type: "text", text: "hi" }] }, + { role: "user", content: [{ type: "text", text: "again" }] }, + ]); + }); + + it("hoists images out of tool results into a user message and serializes prompt-tsx values", () => { + expect( + toChatCompletionMessages([ + message(USER, [toolResult("call_2", [text("screenshot:"), image([9], "image/jpeg"), { value: { node: 1 } }])]), + ]), + ).toEqual([ + { role: "tool", tool_call_id: "call_2", content: 'screenshot:{"node":1}' }, + { role: "user", content: [{ type: "image_url", image_url: { url: "data:image/jpeg;base64,CQ==" } }] }, + ]); + }); + + it("decodes text data parts and ignores unknown parts", () => { + expect(toChatCompletionMessages([message(USER, [{ mimeType: "text/plain", data: Uint8Array.from([104, 105]) }, 42])])).toEqual([ + { role: "user", content: [{ type: "text", text: "hi" }] }, + ]); + }); +}); + +describe("estimateMessageTokens", () => { + it("counts what the gateway will receive, tool results and tool calls included", () => { + const plain = estimateMessageTokens(message(USER, [text("ok")])); + const withToolResult = estimateMessageTokens(message(USER, [toolResult("call_1", [text("y".repeat(800))]), text("ok")])); + const withToolCall = estimateMessageTokens(message(ASSISTANT, [toolCall("call_1", "read_file", { path: "z".repeat(800) })])); + expect(plain).toBeGreaterThan(0); + expect(withToolResult).toBeGreaterThanOrEqual(plain + 200); + expect(withToolCall).toBeGreaterThanOrEqual(200); + }); + + it("charges each image a flat estimate rather than its base64 length", () => { + const withoutImage = estimateMessageTokens(message(USER, [text("see")])); + const withImages = estimateMessageTokens(message(USER, [text("see"), image(new Array(30000).fill(0)), image([1])])); + expect(withImages - withoutImage).toBeGreaterThanOrEqual(2 * ESTIMATED_TOKENS_PER_IMAGE); + expect(withImages - withoutImage).toBeLessThan(2 * ESTIMATED_TOKENS_PER_IMAGE + 20); + }); + + it("counts nothing for a turn the gateway will never see", () => { + expect(estimateMessageTokens(message(ASSISTANT, []))).toBe(0); + }); +}); + +describe("buildChatCompletionParams", () => { + it("streams with usage and forwards only the chosen extras", () => { + expect(buildChatCompletionParams(request())).toEqual({ + model: "gpt-5.6", + messages: [{ role: "user", content: [{ type: "text", text: "hi" }] }], + stream: true, + stream_options: { include_usage: true }, + }); + }); + + it("declares tools as functions and requires a call only when VS Code does", () => { + const tools: readonly vscode.LanguageModelChatTool[] = [ + { name: "read_file", description: "Read a file", inputSchema: { type: "object", properties: { path: { type: "string" } } } }, + { name: "noop", description: "No input" }, + ]; + const auto = buildChatCompletionParams(request({ tools })); + expect(auto.tools).toEqual([ + { + type: "function", + function: { name: "read_file", description: "Read a file", parameters: { type: "object", properties: { path: { type: "string" } } } }, + }, + { type: "function", function: { name: "noop", description: "No input" } }, + ]); + expect(auto.tool_choice).toBeUndefined(); + expect(buildChatCompletionParams(request({ tools, requireToolCall: true })).tool_choice).toBe("required"); + expect(buildChatCompletionParams(request({ requireToolCall: true })).tool_choice).toBeUndefined(); + }); + + it("sends reasoning_effort only when the user picked one", () => { + expect(buildChatCompletionParams(request({ reasoningEffort: "xhigh" })).reasoning_effort).toBe("xhigh"); + expect(buildChatCompletionParams(request()).reasoning_effort).toBeUndefined(); + }); + + it("forwards numeric sampling options and drops everything else", () => { + const params = buildChatCompletionParams( + request({ modelOptions: { temperature: 0.2, max_tokens: 500, seed: "7", foo: "bar", top_p: 0.9 } }), + ); + expect(params).toMatchObject({ temperature: 0.2, max_tokens: 500, top_p: 0.9 }); + expect(params).not.toHaveProperty("seed"); + expect(params).not.toHaveProperty("foo"); + }); +}); diff --git a/vscode-extension/test/models.test.ts b/vscode-extension/test/models.test.ts new file mode 100644 index 00000000000..584b71ef8d4 --- /dev/null +++ b/vscode-extension/test/models.test.ts @@ -0,0 +1,185 @@ +import { describe, expect, it } from "vitest"; +import { + ASSUMED_MAX_INPUT_TOKENS, + ASSUMED_MAX_OUTPUT_TOKENS, + MARKDOWN_LINE_BREAK, + describeModels, + estimateTokens, + formatUsdPerMillionTokens, + parseModelGroups, + reasoningEffortFrom, + type ModelGroupInfo, +} from "../src/models"; + +const gatewayGroup = (overrides: Partial> = {}): Record => ({ + model_group: "gpt-5.6", + providers: ["openai"], + max_input_tokens: 922000, + max_output_tokens: 128000, + input_cost_per_token: 4e-6, + output_cost_per_token: 2e-5, + mode: "chat", + supports_vision: true, + supports_function_calling: true, + supports_reasoning: true, + supported_reasoning_efforts: ["none", "low", "medium", "high", "xhigh"], + ...overrides, +}); + +const parsed = (...groups: readonly Record[]): readonly ModelGroupInfo[] => { + const result = parseModelGroups({ data: groups }); + if (result.kind !== "ok") { + throw new Error(result.reason); + } + return result.groups; +}; + +describe("parseModelGroups", () => { + it("maps the gateway's /model_group/info shape", () => { + expect(parsed(gatewayGroup())).toEqual([ + { + modelGroup: "gpt-5.6", + providers: ["openai"], + mode: "chat", + maxInputTokens: 922000, + maxOutputTokens: 128000, + inputCostPerToken: 4e-6, + outputCostPerToken: 2e-5, + supportsVision: true, + supportsFunctionCalling: true, + supportedReasoningEfforts: ["none", "low", "medium", "high", "xhigh"], + }, + ]); + }); + + it("treats null limits, prices, and efforts as unknown", () => { + const [group] = parsed( + gatewayGroup({ + max_input_tokens: null, + max_output_tokens: null, + input_cost_per_token: null, + output_cost_per_token: null, + supported_reasoning_efforts: null, + supports_vision: null, + }), + ); + expect(group).toMatchObject({ + maxInputTokens: undefined, + inputCostPerToken: undefined, + supportedReasoningEfforts: [], + supportsVision: false, + }); + }); + + it("drops entries without a model_group and rejects payloads without data", () => { + expect(parsed({ providers: ["openai"] }, gatewayGroup()).map((group) => group.modelGroup)).toEqual(["gpt-5.6"]); + expect(parseModelGroups({ detail: "Unauthorized" })).toEqual({ kind: "invalid", reason: "response has no data array" }); + }); +}); + +describe("describeModels", () => { + it("lists chat groups with USD pricing in the detail and a full tooltip", () => { + const [model] = describeModels(parsed(gatewayGroup())); + expect(model).toMatchObject({ + id: "gpt-5.6", + name: "gpt-5.6", + family: "gpt-5.6", + detail: "$4.00 in / $20.00 out per 1M tokens", + maxInputTokens: 922000, + maxOutputTokens: 128000, + imageInput: true, + toolCalling: true, + }); + expect(model?.tooltip).toBe( + [ + "LiteLLM model group gpt-5.6 via openai", + "Input: $4.00 per 1M tokens", + "Output: $20.00 per 1M tokens", + "Context: 922000 in / 128000 out tokens", + "Reasoning effort: none, low, medium, high, xhigh", + ].join(MARKDOWN_LINE_BREAK), + ); + }); + + it("offers the gateway's reasoning efforts behind a gateway default entry", () => { + const [model] = describeModels(parsed(gatewayGroup({ supported_reasoning_efforts: ["low", "high"] }))); + expect(model?.configurationSchema).toEqual({ + properties: { + reasoningEffort: { + type: "string", + title: "Reasoning Effort", + enum: ["default", "low", "high"], + enumItemLabels: ["Gateway default", "Low", "High"], + default: "default", + group: "navigation", + }, + }, + }); + }); + + it("has no configuration schema when the group lists no reasoning efforts", () => { + const [model] = describeModels(parsed(gatewayGroup({ supported_reasoning_efforts: null }))); + expect(model?.configurationSchema).toBeUndefined(); + expect(model?.tooltip).toContain("Reasoning effort: not configurable"); + }); + + it("keeps groups without a mode and skips non-chat groups", () => { + const models = describeModels( + parsed( + gatewayGroup({ model_group: "text-embedding-4", mode: "embedding" }), + gatewayGroup({ model_group: "whisper-3", mode: "audio_transcription" }), + gatewayGroup({ model_group: "gpt-image-2", mode: "image_generation" }), + gatewayGroup({ model_group: "unlabeled", mode: null }), + gatewayGroup(), + ), + ); + expect(models.map((model) => model.id)).toEqual(["unlabeled", "gpt-5.6"]); + }); + + it("falls back to assumed context limits and says so", () => { + const [model] = describeModels(parsed(gatewayGroup({ max_input_tokens: null, max_output_tokens: null }))); + expect(model).toMatchObject({ maxInputTokens: ASSUMED_MAX_INPUT_TOKENS, maxOutputTokens: ASSUMED_MAX_OUTPUT_TOKENS }); + expect(model?.tooltip).toContain(`Context: unknown, assuming ${ASSUMED_MAX_INPUT_TOKENS} in / ${ASSUMED_MAX_OUTPUT_TOKENS} out tokens`); + }); + + it("shows missing prices instead of inventing zeros", () => { + const [both, inputOnly, free] = describeModels( + parsed( + gatewayGroup({ input_cost_per_token: null, output_cost_per_token: null }), + gatewayGroup({ output_cost_per_token: null }), + gatewayGroup({ input_cost_per_token: 0, output_cost_per_token: 0 }), + ), + ); + expect(both?.detail).toBe("No pricing configured"); + expect(both?.tooltip).toContain("Input: no price configured"); + expect(inputOnly?.detail).toBe("$4.00 in / n/a out per 1M tokens"); + expect(free?.detail).toBe("$0.00 in / $0.00 out per 1M tokens"); + }); +}); + +describe("formatUsdPerMillionTokens", () => { + it("renders cents for ordinary prices and two significant digits below a cent", () => { + expect(formatUsdPerMillionTokens(4e-6)).toBe("$4.00"); + expect(formatUsdPerMillionTokens(7.5e-7)).toBe("$0.75"); + expect(formatUsdPerMillionTokens(2.5e-5)).toBe("$25.00"); + expect(formatUsdPerMillionTokens(1e-9)).toBe("$0.0010"); + expect(formatUsdPerMillionTokens(0)).toBe("$0.00"); + }); +}); + +describe("reasoningEffortFrom", () => { + it("forwards a chosen effort and leaves the gateway default unset", () => { + expect(reasoningEffortFrom({ reasoningEffort: "high" })).toBe("high"); + expect(reasoningEffortFrom({ reasoningEffort: "default" })).toBeUndefined(); + expect(reasoningEffortFrom({ reasoningEffort: 3 })).toBeUndefined(); + expect(reasoningEffortFrom(undefined)).toBeUndefined(); + }); +}); + +describe("estimateTokens", () => { + it("rounds four characters per token upward", () => { + expect(estimateTokens("")).toBe(0); + expect(estimateTokens("abcd")).toBe(1); + expect(estimateTokens("abcde")).toBe(2); + }); +}); diff --git a/vscode-extension/test/provider.test.ts b/vscode-extension/test/provider.test.ts new file mode 100644 index 00000000000..07fc4ea5b59 --- /dev/null +++ b/vscode-extension/test/provider.test.ts @@ -0,0 +1,258 @@ +import type { ChatCompletionChunk, ChatCompletionCreateParamsStreaming } from "openai/resources/chat/completions"; +import { describe, expect, it } from "vitest"; +import type * as vscode from "vscode"; +import type { GatewayClient, GatewayConfig, ModelGroupsResult } from "../src/gateway"; +import { ESTIMATED_TOKENS_PER_IMAGE } from "../src/messages"; +import { LiteLLMChatProvider, TRUNCATED_MESSAGE, type LiteLLMModel } from "../src/provider"; +import { CancellationTokenSource, LanguageModelTextPart, LanguageModelToolCallPart } from "./vscode-mock"; + +interface StreamRequest { + readonly config: GatewayConfig; + readonly params: ChatCompletionCreateParamsStreaming; +} + +interface FakeGateway extends GatewayClient { + readonly listCalls: readonly GatewayConfig[]; + readonly streamRequests: readonly StreamRequest[]; +} + +const chunk = (delta: ChatCompletionChunk.Choice.Delta, finishReason: ChatCompletionChunk.Choice["finish_reason"] = null): ChatCompletionChunk => ({ + id: "chatcmpl-1", + object: "chat.completion.chunk", + created: 0, + model: "gpt-5.6", + choices: [{ index: 0, delta, finish_reason: finishReason }], +}); + +const modelGroups: ModelGroupsResult = { + kind: "ok", + groups: [ + { + modelGroup: "gpt-5.6", + providers: ["openai"], + mode: "chat", + maxInputTokens: 922000, + maxOutputTokens: 128000, + inputCostPerToken: 4e-6, + outputCostPerToken: 2e-5, + supportsVision: true, + supportsFunctionCalling: true, + supportedReasoningEfforts: ["low", "high"], + }, + ], +}; + +const fakeGateway = ( + listResult: ModelGroupsResult = modelGroups, + stream: (signal: AbortSignal) => AsyncIterable = () => (async function* () {})(), +): FakeGateway => { + const listCalls: GatewayConfig[] = []; + const streamRequests: StreamRequest[] = []; + return { + listCalls, + streamRequests, + async listModelGroups(config) { + listCalls.push(config); + return listResult; + }, + async streamChatCompletion(config, params, signal) { + streamRequests.push({ config, params }); + return stream(signal); + }, + }; +}; + +const token = (): vscode.CancellationToken => new CancellationTokenSource().token as unknown as vscode.CancellationToken; + +const prepare = (configuration: Record | undefined): vscode.PrepareLanguageModelChatModelOptions => + ({ silent: true, configuration }) as vscode.PrepareLanguageModelChatModelOptions; + +const gateway: GatewayConfig = { baseUrl: "http://127.0.0.1:4000", apiKey: "sk-test" }; + +const model = (overrides: Partial = {}): LiteLLMModel => ({ + id: "gpt-5.6", + name: "gpt-5.6", + family: "gpt-5.6", + version: "1.0", + maxInputTokens: 922000, + maxOutputTokens: 128000, + capabilities: { imageInput: true, toolCalling: true }, + gateway, + ...overrides, +}); + +const userMessage = (parts: readonly unknown[]): vscode.LanguageModelChatRequestMessage => + ({ role: 1, content: parts, name: undefined }) as vscode.LanguageModelChatRequestMessage; + +const responseOptions = (overrides: Partial = {}): vscode.ProvideLanguageModelChatResponseOptions => + ({ toolMode: 1, ...overrides }) as vscode.ProvideLanguageModelChatResponseOptions; + +const collect = ( + provider: LiteLLMChatProvider, + cancellation: CancellationTokenSource = new CancellationTokenSource(), +): { readonly parts: readonly vscode.LanguageModelResponsePart[]; readonly run: Promise } => { + const parts: vscode.LanguageModelResponsePart[] = []; + const run = provider.provideLanguageModelChatResponse( + model(), + [userMessage([{ value: "hi" }])], + responseOptions(), + { report: (part) => parts.push(part) }, + cancellation.token as unknown as vscode.CancellationToken, + ); + return { parts, run }; +}; + +describe("provideLanguageModelChatInformation", () => { + it("returns nothing for the unconfigured probe without touching the gateway", async () => { + const client = fakeGateway(); + expect(await new LiteLLMChatProvider(client).provideLanguageModelChatInformation(prepare(undefined), token())).toEqual([]); + expect(client.listCalls).toEqual([]); + }); + + it("names the API key when the stored secret is gone instead of listing nothing", async () => { + const client = fakeGateway(); + await expect( + new LiteLLMChatProvider(client).provideLanguageModelChatInformation(prepare({ baseUrl: "http://127.0.0.1:4000" }), token()), + ).rejects.toThrow(/missing its API key/); + expect(client.listCalls).toEqual([]); + }); + + it("rejects a gateway URL that is not http or https", async () => { + await expect( + new LiteLLMChatProvider(fakeGateway()).provideLanguageModelChatInformation(prepare({ baseUrl: "litellm.example.com", apiKey: "sk" }), token()), + ).rejects.toThrow(/"litellm.example.com" is not an http or https URL/); + }); + + it("lists the gateway's chat models with pricing, effort choices, and the gateway attached", async () => { + const client = fakeGateway(); + const models = await new LiteLLMChatProvider(client).provideLanguageModelChatInformation( + prepare({ baseUrl: "http://127.0.0.1:4000/v1/", apiKey: "sk-test" }), + token(), + ); + expect(client.listCalls).toEqual([gateway]); + expect(models).toEqual([ + expect.objectContaining({ + id: "gpt-5.6", + detail: "$4.00 in / $20.00 out per 1M tokens", + maxInputTokens: 922000, + capabilities: { imageInput: true, toolCalling: true }, + configurationSchema: expect.objectContaining({ properties: expect.objectContaining({ reasoningEffort: expect.anything() }) }), + gateway, + }), + ]); + }); + + it("shows the gateway's error message, not its whole JSON body, when discovery fails", async () => { + const body = JSON.stringify({ + error: { message: "Authentication Error, Invalid proxy server token passed", type: "auth_error", param: "sk-...abcd", code: "401" }, + }); + await expect( + new LiteLLMChatProvider(fakeGateway({ kind: "http_error", status: 401, body })).provideLanguageModelChatInformation( + prepare({ baseUrl: "http://127.0.0.1:4000", apiKey: "sk-bad" }), + token(), + ), + ).rejects.toThrow("LiteLLM gateway at http://127.0.0.1:4000 answered 401 for /model_group/info: Authentication Error, Invalid proxy server token passed"); + }); +}); + +describe("provideLanguageModelChatResponse", () => { + it("streams text and tool calls with the picked reasoning effort and a required tool choice", async () => { + const client = fakeGateway(modelGroups, () => + (async function* () { + yield chunk({ content: "Reading" }); + yield chunk({ tool_calls: [{ index: 0, id: "call_1", type: "function", function: { name: "read_file", arguments: '{"path":"a"}' } }] }); + yield chunk({}, "tool_calls"); + })(), + ); + const parts: vscode.LanguageModelResponsePart[] = []; + await new LiteLLMChatProvider(client).provideLanguageModelChatResponse( + model(), + [userMessage([{ value: "read a" }])], + responseOptions({ toolMode: 2, tools: [{ name: "read_file", description: "Read" }], modelConfiguration: { reasoningEffort: "high" } }), + { report: (part) => parts.push(part) }, + token(), + ); + expect(parts).toEqual([new LanguageModelTextPart("Reading"), new LanguageModelToolCallPart("call_1", "read_file", { path: "a" })]); + expect(client.streamRequests).toEqual([ + { + config: gateway, + params: expect.objectContaining({ model: "gpt-5.6", reasoning_effort: "high", tool_choice: "required", tools: [expect.anything()] }), + }, + ]); + }); + + it("finishes quietly when the user cancels mid-stream and drops its cancellation listener", async () => { + const cancellation = new CancellationTokenSource(); + const client = fakeGateway(modelGroups, (signal) => + (async function* () { + yield chunk({ content: "partial" }); + await new Promise((resolve) => signal.addEventListener("abort", () => resolve(), { once: true })); + throw new Error("Request was aborted."); + })(), + ); + const { parts, run } = collect(new LiteLLMChatProvider(client), cancellation); + await new Promise((resolve) => setTimeout(resolve, 0)); + cancellation.cancel(); + await expect(run).resolves.toBeUndefined(); + expect(parts).toEqual([new LanguageModelTextPart("partial")]); + expect(cancellation.disposedListeners).toBe(1); + }); + + it("surfaces a gateway failure as an error and still drops its cancellation listener", async () => { + const cancellation = new CancellationTokenSource(); + const client = fakeGateway(modelGroups, () => + (async function* () { + throw new Error("502 Bad Gateway"); + })(), + ); + const { run } = collect(new LiteLLMChatProvider(client), cancellation); + await expect(run).rejects.toThrow("502 Bad Gateway"); + expect(cancellation.disposedListeners).toBe(1); + }); + + it("reports the text it got and then fails when the model hits its output limit", async () => { + const client = fakeGateway(modelGroups, () => + (async function* () { + yield chunk({ content: "half an ans" }); + yield chunk({}, "length"); + })(), + ); + const { parts, run } = collect(new LiteLLMChatProvider(client)); + await expect(run).rejects.toThrow(TRUNCATED_MESSAGE); + expect(parts).toEqual([new LanguageModelTextPart("half an ans")]); + }); + + it("fails on tool arguments that are not JSON", async () => { + const client = fakeGateway(modelGroups, () => + (async function* () { + yield chunk({ tool_calls: [{ index: 0, id: "call_1", type: "function", function: { name: "grep", arguments: "{oops" } }] }); + })(), + ); + const { run } = collect(new LiteLLMChatProvider(client)); + await expect(run).rejects.toThrow("invalid JSON arguments for tool grep"); + }); +}); + +describe("provideTokenCount", () => { + const provider = new LiteLLMChatProvider(fakeGateway()); + + it("estimates plain text at four characters per token", async () => { + expect(await provider.provideTokenCount(model(), "abcdefgh")).toBe(2); + }); + + it("counts tool results and tool calls, not only text parts", async () => { + const textOnly = await provider.provideTokenCount(model(), userMessage([{ value: "ok" }])); + const withToolResult = await provider.provideTokenCount( + model(), + userMessage([{ callId: "call_1", content: [{ value: "x".repeat(400) }] }, { value: "ok" }]), + ); + expect(withToolResult).toBeGreaterThan(textOnly + 100); + }); + + it("charges a flat estimate per image instead of counting its bytes", async () => { + const withImage = await provider.provideTokenCount(model(), userMessage([{ value: "see" }, { mimeType: "image/png", data: new Uint8Array(50000) }])); + const withoutImage = await provider.provideTokenCount(model(), userMessage([{ value: "see" }])); + expect(withImage - withoutImage).toBeGreaterThanOrEqual(ESTIMATED_TOKENS_PER_IMAGE); + expect(withImage - withoutImage).toBeLessThan(ESTIMATED_TOKENS_PER_IMAGE + 20); + }); +}); diff --git a/vscode-extension/test/stream.test.ts b/vscode-extension/test/stream.test.ts new file mode 100644 index 00000000000..eb37013fb10 --- /dev/null +++ b/vscode-extension/test/stream.test.ts @@ -0,0 +1,93 @@ +import { describe, expect, it } from "vitest"; +import type { ChatCompletionChunk } from "openai/resources/chat/completions"; +import { responseParts, type ResponsePart } from "../src/stream"; + +type ToolCallDelta = NonNullable[number]; + +const chunk = (delta: ChatCompletionChunk.Choice.Delta, finishReason: ChatCompletionChunk.Choice["finish_reason"] = null): ChatCompletionChunk => ({ + id: "chatcmpl-1", + object: "chat.completion.chunk", + created: 0, + model: "gpt-5.6", + choices: [{ index: 0, delta, finish_reason: finishReason }], +}); + +const usageChunk: ChatCompletionChunk = { + id: "chatcmpl-1", + object: "chat.completion.chunk", + created: 0, + model: "gpt-5.6", + choices: [], + usage: { prompt_tokens: 3, completion_tokens: 2, total_tokens: 5 }, +}; + +async function* stream(chunks: readonly ChatCompletionChunk[]): AsyncGenerator { + yield* chunks; +} + +const collect = async (chunks: readonly ChatCompletionChunk[]): Promise => { + const parts: ResponsePart[] = []; + for await (const part of responseParts(stream(chunks))) { + parts.push(part); + } + return parts; +}; + +describe("responseParts", () => { + it("yields text deltas as they arrive and ignores empty and usage-only chunks", async () => { + expect(await collect([chunk({ role: "assistant", content: "" }), chunk({ content: "Hel" }), chunk({ content: "lo" }), usageChunk])).toEqual([ + { kind: "text", value: "Hel" }, + { kind: "text", value: "lo" }, + ]); + }); + + it("assembles tool calls split across chunks and emits them after the text, in index order", async () => { + expect( + await collect([ + chunk({ content: "Looking" }), + chunk({ tool_calls: [{ index: 1, id: "call_b", type: "function", function: { name: "grep", arguments: "" } }] }), + chunk({ tool_calls: [{ index: 0, id: "call_a", type: "function", function: { name: "read_file", arguments: '{"pa' } }] }), + chunk({ tool_calls: [{ index: 0, function: { name: "read_file", arguments: 'th":"a"}' } }] }), + chunk({ tool_calls: [{ index: 1, function: { arguments: '{"q":"x"}' } }] }, "tool_calls"), + ]), + ).toEqual([ + { kind: "text", value: "Looking" }, + { kind: "tool_call", callId: "call_a", name: "read_file", input: { path: "a" } }, + { kind: "tool_call", callId: "call_b", name: "grep", input: { q: "x" } }, + ]); + }); + + it("starts a new call when a fresh id reuses an index and appends index-less deltas to the last call", async () => { + expect( + await collect([ + chunk({ tool_calls: [{ index: 0, id: "call_a", type: "function", function: { name: "grep", arguments: '{"q":' } }] }), + chunk({ tool_calls: [{ function: { arguments: '"a"}' } } as ToolCallDelta] }), + chunk({ tool_calls: [{ index: 0, id: "call_b", type: "function", function: { name: "grep", arguments: '{"q":"b"}' } }] }), + ]), + ).toEqual([ + { kind: "tool_call", callId: "call_a", name: "grep", input: { q: "a" } }, + { kind: "tool_call", callId: "call_b", name: "grep", input: { q: "b" } }, + ]); + }); + + it("flags a response cut off at the output token limit after the text it did produce", async () => { + expect(await collect([chunk({ content: "half" }), chunk({}, "length"), usageChunk])).toEqual([ + { kind: "text", value: "half" }, + { kind: "truncated" }, + ]); + }); + + it("treats empty arguments as an empty object and flags malformed JSON", async () => { + expect( + await collect([ + chunk({ tool_calls: [{ index: 0, id: "call_0", type: "function", function: { name: "noop", arguments: "" } }] }), + chunk({ tool_calls: [{ index: 1, id: "call_1", type: "function", function: { name: "bad", arguments: "{oops" } }] }), + chunk({ tool_calls: [{ index: 2, id: "call_2", type: "function", function: { name: "scalar", arguments: "42" } }] }), + ]), + ).toEqual([ + { kind: "tool_call", callId: "call_0", name: "noop", input: {} }, + { kind: "invalid_tool_call", callId: "call_1", name: "bad", arguments: "{oops" }, + { kind: "invalid_tool_call", callId: "call_2", name: "scalar", arguments: "42" }, + ]); + }); +}); diff --git a/vscode-extension/test/vscode-mock.ts b/vscode-extension/test/vscode-mock.ts new file mode 100644 index 00000000000..5d2e8577201 --- /dev/null +++ b/vscode-extension/test/vscode-mock.ts @@ -0,0 +1,82 @@ +export class LanguageModelTextPart { + constructor(readonly value: string) {} +} + +export class LanguageModelToolCallPart { + constructor( + readonly callId: string, + readonly name: string, + readonly input: object, + ) {} +} + +export const LanguageModelChatToolMode = { Auto: 1, Required: 2 } as const; + +type Listener = (value: T) => void; + +interface Subscription { + dispose(): void; +} + +export class EventEmitter { + private readonly listeners: Listener[] = []; + + readonly event = (listener: Listener): Subscription => { + this.listeners.push(listener); + return { + dispose: () => { + const index = this.listeners.indexOf(listener); + if (index >= 0) { + this.listeners.splice(index, 1); + } + }, + }; + }; + + fire(value: T): void { + [...this.listeners].forEach((listener) => listener(value)); + } + + dispose(): void { + this.listeners.splice(0); + } +} + +export interface MockCancellationToken { + readonly isCancellationRequested: boolean; + onCancellationRequested(listener: Listener): Subscription; +} + +export class CancellationTokenSource { + private readonly emitter = new EventEmitter(); + private cancelled = false; + private disposed = 0; + readonly token: MockCancellationToken; + + constructor() { + const source = this; + this.token = { + get isCancellationRequested(): boolean { + return source.cancelled; + }, + onCancellationRequested: (listener) => { + const subscription = source.emitter.event(listener); + return { + dispose: () => { + source.disposed += 1; + subscription.dispose(); + }, + }; + }, + }; + } + + get disposedListeners(): number { + return this.disposed; + } + + cancel(): void { + this.cancelled = true; + this.emitter.fire(); + } +} diff --git a/vscode-extension/tsconfig.json b/vscode-extension/tsconfig.json new file mode 100644 index 00000000000..6835fa9b672 --- /dev/null +++ b/vscode-extension/tsconfig.json @@ -0,0 +1,19 @@ +{ + "compilerOptions": { + "target": "ES2022", + "module": "ESNext", + "moduleResolution": "Bundler", + "lib": ["ES2022"], + "types": ["node"], + "strict": true, + "noUncheckedIndexedAccess": true, + "noImplicitOverride": true, + "isolatedModules": true, + "verbatimModuleSyntax": true, + "esModuleInterop": true, + "resolveJsonModule": true, + "skipLibCheck": true, + "noEmit": true + }, + "include": ["src", "test"] +} diff --git a/vscode-extension/vitest.config.mts b/vscode-extension/vitest.config.mts new file mode 100644 index 00000000000..cf1533f7f16 --- /dev/null +++ b/vscode-extension/vitest.config.mts @@ -0,0 +1,13 @@ +import { fileURLToPath } from "node:url"; +import { defineConfig } from "vitest/config"; + +export default defineConfig({ + resolve: { + alias: { + vscode: fileURLToPath(new URL("./test/vscode-mock.ts", import.meta.url)), + }, + }, + test: { + include: ["test/**/*.test.ts"], + }, +}); From c0b0ba20b23a83a5e382ee2b7de7739cf5dcaa10 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:41:40 -0700 Subject: [PATCH 165/253] fix(azure-ai): keep FLUX.2 tolerant of OpenAI-only image params --- .../image_edit/flux2_transformation.py | 17 +--------- .../image_generation/flux_transformation.py | 34 +++++++++++++++---- ...test_azure_ai_image_edit_transformation.py | 11 ++++++ .../test_azure_ai_flux2_image_generation.py | 34 +++++++++++++++---- 4 files changed, 67 insertions(+), 29 deletions(-) diff --git a/litellm/llms/azure_ai/image_edit/flux2_transformation.py b/litellm/llms/azure_ai/image_edit/flux2_transformation.py index aa8905e5601..f91a87ba0f4 100644 --- a/litellm/llms/azure_ai/image_edit/flux2_transformation.py +++ b/litellm/llms/azure_ai/image_edit/flux2_transformation.py @@ -31,22 +31,7 @@ class AzureFoundryFlux2ImageEditConfig(OpenAIImageEditConfig): """ def get_supported_openai_params(self, model: str) -> list: - """ - FLUX 2 supports a subset of OpenAI image edit params - """ - return [ - "n", - "size", - "width", - "height", - "num_images", - "seed", - "safety_tolerance", - "output_format", - "aspect_ratio", - "guidance", - "steps", - ] + return AzureFoundryFluxImageGenerationConfig().get_supported_openai_params(model) def map_openai_params( self, diff --git a/litellm/llms/azure_ai/image_generation/flux_transformation.py b/litellm/llms/azure_ai/image_generation/flux_transformation.py index 997205b1fc3..ac9ec24420b 100644 --- a/litellm/llms/azure_ai/image_generation/flux_transformation.py +++ b/litellm/llms/azure_ai/image_generation/flux_transformation.py @@ -2,9 +2,18 @@ from collections.abc import Mapping from types import MappingProxyType from typing import Final +from litellm.exceptions import BadRequestError, UnsupportedParamsError from litellm.llms.openai.image_generation import GPTImageGenerationConfig from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams +FLUX2_DROPPED_OPENAI_PARAMS: Final[tuple[OpenAIImageGenerationOptionalParams, ...]] = ( + "background", + "moderation", + "output_compression", + "quality", + "user", +) + class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): """Azure Foundry BFL API configuration for FLUX image generation.""" @@ -81,10 +90,13 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): "num_images", "guidance", "steps", + *FLUX2_DROPPED_OPENAI_PARAMS, ] @staticmethod - def _map_parameter(name: str, value: object) -> tuple[tuple[str, object], ...]: + def _map_parameter(name: str, value: object, model: str) -> tuple[tuple[str, object], ...]: + if name in FLUX2_DROPPED_OPENAI_PARAMS: + return () if isinstance(value, str): if name in ("n", "num_images", "width", "height", "steps", "seed", "safety_tolerance"): return (("num_images" if name == "n" else name, int(value)),) @@ -94,11 +106,17 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): return (("num_images", value),) if name != "size": return ((name, value),) + if str(value).lower() == "auto": + return () try: width, height = (int(dimension) for dimension in str(value).lower().split("x")) except (TypeError, ValueError): - raise ValueError(f"Invalid size format '{value}'. Expected 'WxH', for example '1024x1024'.") + raise BadRequestError( + message=f"Invalid size format '{value}'. Expected 'WxH', for example '1024x1024'.", + model=model, + llm_provider="azure_ai", + ) return (("width", width), ("height", height)) def map_openai_params( @@ -118,9 +136,13 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): supported_params: Final = self.get_supported_openai_params(model) unsupported_params: Final = tuple(name for name in non_default_params if name not in supported_params) if unsupported_params and not drop_params: - raise ValueError( - f"Parameters {unsupported_params} are not supported for model {model}. " - f"Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters." + raise UnsupportedParamsError( + message=( + f"Parameters {unsupported_params} are not supported for model {model}. " + f"Supported parameters are {supported_params}. Set drop_params=True to drop unsupported parameters." + ), + model=model, + llm_provider="azure_ai", ) mapped_params: Final[Mapping[str, object]] = MappingProxyType( @@ -128,7 +150,7 @@ class AzureFoundryFluxImageGenerationConfig(GPTImageGenerationConfig): mapped_name: mapped_value for name, value in non_default_params.items() if name in supported_params - for mapped_name, mapped_value in self._map_parameter(name, value) + for mapped_name, mapped_value in self._map_parameter(name, value, model) } ) return {**optional_params, **mapped_params} # mutable-ok: inherited config contract returns a dict diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py index d74afa88a6b..39001c1795b 100644 --- a/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_azure_ai_image_edit_transformation.py @@ -176,3 +176,14 @@ def test_flux2_image_edit_preserves_controls_and_pixel_cost(dimensions: Mapping[ ) assert response._hidden_params["response_cost"] == pytest.approx(5e-08 * 2048 * 1024 * 2) + + +def test_flux2_image_edit_accepts_and_drops_openai_only_parameters(): + optional_params: Final = ImageEditRequestUtils.get_optional_params_image_edit( + model="FLUX.2-pro", + image_edit_provider_config=AzureFoundryFlux2ImageEditConfig(), + image_edit_optional_params={"n": 1, "size": "auto", "quality": "high", "user": "end-user-1"}, + drop_params=False, + ) + + assert optional_params == {"num_images": 1} diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py index 388cca9f054..512e98b4151 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_azure_ai_flux2_image_generation.py @@ -15,7 +15,7 @@ from litellm.llms.azure_ai.image_generation.flux_transformation import ( AzureFoundryFluxImageGenerationConfig, ) from litellm.types.utils import ImageObject, ImageResponse -from litellm.utils import _invalidate_model_cost_lowercase_map +from litellm.utils import _invalidate_model_cost_lowercase_map, get_optional_params_image_gen @pytest.fixture(autouse=True) @@ -86,15 +86,35 @@ def test_flux2_flex_maps_openai_and_provider_parameters(): } -def test_flux2_flex_rejects_invalid_size(): - with pytest.raises(ValueError, match="Expected 'WxH'"): - AzureFoundryFluxImageGenerationConfig().map_openai_params( - non_default_params={"size": "large"}, - optional_params={}, +def test_flux2_flex_rejects_invalid_size_as_bad_request(): + with pytest.raises(litellm.BadRequestError, match="Expected 'WxH'") as raised: + get_optional_params_image_gen( model="FLUX.2-flex", - drop_params=False, + custom_llm_provider="azure_ai", + provider_config=AzureFoundryFluxImageGenerationConfig(), + size="large", ) + assert raised.value.status_code == 400 + + +@pytest.mark.parametrize("model", ("FLUX.2-pro", "FLUX.2-flex")) +def test_flux2_accepts_and_drops_openai_only_image_parameters(model: str): + optional_params: Final = get_optional_params_image_gen( + model=model, + custom_llm_provider="azure_ai", + provider_config=AzureFoundryFluxImageGenerationConfig(), + n=1, + size="auto", + quality="high", + user="end-user-1", + background="transparent", + moderation="low", + output_compression=50, + ) + + assert optional_params == {"num_images": 1} + def test_flux2_flex_model_info(): model_info = litellm.get_model_info( From 7cfec730d5df0b0994922608984b1660207ba735 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:42:08 -0700 Subject: [PATCH 166/253] fix(proxy): refuse config-owned keys on POST /config/update POST /config/update stored keys the config file owns and answered 200 while the file silently kept winning. Run the config-owned check for general_settings, litellm_settings, and router_settings before the first database read, the way /config/field/update already does, so a refused request stores nothing. The config reload re-loaded the merged settings as yaml settings, which turned every saved router setting read-only after one tick. Read the saved router settings row instead so database-owned values stay writable. The integration harness seeds num_retries through /config/update instead of the config file, which is what the effective-settings and observed routing tests need to keep exercising a database-owned value. --- .circleci/scripts/run_integration.sh | 3 + litellm/proxy/proxy_server.py | 63 ++++++++++++------- tests/integration/proxy_config.yaml | 1 - .../proxy/proxy_server/test_routes_config.py | 37 +++++++++++ tests/test_litellm/proxy/test_proxy_server.py | 53 ++++++++++++---- 5 files changed, 123 insertions(+), 34 deletions(-) diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh index 6fab6dd57db..80f9cb0a4e0 100644 --- a/.circleci/scripts/run_integration.sh +++ b/.circleci/scripts/run_integration.sh @@ -125,6 +125,9 @@ start_proxy() { start_proxy 4000 proxy.log proxy_pid="$launched_pid" .venv/bin/python .circleci/scripts/wait_integration_services.py +curl --noproxy '*' -sSf -X POST "$INTEGRATION_PROXY_URL/config/update" \ + -H "Authorization: Bearer $LITELLM_MASTER_KEY" -H 'Content-Type: application/json' \ + -d '{"router_settings": {"num_retries": 0}}' > "$results/seed-router-settings.json" if [ "$suite" = management ]; then export INTEGRATION_PEER_URL=http://127.0.0.1:4001 start_proxy 4001 peer.log diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bb0e432af51..4654dffcb11 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6869,9 +6869,7 @@ class ProxyConfig: self._add_callbacks_from_db_config(config_data) # router settings - await self._add_router_settings_from_db_config( - config_data=config_data, llm_router=llm_router, prisma_client=prisma_client - ) + await self._add_router_settings_from_db_config(llm_router=llm_router, prisma_client=prisma_client) return still_desired_ids @@ -7079,13 +7077,11 @@ class ProxyConfig: async def _add_router_settings_from_db_config( self, - config_data: Mapping[str, object], llm_router: Router | None, prisma_client: PrismaClient | None, ) -> None: if llm_router is None or prisma_client is None: return - self.router_settings.load_yaml(_as_settings_mapping(config_data.get("router_settings"))) db_router_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": "router_settings"} ) @@ -16944,6 +16940,46 @@ async def update_config( if prisma_client is None: raise Exception("No DB Connected") + requested_general_settings: Final[Mapping[str, JsonValue]] = ( + config_info.general_settings.model_dump(exclude_none=True, exclude_unset=True) + if config_info.general_settings is not None + else {} + ) + raw_litellm_settings: Final[Mapping[str, JsonValue]] = _CONFIG_SECTION_VALUES.validate_python( + config_info.litellm_settings if config_info.litellm_settings is not None else {} + ) + incoming_success_callback: Final = raw_litellm_settings.get("success_callback") + updated_litellm_settings: Final[Mapping[str, JsonValue]] = _CONFIG_SECTION_VALUES.validate_python( + { + **raw_litellm_settings, + **( + {"success_callback": normalize_callback_names(incoming_success_callback)} + if isinstance(incoming_success_callback, list) + else {} + ), + } + ) + typed_router_settings: Final[Mapping[str, JsonValue]] = ( + config_info.router_settings.model_dump(exclude_none=True) if config_info.router_settings is not None else {} + ) + router_settings_updates: Final[Mapping[str, JsonValue]] = { + **typed_router_settings, + **( + { + key: value + for key, value in raw_router_settings.items() + if key not in typed_router_settings and value is not None + } + if isinstance(raw_router_settings, dict) + else {} + ), + } + proxy_config.reject_config_owned_writes( + section_name="general_settings", changed_keys=requested_general_settings + ) + proxy_config.reject_config_owned_writes(section_name="litellm_settings", changed_keys=updated_litellm_settings) + proxy_config.reject_config_owned_writes(section_name="router_settings", changed_keys=router_settings_updates) + async def _read_section(param_name: str) -> dict: row: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first( where={"param_name": param_name} @@ -17014,15 +17050,9 @@ async def update_config( if config_info.litellm_settings is not None: existing = await _read_section("litellm_settings") before_litellm_settings: Final = copy.deepcopy(existing) - updated_litellm_settings: Final = dict(config_info.litellm_settings) - - incoming_cb = updated_litellm_settings.get("success_callback") - if isinstance(incoming_cb, list): - updated_litellm_settings["success_callback"] = normalize_callback_names(incoming_cb) - merged: Final = {**existing, **updated_litellm_settings} - incoming_cb = updated_litellm_settings.get("success_callback") + incoming_cb: Final = updated_litellm_settings.get("success_callback") existing_cb: Final = existing.get("success_callback") if isinstance(incoming_cb, list): if isinstance(existing_cb, list): @@ -17045,15 +17075,6 @@ async def update_config( if isinstance(raw_router_settings, dict): existing = await _read_section("router_settings") before_router_settings: Final = copy.deepcopy(existing) - typed_router_settings: Final = ( - config_info.router_settings.dict(exclude_none=True) if config_info.router_settings is not None else {} - ) - raw_router_settings_without_none: Final = { - key: value - for key, value in raw_router_settings.items() - if key not in typed_router_settings and value is not None - } - router_settings_updates: Final = {**typed_router_settings, **raw_router_settings_without_none} new_router_settings: Final = {**existing, **router_settings_updates} await _upsert_section("router_settings", new_router_settings) asyncio.create_task( diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index b6f9767c210..a3b07f76d2f 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -13,5 +13,4 @@ litellm_settings: host: os.environ/REDIS_HOST port: os.environ/REDIS_PORT router_settings: - num_retries: 0 disable_cooldowns: true diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index 19b36b026f9..0a2daf64463 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -139,6 +139,43 @@ def test_config_update_persists_disable_cooldowns(client, auth_as, mock_prisma, assert persisted["disable_cooldowns"] is True +@pytest.mark.parametrize( + ("section", "store_attr", "yaml_values", "changed_values"), + [ + ("general_settings", "settings", {"alerting": ["slack"]}, {"alerting": ["email"]}), + ("litellm_settings", "litellm_settings", {"success_callback": ["langfuse"]}, {"success_callback": ["otel"]}), + ("router_settings", "router_settings", {"num_retries": 0}, {"num_retries": 2}), + ], +) +def test_config_update_rejects_config_owned_keys_and_accepts_the_same_value( + client, auth_as, mock_prisma, monkeypatch, section, store_attr, yaml_values, changed_values +): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock()) + store = getattr(ps.proxy_config, store_attr) + store.load_yaml(yaml_values) + try: + with auth_as(LitellmUserRoles.PROXY_ADMIN): + rejected = client.post("/config/update", json={section: changed_values}) + rejected_message = rejected.json()["error"]["message"] + table.upsert.assert_not_called() + accepted = client.post("/config/update", json={section: yaml_values}) + finally: + store.load_yaml({}) + + assert rejected.status_code == 400 + assert f"{section} key '{next(iter(yaml_values))}' is set in the config file and cannot be changed here" in ( + rejected_message + ) + assert accepted.status_code == 200 + persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"]) + assert persisted[next(iter(yaml_values))] == yaml_values[next(iter(yaml_values))] + + def test_config_update_rejects_assistants_config(client, auth_as, mock_prisma, monkeypatch): from litellm.proxy import proxy_server as ps from litellm.proxy._types import LitellmUserRoles diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bc7ae556e7b..b6f10ca8f76 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -4919,8 +4919,8 @@ async def test_add_router_settings_from_db_config_merge_logic(): mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) # Call the method under test + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -4980,8 +4980,8 @@ async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_ mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5017,8 +5017,8 @@ async def test_add_router_settings_from_db_config_empty_db_list_still_clears_unc mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5042,8 +5042,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): mock_router.update_settings = MagicMock() # Test Case 1: No router provided + proxy_config.router_settings.load_yaml({"test": "value"}) await proxy_config._add_router_settings_from_db_config( - config_data={"router_settings": {"test": "value"}}, llm_router=None, prisma_client=MagicMock(), ) @@ -5051,8 +5051,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): mock_router.update_settings.assert_not_called() # Test Case 2: No prisma client provided + proxy_config.router_settings.load_yaml({"test": "value"}) await proxy_config._add_router_settings_from_db_config( - config_data={"router_settings": {"test": "value"}}, llm_router=mock_router, prisma_client=None, ) @@ -5065,8 +5065,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): config_data = {"router_settings": {"routing_strategy": "usage-based"}} + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5080,8 +5080,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): mock_db_config.param_value = {"db_setting": "db_value"} mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + proxy_config.router_settings.load_yaml({}) await proxy_config._add_router_settings_from_db_config( - config_data={}, # No router_settings in config llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5093,9 +5093,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): # Test Case 5: Both config and DB router_settings are None/empty mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None) - await proxy_config._add_router_settings_from_db_config( - config_data={}, llm_router=mock_router, prisma_client=mock_prisma_client - ) + proxy_config.router_settings.load_yaml({}) + await proxy_config._add_router_settings_from_db_config(llm_router=mock_router, prisma_client=mock_prisma_client) # Should not call update_settings when no settings exist mock_router.update_settings.assert_not_called() @@ -5107,8 +5106,8 @@ async def test_add_router_settings_from_db_config_edge_cases(): config_data = {"router_settings": {"config_setting": "config_value"}} + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5157,8 +5156,8 @@ async def test_add_router_settings_shallow_merge_behavior(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config) + proxy_config.router_settings.load_yaml(config_data["router_settings"]) await proxy_config._add_router_settings_from_db_config( - config_data=config_data, llm_router=mock_router, prisma_client=mock_prisma_client, ) @@ -5180,6 +5179,36 @@ async def test_add_router_settings_shallow_merge_behavior(): assert merged_settings["top_level"] == "config_top" +@pytest.mark.asyncio +async def test_router_settings_reload_keeps_db_values_writable(tmp_path, monkeypatch): + from litellm.proxy.proxy_server import ProxyConfig + + config_path: Final = tmp_path / "config.yaml" + config_path.write_text(yaml.safe_dump({"model_list": [], "router_settings": {"disable_cooldowns": True}})) + db_row: Final = types.SimpleNamespace(param_value={"num_retries": 0}) + + async def read_config_row(_prisma_client, param_name): + return db_row if param_name == "router_settings" else None + + mock_prisma_client: Final = MagicMock() + mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=db_row) + mock_router: Final = MagicMock() + monkeypatch.setattr(proxy_server_module, "get_config_param", read_config_row) + monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma_client) + monkeypatch.setattr(proxy_server_module, "store_model_in_db", True) + monkeypatch.setattr(proxy_server_module, "user_config_file_path", None) + proxy_config: Final = ProxyConfig() + + for _ in range(2): + await proxy_config.get_config(config_file_path=str(config_path)) + await proxy_config._add_router_settings_from_db_config(llm_router=mock_router, prisma_client=mock_prisma_client) + + assert mock_router.update_settings.call_args.kwargs == {"disable_cooldowns": True, "num_retries": 0} + assert proxy_config.router_settings.source("num_retries") == "db" + assert proxy_config.router_settings.rejected_writes({"num_retries": 3}) == () + assert proxy_config.router_settings.rejected_writes({"disable_cooldowns": False}) == ("disable_cooldowns",) + + @pytest.mark.asyncio async def test_model_info_v1_oci_secrets_not_leaked(): """ From 4f904f69b6235497f784839be5078f1f77d5255d Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 19:44:18 +0000 Subject: [PATCH 167/253] test(proxy): restore app routes after pass-through reload tests The two pass-through reload tests register real FastAPI routes on the shared proxy app and only clean up the internal registry, so any test running after them in the same worker sees stray /v1/kept-* and /v1/deleted-* routes. test_component_allowlists counts those as uncovered and fails. Snapshot app.routes and the registry up front and restore both in a finally block. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/test_proxy_server.py | 52 +++++++++++++------ 1 file changed, 35 insertions(+), 17 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index ce809eda847..19822032d1f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7536,25 +7536,35 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi deleted. The proxy's own registry of live pass-through routes is what decides whether a request is routed upstream or falls through to the auth error, so it has to lose the entry on the reload rather than at the next process restart.""" - from litellm.proxy.pass_through_endpoints.pass_through_endpoints import InitPassThroughEndpointHelpers - from litellm.proxy.proxy_server import ProxyConfig + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + InitPassThroughEndpointHelpers, + _registered_pass_through_routes, + ) + from litellm.proxy.proxy_server import ProxyConfig, app path: Final = f"/v1/deleted-{uuid.uuid4().hex[:8]}" db_endpoint: Final = {"id": "db-1", "path": path, "target": "https://example.com/post"} + prior_routes: Final = list(app.routes) + prior_registry: Final = dict(_registered_pass_through_routes) def live_routes() -> set[str]: return {route for route in InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() if path in route} settings: Final = patch("litellm.proxy.proxy_server.general_settings", {}) # test-quality-ok: the method reads this module global; no injection seam yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", None) # test-quality-ok: module global holding the YAML endpoints; this case has none - with settings, yaml_endpoints: - pc = ProxyConfig() - await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]}) - assert live_routes(), "the stored endpoint should be serving before the row is deleted" + try: + with settings, yaml_endpoints: + pc = ProxyConfig() + await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]}) + assert live_routes(), "the stored endpoint should be serving before the row is deleted" - await pc._update_general_settings(db_general_settings={}) + await pc._update_general_settings(db_general_settings={}) - assert live_routes() == set() + assert live_routes() == set() + finally: + app.routes[:] = prior_routes + _registered_pass_through_routes.clear() + _registered_pass_through_routes.update(prior_registry) @pytest.mark.asyncio @@ -7564,15 +7574,18 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout serving untouched. The stored entry never gets a route of its own.""" from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( InitPassThroughEndpointHelpers, + _registered_pass_through_routes, initialize_pass_through_endpoints, ) - from litellm.proxy.proxy_server import ProxyConfig + from litellm.proxy.proxy_server import ProxyConfig, app marker: Final = uuid.uuid4().hex[:8] config_path: Final = f"/v1/kept-{marker}" db_path: Final = f"/v1/ignored-{marker}" config_endpoint: Final = {"id": f"cfg-{marker}", "path": config_path, "target": "https://example.com/post"} db_endpoint: Final = {"id": f"db-{marker}", "path": db_path, "target": "https://example.com/post"} + prior_routes: Final = list(app.routes) + prior_registry: Final = dict(_registered_pass_through_routes) def live_paths() -> set[str]: registered: Final = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() @@ -7580,17 +7593,22 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]) # test-quality-ok: module global holding the YAML endpoints the reload merges in - with settings, yaml_endpoints: - await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint]) - assert live_paths() == {config_path} + try: + with settings, yaml_endpoints: + await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint]) + assert live_paths() == {config_path} - pc = ProxyConfig() - await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]}) - assert live_paths() == {config_path} + pc = ProxyConfig() + await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]}) + assert live_paths() == {config_path} - await pc._update_general_settings(db_general_settings={}) + await pc._update_general_settings(db_general_settings={}) - assert live_paths() == {config_path} + assert live_paths() == {config_path} + finally: + app.routes[:] = prior_routes + _registered_pass_through_routes.clear() + _registered_pass_through_routes.update(prior_registry) def _fill_user_api_key_cache(cache: DualCache, count: int) -> None: From 0743c97bef983d54d7881254fb7c0d5292b9b743 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 19:44:18 +0000 Subject: [PATCH 168/253] chore(deps): bump anyio to 4.14.2 in uv.lock Addresses GHSA-5p39-cfhj-2xmp and GHSA-82r6-8w77-94w6 flagged by the osv-scan job. Hand-edited the single anyio block so the lockfile stays byte-stable for CI's older uv version; dependency set is unchanged. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- uv.lock | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/uv.lock b/uv.lock index a5e60c68515..e30274c1f6a 100644 --- a/uv.lock +++ b/uv.lock @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.14.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/61/cc/a381afa6efea9f496eff839d4a6a1aed3bfafc7b3ab4b0d1b243a12573dd/anyio-4.14.2.tar.gz", hash = "sha256:cfa139f3ed1a23ee8f88a145ddb5ac7605b8bbfd8592baacd7ce3d8bb4313c7f", size = 260176, upload-time = "2026-07-12T20:29:07.082Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/da/35/f2287558c17e29fafc8ef3daf819bb9834061cfa43bff8014f7df7f63bdc/anyio-4.14.2-py3-none-any.whl", hash = "sha256:9f505dda5ac9f0c8309b5e8bd445a8c2bf7246f3ce950121e45ea15bc41d1494", size = 125813, upload-time = "2026-07-12T20:29:05.763Z" }, ] [[package]] From 8533cb9d5915de44bb8bcd6ca1955d5e0ec51989 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:44:42 -0700 Subject: [PATCH 169/253] test(proxy): drop the pass-through routes the reload tests add to the shared app The two pass-through reload tests register routes on the shared app and never take them off, so a later allowlist test in the same xdist worker finds routes that no component exposes. Restore the app's route list when each of those tests finishes. --- tests/test_litellm/proxy/test_proxy_server.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index b6f10ca8f76..87114e183a9 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7496,7 +7496,15 @@ async def test_update_general_settings_db_pass_through_endpoint_cannot_override_ assert still_open.api_key is None +@pytest.fixture +def app_routes_restored(): + routes_before: Final = tuple(app.router.routes) + yield + app.router.routes[:] = routes_before + + @pytest.mark.asyncio +@pytest.mark.usefixtures("app_routes_restored") async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_service(): """A pass-through route the database declared has to stop serving when that row is deleted. The proxy's own registry of live pass-through routes is what decides whether @@ -7524,6 +7532,7 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi @pytest.mark.asyncio +@pytest.mark.usefixtures("app_routes_restored") async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_routes(): """``pass_through_endpoints`` is config-owned once the file declares it, so writing and then deleting a stored row resolves to the same list both times and the config file's routes keep From 58beea227547cccf54dd70eb1f7ae3024b8a1a03 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:45:54 -0700 Subject: [PATCH 170/253] fix(passthrough): read the stream flag by truthiness in the timeout resolver --- litellm/passthrough/timeout_utils.py | 4 ++-- .../test_pass_through_endpoints.py | 14 ++++++++++++++ 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index 829277105e3..9284f7143b0 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -10,7 +10,6 @@ DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0 class _TimeoutFields(BaseModel): model_config = ConfigDict(frozen=True) - stream: bool = False stream_timeout: float | None = None timeout: float | None = None request_timeout: float | None = None @@ -61,10 +60,11 @@ def resolve_llm_passthrough_timeout( kwargs stream_timeout -> litellm_params stream_timeout -> router_stream_timeout, then the non-streaming chain above. """ + streaming: Final = bool((kwargs or {}).get("stream")) request: Final = _TimeoutFields.model_validate(kwargs or {}) deployment: Final = _TimeoutFields.model_validate(litellm_params or {}) stream_candidates: Final = ( - (request.stream_timeout, deployment.stream_timeout, router_stream_timeout) if request.stream else () + (request.stream_timeout, deployment.stream_timeout, router_stream_timeout) if streaming else () ) candidates: Final = ( *stream_candidates, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 54336800db4..abff897c4f5 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1175,6 +1175,20 @@ def test_resolve_llm_passthrough_timeout_stream_timeout_precedence(): ) +@pytest.mark.parametrize( + "stream, expected", + [(None, 90.0), (0, 90.0), ("", 90.0), (1, 1800.0), ("yes", 1800.0)], +) +def test_resolve_llm_passthrough_timeout_reads_stream_by_truthiness(stream: object, expected: float): + assert ( + resolve_llm_passthrough_timeout( + kwargs={"stream": stream}, + litellm_params={"stream_timeout": 1800, "timeout": 90}, + ) + == expected + ) + + @pytest.mark.asyncio async def test_pass_through_request_uses_resolved_timeout(): with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: From cf22e4b171da3b923765a7e597d78378f03fb9c9 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 19:38:38 +0000 Subject: [PATCH 171/253] fix(proxy): make SettingsStore.clear terminate when the config file owns a key MutableMapping.clear pops items until the mapping is empty, but __delitem__ keeps config-owned keys, so clear spun forever on any store that had loaded a config file. unittest.mock.patch.dict calls clear on exit, which is why the proxy-infra and proxy-endpoints shards hung at 99 percent until the 20 minute job timeout on every run since the store started refusing config-owned writes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/config_resolvers/settings_store.py | 4 ++++ .../config_resolvers/test_settings_store.py | 16 ++++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/litellm/proxy/config_resolvers/settings_store.py b/litellm/proxy/config_resolvers/settings_store.py index 079d262319f..a3d5da4b343 100644 --- a/litellm/proxy/config_resolvers/settings_store.py +++ b/litellm/proxy/config_resolvers/settings_store.py @@ -98,6 +98,10 @@ class SettingsStore(MutableMapping[str, JsonValue]): def __len__(self) -> int: return sum(1 for _ in self) + def clear(self) -> None: + self._runtime_values = _EMPTY_VALUES + self._deleted_runtime_keys = frozenset(key for key in self._keys() if not self.owned_by_config(key)) + def _clear_runtime(self) -> None: self._runtime_values = _EMPTY_VALUES self._deleted_runtime_keys = frozenset() diff --git a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py index aa410deac43..0b472938b74 100644 --- a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py +++ b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py @@ -168,6 +168,22 @@ def test_settings_store_refuses_a_runtime_write_to_a_config_owned_key() -> None: assert store.source("max_parallel_requests") == "config" +@pytest.mark.timeout(10) +def test_settings_store_clear_terminates_and_keeps_config_owned_keys() -> None: + store: Final = SettingsStore("general_settings") + store.load_yaml({"max_parallel_requests": 3}) + store.apply_db_row("general_settings", {"max_file_size_mb": 9}) + store["ui_access_mode"] = "admin_only" + before: Final = dict(store) + + store.clear() + cleared: Final = dict(store) + store.update(before) + + assert cleared == {"max_parallel_requests": 3} + assert dict(store) == before + + def test_settings_store_reports_the_config_owned_keys_a_write_would_change() -> None: store: Final = SettingsStore("general_settings") store.load_yaml({"max_parallel_requests": 3, "ui_access_mode": "admin_only"}) From 32e6a86319ea9967c4af7a00e26411fd91046871 Mon Sep 17 00:00:00 2001 From: jesus Date: Sat, 12 Sep 2026 01:00:44 +0000 Subject: [PATCH 172/253] fix(router): report null cost for unpriced deployments instead of 0 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/proxy_server.py | 5 +- litellm/router.py | 11 ++++- litellm/utils.py | 16 +++++++ .../proxy_server/test_routes_model_info.py | 37 +++++++++++++++ tests/test_litellm/test_router.py | 47 +++++++++++++++++++ 5 files changed, 113 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3b5f4236d22..f21c907c292 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -154,7 +154,7 @@ from litellm.types.utils import ( TextCompletionResponse, TokenCountResponse, ) -from litellm.utils import load_credentials_from_list +from litellm.utils import cost_map_omits_token_price, load_credentials_from_list if TYPE_CHECKING: from aiohttp import ClientSession @@ -13618,9 +13618,10 @@ def _enrich_model_info_with_litellm_data( discovered_model_info: Final = ( llm_router.get_discovered_model_info(model_info.get("id")) if llm_router is not None else MappingProxyType({}) ) + unpriced: Final = cost_map_omits_token_price(model_info.get("id"), litellm_model_info.get("key")) for k, v in MappingProxyType({**litellm_model_info, **discovered_model_info}).items(): if k not in model_info or (model_info[k] is None and k in discovered_model_info): - model_info[k] = v + model_info[k] = None if unpriced and k in ("input_cost_per_token", "output_cost_per_token") else v model["model_info"] = model_info # don't return the api key / vertex credentials # don't return the llm credentials diff --git a/litellm/router.py b/litellm/router.py index eec393d6cef..8a543373692 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10575,11 +10575,12 @@ class Router: 2. If not, check if litellm model name is in model info 3. If not, return None """ - from litellm.utils import _update_dictionary + from litellm.utils import _update_dictionary, cost_map_omits_token_price model_info: ModelInfo | None = None custom_model_info: dict | None = None litellm_model_name_model_info: ModelInfo | None = None + base_model_key: str | None = None try: custom_model_info = ( @@ -10606,6 +10607,7 @@ class Router: ## update litellm model info with base model info base_model_info: Final = copy.deepcopy(litellm.get_model_info(model=base_model)) if base_model_info is not None: + base_model_key = base_model_info.get("key") # Base model provides defaults, custom model info overrides custom_model_info = _update_dictionary( cast(dict, base_model_info), @@ -10633,6 +10635,13 @@ class Router: # custom_model_info already includes base_model defaults at this point, if applicable model_info = cast(ModelInfo, custom_model_info) + if model_info is None: + return None + builtin_key: Final = ( + litellm_model_name_model_info.get("key") if litellm_model_name_model_info is not None else None + ) + if cost_map_omits_token_price(model_id, builtin_key, base_model_key): + return cast(ModelInfo, {**model_info, "input_cost_per_token": None, "output_cost_per_token": None}) return model_info def _set_model_group_info(self, model_group: str, user_facing_model_group_name: str) -> ModelGroupInfo | None: diff --git a/litellm/utils.py b/litellm/utils.py index 2c9200fbad7..ec11895a8b5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3156,6 +3156,22 @@ def reapply_runtime_model_cost_registrations() -> None: register_model(model_cost=dict(_runtime_registered_model_cost)) # mutable-ok: snapshot, replay rewrites it +def cost_map_omits_token_price(*keys: object) -> bool: + """Whether the raw ``litellm.model_cost`` entries under ``keys`` exist but none carries a per-token price. + + ``get_model_info`` substitutes 0 for a missing price, which reads exactly like a declared + zero. Surfaces that report pricing use this to keep an unpriced deployment at ``None``. + """ + entries: Final = tuple( + entry + for entry in (litellm.model_cost.get(key) for key in keys if isinstance(key, str)) + if isinstance(entry, dict) + ) + return len(entries) > 0 and not any( + "input_cost_per_token" in entry or "output_cost_per_token" in entry for entry in entries + ) + + def register_model( model_cost: str | dict, *, diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index a1cf838ab6b..203a237113f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -286,6 +286,43 @@ def test_get_proxy_model_info_surfaces_supports_parallel_function_calling(local_ assert enriched["model_info"]["supports_parallel_function_calling"] is True +def test_model_info_reports_null_cost_for_unpriced_deployment_and_zero_for_declared_zero(): + """A deployment configured with no cost fields must not surface the 0 that ``get_model_info`` + defaults to, since the zero-cost budget bypass only honours a declared zero. The declared zero + and a catalog price still come through.""" + import litellm + + router = litellm.Router( + model_list=[ + { + "model_name": "vllm-unpriced", + "litellm_params": {"model": "openai/vllm-unpriced", "api_key": "x", "api_base": "http://vllm"}, + }, + { + "model_name": "vllm-free", + "litellm_params": { + "model": "openai/vllm-free", + "api_key": "x", + "api_base": "http://vllm", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + }, + {"model_name": "gpt-priced", "litellm_params": {"model": "gpt-4o", "api_key": "x"}}, + ] + ) + + def enriched_cost(model_name: str) -> tuple: + deployment = router.get_model_list(model_name=model_name)[0] + info = proxy_server._enrich_model_info_with_litellm_data({**deployment, "model_info": dict(deployment["model_info"])})["model_info"] + return info.get("input_cost_per_token"), info.get("output_cost_per_token") + + assert enriched_cost("vllm-unpriced") == (None, None) + assert enriched_cost("vllm-free") == (0, 0) + input_cost, output_cost = enriched_cost("gpt-priced") + assert input_cost > 0 and output_cost > 0 + + def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch): from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth from litellm.proxy.auth import model_checks diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 26f9803a022..ceb8fcf6fd1 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1888,6 +1888,53 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost(): assert result.output_cost_per_token is None +def test_model_group_info_cost_none_for_unpriced_deployment_but_zero_when_declared(): + """A deployment with no cost fields anywhere must report None, not the 0 that + get_model_info defaults to, so the reported price matches what the zero-cost + budget bypass accepts. A deployment declaring 0 keeps reporting 0.""" + router = litellm.Router( + model_list=[ + { + "model_name": "vllm-unpriced", + "litellm_params": { + "model": "openai/my-vllm-unpriced", + "api_key": "fake", + "api_base": "http://localhost:8000/v1", + }, + }, + { + "model_name": "vllm-free", + "litellm_params": { + "model": "openai/my-vllm-free", + "api_key": "fake", + "api_base": "http://localhost:8000/v1", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + }, + { + "model_name": "gpt-priced", + "litellm_params": {"model": "gpt-4o", "api_key": "fake"}, + }, + ] + ) + + unpriced = router.get_model_group_info(model_group="vllm-unpriced") + assert unpriced is not None + assert unpriced.input_cost_per_token is None + assert unpriced.output_cost_per_token is None + + free = router.get_model_group_info(model_group="vllm-free") + assert free is not None + assert free.input_cost_per_token == 0 + assert free.output_cost_per_token == 0 + + priced = router.get_model_group_info(model_group="gpt-priced") + assert priced is not None + assert priced.input_cost_per_token is not None and priced.input_cost_per_token > 0 + assert priced.output_cost_per_token is not None and priced.output_cost_per_token > 0 + + @pytest.mark.parametrize( "value,expected", [ From 170581e7101ba821a6bd6a2cb039718f899fc85d Mon Sep 17 00:00:00 2001 From: jesus Date: Sat, 12 Sep 2026 01:04:40 +0000 Subject: [PATCH 173/253] fix(router): annotate cast for unpriced model info Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/router.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/router.py b/litellm/router.py index 8a543373692..9ff23c597f1 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10641,7 +10641,9 @@ class Router: litellm_model_name_model_info.get("key") if litellm_model_name_model_info is not None else None ) if cost_map_omits_token_price(model_id, builtin_key, base_model_key): - return cast(ModelInfo, {**model_info, "input_cost_per_token": None, "output_cost_per_token": None}) + return cast( # cast-ok: TypedDict spread with overridden keys loses its type + ModelInfo, {**model_info, "input_cost_per_token": None, "output_cost_per_token": None} + ) return model_info def _set_model_group_info(self, model_group: str, user_facing_model_group_name: str) -> ModelGroupInfo | None: From 8549189cbbedc85ed067da8a247512cbac54122f Mon Sep 17 00:00:00 2001 From: jesus Date: Fri, 18 Sep 2026 19:59:24 +0000 Subject: [PATCH 174/253] test(proxy): use local_model_cost_map fixture for alias listing test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/test_proxy_utils.py | 37 +++++++++----------- 1 file changed, 16 insertions(+), 21 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index a79e0798e29..25f4886c6f9 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -2243,32 +2243,27 @@ def test_create_model_info_response_resolves_mode_through_deployment_model(): {"team-embeddings": {"model": "my-embeddings", "hidden": False}}, ], ) -def test_create_model_info_response_resolves_model_group_alias_to_target(model_group_alias): +def test_create_model_info_response_resolves_model_group_alias_to_target(model_group_alias, local_model_cost_map): """A `model_group_alias` row must report the metadata of the group it points at, not the cost-map generalization or nothing that the alias name resolves to.""" from litellm import Router - saved_model_cost = dict(litellm.model_cost) - try: - router = Router( - model_list=[ - { - "model_name": "my-embeddings", - "litellm_params": {"model": "openai/text-embedding-3-small"}, - } - ], - model_group_alias=model_group_alias, - ) + router = Router( + model_list=[ + { + "model_name": "my-embeddings", + "litellm_params": {"model": "openai/text-embedding-3-small"}, + } + ], + model_group_alias=model_group_alias, + ) - alias_response = create_model_info_response( - model_id="team-embeddings", provider="openai", llm_router=router - ) - target_response = create_model_info_response( - model_id="my-embeddings", provider="openai", llm_router=router - ) - finally: - litellm.model_cost.clear() - litellm.model_cost.update(saved_model_cost) + alias_response = create_model_info_response( + model_id="team-embeddings", provider="openai", llm_router=router + ) + target_response = create_model_info_response( + model_id="my-embeddings", provider="openai", llm_router=router + ) assert alias_response["id"] == "team-embeddings" for field in ("mode", "max_input_tokens", "max_output_tokens"): From 1e0613d8cae45e088bf122f8adfe42abdba5ab36 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:00:13 -0700 Subject: [PATCH 175/253] chore(vscode): pin @types/vscode to the 1.109 engine minimum --- vscode-extension/package-lock.json | 8 ++++---- vscode-extension/package.json | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/vscode-extension/package-lock.json b/vscode-extension/package-lock.json index 378579dc024..a4f962f2039 100644 --- a/vscode-extension/package-lock.json +++ b/vscode-extension/package-lock.json @@ -13,7 +13,7 @@ }, "devDependencies": { "@types/node": "^22.20.3", - "@types/vscode": "^1.109.0", + "@types/vscode": "1.109.0", "@vscode/vsce": "^4.0.0", "esbuild": "^0.28.2", "typescript": "^5.9.3", @@ -1301,9 +1301,9 @@ } }, "node_modules/@types/vscode": { - "version": "1.138.0", - "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.138.0.tgz", - "integrity": "sha512-DhlvrucJa8n6EHst7qQcXOcjLKs/VRTQy55plLO+RN74jonYsr/t9dd66zBnAyz1tpDig0okQAxuoPHpVNje2g==", + "version": "1.109.0", + "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.109.0.tgz", + "integrity": "sha512-0Pf95rnwEIwDbmXGC08r0B4TQhAbsHQ5UyTIgVgoieDe4cOnf92usuR5dEczb6bTKEp7ziZH4TV1TRGPPCExtw==", "dev": true, "license": "MIT" }, diff --git a/vscode-extension/package.json b/vscode-extension/package.json index a8e52b14eda..67e3bb6c727 100644 --- a/vscode-extension/package.json +++ b/vscode-extension/package.json @@ -78,7 +78,7 @@ }, "devDependencies": { "@types/node": "^22.20.3", - "@types/vscode": "^1.109.0", + "@types/vscode": "1.109.0", "@vscode/vsce": "^4.0.0", "esbuild": "^0.28.2", "typescript": "^5.9.3", From 657cf18aea50d8825b629d95f9e74467d17391d9 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 13:01:55 -0700 Subject: [PATCH 176/253] feat(rust): port exception_type to litellm-core-utils Adds a Rust port of litellm_core_utils/exception_mapping_utils.exception_type with its provider rule tables (OpenAI-compatible, Cohere, Vertex AI), the secret redaction patterns from secret_redaction.py, and a Python repr helper for messages that quote caller values. The public failure shapes are pinned by golden JSON fixtures under tests/test_litellm/rust_bridge/fixtures/public_failures, which the Python side reads too once the bridge is wired to this module. Nothing calls the port yet. --- litellm-rust/Cargo.lock | 29 + litellm-rust/Cargo.toml | 1 + litellm-rust/crates/core-utils/Cargo.toml | 5 + .../src/exception_mapping_utils/cohere.rs | 232 ++++++ .../src/exception_mapping_utils/mod.rs | 734 ++++++++++++++++++ .../src/exception_mapping_utils/openai.rs | 586 ++++++++++++++ .../src/exception_mapping_utils/original.rs | 49 ++ .../src/exception_mapping_utils/public.rs | 262 +++++++ .../src/exception_mapping_utils/rules.rs | 403 ++++++++++ .../src/exception_mapping_utils/status.rs | 153 ++++ .../src/exception_mapping_utils/vertex_ai.rs | 574 ++++++++++++++ litellm-rust/crates/core-utils/src/lib.rs | 3 + .../crates/core-utils/src/python_repr.rs | 93 +++ .../crates/core-utils/src/secret_redaction.rs | 97 +++ .../fixtures/public_failures/api.json | 13 + .../public_failures/api_connection.json | 11 + .../status_authentication.json | 13 + .../public_failures/status_bad_gateway.json | 28 + .../public_failures/status_bad_request.json | 28 + .../status_content_policy_violation.json | 28 + .../status_context_window_exceeded.json | 28 + .../status_internal_server.json | 19 + .../public_failures/status_not_found.json | 28 + .../status_permission_denied.json | 19 + .../public_failures/status_rate_limit.json | 28 + .../status_service_unavailable.json | 28 + .../status_unsupported_params.json | 28 + .../public_failures/timeout_with_status.json | 12 + .../timeout_without_status.json | 12 + 29 files changed, 3544 insertions(+) create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs create mode 100644 litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs create mode 100644 litellm-rust/crates/core-utils/src/python_repr.rs create mode 100644 litellm-rust/crates/core-utils/src/secret_redaction.rs create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/api.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json create mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 3b9ff62fbaa..6c524124550 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -559,6 +559,21 @@ dependencies = [ "vsimd", ] +[[package]] +name = "bit-set" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" +dependencies = [ + "bit-vec", +] + +[[package]] +name = "bit-vec" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" + [[package]] name = "bitflags" version = "2.13.1" @@ -1166,6 +1181,17 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "fancy-regex" +version = "0.19.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d301f5bf187b3c295fce6468d3875037a0bccc5f6b151c63cac2f85babf21912" +dependencies = [ + "bit-set", + "regex-automata", + "regex-syntax", +] + [[package]] name = "fastrand" version = "2.5.0" @@ -2059,11 +2085,14 @@ dependencies = [ name = "litellm-core-utils" version = "0.1.0" dependencies = [ + "fancy-regex", "litellm-types", + "rstest", "serde", "serde_json", "serde_path_to_error", "serde_with", + "strum", "thiserror 2.0.19", "url", ] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index f8377138050..ffdbf64bb49 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -50,6 +50,7 @@ strum = { version = "0.28.0", features = ["derive"] } url = "2.5.8" time = { version = "0.3.53", features = ["parsing"] } criterion = "0.8.2" +fancy-regex = "0.19.2" veil = "0.3.0" [profile.release] diff --git a/litellm-rust/crates/core-utils/Cargo.toml b/litellm-rust/crates/core-utils/Cargo.toml index 109c3312727..59e0ee1a09d 100644 --- a/litellm-rust/crates/core-utils/Cargo.toml +++ b/litellm-rust/crates/core-utils/Cargo.toml @@ -6,10 +6,15 @@ license.workspace = true repository.workspace = true [dependencies] +fancy-regex.workspace = true litellm-types.workspace = true serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" serde_with.workspace = true +strum.workspace = true thiserror.workspace = true url.workspace = true + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs new file mode 100644 index 00000000000..381ebd7ad27 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs @@ -0,0 +1,232 @@ +use super::Mapping; +use super::public::{PublicFailure, StatusClass}; +use super::rules::{Kind, ResponseChoice, Rule, apply, contains_any}; + +const fn with_response(class: StatusClass) -> Kind { + Kind::Status { + class, + response: ResponseChoice::Provider, + } +} + +fn original(mapping: &Mapping<'_>) -> String { + format!("CohereException - {}", mapping.original.message) +} + +fn status_is(mapping: &Mapping<'_>, statuses: &[u16]) -> bool { + mapping + .original + .status + .is_some_and(|status| statuses.contains(&status)) +} + +/// `_map_cohere_exception`, in its branch order. A failure no rule claims falls through to +/// the status table. +const RULES: &[Rule] = &[ + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &["invalid api token", "No API key provided."], + ) + }, + kind: with_response(StatusClass::Authentication), + message: original, + debug: false, + }, + Rule { + when: |mapping| mapping.error_str.contains("invalid type: parameter"), + kind: with_response(StatusClass::BadRequest), + message: original, + debug: false, + }, + Rule { + when: |mapping| mapping.error_str.contains("too many tokens"), + kind: with_response(StatusClass::ContextWindowExceeded), + message: original, + debug: false, + }, + Rule { + when: |mapping| { + mapping + .error_str + .to_lowercase() + .contains("internal server error") + }, + kind: with_response(StatusClass::InternalServer), + message: |mapping| format!("CohereException - {}", mapping.error_str), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, &[400, 498]), + kind: with_response(StatusClass::BadRequest), + message: original, + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, &[408]), + kind: Kind::Timeout(None), + message: original, + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, &[500]), + kind: with_response(StatusClass::InternalServer), + message: original, + debug: false, + }, +]; + +pub(super) fn map(mapping: &Mapping<'_>) -> Option { + apply(RULES, mapping).map(|failure| PublicFailure { + llm_provider: Some("cohere".to_string()), + ..failure + }) +} + +#[cfg(test)] +mod tests { + use super::super::testing::{context, failure, http, status, upstream}; + use super::super::{ExceptionFamily, OriginalException, PublicKind}; + use super::*; + + fn mapped(provider: &str, original: &OriginalException) -> Option { + let context = context(provider, ExceptionFamily::Cohere); + map(&Mapping::new(&context, original)) + } + + fn cohere(class: StatusClass, status_code: u16, body: &str, message: &str) -> PublicFailure { + failure( + status(class, upstream(status_code, body)), + message, + "cohere", + ) + } + + #[rstest::rstest] + #[case::invalid_token( + 500, + "invalid api token", + cohere( + StatusClass::Authentication, + 500, + "invalid api token", + "CohereException - invalid api token" + ) + )] + #[case::no_api_key( + 500, + "No API key provided.", + cohere( + StatusClass::Authentication, + 500, + "No API key provided.", + "CohereException - No API key provided." + ) + )] + #[case::invalid_parameter( + 500, + "invalid type: parameter x", + cohere( + StatusClass::BadRequest, + 500, + "invalid type: parameter x", + "CohereException - invalid type: parameter x" + ) + )] + #[case::too_many_tokens( + 500, + "too many tokens", + cohere( + StatusClass::ContextWindowExceeded, + 500, + "too many tokens", + "CohereException - too many tokens" + ) + )] + #[case::internal_server_text( + 400, + "Internal Server Error", + cohere( + StatusClass::InternalServer, + 400, + "Internal Server Error", + "CohereException - Internal Server Error" + ) + )] + #[case::bad_request( + 400, + "rejected", + cohere(StatusClass::BadRequest, 400, "rejected", "CohereException - rejected") + )] + #[case::invalid_token_status( + 498, + "rejected", + cohere(StatusClass::BadRequest, 498, "rejected", "CohereException - rejected") + )] + #[case::request_timeout(408, "rejected", failure(PublicKind::Timeout { status: None }, "CohereException - rejected", "cohere"))] + #[case::internal_server( + 500, + "rejected", + cohere( + StatusClass::InternalServer, + 500, + "rejected", + "CohereException - rejected" + ) + )] + fn each_rule_maps_and_reports_cohere( + #[case] status_code: u16, + #[case] body: &str, + #[case] expected: PublicFailure, + ) { + assert_eq!(mapped("azure_ai", &http(status_code, body)), Some(expected)); + } + + #[rstest::rstest] + #[case::unmapped_status(409)] + #[case::unauthorized(401)] + fn statuses_without_a_rule_fall_through(#[case] status_code: u16) { + assert_eq!(mapped("cohere", &http(status_code, "rejected")), None); + } + + #[test] + fn the_internal_server_rule_uses_the_redacted_text() { + let body = "internal server error Bearer abcdefghijklmnop"; + assert_eq!( + mapped("cohere", &http(400, body)), + Some(cohere( + StatusClass::InternalServer, + 400, + body, + "CohereException - internal server error REDACTED" + )) + ); + } + + #[rstest::rstest] + #[case::token_before_parameter( + "invalid api token invalid type: parameter", + StatusClass::Authentication + )] + #[case::parameter_before_tokens( + "invalid type: parameter too many tokens", + StatusClass::BadRequest + )] + #[case::tokens_before_internal( + "too many tokens Internal Server Error", + StatusClass::ContextWindowExceeded + )] + #[case::internal_before_status("Internal Server Error", StatusClass::InternalServer)] + fn the_earlier_rule_wins_when_two_apply(#[case] body: &str, #[case] class: StatusClass) { + assert_eq!( + mapped("cohere", &http(400, body)), + Some(cohere( + class, + 400, + body, + &format!("CohereException - {body}") + )) + ); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs new file mode 100644 index 00000000000..8892fba1399 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs @@ -0,0 +1,734 @@ +//! A port of Python's `exception_type` for the routes that run in Rust. +//! +//! KNOWN_GAPS: differences from the Python mapper that no Rust route can reach today. Each +//! one stops being acceptable at its trigger. +//! - The Vertex partner-model API base for "claude" models is not built into +//! `extra_information`. Trigger: a Vertex route whose models include Anthropic partner +//! models; then `api_base` gets that branch and a table row. +//! - Python reports the provider `get_llm_provider` resolves for a stripped model name when +//! that name happens to be in the model cost map. Trigger: a route whose model names +//! overlap the cost map; that needs the provider resolution port, not a classifier change. +//! - The generic `APIConnectionError` fallback appends `traceback.format_exc()` to the +//! message. Rust has no Python traceback and does not invent one; a sweep row that reaches +//! it compares the message before the traceback. + +use super::secret_redaction::{redact_string, secret_redaction_enabled}; + +mod cohere; +mod openai; +mod original; +mod public; +mod rules; +mod status; +mod vertex_ai; + +pub use original::{ExceptionFamily, LocalClass, OriginalException}; +pub use public::{HttpStub, PublicFailure, PublicKind, ResponseArg, StatusClass, UpstreamResponse}; + +const DOCS_URL: &str = "https://docs.litellm.ai/docs"; + +const TIMEOUT_MARKERS: &[&str] = &[ + "Request Timeout Error", + "Request timed out", + "Timed out generating response", + "The read operation timed out", +]; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct ExceptionContext { + pub model: String, + pub custom_llm_provider: Option, + pub family: ExceptionFamily, + pub asynchronous: bool, + pub suppress_debug_info: bool, + pub redact_messages_in_exceptions: bool, + pub vertex_project: Option, + pub vertex_location: Option, + pub model_group: Option, + pub deployment: Option, + pub user_api_key_alias: Option, + pub user_api_key_team_alias: Option, +} + +/// The attributes `exception_type` reads off the Python exception: a provider error +/// (`BaseLLMException`) carries a status, a response and a request, a plain exception +/// carries only its text. +struct Raised { + status: Option, + status_is_synthesized: bool, + message: String, + response: Option, +} + +impl Raised { + fn provider( + status: u16, + message: String, + body: String, + headers: Vec<(String, String)>, + ) -> Self { + Self { + status: Some(status), + status_is_synthesized: false, + message, + response: Some(UpstreamResponse { + status, + body, + headers, + }), + } + } + + fn plain(message: String) -> Self { + Self { + status: None, + status_is_synthesized: false, + message, + response: None, + } + } + + fn new(original: &OriginalException, asynchronous: bool) -> Self { + match original { + OriginalException::Http { + status, + body, + headers, + } => Self::provider(*status, body.clone(), body.clone(), headers.clone()), + OriginalException::Connection { message } => Self { + status_is_synthesized: true, + ..Self::provider(500, message.clone(), String::new(), Vec::new()) + }, + OriginalException::Timeout { + timeout_seconds, + elapsed_seconds, + } => Self::provider( + 408, + timeout_message(asynchronous, *timeout_seconds, *elapsed_seconds), + String::new(), + Vec::new(), + ), + OriginalException::Response { message } + | OriginalException::Local { message, .. } + | OriginalException::Public { message, .. } => Self::plain(message.clone()), + } + } +} + +/// The text `litellm.Timeout` carries when the Python HTTP handler times out: the sync +/// and async handlers word it differently. +fn timeout_message( + asynchronous: bool, + timeout_seconds: Option, + elapsed_seconds: Option, +) -> String { + let timeout = python_float(timeout_seconds); + if asynchronous { + let elapsed = + python_float(elapsed_seconds.map(|seconds| (seconds * 1000.0).round() / 1000.0)); + format!( + "litellm.Timeout: Connection timed out. Timeout passed={timeout}, time taken={elapsed} seconds" + ) + } else { + format!("litellm.Timeout: Connection timed out after {timeout} seconds.") + } +} + +fn python_float(value: Option) -> String { + match value { + None => "None".to_string(), + Some(value) if value.fract() == 0.0 => format!("{value:.1}"), + Some(value) => value.to_string(), + } +} + +/// Everything the rules read: the original as Python sees it and the text `exception_type` +/// derives from the context before any provider mapper runs. +struct Mapping<'a> { + context: &'a ExceptionContext, + original: Raised, + provider: &'a str, + error_str: String, + exception_provider: String, + extra_information: String, +} + +impl<'a> Mapping<'a> { + fn new(context: &'a ExceptionContext, original: &OriginalException) -> Self { + let original = Raised::new(original, context.asynchronous); + let error_str = if secret_redaction_enabled() { + redact_string(&original.message) + } else { + original.message.clone() + }; + Self { + context, + original, + provider: context.custom_llm_provider.as_deref().unwrap_or_default(), + error_str, + exception_provider: match &context.custom_llm_provider { + None => "None".to_string(), + Some(provider) => exception_provider(provider), + }, + extra_information: extra_information(context, api_base(context).as_deref()), + } + } + + fn failure(&self, kind: PublicKind, message: String, debug: bool) -> PublicFailure { + PublicFailure { + kind, + message, + model: self.context.model.clone(), + llm_provider: self.context.custom_llm_provider.clone(), + litellm_debug_info: debug.then(|| self.extra_information.clone()), + litellm_response_headers: None, + print_banner: false, + } + } +} + +pub fn exception_type(context: &ExceptionContext, original: &OriginalException) -> PublicFailure { + if let OriginalException::Public { class, message } = original { + return PublicFailure { + kind: PublicKind::Status { + status_class: *class, + response: None, + }, + message: message.clone(), + model: context.model.clone(), + llm_provider: context.custom_llm_provider.clone(), + litellm_debug_info: None, + litellm_response_headers: None, + print_banner: false, + }; + } + let mapping = Mapping::new(context, original); + let litellm_response_headers = mapping + .original + .response + .as_ref() + .map(|response| response.headers.clone()) + .filter(|headers| !headers.is_empty()); + PublicFailure { + litellm_response_headers, + print_banner: !context.suppress_debug_info, + ..map(&mapping) + } +} + +fn map(mapping: &Mapping<'_>) -> PublicFailure { + if rules::contains_any(&mapping.error_str, TIMEOUT_MARKERS) { + return mapping.failure( + PublicKind::Timeout { status: None }, + format!( + "APITimeoutError - Request timed out. Error_str: {}", + mapping.error_str + ), + true, + ); + } + let provider_failure = match mapping.context.family { + ExceptionFamily::OpenAiCompatible => openai::map(mapping), + ExceptionFamily::VertexAi => vertex_ai::map(mapping), + ExceptionFamily::Cohere => cohere::map(mapping), + ExceptionFamily::Other => None, + }; + provider_failure + .or_else(|| status::map(mapping)) + .unwrap_or_else(|| unmapped(mapping)) +} + +/// The `APIConnectionError` Python raises when no mapper claimed the failure: with the +/// provider prefix for a provider error, with the bare text for a plain exception. +fn unmapped(mapping: &Mapping<'_>) -> PublicFailure { + let message = match mapping.original.status { + Some(_) => format!("{} - {}", mapping.exception_provider, mapping.error_str), + None => mapping.original.message.clone(), + }; + mapping.failure(PublicKind::ApiConnection, message, false) +} + +fn exception_provider(provider: &str) -> String { + let mut characters = provider.chars(); + match characters.next() { + Some(first) => format!("{}{}Exception", first.to_uppercase(), characters.as_str()), + None => String::new(), + } +} + +fn python_capitalize(value: &str) -> String { + let mut characters = value.chars(); + match characters.next() { + Some(first) => format!( + "{}{}", + first.to_uppercase(), + characters.as_str().to_lowercase() + ), + None => String::new(), + } +} + +fn api_base(context: &ExceptionContext) -> Option { + match (&context.vertex_location, &context.vertex_project) { + (Some(location), Some(project)) => Some(format!( + "{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/publishers/google/models/{}:generateContent", + context.model + )), + _ => None, + } +} + +fn extra_information(context: &ExceptionContext, api_base: Option<&str>) -> String { + let lines = [ + Some(format!("\nModel: {}", context.model)), + api_base.map(|api_base| format!("\nAPI Base: `{api_base}`")), + (!context.redact_messages_in_exceptions).then(|| "\nMessages: `None`".to_string()), + context + .model_group + .as_ref() + .map(|value| format!("\nmodel_group: `{value}`\n")), + context + .deployment + .as_ref() + .map(|value| format!("\ndeployment: `{value}`\n")), + context + .vertex_project + .as_ref() + .map(|value| format!("\nvertex_project: `{value}`\n")), + context + .vertex_location + .as_ref() + .map(|value| format!("\nvertex_location: `{value}`\n")), + ]; + let information: String = lines.into_iter().flatten().collect(); + match &context.user_api_key_alias { + Some(alias) => format!( + "\n\nKey Name: `{alias}`\nTeam: `{}`{information}", + context.user_api_key_team_alias.as_deref().unwrap_or("None") + ), + None => information, + } +} + +#[cfg(test)] +mod testing { + use super::*; + + pub(super) const DEBUG: &str = "\nModel: ocr-model\nMessages: `None`"; + + pub(super) fn context(provider: &str, family: ExceptionFamily) -> ExceptionContext { + ExceptionContext { + model: "ocr-model".into(), + custom_llm_provider: Some(provider.into()), + family, + suppress_debug_info: true, + ..ExceptionContext::default() + } + } + + pub(super) fn http(status: u16, body: &str) -> OriginalException { + OriginalException::Http { + status, + body: body.into(), + headers: vec![("retry-after".into(), "7".into())], + } + } + + pub(super) fn upstream(status: u16, body: &str) -> Option { + Some(ResponseArg::Upstream(UpstreamResponse { + status, + body: body.into(), + headers: vec![("retry-after".into(), "7".into())], + })) + } + + pub(super) fn status(class: StatusClass, response: Option) -> PublicKind { + PublicKind::Status { + status_class: class, + response, + } + } + + /// The failure a rule builds before `exception_type` adds the response headers and the + /// banner flag. + pub(super) fn failure(kind: PublicKind, message: &str, provider: &str) -> PublicFailure { + PublicFailure { + kind, + message: message.into(), + model: "ocr-model".into(), + llm_provider: Some(provider.into()), + litellm_debug_info: None, + litellm_response_headers: None, + print_banner: false, + } + } + + pub(super) fn with_debug(failure: PublicFailure) -> PublicFailure { + PublicFailure { + litellm_debug_info: Some(DEBUG.into()), + ..failure + } + } + + /// What `exception_type` returns for an `http` original the rule mapped to `failure`. + pub(super) fn with_headers(failure: PublicFailure) -> PublicFailure { + PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..failure + } + } +} + +#[cfg(test)] +mod tests { + use super::testing::{DEBUG, context, failure, http, status, upstream, with_debug}; + use super::*; + + fn openai() -> ExceptionContext { + context("mistral", ExceptionFamily::OpenAiCompatible) + } + + #[test] + fn a_public_original_passes_through_without_banner_debug_or_prefix() { + let original = OriginalException::Public { + class: StatusClass::UnsupportedParams, + message: "Invalid `req_format`".into(), + }; + let context = ExceptionContext { + suppress_debug_info: false, + ..openai() + }; + assert_eq!( + exception_type(&context, &original), + failure( + status(StatusClass::UnsupportedParams, None), + "Invalid `req_format`", + "mistral" + ) + ); + } + + #[rstest::rstest] + #[case::vertex_family_status_rule(ExceptionFamily::VertexAi, "vertex_ai", PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "Vertex_aiException - rejected", "vertex_ai")) + })] + #[case::cohere_family(ExceptionFamily::Cohere, "cohere", PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "CohereException - rejected", "cohere")) + })] + #[case::other_family(ExceptionFamily::Other, "reducto", PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "ReductoException - rejected", "reducto")) + })] + fn families_without_a_409_rule_reach_the_status_table( + #[case] family: ExceptionFamily, + #[case] provider: &str, + #[case] expected: PublicFailure, + ) { + assert_eq!( + exception_type(&context(provider, family), &http(409, "rejected")), + expected + ); + } + + #[test] + fn the_openai_family_claims_a_409_before_the_status_table() { + assert_eq!( + exception_type(&openai(), &http(409, "rejected")), + PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure( + PublicKind::Api { + status: 409, + request_url: DOCS_URL + }, + "APIError: MistralException - rejected", + "mistral" + )) + } + ); + } + + #[rstest::rstest] + #[case::request_timeout_error("Request Timeout Error")] + #[case::request_timed_out("Request timed out")] + #[case::timed_out_generating("Timed out generating response")] + #[case::read_operation("The read operation timed out")] + fn timeout_markers_win_over_every_family(#[case] marker: &str) { + let body = format!("rate limit {marker}"); + for family in [ + ExceptionFamily::OpenAiCompatible, + ExceptionFamily::VertexAi, + ExceptionFamily::Cohere, + ExceptionFamily::Other, + ] { + assert_eq!( + exception_type(&context("mistral", family), &http(429, &body)), + PublicFailure { + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure( + PublicKind::Timeout { status: None }, + &format!("APITimeoutError - Request timed out. Error_str: {body}"), + "mistral" + )) + } + ); + } + } + + #[rstest::rstest] + #[case::provider_error_keeps_the_prefix(http(409, "rejected"), "ReductoException - rejected")] + #[case::synthesized_status_skips_the_status_table( + OriginalException::Connection { message: "refused".into() }, + "ReductoException - refused" + )] + #[case::plain_exception_keeps_its_text( + OriginalException::Local { class: LocalClass::FileNotFound, message: "File not found: /a".into() }, + "File not found: /a" + )] + fn unmapped_failures_are_connection_errors( + #[case] original: OriginalException, + #[case] message: &str, + ) { + let context = context("reducto", ExceptionFamily::Other); + let expected = failure(PublicKind::ApiConnection, message, "reducto"); + let actual = exception_type(&context, &original); + assert_eq!( + PublicFailure { + litellm_response_headers: None, + ..actual + }, + expected + ); + } + + #[test] + fn a_missing_provider_renders_like_python_none() { + let context = ExceptionContext { + custom_llm_provider: None, + family: ExceptionFamily::Other, + ..openai() + }; + assert_eq!( + exception_type(&context, &http(401, "rejected")), + PublicFailure { + llm_provider: None, + litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), + ..with_debug(failure( + status(StatusClass::Authentication, upstream(401, "rejected")), + "None - rejected", + "unused" + )) + } + ); + } + + #[rstest::rstest] + #[case::suppressed(true, false)] + #[case::printed(false, true)] + fn the_banner_prints_unless_debug_info_is_suppressed( + #[case] suppress_debug_info: bool, + #[case] print_banner: bool, + ) { + let context = ExceptionContext { + suppress_debug_info, + ..openai() + }; + assert_eq!( + exception_type(&context, &http(400, "rejected")).print_banner, + print_banner + ); + } + + #[test] + fn empty_upstream_headers_are_not_reported() { + let original = OriginalException::Http { + status: 400, + body: "rejected".into(), + headers: Vec::new(), + }; + assert_eq!( + exception_type(&openai(), &original).litellm_response_headers, + None + ); + } + + #[test] + fn messages_are_redacted_before_markers_and_prefixes() { + let body = "rejected Bearer abcdefghijklmnop"; + assert_eq!( + exception_type( + &context("reducto", ExceptionFamily::Other), + &http(400, body) + ) + .message, + "ReductoException - rejected REDACTED" + ); + } + + const SYNC_TIMEOUT: &str = "litellm.Timeout: Connection timed out after 0.5 seconds."; + + #[rstest::rstest] + #[case::sync(false, Some(0.5), Some(0.5031), SYNC_TIMEOUT)] + #[case::async_rounds_the_elapsed_time( + true, + Some(0.5), + Some(0.5031), + "litellm.Timeout: Connection timed out. Timeout passed=0.5, time taken=0.503 seconds" + )] + #[case::whole_seconds_keep_a_decimal( + true, + Some(600.0), + Some(2.0), + "litellm.Timeout: Connection timed out. Timeout passed=600.0, time taken=2.0 seconds" + )] + #[case::unknown_values_render_as_none( + true, + None, + None, + "litellm.Timeout: Connection timed out. Timeout passed=None, time taken=None seconds" + )] + fn timeout_text_follows_the_delivery_mode( + #[case] asynchronous: bool, + #[case] timeout_seconds: Option, + #[case] elapsed_seconds: Option, + #[case] expected: &str, + ) { + assert_eq!( + timeout_message(asynchronous, timeout_seconds, elapsed_seconds), + expected + ); + } + + #[rstest::rstest] + #[case::sync(false, SYNC_TIMEOUT)] + #[case::async_( + true, + "litellm.Timeout: Connection timed out. Timeout passed=0.5, time taken=0.503 seconds" + )] + fn a_timeout_is_a_408_carrying_the_handler_text( + #[case] asynchronous: bool, + #[case] text: &str, + ) { + let context = ExceptionContext { + asynchronous, + ..openai() + }; + let original = OriginalException::Timeout { + timeout_seconds: Some(0.5), + elapsed_seconds: Some(0.5031), + }; + assert_eq!( + exception_type(&context, &original), + with_debug(failure( + PublicKind::Timeout { status: None }, + &format!("Timeout Error: MistralException - {text}"), + "mistral" + )) + ); + } + + #[test] + fn a_refused_connection_is_a_500_with_an_empty_response() { + assert_eq!( + exception_type( + &openai(), + &OriginalException::Connection { + message: "refused".into() + } + ), + with_debug(failure( + status( + StatusClass::InternalServer, + Some(ResponseArg::Upstream(UpstreamResponse { + status: 500, + body: String::new(), + headers: Vec::new(), + })) + ), + "InternalServerError: MistralException - refused", + "mistral" + )) + ); + } + + #[test] + fn debug_information_follows_the_python_layout() { + let context = ExceptionContext { + vertex_project: Some("project".into()), + vertex_location: Some("region".into()), + model_group: Some("ocr".into()), + deployment: Some("deployment".into()), + user_api_key_alias: Some("key".into()), + ..openai() + }; + assert_eq!( + extra_information(&context, api_base(&context).as_deref()), + concat!( + "\n\nKey Name: `key`\nTeam: `None`", + "\nModel: ocr-model", + "\nAPI Base: `region-aiplatform.googleapis.com/v1/projects/project/locations/region/publishers/google/models/ocr-model:generateContent`", + "\nMessages: `None`", + "\nmodel_group: `ocr`\n", + "\ndeployment: `deployment`\n", + "\nvertex_project: `project`\n", + "\nvertex_location: `region`\n", + ) + ); + } + + #[rstest::rstest] + #[case::bare(ExceptionContext::default(), "\nModel: ")] + #[case::redacted_messages(ExceptionContext { redact_messages_in_exceptions: true, model: "m".into(), ..ExceptionContext::default() }, "\nModel: m")] + #[case::messages(ExceptionContext { model: "m".into(), ..ExceptionContext::default() }, "\nModel: m\nMessages: `None`")] + #[case::team_alias( + ExceptionContext { model: "m".into(), user_api_key_alias: Some("key".into()), user_api_key_team_alias: Some("team".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + "\n\nKey Name: `key`\nTeam: `team`\nModel: m" + )] + #[case::team_alias_without_key_is_ignored( + ExceptionContext { model: "m".into(), user_api_key_team_alias: Some("team".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + "\nModel: m" + )] + #[case::project_without_location_has_no_api_base( + ExceptionContext { model: "m".into(), vertex_project: Some("p".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + "\nModel: m\nvertex_project: `p`\n" + )] + #[case::location_without_project_has_no_api_base( + ExceptionContext { model: "m".into(), vertex_location: Some("l".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + "\nModel: m\nvertex_location: `l`\n" + )] + fn each_optional_context_field_adds_its_own_line( + #[case] context: ExceptionContext, + #[case] expected: &str, + ) { + assert_eq!( + extra_information(&context, api_base(&context).as_deref()), + expected + ); + } + + #[rstest::rstest] + #[case::lowercase("mistral", "MistralException")] + #[case::keeps_the_rest("azure_ai", "Azure_aiException")] + #[case::empty("", "")] + fn exception_provider_capitalizes_only_the_first_letter( + #[case] provider: &str, + #[case] expected: &str, + ) { + assert_eq!(exception_provider(provider), expected); + } + + #[rstest::rstest] + #[case::lowers_the_rest("vERTEX_AI", "Vertex_ai")] + #[case::empty("", "")] + fn python_capitalize_lowers_the_rest(#[case] value: &str, #[case] expected: &str) { + assert_eq!(python_capitalize(value), expected); + } + + #[test] + fn debug_constant_matches_the_default_test_context() { + let context = openai(); + assert_eq!(extra_information(&context, None), DEBUG); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs new file mode 100644 index 00000000000..ffbb172582b --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs @@ -0,0 +1,586 @@ +use super::public::{PublicFailure, StatusClass}; +use super::rules::{ + ApiStatus, Kind, ResponseChoice, Rule, apply, contains_any, is_context_window_exceeded, + is_rate_limit, +}; +use super::{DOCS_URL, Mapping}; + +const OPENAI_URL: &str = "https://api.openai.com/v1"; + +const ENCRYPTED_CONTENT_HELP: &str = "\n\n This error occurs when load balancing Responses API across deployments with different API keys.\n Encrypted content is tied to the organization that created it and cannot be decrypted by other organizations.\n\n Solution: Enable 'encrypted_content_affinity' to route follow-up requests to the correct deployment:\n\n router_settings:\n enable_pre_call_checks: true\n optional_pre_call_checks:\n - encrypted_content_affinity\n\n Learn more: https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing"; + +const fn with_response(class: StatusClass) -> Kind { + Kind::Status { + class, + response: ResponseChoice::Provider, + } +} + +fn exception_provider(mapping: &Mapping<'_>) -> String { + if mapping.provider == "openai" { + "OpenAIException".to_string() + } else { + super::exception_provider(mapping.provider) + } +} + +/// The raw message with OpenAI's own names swapped for the provider's. +fn message(mapping: &Mapping<'_>) -> String { + let provider = mapping.provider; + mapping + .original + .message + .replace("OPENAI", &provider.to_uppercase()) + .replace("openai.OpenAIError", &format!("{provider}.{provider}Error")) +} + +fn prefixed(mapping: &Mapping<'_>, label: &str) -> String { + format!( + "{label}{} - {}", + exception_provider(mapping), + message(mapping) + ) +} + +fn status_is(mapping: &Mapping<'_>, statuses: &[u16]) -> bool { + mapping + .original + .status + .is_some_and(|status| statuses.contains(&status)) +} + +/// `_map_openai_exception`, in its branch order. +const RULES: &[Rule] = &[ + Rule { + when: |mapping| is_rate_limit(&mapping.error_str, mapping.original.status), + kind: with_response(StatusClass::RateLimit), + message: |mapping| prefixed(mapping, "RateLimitError: "), + debug: false, + }, + Rule { + when: |mapping| is_context_window_exceeded(&mapping.error_str), + kind: with_response(StatusClass::ContextWindowExceeded), + message: |mapping| prefixed(mapping, "ContextWindowExceededError: "), + debug: true, + }, + Rule { + when: |mapping| { + mapping.error_str.contains("invalid_request_error") + && mapping.error_str.contains("model_not_found") + }, + kind: with_response(StatusClass::NotFound), + message: |mapping| prefixed(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| mapping.error_str.contains("A timeout occurred"), + kind: Kind::Timeout(None), + message: |mapping| prefixed(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| { + let error_str = &mapping.error_str; + (error_str.contains("invalid_request_error") + && error_str.contains("content_policy_violation")) + || (error_str.contains("Invalid prompt") + && error_str.contains("violating our usage policy")) + || error_str + .to_lowercase() + .contains("request was rejected as a result of the safety system") + }, + kind: with_response(StatusClass::ContentPolicyViolation), + message: |mapping| prefixed(mapping, "ContentPolicyViolationError: "), + debug: true, + }, + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &["invalid_encrypted_content", "could not be verified"], + ) + }, + kind: with_response(StatusClass::BadRequest), + message: |mapping| { + format!( + "{} - {}{ENCRYPTED_CONTENT_HELP}", + exception_provider(mapping), + message(mapping) + ) + }, + debug: true, + }, + Rule { + when: |mapping| { + mapping.error_str.contains("invalid_request_error") + && !mapping.error_str.contains("Incorrect API key provided") + }, + kind: with_response(StatusClass::BadRequest), + message: |mapping| prefixed(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &[ + "Web server is returning an unknown error", + "The server had an error processing your request.", + ], + ) + }, + kind: Kind::Status { + class: StatusClass::InternalServer, + response: ResponseChoice::Omitted, + }, + message: |mapping| prefixed(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| mapping.error_str.contains("Request too large"), + kind: with_response(StatusClass::RateLimit), + message: |mapping| prefixed(mapping, "RateLimitError: "), + debug: true, + }, + Rule { + when: |mapping| { + mapping.error_str.contains("The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable") + }, + kind: with_response(StatusClass::Authentication), + message: |mapping| prefixed(mapping, "AuthenticationError: "), + debug: true, + }, + Rule { + when: |mapping| { + mapping + .error_str + .contains("Mistral API raised a streaming error") + }, + kind: Kind::Api { + status: ApiStatus::Fixed(500), + request_url: OPENAI_URL, + }, + message: |mapping| prefixed(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| mapping.original.status.is_none(), + kind: Kind::ApiConnection, + message: |mapping| prefixed(mapping, "APIConnectionError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[400, 422]), + kind: with_response(StatusClass::BadRequest), + message: |mapping| prefixed(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[401]), + kind: with_response(StatusClass::Authentication), + message: |mapping| prefixed(mapping, "AuthenticationError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[404]), + kind: with_response(StatusClass::NotFound), + message: |mapping| prefixed(mapping, "NotFoundError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[408]), + kind: Kind::Timeout(None), + message: |mapping| prefixed(mapping, "Timeout Error: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[429]), + kind: with_response(StatusClass::RateLimit), + message: |mapping| prefixed(mapping, "RateLimitError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[500]), + kind: with_response(StatusClass::InternalServer), + message: |mapping| prefixed(mapping, "InternalServerError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[502]), + kind: with_response(StatusClass::BadGateway), + message: |mapping| prefixed(mapping, "BadGatewayError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[503]), + kind: with_response(StatusClass::ServiceUnavailable), + message: |mapping| prefixed(mapping, "ServiceUnavailableError: "), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, &[504]), + kind: Kind::Timeout(Some(504)), + message: |mapping| prefixed(mapping, "Timeout Error: "), + debug: true, + }, + Rule { + when: |_| true, + kind: Kind::Api { + status: ApiStatus::Original, + request_url: DOCS_URL, + }, + message: |mapping| prefixed(mapping, "APIError: "), + debug: true, + }, +]; + +pub(super) fn map(mapping: &Mapping<'_>) -> Option { + apply(RULES, mapping) +} + +#[cfg(test)] +mod tests { + use super::super::testing::{context, failure, http, status, upstream, with_debug}; + use super::super::{ExceptionFamily, OriginalException, PublicKind}; + use super::*; + + fn mapped(provider: &str, original: &OriginalException) -> PublicFailure { + let context = context(provider, ExceptionFamily::OpenAiCompatible); + map(&Mapping::new(&context, original)).expect("the OpenAI table ends in a catch-all") + } + + fn kind(class: StatusClass, status_code: u16, body: &str) -> PublicKind { + status(class, upstream(status_code, body)) + } + + #[rstest::rstest] + #[case::rate_limit_phrase( + 400, + "rate limit reached", + failure( + kind(StatusClass::RateLimit, 400, "rate limit reached"), + "RateLimitError: MistralException - rate limit reached", + "mistral", + ) + )] + #[case::context_window( + 500, + "This model's maximum context length is 10", + with_debug(failure( + kind( + StatusClass::ContextWindowExceeded, + 500, + "This model's maximum context length is 10" + ), + "ContextWindowExceededError: MistralException - This model's maximum context length is 10", + "mistral", + )) + )] + #[case::model_not_found( + 400, + "invalid_request_error model_not_found", + with_debug(failure( + kind(StatusClass::NotFound, 400, "invalid_request_error model_not_found"), + "MistralException - invalid_request_error model_not_found", + "mistral", + )) + )] + #[case::timeout_occurred(400, "A timeout occurred", with_debug(failure( + PublicKind::Timeout { status: None }, + "MistralException - A timeout occurred", + "mistral", + )))] + #[case::content_policy_error_code( + 400, + "invalid_request_error content_policy_violation", + with_debug(failure( + kind( + StatusClass::ContentPolicyViolation, + 400, + "invalid_request_error content_policy_violation" + ), + "ContentPolicyViolationError: MistralException - invalid_request_error content_policy_violation", + "mistral", + )) + )] + #[case::content_policy_usage_policy( + 400, + "Invalid prompt violating our usage policy", + with_debug(failure( + kind( + StatusClass::ContentPolicyViolation, + 400, + "Invalid prompt violating our usage policy" + ), + "ContentPolicyViolationError: MistralException - Invalid prompt violating our usage policy", + "mistral", + )) + )] + #[case::content_policy_safety_system( + 400, + "Request was rejected as a result of the safety system", + with_debug(failure( + kind( + StatusClass::ContentPolicyViolation, + 400, + "Request was rejected as a result of the safety system" + ), + "ContentPolicyViolationError: MistralException - Request was rejected as a result of the safety system", + "mistral", + )) + )] + #[case::encrypted_content(400, "invalid_encrypted_content", with_debug(failure( + kind(StatusClass::BadRequest, 400, "invalid_encrypted_content"), + &format!("MistralException - invalid_encrypted_content{ENCRYPTED_CONTENT_HELP}"), + "mistral", + )))] + #[case::unverifiable_content(400, "could not be verified", with_debug(failure( + kind(StatusClass::BadRequest, 400, "could not be verified"), + &format!("MistralException - could not be verified{ENCRYPTED_CONTENT_HELP}"), + "mistral", + )))] + #[case::invalid_request( + 429, + "invalid_request_error bad field", + with_debug(failure( + kind(StatusClass::BadRequest, 429, "invalid_request_error bad field"), + "MistralException - invalid_request_error bad field", + "mistral", + )) + )] + #[case::unknown_server_error( + 400, + "Web server is returning an unknown error", + failure( + status(StatusClass::InternalServer, None), + "MistralException - Web server is returning an unknown error", + "mistral", + ) + )] + #[case::server_had_an_error( + 400, + "The server had an error processing your request.", + failure( + status(StatusClass::InternalServer, None), + "MistralException - The server had an error processing your request.", + "mistral", + ) + )] + #[case::request_too_large( + 400, + "Request too large", + with_debug(failure( + kind(StatusClass::RateLimit, 400, "Request too large"), + "RateLimitError: MistralException - Request too large", + "mistral", + )) + )] + #[case::missing_client_api_key( + 400, + "The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable", + with_debug(failure( + kind( + StatusClass::Authentication, + 400, + "The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable", + ), + "AuthenticationError: MistralException - The api_key client option must be set either by passing api_key to the client or by setting the MISTRAL_API_KEY environment variable", + "mistral", + )) + )] + #[case::mistral_streaming_error(400, "Mistral API raised a streaming error", with_debug(failure( + PublicKind::Api { status: 500, request_url: OPENAI_URL }, + "MistralException - Mistral API raised a streaming error", + "mistral", + )))] + fn each_text_rule_maps_by_the_body( + #[case] status_code: u16, + #[case] body: &str, + #[case] expected: PublicFailure, + ) { + assert_eq!(mapped("mistral", &http(status_code, body)), expected); + } + + #[rstest::rstest] + #[case::bad_request( + 400, + kind(StatusClass::BadRequest, 400, "rejected"), + "MistralException - rejected" + )] + #[case::unprocessable( + 422, + kind(StatusClass::BadRequest, 422, "rejected"), + "MistralException - rejected" + )] + #[case::authentication( + 401, + kind(StatusClass::Authentication, 401, "rejected"), + "AuthenticationError: MistralException - rejected" + )] + #[case::not_found( + 404, + kind(StatusClass::NotFound, 404, "rejected"), + "NotFoundError: MistralException - rejected" + )] + #[case::request_timeout(408, PublicKind::Timeout { status: None }, "Timeout Error: MistralException - rejected")] + #[case::rate_limited( + 429, + kind(StatusClass::RateLimit, 429, "rejected"), + "RateLimitError: MistralException - rejected" + )] + #[case::internal_server( + 500, + kind(StatusClass::InternalServer, 500, "rejected"), + "InternalServerError: MistralException - rejected" + )] + #[case::bad_gateway( + 502, + kind(StatusClass::BadGateway, 502, "rejected"), + "BadGatewayError: MistralException - rejected" + )] + #[case::service_unavailable( + 503, + kind(StatusClass::ServiceUnavailable, 503, "rejected"), + "ServiceUnavailableError: MistralException - rejected" + )] + #[case::gateway_timeout(504, PublicKind::Timeout { status: Some(504) }, "Timeout Error: MistralException - rejected")] + #[case::any_other_status(409, PublicKind::Api { status: 409, request_url: DOCS_URL }, "APIError: MistralException - rejected")] + fn each_status_rule_maps_by_the_status( + #[case] status_code: u16, + #[case] kind: PublicKind, + #[case] message: &str, + ) { + assert_eq!( + mapped("mistral", &http(status_code, "rejected")), + with_debug(failure(kind, message, "mistral")) + ); + } + + #[test] + fn a_failure_without_a_status_is_a_connection_error() { + let original = OriginalException::Response { + message: "invalid OCR response field: pages".into(), + }; + assert_eq!( + mapped("mistral", &original), + with_debug(failure( + PublicKind::ApiConnection, + "APIConnectionError: MistralException - invalid OCR response field: pages", + "mistral" + )) + ); + } + + #[rstest::rstest] + #[case::rate_limit_before_context_window( + 400, + "rate limit and This model's maximum context length is 10", + kind( + StatusClass::RateLimit, + 400, + "rate limit and This model's maximum context length is 10" + ), + "RateLimitError: MistralException - rate limit and This model's maximum context length is 10", + false + )] + #[case::context_window_before_content_policy( + 400, + "This model's maximum context length is 10 invalid_request_error content_policy_violation", + kind( + StatusClass::ContextWindowExceeded, + 400, + "This model's maximum context length is 10 invalid_request_error content_policy_violation" + ), + "ContextWindowExceededError: MistralException - This model's maximum context length is 10 invalid_request_error content_policy_violation", + true + )] + #[case::model_not_found_before_invalid_request( + 400, + "invalid_request_error model_not_found", + kind(StatusClass::NotFound, 400, "invalid_request_error model_not_found"), + "MistralException - invalid_request_error model_not_found", + true + )] + #[case::timeout_before_invalid_request( + 400, + "A timeout occurred invalid_request_error", + PublicKind::Timeout { status: None }, + "MistralException - A timeout occurred invalid_request_error", + true + )] + #[case::content_policy_before_invalid_request( + 400, + "invalid_request_error content_policy_violation", + kind( + StatusClass::ContentPolicyViolation, + 400, + "invalid_request_error content_policy_violation" + ), + "ContentPolicyViolationError: MistralException - invalid_request_error content_policy_violation", + true + )] + #[case::invalid_request_with_a_bad_key_falls_to_the_status( + 401, + "invalid_request_error Incorrect API key provided", + kind( + StatusClass::Authentication, + 401, + "invalid_request_error Incorrect API key provided" + ), + "AuthenticationError: MistralException - invalid_request_error Incorrect API key provided", + true + )] + #[case::text_rules_before_status( + 429, + "invalid_request_error bad field", + kind(StatusClass::BadRequest, 429, "invalid_request_error bad field"), + "MistralException - invalid_request_error bad field", + true + )] + #[case::echoed_429_is_not_a_rate_limit( + 400, + "token 429 in the prompt", + kind(StatusClass::BadRequest, 400, "token 429 in the prompt"), + "MistralException - token 429 in the prompt", + true + )] + fn the_earlier_rule_wins_when_two_apply( + #[case] status_code: u16, + #[case] body: &str, + #[case] kind: PublicKind, + #[case] message: &str, + #[case] debug: bool, + ) { + let expected = failure(kind, message, "mistral"); + assert_eq!( + mapped("mistral", &http(status_code, body)), + if debug { + with_debug(expected) + } else { + expected + } + ); + } + + #[rstest::rstest] + #[case::provider_names_replace_openai( + "azure_ai", + "OPENAI said openai.OpenAIError", + "Azure_aiException - AZURE_AI said azure_ai.azure_aiError" + )] + #[case::openai_keeps_its_own_name("openai", "rejected", "OpenAIException - rejected")] + fn the_message_names_the_provider( + #[case] provider: &str, + #[case] body: &str, + #[case] message: &str, + ) { + assert_eq!( + mapped(provider, &http(400, body)), + with_debug(failure( + kind(StatusClass::BadRequest, 400, body), + message, + provider + )) + ); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs new file mode 100644 index 00000000000..d64868c1606 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs @@ -0,0 +1,49 @@ +use super::public::StatusClass; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum LocalClass { + ValueError, + FileNotFound, + OsError, +} + +/// A route failure in the shape Python's `exception_type` receives it, before any public +/// class is chosen. +#[derive(Clone, Debug, PartialEq)] +pub enum OriginalException { + Http { + status: u16, + body: String, + headers: Vec<(String, String)>, + }, + Connection { + message: String, + }, + Timeout { + timeout_seconds: Option, + elapsed_seconds: Option, + }, + Response { + message: String, + }, + Local { + class: LocalClass, + message: String, + }, + /// A failure Python raises as a public LiteLLM exception itself, which `exception_type` + /// hands back unchanged. + Public { + class: StatusClass, + message: String, + }, +} + +/// Which of the provider-specific mappers in `exception_type` a route's provider uses. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum ExceptionFamily { + OpenAiCompatible, + VertexAi, + Cohere, + #[default] + Other, +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs new file mode 100644 index 00000000000..a7319c1287b --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs @@ -0,0 +1,262 @@ +use serde::Serialize; + +/// The public LiteLLM classes built from a status code alone: every one takes the same +/// constructor arguments. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, strum::EnumIter, strum::IntoStaticStr)] +#[serde(rename_all = "snake_case")] +#[strum(serialize_all = "snake_case")] +pub enum StatusClass { + BadRequest, + Authentication, + PermissionDenied, + NotFound, + RateLimit, + ContextWindowExceeded, + ContentPolicyViolation, + InternalServer, + BadGateway, + ServiceUnavailable, + UnsupportedParams, +} + +impl StatusClass { + /// The `status_code` the Python class sets on itself. + pub const fn status_code(self) -> u16 { + match self { + Self::BadRequest + | Self::ContextWindowExceeded + | Self::ContentPolicyViolation + | Self::UnsupportedParams => 400, + Self::Authentication => 401, + Self::PermissionDenied => 403, + Self::NotFound => 404, + Self::RateLimit => 429, + Self::InternalServer => 500, + Self::BadGateway => 502, + Self::ServiceUnavailable => 503, + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub struct UpstreamResponse { + pub status: u16, + pub body: String, + pub headers: Vec<(String, String)>, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub struct HttpStub { + pub status: u16, + pub method: &'static str, + pub url: &'static str, + pub content: Option, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ResponseArg { + Upstream(UpstreamResponse), + Stub(HttpStub), +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum PublicKind { + Status { + status_class: StatusClass, + response: Option, + }, + Timeout { + status: Option, + }, + ApiConnection, + Api { + status: u16, + request_url: &'static str, + }, +} + +/// Constructor arguments for the public LiteLLM exception, as `exception_type` passes them. +#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +pub struct PublicFailure { + pub kind: PublicKind, + pub message: String, + pub model: String, + pub llm_provider: Option, + pub litellm_debug_info: Option, + pub litellm_response_headers: Option>, + pub print_banner: bool, +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeSet; + use std::path::PathBuf; + + use serde_json::Value; + use strum::IntoEnumIterator; + + use super::*; + + const REGENERATE: &str = "LITELLM_REGENERATE_PUBLIC_FAILURE_FIXTURES"; + + fn fixture_directory() -> PathBuf { + PathBuf::from(env!("CARGO_MANIFEST_DIR")) + .join("../../../tests/test_litellm/rust_bridge/fixtures/public_failures") + } + + fn upstream(status: u16) -> ResponseArg { + ResponseArg::Upstream(UpstreamResponse { + status, + body: r#"{"message": "rejected"}"#.into(), + headers: vec![("retry-after".into(), "7".into())], + }) + } + + fn status_response(class: StatusClass) -> Option { + match class { + StatusClass::Authentication => None, + StatusClass::PermissionDenied => Some(ResponseArg::Stub(HttpStub { + status: 403, + method: "POST", + url: " https://cloud.google.com/vertex-ai/", + content: None, + })), + StatusClass::InternalServer => Some(ResponseArg::Stub(HttpStub { + status: 500, + method: "completion", + url: "https://github.com/BerriAI/litellm", + content: Some("upstream text".into()), + })), + class => Some(upstream(class.status_code())), + } + } + + fn failure(kind: PublicKind, name: &str) -> PublicFailure { + let headers = matches!( + &kind, + PublicKind::Status { + response: Some(ResponseArg::Upstream(_)), + .. + } + ); + PublicFailure { + kind, + message: format!("MistralException - {name}"), + model: "ocr-model".into(), + llm_provider: Some("mistral".into()), + litellm_debug_info: Some("\nModel: ocr-model".into()), + litellm_response_headers: headers.then(|| vec![("retry-after".into(), "7".into())]), + print_banner: false, + } + } + + /// One payload per public class the constructor can build; `test_failures.py` reads the + /// same files, so a shape change on either side fails there or here. + fn fixtures() -> Vec<(String, PublicFailure)> { + let statuses = StatusClass::iter().map(|class| { + let name: &'static str = class.into(); + let name = format!("status_{name}"); + let built = failure( + PublicKind::Status { + status_class: class, + response: status_response(class), + }, + &name, + ); + (name, built) + }); + let others = [ + ( + "timeout_with_status", + PublicFailure { + print_banner: true, + ..failure( + PublicKind::Timeout { status: Some(504) }, + "timeout_with_status", + ) + }, + ), + ( + "timeout_without_status", + PublicFailure { + litellm_debug_info: None, + ..failure( + PublicKind::Timeout { status: None }, + "timeout_without_status", + ) + }, + ), + ( + "api_connection", + PublicFailure { + llm_provider: None, + ..failure(PublicKind::ApiConnection, "api_connection") + }, + ), + ( + "api", + failure( + PublicKind::Api { + status: 409, + request_url: "https://docs.litellm.ai/docs", + }, + "api", + ), + ), + ] + .map(|(name, built)| (name.to_string(), built)); + statuses.chain(others).collect() + } + + #[test] + fn serialized_payloads_match_the_golden_fixtures_python_reads() { + let directory = fixture_directory(); + let regenerate = std::env::var_os(REGENERATE).is_some(); + let expected = fixtures(); + for (name, built) in &expected { + let path = directory.join(format!("{name}.json")); + let serialized = serde_json::to_value(built).unwrap(); + if regenerate { + std::fs::create_dir_all(&directory).unwrap(); + std::fs::write( + &path, + format!("{}\n", serde_json::to_string_pretty(&serialized).unwrap()), + ) + .unwrap(); + } + let golden: Value = + serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap(); + assert_eq!(serialized, golden, "{name}; set {REGENERATE}=1 to rewrite"); + } + let on_disk: BTreeSet = std::fs::read_dir(&directory) + .unwrap() + .map(|entry| entry.unwrap().file_name().to_string_lossy().into_owned()) + .collect(); + let generated: BTreeSet = expected + .iter() + .map(|(name, _)| format!("{name}.json")) + .collect(); + assert_eq!(on_disk, generated); + } + + #[rstest::rstest] + #[case(StatusClass::BadRequest, 400)] + #[case(StatusClass::Authentication, 401)] + #[case(StatusClass::PermissionDenied, 403)] + #[case(StatusClass::NotFound, 404)] + #[case(StatusClass::RateLimit, 429)] + #[case(StatusClass::ContextWindowExceeded, 400)] + #[case(StatusClass::ContentPolicyViolation, 400)] + #[case(StatusClass::InternalServer, 500)] + #[case(StatusClass::BadGateway, 502)] + #[case(StatusClass::ServiceUnavailable, 503)] + #[case(StatusClass::UnsupportedParams, 400)] + fn status_codes_are_the_ones_the_python_classes_set( + #[case] class: StatusClass, + #[case] status: u16, + ) { + assert_eq!(class.status_code(), status); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs new file mode 100644 index 00000000000..5df8ad991b5 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs @@ -0,0 +1,403 @@ +use std::sync::LazyLock; + +use fancy_regex::Regex; +use serde_json::Value; + +use super::Mapping; +use super::public::{HttpStub, PublicFailure, PublicKind, ResponseArg, StatusClass}; + +const GITHUB_URL: &str = "https://github.com/BerriAI/litellm"; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum ResponseChoice { + Omitted, + Provider, + Stub { status: u16, url: &'static str }, + InternalServerStub, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum ApiStatus { + Fixed(u16), + Original, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(super) enum Kind { + Status { + class: StatusClass, + response: ResponseChoice, + }, + Timeout(Option), + ApiConnection, + Api { + status: ApiStatus, + request_url: &'static str, + }, +} + +/// One branch of a Python `_map_*_exception` function: when it applies, the class it +/// raises, the message it builds, and whether it passes `litellm_debug_info`. +pub(super) struct Rule { + pub(super) when: fn(&Mapping<'_>) -> bool, + pub(super) kind: Kind, + pub(super) message: fn(&Mapping<'_>) -> String, + pub(super) debug: bool, +} + +/// The first rule that applies decides the failure, as the `if`/`elif` chain does in Python. +pub(super) fn apply(rules: &[Rule], mapping: &Mapping<'_>) -> Option { + rules + .iter() + .find(|rule| (rule.when)(mapping)) + .map(|rule| rule.build(mapping)) +} + +impl Rule { + fn build(&self, mapping: &Mapping<'_>) -> PublicFailure { + let kind = match self.kind { + Kind::Status { class, response } => PublicKind::Status { + status_class: class, + response: response.resolve(mapping), + }, + Kind::Timeout(status) => PublicKind::Timeout { status }, + Kind::ApiConnection => PublicKind::ApiConnection, + Kind::Api { + status, + request_url, + } => PublicKind::Api { + status: match status { + ApiStatus::Fixed(status) => status, + ApiStatus::Original => mapping.original.status.unwrap_or(500), + }, + request_url, + }, + }; + PublicFailure { + kind, + message: (self.message)(mapping), + model: mapping.context.model.clone(), + llm_provider: mapping.context.custom_llm_provider.clone(), + litellm_debug_info: self.debug.then(|| mapping.extra_information.clone()), + litellm_response_headers: None, + print_banner: false, + } + } +} + +impl ResponseChoice { + fn resolve(self, mapping: &Mapping<'_>) -> Option { + match self { + Self::Omitted => None, + Self::Provider => mapping.original.response.clone().map(ResponseArg::Upstream), + Self::Stub { status, url } => Some(ResponseArg::Stub(HttpStub { + status, + method: "POST", + url, + content: None, + })), + Self::InternalServerStub => Some(ResponseArg::Stub(HttpStub { + status: 500, + method: "completion", + url: GITHUB_URL, + content: Some(mapping.original.message.clone()), + })), + } + } +} + +pub(super) fn contains_any(text: &str, markers: &[&str]) -> bool { + markers.iter().any(|marker| text.contains(marker)) +} + +static STANDALONE_429: LazyLock = + LazyLock::new(|| Regex::new(r"\b429\b").expect("valid regex")); +static RATE_LIMIT_PHRASE: LazyLock = + LazyLock::new(|| Regex::new(r"rate[\s_\-]*limit").expect("valid regex")); + +/// `ExceptionCheckers.is_error_str_rate_limit`. +pub(super) fn is_rate_limit(error_str: &str, status: Option) -> bool { + if STANDALONE_429.is_match(error_str).unwrap_or(false) && status == Some(429) { + return true; + } + let lower = error_str.to_lowercase(); + RATE_LIMIT_PHRASE.is_match(&lower).unwrap_or(false) + || lower.contains("service tier capacity exceeded") +} + +/// `ExceptionCheckers.is_error_str_context_window_exceeded`. +pub(super) fn is_context_window_exceeded(error_str: &str) -> bool { + let lower = error_str.to_lowercase(); + if lower.contains("string_above_max_length") { + return false; + } + if lower.contains("invalid 'user'") && lower.contains("string too long") { + return false; + } + contains_any( + &lower, + &[ + "exceed context limit", + "this model's maximum context length is", + "string too long. expected a string with maximum length", + "model's maximum context limit", + "is longer than the model's context length", + "input tokens exceed the configured limit", + "`inputs` tokens + `max_new_tokens` must be", + "exceeds the available context size", + "exceeds the maximum number of tokens allowed", + ], + ) || (lower.contains("current length is") && lower.contains("while limit is")) + || (lower.contains("maximum input length is") && lower.contains("tokens")) +} + +/// The integer `error.code` of a JSON error body, read the way Python's `int()` would. +pub(super) fn body_error_code(error_str: &str) -> Option { + let body: Value = serde_json::from_str(error_str).ok()?; + let Some(Value::Object(error)) = body.as_object()?.get("error") else { + return None; + }; + match error.get("code")? { + Value::Number(number) => number + .as_i64() + .or_else(|| number.as_f64().map(|value| value.trunc() as i64)), + Value::String(code) => code.trim().replace('_', "").parse().ok(), + Value::Bool(flag) => Some(i64::from(*flag)), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::super::testing::{context, failure, http}; + use super::super::{ExceptionFamily, OriginalException, UpstreamResponse}; + use super::*; + + fn first_marker(mapping: &Mapping<'_>) -> bool { + mapping.error_str.contains("first") + } + + fn always(_: &Mapping<'_>) -> bool { + true + } + + fn text(mapping: &Mapping<'_>) -> String { + format!("seen {}", mapping.error_str) + } + + const ORDERED: &[Rule] = &[ + Rule { + when: first_marker, + kind: Kind::Status { + class: StatusClass::NotFound, + response: ResponseChoice::Omitted, + }, + message: text, + debug: false, + }, + Rule { + when: always, + kind: Kind::ApiConnection, + message: text, + debug: true, + }, + ]; + + fn apply_one(kind: Kind, debug: bool, original: &OriginalException) -> Option { + let context = context("mistral", ExceptionFamily::OpenAiCompatible); + let mapping = Mapping::new(&context, original); + apply( + &[Rule { + when: always, + kind, + message: text, + debug, + }], + &mapping, + ) + } + + #[rstest::rstest] + #[case::earlier_rule_wins("first and second", failure( + PublicKind::Status { status_class: StatusClass::NotFound, response: None }, + "seen first and second", + "mistral", + ))] + #[case::later_rule_when_the_earlier_does_not_apply("second", PublicFailure { + litellm_debug_info: Some("\nModel: ocr-model\nMessages: `None`".into()), + ..failure(PublicKind::ApiConnection, "seen second", "mistral") + })] + fn the_first_applicable_rule_decides(#[case] body: &str, #[case] expected: PublicFailure) { + let context = context("mistral", ExceptionFamily::OpenAiCompatible); + let original = http(400, body); + assert_eq!( + apply(ORDERED, &Mapping::new(&context, &original)), + Some(expected) + ); + } + + #[test] + fn no_applicable_rule_leaves_the_failure_to_the_caller() { + let context = context("mistral", ExceptionFamily::OpenAiCompatible); + let original = http(400, "second"); + assert_eq!( + apply(&ORDERED[..1], &Mapping::new(&context, &original)), + None + ); + } + + #[rstest::rstest] + #[case::omitted(ResponseChoice::Omitted, None)] + #[case::provider(ResponseChoice::Provider, Some(ResponseArg::Upstream(UpstreamResponse { + status: 400, + body: "body".into(), + headers: vec![("retry-after".into(), "7".into())], + })))] + #[case::stub( + ResponseChoice::Stub { status: 429, url: "https://stub.test" }, + Some(ResponseArg::Stub(HttpStub { status: 429, method: "POST", url: "https://stub.test", content: None })) + )] + #[case::internal_server_stub( + ResponseChoice::InternalServerStub, + Some(ResponseArg::Stub(HttpStub { + status: 500, + method: "completion", + url: GITHUB_URL, + content: Some("body".into()), + })) + )] + fn response_choices_resolve_against_the_original( + #[case] response: ResponseChoice, + #[case] expected: Option, + ) { + let built = apply_one( + Kind::Status { + class: StatusClass::BadRequest, + response, + }, + false, + &http(400, "body"), + ) + .unwrap(); + assert_eq!( + built.kind, + PublicKind::Status { + status_class: StatusClass::BadRequest, + response: expected, + } + ); + } + + #[rstest::rstest] + #[case::fixed(ApiStatus::Fixed(500), http(409, "body"), 500)] + #[case::original(ApiStatus::Original, http(409, "body"), 409)] + #[case::original_without_a_status( + ApiStatus::Original, + OriginalException::Response { message: "body".into() }, + 500 + )] + fn api_status_is_fixed_or_the_originals( + #[case] status: ApiStatus, + #[case] original: OriginalException, + #[case] expected: u16, + ) { + let built = apply_one( + Kind::Api { + status, + request_url: "https://api.test", + }, + false, + &original, + ) + .unwrap(); + assert_eq!( + built, + failure( + PublicKind::Api { + status: expected, + request_url: "https://api.test" + }, + "seen body", + "mistral" + ) + ); + } + + #[rstest::rstest] + #[case::with_debug(true, Some("\nModel: ocr-model\nMessages: `None`"))] + #[case::without_debug(false, None)] + fn debug_rules_carry_the_extra_information( + #[case] debug: bool, + #[case] expected: Option<&str>, + ) { + let built = apply_one(Kind::Timeout(Some(504)), debug, &http(504, "body")).unwrap(); + assert_eq!( + built, + PublicFailure { + litellm_debug_info: expected.map(str::to_string), + ..failure( + PublicKind::Timeout { status: Some(504) }, + "seen body", + "mistral" + ) + } + ); + } + + #[rstest::rstest] + #[case::standalone_429_with_429_status("got 429 back", Some(429), true)] + #[case::standalone_429_with_other_status("got 429 back", Some(400), false)] + #[case::embedded_429("token4290", Some(429), false)] + #[case::phrase_spaced("Rate Limit reached", None, true)] + #[case::phrase_underscored("rate_limit", None, true)] + #[case::phrase_hyphenated("rate-limit", None, true)] + #[case::service_tier("Service tier capacity exceeded", None, true)] + #[case::unrelated("rejected", Some(429), false)] + fn rate_limit_detection( + #[case] text: &str, + #[case] status: Option, + #[case] expected: bool, + ) { + assert_eq!(is_rate_limit(text, status), expected); + } + + #[rstest::rstest] + #[case::exceed_context_limit("Exceed context limit", true)] + #[case::maximum_context_length("This model's maximum context length is 10", true)] + #[case::string_too_long("string too long. Expected a string with maximum length 5", true)] + #[case::maximum_context_limit("the model's maximum context limit", true)] + #[case::longer_than_context("prompt is longer than the model's context length", true)] + #[case::configured_limit("input tokens exceed the configured limit", true)] + #[case::max_new_tokens("`inputs` tokens + `max_new_tokens` must be <= 10", true)] + #[case::available_context("exceeds the available context size", true)] + #[case::maximum_tokens("exceeds the maximum number of tokens allowed", true)] + #[case::current_and_limit("current length is 9 while limit is 8", true)] + #[case::current_without_limit("current length is 9", false)] + #[case::maximum_input_tokens("maximum input length is 8 tokens", true)] + #[case::maximum_input_without_tokens("maximum input length is 8", false)] + #[case::string_above_max_length_wins("string_above_max_length exceed context limit", false)] + #[case::user_field_is_not_context( + "invalid 'user': string too long. expected a string with maximum length", + false + )] + #[case::unrelated("rejected", false)] + fn context_window_detection(#[case] text: &str, #[case] expected: bool) { + assert_eq!(is_context_window_exceeded(text), expected); + } + + #[rstest::rstest] + #[case::integer(r#"{"error": {"code": 429}}"#, Some(429))] + #[case::float(r#"{"error": {"code": 429.9}}"#, Some(429))] + #[case::string(r#"{"error": {"code": " 4_29 "}}"#, Some(429))] + #[case::boolean(r#"{"error": {"code": true}}"#, Some(1))] + #[case::unparseable_string(r#"{"error": {"code": "slow"}}"#, None)] + #[case::null(r#"{"error": {"code": null}}"#, None)] + #[case::no_code(r#"{"error": {}}"#, None)] + #[case::error_not_an_object(r#"{"error": "429"}"#, None)] + #[case::no_error(r#"{"code": 429}"#, None)] + #[case::not_an_object("[429]", None)] + #[case::not_json("429", None)] + fn body_error_code_reads_the_nested_code(#[case] body: &str, #[case] expected: Option) { + assert_eq!(body_error_code(body), expected); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs new file mode 100644 index 00000000000..3d817fe9ae2 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs @@ -0,0 +1,153 @@ +use super::public::{PublicFailure, StatusClass}; +use super::rules::{ApiStatus, Kind, ResponseChoice, Rule, apply}; +use super::{DOCS_URL, Mapping}; + +const fn with_response(class: StatusClass) -> Kind { + Kind::Status { + class, + response: ResponseChoice::Provider, + } +} + +fn message(mapping: &Mapping<'_>) -> String { + format!("{} - {}", mapping.exception_provider, mapping.error_str) +} + +fn status(mapping: &Mapping<'_>) -> u16 { + mapping.original.status.unwrap_or_default() +} + +/// `_map_exception_by_status`, the fallback for a provider error no provider mapper claimed. +const RULES: &[Rule] = &[ + Rule { + when: |mapping| status(mapping) == 401, + kind: with_response(StatusClass::Authentication), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 403, + kind: with_response(StatusClass::PermissionDenied), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 404, + kind: with_response(StatusClass::NotFound), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 408, + kind: Kind::Timeout(None), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 429, + kind: with_response(StatusClass::RateLimit), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 500, + kind: with_response(StatusClass::InternalServer), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 502, + kind: with_response(StatusClass::BadGateway), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 503, + kind: with_response(StatusClass::ServiceUnavailable), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) == 504, + kind: Kind::Timeout(Some(504)), + message, + debug: true, + }, + Rule { + when: |mapping| status(mapping) < 500, + kind: with_response(StatusClass::BadRequest), + message, + debug: true, + }, + Rule { + when: |_| true, + kind: Kind::Api { + status: ApiStatus::Original, + request_url: DOCS_URL, + }, + message, + debug: true, + }, +]; + +/// Only a real provider status of 400 or more reaches the table; a status the HTTP handler +/// synthesized for a failure without a response does not. +pub(super) fn map(mapping: &Mapping<'_>) -> Option { + let status = mapping.original.status?; + if status < 400 || mapping.original.status_is_synthesized { + return None; + } + apply(RULES, mapping) +} + +#[cfg(test)] +mod tests { + use super::super::testing::{context, failure, http, upstream, with_debug}; + use super::super::{ExceptionFamily, OriginalException, PublicKind}; + use super::*; + + fn mapped(original: &OriginalException) -> Option { + let context = context("reducto", ExceptionFamily::Other); + map(&Mapping::new(&context, original)) + } + + fn classified(class: StatusClass, status_code: u16) -> PublicKind { + PublicKind::Status { + status_class: class, + response: upstream(status_code, "rejected"), + } + } + + #[rstest::rstest] + #[case::authentication(401, classified(StatusClass::Authentication, 401))] + #[case::permission_denied(403, classified(StatusClass::PermissionDenied, 403))] + #[case::not_found(404, classified(StatusClass::NotFound, 404))] + #[case::request_timeout(408, PublicKind::Timeout { status: None })] + #[case::rate_limited(429, classified(StatusClass::RateLimit, 429))] + #[case::internal_server(500, classified(StatusClass::InternalServer, 500))] + #[case::bad_gateway(502, classified(StatusClass::BadGateway, 502))] + #[case::service_unavailable(503, classified(StatusClass::ServiceUnavailable, 503))] + #[case::gateway_timeout(504, PublicKind::Timeout { status: Some(504) })] + #[case::lowest_client_error(400, classified(StatusClass::BadRequest, 400))] + #[case::other_client_error(409, classified(StatusClass::BadRequest, 409))] + #[case::highest_client_error(499, classified(StatusClass::BadRequest, 499))] + #[case::other_server_error(501, PublicKind::Api { status: 501, request_url: DOCS_URL })] + fn every_mapped_status_and_the_fallback(#[case] status_code: u16, #[case] kind: PublicKind) { + assert_eq!( + mapped(&http(status_code, "rejected")), + Some(with_debug(failure( + kind, + "ReductoException - rejected", + "reducto" + ))) + ); + } + + #[rstest::rstest] + #[case::below_client_errors(http(399, "rejected"))] + #[case::synthesized(OriginalException::Connection { message: "refused".into() })] + #[case::no_status(OriginalException::Response { message: "bad body".into() })] + fn failures_the_table_does_not_claim(#[case] original: OriginalException) { + assert_eq!(mapped(&original), None); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs new file mode 100644 index 00000000000..0fa8c19c4f6 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs @@ -0,0 +1,574 @@ +use super::public::{PublicFailure, PublicKind, ResponseArg, StatusClass, UpstreamResponse}; +use super::rules::{ + Kind, ResponseChoice, Rule, apply, body_error_code, contains_any, is_context_window_exceeded, +}; +use super::{Mapping, python_capitalize}; + +const VERTEX_URL: &str = "https://cloud.google.com/vertex-ai/"; +const VERTEX_URL_WITH_SPACE: &str = " https://cloud.google.com/vertex-ai/"; + +const QUOTA_MARKERS: &[&str] = &[ + "429 Quota exceeded", + "Quota exceeded for", + "Resource exhausted", + "IndexError: list index out of range", + "429 Unable to submit request because the service is temporarily out of capacity.", +]; + +const fn stubbed(class: StatusClass, status: u16, url: &'static str) -> Kind { + Kind::Status { + class, + response: ResponseChoice::Stub { status, url }, + } +} + +const fn bare(class: StatusClass) -> Kind { + Kind::Status { + class, + response: ResponseChoice::Omitted, + } +} + +/// `{Provider}Exception{label} - {error_str}` with Python's `str.capitalize()`. +fn capitalized(mapping: &Mapping<'_>, label: &str) -> String { + format!( + "{}Exception{label} - {}", + python_capitalize(mapping.provider), + mapping.error_str + ) +} + +/// `litellm.{Class}: {provider}Exception - {error_str}` with the provider as given. +fn litellm_prefixed(mapping: &Mapping<'_>, class: &str) -> String { + format!( + "litellm.{class}: {}Exception - {}", + mapping.provider, mapping.error_str + ) +} + +fn status_is(mapping: &Mapping<'_>, status: u16) -> bool { + mapping.original.status == Some(status) +} + +/// `_map_vertex_exception`, in its branch order. A failure no rule claims falls through +/// to the status table. +const RULES: &[Rule] = &[ + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &[ + "Vertex AI API has not been used in project", + "Unable to find your project", + ], + ) + }, + kind: stubbed(StatusClass::BadRequest, 400, VERTEX_URL_WITH_SPACE), + message: |mapping| litellm_prefixed(mapping, "BadRequestError"), + debug: true, + }, + Rule { + when: |mapping| { + mapping + .error_str + .contains("400 Request payload size exceeds") + }, + kind: bare(StatusClass::ContextWindowExceeded), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| is_context_window_exceeded(&mapping.error_str), + kind: bare(StatusClass::ContextWindowExceeded), + message: |mapping| format!("ContextWindowExceededError: {}", capitalized(mapping, "")), + debug: true, + }, + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &["None Unknown Error.", "Content has no parts."], + ) + }, + kind: Kind::Status { + class: StatusClass::InternalServer, + response: ResponseChoice::InternalServerStub, + }, + message: |mapping| litellm_prefixed(mapping, "InternalServerError"), + debug: true, + }, + Rule { + when: |mapping| mapping.error_str.contains("API key not valid."), + kind: bare(StatusClass::Authentication), + message: |mapping| capitalized(mapping, ""), + debug: true, + }, + Rule { + when: |mapping| mapping.error_str.contains("403"), + kind: stubbed(StatusClass::BadRequest, 403, VERTEX_URL_WITH_SPACE), + message: |mapping| capitalized(mapping, " BadRequestError"), + debug: true, + }, + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &[ + "The response was blocked.", + "Output blocked by content filtering policy", + ], + ) + }, + kind: stubbed( + StatusClass::ContentPolicyViolation, + 400, + VERTEX_URL_WITH_SPACE, + ), + message: |mapping| capitalized(mapping, " ContentPolicyViolationError"), + debug: true, + }, + Rule { + when: |mapping| { + contains_any(&mapping.error_str, QUOTA_MARKERS) + || (mapping + .original + .status + .is_some_and(|status| (500..600).contains(&status)) + && body_error_code(&mapping.error_str) == Some(429)) + }, + kind: stubbed(StatusClass::RateLimit, 429, VERTEX_URL_WITH_SPACE), + message: |mapping| litellm_prefixed(mapping, "RateLimitError"), + debug: true, + }, + Rule { + when: |mapping| { + contains_any( + &mapping.error_str, + &["500 Internal Server Error", "The model is overloaded."], + ) + }, + kind: bare(StatusClass::InternalServer), + message: |mapping| litellm_prefixed(mapping, "InternalServerError"), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, 400), + kind: stubbed(StatusClass::BadRequest, 400, VERTEX_URL), + message: |mapping| capitalized(mapping, " BadRequestError"), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, 401), + kind: bare(StatusClass::Authentication), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, 403), + kind: stubbed(StatusClass::PermissionDenied, 403, VERTEX_URL), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, 404), + kind: bare(StatusClass::NotFound), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, 408), + kind: Kind::Timeout(None), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, 429), + kind: stubbed(StatusClass::RateLimit, 429, VERTEX_URL_WITH_SPACE), + message: |mapping| format!("litellm.RateLimitError: {}", capitalized(mapping, "")), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, 500), + kind: Kind::Status { + class: StatusClass::InternalServer, + response: ResponseChoice::InternalServerStub, + }, + message: |mapping| capitalized(mapping, " InternalServerError"), + debug: true, + }, + Rule { + when: |mapping| status_is(mapping, 502), + kind: Kind::ApiConnection, + message: |mapping| capitalized(mapping, ""), + debug: false, + }, + Rule { + when: |mapping| status_is(mapping, 503), + kind: bare(StatusClass::ServiceUnavailable), + message: |mapping| capitalized(mapping, ""), + debug: false, + }, +]; + +pub(super) fn map(mapping: &Mapping<'_>) -> Option { + apply(RULES, mapping).map(|failure| keep_upstream_response(mapping, failure)) +} + +/// Deliberate divergence from `_map_vertex_exception`, which replaces the provider response +/// with a stub and so drops the upstream body and `retry-after`. The response keeps the +/// status the public class carries. +fn keep_upstream_response(mapping: &Mapping<'_>, failure: PublicFailure) -> PublicFailure { + let (PublicKind::Status { status_class, .. }, Some(upstream), false) = ( + &failure.kind, + &mapping.original.response, + mapping.original.status_is_synthesized, + ) else { + return failure; + }; + PublicFailure { + kind: PublicKind::Status { + status_class: *status_class, + response: Some(ResponseArg::Upstream(UpstreamResponse { + status: status_class.status_code(), + ..upstream.clone() + })), + }, + ..failure + } +} + +#[cfg(test)] +mod tests { + use super::super::testing::{context, failure, http, status, upstream, with_debug}; + use super::super::{ExceptionFamily, HttpStub, OriginalException}; + use super::*; + + fn mapped(original: &OriginalException) -> Option { + let context = context("vertex_ai", ExceptionFamily::VertexAi); + map(&Mapping::new(&context, original)) + } + + fn kept(class: StatusClass, body: &str) -> PublicKind { + status(class, upstream(class.status_code(), body)) + } + + #[rstest::rstest] + #[case::api_not_enabled( + 400, + "Vertex AI API has not been used in project x", + with_debug(failure( + kept( + StatusClass::BadRequest, + "Vertex AI API has not been used in project x" + ), + "litellm.BadRequestError: vertex_aiException - Vertex AI API has not been used in project x", + "vertex_ai", + )) + )] + #[case::project_not_found( + 400, + "Unable to find your project", + with_debug(failure( + kept(StatusClass::BadRequest, "Unable to find your project"), + "litellm.BadRequestError: vertex_aiException - Unable to find your project", + "vertex_ai", + )) + )] + #[case::payload_too_large( + 400, + "400 Request payload size exceeds the limit", + failure( + kept( + StatusClass::ContextWindowExceeded, + "400 Request payload size exceeds the limit" + ), + "Vertex_aiException - 400 Request payload size exceeds the limit", + "vertex_ai", + ) + )] + #[case::context_window( + 500, + "This model's maximum context length is 10", + with_debug(failure( + kept( + StatusClass::ContextWindowExceeded, + "This model's maximum context length is 10" + ), + "ContextWindowExceededError: Vertex_aiException - This model's maximum context length is 10", + "vertex_ai", + )) + )] + #[case::unknown_error( + 400, + "None Unknown Error.", + with_debug(failure( + kept(StatusClass::InternalServer, "None Unknown Error."), + "litellm.InternalServerError: vertex_aiException - None Unknown Error.", + "vertex_ai", + )) + )] + #[case::no_parts( + 400, + "Content has no parts.", + with_debug(failure( + kept(StatusClass::InternalServer, "Content has no parts."), + "litellm.InternalServerError: vertex_aiException - Content has no parts.", + "vertex_ai", + )) + )] + #[case::api_key_not_valid( + 400, + "API key not valid.", + with_debug(failure( + kept(StatusClass::Authentication, "API key not valid."), + "Vertex_aiException - API key not valid.", + "vertex_ai", + )) + )] + #[case::forbidden_text( + 400, + "got a 403", + with_debug(failure( + kept(StatusClass::BadRequest, "got a 403"), + "Vertex_aiException BadRequestError - got a 403", + "vertex_ai", + )) + )] + #[case::response_blocked( + 400, + "The response was blocked.", + with_debug(failure( + kept(StatusClass::ContentPolicyViolation, "The response was blocked."), + "Vertex_aiException ContentPolicyViolationError - The response was blocked.", + "vertex_ai", + )) + )] + #[case::output_blocked( + 400, + "Output blocked by content filtering policy", + with_debug(failure( + kept( + StatusClass::ContentPolicyViolation, + "Output blocked by content filtering policy" + ), + "Vertex_aiException ContentPolicyViolationError - Output blocked by content filtering policy", + "vertex_ai", + )) + )] + #[case::quota_marker( + 400, + "Quota exceeded for aiplatform", + with_debug(failure( + kept(StatusClass::RateLimit, "Quota exceeded for aiplatform"), + "litellm.RateLimitError: vertex_aiException - Quota exceeded for aiplatform", + "vertex_ai", + )) + )] + #[case::wrapped_429( + 503, + r#"{"error": {"code": "429"}}"#, + with_debug(failure( + kept(StatusClass::RateLimit, r#"{"error": {"code": "429"}}"#), + r#"litellm.RateLimitError: vertex_aiException - {"error": {"code": "429"}}"#, + "vertex_ai", + )) + )] + #[case::overloaded( + 400, + "The model is overloaded.", + with_debug(failure( + kept(StatusClass::InternalServer, "The model is overloaded."), + "litellm.InternalServerError: vertex_aiException - The model is overloaded.", + "vertex_ai", + )) + )] + #[case::internal_server_text( + 400, + "500 Internal Server Error", + with_debug(failure( + kept(StatusClass::InternalServer, "500 Internal Server Error"), + "litellm.InternalServerError: vertex_aiException - 500 Internal Server Error", + "vertex_ai", + )) + )] + fn each_text_rule_maps_by_the_body( + #[case] status_code: u16, + #[case] body: &str, + #[case] expected: PublicFailure, + ) { + assert_eq!(mapped(&http(status_code, body)), Some(expected)); + } + + #[rstest::rstest] + #[case::bad_request( + 400, + with_debug(failure( + kept(StatusClass::BadRequest, "rejected"), + "Vertex_aiException BadRequestError - rejected", + "vertex_ai" + )) + )] + #[case::authentication( + 401, + failure( + kept(StatusClass::Authentication, "rejected"), + "Vertex_aiException - rejected", + "vertex_ai" + ) + )] + #[case::permission_denied( + 403, + failure( + kept(StatusClass::PermissionDenied, "rejected"), + "Vertex_aiException - rejected", + "vertex_ai" + ) + )] + #[case::not_found( + 404, + failure( + kept(StatusClass::NotFound, "rejected"), + "Vertex_aiException - rejected", + "vertex_ai" + ) + )] + #[case::request_timeout(408, failure(PublicKind::Timeout { status: None }, "Vertex_aiException - rejected", "vertex_ai"))] + #[case::rate_limited( + 429, + with_debug(failure( + kept(StatusClass::RateLimit, "rejected"), + "litellm.RateLimitError: Vertex_aiException - rejected", + "vertex_ai" + )) + )] + #[case::internal_server( + 500, + with_debug(failure( + kept(StatusClass::InternalServer, "rejected"), + "Vertex_aiException InternalServerError - rejected", + "vertex_ai" + )) + )] + #[case::bad_gateway( + 502, + failure( + PublicKind::ApiConnection, + "Vertex_aiException - rejected", + "vertex_ai" + ) + )] + #[case::service_unavailable( + 503, + failure( + kept(StatusClass::ServiceUnavailable, "rejected"), + "Vertex_aiException - rejected", + "vertex_ai" + ) + )] + fn each_status_rule_maps_by_the_status( + #[case] status_code: u16, + #[case] expected: PublicFailure, + ) { + assert_eq!(mapped(&http(status_code, "rejected")), Some(expected)); + } + + #[rstest::rstest] + #[case::unmapped_status(409)] + #[case::gateway_timeout(504)] + fn statuses_without_a_rule_fall_through(#[case] status_code: u16) { + assert_eq!(mapped(&http(status_code, "rejected")), None); + } + + #[rstest::rstest] + #[case::stub_without_an_upstream_response( + OriginalException::Response { message: "got a 403".into() }, + status(StatusClass::BadRequest, Some(ResponseArg::Stub(HttpStub { status: 403, method: "POST", url: VERTEX_URL_WITH_SPACE, content: None }))) + )] + #[case::stub_for_a_synthesized_status( + OriginalException::Connection { message: "got a 403".into() }, + status(StatusClass::BadRequest, Some(ResponseArg::Stub(HttpStub { status: 403, method: "POST", url: VERTEX_URL_WITH_SPACE, content: None }))) + )] + fn the_rule_response_stays_when_there_is_no_real_upstream_response( + #[case] original: OriginalException, + #[case] kind: PublicKind, + ) { + assert_eq!(mapped(&original).map(|failure| failure.kind), Some(kind)); + } + + #[test] + fn a_synthesized_500_keeps_the_internal_server_stub() { + let original = OriginalException::Connection { + message: "refused".into(), + }; + assert_eq!( + mapped(&original), + Some(with_debug(failure( + status( + StatusClass::InternalServer, + Some(ResponseArg::Stub(HttpStub { + status: 500, + method: "completion", + url: "https://github.com/BerriAI/litellm", + content: Some("refused".into()), + })) + ), + "Vertex_aiException InternalServerError - refused", + "vertex_ai" + ))) + ); + } + + #[rstest::rstest] + #[case::project_before_payload_size( + "Unable to find your project 400 Request payload size exceeds", + StatusClass::BadRequest, + "litellm.BadRequestError: vertex_aiException - Unable to find your project 400 Request payload size exceeds", + true + )] + #[case::payload_size_before_context_window( + "400 Request payload size exceeds; This model's maximum context length is 10", + StatusClass::ContextWindowExceeded, + "Vertex_aiException - 400 Request payload size exceeds; This model's maximum context length is 10", + false + )] + #[case::api_key_before_forbidden( + "API key not valid. 403", + StatusClass::Authentication, + "Vertex_aiException - API key not valid. 403", + true + )] + #[case::forbidden_before_blocked( + "403 The response was blocked.", + StatusClass::BadRequest, + "Vertex_aiException BadRequestError - 403 The response was blocked.", + true + )] + #[case::blocked_before_quota( + "The response was blocked. Resource exhausted", + StatusClass::ContentPolicyViolation, + "Vertex_aiException ContentPolicyViolationError - The response was blocked. Resource exhausted", + true + )] + #[case::quota_before_overloaded( + "Resource exhausted The model is overloaded.", + StatusClass::RateLimit, + "litellm.RateLimitError: vertex_aiException - Resource exhausted The model is overloaded.", + true + )] + fn the_earlier_rule_wins_when_two_apply( + #[case] body: &str, + #[case] class: StatusClass, + #[case] message: &str, + #[case] debug: bool, + ) { + let expected = failure(kept(class, body), message, "vertex_ai"); + assert_eq!( + mapped(&http(401, body)), + Some(if debug { + with_debug(expected) + } else { + expected + }) + ); + } +} diff --git a/litellm-rust/crates/core-utils/src/lib.rs b/litellm-rust/crates/core-utils/src/lib.rs index a8895cccf2c..0c26aa50cc3 100644 --- a/litellm-rust/crates/core-utils/src/lib.rs +++ b/litellm-rust/crates/core-utils/src/lib.rs @@ -1,7 +1,10 @@ pub mod call_arguments; pub mod core_helpers; +pub mod exception_mapping_utils; pub mod get_llm_provider_logic; pub mod params; pub mod prompt_templates; +pub mod python_repr; +pub mod secret_redaction; pub mod serde_compat; pub mod url_utils; diff --git a/litellm-rust/crates/core-utils/src/python_repr.rs b/litellm-rust/crates/core-utils/src/python_repr.rs new file mode 100644 index 00000000000..7ccfb826377 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/python_repr.rs @@ -0,0 +1,93 @@ +/// `repr()` of a Python `str`: single quotes unless the text holds a single quote and no +/// double quote, with backslashes, the chosen quote and control characters escaped. +pub fn python_str_repr(value: &str) -> String { + let quote = if value.contains('\'') && !value.contains('"') { + '"' + } else { + '\'' + }; + let escaped: String = value + .chars() + .map(|character| match character { + '\\' => "\\\\".to_string(), + '\t' => "\\t".to_string(), + '\n' => "\\n".to_string(), + '\r' => "\\r".to_string(), + character if character == quote => format!("\\{character}"), + character + if (character as u32) < 0x20 || (0x7f..0xa0).contains(&(character as u32)) => + { + format!("\\x{:02x}", character as u32) + } + character => character.to_string(), + }) + .collect(); + format!("{quote}{escaped}{quote}") +} + +/// `repr()` of the Python value a JSON value decodes to. +pub fn python_value_repr(value: &serde_json::Value) -> String { + use serde_json::Value; + match value { + Value::Null => "None".to_string(), + Value::Bool(true) => "True".to_string(), + Value::Bool(false) => "False".to_string(), + Value::Number(number) => number.to_string(), + Value::String(text) => python_str_repr(text), + Value::Array(items) => format!( + "[{}]", + items + .iter() + .map(python_value_repr) + .collect::>() + .join(", ") + ), + Value::Object(fields) => format!( + "{{{}}}", + fields + .iter() + .map(|(key, value)| format!( + "{}: {}", + python_str_repr(key), + python_value_repr(value) + )) + .collect::>() + .join(", ") + ), + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::{python_str_repr, python_value_repr}; + + #[rstest::rstest] + #[case::null(json!(null), "None")] + #[case::true_(json!(true), "True")] + #[case::false_(json!(false), "False")] + #[case::integer(json!(5), "5")] + #[case::float(json!(1.5), "1.5")] + #[case::string(json!("it's"), "\"it's\"")] + #[case::list(json!(["a", 1]), "['a', 1]")] + #[case::dict(json!({"format": "native"}), "{'format': 'native'}")] + #[case::empty_list(json!([]), "[]")] + fn value_repr_matches_python(#[case] value: serde_json::Value, #[case] expected: &str) { + assert_eq!(python_value_repr(&value), expected); + } + + #[rstest::rstest] + #[case::plain("native", "'native'")] + #[case::single_quote("it's", "\"it's\"")] + #[case::both_quotes("it's \"x\"", "'it\\'s \"x\"'")] + #[case::double_quote("say \"x\"", "'say \"x\"'")] + #[case::backslash("a\\b", "'a\\\\b'")] + #[case::whitespace("a\tb\nc\rd", "'a\\tb\\nc\\rd'")] + #[case::control("a\u{1}b\u{7f}c\u{85}", "'a\\x01b\\x7fc\\x85'")] + #[case::unicode("café", "'café'")] + #[case::empty("", "''")] + fn matches_python_repr(#[case] value: &str, #[case] expected: &str) { + assert_eq!(python_str_repr(value), expected); + } +} diff --git a/litellm-rust/crates/core-utils/src/secret_redaction.rs b/litellm-rust/crates/core-utils/src/secret_redaction.rs new file mode 100644 index 00000000000..32922d7430c --- /dev/null +++ b/litellm-rust/crates/core-utils/src/secret_redaction.rs @@ -0,0 +1,97 @@ +use std::sync::LazyLock; + +use fancy_regex::Regex; + +pub const REDACTED: &str = "REDACTED"; + +const DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH: usize = 16; + +fn minimum_custom_key_length() -> usize { + std::env::var("MINIMUM_CUSTOM_KEY_LENGTH") + .ok() + .and_then(|value| value.trim().parse().ok()) + .unwrap_or(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH) +} + +fn secret_patterns(minimum_custom_key_length: usize) -> String { + let sk_suffix_length = minimum_custom_key_length.saturating_sub("sk-".len()); + [ + r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----", + r"\bya29\.[A-Za-z0-9_.~+/-]+", + r#"(?:client_secret|azure_password|azure_username)\s+[^\s,'"})\]{}>]+"#, + r"(?:AKIA|ASIA)[0-9A-Z]{16}", + r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*", + r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}", + &format!(r"sk-[A-Za-z0-9\-_]{{{sk_suffix_length},}}"), + r#"(?<=[?&])(?:api[_-]?key|\w*(?:token|password|passwd|client_secret|secret_key|_secret))=[^\s&'"]+"#, + r#"(?:api[_-]?key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]{8,}"#, + r#"(?:x-api-key|api-key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#, + r"x-ak-[A-Za-z0-9\-_]{20,}", + r"AIza[0-9A-Za-z\-_]{35}", + r#"(?<=[?&])key=[^\s&'"]{8,}"#, + r#"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#, + r#"(?<=://)[^\s'":]{0,4096}:[^\s'"]{1,4096}(?=@)"#, + r"dapi[0-9a-f]{32}", + r#"litellm\.[A-Za-z0-9_]*_key['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#, + r#"private_key['"]?\s*[:=]\s*['"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'"})\]{}>]+)"#, + concat!( + r"(?:master_key|xai_key|database_url|db_url|connection_string|", + r"aws_secret_access_key|aws_session_token|aws_access_key_id|", + r"signing_key|encryption_key|", + r"auth_token|access_token|refresh_token|", + r"slack_webhook_url|webhook_url|", + r"database_connection_string|", + r"huggingface_token|jwt_secret)", + r#"['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#, + ), + r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*", + r"(?<=[?&])sig=[A-Za-z0-9%+/=]+", + r#"\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}"#, + ] + .join("|") +} + +static SECRET_RE: LazyLock = LazyLock::new(|| { + Regex::new(&format!( + "(?i){}", + secret_patterns(minimum_custom_key_length()) + )) + .expect("secret redaction patterns compile") +}); + +pub fn redact_string(value: &str) -> String { + SECRET_RE.replace_all(value, REDACTED).into_owned() +} + +pub fn secret_redaction_enabled() -> bool { + !std::env::var("LITELLM_DISABLE_REDACT_SECRETS") + .is_ok_and(|value| value.eq_ignore_ascii_case("true")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[rstest::rstest] + #[case::bearer("auth failed: Bearer abcdefghijklmnop", "auth failed: REDACTED")] + #[case::sk_key("key sk-abcdefghijklmnopqrstuvwxyz rejected", "key REDACTED rejected")] + #[case::short_sk_key_is_kept("sk-abc", "sk-abc")] + #[case::query_param("GET /v1?api_key=secret123&x=1", "GET /v1?REDACTED&x=1")] + #[case::dict_repr("{'api_key': 'abcdefghij'}", "{'REDACTED'}")] + #[case::url_credentials("postgres://user:pass@host/db", "postgres://REDACTED@host/db")] + #[case::case_insensitive("BEARER ABCDEFGHIJKLMNOP", "REDACTED")] + #[case::aws_key("AKIAABCDEFGHIJKLMNOP", "REDACTED")] + #[case::sas_signature("https://x.blob/a?sv=1&sig=abc%2B=", "https://x.blob/a?sv=1&REDACTED")] + #[case::password_needs_word_boundary("db_password=hunter2", "REDACTED")] + #[case::plain_text_is_kept(r#"{"message": "rejected"}"#, r#"{"message": "rejected"}"#)] + fn redacts_the_same_spans_as_the_python_patterns(#[case] input: &str, #[case] expected: &str) { + assert_eq!(redact_string(input), expected); + } + + #[test] + fn sk_threshold_follows_the_minimum_custom_key_length() { + let patterns = Regex::new(&format!("(?i){}", secret_patterns(8))).unwrap(); + assert_eq!(patterns.replace_all("sk-abcde", REDACTED), REDACTED); + assert_eq!(patterns.replace_all("sk-abcd", REDACTED), "sk-abcd"); + } +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json new file mode 100644 index 00000000000..ae629114589 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json @@ -0,0 +1,13 @@ +{ + "kind": { + "type": "api", + "status": 409, + "request_url": "https://docs.litellm.ai/docs" + }, + "message": "MistralException - api", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json new file mode 100644 index 00000000000..ab400ed9e01 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json @@ -0,0 +1,11 @@ +{ + "kind": { + "type": "api_connection" + }, + "message": "MistralException - api_connection", + "model": "ocr-model", + "llm_provider": null, + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json new file mode 100644 index 00000000000..057392575a1 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json @@ -0,0 +1,13 @@ +{ + "kind": { + "type": "status", + "status_class": "authentication", + "response": null + }, + "message": "MistralException - status_authentication", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json new file mode 100644 index 00000000000..abb1425f686 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "bad_gateway", + "response": { + "type": "upstream", + "status": 502, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_bad_gateway", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json new file mode 100644 index 00000000000..171a994cd35 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "bad_request", + "response": { + "type": "upstream", + "status": 400, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_bad_request", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json new file mode 100644 index 00000000000..ae9c1e145de --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "content_policy_violation", + "response": { + "type": "upstream", + "status": 400, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_content_policy_violation", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json new file mode 100644 index 00000000000..61e1a56a622 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "context_window_exceeded", + "response": { + "type": "upstream", + "status": 400, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_context_window_exceeded", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json new file mode 100644 index 00000000000..b3c5c51a785 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json @@ -0,0 +1,19 @@ +{ + "kind": { + "type": "status", + "status_class": "internal_server", + "response": { + "type": "stub", + "status": 500, + "method": "completion", + "url": "https://github.com/BerriAI/litellm", + "content": "upstream text" + } + }, + "message": "MistralException - status_internal_server", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json new file mode 100644 index 00000000000..26ffd872961 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "not_found", + "response": { + "type": "upstream", + "status": 404, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_not_found", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json new file mode 100644 index 00000000000..42772f98830 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json @@ -0,0 +1,19 @@ +{ + "kind": { + "type": "status", + "status_class": "permission_denied", + "response": { + "type": "stub", + "status": 403, + "method": "POST", + "url": " https://cloud.google.com/vertex-ai/", + "content": null + } + }, + "message": "MistralException - status_permission_denied", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json new file mode 100644 index 00000000000..c9b88822ebd --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "rate_limit", + "response": { + "type": "upstream", + "status": 429, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_rate_limit", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json new file mode 100644 index 00000000000..5af65202ef1 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "service_unavailable", + "response": { + "type": "upstream", + "status": 503, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_service_unavailable", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json new file mode 100644 index 00000000000..a1b318ce03c --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json @@ -0,0 +1,28 @@ +{ + "kind": { + "type": "status", + "status_class": "unsupported_params", + "response": { + "type": "upstream", + "status": 400, + "body": "{\"message\": \"rejected\"}", + "headers": [ + [ + "retry-after", + "7" + ] + ] + } + }, + "message": "MistralException - status_unsupported_params", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": [ + [ + "retry-after", + "7" + ] + ], + "print_banner": false +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json new file mode 100644 index 00000000000..8b215bdb7e6 --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json @@ -0,0 +1,12 @@ +{ + "kind": { + "type": "timeout", + "status": 504 + }, + "message": "MistralException - timeout_with_status", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": "\nModel: ocr-model", + "litellm_response_headers": null, + "print_banner": true +} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json new file mode 100644 index 00000000000..4a79169776e --- /dev/null +++ b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json @@ -0,0 +1,12 @@ +{ + "kind": { + "type": "timeout", + "status": null + }, + "message": "MistralException - timeout_without_status", + "model": "ocr-model", + "llm_provider": "mistral", + "litellm_debug_info": null, + "litellm_response_headers": null, + "print_banner": false +} From 5bbdc72879cbbfd2c3970b09fb189872964484c3 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 20:02:48 +0000 Subject: [PATCH 177/253] fix(gateway): expose /deepgram/ passthrough routes on the gateway allowlist Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- gateway/routes/allowlist.py | 1 + 1 file changed, 1 insertion(+) diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 1c503a083f1..30a5374e91d 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -94,6 +94,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = ( "/vertex-ai/", "/assemblyai/", "/eu.assemblyai/", + "/deepgram/", "/langfuse/", "/vllm/", "/mistral/", From 74e9fb23233c8d3bfeed73485bc0270da79370d3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:07:33 -0700 Subject: [PATCH 178/253] fix(passthrough): validate only the winning timeout value in the resolver --- litellm/passthrough/timeout_utils.py | 36 +++++++++---------- .../test_pass_through_endpoints.py | 20 +++++++++++ 2 files changed, 38 insertions(+), 18 deletions(-) diff --git a/litellm/passthrough/timeout_utils.py b/litellm/passthrough/timeout_utils.py index 9284f7143b0..fc67aa8c553 100644 --- a/litellm/passthrough/timeout_utils.py +++ b/litellm/passthrough/timeout_utils.py @@ -1,18 +1,14 @@ import sys from collections.abc import Mapping +from types import MappingProxyType from typing import Final -from pydantic import BaseModel, ConfigDict +from pydantic import TypeAdapter DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0 - -class _TimeoutFields(BaseModel): - model_config = ConfigDict(frozen=True) - - stream_timeout: float | None = None - timeout: float | None = None - request_timeout: float | None = None +_SECONDS: Final = TypeAdapter(float) +_NO_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) def resolve_pass_through_request_timeout( @@ -59,20 +55,24 @@ def resolve_llm_passthrough_timeout( any generic timeout, matching ``Router._get_stream_timeout`` on the completion route: kwargs stream_timeout -> litellm_params stream_timeout -> router_stream_timeout, then the non-streaming chain above. + + Only the first set value is validated as seconds, so a value in a lower-precedence + field never fails the call. """ - streaming: Final = bool((kwargs or {}).get("stream")) - request: Final = _TimeoutFields.model_validate(kwargs or {}) - deployment: Final = _TimeoutFields.model_validate(litellm_params or {}) + request: Final = kwargs if kwargs is not None else _NO_PARAMS + deployment: Final = litellm_params if litellm_params is not None else _NO_PARAMS stream_candidates: Final = ( - (request.stream_timeout, deployment.stream_timeout, router_stream_timeout) if streaming else () + (request.get("stream_timeout"), deployment.get("stream_timeout"), router_stream_timeout) + if request.get("stream") + else () ) candidates: Final = ( *stream_candidates, - request.timeout, - request.request_timeout, - deployment.timeout, - deployment.request_timeout, + request.get("timeout"), + request.get("request_timeout"), + deployment.get("timeout"), + deployment.get("request_timeout"), router_timeout, ) - resolved: Final = next((float(val) for val in candidates if val is not None), None) - return resolved if resolved is not None else resolve_pass_through_request_timeout() + winner: Final = next((val for val in candidates if val is not None), None) + return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index abff897c4f5..13cd41278b9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from fastapi import Request, Response, UploadFile +from pydantic import ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -1189,6 +1190,25 @@ def test_resolve_llm_passthrough_timeout_reads_stream_by_truthiness(stream: obje ) +@pytest.mark.parametrize( + "kwargs, litellm_params, expected", + [ + ({"stream": True, "stream_timeout": 1800, "timeout": httpx.Timeout(30.0)}, {}, 1800.0), + ({"stream": False}, {"stream_timeout": httpx.Timeout(30.0), "timeout": 90}, 90.0), + ({"timeout": 45}, {"request_timeout": httpx.Timeout(30.0)}, 45.0), + ], +) +def test_resolve_llm_passthrough_timeout_validates_only_the_winning_value( + kwargs: dict[str, object], litellm_params: dict[str, object], expected: float +): + assert resolve_llm_passthrough_timeout(kwargs=kwargs, litellm_params=litellm_params) == expected + + +def test_resolve_llm_passthrough_timeout_rejects_a_non_numeric_winner(): + with pytest.raises(ValidationError): + resolve_llm_passthrough_timeout(kwargs={"timeout": httpx.Timeout(30.0)}) + + @pytest.mark.asyncio async def test_pass_through_request_uses_resolved_timeout(): with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: From d37320f8cf242f310bad3e759c9a7dbbdc27b699 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 13:08:55 -0700 Subject: [PATCH 179/253] fix(proxy): report null cost for unpriced deployments on the /model/info id lookup GET /model/info?litellm_model_id= went through _get_proxy_model_info, a copy of _enrich_model_info_with_litellm_data that missed the unpriced check, so the same deployment read as null in the list and as 0 on the id lookup. The id lookup now delegates to the shared helper, so the two paths cannot drift again --- litellm/proxy/proxy_server.py | 40 +-------------- .../proxy_server/test_routes_model_info.py | 50 +++++++++++++++++++ 2 files changed, 51 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f21c907c292..f3b2dbd0b03 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15057,45 +15057,7 @@ def _translate_model_name_for_response(model: dict) -> dict: def _get_proxy_model_info(model: dict) -> dict: - # provided model_info in config.yaml - model_info: Final = model.get("model_info", {}) - - # read litellm model_prices_and_context_window.json to get the following: - # input_cost_per_token, output_cost_per_token, max_tokens - litellm_model_info = get_litellm_model_info(model=model) - - # 2nd pass on the model, try seeing if we can find model in litellm model_cost map - if litellm_model_info == {}: - # use litellm_param model_name to get model_info - litellm_params = model.get("litellm_params", {}) - litellm_model = litellm_params.get("model", None) - try: - litellm_model_info = litellm.get_model_info(model=litellm_model) - except Exception: - litellm_model_info = {} - # 3rd pass on the model, try seeing if we can find model but without the "/" in model cost map - if litellm_model_info == {}: - # use litellm_param model_name to get model_info - litellm_params = model.get("litellm_params", {}) - litellm_model = litellm_params.get("model", None) - split_model: Final = litellm_model.split("/") - if len(split_model) > 0: - litellm_model = split_model[-1] - try: - litellm_model_info = litellm.get_model_info(model=litellm_model, custom_llm_provider=split_model[0]) - except Exception: - litellm_model_info = {} - discovered_model_info: Final = ( - llm_router.get_discovered_model_info(model_info.get("id")) if llm_router is not None else MappingProxyType({}) - ) - for k, v in MappingProxyType({**litellm_model_info, **discovered_model_info}).items(): - if k not in model_info or (model_info[k] is None and k in discovered_model_info): - model_info[k] = v - model["model_info"] = model_info - # don't return the llm credentials - model = remove_sensitive_info_from_deployment(deployment_dict=model, excluded_keys={"litellm_credential_name"}) - - return _translate_model_name_for_response(model) + return _translate_model_name_for_response(_enrich_model_info_with_litellm_data(model=model, llm_router=llm_router)) def _model_info_json_response(data: Sequence[Mapping[str, object]] | Mapping[str, object]) -> Response: diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index 203a237113f..75a8657356a 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -323,6 +323,56 @@ def test_model_info_reports_null_cost_for_unpriced_deployment_and_zero_for_decla assert input_cost > 0 and output_cost > 0 +def test_model_info_id_lookup_reports_the_same_cost_as_the_list( + client: TestClient, + auth_as: Callable[[], AbstractContextManager[object]], + monkeypatch: pytest.MonkeyPatch, +) -> None: + """``GET /model/info?litellm_model_id=`` must agree with the ``GET /model/info`` list, so an + unpriced deployment cannot read as null in the list and as free on the id lookup.""" + monkeypatch.setattr(litellm, "model_cost", copy.deepcopy(litellm.model_cost)) + declared: Final = {"input_cost_per_token": 0.001, "output_cost_per_token": 0.002} + free: Final = {"input_cost_per_token": 0, "output_cost_per_token": 0} + router: Final = litellm.Router( + model_list=[ + { + "model_name": name, + "litellm_params": {"model": f"openai/{name}", "api_key": "x", "api_base": "http://vllm", **costs}, + "model_info": {"id": f"{name}-id"}, + } + for name, costs in (("vllm-unpriced", {}), ("vllm-free", free), ("vllm-priced", declared)) + ] + ) + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "llm_model_list", router.get_model_list()) + monkeypatch.setattr(proxy_server, "user_model", None) + + def costs_of(response: httpx.Response) -> dict[str, tuple[object, object]]: + assert response.status_code == 200, response.text + return { + row["model_info"]["id"]: ( + row["model_info"].get("input_cost_per_token"), + row["model_info"].get("output_cost_per_token"), + ) + for row in response.json()["data"] + } + + with auth_as(): + listed: Final = costs_of(client.get("/model/info")) + by_id: Final = { + model_id: costs_of(client.get("/model/info", params={"litellm_model_id": model_id}))[model_id] + for model_id in listed + } + + assert listed == { + "vllm-unpriced-id": (None, None), + "vllm-free-id": (0, 0), + "vllm-priced-id": (declared["input_cost_per_token"], declared["output_cost_per_token"]), + } + assert by_id == listed + _invalidate_model_cost_lowercase_map() + + def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch): from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth from litellm.proxy.auth import model_checks From d1563e0b552e79eea3b8c1a84bbd7808d467cae0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:09:13 -0700 Subject: [PATCH 180/253] feat: honor eager_input_streaming on Bedrock and Anthropic Claude tools --- litellm/llms/anthropic/chat/transformation.py | 19 +++- litellm/llms/anthropic/common_utils.py | 27 ++++- .../adapters/transformation.py | 11 +- .../bedrock/chat/converse_transformation.py | 12 +- .../anthropic_claude3_transformation.py | 12 +- litellm/llms/bedrock/common_utils.py | 14 ++- .../anthropic_claude3_transformation.py | 9 ++ litellm/proxy/_lazy_openapi_snapshot.json | 10 +- litellm/types/llms/anthropic.py | 3 + litellm/types/llms/openai.py | 2 + .../test_anthropic_chat_transformation.py | 71 ++++++++++++ ...al_pass_through_adapters_transformation.py | 42 +++++++ ...ations_anthropic_claude3_transformation.py | 53 +++++++++ .../chat/test_converse_transformation.py | 106 ++++++++++++++++++ .../test_anthropic_claude3_transformation.py | 62 ++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 + 16 files changed, 444 insertions(+), 13 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 0f99441a115..1f90d375bc2 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -94,6 +94,7 @@ from litellm.utils import ( from ..common_utils import ( AnthropicError, AnthropicModelInfo, + eager_input_streaming_flag, process_anthropic_headers, strip_advisor_blocks_from_messages, ) @@ -732,10 +733,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): input_anthropic_schema: Final = sanitize_input_schema_for_anthropic(_input_schema) - _tool: Final = AnthropicMessagesTool( - name=tool["function"]["name"], - input_schema=input_anthropic_schema, - type="custom", + _eager_input_streaming: Final = eager_input_streaming_flag(tool) + _tool: Final = ( + AnthropicMessagesTool( + name=tool["function"]["name"], + input_schema=input_anthropic_schema, + type="custom", + ) + if _eager_input_streaming is None + else AnthropicMessagesTool( + name=tool["function"]["name"], + input_schema=input_anthropic_schema, + type="custom", + eager_input_streaming=_eager_input_streaming, + ) ) _description: Final = tool["function"].get("description") diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index d35a9372058..2de9ab41d95 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -10,7 +10,7 @@ from types import MappingProxyType from typing import Any, Final, Literal import httpx -from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, StrictBool, TypeAdapter, ValidationError import litellm from litellm.constants import ( @@ -19,6 +19,7 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET, ) +from litellm.exceptions import UnsupportedParamsError from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_file_ids_from_messages, is_encrypted_reasoning_block, @@ -231,6 +232,27 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup return headers, api_key +class _EagerInputStreamingFunction(BaseModel): + eager_input_streaming: StrictBool | None = None + + +class _EagerInputStreamingTool(BaseModel): + eager_input_streaming: StrictBool | None = None + function: _EagerInputStreamingFunction | None = None + + +def eager_input_streaming_flag(tool: object) -> bool | None: + try: + parsed: Final = _EagerInputStreamingTool.model_validate(tool) + except ValidationError as error: + if isinstance(tool, Mapping): + raise UnsupportedParamsError(message="eager_input_streaming must be a boolean") from error + return None + if parsed.eager_input_streaming is not None: + return parsed.eager_input_streaming + return parsed.function.eager_input_streaming if parsed.function is not None else None + + class AnthropicError(BaseLLMException): def __init__( self, @@ -373,6 +395,9 @@ class AnthropicModelInfo(BaseLLMModelInfo): return False + def is_eager_input_streaming_used(self, tools: Sequence[object] | None) -> bool: + return any(eager_input_streaming_flag(tool) is True for tool in tools or ()) + @staticmethod def _supports_sampling_params(model: str) -> bool: """Claude 4.7+ (Opus 4.7/4.8, Fable 5) removed sampling params: the API diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 7eb56ae55d3..d1047c8d86c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -111,6 +111,7 @@ from litellm.litellm_core_utils.reasoning_effort_utils import ( reasoning_effort_from_thinking_budget, ) from litellm.llms.anthropic.common_utils import ( + eager_input_streaming_flag, is_empty_unsigned_thinking_block, normalize_anthropic_tool_use_id, strip_encrypted_reasoning_blocks_from_anthropic_messages, @@ -770,6 +771,7 @@ class LiteLLMAnthropicMessagesAdapter: "cache_control", "strict", "type", + "eager_input_streaming", ] for idx, tool in enumerate(tools): @@ -808,7 +810,14 @@ class LiteLLMAnthropicMessagesAdapter: for k, v in tool.items(): if k not in mapped_tool_params: # pass additional computer kwargs function_chunk.setdefault("parameters", {}).update({k: v}) - tool_param = ChatCompletionToolParam(type="function", function=function_chunk) + eager_input_streaming = eager_input_streaming_flag(tool) + tool_param = ( + ChatCompletionToolParam(type="function", function=function_chunk) + if eager_input_streaming is None + else ChatCompletionToolParam( + type="function", function=function_chunk, eager_input_streaming=eager_input_streaming + ) + ) self._add_cache_control_if_applicable(tool, tool_param, model) new_tools.append(tool_param) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 1a4f27e2ddd..87e28054817 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -48,6 +48,7 @@ from litellm.llms.bedrock.request_metadata import ( merge_bedrock_invoke_headers, resolve_bedrock_request_metadata, ) +from litellm.types.llms.anthropic import ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER from litellm.types.llms.bedrock import * from litellm.types.llms.openai import ( AllMessageValues, @@ -1518,10 +1519,7 @@ class AmazonConverseConfig(BaseConfig): bedrock_tools: list[ToolBlock] = [] # Collect anthropic_beta values from user headers - anthropic_beta_list: Final = [] - if headers: - user_betas: Final = get_anthropic_beta_from_headers(headers) - anthropic_beta_list.extend(user_betas) + anthropic_beta_list: Final = list(get_anthropic_beta_from_headers(headers or {})) # Separate pre-formatted Bedrock tools (e.g. systemTool from web_search_options) # from OpenAI-format tools that need transformation via _bedrock_tools_pt @@ -1634,6 +1632,12 @@ class AmazonConverseConfig(BaseConfig): if ANTHROPIC_EFFORT_BETA_HEADER not in anthropic_beta_list: anthropic_beta_list.append(ANTHROPIC_EFFORT_BETA_HEADER) + if ( + AnthropicModelInfo().is_eager_input_streaming_used(filtered_tools) + and ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER not in anthropic_beta_list + ): + anthropic_beta_list.append(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER) + # Bedrock Converse: compact_20260112 edits only (+ beta header). AmazonConverseConfig._filter_context_management_for_bedrock_converse( additional_request_params, anthropic_beta_list diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index 72bc43ba938..1326dc22ca0 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -24,8 +24,12 @@ from litellm.llms.bedrock.common_utils import ( normalize_custom_field_on_tools, normalize_tool_input_schema_types_for_bedrock_invoke, strip_unsupported_bedrock_invoke_output_config_keys, + tools_without_eager_input_streaming, +) +from litellm.types.llms.anthropic import ( + ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER, + ANTHROPIC_TOOL_SEARCH_BETA_HEADER, ) -from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse @@ -237,6 +241,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): # Hoist `custom.defer_loading` then drop `custom` (Bedrock doesn't support it) normalize_custom_field_on_tools(anthropic_request) normalize_tool_input_schema_types_for_bedrock_invoke(anthropic_request) + outbound_tools: Final = tools_without_eager_input_streaming(anthropic_request) + if outbound_tools is not None: + anthropic_request["tools"] = outbound_tools return anthropic_request def _compute_bedrock_invoke_beta_headers( @@ -269,6 +276,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): if bedrock_supports_tool_search(model): beta_set.add("tool-search-tool-2025-10-19") + if self.is_eager_input_streaming_used(tools): + beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER) + auto_beta_list: Final = filter_and_transform_beta_headers( beta_headers=list(beta_set - user_beta_set), provider="bedrock", diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 62b485588ac..7e24292a87e 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -9,7 +9,7 @@ import functools import json import os import re -from collections.abc import Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict if TYPE_CHECKING: @@ -18,6 +18,7 @@ if TYPE_CHECKING: from litellm.types.llms.bedrock import BedrockCreateBatchRequest import httpx +from pydantic import TypeAdapter, ValidationError import litellm from litellm import verbose_logger @@ -330,6 +331,17 @@ def normalize_custom_field_on_tools(request_body: dict) -> None: tool["defer_loading"] = deferred +_TOOL_DICTS_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...]) + + +def tools_without_eager_input_streaming(request_body: Mapping[str, object]) -> Sequence[object] | None: + try: + tools: Final = _TOOL_DICTS_ADAPTER.validate_python(request_body.get("tools")) + except ValidationError: + return None + return [{key: value for key, value in tool.items() if key != "eager_input_streaming"} for tool in tools] + + def normalize_json_schema_custom_types_to_object(schema: dict) -> None: """ In-place: replace JSON Schema ``type: \"custom\"`` with ``\"object\"`` (iterative walk). diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 4aa2afdbc78..d2be1ad9156 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -39,6 +39,7 @@ from litellm.llms.bedrock.common_utils import ( normalize_custom_field_on_tools, normalize_tool_input_schema_types_for_bedrock_invoke, strip_unsupported_bedrock_invoke_output_config_keys, + tools_without_eager_input_streaming, ) from litellm.llms.bedrock.request_metadata import ( bedrock_request_metadata_headers, @@ -46,6 +47,7 @@ from litellm.llms.bedrock.request_metadata import ( ) from litellm.types.llms.anthropic import ( ANTHROPIC_BETA_HEADER_VALUES, + ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER, ANTHROPIC_TOOL_SEARCH_BETA_HEADER, ) from litellm.types.llms.bedrock import BedrockInvokeAnthropicMessagesRequest @@ -525,6 +527,9 @@ class AmazonAnthropicClaudeMessagesConfig( if injected_thinking_for_clear_thinking: beta_set.add("interleaved-thinking-2025-05-14") + if anthropic_model_info.is_eager_input_streaming_used(tools): + beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER) + self._filter_context_management_for_bedrock_invoke( anthropic_messages_request=anthropic_messages_request, beta_set=beta_set, @@ -719,6 +724,10 @@ class AmazonAnthropicClaudeMessagesConfig( if filtered_betas: anthropic_messages_request["anthropic_beta"] = filtered_betas + outbound_tools: Final = tools_without_eager_input_streaming(anthropic_messages_request) + if outbound_tools is not None: + anthropic_messages_request["tools"] = outbound_tools + remaining_output_config: Final = anthropic_messages_request.get("output_config") if ( litellm.drop_params is True diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 5270df7b468..64daa695316 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19370,7 +19370,7 @@ } } }, - "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" }, "500": { "content": { @@ -32899,6 +32899,10 @@ "cache_control": { "$ref": "#/components/schemas/ChatCompletionCachedContent" }, + "eager_input_streaming": { + "title": "Eager Input Streaming", + "type": "boolean" + }, "function": { "$ref": "#/components/schemas/ChatCompletionToolParamFunctionChunk" }, @@ -32928,6 +32932,10 @@ "title": "Description", "type": "string" }, + "eager_input_streaming": { + "title": "Eager Input Streaming", + "type": "boolean" + }, "name": { "title": "Name", "type": "string" diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index bcdee86360e..bcd24695f25 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -56,6 +56,7 @@ class AnthropicMessagesTool(TypedDict, total=False): defer_loading: bool allowed_callers: list[str] | None input_examples: list[dict[str, Any]] | None + eager_input_streaming: ReadOnly[bool] class AnthropicComputerTool(TypedDict, total=False): @@ -755,6 +756,8 @@ ANTHROPIC_TOOL_SEARCH_BETA_HEADER: Final = "advanced-tool-use-2025-11-20" # Effort beta header constant ANTHROPIC_EFFORT_BETA_HEADER: Final = "effort-2025-11-24" +ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER: Final = "fine-grained-tool-streaming-2025-05-14" + # OAuth constants ANTHROPIC_OAUTH_TOKEN_PREFIX: Final = "sk-ant-oat" ANTHROPIC_OAUTH_BETA_HEADER: Final = "oauth-2025-04-20" diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index e3eac9b9205..abb76527e82 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -992,6 +992,7 @@ class ChatCompletionToolParamFunctionChunk(TypedDict, total=False): description: str parameters: dict strict: bool + eager_input_streaming: ReadOnly[bool] class OpenAIChatCompletionToolParam(TypedDict): @@ -1002,6 +1003,7 @@ class OpenAIChatCompletionToolParam(TypedDict): class ChatCompletionToolParam(OpenAIChatCompletionToolParam, total=False): cache_control: ChatCompletionCachedContent allowed_callers: list[str] + eager_input_streaming: ReadOnly[bool] class Function(TypedDict, total=False): 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 269c351f866..1359bf40dcd 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 @@ -6370,3 +6370,74 @@ def test_response_format_tool_path_skips_forced_tool_choice_when_unsupported(loc assert "tools" in result assert "tool_choice" not in result + + +def _eager_chat_tool(**extra): + return { + "type": "function", + "function": { + "name": "write_file", + "description": "Write a file", + "parameters": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}, + }, + **extra, + } + + +@pytest.mark.parametrize("flag", [True, False]) +def test_eager_input_streaming_passed_through_from_tool_top_level(flag): + mapped_tool, _ = AnthropicConfig()._map_tool_helper(_eager_chat_tool(eager_input_streaming=flag)) + + assert mapped_tool == { + "name": "write_file", + "description": "Write a file", + "input_schema": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}, + "type": "custom", + "eager_input_streaming": flag, + } + + +def test_eager_input_streaming_passed_through_from_function(): + tool = _eager_chat_tool() + tool["function"]["eager_input_streaming"] = True + + mapped_tool, _ = AnthropicConfig()._map_tool_helper(tool) + + assert mapped_tool["eager_input_streaming"] is True + assert "eager_input_streaming" not in mapped_tool["input_schema"] + + +def test_eager_input_streaming_absent_stays_absent(): + mapped_tool, _ = AnthropicConfig()._map_tool_helper(_eager_chat_tool()) + + assert "eager_input_streaming" not in mapped_tool + + +def test_eager_input_streaming_rejects_non_boolean(): + with pytest.raises(litellm.BadRequestError, match="eager_input_streaming must be a boolean"): + AnthropicConfig()._map_tool_helper(_eager_chat_tool(eager_input_streaming="true")) + + +def test_eager_input_streaming_not_set_on_computer_use_tool(): + computer_tool = { + "type": "computer_20250124", + "function": {"name": "computer", "parameters": {"display_width_px": 1024, "display_height_px": 768}}, + "eager_input_streaming": True, + } + + mapped_tool, _ = AnthropicConfig()._map_tool_helper(computer_tool) + + assert mapped_tool["type"] == "computer_20250124" + assert "eager_input_streaming" not in mapped_tool + + +def test_eager_input_streaming_reaches_anthropic_request_tools(): + result = AnthropicConfig().map_openai_params( + non_default_params={"tools": [_eager_chat_tool(eager_input_streaming=True)], "stream": True}, + optional_params={}, + model="claude-sonnet-5", + drop_params=False, + ) + + assert result["tools"][0]["eager_input_streaming"] is True + assert result["tools"][0]["name"] == "write_file" 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 e6782b70d3e..2318ceb937a 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 @@ -5017,3 +5017,45 @@ def test_redacted_thinking_blocks_never_carry_cache_control(): replayed: Final = outbound["messages"][1]["content"][0] assert replayed["type"] == "redacted_thinking" assert "cache_control" not in replayed + + +EAGER_INPUT_SCHEMA: Final = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} + + +@pytest.mark.parametrize("flag", [True, False]) +def test_translate_anthropic_tools_to_openai_carries_eager_input_streaming_onto_tool(flag): + """The per-tool flag lands on the OpenAI tool object, never inside the JSON schema Bedrock sends as inputSchema.""" + tools: Final = [{"name": "write_file", "input_schema": EAGER_INPUT_SCHEMA, "eager_input_streaming": flag}] + + new_tools, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_tools_to_openai(tools=tools) + + assert new_tools[0]["eager_input_streaming"] is flag + assert new_tools[0]["function"]["parameters"] == EAGER_INPUT_SCHEMA + assert "eager_input_streaming" not in new_tools[0]["function"] + + +def test_translate_anthropic_tools_to_openai_omits_unset_eager_input_streaming(): + tools: Final = [{"name": "write_file", "input_schema": EAGER_INPUT_SCHEMA}] + + new_tools, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_tools_to_openai(tools=tools) + + assert "eager_input_streaming" not in new_tools[0] + assert "eager_input_streaming" not in new_tools[0]["function"]["parameters"] + + +def test_eager_input_streaming_tool_reaches_bedrock_converse_as_beta(): + """An Anthropic Messages request routed to bedrock/converse/ turns the flag into the fine-grained streaming beta.""" + from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + + tools: Final = [{"name": "write_file", "input_schema": EAGER_INPUT_SCHEMA, "eager_input_streaming": True}] + new_tools, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_tools_to_openai(tools=tools) + + data: Final = AmazonConverseConfig()._transform_request_helper( + model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", + system_content_blocks=[], + optional_params={"tools": new_tools}, + messages=[{"role": "user", "content": "write a big file"}], + ) + + assert data["additionalModelRequestFields"]["anthropic_beta"] == ["fine-grained-tool-streaming-2025-05-14"] + assert data["toolConfig"]["tools"][0]["toolSpec"]["inputSchema"]["json"] == EAGER_INPUT_SCHEMA diff --git a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py index ec0bf6b842a..1395489334e 100644 --- a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py @@ -859,3 +859,56 @@ def test_bedrock_chat_invoke_tool_search_beta_follows_model_map( ) assert result.get("anthropic_beta") == expected_betas + + +FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" +EAGER_TOOL_SCHEMA = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} + + +def _chat_invoke_request_with_tools(tools, headers=None): + config = AmazonAnthropicClaudeConfig() + model = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" + optional_params = config.map_openai_params( + non_default_params={"max_tokens": 64, "stream": True, "tools": tools}, + optional_params={}, + model=model, + drop_params=False, + ) + return config.transform_request( + model=model, + messages=[{"role": "user", "content": "write a big file"}], + optional_params=optional_params, + litellm_params={}, + headers=headers or {}, + ) + + +def _eager_openai_tool(name, **extra): + return {"type": "function", "function": {"name": name, "parameters": EAGER_TOOL_SCHEMA}, **extra} + + +def test_bedrock_chat_invoke_eager_input_streaming_tool_adds_beta_and_strips_key(): + result = _chat_invoke_request_with_tools( + [_eager_openai_tool("write_file", eager_input_streaming=True), _eager_openai_tool("read_file")] + ) + + assert result["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] + assert [tool["name"] for tool in result["tools"]] == ["write_file", "read_file"] + assert all("eager_input_streaming" not in tool for tool in result["tools"]) + assert result["tools"][0]["input_schema"] == EAGER_TOOL_SCHEMA + + +def test_bedrock_chat_invoke_eager_input_streaming_false_strips_key_without_beta(): + result = _chat_invoke_request_with_tools([_eager_openai_tool("write_file", eager_input_streaming=False)]) + + assert "anthropic_beta" not in result + assert "eager_input_streaming" not in result["tools"][0] + + +def test_bedrock_chat_invoke_eager_input_streaming_beta_not_duplicated_with_client_header(): + result = _chat_invoke_request_with_tools( + [_eager_openai_tool("write_file", eager_input_streaming=True)], + headers={"anthropic-beta": FINE_GRAINED_TOOL_STREAMING_BETA}, + ) + + assert result["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] 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 96c78c1cf75..650427f2ae3 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -7400,3 +7400,109 @@ def test_transform_response_honors_json_mode_kwarg_when_optional_params_lack_it( ) assert result.choices[0].message.tool_calls is None assert json.loads(result.choices[0].message.content) == {"city": "Paris", "population": 2100000} + + +FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" +EAGER_TOOL_SCHEMA = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} + + +def _eager_openai_tool(**extra): + return {"type": "function", "function": {"name": "write_file", "parameters": EAGER_TOOL_SCHEMA}, **extra} + + +def _eager_openai_function_tool(**extra): + return {"type": "function", "function": {"name": "write_file", "parameters": EAGER_TOOL_SCHEMA, **extra}} + + +def _eager_anthropic_tool(**extra): + return {"name": "write_file", "input_schema": EAGER_TOOL_SCHEMA, **extra} + + +def _converse_request(model, tools, headers=None): + return AmazonConverseConfig()._transform_request_helper( + model=model, + system_content_blocks=[], + optional_params={"tools": tools}, + messages=[{"role": "user", "content": "write a big file"}], + headers=headers, + ) + + +@pytest.mark.parametrize( + "tool", + [ + _eager_openai_tool(eager_input_streaming=True), + _eager_openai_function_tool(eager_input_streaming=True), + _eager_anthropic_tool(eager_input_streaming=True), + ], + ids=["openai_top_level", "openai_under_function", "anthropic_shape"], +) +def test_eager_input_streaming_tool_adds_fine_grained_tool_streaming_beta(tool): + data = _converse_request("us.anthropic.claude-sonnet-4-5-20250929-v1:0", [tool]) + + assert data["additionalModelRequestFields"]["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] + tool_spec = data["toolConfig"]["tools"][0]["toolSpec"] + assert tool_spec["name"] == "write_file" + assert "eager_input_streaming" not in tool_spec + assert "eager_input_streaming" not in tool_spec["inputSchema"]["json"] + + +@pytest.mark.parametrize( + "tool", + [ + _eager_openai_tool(eager_input_streaming=False), + _eager_openai_function_tool(eager_input_streaming=False), + _eager_anthropic_tool(eager_input_streaming=False), + _eager_openai_tool(), + ], + ids=["openai_false", "function_false", "anthropic_false", "absent"], +) +def test_eager_input_streaming_false_or_absent_adds_no_beta(tool): + data = _converse_request("us.anthropic.claude-sonnet-4-5-20250929-v1:0", [tool]) + + assert "anthropic_beta" not in data.get("additionalModelRequestFields", {}) + assert "eager_input_streaming" not in data["toolConfig"]["tools"][0]["toolSpec"] + + +def test_eager_input_streaming_beta_only_on_anthropic_models(): + data = _converse_request("amazon.nova-pro-v1:0", [_eager_openai_tool(eager_input_streaming=True)]) + + assert "anthropic_beta" not in data.get("additionalModelRequestFields", {}) + assert data["toolConfig"]["tools"][0]["toolSpec"]["name"] == "write_file" + + +def test_eager_input_streaming_beta_not_duplicated_with_client_header(): + data = _converse_request( + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + [_eager_openai_tool(eager_input_streaming=True)], + headers={"anthropic-beta": f"{FINE_GRAINED_TOOL_STREAMING_BETA},interleaved-thinking-2025-05-14"}, + ) + + assert data["additionalModelRequestFields"]["anthropic_beta"] == [ + FINE_GRAINED_TOOL_STREAMING_BETA, + "interleaved-thinking-2025-05-14", + ] + + +def test_eager_input_streaming_beta_never_written_back_into_client_header_list(): + headers = {"anthropic-beta": ["interleaved-thinking-2025-05-14"]} + + data = _converse_request( + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + [_eager_openai_tool(eager_input_streaming=True)], + headers=headers, + ) + + assert data["additionalModelRequestFields"]["anthropic_beta"] == [ + "interleaved-thinking-2025-05-14", + FINE_GRAINED_TOOL_STREAMING_BETA, + ] + assert headers == {"anthropic-beta": ["interleaved-thinking-2025-05-14"]} + + +def test_eager_input_streaming_non_boolean_is_a_bad_request(): + with pytest.raises(litellm.BadRequestError, match="eager_input_streaming must be a boolean"): + _converse_request( + "us.anthropic.claude-sonnet-4-5-20250929-v1:0", + [_eager_openai_tool(eager_input_streaming="true")], + ) diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index e43accdb835..b71e8e521f2 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -3244,3 +3244,65 @@ def test_bedrock_messages_strips_effort_but_keeps_format_for_sonnet_4_5(local_mo ) assert result.get("output_config") == {"format": schema_format} + + +FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" + + +def _invoke_request_with_tools(tools, headers=None): + from litellm.types.router import GenericLiteLLMParams + + return AmazonAnthropicClaudeMessagesConfig().transform_anthropic_messages_request( + model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=[{"role": "user", "content": "write a big file"}], + anthropic_messages_optional_request_params={"max_tokens": 4096, "tools": copy.deepcopy(tools), "stream": True}, + litellm_params=GenericLiteLLMParams(), + headers=headers or {}, + ) + + +def _eager_invoke_tool(name, eager_input_streaming): + return { + "name": name, + "description": f"{name} tool", + "input_schema": {"type": "object", "properties": {"path": {"type": "string"}}}, + "eager_input_streaming": eager_input_streaming, + } + + +def test_bedrock_invoke_eager_input_streaming_tool_adds_beta_and_strips_key(): + result = _invoke_request_with_tools( + [ + _eager_invoke_tool("write_file", True), + _eager_invoke_tool("read_file", False), + {"name": "list_files", "input_schema": {"type": "object", "properties": {}}}, + ] + ) + + assert result["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] + assert [tool["name"] for tool in result["tools"]] == ["write_file", "read_file", "list_files"] + assert all("eager_input_streaming" not in tool for tool in result["tools"]) + assert result["tools"][0]["description"] == "write_file tool" + assert result["tools"][0]["input_schema"] == {"type": "object", "properties": {"path": {"type": "string"}}} + + +def test_bedrock_invoke_eager_input_streaming_false_strips_key_without_beta(): + result = _invoke_request_with_tools([_eager_invoke_tool("write_file", False)]) + + assert "anthropic_beta" not in result + assert result["tools"] == [ + { + "name": "write_file", + "description": "write_file tool", + "input_schema": {"type": "object", "properties": {"path": {"type": "string"}}}, + } + ] + + +def test_bedrock_invoke_eager_input_streaming_beta_not_duplicated_with_client_header(): + result = _invoke_request_with_tools( + [_eager_invoke_tool("write_file", True)], + headers={"anthropic-beta": FINE_GRAINED_TOOL_STREAMING_BETA}, + ) + + assert result["anthropic_beta"] == [FINE_GRAINED_TOOL_STREAMING_BETA] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b1b0081cb9c..cc5169dd84c 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -25970,6 +25970,8 @@ export interface components { /** Allowed Callers */ allowed_callers?: string[]; cache_control?: components["schemas"]["ChatCompletionCachedContent"]; + /** Eager Input Streaming */ + eager_input_streaming?: boolean; function: components["schemas"]["ChatCompletionToolParamFunctionChunk"]; /** Type */ type: "function" | string; @@ -25978,6 +25980,8 @@ export interface components { ChatCompletionToolParamFunctionChunk: { /** Description */ description?: string; + /** Eager Input Streaming */ + eager_input_streaming?: boolean; /** Name */ name: string; /** Parameters */ From 04193745a65cc18a36820be8ef1d901d6cbe19db Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 20:10:20 +0000 Subject: [PATCH 181/253] fix(docs-test): restore line-anchored table regex in router settings check A merge on this branch dropped the ^ anchor and MULTILINE flag from the doc_key_pattern, so the unanchored match started swallowing key names into captured fields whenever a row's description cell itself contains a pipe (Literal unions and similar). The check then reported 27 long-documented keys as undocumented. Restores the exact pattern used on main, verified against a live litellm-docs checkout: all 58 Router init params resolve as documented. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/documentation_tests/test_router_settings.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tests/documentation_tests/test_router_settings.py b/tests/documentation_tests/test_router_settings.py index 75032f80dfa..7e3d0c07459 100644 --- a/tests/documentation_tests/test_router_settings.py +++ b/tests/documentation_tests/test_router_settings.py @@ -51,9 +51,7 @@ try: if general_settings_section: # Extract the table rows, which contain the documented keys table_content = general_settings_section.group(1) - doc_key_pattern = re.compile( - r"\|\s*([^\|]+?)\s*\|" - ) # Capture the key from each row of the table + doc_key_pattern = re.compile(r"^\|\s*([^\|]+?)\s*\|", re.MULTILINE) documented_keys.update(doc_key_pattern.findall(table_content)) except Exception as e: raise Exception( From 1581c45615ccd841560b5123b3843f7b864c6c6f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:14:22 -0700 Subject: [PATCH 182/253] fix(proxy): check and persist only the settings the request sent POST /config/update compared litellm_settings after lowercasing the callback list, so a config file spelling a callback in mixed case refused the same list sent back, and it stored every general_settings model default next to the keys the request set. Both now use the request as sent; only the stored callback list is lowercased. Also drops config_data from the router settings reload callers the previous commit left behind and teaches the legacy MockProxyConfig the ownership check. --- litellm/proxy/proxy_server.py | 5 +-- tests/proxy_unit_tests/test_proxy_server.py | 3 ++ .../proxy/proxy_server/test_proxy_config.py | 5 +-- .../proxy/proxy_server/test_routes_config.py | 42 +++++++++++++++++++ 4 files changed, 49 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4654dffcb11..ddee8337609 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -16977,7 +16977,7 @@ async def update_config( proxy_config.reject_config_owned_writes( section_name="general_settings", changed_keys=requested_general_settings ) - proxy_config.reject_config_owned_writes(section_name="litellm_settings", changed_keys=updated_litellm_settings) + proxy_config.reject_config_owned_writes(section_name="litellm_settings", changed_keys=raw_litellm_settings) proxy_config.reject_config_owned_writes(section_name="router_settings", changed_keys=router_settings_updates) async def _read_section(param_name: str) -> dict: @@ -17006,8 +17006,7 @@ async def update_config( if config_info.general_settings is not None: existing = await _read_section("general_settings") before_general_settings: Final = copy.deepcopy(existing) - updates: Mapping[str, JsonValue] = config_info.general_settings.dict(exclude_none=True) - for k, v in updates.items(): + for k, v in requested_general_settings.items(): if k == "alert_to_webhook_url": if "alerting" not in existing: existing["alerting"] = ["slack"] diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 1fcdaa67143..8d1daa8a3d5 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -3069,6 +3069,9 @@ async def test_update_config_success_callback_normalization(): async def add_deployment(self, prisma_client=None, proxy_logging_obj=None): # noqa: F811 # pytest fixture, not a redefinition return None + def reject_config_owned_writes(self, *, section_name, changed_keys): + return None + setattr(proxy_server, "proxy_config", MockProxyConfig()) config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]}) diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 42cd6e4ed78..1e9a93bd019 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -3399,9 +3399,8 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): fake_prisma.db.litellm_config.find_first = AsyncMock( return_value=SimpleNamespace(param_value={"timeout": 30, "retries": 2, "fallbacks": []}) ) - config_data = {"router_settings": {"timeout": 10}} + pc.router_settings.load_yaml({"timeout": 10}) await pc._add_router_settings_from_db_config( - config_data=config_data, llm_router=fake_router, prisma_client=fake_prisma, ) @@ -3421,7 +3420,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): async def test_ProxyConfig__add_router_settings_from_db_config_none_router_noop(): pc = ProxyConfig() # No router and no prisma — should silently return. - await pc._add_router_settings_from_db_config(config_data={}, llm_router=None, prisma_client=None) + await pc._add_router_settings_from_db_config(llm_router=None, prisma_client=None) # Error-style: bad call signature raises. with pytest.raises(TypeError): await pc._add_router_settings_from_db_config() # type: ignore[call-arg] diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index 0a2daf64463..e4d80c86f57 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -176,6 +176,48 @@ def test_config_update_rejects_config_owned_keys_and_accepts_the_same_value( assert persisted[next(iter(yaml_values))] == yaml_values[next(iter(yaml_values))] +def test_config_update_persists_only_the_general_settings_keys_the_request_set( + client, auth_as, mock_prisma, monkeypatch +): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock()) + ps.proxy_config.settings.load_yaml({"health_check_interval": 60}) + try: + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/config/update", json={"general_settings": {"alerting_threshold": 600}}) + finally: + ps.proxy_config.settings.load_yaml({}) + + assert response.status_code == 200 + persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"]) + assert persisted == {"alerting_threshold": 600} + + +def test_config_update_accepts_a_config_owned_success_callback_the_file_spells_in_mixed_case( + client, auth_as, mock_prisma, monkeypatch +): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock()) + ps.proxy_config.litellm_settings.load_yaml({"success_callback": ["Langfuse"]}) + try: + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post("/config/update", json={"litellm_settings": {"success_callback": ["Langfuse"]}}) + finally: + ps.proxy_config.litellm_settings.load_yaml({}) + + assert response.status_code == 200 + persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"]) + assert persisted["success_callback"] == ["langfuse"] + + def test_config_update_rejects_assistants_config(client, auth_as, mock_prisma, monkeypatch): from litellm.proxy import proxy_server as ps from litellm.proxy._types import LitellmUserRoles From d16d17d2d28f56cb0775a8c42ebc95b8ccf4f032 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:15:17 -0700 Subject: [PATCH 183/253] ci(osv-scan): scan the VS Code extension lockfile --- .github/workflows/osv-scan.yml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.github/workflows/osv-scan.yml b/.github/workflows/osv-scan.yml index 31104002dab..0aedeaec12a 100644 --- a/.github/workflows/osv-scan.yml +++ b/.github/workflows/osv-scan.yml @@ -41,4 +41,5 @@ jobs: "$RUNNER_TEMP/osv-scanner" scan source \ --config osv-scanner.toml \ -L uv.lock \ - -L ui/litellm-dashboard/package-lock.json + -L ui/litellm-dashboard/package-lock.json \ + -L vscode-extension/package-lock.json From 817aefc41348aa9f47b246b554afc5efbf6e2523 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 20:15:48 +0000 Subject: [PATCH 184/253] test(proxy): keep stored pass-through route tests off the shared FastAPI app so the allowlist coverage test stays order independent Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/proxy/test_proxy_server.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index bc7ae556e7b..5246463f616 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7484,7 +7484,8 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi settings: Final = patch("litellm.proxy.proxy_server.general_settings", {}) # test-quality-ok: the method reads this module global; no injection seam yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", None) # test-quality-ok: module global holding the YAML endpoints; this case has none - with settings, yaml_endpoints: + app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker + with settings, yaml_endpoints, app_routes: pc = ProxyConfig() await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]}) assert live_routes(), "the stored endpoint should be serving before the row is deleted" @@ -7517,7 +7518,8 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]) # test-quality-ok: module global holding the YAML endpoints the reload merges in - with settings, yaml_endpoints: + app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker + with settings, yaml_endpoints, app_routes: await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint]) assert live_paths() == {config_path} From c54c0b049daf75fd761723d388e81b9b9cbb8f9a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:16:56 -0700 Subject: [PATCH 185/253] fix(vertex_ai): keep reasoning_effort unsupported on Vertex AI Mistral partner models Vertex AI Mistral models reused MistralConfig, whose reasoning_effort advertisement checks the mistral provider entry of the cost map, so vertex_ai/mistral-medium-3 started advertising reasoning_effort and drop_params stopped dropping it, turning a 200 into a Vertex 400. VertexAIMistralConfig scopes that lookup to the vertex_ai provider, and MistralConfig now reads the provider from its custom_llm_provider property instead of a hardcoded "mistral". --- litellm/__init__.py | 3 +++ litellm/_lazy_imports_registry.py | 5 +++++ .../get_supported_openai_params.py | 2 +- litellm/llms/mistral/chat/transformation.py | 21 ++++++++++++++----- .../mistral/transformation.py | 7 +++++++ litellm/utils.py | 4 ++-- ...i_partner_models_mistral_transformation.py | 17 +++++++++++++++ 7 files changed, 51 insertions(+), 8 deletions(-) create mode 100644 litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/transformation.py create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/test_vertex_ai_partner_models_mistral_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 71857877e53..2e7a2a4efd8 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1704,6 +1704,9 @@ if TYPE_CHECKING: from .llms.vertex_ai.vertex_ai_partner_models.ai21.transformation import ( VertexAIAi21Config as VertexAIAi21Config, ) + from .llms.vertex_ai.vertex_ai_partner_models.mistral.transformation import ( + VertexAIMistralConfig as VertexAIMistralConfig, + ) from .llms.bedrock.chat.invoke_handler import ( AmazonCohereChatConfig as AmazonCohereChatConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 4c478b51ed1..9cfcb9e41f7 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -184,6 +184,7 @@ LLM_CONFIG_NAMES: Final = ( "VertexAIAnthropicConfig", "VertexAILlama3Config", "VertexAIAi21Config", + "VertexAIMistralConfig", "AmazonCohereChatConfig", "AmazonBedrockGlobalConfig", "AmazonAI21Config", @@ -771,6 +772,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = { ".llms.vertex_ai.vertex_ai_partner_models.ai21.transformation", "VertexAIAi21Config", ), + "VertexAIMistralConfig": ( + ".llms.vertex_ai.vertex_ai_partner_models.mistral.transformation", + "VertexAIMistralConfig", + ), "AmazonCohereChatConfig": ( ".llms.bedrock.chat.invoke_handler", "AmazonCohereChatConfig", diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 915a03025d9..08b8816e17d 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -190,7 +190,7 @@ def get_supported_openai_params( elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta": if request_type == "chat_completion": if model.startswith("mistral"): - return litellm.MistralConfig().get_supported_openai_params(model=model) + return litellm.VertexAIMistralConfig().get_supported_openai_params(model=model) elif model.startswith("codestral"): return litellm.CodestralTextCompletionConfig().get_supported_openai_params(model=model) elif model.startswith("claude"): diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index 8cbe4291054..f77e828b59a 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -35,14 +35,19 @@ if TYPE_CHECKING: import tiktoken -def _accepted_reasoning_effort(model: str, requested: str) -> str: - declared: Final = declared_reasoning_efforts_for_model(model, "mistral") +def _accepted_reasoning_effort(model: str, requested: str, custom_llm_provider: str) -> str: + declared: Final = declared_reasoning_efforts_for_model(model, custom_llm_provider) if declared is None: return requested accepted: Final = nearest_declared_reasoning_effort(requested, declared) if accepted != requested: verbose_logger.debug( - "mistral: %s takes reasoning_effort %s, sending %s in place of %s", model, declared, accepted, requested + "%s: %s takes reasoning_effort %s, sending %s in place of %s", + custom_llm_provider, + model, + declared, + accepted, + requested, ) return accepted @@ -103,9 +108,15 @@ class MistralConfig(OpenAIGPTConfig): def get_config(cls): return super().get_config() + @property + def custom_llm_provider(self) -> str: + return "mistral" + def get_supported_openai_params(self, model: str) -> list[str]: is_magistral: Final = "magistral" in model.lower() - accepts_reasoning_effort: Final = is_magistral or supports_reasoning(model=model, custom_llm_provider="mistral") + accepts_reasoning_effort: Final = is_magistral or supports_reasoning( + model=model, custom_llm_provider=self.custom_llm_provider + ) return [ "stream", "temperature", @@ -187,7 +198,7 @@ class MistralConfig(OpenAIGPTConfig): if param == "response_format": optional_params["response_format"] = value if param == "reasoning_effort" and "magistral" not in model.lower(): - optional_params["reasoning_effort"] = _accepted_reasoning_effort(model, value) + optional_params["reasoning_effort"] = _accepted_reasoning_effort(model, value, self.custom_llm_provider) if param in ("reasoning_effort", "thinking") and "magistral" in model.lower(): # Flag that we need to add reasoning system prompt optional_params["_add_reasoning_prompt"] = True diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/transformation.py new file mode 100644 index 00000000000..18321c8768a --- /dev/null +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/transformation.py @@ -0,0 +1,7 @@ +from litellm.llms.mistral.chat.transformation import MistralConfig + + +class VertexAIMistralConfig(MistralConfig): + @property + def custom_llm_provider(self) -> str: + return "vertex_ai" diff --git a/litellm/utils.py b/litellm/utils.py index 2c9200fbad7..56985a11b8c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4510,7 +4510,7 @@ def get_optional_params( drop_params=bool(drop_params), ) else: - optional_params = litellm.MistralConfig().map_openai_params( + optional_params = litellm.VertexAIMistralConfig().map_openai_params( model=model, non_default_params=non_default_params, optional_params=optional_params, @@ -8380,7 +8380,7 @@ class ProviderConfigManager: elif model in litellm.vertex_mistral_models: if "codestral" in model: return litellm.CodestralTextCompletionConfig() - return litellm.MistralConfig() + return litellm.VertexAIMistralConfig() elif model in litellm.vertex_ai_ai21_models: return litellm.VertexAIAi21Config() else: diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/test_vertex_ai_partner_models_mistral_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/test_vertex_ai_partner_models_mistral_transformation.py new file mode 100644 index 00000000000..f7df4507651 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/mistral/test_vertex_ai_partner_models_mistral_transformation.py @@ -0,0 +1,17 @@ +import litellm + + +def test_reasoning_effort_stays_unsupported_on_vertex_partner_models(local_model_cost_map): + assert "reasoning_effort" in litellm.get_supported_openai_params( + model="mistral-medium-3", custom_llm_provider="mistral" + ) + assert "reasoning_effort" not in litellm.get_supported_openai_params( + model="mistral-medium-3", custom_llm_provider="vertex_ai" + ) + dropped = litellm.get_optional_params( + model="mistral-medium-3", + custom_llm_provider="vertex_ai", + reasoning_effort="high", + drop_params=True, + ) + assert "reasoning_effort" not in dropped From 0f59e61f24dc72c6f4a9af15f1618d1b13591cde Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 20:17:11 +0000 Subject: [PATCH 186/253] fix(rust): align exception mapping tests with Python Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/exception_mapping_utils/mod.rs | 25 +++++++++---------- .../src/exception_mapping_utils/rules.rs | 4 +-- 2 files changed, 14 insertions(+), 15 deletions(-) diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs index 8892fba1399..acc7820c770 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs @@ -282,7 +282,6 @@ fn extra_information(context: &ExceptionContext, api_base: Option<&str>) -> Stri let lines = [ Some(format!("\nModel: {}", context.model)), api_base.map(|api_base| format!("\nAPI Base: `{api_base}`")), - (!context.redact_messages_in_exceptions).then(|| "\nMessages: `None`".to_string()), context .model_group .as_ref() @@ -314,7 +313,7 @@ fn extra_information(context: &ExceptionContext, api_base: Option<&str>) -> Stri mod testing { use super::*; - pub(super) const DEBUG: &str = "\nModel: ocr-model\nMessages: `None`"; + pub(super) const DEBUG: &str = "\nModel: ocr-model"; pub(super) fn context(provider: &str, family: ExceptionFamily) -> ExceptionContext { ExceptionContext { @@ -369,14 +368,6 @@ mod testing { ..failure } } - - /// What `exception_type` returns for an `http` original the rule mapped to `failure`. - pub(super) fn with_headers(failure: PublicFailure) -> PublicFailure { - PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..failure - } - } } #[cfg(test)] @@ -492,7 +483,17 @@ mod tests { #[case] message: &str, ) { let context = context("reducto", ExceptionFamily::Other); - let expected = failure(PublicKind::ApiConnection, message, "reducto"); + let expected = match original { + OriginalException::Http { .. } => with_debug(failure( + status(StatusClass::BadRequest, upstream(409, "rejected")), + message, + "reducto", + )), + OriginalException::Connection { .. } | OriginalException::Local { .. } => { + failure(PublicKind::ApiConnection, message, "reducto") + } + _ => unreachable!(), + }; let actual = exception_type(&context, &original); assert_eq!( PublicFailure { @@ -669,7 +670,6 @@ mod tests { "\n\nKey Name: `key`\nTeam: `None`", "\nModel: ocr-model", "\nAPI Base: `region-aiplatform.googleapis.com/v1/projects/project/locations/region/publishers/google/models/ocr-model:generateContent`", - "\nMessages: `None`", "\nmodel_group: `ocr`\n", "\ndeployment: `deployment`\n", "\nvertex_project: `project`\n", @@ -681,7 +681,6 @@ mod tests { #[rstest::rstest] #[case::bare(ExceptionContext::default(), "\nModel: ")] #[case::redacted_messages(ExceptionContext { redact_messages_in_exceptions: true, model: "m".into(), ..ExceptionContext::default() }, "\nModel: m")] - #[case::messages(ExceptionContext { model: "m".into(), ..ExceptionContext::default() }, "\nModel: m\nMessages: `None`")] #[case::team_alias( ExceptionContext { model: "m".into(), user_api_key_alias: Some("key".into()), user_api_key_team_alias: Some("team".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, "\n\nKey Name: `key`\nTeam: `team`\nModel: m" diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs index 5df8ad991b5..34cf2c390c3 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs @@ -224,7 +224,7 @@ mod tests { "mistral", ))] #[case::later_rule_when_the_earlier_does_not_apply("second", PublicFailure { - litellm_debug_info: Some("\nModel: ocr-model\nMessages: `None`".into()), + litellm_debug_info: Some("\nModel: ocr-model".into()), ..failure(PublicKind::ApiConnection, "seen second", "mistral") })] fn the_first_applicable_rule_decides(#[case] body: &str, #[case] expected: PublicFailure) { @@ -324,7 +324,7 @@ mod tests { } #[rstest::rstest] - #[case::with_debug(true, Some("\nModel: ocr-model\nMessages: `None`"))] + #[case::with_debug(true, Some("\nModel: ocr-model"))] #[case::without_debug(false, None)] fn debug_rules_carry_the_extra_information( #[case] debug: bool, From 6464c1fd6c9a58efcf75f4b73650d541aa4c76bd Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 19:17:08 +0000 Subject: [PATCH 187/253] fix(proxy): make SettingsStore.clear() terminate when the config file owns a key Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/config_resolvers/settings_store.py | 4 ++ .../config_resolvers/test_settings_store.py | 49 +++++++++++++++++++ 2 files changed, 53 insertions(+) diff --git a/litellm/proxy/config_resolvers/settings_store.py b/litellm/proxy/config_resolvers/settings_store.py index 079d262319f..49f5a833b8f 100644 --- a/litellm/proxy/config_resolvers/settings_store.py +++ b/litellm/proxy/config_resolvers/settings_store.py @@ -87,6 +87,10 @@ class SettingsStore(MutableMapping[str, JsonValue]): ) self._deleted_runtime_keys = self._deleted_runtime_keys | frozenset((key,)) + def clear(self) -> None: + self._deleted_runtime_keys = frozenset(key for key in self._keys() if not self.owned_by_config(key)) + self._runtime_values = _EMPTY_VALUES + def __iter__(self) -> Iterator[str]: return iter( key diff --git a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py index aa410deac43..aec78933c89 100644 --- a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py +++ b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py @@ -1,6 +1,7 @@ from __future__ import annotations from typing import Final +from unittest.mock import patch import pytest @@ -168,6 +169,54 @@ def test_settings_store_refuses_a_runtime_write_to_a_config_owned_key() -> None: assert store.source("max_parallel_requests") == "config" +@pytest.mark.timeout(10) +def test_settings_store_clear_removes_every_key_the_config_file_does_not_own() -> None: + store: Final = SettingsStore("general_settings") + store.load_yaml({"master_key": "sk-config"}) + store.apply_db_row("general_settings", {"max_parallel_requests": 3, "alerting": ["slack"]}) + store["allow_requests_on_db_unavailable"] = True + del store["alerting"] + + store.clear() + + assert dict(store) == {"master_key": "sk-config"} + assert "alerting" not in store + with pytest.raises(KeyError): + store["max_parallel_requests"] + + +@pytest.mark.timeout(10) +def test_settings_store_clear_then_refill_matches_a_plain_dict() -> None: + expected: Final[dict[str, JsonValue]] = {"max_parallel_requests": 3, "alerting": ["slack"]} + store: Final = SettingsStore("general_settings") + store.update(expected) + + expected.clear() + store.clear() + expected.update({"alerting": ["email"], "max_parallel_requests": 11}) + store.update({"alerting": ["email"], "max_parallel_requests": 11}) + + assert dict(store) == expected + assert tuple(store) == tuple(expected) + assert len(store) == len(expected) + + +@pytest.mark.timeout(10) +@pytest.mark.parametrize("clear", (False, True)) +def test_settings_store_survives_a_patch_dict_round_trip_when_the_config_file_owns_a_key(clear: bool) -> None: + store: Final = SettingsStore("general_settings") + store.load_yaml({"master_key": "sk-config"}) + store.apply_db_row("general_settings", {"max_parallel_requests": 3}) + before: Final = dict(store) + + with patch.dict(store, {"allow_requests_on_db_unavailable": True}, clear=clear): + assert store["allow_requests_on_db_unavailable"] is True + assert store["master_key"] == "sk-config" + assert ("max_parallel_requests" in store) is not clear + + assert dict(store) == before + + def test_settings_store_reports_the_config_owned_keys_a_write_would_change() -> None: store: Final = SettingsStore("general_settings") store.load_yaml({"max_parallel_requests": 3, "ui_access_mode": "admin_only"}) From d6b6ab31b75992150f800ec5a2bc71cf78453e8c Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 19:27:23 +0000 Subject: [PATCH 188/253] test(proxy): compare the cleared settings store against an unmutated refill mapping Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/config_resolvers/test_settings_store.py | 14 ++++++-------- 1 file changed, 6 insertions(+), 8 deletions(-) diff --git a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py index aec78933c89..d59fb3e5b14 100644 --- a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py +++ b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py @@ -187,18 +187,16 @@ def test_settings_store_clear_removes_every_key_the_config_file_does_not_own() - @pytest.mark.timeout(10) def test_settings_store_clear_then_refill_matches_a_plain_dict() -> None: - expected: Final[dict[str, JsonValue]] = {"max_parallel_requests": 3, "alerting": ["slack"]} + refilled: Final[dict[str, JsonValue]] = {"alerting": ["email"], "max_parallel_requests": 11} store: Final = SettingsStore("general_settings") - store.update(expected) + store.update({"max_parallel_requests": 3, "alerting": ["slack"]}) - expected.clear() store.clear() - expected.update({"alerting": ["email"], "max_parallel_requests": 11}) - store.update({"alerting": ["email"], "max_parallel_requests": 11}) + store.update(refilled) - assert dict(store) == expected - assert tuple(store) == tuple(expected) - assert len(store) == len(expected) + assert dict(store) == refilled + assert tuple(store) == tuple(refilled) + assert len(store) == len(refilled) @pytest.mark.timeout(10) From 08d08191592b822396ad655b55d5b57a47449244 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 19:39:15 +0000 Subject: [PATCH 189/253] fix(proxy): keep the resolved values of config-owned keys across SettingsStore.clear() Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/config_resolvers/settings_store.py | 4 +++- .../proxy/config_resolvers/test_settings_store.py | 10 ++++++---- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/config_resolvers/settings_store.py b/litellm/proxy/config_resolvers/settings_store.py index 49f5a833b8f..d4ca0e87d2b 100644 --- a/litellm/proxy/config_resolvers/settings_store.py +++ b/litellm/proxy/config_resolvers/settings_store.py @@ -89,7 +89,9 @@ class SettingsStore(MutableMapping[str, JsonValue]): def clear(self) -> None: self._deleted_runtime_keys = frozenset(key for key in self._keys() if not self.owned_by_config(key)) - self._runtime_values = _EMPTY_VALUES + self._runtime_values = MappingProxyType( + {key: value for key, value in self._runtime_values.items() if self.owned_by_config(key)} + ) def __iter__(self) -> Iterator[str]: return iter( diff --git a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py index d59fb3e5b14..88ec382b013 100644 --- a/tests/test_litellm/proxy/config_resolvers/test_settings_store.py +++ b/tests/test_litellm/proxy/config_resolvers/test_settings_store.py @@ -172,14 +172,15 @@ def test_settings_store_refuses_a_runtime_write_to_a_config_owned_key() -> None: @pytest.mark.timeout(10) def test_settings_store_clear_removes_every_key_the_config_file_does_not_own() -> None: store: Final = SettingsStore("general_settings") - store.load_yaml({"master_key": "sk-config"}) + store.load_yaml({"master_key": "os.environ/MASTER_KEY"}) store.apply_db_row("general_settings", {"max_parallel_requests": 3, "alerting": ["slack"]}) + store.apply_runtime_values({"master_key": "sk-resolved", "alerting": ["slack"]}) store["allow_requests_on_db_unavailable"] = True del store["alerting"] store.clear() - assert dict(store) == {"master_key": "sk-config"} + assert dict(store) == {"master_key": "sk-resolved"} assert "alerting" not in store with pytest.raises(KeyError): store["max_parallel_requests"] @@ -203,13 +204,14 @@ def test_settings_store_clear_then_refill_matches_a_plain_dict() -> None: @pytest.mark.parametrize("clear", (False, True)) def test_settings_store_survives_a_patch_dict_round_trip_when_the_config_file_owns_a_key(clear: bool) -> None: store: Final = SettingsStore("general_settings") - store.load_yaml({"master_key": "sk-config"}) + store.load_yaml({"master_key": "os.environ/MASTER_KEY"}) store.apply_db_row("general_settings", {"max_parallel_requests": 3}) + store.apply_runtime_values({"master_key": "sk-resolved", "max_parallel_requests": 3}) before: Final = dict(store) with patch.dict(store, {"allow_requests_on_db_unavailable": True}, clear=clear): assert store["allow_requests_on_db_unavailable"] is True - assert store["master_key"] == "sk-config" + assert store["master_key"] == "sk-resolved" assert ("max_parallel_requests" in store) is not clear assert dict(store) == before From 2703fab3c810f1c82b65c035acf4920ad3ac85e3 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 20:21:21 +0000 Subject: [PATCH 190/253] fix(deps): update anyio for osv scan Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- uv.lock | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/uv.lock b/uv.lock index a5e60c68515..29027c2a75f 100644 --- a/uv.lock +++ b/uv.lock @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.13.0" +version = "4.15.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, - { name = "typing-extensions", marker = "python_full_version < '3.13'" }, + { name = "typing-extensions" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } +sdist = { url = "https://files.pythonhosted.org/packages/a9/d2/f4d173e22df740bc37b1db102b386ba719b66e95b0f0d751f556b387e6d2/anyio-4.15.1.tar.gz", hash = "sha256:9f28306018cbd6d329e64a36d58256edff76dd996fe423bc957326e578b82a94", size = 276966, upload-time = "2026-09-05T10:42:39.44Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, + { url = "https://files.pythonhosted.org/packages/12/b8/4bd346e22b28902df4d651910f5242c28d84e4a5c2435ca5c3f797ed7e2e/anyio-4.15.1-py3-none-any.whl", hash = "sha256:6152fdbbf9a77fdec97731721bebf7c4c44f7c29b424b0065826173efc7ed101", size = 132079, upload-time = "2026-09-05T10:42:37.923Z" }, ] [[package]] From e036e256ef00be1afbed0cfb0a9ce87f67d18fa5 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 20:23:53 +0000 Subject: [PATCH 191/253] revert: drop redundant anyio lockfile bump Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- uv.lock | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/uv.lock b/uv.lock index 29027c2a75f..a5e60c68515 100644 --- a/uv.lock +++ b/uv.lock @@ -315,16 +315,16 @@ vertex = [ [[package]] name = "anyio" -version = "4.15.1" +version = "4.13.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "exceptiongroup", marker = "python_full_version < '3.11'" }, { name = "idna" }, - { name = "typing-extensions" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/a9/d2/f4d173e22df740bc37b1db102b386ba719b66e95b0f0d751f556b387e6d2/anyio-4.15.1.tar.gz", hash = "sha256:9f28306018cbd6d329e64a36d58256edff76dd996fe423bc957326e578b82a94", size = 276966, upload-time = "2026-09-05T10:42:39.44Z" } +sdist = { url = "https://files.pythonhosted.org/packages/19/14/2c5dd9f512b66549ae92767a9c7b330ae88e1932ca57876909410251fe13/anyio-4.13.0.tar.gz", hash = "sha256:334b70e641fd2221c1505b3890c69882fe4a2df910cba14d97019b90b24439dc", size = 231622, upload-time = "2026-03-24T12:59:09.671Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/12/b8/4bd346e22b28902df4d651910f5242c28d84e4a5c2435ca5c3f797ed7e2e/anyio-4.15.1-py3-none-any.whl", hash = "sha256:6152fdbbf9a77fdec97731721bebf7c4c44f7c29b424b0065826173efc7ed101", size = 132079, upload-time = "2026-09-05T10:42:37.923Z" }, + { url = "https://files.pythonhosted.org/packages/da/42/e921fccf5015463e32a3cf6ee7f980a6ed0f395ceeaa45060b61d86486c2/anyio-4.13.0-py3-none-any.whl", hash = "sha256:08b310f9e24a9594186fd75b4f73f4a4152069e3853f1ed8bfbf58369f4ad708", size = 114353, upload-time = "2026-03-24T12:59:08.246Z" }, ] [[package]] From eeca8b682b4c681d8fa0489959a63dd9ba294550 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:28:28 -0700 Subject: [PATCH 192/253] fix(vscode): raise the VS Code minimum to 1.115 for per-model configuration --- vscode-extension/README.md | 2 +- vscode-extension/package-lock.json | 10 +++++----- vscode-extension/package.json | 4 ++-- 3 files changed, 8 insertions(+), 8 deletions(-) diff --git a/vscode-extension/README.md b/vscode-extension/README.md index 47c1a3afd27..cea8af6b578 100644 --- a/vscode-extension/README.md +++ b/vscode-extension/README.md @@ -17,7 +17,7 @@ Add the same provider again with another name to reach a second gateway or a sec ## Requirements -VS Code 1.109 or newer and a LiteLLM AI Gateway the key can reach. The key needs access to at least one model group whose mode is `chat` +VS Code 1.115 or newer and a LiteLLM AI Gateway the key can reach. The key needs access to at least one model group whose mode is `chat` ## Development diff --git a/vscode-extension/package-lock.json b/vscode-extension/package-lock.json index a4f962f2039..453dd8e1ff8 100644 --- a/vscode-extension/package-lock.json +++ b/vscode-extension/package-lock.json @@ -13,14 +13,14 @@ }, "devDependencies": { "@types/node": "^22.20.3", - "@types/vscode": "1.109.0", + "@types/vscode": "1.115.0", "@vscode/vsce": "^4.0.0", "esbuild": "^0.28.2", "typescript": "^5.9.3", "vitest": "^4.1.11" }, "engines": { - "vscode": "^1.109.0" + "vscode": "^1.115.0" } }, "node_modules/@azure/abort-controller": { @@ -1301,9 +1301,9 @@ } }, "node_modules/@types/vscode": { - "version": "1.109.0", - "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.109.0.tgz", - "integrity": "sha512-0Pf95rnwEIwDbmXGC08r0B4TQhAbsHQ5UyTIgVgoieDe4cOnf92usuR5dEczb6bTKEp7ziZH4TV1TRGPPCExtw==", + "version": "1.115.0", + "resolved": "https://registry.npmjs.org/@types/vscode/-/vscode-1.115.0.tgz", + "integrity": "sha512-/M8cdznOlqtMqduHKKlIF00v4eum4ZWKgn8YoPRKcN6PDdvoWeeqDaQSnw63ipDbq1Uzz78Wndk/d0uSPwORfA==", "dev": true, "license": "MIT" }, diff --git a/vscode-extension/package.json b/vscode-extension/package.json index 67e3bb6c727..431dad316a1 100644 --- a/vscode-extension/package.json +++ b/vscode-extension/package.json @@ -15,7 +15,7 @@ "url": "https://github.com/BerriAI/litellm/issues" }, "engines": { - "vscode": "^1.109.0" + "vscode": "^1.115.0" }, "categories": [ "AI", @@ -78,7 +78,7 @@ }, "devDependencies": { "@types/node": "^22.20.3", - "@types/vscode": "1.109.0", + "@types/vscode": "1.115.0", "@vscode/vsce": "^4.0.0", "esbuild": "^0.28.2", "typescript": "^5.9.3", From a37d1e8b9b752f2c0672bc394e14095ef7c69b31 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:30:08 -0700 Subject: [PATCH 193/253] test: type the eager_input_streaming test helpers and mark constants Final --- .../test_anthropic_chat_transformation.py | 22 +++++++++---------- ...ations_anthropic_claude3_transformation.py | 17 ++++++++------ .../chat/test_converse_transformation.py | 15 ++++++++----- .../test_anthropic_claude3_transformation.py | 9 +++++--- 4 files changed, 36 insertions(+), 27 deletions(-) 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 1359bf40dcd..5e179f950a0 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 @@ -6372,18 +6372,19 @@ def test_response_format_tool_path_skips_forced_tool_choice_when_unsupported(loc assert "tool_choice" not in result -def _eager_chat_tool(**extra): +def _eager_chat_function(**extra: object) -> dict[str, object]: return { - "type": "function", - "function": { - "name": "write_file", - "description": "Write a file", - "parameters": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}, - }, + "name": "write_file", + "description": "Write a file", + "parameters": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}, **extra, } +def _eager_chat_tool(**extra: object) -> dict[str, object]: + return {"type": "function", "function": _eager_chat_function(), **extra} + + @pytest.mark.parametrize("flag", [True, False]) def test_eager_input_streaming_passed_through_from_tool_top_level(flag): mapped_tool, _ = AnthropicConfig()._map_tool_helper(_eager_chat_tool(eager_input_streaming=flag)) @@ -6398,10 +6399,9 @@ def test_eager_input_streaming_passed_through_from_tool_top_level(flag): def test_eager_input_streaming_passed_through_from_function(): - tool = _eager_chat_tool() - tool["function"]["eager_input_streaming"] = True - - mapped_tool, _ = AnthropicConfig()._map_tool_helper(tool) + mapped_tool, _ = AnthropicConfig()._map_tool_helper( + {"type": "function", "function": _eager_chat_function(eager_input_streaming=True)} + ) assert mapped_tool["eager_input_streaming"] is True assert "eager_input_streaming" not in mapped_tool["input_schema"] diff --git a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py index 1395489334e..84db0733227 100644 --- a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py @@ -1,6 +1,7 @@ import asyncio import json import uuid +from typing import Final from unittest.mock import patch import httpx @@ -861,14 +862,16 @@ def test_bedrock_chat_invoke_tool_search_beta_follows_model_map( assert result.get("anthropic_beta") == expected_betas -FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" -EAGER_TOOL_SCHEMA = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} +FINE_GRAINED_TOOL_STREAMING_BETA: Final = "fine-grained-tool-streaming-2025-05-14" +EAGER_TOOL_SCHEMA: Final = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} -def _chat_invoke_request_with_tools(tools, headers=None): - config = AmazonAnthropicClaudeConfig() - model = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" - optional_params = config.map_openai_params( +def _chat_invoke_request_with_tools( + tools: list[dict[str, object]], headers: dict[str, str] | None = None +) -> dict[str, object]: + config: Final = AmazonAnthropicClaudeConfig() + model: Final = "us.anthropic.claude-sonnet-4-5-20250929-v1:0" + optional_params: Final = config.map_openai_params( non_default_params={"max_tokens": 64, "stream": True, "tools": tools}, optional_params={}, model=model, @@ -883,7 +886,7 @@ def _chat_invoke_request_with_tools(tools, headers=None): ) -def _eager_openai_tool(name, **extra): +def _eager_openai_tool(name: str, **extra: object) -> dict[str, object]: return {"type": "function", "function": {"name": name, "parameters": EAGER_TOOL_SCHEMA}, **extra} 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 650427f2ae3..cd33e8d34d8 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -4,6 +4,7 @@ import os import httpx import pytest +from typing import Final from unittest.mock import MagicMock, patch import litellm @@ -7402,23 +7403,25 @@ def test_transform_response_honors_json_mode_kwarg_when_optional_params_lack_it( assert json.loads(result.choices[0].message.content) == {"city": "Paris", "population": 2100000} -FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" -EAGER_TOOL_SCHEMA = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} +FINE_GRAINED_TOOL_STREAMING_BETA: Final = "fine-grained-tool-streaming-2025-05-14" +EAGER_TOOL_SCHEMA: Final = {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]} -def _eager_openai_tool(**extra): +def _eager_openai_tool(**extra: object) -> dict[str, object]: return {"type": "function", "function": {"name": "write_file", "parameters": EAGER_TOOL_SCHEMA}, **extra} -def _eager_openai_function_tool(**extra): +def _eager_openai_function_tool(**extra: object) -> dict[str, object]: return {"type": "function", "function": {"name": "write_file", "parameters": EAGER_TOOL_SCHEMA, **extra}} -def _eager_anthropic_tool(**extra): +def _eager_anthropic_tool(**extra: object) -> dict[str, object]: return {"name": "write_file", "input_schema": EAGER_TOOL_SCHEMA, **extra} -def _converse_request(model, tools, headers=None): +def _converse_request( + model: str, tools: list[dict[str, object]], headers: dict[str, object] | None = None +) -> dict[str, object]: return AmazonConverseConfig()._transform_request_helper( model=model, system_content_blocks=[], diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index b71e8e521f2..7be005c0efe 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -4,6 +4,7 @@ import json import os from datetime import datetime from types import SimpleNamespace +from typing import Final from unittest.mock import Mock import pytest @@ -3246,10 +3247,12 @@ def test_bedrock_messages_strips_effort_but_keeps_format_for_sonnet_4_5(local_mo assert result.get("output_config") == {"format": schema_format} -FINE_GRAINED_TOOL_STREAMING_BETA = "fine-grained-tool-streaming-2025-05-14" +FINE_GRAINED_TOOL_STREAMING_BETA: Final = "fine-grained-tool-streaming-2025-05-14" -def _invoke_request_with_tools(tools, headers=None): +def _invoke_request_with_tools( + tools: list[dict[str, object]], headers: dict[str, str] | None = None +) -> dict[str, object]: from litellm.types.router import GenericLiteLLMParams return AmazonAnthropicClaudeMessagesConfig().transform_anthropic_messages_request( @@ -3261,7 +3264,7 @@ def _invoke_request_with_tools(tools, headers=None): ) -def _eager_invoke_tool(name, eager_input_streaming): +def _eager_invoke_tool(name: str, eager_input_streaming: bool) -> dict[str, object]: return { "name": name, "description": f"{name} tool", From e75ad61fceaf08b49830dc4c737db2ea7a28f49c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:32:43 -0700 Subject: [PATCH 194/253] fix(router): rank litellm_settings.request_timeout on the passthrough route like the completion route --- litellm/router.py | 8 +++++++- tests/test_litellm/test_router.py | 10 ++++++++++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/litellm/router.py b/litellm/router.py index a5523d7af79..e3e9454ddce 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -3877,11 +3877,17 @@ class Router: ) _router_timeout: Final = ( - float(self._explicit_timeout) if isinstance(self._explicit_timeout, (int, float)) else None + self.request_timeout + if self.request_timeout is not None + else float(self._explicit_timeout) + if isinstance(self._explicit_timeout, (int, float)) + else None ) _router_stream_timeout: Final = ( self.stream_timeout if self.stream_timeout is not None + else self.request_timeout + if self.request_timeout is not None else self.default_litellm_params.get("stream_timeout") ) kwargs["timeout"] = resolve_llm_passthrough_timeout( diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e92e5de23dc..148ca0b0ffa 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -8230,6 +8230,16 @@ class TestRouterRequestTimeoutPropagation: == 60 ) + def test_passthrough_prefers_request_timeout_over_router_timeout(self, explicit_request_timeout): + router = self._make_router(timeout=330) + deployment: Final = router.model_list[0] + assert _passthrough_timeout(router, deployment, stream=False) == 300.0 + assert _passthrough_timeout(router, deployment, stream=True) == 300.0 + + def test_passthrough_stream_timeout_still_wins_over_request_timeout(self, explicit_request_timeout): + router = self._make_router(timeout=330, stream_timeout=45) + assert _passthrough_timeout(router, router.model_list[0], stream=True) == 45.0 + # --------------------------------------------------------------------------- # Deferred-stream eager-fetch tests From e5fd2bec832be2ddf76851f73859d8bc2fb2b957 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:34:29 -0700 Subject: [PATCH 195/253] fix: forward eager_input_streaming through the Responses API tool bridge --- .../transformation.py | 4 +++ .../test_litellm_completion_responses.py | 32 +++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 01fb6cb483d..0d1297f63e2 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -2020,6 +2020,8 @@ class LiteLLMCompletionResponsesConfig: chat_completion_tool["allowed_callers"] = tool.get("allowed_callers") if tool.get("input_examples"): chat_completion_tool["input_examples"] = tool.get("input_examples") + if tool.get("eager_input_streaming") is not None: + chat_completion_tool["eager_input_streaming"] = tool.get("eager_input_streaming") return ResponsesToolChatForm( chat_tools=(cast(ChatCompletionToolParam, chat_completion_tool),), web_search_options=None ) @@ -2096,6 +2098,8 @@ class LiteLLMCompletionResponsesConfig: responses_tool["allowed_callers"] = tool.get("allowed_callers") if tool.get("input_examples") is not None: responses_tool["input_examples"] = tool.get("input_examples") + if tool.get("eager_input_streaming") is not None: + responses_tool["eager_input_streaming"] = tool.get("eager_input_streaming") result.append(responses_tool) else: # mcp or other: pass through unchanged diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 0ed101952be..f1924b87a20 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -1900,6 +1900,38 @@ class TestToolTransformation: assert "defer_loading" not in result_tool assert "allowed_callers" not in result_tool assert "input_examples" not in result_tool + assert "eager_input_streaming" not in result_tool + + @pytest.mark.parametrize("eager_input_streaming", [True, False]) + def test_transform_function_tools_forwards_eager_input_streaming(self, eager_input_streaming: bool) -> None: + function_tool: Final = { + "type": "function", + "name": "write_file", + "parameters": {"type": "object", "properties": {"path": {"type": "string"}}}, + "eager_input_streaming": eager_input_streaming, + } + + result_tools, _ = LiteLLMCompletionResponsesConfig.transform_responses_api_tools_to_chat_completion_tools( + tools=[function_tool] + ) + + assert result_tools[0]["eager_input_streaming"] is eager_input_streaming + + @pytest.mark.parametrize("eager_input_streaming", [True, False]) + def test_chat_completion_tools_to_responses_tools_keeps_eager_input_streaming( + self, eager_input_streaming: bool + ) -> None: + chat_tool: Final = { + "type": "function", + "function": {"name": "write_file", "parameters": {"type": "object"}}, + "eager_input_streaming": eager_input_streaming, + } + + result_tools: Final = LiteLLMCompletionResponsesConfig.transform_chat_completion_tool_params_to_responses_api_tools( + [chat_tool] + ) + + assert result_tools[0]["eager_input_streaming"] is eager_input_streaming def test_transform_code_execution_tools(self): """Test that code_execution tools are passed through as-is""" From 72a14246a6ffef7724492b8c0da2798c30a32148 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 13:50:56 -0700 Subject: [PATCH 196/253] fix(proxy): requeue daily spend rows when the commit fails without the Redis buffer With the Redis transaction buffer off, each daily spend queue (user, team, org, end user, agent) was drained into a dict and handed to the bulk upsert. When the upsert raised after its retries, the drained dict was discarded and the exception escaped update_spend, so those rows never reached the daily rollup tables and /user/daily/activity stayed short forever while /spend/logs had every request. Each daily queue now flushes through one helper that puts the uncommitted remainder back on the queue for the next tick and moves on to the next table, the same shape the window-spend step already used. --- litellm/proxy/db/db_spend_update_writer.py | 99 ++++++++++++------- .../proxy/db/test_db_spend_update_writer.py | 55 +++++++++++ 2 files changed, 118 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index c13b852484e..a4f1713d61d 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -15,7 +15,7 @@ import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast, overload from urllib.parse import quote, unquote from typing_extensions import ReadOnly, TypedDict @@ -143,6 +143,20 @@ class _SpendTransactionManager(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... +_DailySpendTransactionT = TypeVar("_DailySpendTransactionT", bound=BaseDailySpendTransaction) + + +class _DailySpendCommit(Protocol[_DailySpendTransactionT]): + async def __call__( + self, + *, + n_retry_times: int, + prisma_client: PrismaClient, + proxy_logging_obj: ProxyLogging, + daily_spend_transactions: dict[str, _DailySpendTransactionT], + ) -> None: ... + + def _timed_request_duration_ms( payload: dict | SpendLogsPayload, request_status: Literal["success", "failure"], @@ -1288,6 +1302,34 @@ class DBSpendUpdateWriter: cronjob_id=DB_SPEND_UPDATE_JOB_NAME, ) + async def _flush_daily_spend_queue( + self, + queue: DailySpendUpdateQueue, + entity_type: Literal["user", "team", "org", "end_user", "agent"], + commit: _DailySpendCommit[_DailySpendTransactionT], + n_retry_times: int, + prisma_client: PrismaClient, + proxy_logging_obj: ProxyLogging, + ) -> None: + transactions: Final = await queue.flush_and_get_aggregated_daily_spend_update_transactions() + try: + await commit( + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions), + ) + except Exception as e: # noqa: BLE001 # the uncommitted rows go back on the queue; the other tables must still flush + spend_log_error( + "Spend tracking - failed to commit daily %s spend updates. " + "Re-queued %d rows for retry on next tick. Error: %s", + entity_type, + len(transactions), + str(e), + exc=e, + ) + await queue.add_update(transactions) + async def _commit_spend_updates_to_db_without_redis_buffer( self, prisma_client: PrismaClient, @@ -1316,74 +1358,59 @@ class DBSpendUpdateWriter: ################## Daily Spend Update Transactions ################## # Aggregate all in memory daily spend transactions and commit to db - daily_spend_update_transactions: Final = cast( - dict[str, DailyUserSpendTransaction], - await self.daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), - ) - - await DBSpendUpdateWriter.update_daily_user_spend( + await self._flush_daily_spend_queue( + queue=self.daily_spend_update_queue, + entity_type="user", + commit=DBSpendUpdateWriter.update_daily_user_spend, n_retry_times=n_retry_times, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_spend_update_transactions, ) ################## Daily Team Spend Update Transactions ################## # Aggregate all in memory daily team spend transactions and commit to db - daily_team_spend_update_transactions: Final = cast( - dict[str, DailyTeamSpendTransaction], - await self.daily_team_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), - ) - - await DBSpendUpdateWriter.update_daily_team_spend( + await self._flush_daily_spend_queue( + queue=self.daily_team_spend_update_queue, + entity_type="team", + commit=DBSpendUpdateWriter.update_daily_team_spend, n_retry_times=n_retry_times, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_team_spend_update_transactions, ) ################## Daily Organization Spend Update Transactions ################## # Aggregate all in memory daily org spend transactions and commit to db - daily_org_spend_update_transactions: Final = cast( - dict[str, DailyOrganizationSpendTransaction], - await self.daily_org_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), - ) - - await DBSpendUpdateWriter.update_daily_org_spend( + await self._flush_daily_spend_queue( + queue=self.daily_org_spend_update_queue, + entity_type="org", + commit=DBSpendUpdateWriter.update_daily_org_spend, n_retry_times=n_retry_times, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_org_spend_update_transactions, ) # NOTE: Daily tag spend is committed by a separate scheduler job. ################## Daily End-User Spend Update Transactions ################## # Aggregate all in memory daily end-user spend transactions and commit to db - daily_end_user_spend_update_transactions: Final = cast( - dict[str, DailyEndUserSpendTransaction], - await self.daily_end_user_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), - ) - - await DBSpendUpdateWriter.update_daily_end_user_spend( + await self._flush_daily_spend_queue( + queue=self.daily_end_user_spend_update_queue, + entity_type="end_user", + commit=DBSpendUpdateWriter.update_daily_end_user_spend, n_retry_times=n_retry_times, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_end_user_spend_update_transactions, ) ################## Daily Agent Spend Update Transactions ################## # Aggregate all in memory daily agent spend transactions and commit to db - daily_agent_spend_update_transactions: Final = cast( - dict[str, DailyAgentSpendTransaction], - await self.daily_agent_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), - ) - - await DBSpendUpdateWriter.update_daily_agent_spend( + await self._flush_daily_spend_queue( + queue=self.daily_agent_spend_update_queue, + entity_type="agent", + commit=DBSpendUpdateWriter.update_daily_agent_spend, n_retry_times=n_retry_times, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_agent_spend_update_transactions, ) ################## Budget Window Spend Update Transactions ################## diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b3f5a60877d..3feb117edb7 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -2809,6 +2809,61 @@ async def test_failed_window_spend_commit_requeues_the_increments_and_continues_ assert requeued == (transaction,) +class _DailySpendFakeDB(_WindowSpendFakeDB): + """Records the daily rollup upserts it is handed and fails the ones aimed at one table.""" + + def __init__(self, failing_table: str | None) -> None: + super().__init__() + self.failing_table = failing_table + self.execute_raw_calls: list[Statement] = [] + + async def execute_raw(self, query: str, *args: object) -> int: + if self.failing_table is not None and self.failing_table in query: + raise Exception("connection reset") + self.execute_raw_calls.append((query, args)) + return len(args) + + +def _daily_upserts(db: _DailySpendFakeDB, table: str) -> list[Statement]: + return [statement for statement in db.execute_raw_calls if table in statement[0]] + + +@pytest.mark.asyncio +async def test_failed_daily_spend_commit_requeues_the_rows_and_flushes_the_other_tables(): + """With the Redis buffer off, a daily batch that failed to commit was discarded along + with the tick's exception, so the Usage page stayed short of LiteLLM_SpendLogs for good. + The uncommitted rows must go back on their queue and land on the next tick, and the + other daily tables must still be flushed on the failing tick.""" + db_writer = DBSpendUpdateWriter() + await db_writer.daily_spend_update_queue.add_update({"user-key": _daily_txn(user_id="user-1")}) + team_txn = {key: value for key, value in _daily_txn().items() if key != "user_id"} | {"team_id": "team-1"} + await db_writer.daily_team_spend_update_queue.add_update({"team-key": team_txn}) + db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend") + db_writer._flush_tool_discovery_queue = AsyncMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.failure_handler = AsyncMock() + + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + assert _daily_upserts(db, "LiteLLM_DailyUserSpend") == [] + (team_upsert,) = _daily_upserts(db, "LiteLLM_DailyTeamSpend") + assert _row_values(team_upsert, "team_id") == ["team-1"] + db_writer._flush_tool_discovery_queue.assert_called_once() + + db.failing_table = None + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + (user_upsert,) = _daily_upserts(db, "LiteLLM_DailyUserSpend") + assert _row_values(user_upsert, "user_id") == ["user-1"] + assert _row_values(user_upsert, "spend") == [0.1] + assert len(_daily_upserts(db, "LiteLLM_DailyTeamSpend")) == 1 + assert db_writer.daily_spend_update_queue.update_queue.empty() + + @pytest.mark.asyncio async def test_failed_window_spend_commit_from_redis_is_restored_to_redis(): """The Redis drain is destructive, so a failed window commit has to push From de4b520153718864dddfa53a37dd6b56825c958b Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:12:56 +0000 Subject: [PATCH 197/253] fix(proxy): track team member spend when the member has no budget add_new_member only wrote a LiteLLM_TeamMembership row when a budget id resolved, and the spend writer used update_many so a missing row failed silently. Members without a budget therefore never accrued per-member spend. The membership row is now always upserted (budget_id NULL when no budget applies) and the spend write is an upsert so members added before this fix start accruing on their next request. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 16 +++- litellm/proxy/management_helpers/utils.py | 23 +++--- litellm/repositories/prisma_protocols.py | 2 + .../proxy/db/test_db_spend_update_writer.py | 37 ++++++--- .../test_team_endpoints.py | 10 +-- .../test_management_helpers_utils.py | 75 ++++++++++++++----- 6 files changed, 114 insertions(+), 49 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index c13b852484e..913675705e3 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1693,11 +1693,19 @@ class DBSpendUpdateWriter: team_id = key.split("::")[1] user_id = key.split("::")[3] - batcher.litellm_teammembership.update_many( # 'update_many' prevents error from being raised if no row exists - where={"team_id": team_id, "user_id": user_id}, + batcher.litellm_teammembership.upsert( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, data={ - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, }, ) # Transaction succeeded, break out of retry loop diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index f3bd4b0f6dd..d940eba86d3 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -86,7 +86,9 @@ class _PrismaUserTable(Protocol): class _PrismaTeamMembershipTable(Protocol): """Team membership table actions the management helpers issue.""" - async def create(self, *, data: Mapping[str, object], include: Mapping[str, bool]) -> _PrismaRecord: ... + async def upsert( + self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]], include: Mapping[str, bool] + ) -> _PrismaRecord: ... class MemberWriteTx(Protocol): @@ -348,7 +350,7 @@ async def _resolve_member_budget_id( default member budget is cloned (with ``budget_duration`` overriding its reset window while keeping its other limits). A lone ``budget_duration`` with no team default creates a window-only budget. With nothing set the - member gets no budget. + member gets no budget, though ``add_new_member`` still writes its membership row. """ has_explicit_limit: Final = max_budget_in_team is not None or allowed_models is not None @@ -415,9 +417,9 @@ async def add_new_member( Add a new member to a team - add team id to user table - - add team member w/ budget to team member table + - add team member to team member table, linked to a budget when one resolves - Returns created/existing user + team membership w/ budget id + Returns created/existing user + team membership (``budget_id`` is ``None`` when no budget applies) Callers already inside a transaction pass it as ``tx`` so every write here runs on that connection instead of borrowing more from the pool while the caller's locks are held. @@ -471,14 +473,13 @@ async def add_new_member( tx=tx, ) - if _budget_id and returned_user is not None and returned_user.user_id is not None: + if returned_user is not None and returned_user.user_id is not None: membership_table: Final[_PrismaTeamMembershipTable] = _team_membership_table(prisma_client, tx) - _returned_team_membership: Final = await membership_table.create( - data={ - "team_id": team_id, - "user_id": returned_user.user_id, - "budget_id": _budget_id, - }, + membership_key: Final[Mapping[str, object]] = {"user_id": returned_user.user_id, "team_id": team_id} + budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} + _returned_team_membership: Final = await membership_table.upsert( + where={"user_id_team_id": membership_key}, + data={"create": {**membership_key, **budget_link}, "update": {**budget_link}}, include={"litellm_budget_table": True}, ) diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 93b8c5c7cd7..0919ae9f808 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -123,6 +123,8 @@ class BatchTable(Protocol): def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... + def upsert(self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]]) -> None: ... + class PrismaBatch(Protocol): @property diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b3f5a60877d..bed903d7a1f 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -918,7 +918,11 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total """ Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - update_many call, using the same response_cost. + upsert call, using the same response_cost. + + Regression (LIT-5502): members added without a budget had no membership row, and + the previous update_many matched zero rows, so their spend was silently dropped. + The upsert has to create the row seeded with this call's cost in that case. """ db_writer = DBSpendUpdateWriter() @@ -930,7 +934,7 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total mock_batcher.litellm_teamtable = MagicMock() mock_batcher.litellm_teamtable.update_many = MagicMock() mock_batcher.litellm_teammembership = MagicMock() - mock_batcher.litellm_teammembership.update_many = MagicMock() + mock_batcher.litellm_teammembership.upsert = MagicMock() mock_batcher.litellm_organizationtable = MagicMock() mock_batcher.litellm_organizationtable.update_many = MagicMock() mock_batcher.litellm_tagtable = MagicMock() @@ -979,12 +983,21 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total db_spend_update_transactions=db_spend_update_transactions, ) - mock_batcher.litellm_teammembership.update_many.assert_called_once() - call_kwargs = mock_batcher.litellm_teammembership.update_many.call_args[1] - assert call_kwargs["where"] == {"team_id": team_id, "user_id": user_id} + mock_batcher.litellm_teammembership.upsert.assert_called_once() + mock_batcher.litellm_teammembership.update_many.assert_not_called() + call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs + assert call_kwargs["where"] == {"user_id_team_id": {"user_id": user_id, "team_id": team_id}} assert call_kwargs["data"] == { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, } @@ -2219,9 +2232,13 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): "team_id::team_b::user_id::user_x": 0.3, }, "litellm_teammembership", - "update_many", - "team_id", - ["team_a", "team_b", "team_c"], + "upsert", + "user_id_team_id", + [ + {"user_id": "user_x", "team_id": "team_a"}, + {"user_id": "user_x", "team_id": "team_b"}, + {"user_id": "user_x", "team_id": "team_c"}, + ], id="team_member", ), pytest.param( 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 4f5c5066367..f7157da3100 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -2086,7 +2086,7 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti tx.litellm_usertable.upsert = AsyncMock(return_value=added_user) tx.litellm_usertable.update_many = AsyncMock() tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) - tx.litellm_teammembership.create = AsyncMock(return_value=membership) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) tx_cm = MagicMock() tx_cm.__aenter__ = AsyncMock(return_value=tx) @@ -5772,7 +5772,7 @@ async def test_new_team_max_budget_within_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -5915,7 +5915,7 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -6063,7 +6063,7 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) @@ -9525,7 +9525,7 @@ async def test_new_team_soft_budget_validation( "budget_id": None, } mock_prisma.db.litellm_teammembership = MagicMock() - mock_prisma.db.litellm_teammembership.create = AsyncMock( + mock_prisma.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_membership ) diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index a6b1fc32eda..0dcf7e1b8ad 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -234,7 +234,7 @@ async def test_add_new_member_clones_default_team_budget_id(): "budget_id": test_cloned_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -257,7 +257,7 @@ async def test_add_new_member_clones_default_team_budget_id(): assert result_team_membership.budget_id != test_default_budget_id mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() - mock_prisma_client.db.litellm_teammembership.create.assert_called_once() + mock_prisma_client.db.litellm_teammembership.upsert.assert_called_once() # The clone must have happened: find_unique on the default, create for the clone. mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with( @@ -274,9 +274,9 @@ async def test_add_new_member_clones_default_team_budget_id(): assert cloned_create_data["created_by"] == user_api_key_dict.user_id team_membership_call_args = ( - mock_prisma_client.db.litellm_teammembership.create.call_args + mock_prisma_client.db.litellm_teammembership.upsert.call_args ) - create_data = team_membership_call_args.kwargs["data"] + create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_cloned_budget_id @@ -332,7 +332,7 @@ async def test_add_new_member_budget_duration_only_clones_default_max_budget(): "budget_id": "cloned-dc", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -362,7 +362,8 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): Test that add_new_member links no budget to the team membership when neither max_budget_in_team nor default_team_budget_id is provided. - When the team has no default member budget, new members get nothing. + When the team has no default member budget, no budget row is created, but the + membership row still is, otherwise the member's spend has nowhere to accrue. """ from litellm.proxy._types import LitellmUserRoles @@ -393,7 +394,19 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): # Even though we mock these, they must NOT be called on the no-budget path. mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock() mock_prisma_client.db.litellm_budgettable.create = AsyncMock() - mock_prisma_client.db.litellm_teammembership.create = AsyncMock() + + mock_team_membership_response = MagicMock() + mock_team_membership_response.model_dump.return_value = { + "team_id": test_team_id, + "user_id": test_user_id, + "budget_id": None, + "spend": 0.0, + "total_spend": 0.0, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( + return_value=mock_team_membership_response + ) result_user, result_team_membership = await add_new_member( new_member=new_member, @@ -408,11 +421,20 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): assert result_user is not None assert result_user.user_id == test_user_id - # No budget id, so no team membership row is created. - assert result_team_membership is None mock_prisma_client.db.litellm_budgettable.find_unique.assert_not_called() mock_prisma_client.db.litellm_budgettable.create.assert_not_called() - mock_prisma_client.db.litellm_teammembership.create.assert_not_called() + + # Regression (LIT-5502): the membership row is what per-member spend increments land on, + # so it has to exist even when the member has no budget. Skipping it silently dropped spend. + assert result_team_membership is not None + assert result_team_membership.budget_id is None + mock_prisma_client.db.litellm_teammembership.upsert.assert_awaited_once() + upsert_kwargs = mock_prisma_client.db.litellm_teammembership.upsert.call_args.kwargs + assert upsert_kwargs["where"] == { + "user_id_team_id": {"user_id": test_user_id, "team_id": test_team_id} + } + assert upsert_kwargs["data"]["create"] == {"user_id": test_user_id, "team_id": test_team_id} + assert "budget_id" not in upsert_kwargs["data"]["update"] @pytest.mark.asyncio @@ -473,7 +495,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): "budget_id": test_new_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -502,10 +524,10 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): # Verify the team membership was created with the correct budget_id team_membership_call_args = ( - mock_prisma_client.db.litellm_teammembership.create.call_args + mock_prisma_client.db.litellm_teammembership.upsert.call_args ) assert team_membership_call_args is not None - create_data = team_membership_call_args.kwargs["data"] + create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_new_budget_id @@ -546,7 +568,7 @@ async def test_add_new_member_persists_budget_duration(): "budget_id": "budget-dur", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -610,7 +632,7 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): "budget_id": "budget-dur2", "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -700,7 +722,7 @@ async def test_add_new_member_with_user_email_clones_default_budget(): "budget_id": test_cloned_budget_id, "litellm_budget_table": None, } - mock_prisma_client.db.litellm_teammembership.create = AsyncMock( + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock( return_value=mock_team_membership_response ) @@ -1031,8 +1053,15 @@ async def test_add_new_member_appends_team_only_if_absent_for_existing_user(): } mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_after) mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() - # no team default budget and no explicit budget -> no team membership row mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + mock_membership = MagicMock() + mock_membership.model_dump.return_value = { + "team_id": "team-1", + "user_id": "existing-user", + "budget_id": None, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership) result_user, _ = await add_new_member( new_member=new_member, @@ -1099,6 +1128,14 @@ async def test_add_new_member_creates_missing_user_atomically_via_upsert(): mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() mock_prisma_client.db.litellm_usertable.create = AsyncMock() mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + mock_membership = MagicMock() + mock_membership.model_dump.return_value = { + "team_id": "team-1", + "user_id": "brand-new-user", + "budget_id": None, + "litellm_budget_table": None, + } + mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership) result_user, _ = await add_new_member( new_member=new_member, @@ -1147,7 +1184,7 @@ def _member_write_tx() -> MagicMock: tx.litellm_usertable.find_many = AsyncMock(return_value=[]) tx.litellm_budgettable.find_unique = AsyncMock(return_value=None) tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) - tx.litellm_teammembership.create = AsyncMock(return_value=membership) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) return tx @@ -1192,7 +1229,7 @@ async def test_add_new_member_runs_every_write_on_the_caller_transaction(new_mem assert result_membership.budget_id == "budget-pool" assert tx.litellm_budgettable.create.await_count == 1 - assert tx.litellm_teammembership.create.await_count == 1 + assert tx.litellm_teammembership.upsert.await_count == 1 assert tx.litellm_usertable.upsert.await_count + tx.litellm_usertable.create.await_count == 1 prisma_client.db.assert_not_called() From d9b48ac9414adbfbcf5c5f15521a7cb4293d278d Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 01:35:53 +0000 Subject: [PATCH 198/253] test(proxy): mock team membership upsert in team admin member add test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/proxy_unit_tests/test_proxy_server.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 1fcdaa67143..659097abbb9 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -1372,6 +1372,7 @@ async def test_create_team_member_add_team_admin( from fastapi import Request from litellm.proxy._types import ( + LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, Member, @@ -1454,6 +1455,10 @@ async def test_create_team_member_add_team_admin( team_mock_client.update = AsyncMock( return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) + membership_mock_client = AsyncMock() + membership_mock_client.upsert = AsyncMock( + return_value=LiteLLM_TeamMembership(user_id="1234", team_id=_team_id) + ) tx_cm = _member_add_tx_cm(team_mock_client) @@ -1463,6 +1468,11 @@ async def test_create_team_member_add_team_admin( "litellm_teamtable", team_mock_client, ), + patch.object( # test-quality-ok: legacy test swaps the prisma table on the module-level client + litellm.proxy.proxy_server.prisma_client.db, + "litellm_teammembership", + membership_mock_client, + ), patch.object( litellm.proxy.proxy_server.prisma_client, "tx", From 6dce85c7285e7e7a6ea1b0b450ff4041815470ba Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:10:29 +0000 Subject: [PATCH 199/253] fix(proxy): skip recreating membership rows for members removed before a spend flush Take the team advisory lock in the spend flush transaction and read the roster through it, so a delayed flush after /team/member_delete cannot recreate the deleted LiteLLM_TeamMembership row. TEAM_ADVISORY_LOCK_SQL moves to team_repository so the spend writer can import it without a circular import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 72 +++++++--- .../access_group_team_sync.py | 8 +- litellm/repositories/prisma_protocols.py | 8 +- litellm/repositories/team_repository.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 134 ++++++++++++------ 5 files changed, 160 insertions(+), 74 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 913675705e3..5c86c52c539 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -75,7 +75,8 @@ from litellm.proxy.spend_tracking.savings import ( marks_gateway_injection, ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error -from litellm.repositories.prisma_protocols import BatchTable +from litellm.repositories.prisma_protocols import BatchTable, RawQueryTransaction +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL, TeamRepository from litellm.types.utils import CallTypes if TYPE_CHECKING: @@ -133,7 +134,7 @@ class _SpendBatchManager(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... -class _SpendTransaction(Protocol): +class _SpendTransaction(RawQueryTransaction, Protocol): def batch_(self) -> _SpendBatchManager: ... @@ -161,6 +162,50 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx +async def _lock_and_read_rosters( + prisma_client: PrismaClient, transaction: _SpendTransaction, team_ids: Sequence[str] +) -> frozenset[tuple[str, str]]: + """Take each team's advisory lock on ``transaction`` and return its rostered (user_id, team_id) pairs. + + A spend flush can land after ``/team/member_delete`` removed the member. Holding the same + lock that endpoint takes, until this transaction commits, means a member read here is on + the team for the whole flush, so only they may have a missing membership row created. + """ + repository: Final = TeamRepository(prisma_client) + + async def locked_roster(team_id: str) -> tuple[tuple[str, str], ...]: + await transaction.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) + roster: Final = await repository.get_members_with_roles_locked(transaction, team_id) + return tuple((member.user_id, team_id) for member in roster or () if member.user_id is not None) + + rosters: Final = tuple([await locked_roster(team_id) for team_id in sorted(frozenset(team_ids))]) + return frozenset(pair for roster in rosters for pair in roster) + + +def _queue_team_member_spend( + memberships: BatchTable, user_id: str, team_id: str, response_cost: float, rostered: bool +) -> None: + increments: Final = { + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + } + if not rostered: + memberships.update_many(where={"team_id": team_id, "user_id": user_id}, data=increments) + return + memberships.upsert( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={ + "create": { + "team_id": team_id, + "user_id": user_id, + "spend": response_cost, + "total_spend": response_cost, + }, + "update": increments, + }, + ) + + def get_llm_router(): """The proxy's router, or None outside a running proxy. @@ -1685,6 +1730,9 @@ class DBSpendUpdateWriter: start_time = time.time() try: async with _spend_update_tx(prisma_client) as transaction: + rostered_members = await _lock_and_read_rosters( + prisma_client, transaction, tuple(team_id for _, team_id in team_memberships_to_invalidate) + ) async with transaction.batch_() as batcher: # Sort by composite key for consistent lock ordering across pods to prevent deadlocks. # Key format "team_id::::user_id::" makes the string sort equivalent to sorting by (team_id, user_id). @@ -1693,20 +1741,12 @@ class DBSpendUpdateWriter: team_id = key.split("::")[1] user_id = key.split("::")[3] - batcher.litellm_teammembership.upsert( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={ - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - }, + _queue_team_member_spend( + batcher.litellm_teammembership, + user_id, + team_id, + response_cost, + (user_id, team_id) in rostered_members, ) # Transaction succeeded, break out of retry loop break diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index 664e36c9f10..cfe207e66ea 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -21,13 +21,7 @@ from typing import Final, Protocol from pydantic import BaseModel, TypeAdapter from litellm.proxy.auth.auth_checks import _delete_cache_access_object - -# hashtext collisions only cost two unrelated teams a little serialization, and the -# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, -# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints -# reuses this exact statement to serialize /team/member_add and /team/delete against each -# other and against this mirror, rather than defining a second, divergent lock on the same key. -TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL _READ_TEAM_SQL: Final = 'SELECT access_group_ids FROM "LiteLLM_TeamTable" WHERE team_id = $1' diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index 0919ae9f808..c742f3f5f3f 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -7,7 +7,7 @@ private ones per file. """ from collections.abc import Mapping, Sequence -from typing import Protocol, TypeVar +from typing import LiteralString, Protocol, TypeVar RowT_co = TypeVar("RowT_co", covariant=True) @@ -108,6 +108,12 @@ class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... +class RawQueryTransaction(Protocol): + """A prisma transaction handle that can run raw SQL, e.g. an advisory lock or a locked read.""" + + async def query_raw(self, query: LiteralString, *args: str) -> Sequence[Mapping[str, object]]: ... + + class ReadOnlyTable(Protocol): async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ... diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 5ff07d76b5d..1b9cba48e41 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -15,10 +15,9 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) -from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.prisma_protocols import RawQueryTransaction, TableActions if TYPE_CHECKING: - from prisma import Prisma from prisma import models as prisma_models @@ -40,6 +39,13 @@ def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays: return team +# hashtext collisions only cost two unrelated teams a little serialization, and the +# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, +# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints, +# the access-group mirror and the team member spend flush all reuse this exact statement to +# serialize against each other, rather than defining a second, divergent lock on the same key. +TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" + _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( "metadata", @@ -78,7 +84,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) - async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member] | None: + async def get_members_with_roles_locked(self, tx: RawQueryTransaction, team_id: str) -> list[Member] | None: """Return the team's members_with_roles. The caller must already hold ``TEAM_ADVISORY_LOCK_SQL`` for this team_id on ``tx`` before calling this. diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index bed903d7a1f..b12dc68b7b1 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -20,6 +20,7 @@ from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) +from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL @pytest.mark.asyncio @@ -913,38 +914,22 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): assert call_kwargs["data"] == {"spend": {"increment": response_cost}} -@pytest.mark.asyncio -async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): - """ - Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) - and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - upsert call, using the same response_cost. - - Regression (LIT-5502): members added without a budget had no membership row, and - the previous update_many matched zero rows, so their spend was silently dropped. - The upsert has to create the row seeded with this call's cost in that case. - """ - db_writer = DBSpendUpdateWriter() - +def _team_member_flush_fixtures(team_id: str, rostered_user_ids: list[str]) -> tuple[MagicMock, AsyncMock, MagicMock]: + """A batcher, transaction and prisma client whose locked roster read for `team_id` lists `rostered_user_ids`.""" mock_batcher = MagicMock() - mock_batcher.litellm_verificationtoken = MagicMock() - mock_batcher.litellm_verificationtoken.update_many = MagicMock() - mock_batcher.litellm_usertable = MagicMock() - mock_batcher.litellm_usertable.update_many = MagicMock() - mock_batcher.litellm_teamtable = MagicMock() - mock_batcher.litellm_teamtable.update_many = MagicMock() mock_batcher.litellm_teammembership = MagicMock() mock_batcher.litellm_teammembership.upsert = MagicMock() - mock_batcher.litellm_organizationtable = MagicMock() - mock_batcher.litellm_organizationtable.update_many = MagicMock() - mock_batcher.litellm_tagtable = MagicMock() - mock_batcher.litellm_tagtable.update_many = MagicMock() - mock_batcher.litellm_agentstable = MagicMock() - mock_batcher.litellm_agentstable.update_many = MagicMock() + mock_batcher.litellm_teammembership.update_many = MagicMock() + + roster_row = {"members_with_roles": json.dumps([{"user_id": uid, "role": "user"} for uid in rostered_user_ids])} + + async def query_raw(query: str, *args: str) -> list[dict[str, object]]: + return [] if query == TEAM_ADVISORY_LOCK_SQL else [roster_row] mock_transaction = AsyncMock() mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) mock_transaction.__aexit__ = AsyncMock(return_value=False) + mock_transaction.query_raw = AsyncMock(side_effect=query_raw) mock_transaction.batch_ = MagicMock( return_value=AsyncMock( __aenter__=AsyncMock(return_value=mock_batcher), @@ -955,16 +940,11 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) + return mock_batcher, mock_transaction, mock_prisma_client - mock_proxy_logging = MagicMock() - # Skip team-membership cache invalidation — out of scope for this test. - mock_proxy_logging.call_details.get = MagicMock(return_value=None) - team_id = "team-abc" - user_id = "user-xyz" - response_cost = 0.75 - entity_id = f"team_id::{team_id}::user_id::{user_id}" - db_spend_update_transactions = { +def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict[str, dict[str, float]]: + return { "user_list_transactions": {}, "end_user_list_transactions": {}, "key_list_transactions": {}, @@ -975,14 +955,38 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total "agent_list_transactions": {}, } - with patch("litellm.proxy.utils._raise_failed_update_spend_exception"): - await db_writer._commit_spend_updates_to_db( - prisma_client=mock_prisma_client, - n_retry_times=0, - proxy_logging_obj=mock_proxy_logging, - db_spend_update_transactions=db_spend_update_transactions, - ) +@pytest.mark.asyncio +async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): + """ + Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) + and total_spend (non-resetting) on LiteLLM_TeamMembership in a single + upsert call, using the same response_cost. + + Regression (LIT-5502): members added without a budget had no membership row, and + the previous update_many matched zero rows, so their spend was silently dropped. + For a member still on the team roster the upsert has to create the row seeded + with this call's cost in that case. + """ + db_writer = DBSpendUpdateWriter() + team_id = "team-abc" + user_id = "user-xyz" + response_cost = 0.75 + mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, [user_id]) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details.get = MagicMock(return_value=None) + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=0, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=_team_member_only_transactions( + f"team_id::{team_id}::user_id::{user_id}", response_cost + ), + ) + + assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) mock_batcher.litellm_teammembership.upsert.assert_called_once() mock_batcher.litellm_teammembership.update_many.assert_not_called() call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs @@ -1001,6 +1005,43 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total } +@pytest.mark.asyncio +async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_removed_team_member(): + """ + A spend flush that lands after /team/member_delete must not resurrect the deleted + membership row: a user missing from the team roster, read under the team's advisory + lock inside the flush transaction, only gets an increment on whatever row still + exists, never a create. + """ + db_writer = DBSpendUpdateWriter() + team_id = "team-abc" + removed_user_id = "user-removed" + response_cost = 0.75 + mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, ["user-still-here"]) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.call_details.get = MagicMock(return_value=None) + + await db_writer._commit_spend_updates_to_db( + prisma_client=mock_prisma_client, + n_retry_times=0, + proxy_logging_obj=mock_proxy_logging, + db_spend_update_transactions=_team_member_only_transactions( + f"team_id::{team_id}::user_id::{removed_user_id}", response_cost + ), + ) + + assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) + mock_batcher.litellm_teammembership.upsert.assert_not_called() + mock_batcher.litellm_teammembership.update_many.assert_called_once_with( + where={"team_id": team_id, "user_id": removed_user_id}, + data={ + "spend": {"increment": response_cost}, + "total_spend": {"increment": response_cost}, + }, + ) + + @pytest.mark.asyncio async def test_org_spend_increments_organization_membership_row_for_the_calling_user(): """A request made with a user_id inside an org must increment that user's @@ -2232,13 +2273,9 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): "team_id::team_b::user_id::user_x": 0.3, }, "litellm_teammembership", - "upsert", - "user_id_team_id", - [ - {"user_id": "user_x", "team_id": "team_a"}, - {"user_id": "user_x", "team_id": "team_b"}, - {"user_id": "user_x", "team_id": "team_c"}, - ], + "update_many", + "team_id", + ["team_a", "team_b", "team_c"], id="team_member", ), pytest.param( @@ -2312,6 +2349,8 @@ async def test_commit_spend_updates_iterates_in_sorted_order( ) ) + mock_transaction.query_raw = AsyncMock(return_value=[]) + mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) @@ -3071,6 +3110,7 @@ def _good_tx(mock_batcher): tx = AsyncMock() tx.__aenter__ = AsyncMock(return_value=tx) tx.__aexit__ = AsyncMock(return_value=False) + tx.query_raw = AsyncMock(return_value=[]) tx.batch_ = MagicMock( return_value=AsyncMock( __aenter__=AsyncMock(return_value=mock_batcher), From 7c068a4cf7772b45246568346e36cef7e4f27ec2 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:16:11 +0000 Subject: [PATCH 200/253] fix(repositories): import LiteralString from typing_extensions for python 3.10 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/repositories/prisma_protocols.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index c742f3f5f3f..c7421c9284a 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -7,7 +7,9 @@ private ones per file. """ from collections.abc import Mapping, Sequence -from typing import LiteralString, Protocol, TypeVar +from typing import Protocol, TypeVar + +from typing_extensions import LiteralString RowT_co = TypeVar("RowT_co", covariant=True) From efef6ab68491299a550739286990cf922330dd89 Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 02:53:47 +0000 Subject: [PATCH 201/253] fix(proxy): write team member spend as one roster checked upsert statement Replaces the per team advisory lock and Pydantic roster parse in the spend flush with a single INSERT ... ON CONFLICT statement that checks the stored roster in SQL, so malformed roster JSON cannot fail the whole flush and large batches no longer issue two queries per team inside the fixed transaction deadline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 95 +++++-------- .../access_group_team_sync.py | 8 +- litellm/repositories/prisma_protocols.py | 10 -- litellm/repositories/team_repository.py | 12 +- .../proxy/db/test_db_spend_update_writer.py | 126 ++++++------------ 5 files changed, 87 insertions(+), 164 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 5c86c52c539..ecf107ef4be 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -18,7 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload from urllib.parse import quote, unquote -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import LiteralString, ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -75,8 +75,7 @@ from litellm.proxy.spend_tracking.savings import ( marks_gateway_injection, ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error -from litellm.repositories.prisma_protocols import BatchTable, RawQueryTransaction -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL, TeamRepository +from litellm.repositories.prisma_protocols import BatchTable from litellm.types.utils import CallTypes if TYPE_CHECKING: @@ -134,9 +133,11 @@ class _SpendBatchManager(Protocol): async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ... -class _SpendTransaction(RawQueryTransaction, Protocol): +class _SpendTransaction(Protocol): def batch_(self) -> _SpendBatchManager: ... + async def execute_raw(self, query: LiteralString, *args: object) -> int: ... + class _SpendTransactionManager(Protocol): async def __aenter__(self) -> _SpendTransaction: ... @@ -162,47 +163,36 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx -async def _lock_and_read_rosters( - prisma_client: PrismaClient, transaction: _SpendTransaction, team_ids: Sequence[str] -) -> frozenset[tuple[str, str]]: - """Take each team's advisory lock on ``transaction`` and return its rostered (user_id, team_id) pairs. - - A spend flush can land after ``/team/member_delete`` removed the member. Holding the same - lock that endpoint takes, until this transaction commits, means a member read here is on - the team for the whole flush, so only they may have a missing membership row created. - """ - repository: Final = TeamRepository(prisma_client) - - async def locked_roster(team_id: str) -> tuple[tuple[str, str], ...]: - await transaction.query_raw(TEAM_ADVISORY_LOCK_SQL, team_id) - roster: Final = await repository.get_members_with_roles_locked(transaction, team_id) - return tuple((member.user_id, team_id) for member in roster or () if member.user_id is not None) - - rosters: Final = tuple([await locked_roster(team_id) for team_id in sorted(frozenset(team_ids))]) - return frozenset(pair for roster in rosters for pair in roster) +# One statement adds every member's cost to their membership row. A missing row is created only +# while the user is still on the team's roster; FOR SHARE on the team row makes that check and +# the insert atomic against /team/member_delete and /team/delete, which update or delete that row +# before removing memberships, so a spend flush landing after a removal never recreates the member. +_TEAM_MEMBER_SPEND_SQL: Final = """ +INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend) +SELECT p.user_id, p.team_id, p.cost, p.cost +FROM unnest($1::text[], $2::text[], $3::float8[]) AS p(user_id, team_id, cost) +WHERE EXISTS ( + SELECT 1 FROM "LiteLLM_TeamTable" t + WHERE t.team_id = p.team_id + AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id)) + FOR SHARE +) + OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id) +ON CONFLICT (user_id, team_id) DO UPDATE +SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, + total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend +""" -def _queue_team_member_spend( - memberships: BatchTable, user_id: str, team_id: str, response_cost: float, rostered: bool -) -> None: - increments: Final = { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - } - if not rostered: - memberships.update_many(where={"team_id": team_id, "user_id": user_id}, data=increments) - return - memberships.upsert( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={ - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": increments, - }, +async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: + # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), + # keeping lock order consistent across pods to prevent deadlocks + keys: Final = tuple(sorted(spend_by_member_key)) + _ = await transaction.execute_raw( + _TEAM_MEMBER_SPEND_SQL, + [key.split("::")[3] for key in keys], + [key.split("::")[1] for key in keys], + [spend_by_member_key[key] for key in keys], ) @@ -1730,24 +1720,7 @@ class DBSpendUpdateWriter: start_time = time.time() try: async with _spend_update_tx(prisma_client) as transaction: - rostered_members = await _lock_and_read_rosters( - prisma_client, transaction, tuple(team_id for _, team_id in team_memberships_to_invalidate) - ) - async with transaction.batch_() as batcher: - # Sort by composite key for consistent lock ordering across pods to prevent deadlocks. - # Key format "team_id::::user_id::" makes the string sort equivalent to sorting by (team_id, user_id). - for key, response_cost in sorted(team_member_list_transactions.items()): - # key is "team_id::::user_id::" - team_id = key.split("::")[1] - user_id = key.split("::")[3] - - _queue_team_member_spend( - batcher.litellm_teammembership, - user_id, - team_id, - response_cost, - (user_id, team_id) in rostered_members, - ) + await _write_team_member_spend(transaction, team_member_list_transactions) # Transaction succeeded, break out of retry loop break except Exception as e: diff --git a/litellm/proxy/management_helpers/access_group_team_sync.py b/litellm/proxy/management_helpers/access_group_team_sync.py index cfe207e66ea..664e36c9f10 100644 --- a/litellm/proxy/management_helpers/access_group_team_sync.py +++ b/litellm/proxy/management_helpers/access_group_team_sync.py @@ -21,7 +21,13 @@ from typing import Final, Protocol from pydantic import BaseModel, TypeAdapter from litellm.proxy.auth.auth_checks import _delete_cache_access_object -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL + +# hashtext collisions only cost two unrelated teams a little serialization, and the +# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, +# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints +# reuses this exact statement to serialize /team/member_add and /team/delete against each +# other and against this mirror, rather than defining a second, divergent lock on the same key. +TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" _READ_TEAM_SQL: Final = 'SELECT access_group_ids FROM "LiteLLM_TeamTable" WHERE team_id = $1' diff --git a/litellm/repositories/prisma_protocols.py b/litellm/repositories/prisma_protocols.py index c7421c9284a..93b8c5c7cd7 100644 --- a/litellm/repositories/prisma_protocols.py +++ b/litellm/repositories/prisma_protocols.py @@ -9,8 +9,6 @@ private ones per file. from collections.abc import Mapping, Sequence from typing import Protocol, TypeVar -from typing_extensions import LiteralString - RowT_co = TypeVar("RowT_co", covariant=True) @@ -110,12 +108,6 @@ class PrismaRecord(Protocol): def dict(self) -> Mapping[str, object]: ... -class RawQueryTransaction(Protocol): - """A prisma transaction handle that can run raw SQL, e.g. an advisory lock or a locked read.""" - - async def query_raw(self, query: LiteralString, *args: str) -> Sequence[Mapping[str, object]]: ... - - class ReadOnlyTable(Protocol): async def find_many(self, *, where: Mapping[str, object]) -> Sequence[PrismaRecord]: ... @@ -131,8 +123,6 @@ class BatchTable(Protocol): def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> None: ... - def upsert(self, *, where: Mapping[str, object], data: Mapping[str, Mapping[str, object]]) -> None: ... - class PrismaBatch(Protocol): @property diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 1b9cba48e41..5ff07d76b5d 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -15,9 +15,10 @@ from litellm.repositories.base_repository import ( DbRecord, record_to_dict, ) -from litellm.repositories.prisma_protocols import RawQueryTransaction, TableActions +from litellm.repositories.prisma_protocols import TableActions if TYPE_CHECKING: + from prisma import Prisma from prisma import models as prisma_models @@ -39,13 +40,6 @@ def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays: return team -# hashtext collisions only cost two unrelated teams a little serialization, and the -# lock is never taken by the access-group endpoints as a SELECT ... FOR UPDATE row lock, -# so it cannot join their access-group-then-team lock order to form a cycle. team_endpoints, -# the access-group mirror and the team member spend flush all reuse this exact statement to -# serialize against each other, rather than defining a second, divergent lock on the same key. -TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" - _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( "metadata", @@ -84,7 +78,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable.model_validate(data) - async def get_members_with_roles_locked(self, tx: RawQueryTransaction, team_id: str) -> list[Member] | None: + async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> list[Member] | None: """Return the team's members_with_roles. The caller must already hold ``TEAM_ADVISORY_LOCK_SQL`` for this team_id on ``tx`` before calling this. diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b12dc68b7b1..05769eea343 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -16,11 +16,10 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType -from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter +from litellm.proxy.db.db_spend_update_writer import _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) -from litellm.repositories.team_repository import TEAM_ADVISORY_LOCK_SQL @pytest.mark.asyncio @@ -914,42 +913,26 @@ async def test_commit_spend_updates_to_db_increments_agent_spend(): assert call_kwargs["data"] == {"spend": {"increment": response_cost}} -def _team_member_flush_fixtures(team_id: str, rostered_user_ids: list[str]) -> tuple[MagicMock, AsyncMock, MagicMock]: - """A batcher, transaction and prisma client whose locked roster read for `team_id` lists `rostered_user_ids`.""" - mock_batcher = MagicMock() - mock_batcher.litellm_teammembership = MagicMock() - mock_batcher.litellm_teammembership.upsert = MagicMock() - mock_batcher.litellm_teammembership.update_many = MagicMock() - - roster_row = {"members_with_roles": json.dumps([{"user_id": uid, "role": "user"} for uid in rostered_user_ids])} - - async def query_raw(query: str, *args: str) -> list[dict[str, object]]: - return [] if query == TEAM_ADVISORY_LOCK_SQL else [roster_row] - +def _team_member_flush_fixtures() -> tuple[AsyncMock, MagicMock]: + """A transaction and prisma client that record the raw statement the member spend flush runs.""" mock_transaction = AsyncMock() mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction) mock_transaction.__aexit__ = AsyncMock(return_value=False) - mock_transaction.query_raw = AsyncMock(side_effect=query_raw) - mock_transaction.batch_ = MagicMock( - return_value=AsyncMock( - __aenter__=AsyncMock(return_value=mock_batcher), - __aexit__=AsyncMock(return_value=False), - ) - ) + mock_transaction.execute_raw = AsyncMock(return_value=1) mock_prisma_client = MagicMock() mock_prisma_client.db = MagicMock() mock_prisma_client.db.tx = MagicMock(return_value=mock_transaction) - return mock_batcher, mock_transaction, mock_prisma_client + return mock_transaction, mock_prisma_client -def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict[str, dict[str, float]]: +def _team_member_only_transactions(spend_by_member_key: dict[str, float]) -> dict[str, dict[str, float]]: return { "user_list_transactions": {}, "end_user_list_transactions": {}, "key_list_transactions": {}, "team_list_transactions": {}, - "team_member_list_transactions": {entity_id: response_cost}, + "team_member_list_transactions": spend_by_member_key, "org_list_transactions": {}, "tag_list_transactions": {}, "agent_list_transactions": {}, @@ -957,22 +940,20 @@ def _team_member_only_transactions(entity_id: str, response_cost: float) -> dict @pytest.mark.asyncio -async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total_spend(): +async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster_checked_upsert(): """ - Verify that _commit_spend_updates_to_db increments BOTH spend (cycle-scoped) - and total_spend (non-resetting) on LiteLLM_TeamMembership in a single - upsert call, using the same response_cost. + Regression (LIT-5502): members added without a budget had no membership row, and the + previous update_many matched zero rows, so their spend was silently dropped. - Regression (LIT-5502): members added without a budget had no membership row, and - the previous update_many matched zero rows, so their spend was silently dropped. - For a member still on the team roster the upsert has to create the row seeded - with this call's cost in that case. + The flush now runs one INSERT ... ON CONFLICT statement for the whole batch that adds + the cost to both spend and total_spend and creates the missing row for a user still on + the team roster, so no per-team read can fail or time out ahead of the writes. """ db_writer = DBSpendUpdateWriter() team_id = "team-abc" user_id = "user-xyz" response_cost = 0.75 - mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, [user_id]) + mock_transaction, mock_prisma_client = _team_member_flush_fixtures() mock_proxy_logging = MagicMock() mock_proxy_logging.call_details.get = MagicMock(return_value=None) @@ -982,42 +963,31 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total n_retry_times=0, proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( - f"team_id::{team_id}::user_id::{user_id}", response_cost + {f"team_id::{team_id}::user_id::{user_id}": response_cost} ), ) - assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) - mock_batcher.litellm_teammembership.upsert.assert_called_once() - mock_batcher.litellm_teammembership.update_many.assert_not_called() - call_kwargs = mock_batcher.litellm_teammembership.upsert.call_args.kwargs - assert call_kwargs["where"] == {"user_id_team_id": {"user_id": user_id, "team_id": team_id}} - assert call_kwargs["data"] == { - "create": { - "team_id": team_id, - "user_id": user_id, - "spend": response_cost, - "total_spend": response_cost, - }, - "update": { - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - } + mock_transaction.execute_raw.assert_awaited_once() + statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + assert statement is _TEAM_MEMBER_SPEND_SQL + assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) + assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement + assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement + assert "FOR SHARE" in statement + assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement + assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement + assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement @pytest.mark.asyncio -async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_removed_team_member(): +async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user(): """ - A spend flush that lands after /team/member_delete must not resurrect the deleted - membership row: a user missing from the team roster, read under the team's advisory - lock inside the flush transaction, only gets an increment on whatever row still - exists, never a create. + The single member spend statement locks rows in the order of its input arrays, so the + batch is handed over sorted by (team_id, user_id), with each cost kept next to its + member, so concurrent pods lock in the same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() - team_id = "team-abc" - removed_user_id = "user-removed" - response_cost = 0.75 - mock_batcher, mock_transaction, mock_prisma_client = _team_member_flush_fixtures(team_id, ["user-still-here"]) + mock_transaction, mock_prisma_client = _team_member_flush_fixtures() mock_proxy_logging = MagicMock() mock_proxy_logging.call_details.get = MagicMock(return_value=None) @@ -1027,19 +997,22 @@ async def test_commit_spend_updates_to_db_does_not_recreate_membership_of_remove n_retry_times=0, proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( - f"team_id::{team_id}::user_id::{removed_user_id}", response_cost + { + "team_id::team_c::user_id::user_x": 0.1, + "team_id::team_a::user_id::user_y": 0.2, + "team_id::team_a::user_id::user_x": 0.3, + "team_id::team_b::user_id::user_x": 0.4, + } ), ) - assert mock_transaction.query_raw.await_args_list[0] == call(TEAM_ADVISORY_LOCK_SQL, team_id) - mock_batcher.litellm_teammembership.upsert.assert_not_called() - mock_batcher.litellm_teammembership.update_many.assert_called_once_with( - where={"team_id": team_id, "user_id": removed_user_id}, - data={ - "spend": {"increment": response_cost}, - "total_spend": {"increment": response_cost}, - }, - ) + _statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + assert list(zip(team_ids, user_ids, costs)) == [ + ("team_a", "user_x", 0.3), + ("team_a", "user_y", 0.2), + ("team_b", "user_x", 0.4), + ("team_c", "user_x", 0.1), + ] @pytest.mark.asyncio @@ -2265,19 +2238,6 @@ async def test_commit_daily_tag_spend_no_requeue_on_success(): ["team_a", "team_b", "team_c"], id="team", ), - pytest.param( - "team_member_list_transactions", - { - "team_id::team_c::user_id::user_x": 0.1, - "team_id::team_a::user_id::user_x": 0.2, - "team_id::team_b::user_id::user_x": 0.3, - }, - "litellm_teammembership", - "update_many", - "team_id", - ["team_a", "team_b", "team_c"], - id="team_member", - ), pytest.param( "org_list_transactions", {"org_c": 0.1, "org_a": 0.2, "org_b": 0.3}, From 60077e90aabfcb52c9babeb19c28efb3ad0c1fdd Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 03:16:54 +0000 Subject: [PATCH 202/253] fix(proxy): take the team advisory lock before the member spend flush Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 18 +++++++--- .../proxy/db/test_db_spend_update_writer.py | 34 +++++++++++++------ 2 files changed, 36 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index ecf107ef4be..f617302d8ad 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -163,10 +163,17 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: return tx +# The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL), +# taken in sorted order so the roster check below cannot interleave with their writes. A row lock would +# deadlock with the access-group endpoints, which lock a team row after an access-group lock. +_TEAM_ADVISORY_LOCKS_SQL: Final = """ +SELECT pg_advisory_xact_lock(hashtext(teams.team_id)) +FROM (SELECT DISTINCT team_id FROM unnest($1::text[]) AS team_id ORDER BY team_id) AS teams +""" + # One statement adds every member's cost to their membership row. A missing row is created only -# while the user is still on the team's roster; FOR SHARE on the team row makes that check and -# the insert atomic against /team/member_delete and /team/delete, which update or delete that row -# before removing memberships, so a spend flush landing after a removal never recreates the member. +# while the user is still on the team's roster, so a spend flush landing after a removal never +# recreates the member. _TEAM_MEMBER_SPEND_SQL: Final = """ INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend) SELECT p.user_id, p.team_id, p.cost, p.cost @@ -175,7 +182,6 @@ WHERE EXISTS ( SELECT 1 FROM "LiteLLM_TeamTable" t WHERE t.team_id = p.team_id AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id)) - FOR SHARE ) OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id) ON CONFLICT (user_id, team_id) DO UPDATE @@ -188,10 +194,12 @@ async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_memb # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), # keeping lock order consistent across pods to prevent deadlocks keys: Final = tuple(sorted(spend_by_member_key)) + team_ids: Final = [key.split("::")[1] for key in keys] + _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCKS_SQL, team_ids) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, [key.split("::")[3] for key in keys], - [key.split("::")[1] for key in keys], + team_ids, [spend_by_member_key[key] for key in keys], ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 05769eea343..9f3bc5acd92 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -16,7 +16,11 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType -from litellm.proxy.db.db_spend_update_writer import _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter +from litellm.proxy.db.db_spend_update_writer import ( + _TEAM_ADVISORY_LOCKS_SQL, + _TEAM_MEMBER_SPEND_SQL, + DBSpendUpdateWriter, +) from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( build_window_spend_transaction, ) @@ -945,9 +949,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster Regression (LIT-5502): members added without a budget had no membership row, and the previous update_many matched zero rows, so their spend was silently dropped. - The flush now runs one INSERT ... ON CONFLICT statement for the whole batch that adds - the cost to both spend and total_spend and creates the missing row for a user still on - the team roster, so no per-team read can fail or time out ahead of the writes. + The flush now takes the same per-team advisory lock the team endpoints hold, then runs + one INSERT ... ON CONFLICT statement for the whole batch that adds the cost to both spend + and total_spend and creates the missing row for a user still on the team roster, so no + per-team read can fail or time out ahead of the writes. """ db_writer = DBSpendUpdateWriter() team_id = "team-abc" @@ -967,13 +972,17 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster ), ) - mock_transaction.execute_raw.assert_awaited_once() - statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + lock_call, spend_call = mock_transaction.execute_raw.await_args_list + lock_statement, locked_team_ids = lock_call.args + assert lock_statement is _TEAM_ADVISORY_LOCKS_SQL + assert locked_team_ids == [team_id] + assert "pg_advisory_xact_lock(hashtext(teams.team_id))" in lock_statement + assert "ORDER BY team_id" in lock_statement + statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement - assert "FOR SHARE" in statement assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement @@ -982,9 +991,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster @pytest.mark.asyncio async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user(): """ - The single member spend statement locks rows in the order of its input arrays, so the - batch is handed over sorted by (team_id, user_id), with each cost kept next to its - member, so concurrent pods lock in the same order and cannot deadlock. + The member spend statement touches rows in the order of its input arrays, so the batch + is handed over sorted by (team_id, user_id), with each cost kept next to its member, and + the advisory lock statement receives the same team ids, so concurrent pods lock in the + same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() mock_transaction, mock_prisma_client = _team_member_flush_fixtures() @@ -1006,7 +1016,9 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u ), ) - _statement, user_ids, team_ids, costs = mock_transaction.execute_raw.await_args.args + lock_call, spend_call = mock_transaction.execute_raw.await_args_list + _statement, user_ids, team_ids, costs = spend_call.args + assert lock_call.args == (_TEAM_ADVISORY_LOCKS_SQL, ["team_a", "team_a", "team_b", "team_c"]) assert list(zip(team_ids, user_ids, costs)) == [ ("team_a", "user_x", 0.3), ("team_a", "user_y", 0.2), From 499c334fce491fa67fc1db92e5e1c160cd00b33d Mon Sep 17 00:00:00 2001 From: ryan Date: Wed, 16 Sep 2026 03:35:34 +0000 Subject: [PATCH 203/253] fix(proxy): lock each team in sorted order before the member spend upsert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 12 +++++------ .../proxy/db/test_db_spend_update_writer.py | 21 +++++++++++-------- 2 files changed, 17 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index f617302d8ad..09ce6b6f451 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -164,12 +164,9 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager: # The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL), -# taken in sorted order so the roster check below cannot interleave with their writes. A row lock would -# deadlock with the access-group endpoints, which lock a team row after an access-group lock. -_TEAM_ADVISORY_LOCKS_SQL: Final = """ -SELECT pg_advisory_xact_lock(hashtext(teams.team_id)) -FROM (SELECT DISTINCT team_id FROM unnest($1::text[]) AS team_id ORDER BY team_id) AS teams -""" +# so the roster check below cannot interleave with their writes. A row lock would deadlock with the +# access-group endpoints, which lock a team row after an access-group lock. +_TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked" # One statement adds every member's cost to their membership row. A missing row is created only # while the user is still on the team's roster, so a spend flush landing after a removal never @@ -195,7 +192,8 @@ async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_memb # keeping lock order consistent across pods to prevent deadlocks keys: Final = tuple(sorted(spend_by_member_key)) team_ids: Final = [key.split("::")[1] for key in keys] - _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCKS_SQL, team_ids) + for team_id in dict.fromkeys(team_ids): + _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, [key.split("::")[3] for key in keys], diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 9f3bc5acd92..2ac9fc8eccd 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -17,7 +17,7 @@ from redis.exceptions import DataError import litellm from litellm.proxy._types import Litellm_EntityType from litellm.proxy.db.db_spend_update_writer import ( - _TEAM_ADVISORY_LOCKS_SQL, + _TEAM_ADVISORY_LOCK_SQL, _TEAM_MEMBER_SPEND_SQL, DBSpendUpdateWriter, ) @@ -973,11 +973,10 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster ) lock_call, spend_call = mock_transaction.execute_raw.await_args_list - lock_statement, locked_team_ids = lock_call.args - assert lock_statement is _TEAM_ADVISORY_LOCKS_SQL - assert locked_team_ids == [team_id] - assert "pg_advisory_xact_lock(hashtext(teams.team_id))" in lock_statement - assert "ORDER BY team_id" in lock_statement + lock_statement, locked_team_id = lock_call.args + assert lock_statement is _TEAM_ADVISORY_LOCK_SQL + assert locked_team_id == team_id + assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) @@ -993,7 +992,7 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u """ The member spend statement touches rows in the order of its input arrays, so the batch is handed over sorted by (team_id, user_id), with each cost kept next to its member, and - the advisory lock statement receives the same team ids, so concurrent pods lock in the + each distinct team is locked once, in that same order, so concurrent pods lock in the same order and cannot deadlock. """ db_writer = DBSpendUpdateWriter() @@ -1016,9 +1015,13 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u ), ) - lock_call, spend_call = mock_transaction.execute_raw.await_args_list + *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list _statement, user_ids, team_ids, costs = spend_call.args - assert lock_call.args == (_TEAM_ADVISORY_LOCKS_SQL, ["team_a", "team_a", "team_b", "team_c"]) + assert [lock_call.args for lock_call in lock_calls] == [ + (_TEAM_ADVISORY_LOCK_SQL, "team_a"), + (_TEAM_ADVISORY_LOCK_SQL, "team_b"), + (_TEAM_ADVISORY_LOCK_SQL, "team_c"), + ] assert list(zip(team_ids, user_ids, costs)) == [ ("team_a", "user_x", 0.3), ("team_a", "user_y", 0.2), From 2d61fa66b165d8e5e9b4295bba2587346709aab9 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 00:37:03 +0000 Subject: [PATCH 204/253] fix(proxy): lock teams in sorted team id order and keep existing member budgets on re-add Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 12 ++++----- litellm/proxy/management_helpers/utils.py | 2 +- .../proxy/db/test_db_spend_update_writer.py | 27 ++++++++++--------- .../test_management_helpers_utils.py | 3 +++ 4 files changed, 24 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 09ce6b6f451..77d16e90706 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -188,17 +188,17 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: - # key is "team_id::::user_id::"; the string sort orders rows by (team_id, user_id), - # keeping lock order consistent across pods to prevent deadlocks - keys: Final = tuple(sorted(spend_by_member_key)) - team_ids: Final = [key.split("::")[1] for key in keys] + # key is "team_id::::user_id::"; rows are sorted by (team_id, user_id) so the teams are + # locked in the same `sorted(team_ids)` order the team endpoints use, preventing deadlocks + rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items()) + team_ids: Final = [team_id for team_id, _user_id, _cost in rows] for team_id in dict.fromkeys(team_ids): _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, - [key.split("::")[3] for key in keys], + [user_id for _team_id, user_id, _cost in rows], team_ids, - [spend_by_member_key[key] for key in keys], + [cost for _team_id, _user_id, cost in rows], ) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index d940eba86d3..d84d32431ec 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -479,7 +479,7 @@ async def add_new_member( budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} _returned_team_membership: Final = await membership_table.upsert( where={"user_id_team_id": membership_key}, - data={"create": {**membership_key, **budget_link}, "update": {**budget_link}}, + data={"create": {**membership_key, **budget_link}, "update": {}}, include={"litellm_budget_table": True}, ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 2ac9fc8eccd..70199ee0c73 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -992,8 +992,9 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u """ The member spend statement touches rows in the order of its input arrays, so the batch is handed over sorted by (team_id, user_id), with each cost kept next to its member, and - each distinct team is locked once, in that same order, so concurrent pods lock in the - same order and cannot deadlock. + each distinct team is locked once, in `sorted(team_ids)` order, the order /team/delete + locks in, so a concurrent flush and delete cannot deadlock. `eng` and `eng2` pin that: + sorting the composite keys instead would lock `eng2` first because `2` < `:`. """ db_writer = DBSpendUpdateWriter() mock_transaction, mock_prisma_client = _team_member_flush_fixtures() @@ -1007,10 +1008,10 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u proxy_logging_obj=mock_proxy_logging, db_spend_update_transactions=_team_member_only_transactions( { - "team_id::team_c::user_id::user_x": 0.1, - "team_id::team_a::user_id::user_y": 0.2, - "team_id::team_a::user_id::user_x": 0.3, - "team_id::team_b::user_id::user_x": 0.4, + "team_id::eng2::user_id::user_x": 0.1, + "team_id::eng::user_id::user_y": 0.2, + "team_id::eng::user_id::user_x": 0.3, + "team_id::eng-b::user_id::user_x": 0.4, } ), ) @@ -1018,15 +1019,15 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u *lock_calls, spend_call = mock_transaction.execute_raw.await_args_list _statement, user_ids, team_ids, costs = spend_call.args assert [lock_call.args for lock_call in lock_calls] == [ - (_TEAM_ADVISORY_LOCK_SQL, "team_a"), - (_TEAM_ADVISORY_LOCK_SQL, "team_b"), - (_TEAM_ADVISORY_LOCK_SQL, "team_c"), + (_TEAM_ADVISORY_LOCK_SQL, "eng"), + (_TEAM_ADVISORY_LOCK_SQL, "eng-b"), + (_TEAM_ADVISORY_LOCK_SQL, "eng2"), ] assert list(zip(team_ids, user_ids, costs)) == [ - ("team_a", "user_x", 0.3), - ("team_a", "user_y", 0.2), - ("team_b", "user_x", 0.4), - ("team_c", "user_x", 0.1), + ("eng", "user_x", 0.3), + ("eng", "user_y", 0.2), + ("eng-b", "user_x", 0.4), + ("eng2", "user_x", 0.1), ] diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index 0dcf7e1b8ad..09d5f684f3d 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -446,6 +446,8 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): 1. When max_budget_in_team is provided 2. A new budget is created in the litellm_budgettable 3. The new budget_id is used for the team membership + 4. The upsert's update branch stays empty, so a bulk /team/member_add that names a member + already on the team does not replace the budget_id (and the spend) their existing row carries """ from litellm.proxy._types import LitellmUserRoles @@ -529,6 +531,7 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): assert team_membership_call_args is not None create_data = team_membership_call_args.kwargs["data"]["create"] assert create_data["budget_id"] == test_new_budget_id + assert team_membership_call_args.kwargs["data"]["update"] == {} @pytest.mark.asyncio From b6f4ad190e00a21935479091a423e16e94530778 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 00:49:48 +0000 Subject: [PATCH 205/253] refactor(proxy): freeze the member spend arrays and budget link to stay within the type discipline budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 9 ++++----- litellm/proxy/management_helpers/utils.py | 4 +++- .../test_litellm/proxy/db/test_db_spend_update_writer.py | 2 +- 3 files changed, 8 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 77d16e90706..8bcfe28488e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -188,17 +188,16 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend, async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None: - # key is "team_id::::user_id::"; rows are sorted by (team_id, user_id) so the teams are - # locked in the same `sorted(team_ids)` order the team endpoints use, preventing deadlocks + # key is "team_id::::user_id::"; locks are taken in sorted team_id order like the team endpoints rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items()) - team_ids: Final = [team_id for team_id, _user_id, _cost in rows] + team_ids: Final = tuple(team_id for team_id, _user_id, _cost in rows) for team_id in dict.fromkeys(team_ids): _ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id) _ = await transaction.execute_raw( _TEAM_MEMBER_SPEND_SQL, - [user_id for _team_id, user_id, _cost in rows], + tuple(user_id for _team_id, user_id, _cost in rows), team_ids, - [cost for _team_id, _user_id, cost in rows], + tuple(cost for _team_id, _user_id, cost in rows), ) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index d84d32431ec..c10e5f9b23d 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -476,7 +476,9 @@ async def add_new_member( if returned_user is not None and returned_user.user_id is not None: membership_table: Final[_PrismaTeamMembershipTable] = _team_membership_table(prisma_client, tx) membership_key: Final[Mapping[str, object]] = {"user_id": returned_user.user_id, "team_id": team_id} - budget_link: Final[Mapping[str, object]] = {"budget_id": _budget_id} if _budget_id is not None else {} + budget_link: Final[Mapping[str, str]] = ( + MappingProxyType({"budget_id": _budget_id}) if _budget_id is not None else MappingProxyType({}) + ) _returned_team_membership: Final = await membership_table.upsert( where={"user_id_team_id": membership_key}, data={"create": {**membership_key, **budget_link}, "update": {}}, diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 70199ee0c73..e38a65f2b12 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -979,7 +979,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement statement, user_ids, team_ids, costs = spend_call.args assert statement is _TEAM_MEMBER_SPEND_SQL - assert (user_ids, team_ids, costs) == ([user_id], [team_id], [response_cost]) + assert (list(user_ids), list(team_ids), list(costs)) == ([user_id], [team_id], [response_cost]) assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement From 61d4c5b9b5b427929f2d97d216f3aafb611f0abb Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 17 Sep 2026 01:12:54 +0000 Subject: [PATCH 206/253] fix(proxy): skip members already on the team before resolving a per-member budget A mixed /team/member_add list that names an existing member used to run add_new_member for them, which created or cloned a budget that the empty upsert update branch never linked to their membership row. Filter the requested members against the freshly locked roster first so budgets and membership rows are only written for members who are actually new Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/team_endpoints.py | 34 +++------ .../test_team_endpoints.py | 71 +++++++++++++++++++ 2 files changed, 79 insertions(+), 26 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 4957b3bd925..28c12173ea7 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2904,10 +2904,15 @@ async def _process_team_members( if member_allowed_models is None and team_default_member_models: member_allowed_models = team_default_member_models - if isinstance(data.member, Member): + requested_members: Final[Sequence[Member]] = ( + (data.member,) if isinstance(data.member, Member) else tuple(data.member) + ) + for m in requested_members: + if _member_already_in_team(m, complete_team_data): + continue try: updated_user, updated_tm = await add_new_member( - new_member=data.member, + new_member=m, max_budget_in_team=data.max_budget_in_team, prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, @@ -2921,34 +2926,11 @@ async def _process_team_members( except Exception as e: raise HTTPException( status_code=500, - detail={"error": f"Unable to add user - {data.member}, to team - {data.team_id}, for reason - {e}"}, + detail={"error": f"Unable to add user - {m}, to team - {data.team_id}, for reason - {e}"}, ) updated_users.append(updated_user) if updated_tm is not None: updated_team_memberships.append(updated_tm) - elif isinstance(data.member, list): - for m in data.member: - try: - updated_user, updated_tm = await add_new_member( - new_member=m, - max_budget_in_team=data.max_budget_in_team, - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - litellm_proxy_admin_name=litellm_proxy_admin_name, - team_id=data.team_id, - default_team_budget_id=default_team_budget_id, - allowed_models=member_allowed_models, - budget_duration=data.budget_duration, - tx=tx, - ) - except Exception as e: - raise HTTPException( - status_code=500, - detail={"error": f"Unable to add user - {m}, to team - {data.team_id}, for reason - {e}"}, - ) - updated_users.append(updated_user) - if updated_tm is not None: - updated_team_memberships.append(updated_tm) return updated_users, updated_team_memberships 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 f7157da3100..690b5ae80b6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -1794,6 +1794,7 @@ async def test_process_team_members_single_member(): mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = {"team_member_budget_id": "budget-123"} mock_team.default_team_member_models = None + mock_team.members_with_roles = [] # Mock user and membership objects mock_user = MagicMock(spec=LiteLLM_UserTable) @@ -1854,6 +1855,7 @@ async def test_process_team_members_multiple_members(): mock_team = MagicMock(spec=LiteLLM_TeamTable) mock_team.metadata = None mock_team.default_team_member_models = None + mock_team.members_with_roles = [] # Create multiple members as dictionaries (they will be converted to Member objects) members = [ @@ -2114,6 +2116,75 @@ async def test_add_team_members_runs_member_writes_on_the_lock_holding_transacti assert [tm.budget_id for tm in updated_team_memberships] == ["budget-pool"] +@pytest.mark.asyncio +async def test_add_team_members_skips_budget_and_membership_writes_for_members_already_on_the_roster(): + """ + Regression pin for orphaned budgets on a mixed /team/member_add list. + + A list naming one member already on the team and one new member must only create a + budget and membership row for the new member. Running add_new_member for the existing + member would create a per-member budget that nothing links to, since their membership + row (and the budget it already carries) is left untouched. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + _add_team_members_to_team, + ) + + added_user = MagicMock() + added_user.user_id = "bob" + added_user.model_dump.return_value = {"user_id": "bob", "teams": ["team-mixed"]} + created_budget = MagicMock() + created_budget.budget_id = "budget-bob" + membership = MagicMock() + membership.model_dump.return_value = { + "team_id": "team-mixed", + "user_id": "bob", + "budget_id": "budget-bob", + "litellm_budget_table": None, + } + + tx = MagicMock() + tx.query_raw = AsyncMock( + return_value=[{"members_with_roles": [{"user_id": "alice", "user_email": None, "role": "user"}]}] + ) + tx.litellm_teamtable.update = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="team-mixed", members_with_roles=[]) + ) + tx.litellm_usertable.upsert = AsyncMock(return_value=added_user) + tx.litellm_usertable.update_many = AsyncMock() + tx.litellm_budgettable.create = AsyncMock(return_value=created_budget) + tx.litellm_teammembership.upsert = AsyncMock(return_value=membership) + + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx) + tx_cm.__aexit__ = AsyncMock(return_value=None) + + prisma_client = MagicMock() + prisma_client.tx = MagicMock(return_value=tx_cm) + + _, updated_users, updated_team_memberships = await _add_team_members_to_team( + data=TeamMemberAddRequest( + team_id="team-mixed", + member=[Member(user_id="alice", role="user"), Member(user_id="bob", role="user")], + max_budget_in_team=50.0, + ), + complete_team_data=LiteLLM_TeamTable(team_id="team-mixed", members_with_roles=[]), + prisma_client=cast(object, prisma_client), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + litellm_proxy_admin_name="admin", + ) + + tx.litellm_budgettable.create.assert_awaited_once() + tx.litellm_teammembership.upsert.assert_awaited_once() + assert tx.litellm_teammembership.upsert.call_args.kwargs["where"] == { + "user_id_team_id": {"user_id": "bob", "team_id": "team-mixed"} + } + assert [user.user_id for user in updated_users] == ["bob"] + assert [tm.user_id for tm in updated_team_memberships] == ["bob"] + written_ids = [m["user_id"] for m in json.loads(tx.litellm_teamtable.update.call_args.kwargs["data"]["members_with_roles"])] + assert written_ids == ["alice", "bob"] + + @pytest.mark.asyncio async def test_add_team_members_writes_nothing_when_the_team_is_deleted_mid_request(): """ From c21e86a4436abdd48e0ad83cd3284adc35a50ea7 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 12:27:39 -0700 Subject: [PATCH 207/253] feat(ui): reset a team member's spend from the Members tab Every member now carries a membership row, so a member who spent with no budget is over budget the moment a member budget is added later. The only fix was POST /team/{team_id}/member/{user_id}/reset_spend, which had no UI. The Members tab gets a Reset spend action on rows that have current cycle spend. It confirms in a dialog, posts reset_to 0 through the typed client, and refreshes the team without remounting the page so the tab stays open. A team admin does not see it on their own row because the backend rejects that reset --- .../hooks/teams/useResetTeamMemberSpend.ts | 17 +++ .../TableIconActionButton.tsx | 1 + .../common_components/MemberTable.tsx | 18 +++ .../src/components/team/TeamInfo.tsx | 10 ++ .../components/team/TeamMemberTab.test.tsx | 114 +++++++++++++++++- .../src/components/team/TeamMemberTab.tsx | 94 ++++++++++++--- 6 files changed, 234 insertions(+), 20 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts new file mode 100644 index 00000000000..1fdf8fbd98d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useResetTeamMemberSpend.ts @@ -0,0 +1,17 @@ +import { useMutation } from "@tanstack/react-query"; +import { fetchClient } from "@/lib/http/api"; + +export interface ResetTeamMemberSpendParams { + teamId: string; + userId: string; +} + +export const resetTeamMemberSpend = async ({ teamId, userId }: ResetTeamMemberSpendParams): Promise => { + await fetchClient.POST("/team/{team_id}/member/{user_id}/reset_spend", { + params: { path: { team_id: teamId, user_id: userId } }, + body: { reset_to: 0 }, + }); +}; + +export const useResetTeamMemberSpend = () => + useMutation({ mutationFn: resetTeamMemberSpend }); diff --git a/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx b/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx index 396355e1c0b..9da9c2dc702 100644 --- a/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx +++ b/ui/litellm-dashboard/src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx @@ -30,6 +30,7 @@ export const TableIconActionButtonMap: Record boolean; + onResetSpend?: (member: Member) => void; + showResetSpendForMember?: (member: Member) => boolean; emptyText?: string; } @@ -73,6 +75,8 @@ interface MemberColumnDeps { roleTooltip?: string; extraColumns: MemberTableColumn[]; showDeleteForMember?: (member: Member) => boolean; + onResetSpend?: (member: Member) => void; + showResetSpendForMember?: (member: Member) => boolean; } const extraColumnDef = (column: MemberTableColumn): ColumnDef => { @@ -105,6 +109,8 @@ const buildColumns = ({ roleTooltip, extraColumns, showDeleteForMember, + onResetSpend, + showResetSpendForMember, }: MemberColumnDeps): ColumnDef[] => [ { id: "user_alias", @@ -173,6 +179,14 @@ const buildColumns = ({ dataTestId="edit-member" onClick={() => onEdit(row.original)} /> + {onResetSpend && (showResetSpendForMember?.(row.original) ?? true) && ( + onResetSpend(row.original)} + /> + )} {(!showDeleteForMember || showDeleteForMember(row.original)) && ( = ({ } }; + const refreshTeamData = async () => { + if (!accessToken) return; + try { + setTeamData(await teamInfoCall(accessToken, teamId)); + } catch { + toast.fromError("Failed to load team information"); + } + }; + useEffect(() => { fetchTeamInfo(); }, [teamId, accessToken]); @@ -1351,6 +1360,7 @@ const TeamInfoView: React.FC = ({ teamData={teamData} canEditTeam={canEditTeam} handleMemberDelete={handleMemberDelete} + onMemberSpendReset={refreshTeamData} setSelectedEditMember={setSelectedEditMember} setIsEditMemberModalVisible={setIsEditMemberModalVisible} setIsAddMemberModalVisible={setIsAddMemberModalVisible} diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx index 760074d5dc9..52cba1e6330 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx @@ -1,4 +1,4 @@ -import { fireEvent, screen, within } from "@testing-library/react"; +import { fireEvent, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../tests/test-utils"; @@ -13,6 +13,9 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: vi.fn(), })); +const { POST } = vi.hoisted(() => ({ POST: vi.fn() })); +vi.mock("@/lib/http/api", () => ({ fetchClient: { POST } })); + vi.mock("@/utils/roles", () => ({ isUserTeamAdminForSingleTeam: vi.fn(() => false), isProxyAdminRole: vi.fn(() => false), @@ -26,6 +29,7 @@ const mockHandleMemberDelete = vi.fn(); const mockSetSelectedEditMember = vi.fn(); const mockSetIsEditMemberModalVisible = vi.fn(); const mockSetIsAddMemberModalVisible = vi.fn(); +const mockOnMemberSpendReset = vi.fn(); const budgetResetIso = new Date(2026, 6, 15, 12, 0, 0).toISOString(); @@ -121,6 +125,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -136,6 +141,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -154,6 +160,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -172,6 +179,7 @@ describe("TeamMembersComponent", () => { const props = { canEditTeam: false, handleMemberDelete: mockHandleMemberDelete, + onMemberSpendReset: mockOnMemberSpendReset, setSelectedEditMember: mockSetSelectedEditMember, setIsEditMemberModalVisible: mockSetIsEditMemberModalVisible, setIsAddMemberModalVisible: mockSetIsAddMemberModalVisible, @@ -195,6 +203,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -221,6 +230,7 @@ describe("TeamMembersComponent", () => { })} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -247,6 +257,7 @@ describe("TeamMembersComponent", () => { })} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -262,6 +273,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -280,6 +292,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -295,6 +308,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -311,6 +325,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -330,6 +345,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -364,6 +380,7 @@ describe("TeamMembersComponent", () => { teamData={teamData} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -417,6 +434,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -447,6 +465,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -466,6 +485,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={true} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -482,6 +502,7 @@ describe("TeamMembersComponent", () => { teamData={createMockTeamData()} canEditTeam={false} handleMemberDelete={mockHandleMemberDelete} + onMemberSpendReset={mockOnMemberSpendReset} setSelectedEditMember={mockSetSelectedEditMember} setIsEditMemberModalVisible={mockSetIsEditMemberModalVisible} setIsAddMemberModalVisible={mockSetIsAddMemberModalVisible} @@ -491,4 +512,95 @@ describe("TeamMembersComponent", () => { expect(screen.queryByTestId("edit-member")).not.toBeInTheDocument(); expect(screen.queryByTestId("delete-member")).not.toBeInTheDocument(); }); + + describe("reset spend", () => { + const renderEditableTab = () => + renderWithProviders( + , + ); + + it("resets the member's current cycle spend to $0 after confirming, then refreshes the team", async () => { + const user = userEvent.setup(); + POST.mockResolvedValue({ data: {} }); + renderEditableTab(); + + const memberRow = screen.getByRole("row", { name: /user1@test\.com/ }); + await user.click(within(memberRow).getByTestId("reset-member-spend")); + + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + expect(dialog).toHaveTextContent("user1@test.com"); + expect(dialog).toHaveTextContent("$100.5000"); + expect(POST).not.toHaveBeenCalled(); + + await user.click(within(dialog).getByRole("button", { name: "Reset" })); + + await waitFor(() => expect(mockOnMemberSpendReset).toHaveBeenCalledTimes(1)); + expect(POST).toHaveBeenCalledExactlyOnceWith("/team/{team_id}/member/{user_id}/reset_spend", { + params: { path: { team_id: "team-123", user_id: "user1@test.com" } }, + body: { reset_to: 0 }, + }); + expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); + }); + + it("keeps the dialog open and does not refresh the team when the reset fails", async () => { + const user = userEvent.setup(); + POST.mockRejectedValue(new Error("Cannot reset your own spend. Ask a proxy admin.")); + renderEditableTab(); + + await user.click(screen.getByTestId("reset-member-spend")); + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + await user.click(within(dialog).getByRole("button", { name: "Reset" })); + + await waitFor(() => expect(POST).toHaveBeenCalledTimes(1)); + expect(mockOnMemberSpendReset).not.toHaveBeenCalled(); + expect(screen.getByRole("dialog", { name: "Reset Team Member Spend" })).toBeInTheDocument(); + }); + + it("does not call the API when the dialog is cancelled", async () => { + const user = userEvent.setup(); + renderEditableTab(); + + await user.click(screen.getByTestId("reset-member-spend")); + const dialog = await screen.findByRole("dialog", { name: "Reset Team Member Spend" }); + await user.click(within(dialog).getByRole("button", { name: "Cancel" })); + + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + expect(POST).not.toHaveBeenCalled(); + }); + + it("only offers the reset on members that have current cycle spend", () => { + renderEditableTab(); + + expect( + within(screen.getByRole("row", { name: /user1@test\.com/ })).getByTestId("reset-member-spend"), + ).toBeVisible(); + expect( + within(screen.getByRole("row", { name: /user2@test\.com/ })).queryByTestId("reset-member-spend"), + ).not.toBeInTheDocument(); + }); + + it("hides the reset on the caller's own row for a team admin, since the backend rejects it", () => { + vi.mocked(useAuthorized).mockReturnValue({ userId: "user1@test.com", userRole: "Internal User" } as never); + vi.mocked(isProxyAdminRole).mockReturnValue(false); + renderEditableTab(); + + expect(screen.queryByTestId("reset-member-spend")).not.toBeInTheDocument(); + }); + + it("shows the reset on the caller's own row for a proxy admin", () => { + vi.mocked(useAuthorized).mockReturnValue({ userId: "user1@test.com", userRole: "Admin" } as never); + vi.mocked(isProxyAdminRole).mockReturnValue(true); + renderEditableTab(); + + expect(screen.getByTestId("reset-member-spend")).toBeVisible(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx index 16f3d12d71c..a869c1ad624 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx @@ -1,13 +1,18 @@ +import { useResetTeamMemberSpend } from "@/app/(dashboard)/hooks/teams/useResetTeamMemberSpend"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { Button } from "@/components/ui/button"; +import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { SimpleTooltip } from "@/components/ui/tooltip"; import MemberTable from "@/components/common_components/MemberTable"; import { Member } from "@/components/networking"; +import { parseErrorMessage } from "@/components/shared/errorUtils"; import { DateCell, MoneyCell } from "@/components/shared/table_cells"; +import { toast } from "@/lib/toast"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { isProxyAdminRole, isUserTeamAdminForSingleTeam } from "@/utils/roles"; import { CircleHelp } from "lucide-react"; -import type { ComponentProps } from "react"; +import { useState, type ComponentProps } from "react"; import { TeamData, TeamMembership } from "./TeamInfo"; export const seedMemberBudgetFields = ( @@ -31,6 +36,7 @@ interface TeamMemberTabProps { setSelectedEditMember: (member: Member) => void; setIsEditMemberModalVisible: (visible: boolean) => void; setIsAddMemberModalVisible: (visible: boolean) => void; + onMemberSpendReset: () => void; } export default function TeamMemberTab({ @@ -40,7 +46,11 @@ export default function TeamMemberTab({ setSelectedEditMember, setIsEditMemberModalVisible, setIsAddMemberModalVisible, + onMemberSpendReset, }: TeamMemberTabProps) { + const [memberToResetSpend, setMemberToResetSpend] = useState(null); + const { mutate: resetMemberSpend, isPending: isResettingSpend } = useResetTeamMemberSpend(); + const formatNumber = (value: number | null): string => { if (value === null || value === undefined) return "0"; @@ -199,24 +209,70 @@ export default function TeamMemberTab({ }, ]; + const handleResetSpend = () => { + if (!memberToResetSpend?.user_id) return; + resetMemberSpend( + { teamId: teamData.team_id, userId: memberToResetSpend.user_id }, + { + onSuccess: () => { + toast.success("Team member spend reset to $0"); + setMemberToResetSpend(null); + onMemberSpendReset(); + }, + onError: (error) => toast.fromError(parseErrorMessage(error)), + }, + ); + }; + return ( - { - const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); - setSelectedEditMember(seedMemberBudgetFields(record, membership?.litellm_budget_table)); - setIsEditMemberModalVisible(true); - }} - onDelete={handleMemberDelete} - onAddMember={() => setIsAddMemberModalVisible(true)} - roleColumnTitle="Team Role" - roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." - extraColumns={extraColumns} - showDeleteForMember={() => - isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) - } - /> + <> + { + const membership = teamData.team_memberships.find((tm) => tm.user_id === record.user_id); + setSelectedEditMember(seedMemberBudgetFields(record, membership?.litellm_budget_table)); + setIsEditMemberModalVisible(true); + }} + onDelete={handleMemberDelete} + onAddMember={() => setIsAddMemberModalVisible(true)} + roleColumnTitle="Team Role" + roleTooltip="This role applies only to this team and is independent from the user's proxy-level role." + extraColumns={extraColumns} + showDeleteForMember={() => + isProxyAdmin || (canEditTeam && !isUserTeamAdmin) || (isUserTeamAdmin && !disableTeamAdminDeleteTeamUser) + } + onResetSpend={setMemberToResetSpend} + showResetSpendForMember={(record) => + getUserCurrentCycleSpend(record.user_id) > 0 && (isProxyAdmin || record.user_id !== userId) + } + /> + !open && setMemberToResetSpend(null)}> + + + Reset Team Member Spend + +

+ Reset current cycle spend for{" "} + {memberToResetSpend?.user_email || memberToResetSpend?.user_id} in this team to{" "} + $0? +

+

+ Current cycle spend:{" "} + ${formatNumberWithCommas(getUserCurrentCycleSpend(memberToResetSpend?.user_id ?? null), 4)} + . This is the value checked against the member's budget. Total spend and logs are preserved. +

+ + + + +
+
+ ); } From abf530fbeb0563d5877ceb883261b25161855b29 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 20:55:11 +0000 Subject: [PATCH 208/253] fix(proxy): never rewind the daily global spend marker from an overlapping reconcile run Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../daily_global_spend_rollup.py | 31 +++++++++++++------ .../test_daily_global_spend_rollup.py | 30 +++++++++++++++++- 2 files changed, 51 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index c135c7d1d9c..d50feb6f16b 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -161,6 +161,24 @@ async def _record_marker(prisma_client: "PrismaClient", marker: ReconciledThroug await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) +async def _stored_marker(prisma_client: "PrismaClient") -> ReconciledThrough | None: + """The marker as another pod may have just written it, bypassing this pod's config cache.""" + param: Final = await ConfigRepository(prisma_client).get_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) + return None if param is None else _marker_from_param_value(param.param_value) + + +async def _record_advanced(prisma_client: "PrismaClient", days: tuple[str, ...], *, scanned_at: str | None) -> None: + """Advance the stored marker by ``days``. Two runs can overlap (Redis unreachable, lock expired + on a long backfill), so the base is what is stored now, not the snapshot this run scanned from: + a slower run may then only add to the faster run's marker, never rewind it. Without a new scan + time the stored one is kept.""" + stored: Final = await _stored_marker(prisma_client) + kept_scanned_at: Final = None if stored is None else stored.scanned_at + await _record_marker( + prisma_client, _advanced(stored, days, scanned_at=scanned_at if scanned_at is not None else kept_scanned_at) + ) + + async def _db_now(prisma_client: "PrismaClient") -> _NowRow: rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL) return _NowRow.model_validate(rows[0]) @@ -203,7 +221,7 @@ async def run_daily_global_spend_reconcile(prisma_client: "PrismaClient") -> Rec marker: Final = await reconciled_through(prisma_client) return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=scan.days[len(done)]) if scan.marker is not None or done: - await _record_marker(prisma_client, _advanced(scan.marker, done, scanned_at=scan.scanned_at)) + await _record_advanced(prisma_client, done, scanned_at=scan.scanned_at) return ReconcileResult(days_reconciled=done, reconciled_through=await reconciled_through(prisma_client)) @@ -215,21 +233,16 @@ def _advanced(marker: ReconciledThrough | None, days: tuple[str, ...], *, scanne async def _reconcile_until_failure(prisma_client: "PrismaClient", scan: _PendingScan) -> tuple[str, ...]: for index, day in enumerate(scan.days): - if not await _reconcile_and_record(prisma_client, scan.marker, scan.days[: index + 1]): + if not await _reconcile_and_record(prisma_client, scan.days[: index + 1]): return scan.days[:index] return scan.days -async def _reconcile_and_record( - prisma_client: "PrismaClient", marker: ReconciledThrough | None, done_with_this: tuple[str, ...] -) -> bool: +async def _reconcile_and_record(prisma_client: "PrismaClient", done_with_this: tuple[str, ...]) -> bool: day: Final = done_with_this[-1] try: await reconcile_day(prisma_client, day) - await _record_marker( - prisma_client, - _advanced(marker, done_with_this, scanned_at=None if marker is None else marker.scanned_at), - ) + await _record_advanced(prisma_client, done_with_this, scanned_at=None) except Exception as exc: # noqa: BLE001 # one bad day must not lose the days already done verbose_proxy_logger.exception("Daily global spend reconcile: day %s failed: %s", day, exc) return False diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 9655953134a..69f06fad081 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -40,6 +40,10 @@ class _FakeConfigTable: self.rows[where["param_name"]] = data["update"]["param_value"] return _FakeConfigRow(where["param_name"], data["update"]["param_value"]) + async def find_unique(self, *, where: dict[str, str]) -> _FakeConfigRow | None: + stored = self.rows.get(where["param_name"]) + return None if stored is None else _FakeConfigRow(where["param_name"], stored) + class _FakeDb: """Per-key rows are ``{date: updated_at}`` with a fake database clock that ticks per query, @@ -68,11 +72,15 @@ class _FakeDb: if day in self._prisma.failing_days: raise RuntimeError(f"day {day} exploded") self._prisma.reconciled.append(day) + landing = self._prisma.marker_landing_on_day.get(day) + if landing is not None: + self.litellm_config.rows[DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM] = landing return 1 class _FakePrisma: - """Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw.""" + """Enough of PrismaClient for the reconcile: per-key dates, a config table, and execute_raw. + ``marker_landing_on_day`` stores another pod's marker the moment this run rewrites that day.""" def __init__( self, user_days: tuple[str, ...], failing_days: frozenset[str] = frozenset(), today: date = TODAY @@ -81,6 +89,7 @@ class _FakePrisma: self.today = today self.user_rows: dict[str, str] = {d: "clock-0000" for d in user_days} self.failing_days = failing_days + self.marker_landing_on_day: dict[str, str] = {} self.reconciled: list[str] = [] self.db = _FakeDb(self) @@ -221,6 +230,25 @@ async def test_the_next_run_resumes_from_the_failed_day(): assert await reconciled_through(prisma) == "2026-09-03" +@pytest.mark.asyncio +async def test_a_slower_overlapping_run_never_rewinds_the_marker_a_faster_run_stored(): + """Two pods can reconcile at once (Redis unreachable, or the lock expired on a long backfill). + When the faster one has already stored a later marker, the slower one may only add to it. Putting + its own older prefix back, or dropping the scan time, would send usage reads for every day in + between back to the per-key table until the next run.""" + prisma = _FakePrisma(user_days=("2026-09-01", "2026-09-02", "2026-09-03"), failing_days=frozenset({"2026-09-03"})) + prisma.marker_landing_on_day = { + "2026-09-02": '{"reconciled_through": "2026-09-14", "scanned_at": "clock-0009"}', + } + + result = await run_daily_global_spend_reconcile(prisma) + + assert result.days_reconciled == ("2026-09-01", "2026-09-02") + assert result.reconciled_through == "2026-09-14" + marker = await read_marker(prisma) + assert marker is not None and (marker.reconciled_through, marker.scanned_at) == ("2026-09-14", "clock-0009") + + @pytest.mark.asyncio async def test_a_failure_with_nothing_done_reports_the_previous_marker_and_alerts(): """When the rewrite of a late day fails the marker must stay put and the operator must hear about it.""" From 836bf7d8974bc3fc3cdd71a3fe6d9565201c2f2e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 13:59:36 -0700 Subject: [PATCH 209/253] refactor(rust): make the exception mapper a pure rule table Rebuild exception_type around text rules per provider family and one shared status table. The mapper takes the context, an injected redactor and the original failure, and returns a PublicError with the message, the real upstream response and the debug text. Every divergence from the Python mapper and every known gap is listed in the module header Match Python on a standalone 429 with an unknown status and on Cohere's rules for failures without a status. Drop python_repr and the unread public_failures fixtures --- litellm-rust/Cargo.lock | 1 - litellm-rust/crates/core-utils/Cargo.toml | 1 - .../src/exception_mapping_utils/cohere.rs | 275 ++----- .../src/exception_mapping_utils/mod.rs | 757 +++++++----------- .../src/exception_mapping_utils/openai.rs | 600 +++----------- .../src/exception_mapping_utils/original.rs | 145 +++- .../src/exception_mapping_utils/public.rs | 257 +----- .../src/exception_mapping_utils/rules.rs | 283 +------ .../src/exception_mapping_utils/status.rs | 182 +---- .../src/exception_mapping_utils/vertex_ai.rs | 573 ++----------- litellm-rust/crates/core-utils/src/lib.rs | 1 - .../crates/core-utils/src/python_repr.rs | 93 --- .../crates/core-utils/src/secret_redaction.rs | 50 +- .../fixtures/public_failures/api.json | 13 - .../public_failures/api_connection.json | 11 - .../status_authentication.json | 13 - .../public_failures/status_bad_gateway.json | 28 - .../public_failures/status_bad_request.json | 28 - .../status_content_policy_violation.json | 28 - .../status_context_window_exceeded.json | 28 - .../status_internal_server.json | 19 - .../public_failures/status_not_found.json | 28 - .../status_permission_denied.json | 19 - .../public_failures/status_rate_limit.json | 28 - .../status_service_unavailable.json | 28 - .../status_unsupported_params.json | 28 - .../public_failures/timeout_with_status.json | 12 - .../timeout_without_status.json | 12 - 28 files changed, 807 insertions(+), 2734 deletions(-) delete mode 100644 litellm-rust/crates/core-utils/src/python_repr.rs delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/api.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json delete mode 100644 tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 6c524124550..c359ca19986 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2092,7 +2092,6 @@ dependencies = [ "serde_json", "serde_path_to_error", "serde_with", - "strum", "thiserror 2.0.19", "url", ] diff --git a/litellm-rust/crates/core-utils/Cargo.toml b/litellm-rust/crates/core-utils/Cargo.toml index 59e0ee1a09d..eb353bc060c 100644 --- a/litellm-rust/crates/core-utils/Cargo.toml +++ b/litellm-rust/crates/core-utils/Cargo.toml @@ -12,7 +12,6 @@ serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" serde_with.workspace = true -strum.workspace = true thiserror.workspace = true url.workspace = true diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs index 381ebd7ad27..c2a391ee223 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/cohere.rs @@ -1,232 +1,115 @@ -use super::Mapping; -use super::public::{PublicFailure, StatusClass}; -use super::rules::{Kind, ResponseChoice, Rule, apply, contains_any}; +use super::public::PublicError; +use super::rules::{Rule, contains_any}; -const fn with_response(class: StatusClass) -> Kind { - Kind::Status { - class, - response: ResponseChoice::Provider, - } -} - -fn original(mapping: &Mapping<'_>) -> String { - format!("CohereException - {}", mapping.original.message) -} - -fn status_is(mapping: &Mapping<'_>, statuses: &[u16]) -> bool { - mapping - .original - .status - .is_some_and(|status| statuses.contains(&status)) -} - -/// `_map_cohere_exception`, in its branch order. A failure no rule claims falls through to -/// the status table. -const RULES: &[Rule] = &[ - Rule { - when: |mapping| { +/// The text branches of `_map_cohere_exception`, in its order. +pub(super) const RULES: &[Rule] = &[ + Rule::new( + |mapping| { contains_any( &mapping.error_str, &["invalid api token", "No API key provided."], ) }, - kind: with_response(StatusClass::Authentication), - message: original, - debug: false, - }, - Rule { - when: |mapping| mapping.error_str.contains("invalid type: parameter"), - kind: with_response(StatusClass::BadRequest), - message: original, - debug: false, - }, - Rule { - when: |mapping| mapping.error_str.contains("too many tokens"), - kind: with_response(StatusClass::ContextWindowExceeded), - message: original, - debug: false, - }, - Rule { - when: |mapping| { + PublicError::Authentication, + ), + Rule::new( + |mapping| mapping.error_str.contains("invalid type: parameter"), + PublicError::BadRequest, + ), + Rule::new( + |mapping| mapping.error_str.contains("too many tokens"), + PublicError::ContextWindowExceeded, + ), + Rule::new( + |mapping| { mapping .error_str .to_lowercase() .contains("internal server error") }, - kind: with_response(StatusClass::InternalServer), - message: |mapping| format!("CohereException - {}", mapping.error_str), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, &[400, 498]), - kind: with_response(StatusClass::BadRequest), - message: original, - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, &[408]), - kind: Kind::Timeout(None), - message: original, - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, &[500]), - kind: with_response(StatusClass::InternalServer), - message: original, - debug: false, - }, + PublicError::InternalServer, + ), + Rule::new( + |mapping| mapping.status.is_none() && mapping.error_str.contains("invalid type:"), + PublicError::BadRequest, + ), + Rule::new( + |mapping| mapping.status.is_none() && mapping.error_str.contains("Unexpected server error"), + PublicError::InternalServer, + ), ]; -pub(super) fn map(mapping: &Mapping<'_>) -> Option { - apply(RULES, mapping).map(|failure| PublicFailure { - llm_provider: Some("cohere".to_string()), - ..failure - }) -} - #[cfg(test)] mod tests { - use super::super::testing::{context, failure, http, status, upstream}; - use super::super::{ExceptionFamily, OriginalException, PublicKind}; + use super::super::rules::first_match; + use super::super::testing::mapping; use super::*; - fn mapped(provider: &str, original: &OriginalException) -> Option { - let context = context(provider, ExceptionFamily::Cohere); - map(&Mapping::new(&context, original)) + fn classified(text: &str) -> Option { + classified_with(Some(400), text) } - fn cohere(class: StatusClass, status_code: u16, body: &str, message: &str) -> PublicFailure { - failure( - status(class, upstream(status_code, body)), - message, - "cohere", - ) + fn classified_with(status: Option, text: &str) -> Option { + first_match(RULES, &mapping(status, text)).map(|rule| rule.error) } #[rstest::rstest] - #[case::invalid_token( - 500, - "invalid api token", - cohere( - StatusClass::Authentication, - 500, - "invalid api token", - "CohereException - invalid api token" - ) - )] - #[case::no_api_key( - 500, - "No API key provided.", - cohere( - StatusClass::Authentication, - 500, - "No API key provided.", - "CohereException - No API key provided." - ) - )] - #[case::invalid_parameter( - 500, - "invalid type: parameter x", - cohere( - StatusClass::BadRequest, - 500, - "invalid type: parameter x", - "CohereException - invalid type: parameter x" - ) - )] - #[case::too_many_tokens( - 500, - "too many tokens", - cohere( - StatusClass::ContextWindowExceeded, - 500, - "too many tokens", - "CohereException - too many tokens" - ) - )] - #[case::internal_server_text( - 400, - "Internal Server Error", - cohere( - StatusClass::InternalServer, - 400, - "Internal Server Error", - "CohereException - Internal Server Error" - ) - )] - #[case::bad_request( - 400, - "rejected", - cohere(StatusClass::BadRequest, 400, "rejected", "CohereException - rejected") - )] - #[case::invalid_token_status( - 498, - "rejected", - cohere(StatusClass::BadRequest, 498, "rejected", "CohereException - rejected") - )] - #[case::request_timeout(408, "rejected", failure(PublicKind::Timeout { status: None }, "CohereException - rejected", "cohere"))] - #[case::internal_server( - 500, - "rejected", - cohere( - StatusClass::InternalServer, - 500, - "rejected", - "CohereException - rejected" - ) - )] - fn each_rule_maps_and_reports_cohere( - #[case] status_code: u16, - #[case] body: &str, - #[case] expected: PublicFailure, - ) { - assert_eq!(mapped("azure_ai", &http(status_code, body)), Some(expected)); - } - - #[rstest::rstest] - #[case::unmapped_status(409)] - #[case::unauthorized(401)] - fn statuses_without_a_rule_fall_through(#[case] status_code: u16) { - assert_eq!(mapped("cohere", &http(status_code, "rejected")), None); - } - - #[test] - fn the_internal_server_rule_uses_the_redacted_text() { - let body = "internal server error Bearer abcdefghijklmnop"; - assert_eq!( - mapped("cohere", &http(400, body)), - Some(cohere( - StatusClass::InternalServer, - 400, - body, - "CohereException - internal server error REDACTED" - )) - ); + #[case::invalid_token("invalid api token", PublicError::Authentication)] + #[case::no_api_key("No API key provided.", PublicError::Authentication)] + #[case::invalid_parameter("invalid type: parameter x", PublicError::BadRequest)] + #[case::too_many_tokens("too many tokens", PublicError::ContextWindowExceeded)] + #[case::internal_server_text("Internal Server Error", PublicError::InternalServer)] + #[case::internal_server_any_case("INTERNAL server ERROR", PublicError::InternalServer)] + fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(text), Some(expected)); } #[rstest::rstest] #[case::token_before_parameter( "invalid api token invalid type: parameter", - StatusClass::Authentication + PublicError::Authentication )] #[case::parameter_before_tokens( "invalid type: parameter too many tokens", - StatusClass::BadRequest + PublicError::BadRequest )] #[case::tokens_before_internal( "too many tokens Internal Server Error", - StatusClass::ContextWindowExceeded + PublicError::ContextWindowExceeded )] - #[case::internal_before_status("Internal Server Error", StatusClass::InternalServer)] - fn the_earlier_rule_wins_when_two_apply(#[case] body: &str, #[case] class: StatusClass) { - assert_eq!( - mapped("cohere", &http(400, body)), - Some(cohere( - class, - 400, - body, - &format!("CohereException - {body}") - )) - ); + fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(text), Some(expected)); + } + + #[rstest::rstest] + #[case::invalid_type(None, "invalid type: x", Some(PublicError::BadRequest))] + #[case::unexpected_server_error( + None, + "Unexpected server error", + Some(PublicError::InternalServer) + )] + #[case::invalid_type_before_unexpected( + None, + "invalid type: x Unexpected server error", + Some(PublicError::BadRequest) + )] + #[case::internal_before_invalid_type( + None, + "internal server error invalid type: x", + Some(PublicError::InternalServer) + )] + #[case::invalid_type_with_a_status(Some(500), "invalid type: x", None)] + #[case::unexpected_with_a_status(Some(400), "Unexpected server error", None)] + fn the_trailing_rules_only_claim_failures_without_a_status( + #[case] status: Option, + #[case] text: &str, + #[case] expected: Option, + ) { + assert_eq!(classified_with(status, text), expected); + } + + #[test] + fn text_without_a_marker_is_left_to_the_status_table() { + assert_eq!(classified("rejected"), None); } } diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs index acc7820c770..162d325e4f4 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/mod.rs @@ -1,18 +1,47 @@ -//! A port of Python's `exception_type` for the routes that run in Rust. +//! A port of Python's `exception_type` for the routes that run in Rust. Rust decides the +//! public class, the message and the debug text; Python only builds the class. +//! +//! DIVERGENCES: where the Python mapper is inconsistent, the port follows one rule instead. +//! - The message is always `{Provider}Exception - {redacted text}`. Python's per-branch +//! labels (`RateLimitError: `, `litellm.RateLimitError: `, `Vertex_aiException BadRequestError`) +//! are dropped because every public class already prefixes `litellm.{Class}: `. +//! - The upstream response is always the real one. Python swaps in made-up `httpx.Response` +//! stubs on some Vertex branches, losing the body and `retry-after`. +//! - The debug text is always attached; Python passes it on some branches only. +//! - No family rule turns a status into a class; the shared status table owns that. So a +//! Vertex 502 is a `BadGatewayError` and an OpenAI-family 403 is a `PermissionDeniedError`. +//! Three rules read the status only to gate a text match, as Python does: the standalone +//! `429`, Vertex's wrapped 429 behind a 5xx, and Cohere's rules for failures with no status. +//! - A timeout text marker on an HTTP failure keeps the upstream response. Python's `Timeout` +//! carries none. +//! - Every family matches and reports the redacted text. Python's OpenAI mapper builds the +//! message from the unredacted text. +//! - A refused connection is an `APIConnectionError`, not the 500 Python's HTTP handler +//! synthesizes. +//! - Dropped Python rules: Vertex's bare `403` substring (it matches `4031 tokens`), Vertex's +//! `IndexError` quota marker (a Python client crash), the OpenAI SDK's missing-`api_key` +//! text and its `OPENAI` renaming, Cohere's `llm_provider="cohere"` override, and Cohere's +//! `CohereConnectionError` check (a Python SDK class name). //! //! KNOWN_GAPS: differences from the Python mapper that no Rust route can reach today. Each //! one stops being acceptable at its trigger. -//! - The Vertex partner-model API base for "claude" models is not built into -//! `extra_information`. Trigger: a Vertex route whose models include Anthropic partner -//! models; then `api_base` gets that branch and a table row. +//! - The Vertex partner-model API base for "claude" models is not built into the debug text. +//! Trigger: a Vertex route whose models include Anthropic partner models. +//! - The debug text's `API Base` line is only the non-streaming Vertex URL. Python prefers an +//! explicit or provider-resolved `api_base`, uses `:streamGenerateContent` when streaming, +//! and has Gemini and OpenAI defaults. Trigger: the first route wired to this mapper, since +//! every route knows its `api_base`. +//! - The debug text has no `Messages:` line, which Python adds when +//! `redact_messages_in_exceptions` is off. Trigger: a wired route that carries messages. //! - Python reports the provider `get_llm_provider` resolves for a stripped model name when //! that name happens to be in the model cost map. Trigger: a route whose model names //! overlap the cost map; that needs the provider resolution port, not a classifier change. -//! - The generic `APIConnectionError` fallback appends `traceback.format_exc()` to the -//! message. Rust has no Python traceback and does not invent one; a sweep row that reaches -//! it compares the message before the traceback. +//! - `litellm_proxy` errors are not unwrapped into the proxied exception. Trigger: a Rust +//! route that calls a LiteLLM proxy. +//! - Only the OpenAI-compatible, Vertex AI and Cohere mappers are ported; every other +//! provider goes straight to the status table. Trigger: a Rust route for such a provider. -use super::secret_redaction::{redact_string, secret_redaction_enabled}; +use super::secret_redaction::SecretRedactor; mod cohere; mod openai; @@ -22,10 +51,10 @@ mod rules; mod status; mod vertex_ai; -pub use original::{ExceptionFamily, LocalClass, OriginalException}; -pub use public::{HttpStub, PublicFailure, PublicKind, ResponseArg, StatusClass, UpstreamResponse}; +pub use original::{ExceptionFamily, OriginalException}; +pub use public::{MappedFailure, PublicError, UpstreamResponse}; -const DOCS_URL: &str = "https://docs.litellm.ai/docs"; +use rules::{Rule, contains_any, first_match}; const TIMEOUT_MARKERS: &[&str] = &[ "Request Timeout Error", @@ -37,11 +66,8 @@ const TIMEOUT_MARKERS: &[&str] = &[ #[derive(Clone, Debug, Default, PartialEq, Eq)] pub struct ExceptionContext { pub model: String, - pub custom_llm_provider: Option, - pub family: ExceptionFamily, + pub custom_llm_provider: String, pub asynchronous: bool, - pub suppress_debug_info: bool, - pub redact_messages_in_exceptions: bool, pub vertex_project: Option, pub vertex_location: Option, pub model_group: Option, @@ -50,73 +76,93 @@ pub struct ExceptionContext { pub user_api_key_team_alias: Option, } -/// The attributes `exception_type` reads off the Python exception: a provider error -/// (`BaseLLMException`) carries a status, a response and a request, a plain exception -/// carries only its text. -struct Raised { +/// What the rules read: the status of a provider response, if any, and the redacted text. +struct Mapping { status: Option, - status_is_synthesized: bool, - message: String, - response: Option, + error_str: String, } -impl Raised { - fn provider( - status: u16, - message: String, - body: String, - headers: Vec<(String, String)>, - ) -> Self { - Self { - status: Some(status), - status_is_synthesized: false, - message, - response: Some(UpstreamResponse { - status, - body, - headers, +pub fn exception_type( + context: &ExceptionContext, + redactor: Option<&SecretRedactor>, + original: &OriginalException, +) -> MappedFailure { + let (status, text, upstream) = match original { + OriginalException::Http { + status, + body, + headers, + } => ( + Some(*status), + body.clone(), + Some(UpstreamResponse { + status: *status, + body: body.clone(), + headers: headers.clone(), }), + ), + OriginalException::Connection { message } | OriginalException::Plain { message } => { + (None, message.clone(), None) } - } - - fn plain(message: String) -> Self { - Self { - status: None, - status_is_synthesized: false, - message, - response: None, - } - } - - fn new(original: &OriginalException, asynchronous: bool) -> Self { - match original { - OriginalException::Http { - status, - body, - headers, - } => Self::provider(*status, body.clone(), body.clone(), headers.clone()), - OriginalException::Connection { message } => Self { - status_is_synthesized: true, - ..Self::provider(500, message.clone(), String::new(), Vec::new()) - }, - OriginalException::Timeout { - timeout_seconds, - elapsed_seconds, - } => Self::provider( - 408, - timeout_message(asynchronous, *timeout_seconds, *elapsed_seconds), - String::new(), - Vec::new(), - ), - OriginalException::Response { message } - | OriginalException::Local { message, .. } - | OriginalException::Public { message, .. } => Self::plain(message.clone()), - } + OriginalException::Timeout { + timeout_seconds, + elapsed_seconds, + } => ( + None, + timeout_message(context.asynchronous, *timeout_seconds, *elapsed_seconds), + None, + ), + }; + let mapping = Mapping { + status, + error_str: match redactor { + Some(redactor) => redactor.redact(&text), + None => text, + }, + }; + let family = ExceptionFamily::for_provider(&context.custom_llm_provider); + let (error, hint) = classify(family, original, &mapping); + MappedFailure { + error, + message: format!( + "{} - {}{hint}", + exception_provider(&context.custom_llm_provider), + mapping.error_str + ), + upstream, + debug_info: extra_information(context, api_base(context).as_deref()), } } -/// The text `litellm.Timeout` carries when the Python HTTP handler times out: the sync -/// and async handlers word it differently. +fn classify( + family: ExceptionFamily, + original: &OriginalException, + mapping: &Mapping, +) -> (PublicError, &'static str) { + const TIMEOUT: PublicError = PublicError::Timeout { status: 408 }; + if matches!(original, OriginalException::Timeout { .. }) + || contains_any(&mapping.error_str, TIMEOUT_MARKERS) + { + return (TIMEOUT, ""); + } + if let Some(rule) = first_match(family_rules(family), mapping) { + return (rule.error, rule.hint); + } + let by_status = mapping.status.and_then(status::classify); + (by_status.unwrap_or(PublicError::ApiConnection), "") +} + +fn family_rules(family: ExceptionFamily) -> &'static [Rule] { + match family { + ExceptionFamily::OpenAiCompatible => openai::RULES, + ExceptionFamily::VertexAi => vertex_ai::RULES, + ExceptionFamily::Cohere => cohere::RULES, + ExceptionFamily::Other => &[], + } +} + +/// The text the Python HTTP handler's timeout carries: the sync and async handlers word it +/// differently. fn timeout_message( asynchronous: bool, timeout_seconds: Option, @@ -126,11 +172,9 @@ fn timeout_message( if asynchronous { let elapsed = python_float(elapsed_seconds.map(|seconds| (seconds * 1000.0).round() / 1000.0)); - format!( - "litellm.Timeout: Connection timed out. Timeout passed={timeout}, time taken={elapsed} seconds" - ) + format!("Connection timed out. Timeout passed={timeout}, time taken={elapsed} seconds") } else { - format!("litellm.Timeout: Connection timed out after {timeout} seconds.") + format!("Connection timed out after {timeout} seconds.") } } @@ -142,113 +186,10 @@ fn python_float(value: Option) -> String { } } -/// Everything the rules read: the original as Python sees it and the text `exception_type` -/// derives from the context before any provider mapper runs. -struct Mapping<'a> { - context: &'a ExceptionContext, - original: Raised, - provider: &'a str, - error_str: String, - exception_provider: String, - extra_information: String, -} - -impl<'a> Mapping<'a> { - fn new(context: &'a ExceptionContext, original: &OriginalException) -> Self { - let original = Raised::new(original, context.asynchronous); - let error_str = if secret_redaction_enabled() { - redact_string(&original.message) - } else { - original.message.clone() - }; - Self { - context, - original, - provider: context.custom_llm_provider.as_deref().unwrap_or_default(), - error_str, - exception_provider: match &context.custom_llm_provider { - None => "None".to_string(), - Some(provider) => exception_provider(provider), - }, - extra_information: extra_information(context, api_base(context).as_deref()), - } - } - - fn failure(&self, kind: PublicKind, message: String, debug: bool) -> PublicFailure { - PublicFailure { - kind, - message, - model: self.context.model.clone(), - llm_provider: self.context.custom_llm_provider.clone(), - litellm_debug_info: debug.then(|| self.extra_information.clone()), - litellm_response_headers: None, - print_banner: false, - } - } -} - -pub fn exception_type(context: &ExceptionContext, original: &OriginalException) -> PublicFailure { - if let OriginalException::Public { class, message } = original { - return PublicFailure { - kind: PublicKind::Status { - status_class: *class, - response: None, - }, - message: message.clone(), - model: context.model.clone(), - llm_provider: context.custom_llm_provider.clone(), - litellm_debug_info: None, - litellm_response_headers: None, - print_banner: false, - }; - } - let mapping = Mapping::new(context, original); - let litellm_response_headers = mapping - .original - .response - .as_ref() - .map(|response| response.headers.clone()) - .filter(|headers| !headers.is_empty()); - PublicFailure { - litellm_response_headers, - print_banner: !context.suppress_debug_info, - ..map(&mapping) - } -} - -fn map(mapping: &Mapping<'_>) -> PublicFailure { - if rules::contains_any(&mapping.error_str, TIMEOUT_MARKERS) { - return mapping.failure( - PublicKind::Timeout { status: None }, - format!( - "APITimeoutError - Request timed out. Error_str: {}", - mapping.error_str - ), - true, - ); - } - let provider_failure = match mapping.context.family { - ExceptionFamily::OpenAiCompatible => openai::map(mapping), - ExceptionFamily::VertexAi => vertex_ai::map(mapping), - ExceptionFamily::Cohere => cohere::map(mapping), - ExceptionFamily::Other => None, - }; - provider_failure - .or_else(|| status::map(mapping)) - .unwrap_or_else(|| unmapped(mapping)) -} - -/// The `APIConnectionError` Python raises when no mapper claimed the failure: with the -/// provider prefix for a provider error, with the bare text for a plain exception. -fn unmapped(mapping: &Mapping<'_>) -> PublicFailure { - let message = match mapping.original.status { - Some(_) => format!("{} - {}", mapping.exception_provider, mapping.error_str), - None => mapping.original.message.clone(), - }; - mapping.failure(PublicKind::ApiConnection, message, false) -} - fn exception_provider(provider: &str) -> String { + if provider == "openai" { + return "OpenAIException".to_string(); + } let mut characters = provider.chars(); match characters.next() { Some(first) => format!("{}{}Exception", first.to_uppercase(), characters.as_str()), @@ -256,18 +197,6 @@ fn exception_provider(provider: &str) -> String { } } -fn python_capitalize(value: &str) -> String { - let mut characters = value.chars(); - match characters.next() { - Some(first) => format!( - "{}{}", - first.to_uppercase(), - characters.as_str().to_lowercase() - ), - None => String::new(), - } -} - fn api_base(context: &ExceptionContext) -> Option { match (&context.vertex_location, &context.vertex_project) { (Some(location), Some(project)) => Some(format!( @@ -311,132 +240,99 @@ fn extra_information(context: &ExceptionContext, api_base: Option<&str>) -> Stri #[cfg(test)] mod testing { - use super::*; + use super::Mapping; - pub(super) const DEBUG: &str = "\nModel: ocr-model"; - - pub(super) fn context(provider: &str, family: ExceptionFamily) -> ExceptionContext { - ExceptionContext { - model: "ocr-model".into(), - custom_llm_provider: Some(provider.into()), - family, - suppress_debug_info: true, - ..ExceptionContext::default() - } - } - - pub(super) fn http(status: u16, body: &str) -> OriginalException { - OriginalException::Http { + pub(super) fn mapping(status: Option, text: &str) -> Mapping { + Mapping { status, - body: body.into(), - headers: vec![("retry-after".into(), "7".into())], - } - } - - pub(super) fn upstream(status: u16, body: &str) -> Option { - Some(ResponseArg::Upstream(UpstreamResponse { - status, - body: body.into(), - headers: vec![("retry-after".into(), "7".into())], - })) - } - - pub(super) fn status(class: StatusClass, response: Option) -> PublicKind { - PublicKind::Status { - status_class: class, - response, - } - } - - /// The failure a rule builds before `exception_type` adds the response headers and the - /// banner flag. - pub(super) fn failure(kind: PublicKind, message: &str, provider: &str) -> PublicFailure { - PublicFailure { - kind, - message: message.into(), - model: "ocr-model".into(), - llm_provider: Some(provider.into()), - litellm_debug_info: None, - litellm_response_headers: None, - print_banner: false, - } - } - - pub(super) fn with_debug(failure: PublicFailure) -> PublicFailure { - PublicFailure { - litellm_debug_info: Some(DEBUG.into()), - ..failure + error_str: text.into(), } } } #[cfg(test)] mod tests { - use super::testing::{DEBUG, context, failure, http, status, upstream, with_debug}; use super::*; - fn openai() -> ExceptionContext { - context("mistral", ExceptionFamily::OpenAiCompatible) + const DEBUG: &str = "\nModel: ocr-model"; + + fn context(provider: &str) -> ExceptionContext { + ExceptionContext { + model: "ocr-model".into(), + custom_llm_provider: provider.into(), + ..ExceptionContext::default() + } + } + + fn redactor() -> SecretRedactor { + SecretRedactor::new(16) + } + + fn headers() -> Vec<(String, String)> { + vec![("retry-after".into(), "7".into())] + } + + fn http(status: u16, body: &str) -> OriginalException { + OriginalException::Http { + status, + body: body.into(), + headers: headers(), + } + } + + fn upstream(status: u16, body: &str) -> Option { + Some(UpstreamResponse { + status, + body: body.into(), + headers: headers(), + }) + } + + fn mapped(provider: &str, original: &OriginalException) -> MappedFailure { + exception_type(&context(provider), Some(&redactor()), original) + } + + #[rstest::rstest] + #[case::openai_family("mistral", "rate limit reached", PublicError::RateLimit)] + #[case::vertex_family("vertex_ai", "Resource exhausted", PublicError::RateLimit)] + #[case::cohere_family("cohere", "too many tokens", PublicError::ContextWindowExceeded)] + fn a_family_text_rule_beats_the_status_and_keeps_the_real_response( + #[case] provider: &str, + #[case] body: &str, + #[case] expected: PublicError, + ) { + let failure = mapped(provider, &http(401, body)); + assert_eq!(failure.error, expected); + assert_eq!(failure.upstream, upstream(401, body)); } #[test] - fn a_public_original_passes_through_without_banner_debug_or_prefix() { - let original = OriginalException::Public { - class: StatusClass::UnsupportedParams, - message: "Invalid `req_format`".into(), - }; - let context = ExceptionContext { - suppress_debug_info: false, - ..openai() - }; + fn the_other_family_has_no_text_rules() { assert_eq!( - exception_type(&context, &original), - failure( - status(StatusClass::UnsupportedParams, None), - "Invalid `req_format`", - "mistral" - ) + mapped("reducto", &http(401, "rate limit reached")).error, + PublicError::Authentication ); } #[rstest::rstest] - #[case::vertex_family_status_rule(ExceptionFamily::VertexAi, "vertex_ai", PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "Vertex_aiException - rejected", "vertex_ai")) - })] - #[case::cohere_family(ExceptionFamily::Cohere, "cohere", PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "CohereException - rejected", "cohere")) - })] - #[case::other_family(ExceptionFamily::Other, "reducto", PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure(status(StatusClass::BadRequest, upstream(409, "rejected")), "ReductoException - rejected", "reducto")) - })] - fn families_without_a_409_rule_reach_the_status_table( - #[case] family: ExceptionFamily, + #[case::openai_403_is_permission_denied("mistral", 403, PublicError::PermissionDenied)] + #[case::openai_409_is_bad_request("mistral", 409, PublicError::BadRequest)] + #[case::vertex_502_is_bad_gateway("vertex_ai", 502, PublicError::BadGateway)] + #[case::vertex_504_is_a_timeout("vertex_ai", 504, PublicError::Timeout { status: 504 })] + #[case::cohere_498_is_bad_request("cohere", 498, PublicError::BadRequest)] + #[case::other_503("reducto", 503, PublicError::ServiceUnavailable)] + fn without_a_text_rule_every_family_uses_the_status_table( #[case] provider: &str, - #[case] expected: PublicFailure, + #[case] status: u16, + #[case] expected: PublicError, ) { assert_eq!( - exception_type(&context(provider, family), &http(409, "rejected")), - expected - ); - } - - #[test] - fn the_openai_family_claims_a_409_before_the_status_table() { - assert_eq!( - exception_type(&openai(), &http(409, "rejected")), - PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure( - PublicKind::Api { - status: 409, - request_url: DOCS_URL - }, - "APIError: MistralException - rejected", - "mistral" - )) + mapped(provider, &http(status, "rejected")), + MappedFailure { + error: expected, + message: format!("{} - rejected", exception_provider(provider)), + upstream: upstream(status, "rejected"), + debug_info: DEBUG.into(), } ); } @@ -447,148 +343,126 @@ mod tests { #[case::timed_out_generating("Timed out generating response")] #[case::read_operation("The read operation timed out")] fn timeout_markers_win_over_every_family(#[case] marker: &str) { - let body = format!("rate limit {marker}"); - for family in [ - ExceptionFamily::OpenAiCompatible, - ExceptionFamily::VertexAi, - ExceptionFamily::Cohere, - ExceptionFamily::Other, - ] { + let body = format!("rate limit invalid api token {marker}"); + for provider in ["mistral", "vertex_ai", "cohere", "reducto"] { assert_eq!( - exception_type(&context("mistral", family), &http(429, &body)), - PublicFailure { - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure( - PublicKind::Timeout { status: None }, - &format!("APITimeoutError - Request timed out. Error_str: {body}"), - "mistral" - )) - } + mapped(provider, &http(429, &body)).error, + PublicError::Timeout { status: 408 }, + "{provider}" ); } } + #[test] + fn a_handler_timeout_is_a_408_without_a_response() { + let original = OriginalException::Timeout { + timeout_seconds: Some(0.5), + elapsed_seconds: Some(0.5031), + }; + assert_eq!( + mapped("mistral", &original), + MappedFailure { + error: PublicError::Timeout { status: 408 }, + message: "MistralException - Connection timed out after 0.5 seconds.".into(), + upstream: None, + debug_info: DEBUG.into(), + } + ); + } + #[rstest::rstest] - #[case::provider_error_keeps_the_prefix(http(409, "rejected"), "ReductoException - rejected")] - #[case::synthesized_status_skips_the_status_table( - OriginalException::Connection { message: "refused".into() }, - "ReductoException - refused" - )] - #[case::plain_exception_keeps_its_text( - OriginalException::Local { class: LocalClass::FileNotFound, message: "File not found: /a".into() }, - "File not found: /a" - )] - fn unmapped_failures_are_connection_errors( + #[case::refused_connection(OriginalException::Connection { message: "refused".into() })] + #[case::unparseable_response(OriginalException::Plain { message: "refused".into() })] + #[case::informational_status(OriginalException::Http { status: 399, body: "refused".into(), headers: Vec::new() })] + fn a_failure_no_rule_or_status_claims_is_a_connection_error( #[case] original: OriginalException, - #[case] message: &str, ) { - let context = context("reducto", ExceptionFamily::Other); - let expected = match original { - OriginalException::Http { .. } => with_debug(failure( - status(StatusClass::BadRequest, upstream(409, "rejected")), - message, - "reducto", - )), - OriginalException::Connection { .. } | OriginalException::Local { .. } => { - failure(PublicKind::ApiConnection, message, "reducto") - } - _ => unreachable!(), - }; - let actual = exception_type(&context, &original); - assert_eq!( - PublicFailure { - litellm_response_headers: None, - ..actual - }, - expected - ); + let failure = mapped("reducto", &original); + assert_eq!(failure.error, PublicError::ApiConnection); + assert_eq!(failure.message, "ReductoException - refused"); } #[test] - fn a_missing_provider_renders_like_python_none() { - let context = ExceptionContext { - custom_llm_provider: None, - family: ExceptionFamily::Other, - ..openai() + fn a_timeout_marker_on_a_response_keeps_the_response() { + let failure = mapped("reducto", &http(429, "Request timed out")); + assert_eq!(failure.error, PublicError::Timeout { status: 408 }); + assert_eq!(failure.upstream, upstream(429, "Request timed out")); + } + + #[test] + fn family_text_rules_also_classify_failures_without_a_response() { + let original = OriginalException::Plain { + message: "Request too large".into(), }; - assert_eq!( - exception_type(&context, &http(401, "rejected")), - PublicFailure { - llm_provider: None, - litellm_response_headers: Some(vec![("retry-after".into(), "7".into())]), - ..with_debug(failure( - status(StatusClass::Authentication, upstream(401, "rejected")), - "None - rejected", - "unused" - )) - } - ); + assert_eq!(mapped("mistral", &original).error, PublicError::RateLimit); } #[rstest::rstest] - #[case::suppressed(true, false)] - #[case::printed(false, true)] - fn the_banner_prints_unless_debug_info_is_suppressed( - #[case] suppress_debug_info: bool, - #[case] print_banner: bool, - ) { - let context = ExceptionContext { - suppress_debug_info, - ..openai() - }; + #[case::openai_family("mistral", "MistralException - rejected REDACTED")] + #[case::vertex_family("vertex_ai", "Vertex_aiException - rejected REDACTED")] + #[case::other_family("reducto", "ReductoException - rejected REDACTED")] + fn every_family_reports_the_redacted_text(#[case] provider: &str, #[case] message: &str) { + let failure = mapped(provider, &http(400, "rejected Bearer abcdefghijklmnop")); + assert_eq!(failure.message, message); + } + + #[test] + fn redaction_runs_before_the_rules_see_the_text() { + let body = "db_password=rate_limit"; assert_eq!( - exception_type(&context, &http(400, "rejected")).print_banner, - print_banner + mapped("mistral", &http(400, body)).error, + PublicError::BadRequest + ); + assert_eq!( + exception_type(&context("mistral"), None, &http(400, body)).error, + PublicError::RateLimit ); } #[test] - fn empty_upstream_headers_are_not_reported() { - let original = OriginalException::Http { - status: 400, - body: "rejected".into(), - headers: Vec::new(), - }; - assert_eq!( - exception_type(&openai(), &original).litellm_response_headers, - None - ); - } - - #[test] - fn messages_are_redacted_before_markers_and_prefixes() { + fn without_a_redactor_the_text_is_kept() { let body = "rejected Bearer abcdefghijklmnop"; assert_eq!( - exception_type( - &context("reducto", ExceptionFamily::Other), - &http(400, body) - ) - .message, - "ReductoException - rejected REDACTED" + exception_type(&context("reducto"), None, &http(400, body)).message, + format!("ReductoException - {body}") ); } - const SYNC_TIMEOUT: &str = "litellm.Timeout: Connection timed out after 0.5 seconds."; + #[test] + fn a_rule_hint_follows_the_message() { + let failure = mapped("mistral", &http(400, "invalid_encrypted_content")); + assert_eq!(failure.error, PublicError::BadRequest); + assert!( + failure + .message + .starts_with("MistralException - invalid_encrypted_content\n\n This error occurs") + ); + } #[rstest::rstest] - #[case::sync(false, Some(0.5), Some(0.5031), SYNC_TIMEOUT)] + #[case::sync( + false, + Some(0.5), + Some(0.5031), + "Connection timed out after 0.5 seconds." + )] #[case::async_rounds_the_elapsed_time( true, Some(0.5), Some(0.5031), - "litellm.Timeout: Connection timed out. Timeout passed=0.5, time taken=0.503 seconds" + "Connection timed out. Timeout passed=0.5, time taken=0.503 seconds" )] #[case::whole_seconds_keep_a_decimal( true, Some(600.0), Some(2.0), - "litellm.Timeout: Connection timed out. Timeout passed=600.0, time taken=2.0 seconds" + "Connection timed out. Timeout passed=600.0, time taken=2.0 seconds" )] #[case::unknown_values_render_as_none( true, None, None, - "litellm.Timeout: Connection timed out. Timeout passed=None, time taken=None seconds" + "Connection timed out. Timeout passed=None, time taken=None seconds" )] fn timeout_text_follows_the_delivery_mode( #[case] asynchronous: bool, @@ -602,58 +476,6 @@ mod tests { ); } - #[rstest::rstest] - #[case::sync(false, SYNC_TIMEOUT)] - #[case::async_( - true, - "litellm.Timeout: Connection timed out. Timeout passed=0.5, time taken=0.503 seconds" - )] - fn a_timeout_is_a_408_carrying_the_handler_text( - #[case] asynchronous: bool, - #[case] text: &str, - ) { - let context = ExceptionContext { - asynchronous, - ..openai() - }; - let original = OriginalException::Timeout { - timeout_seconds: Some(0.5), - elapsed_seconds: Some(0.5031), - }; - assert_eq!( - exception_type(&context, &original), - with_debug(failure( - PublicKind::Timeout { status: None }, - &format!("Timeout Error: MistralException - {text}"), - "mistral" - )) - ); - } - - #[test] - fn a_refused_connection_is_a_500_with_an_empty_response() { - assert_eq!( - exception_type( - &openai(), - &OriginalException::Connection { - message: "refused".into() - } - ), - with_debug(failure( - status( - StatusClass::InternalServer, - Some(ResponseArg::Upstream(UpstreamResponse { - status: 500, - body: String::new(), - headers: Vec::new(), - })) - ), - "InternalServerError: MistralException - refused", - "mistral" - )) - ); - } - #[test] fn debug_information_follows_the_python_layout() { let context = ExceptionContext { @@ -662,10 +484,10 @@ mod tests { model_group: Some("ocr".into()), deployment: Some("deployment".into()), user_api_key_alias: Some("key".into()), - ..openai() + ..context("vertex_ai") }; assert_eq!( - extra_information(&context, api_base(&context).as_deref()), + exception_type(&context, None, &http(400, "rejected")).debug_info, concat!( "\n\nKey Name: `key`\nTeam: `None`", "\nModel: ocr-model", @@ -680,21 +502,20 @@ mod tests { #[rstest::rstest] #[case::bare(ExceptionContext::default(), "\nModel: ")] - #[case::redacted_messages(ExceptionContext { redact_messages_in_exceptions: true, model: "m".into(), ..ExceptionContext::default() }, "\nModel: m")] #[case::team_alias( - ExceptionContext { model: "m".into(), user_api_key_alias: Some("key".into()), user_api_key_team_alias: Some("team".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + ExceptionContext { model: "m".into(), user_api_key_alias: Some("key".into()), user_api_key_team_alias: Some("team".into()), ..ExceptionContext::default() }, "\n\nKey Name: `key`\nTeam: `team`\nModel: m" )] #[case::team_alias_without_key_is_ignored( - ExceptionContext { model: "m".into(), user_api_key_team_alias: Some("team".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + ExceptionContext { model: "m".into(), user_api_key_team_alias: Some("team".into()), ..ExceptionContext::default() }, "\nModel: m" )] #[case::project_without_location_has_no_api_base( - ExceptionContext { model: "m".into(), vertex_project: Some("p".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + ExceptionContext { model: "m".into(), vertex_project: Some("p".into()), ..ExceptionContext::default() }, "\nModel: m\nvertex_project: `p`\n" )] #[case::location_without_project_has_no_api_base( - ExceptionContext { model: "m".into(), vertex_location: Some("l".into()), redact_messages_in_exceptions: true, ..ExceptionContext::default() }, + ExceptionContext { model: "m".into(), vertex_location: Some("l".into()), ..ExceptionContext::default() }, "\nModel: m\nvertex_location: `l`\n" )] fn each_optional_context_field_adds_its_own_line( @@ -708,6 +529,7 @@ mod tests { } #[rstest::rstest] + #[case::openai_keeps_its_brand("openai", "OpenAIException")] #[case::lowercase("mistral", "MistralException")] #[case::keeps_the_rest("azure_ai", "Azure_aiException")] #[case::empty("", "")] @@ -717,17 +539,4 @@ mod tests { ) { assert_eq!(exception_provider(provider), expected); } - - #[rstest::rstest] - #[case::lowers_the_rest("vERTEX_AI", "Vertex_ai")] - #[case::empty("", "")] - fn python_capitalize_lowers_the_rest(#[case] value: &str, #[case] expected: &str) { - assert_eq!(python_capitalize(value), expected); - } - - #[test] - fn debug_constant_matches_the_default_test_context() { - let context = openai(); - assert_eq!(extra_information(&context, None), DEBUG); - } } diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs index ffbb172582b..b4f4496dbf0 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs @@ -1,85 +1,31 @@ -use super::public::{PublicFailure, StatusClass}; -use super::rules::{ - ApiStatus, Kind, ResponseChoice, Rule, apply, contains_any, is_context_window_exceeded, - is_rate_limit, -}; -use super::{DOCS_URL, Mapping}; - -const OPENAI_URL: &str = "https://api.openai.com/v1"; +use super::public::PublicError; +use super::rules::{Rule, contains_any, is_context_window_exceeded, is_rate_limit}; const ENCRYPTED_CONTENT_HELP: &str = "\n\n This error occurs when load balancing Responses API across deployments with different API keys.\n Encrypted content is tied to the organization that created it and cannot be decrypted by other organizations.\n\n Solution: Enable 'encrypted_content_affinity' to route follow-up requests to the correct deployment:\n\n router_settings:\n enable_pre_call_checks: true\n optional_pre_call_checks:\n - encrypted_content_affinity\n\n Learn more: https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing"; -const fn with_response(class: StatusClass) -> Kind { - Kind::Status { - class, - response: ResponseChoice::Provider, - } -} - -fn exception_provider(mapping: &Mapping<'_>) -> String { - if mapping.provider == "openai" { - "OpenAIException".to_string() - } else { - super::exception_provider(mapping.provider) - } -} - -/// The raw message with OpenAI's own names swapped for the provider's. -fn message(mapping: &Mapping<'_>) -> String { - let provider = mapping.provider; - mapping - .original - .message - .replace("OPENAI", &provider.to_uppercase()) - .replace("openai.OpenAIError", &format!("{provider}.{provider}Error")) -} - -fn prefixed(mapping: &Mapping<'_>, label: &str) -> String { - format!( - "{label}{} - {}", - exception_provider(mapping), - message(mapping) - ) -} - -fn status_is(mapping: &Mapping<'_>, statuses: &[u16]) -> bool { - mapping - .original - .status - .is_some_and(|status| statuses.contains(&status)) -} - -/// `_map_openai_exception`, in its branch order. -const RULES: &[Rule] = &[ - Rule { - when: |mapping| is_rate_limit(&mapping.error_str, mapping.original.status), - kind: with_response(StatusClass::RateLimit), - message: |mapping| prefixed(mapping, "RateLimitError: "), - debug: false, - }, - Rule { - when: |mapping| is_context_window_exceeded(&mapping.error_str), - kind: with_response(StatusClass::ContextWindowExceeded), - message: |mapping| prefixed(mapping, "ContextWindowExceededError: "), - debug: true, - }, - Rule { - when: |mapping| { +/// The text branches of `_map_openai_exception`, in its order. +pub(super) const RULES: &[Rule] = &[ + Rule::new( + |mapping| is_rate_limit(&mapping.error_str, mapping.status), + PublicError::RateLimit, + ), + Rule::new( + |mapping| is_context_window_exceeded(&mapping.error_str), + PublicError::ContextWindowExceeded, + ), + Rule::new( + |mapping| { mapping.error_str.contains("invalid_request_error") && mapping.error_str.contains("model_not_found") }, - kind: with_response(StatusClass::NotFound), - message: |mapping| prefixed(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| mapping.error_str.contains("A timeout occurred"), - kind: Kind::Timeout(None), - message: |mapping| prefixed(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::NotFound, + ), + Rule::new( + |mapping| mapping.error_str.contains("A timeout occurred"), + PublicError::Timeout { status: 408 }, + ), + Rule::new( + |mapping| { let error_str = &mapping.error_str; (error_str.contains("invalid_request_error") && error_str.contains("content_policy_violation")) @@ -89,38 +35,29 @@ const RULES: &[Rule] = &[ .to_lowercase() .contains("request was rejected as a result of the safety system") }, - kind: with_response(StatusClass::ContentPolicyViolation), - message: |mapping| prefixed(mapping, "ContentPolicyViolationError: "), - debug: true, - }, + PublicError::ContentPolicyViolation, + ), Rule { - when: |mapping| { - contains_any( - &mapping.error_str, - &["invalid_encrypted_content", "could not be verified"], - ) - }, - kind: with_response(StatusClass::BadRequest), - message: |mapping| { - format!( - "{} - {}{ENCRYPTED_CONTENT_HELP}", - exception_provider(mapping), - message(mapping) - ) - }, - debug: true, + hint: ENCRYPTED_CONTENT_HELP, + ..Rule::new( + |mapping| { + contains_any( + &mapping.error_str, + &["invalid_encrypted_content", "could not be verified"], + ) + }, + PublicError::BadRequest, + ) }, - Rule { - when: |mapping| { + Rule::new( + |mapping| { mapping.error_str.contains("invalid_request_error") && !mapping.error_str.contains("Incorrect API key provided") }, - kind: with_response(StatusClass::BadRequest), - message: |mapping| prefixed(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::BadRequest, + ), + Rule::new( + |mapping| { contains_any( &mapping.error_str, &[ @@ -129,458 +66,127 @@ const RULES: &[Rule] = &[ ], ) }, - kind: Kind::Status { - class: StatusClass::InternalServer, - response: ResponseChoice::Omitted, - }, - message: |mapping| prefixed(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| mapping.error_str.contains("Request too large"), - kind: with_response(StatusClass::RateLimit), - message: |mapping| prefixed(mapping, "RateLimitError: "), - debug: true, - }, - Rule { - when: |mapping| { - mapping.error_str.contains("The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable") - }, - kind: with_response(StatusClass::Authentication), - message: |mapping| prefixed(mapping, "AuthenticationError: "), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::InternalServer, + ), + Rule::new( + |mapping| mapping.error_str.contains("Request too large"), + PublicError::RateLimit, + ), + Rule::new( + |mapping| { mapping .error_str .contains("Mistral API raised a streaming error") }, - kind: Kind::Api { - status: ApiStatus::Fixed(500), - request_url: OPENAI_URL, - }, - message: |mapping| prefixed(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| mapping.original.status.is_none(), - kind: Kind::ApiConnection, - message: |mapping| prefixed(mapping, "APIConnectionError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[400, 422]), - kind: with_response(StatusClass::BadRequest), - message: |mapping| prefixed(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[401]), - kind: with_response(StatusClass::Authentication), - message: |mapping| prefixed(mapping, "AuthenticationError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[404]), - kind: with_response(StatusClass::NotFound), - message: |mapping| prefixed(mapping, "NotFoundError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[408]), - kind: Kind::Timeout(None), - message: |mapping| prefixed(mapping, "Timeout Error: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[429]), - kind: with_response(StatusClass::RateLimit), - message: |mapping| prefixed(mapping, "RateLimitError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[500]), - kind: with_response(StatusClass::InternalServer), - message: |mapping| prefixed(mapping, "InternalServerError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[502]), - kind: with_response(StatusClass::BadGateway), - message: |mapping| prefixed(mapping, "BadGatewayError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[503]), - kind: with_response(StatusClass::ServiceUnavailable), - message: |mapping| prefixed(mapping, "ServiceUnavailableError: "), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, &[504]), - kind: Kind::Timeout(Some(504)), - message: |mapping| prefixed(mapping, "Timeout Error: "), - debug: true, - }, - Rule { - when: |_| true, - kind: Kind::Api { - status: ApiStatus::Original, - request_url: DOCS_URL, - }, - message: |mapping| prefixed(mapping, "APIError: "), - debug: true, - }, + PublicError::Api { status: 500 }, + ), ]; -pub(super) fn map(mapping: &Mapping<'_>) -> Option { - apply(RULES, mapping) -} - #[cfg(test)] mod tests { - use super::super::testing::{context, failure, http, status, upstream, with_debug}; - use super::super::{ExceptionFamily, OriginalException, PublicKind}; + use super::super::rules::first_match; + use super::super::testing::mapping; use super::*; - fn mapped(provider: &str, original: &OriginalException) -> PublicFailure { - let context = context(provider, ExceptionFamily::OpenAiCompatible); - map(&Mapping::new(&context, original)).expect("the OpenAI table ends in a catch-all") - } - - fn kind(class: StatusClass, status_code: u16, body: &str) -> PublicKind { - status(class, upstream(status_code, body)) + fn classified(status: Option, text: &str) -> Option { + first_match(RULES, &mapping(status, text)).map(|rule| rule.error) } #[rstest::rstest] - #[case::rate_limit_phrase( - 400, - "rate limit reached", - failure( - kind(StatusClass::RateLimit, 400, "rate limit reached"), - "RateLimitError: MistralException - rate limit reached", - "mistral", - ) - )] + #[case::rate_limit_phrase("rate limit reached", PublicError::RateLimit)] #[case::context_window( - 500, "This model's maximum context length is 10", - with_debug(failure( - kind( - StatusClass::ContextWindowExceeded, - 500, - "This model's maximum context length is 10" - ), - "ContextWindowExceededError: MistralException - This model's maximum context length is 10", - "mistral", - )) + PublicError::ContextWindowExceeded )] - #[case::model_not_found( - 400, - "invalid_request_error model_not_found", - with_debug(failure( - kind(StatusClass::NotFound, 400, "invalid_request_error model_not_found"), - "MistralException - invalid_request_error model_not_found", - "mistral", - )) - )] - #[case::timeout_occurred(400, "A timeout occurred", with_debug(failure( - PublicKind::Timeout { status: None }, - "MistralException - A timeout occurred", - "mistral", - )))] + #[case::model_not_found("invalid_request_error model_not_found", PublicError::NotFound)] + #[case::timeout_occurred("A timeout occurred", PublicError::Timeout { status: 408 })] #[case::content_policy_error_code( - 400, "invalid_request_error content_policy_violation", - with_debug(failure( - kind( - StatusClass::ContentPolicyViolation, - 400, - "invalid_request_error content_policy_violation" - ), - "ContentPolicyViolationError: MistralException - invalid_request_error content_policy_violation", - "mistral", - )) + PublicError::ContentPolicyViolation )] #[case::content_policy_usage_policy( - 400, "Invalid prompt violating our usage policy", - with_debug(failure( - kind( - StatusClass::ContentPolicyViolation, - 400, - "Invalid prompt violating our usage policy" - ), - "ContentPolicyViolationError: MistralException - Invalid prompt violating our usage policy", - "mistral", - )) + PublicError::ContentPolicyViolation )] #[case::content_policy_safety_system( - 400, "Request was rejected as a result of the safety system", - with_debug(failure( - kind( - StatusClass::ContentPolicyViolation, - 400, - "Request was rejected as a result of the safety system" - ), - "ContentPolicyViolationError: MistralException - Request was rejected as a result of the safety system", - "mistral", - )) - )] - #[case::encrypted_content(400, "invalid_encrypted_content", with_debug(failure( - kind(StatusClass::BadRequest, 400, "invalid_encrypted_content"), - &format!("MistralException - invalid_encrypted_content{ENCRYPTED_CONTENT_HELP}"), - "mistral", - )))] - #[case::unverifiable_content(400, "could not be verified", with_debug(failure( - kind(StatusClass::BadRequest, 400, "could not be verified"), - &format!("MistralException - could not be verified{ENCRYPTED_CONTENT_HELP}"), - "mistral", - )))] - #[case::invalid_request( - 429, - "invalid_request_error bad field", - with_debug(failure( - kind(StatusClass::BadRequest, 429, "invalid_request_error bad field"), - "MistralException - invalid_request_error bad field", - "mistral", - )) + PublicError::ContentPolicyViolation )] + #[case::encrypted_content("invalid_encrypted_content", PublicError::BadRequest)] + #[case::unverifiable_content("could not be verified", PublicError::BadRequest)] + #[case::invalid_request("invalid_request_error bad field", PublicError::BadRequest)] #[case::unknown_server_error( - 400, "Web server is returning an unknown error", - failure( - status(StatusClass::InternalServer, None), - "MistralException - Web server is returning an unknown error", - "mistral", - ) + PublicError::InternalServer )] #[case::server_had_an_error( - 400, "The server had an error processing your request.", - failure( - status(StatusClass::InternalServer, None), - "MistralException - The server had an error processing your request.", - "mistral", - ) + PublicError::InternalServer )] - #[case::request_too_large( - 400, - "Request too large", - with_debug(failure( - kind(StatusClass::RateLimit, 400, "Request too large"), - "RateLimitError: MistralException - Request too large", - "mistral", - )) + #[case::request_too_large("Request too large", PublicError::RateLimit)] + #[case::mistral_streaming_error( + "Mistral API raised a streaming error", + PublicError::Api { status: 500 } )] - #[case::missing_client_api_key( - 400, - "The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable", - with_debug(failure( - kind( - StatusClass::Authentication, - 400, - "The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable", - ), - "AuthenticationError: MistralException - The api_key client option must be set either by passing api_key to the client or by setting the MISTRAL_API_KEY environment variable", - "mistral", - )) - )] - #[case::mistral_streaming_error(400, "Mistral API raised a streaming error", with_debug(failure( - PublicKind::Api { status: 500, request_url: OPENAI_URL }, - "MistralException - Mistral API raised a streaming error", - "mistral", - )))] - fn each_text_rule_maps_by_the_body( - #[case] status_code: u16, - #[case] body: &str, - #[case] expected: PublicFailure, - ) { - assert_eq!(mapped("mistral", &http(status_code, body)), expected); - } - - #[rstest::rstest] - #[case::bad_request( - 400, - kind(StatusClass::BadRequest, 400, "rejected"), - "MistralException - rejected" - )] - #[case::unprocessable( - 422, - kind(StatusClass::BadRequest, 422, "rejected"), - "MistralException - rejected" - )] - #[case::authentication( - 401, - kind(StatusClass::Authentication, 401, "rejected"), - "AuthenticationError: MistralException - rejected" - )] - #[case::not_found( - 404, - kind(StatusClass::NotFound, 404, "rejected"), - "NotFoundError: MistralException - rejected" - )] - #[case::request_timeout(408, PublicKind::Timeout { status: None }, "Timeout Error: MistralException - rejected")] - #[case::rate_limited( - 429, - kind(StatusClass::RateLimit, 429, "rejected"), - "RateLimitError: MistralException - rejected" - )] - #[case::internal_server( - 500, - kind(StatusClass::InternalServer, 500, "rejected"), - "InternalServerError: MistralException - rejected" - )] - #[case::bad_gateway( - 502, - kind(StatusClass::BadGateway, 502, "rejected"), - "BadGatewayError: MistralException - rejected" - )] - #[case::service_unavailable( - 503, - kind(StatusClass::ServiceUnavailable, 503, "rejected"), - "ServiceUnavailableError: MistralException - rejected" - )] - #[case::gateway_timeout(504, PublicKind::Timeout { status: Some(504) }, "Timeout Error: MistralException - rejected")] - #[case::any_other_status(409, PublicKind::Api { status: 409, request_url: DOCS_URL }, "APIError: MistralException - rejected")] - fn each_status_rule_maps_by_the_status( - #[case] status_code: u16, - #[case] kind: PublicKind, - #[case] message: &str, - ) { - assert_eq!( - mapped("mistral", &http(status_code, "rejected")), - with_debug(failure(kind, message, "mistral")) - ); - } - - #[test] - fn a_failure_without_a_status_is_a_connection_error() { - let original = OriginalException::Response { - message: "invalid OCR response field: pages".into(), - }; - assert_eq!( - mapped("mistral", &original), - with_debug(failure( - PublicKind::ApiConnection, - "APIConnectionError: MistralException - invalid OCR response field: pages", - "mistral" - )) - ); + fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(Some(400), text), Some(expected)); } #[rstest::rstest] #[case::rate_limit_before_context_window( - 400, "rate limit and This model's maximum context length is 10", - kind( - StatusClass::RateLimit, - 400, - "rate limit and This model's maximum context length is 10" - ), - "RateLimitError: MistralException - rate limit and This model's maximum context length is 10", - false + PublicError::RateLimit )] #[case::context_window_before_content_policy( - 400, "This model's maximum context length is 10 invalid_request_error content_policy_violation", - kind( - StatusClass::ContextWindowExceeded, - 400, - "This model's maximum context length is 10 invalid_request_error content_policy_violation" - ), - "ContextWindowExceededError: MistralException - This model's maximum context length is 10 invalid_request_error content_policy_violation", - true + PublicError::ContextWindowExceeded )] #[case::model_not_found_before_invalid_request( - 400, "invalid_request_error model_not_found", - kind(StatusClass::NotFound, 400, "invalid_request_error model_not_found"), - "MistralException - invalid_request_error model_not_found", - true + PublicError::NotFound )] #[case::timeout_before_invalid_request( - 400, "A timeout occurred invalid_request_error", - PublicKind::Timeout { status: None }, - "MistralException - A timeout occurred invalid_request_error", - true + PublicError::Timeout { status: 408 } )] #[case::content_policy_before_invalid_request( - 400, "invalid_request_error content_policy_violation", - kind( - StatusClass::ContentPolicyViolation, - 400, - "invalid_request_error content_policy_violation" - ), - "ContentPolicyViolationError: MistralException - invalid_request_error content_policy_violation", - true + PublicError::ContentPolicyViolation )] - #[case::invalid_request_with_a_bad_key_falls_to_the_status( - 401, - "invalid_request_error Incorrect API key provided", - kind( - StatusClass::Authentication, - 401, - "invalid_request_error Incorrect API key provided" - ), - "AuthenticationError: MistralException - invalid_request_error Incorrect API key provided", - true + #[case::encrypted_content_before_invalid_request( + "invalid_request_error invalid_encrypted_content", + PublicError::BadRequest )] - #[case::text_rules_before_status( - 429, - "invalid_request_error bad field", - kind(StatusClass::BadRequest, 429, "invalid_request_error bad field"), - "MistralException - invalid_request_error bad field", - true - )] - #[case::echoed_429_is_not_a_rate_limit( - 400, - "token 429 in the prompt", - kind(StatusClass::BadRequest, 400, "token 429 in the prompt"), - "MistralException - token 429 in the prompt", - true - )] - fn the_earlier_rule_wins_when_two_apply( - #[case] status_code: u16, - #[case] body: &str, - #[case] kind: PublicKind, - #[case] message: &str, - #[case] debug: bool, + fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(Some(400), text), Some(expected)); + } + + #[rstest::rstest] + #[case::encrypted_content("invalid_encrypted_content", ENCRYPTED_CONTENT_HELP)] + #[case::plain_invalid_request("invalid_request_error bad field", "")] + fn only_encrypted_content_failures_carry_the_affinity_help( + #[case] text: &str, + #[case] hint: &str, ) { - let expected = failure(kind, message, "mistral"); assert_eq!( - mapped("mistral", &http(status_code, body)), - if debug { - with_debug(expected) - } else { - expected - } + first_match(RULES, &mapping(Some(400), text)).map(|rule| rule.hint), + Some(hint) ); } #[rstest::rstest] - #[case::provider_names_replace_openai( - "azure_ai", - "OPENAI said openai.OpenAIError", - "Azure_aiException - AZURE_AI said azure_ai.azure_aiError" - )] - #[case::openai_keeps_its_own_name("openai", "rejected", "OpenAIException - rejected")] - fn the_message_names_the_provider( - #[case] provider: &str, - #[case] body: &str, - #[case] message: &str, - ) { + #[case::bad_key_is_left_to_the_status("invalid_request_error Incorrect API key provided")] + #[case::echoed_429_is_not_a_rate_limit("token 429 in the prompt")] + #[case::unmarked("rejected")] + fn text_without_a_marker_is_left_to_the_status_table(#[case] text: &str) { + assert_eq!(classified(Some(400), text), None); + } + + #[test] + fn a_standalone_429_counts_only_with_a_429_status() { assert_eq!( - mapped(provider, &http(400, body)), - with_debug(failure( - kind(StatusClass::BadRequest, 400, body), - message, - provider - )) + classified(Some(429), "got 429 back"), + Some(PublicError::RateLimit) ); } } diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs index d64868c1606..82392cbd7ee 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/original.rs @@ -1,14 +1,4 @@ -use super::public::StatusClass; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum LocalClass { - ValueError, - FileNotFound, - OsError, -} - -/// A route failure in the shape Python's `exception_type` receives it, before any public -/// class is chosen. +/// A failure a Rust route produced, before any public class is chosen. #[derive(Clone, Debug, PartialEq)] pub enum OriginalException { Http { @@ -23,27 +13,132 @@ pub enum OriginalException { timeout_seconds: Option, elapsed_seconds: Option, }, - Response { - message: String, - }, - Local { - class: LocalClass, - message: String, - }, - /// A failure Python raises as a public LiteLLM exception itself, which `exception_type` - /// hands back unchanged. - Public { - class: StatusClass, + /// A failure with no HTTP response behind it, such as an unparseable body or a local + /// file error. + Plain { message: String, }, } -/// Which of the provider-specific mappers in `exception_type` a route's provider uses. -#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +/// Which provider-specific text rules apply before the shared status table. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum ExceptionFamily { OpenAiCompatible, VertexAi, Cohere, - #[default] Other, } + +/// `openai_compatible_providers` in `litellm/constants.py`. +const OPENAI_COMPATIBLE_PROVIDERS: &[&str] = &[ + "anyscale", + "groq", + "nvidia_nim", + "cerebras", + "baseten", + "sambanova", + "ai21_chat", + "ai21", + "volcengine", + "codestral", + "deepseek", + "tencent", + "deepinfra", + "perplexity", + "xinference", + "xai", + "zai", + "together_ai", + "fireworks_ai", + "empower", + "friendliai", + "azure_ai", + "github", + "litellm_proxy", + "hosted_vllm", + "llamafile", + "lm_studio", + "galadriel", + "github_copilot", + "chatgpt", + "novita", + "meta_llama", + "publicai", + "synthetic", + "tensormesh", + "apertis", + "nano-gpt", + "poe", + "chutes", + "parasail", + "libertai", + "featherless_ai", + "nscale", + "nebius", + "dashscope", + "qwencloud", + "qwen_ai_platform", + "modelscope", + "moonshot", + "v0", + "helicone", + "morph", + "lambda_ai", + "inception", + "hyperbolic", + "vercel_ai_gateway", + "aiml", + "wandb", + "cometapi", + "clarifai", + "docker_model_runner", + "ragflow", + "pinstripes", + "darkbloom", + "meta", + "cognition", + "scx-ai", +]; + +impl ExceptionFamily { + /// The provider dispatch at the top of Python's `exception_type`, in its order. + pub fn for_provider(provider: &str) -> Self { + match provider { + "openai" | "text-completion-openai" | "custom_openai" | "mistral" | "runwayml" => { + Self::OpenAiCompatible + } + provider if OPENAI_COMPATIBLE_PROVIDERS.contains(&provider) => Self::OpenAiCompatible, + "vertex_ai" | "vertex_ai_beta" | "gemini" => Self::VertexAi, + "cohere" | "cohere_chat" => Self::Cohere, + _ => Self::Other, + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[rstest::rstest] + #[case::openai("openai", ExceptionFamily::OpenAiCompatible)] + #[case::text_completion_openai("text-completion-openai", ExceptionFamily::OpenAiCompatible)] + #[case::custom_openai("custom_openai", ExceptionFamily::OpenAiCompatible)] + #[case::mistral("mistral", ExceptionFamily::OpenAiCompatible)] + #[case::runwayml("runwayml", ExceptionFamily::OpenAiCompatible)] + #[case::listed_compatible("azure_ai", ExceptionFamily::OpenAiCompatible)] + #[case::compatible_list_wins_over_its_own_mapper( + "together_ai", + ExceptionFamily::OpenAiCompatible + )] + #[case::vertex_ai("vertex_ai", ExceptionFamily::VertexAi)] + #[case::vertex_ai_beta("vertex_ai_beta", ExceptionFamily::VertexAi)] + #[case::gemini("gemini", ExceptionFamily::VertexAi)] + #[case::cohere("cohere", ExceptionFamily::Cohere)] + #[case::cohere_chat("cohere_chat", ExceptionFamily::Cohere)] + #[case::unported_mapper("anthropic", ExceptionFamily::Other)] + #[case::unknown("reducto", ExceptionFamily::Other)] + #[case::empty("", ExceptionFamily::Other)] + fn provider_selects_the_family(#[case] provider: &str, #[case] family: ExceptionFamily) { + assert_eq!(ExceptionFamily::for_provider(provider), family); + } +} diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs index a7319c1287b..c567185aa6d 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/public.rs @@ -1,262 +1,77 @@ -use serde::Serialize; - -/// The public LiteLLM classes built from a status code alone: every one takes the same -/// constructor arguments. -#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, strum::EnumIter, strum::IntoStaticStr)] -#[serde(rename_all = "snake_case")] -#[strum(serialize_all = "snake_case")] -pub enum StatusClass { +/// The public LiteLLM exception classes a Rust route failure can become. Python builds the +/// class; Rust decides which one. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum PublicError { BadRequest, + ContextWindowExceeded, + ContentPolicyViolation, Authentication, PermissionDenied, NotFound, + Timeout { status: u16 }, RateLimit, - ContextWindowExceeded, - ContentPolicyViolation, InternalServer, BadGateway, ServiceUnavailable, - UnsupportedParams, + ApiConnection, + Api { status: u16 }, } -impl StatusClass { - /// The `status_code` the Python class sets on itself. +impl PublicError { + /// The `status_code` the Python class carries. pub const fn status_code(self) -> u16 { match self { - Self::BadRequest - | Self::ContextWindowExceeded - | Self::ContentPolicyViolation - | Self::UnsupportedParams => 400, + Self::BadRequest | Self::ContextWindowExceeded | Self::ContentPolicyViolation => 400, Self::Authentication => 401, Self::PermissionDenied => 403, Self::NotFound => 404, Self::RateLimit => 429, - Self::InternalServer => 500, + Self::InternalServer | Self::ApiConnection => 500, Self::BadGateway => 502, Self::ServiceUnavailable => 503, + Self::Timeout { status } | Self::Api { status } => status, } } } -#[derive(Clone, Debug, PartialEq, Eq, Serialize)] +#[derive(Clone, Debug, PartialEq, Eq)] pub struct UpstreamResponse { pub status: u16, pub body: String, pub headers: Vec<(String, String)>, } -#[derive(Clone, Debug, PartialEq, Eq, Serialize)] -pub struct HttpStub { - pub status: u16, - pub method: &'static str, - pub url: &'static str, - pub content: Option, -} - -#[derive(Clone, Debug, PartialEq, Eq, Serialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum ResponseArg { - Upstream(UpstreamResponse), - Stub(HttpStub), -} - -#[derive(Clone, Debug, PartialEq, Eq, Serialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum PublicKind { - Status { - status_class: StatusClass, - response: Option, - }, - Timeout { - status: Option, - }, - ApiConnection, - Api { - status: u16, - request_url: &'static str, - }, -} - -/// Constructor arguments for the public LiteLLM exception, as `exception_type` passes them. -#[derive(Clone, Debug, PartialEq, Eq, Serialize)] -pub struct PublicFailure { - pub kind: PublicKind, +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct MappedFailure { + pub error: PublicError, pub message: String, - pub model: String, - pub llm_provider: Option, - pub litellm_debug_info: Option, - pub litellm_response_headers: Option>, - pub print_banner: bool, + pub upstream: Option, + pub debug_info: String, } #[cfg(test)] mod tests { - use std::collections::BTreeSet; - use std::path::PathBuf; - - use serde_json::Value; - use strum::IntoEnumIterator; - use super::*; - const REGENERATE: &str = "LITELLM_REGENERATE_PUBLIC_FAILURE_FIXTURES"; - - fn fixture_directory() -> PathBuf { - PathBuf::from(env!("CARGO_MANIFEST_DIR")) - .join("../../../tests/test_litellm/rust_bridge/fixtures/public_failures") - } - - fn upstream(status: u16) -> ResponseArg { - ResponseArg::Upstream(UpstreamResponse { - status, - body: r#"{"message": "rejected"}"#.into(), - headers: vec![("retry-after".into(), "7".into())], - }) - } - - fn status_response(class: StatusClass) -> Option { - match class { - StatusClass::Authentication => None, - StatusClass::PermissionDenied => Some(ResponseArg::Stub(HttpStub { - status: 403, - method: "POST", - url: " https://cloud.google.com/vertex-ai/", - content: None, - })), - StatusClass::InternalServer => Some(ResponseArg::Stub(HttpStub { - status: 500, - method: "completion", - url: "https://github.com/BerriAI/litellm", - content: Some("upstream text".into()), - })), - class => Some(upstream(class.status_code())), - } - } - - fn failure(kind: PublicKind, name: &str) -> PublicFailure { - let headers = matches!( - &kind, - PublicKind::Status { - response: Some(ResponseArg::Upstream(_)), - .. - } - ); - PublicFailure { - kind, - message: format!("MistralException - {name}"), - model: "ocr-model".into(), - llm_provider: Some("mistral".into()), - litellm_debug_info: Some("\nModel: ocr-model".into()), - litellm_response_headers: headers.then(|| vec![("retry-after".into(), "7".into())]), - print_banner: false, - } - } - - /// One payload per public class the constructor can build; `test_failures.py` reads the - /// same files, so a shape change on either side fails there or here. - fn fixtures() -> Vec<(String, PublicFailure)> { - let statuses = StatusClass::iter().map(|class| { - let name: &'static str = class.into(); - let name = format!("status_{name}"); - let built = failure( - PublicKind::Status { - status_class: class, - response: status_response(class), - }, - &name, - ); - (name, built) - }); - let others = [ - ( - "timeout_with_status", - PublicFailure { - print_banner: true, - ..failure( - PublicKind::Timeout { status: Some(504) }, - "timeout_with_status", - ) - }, - ), - ( - "timeout_without_status", - PublicFailure { - litellm_debug_info: None, - ..failure( - PublicKind::Timeout { status: None }, - "timeout_without_status", - ) - }, - ), - ( - "api_connection", - PublicFailure { - llm_provider: None, - ..failure(PublicKind::ApiConnection, "api_connection") - }, - ), - ( - "api", - failure( - PublicKind::Api { - status: 409, - request_url: "https://docs.litellm.ai/docs", - }, - "api", - ), - ), - ] - .map(|(name, built)| (name.to_string(), built)); - statuses.chain(others).collect() - } - - #[test] - fn serialized_payloads_match_the_golden_fixtures_python_reads() { - let directory = fixture_directory(); - let regenerate = std::env::var_os(REGENERATE).is_some(); - let expected = fixtures(); - for (name, built) in &expected { - let path = directory.join(format!("{name}.json")); - let serialized = serde_json::to_value(built).unwrap(); - if regenerate { - std::fs::create_dir_all(&directory).unwrap(); - std::fs::write( - &path, - format!("{}\n", serde_json::to_string_pretty(&serialized).unwrap()), - ) - .unwrap(); - } - let golden: Value = - serde_json::from_str(&std::fs::read_to_string(&path).unwrap()).unwrap(); - assert_eq!(serialized, golden, "{name}; set {REGENERATE}=1 to rewrite"); - } - let on_disk: BTreeSet = std::fs::read_dir(&directory) - .unwrap() - .map(|entry| entry.unwrap().file_name().to_string_lossy().into_owned()) - .collect(); - let generated: BTreeSet = expected - .iter() - .map(|(name, _)| format!("{name}.json")) - .collect(); - assert_eq!(on_disk, generated); - } - #[rstest::rstest] - #[case(StatusClass::BadRequest, 400)] - #[case(StatusClass::Authentication, 401)] - #[case(StatusClass::PermissionDenied, 403)] - #[case(StatusClass::NotFound, 404)] - #[case(StatusClass::RateLimit, 429)] - #[case(StatusClass::ContextWindowExceeded, 400)] - #[case(StatusClass::ContentPolicyViolation, 400)] - #[case(StatusClass::InternalServer, 500)] - #[case(StatusClass::BadGateway, 502)] - #[case(StatusClass::ServiceUnavailable, 503)] - #[case(StatusClass::UnsupportedParams, 400)] + #[case::bad_request(PublicError::BadRequest, 400)] + #[case::context_window(PublicError::ContextWindowExceeded, 400)] + #[case::content_policy(PublicError::ContentPolicyViolation, 400)] + #[case::authentication(PublicError::Authentication, 401)] + #[case::permission_denied(PublicError::PermissionDenied, 403)] + #[case::not_found(PublicError::NotFound, 404)] + #[case::request_timeout(PublicError::Timeout { status: 408 }, 408)] + #[case::gateway_timeout(PublicError::Timeout { status: 504 }, 504)] + #[case::rate_limit(PublicError::RateLimit, 429)] + #[case::internal_server(PublicError::InternalServer, 500)] + #[case::api_connection(PublicError::ApiConnection, 500)] + #[case::bad_gateway(PublicError::BadGateway, 502)] + #[case::service_unavailable(PublicError::ServiceUnavailable, 503)] + #[case::api(PublicError::Api { status: 501 }, 501)] fn status_codes_are_the_ones_the_python_classes_set( - #[case] class: StatusClass, + #[case] error: PublicError, #[case] status: u16, ) { - assert_eq!(class.status_code(), status); + assert_eq!(error.status_code(), status); } } diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs index 34cf2c390c3..0346a8bc718 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/rules.rs @@ -4,106 +4,29 @@ use fancy_regex::Regex; use serde_json::Value; use super::Mapping; -use super::public::{HttpStub, PublicFailure, PublicKind, ResponseArg, StatusClass}; +use super::public::PublicError; -const GITHUB_URL: &str = "https://github.com/BerriAI/litellm"; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(super) enum ResponseChoice { - Omitted, - Provider, - Stub { status: u16, url: &'static str }, - InternalServerStub, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(super) enum ApiStatus { - Fixed(u16), - Original, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(super) enum Kind { - Status { - class: StatusClass, - response: ResponseChoice, - }, - Timeout(Option), - ApiConnection, - Api { - status: ApiStatus, - request_url: &'static str, - }, -} - -/// One branch of a Python `_map_*_exception` function: when it applies, the class it -/// raises, the message it builds, and whether it passes `litellm_debug_info`. +/// One text branch of a Python `_map_*_exception` function: when it applies, the class it +/// raises, and any help text appended to the message. pub(super) struct Rule { - pub(super) when: fn(&Mapping<'_>) -> bool, - pub(super) kind: Kind, - pub(super) message: fn(&Mapping<'_>) -> String, - pub(super) debug: bool, -} - -/// The first rule that applies decides the failure, as the `if`/`elif` chain does in Python. -pub(super) fn apply(rules: &[Rule], mapping: &Mapping<'_>) -> Option { - rules - .iter() - .find(|rule| (rule.when)(mapping)) - .map(|rule| rule.build(mapping)) + pub(super) when: fn(&Mapping) -> bool, + pub(super) error: PublicError, + pub(super) hint: &'static str, } impl Rule { - fn build(&self, mapping: &Mapping<'_>) -> PublicFailure { - let kind = match self.kind { - Kind::Status { class, response } => PublicKind::Status { - status_class: class, - response: response.resolve(mapping), - }, - Kind::Timeout(status) => PublicKind::Timeout { status }, - Kind::ApiConnection => PublicKind::ApiConnection, - Kind::Api { - status, - request_url, - } => PublicKind::Api { - status: match status { - ApiStatus::Fixed(status) => status, - ApiStatus::Original => mapping.original.status.unwrap_or(500), - }, - request_url, - }, - }; - PublicFailure { - kind, - message: (self.message)(mapping), - model: mapping.context.model.clone(), - llm_provider: mapping.context.custom_llm_provider.clone(), - litellm_debug_info: self.debug.then(|| mapping.extra_information.clone()), - litellm_response_headers: None, - print_banner: false, + pub(super) const fn new(when: fn(&Mapping) -> bool, error: PublicError) -> Self { + Self { + when, + error, + hint: "", } } } -impl ResponseChoice { - fn resolve(self, mapping: &Mapping<'_>) -> Option { - match self { - Self::Omitted => None, - Self::Provider => mapping.original.response.clone().map(ResponseArg::Upstream), - Self::Stub { status, url } => Some(ResponseArg::Stub(HttpStub { - status, - method: "POST", - url, - content: None, - })), - Self::InternalServerStub => Some(ResponseArg::Stub(HttpStub { - status: 500, - method: "completion", - url: GITHUB_URL, - content: Some(mapping.original.message.clone()), - })), - } - } +/// The first rule that applies decides the class, as the `if`/`elif` chain does in Python. +pub(super) fn first_match<'r>(rules: &'r [Rule], mapping: &Mapping) -> Option<&'r Rule> { + rules.iter().find(|rule| (rule.when)(mapping)) } pub(super) fn contains_any(text: &str, markers: &[&str]) -> bool { @@ -117,7 +40,7 @@ static RATE_LIMIT_PHRASE: LazyLock = /// `ExceptionCheckers.is_error_str_rate_limit`. pub(super) fn is_rate_limit(error_str: &str, status: Option) -> bool { - if STANDALONE_429.is_match(error_str).unwrap_or(false) && status == Some(429) { + if STANDALONE_429.is_match(error_str).unwrap_or(false) && matches!(status, None | Some(429)) { return true; } let lower = error_str.to_lowercase(); @@ -169,184 +92,34 @@ pub(super) fn body_error_code(error_str: &str) -> Option { #[cfg(test)] mod tests { - use super::super::testing::{context, failure, http}; - use super::super::{ExceptionFamily, OriginalException, UpstreamResponse}; + use super::super::testing::mapping; use super::*; - fn first_marker(mapping: &Mapping<'_>) -> bool { - mapping.error_str.contains("first") - } - - fn always(_: &Mapping<'_>) -> bool { - true - } - - fn text(mapping: &Mapping<'_>) -> String { - format!("seen {}", mapping.error_str) - } - const ORDERED: &[Rule] = &[ - Rule { - when: first_marker, - kind: Kind::Status { - class: StatusClass::NotFound, - response: ResponseChoice::Omitted, - }, - message: text, - debug: false, - }, - Rule { - when: always, - kind: Kind::ApiConnection, - message: text, - debug: true, - }, + Rule::new( + |mapping| mapping.error_str.contains("first"), + PublicError::NotFound, + ), + Rule::new(|_| true, PublicError::ApiConnection), ]; - fn apply_one(kind: Kind, debug: bool, original: &OriginalException) -> Option { - let context = context("mistral", ExceptionFamily::OpenAiCompatible); - let mapping = Mapping::new(&context, original); - apply( - &[Rule { - when: always, - kind, - message: text, - debug, - }], - &mapping, - ) - } - #[rstest::rstest] - #[case::earlier_rule_wins("first and second", failure( - PublicKind::Status { status_class: StatusClass::NotFound, response: None }, - "seen first and second", - "mistral", - ))] - #[case::later_rule_when_the_earlier_does_not_apply("second", PublicFailure { - litellm_debug_info: Some("\nModel: ocr-model".into()), - ..failure(PublicKind::ApiConnection, "seen second", "mistral") - })] - fn the_first_applicable_rule_decides(#[case] body: &str, #[case] expected: PublicFailure) { - let context = context("mistral", ExceptionFamily::OpenAiCompatible); - let original = http(400, body); - assert_eq!( - apply(ORDERED, &Mapping::new(&context, &original)), - Some(expected) - ); + #[case::earlier_rule_wins("first and second", PublicError::NotFound)] + #[case::later_rule_when_the_earlier_does_not_apply("second", PublicError::ApiConnection)] + fn the_first_applicable_rule_decides(#[case] text: &str, #[case] expected: PublicError) { + let rule = first_match(ORDERED, &mapping(Some(400), text)); + assert_eq!(rule.map(|rule| rule.error), Some(expected)); } #[test] fn no_applicable_rule_leaves_the_failure_to_the_caller() { - let context = context("mistral", ExceptionFamily::OpenAiCompatible); - let original = http(400, "second"); - assert_eq!( - apply(&ORDERED[..1], &Mapping::new(&context, &original)), - None - ); - } - - #[rstest::rstest] - #[case::omitted(ResponseChoice::Omitted, None)] - #[case::provider(ResponseChoice::Provider, Some(ResponseArg::Upstream(UpstreamResponse { - status: 400, - body: "body".into(), - headers: vec![("retry-after".into(), "7".into())], - })))] - #[case::stub( - ResponseChoice::Stub { status: 429, url: "https://stub.test" }, - Some(ResponseArg::Stub(HttpStub { status: 429, method: "POST", url: "https://stub.test", content: None })) - )] - #[case::internal_server_stub( - ResponseChoice::InternalServerStub, - Some(ResponseArg::Stub(HttpStub { - status: 500, - method: "completion", - url: GITHUB_URL, - content: Some("body".into()), - })) - )] - fn response_choices_resolve_against_the_original( - #[case] response: ResponseChoice, - #[case] expected: Option, - ) { - let built = apply_one( - Kind::Status { - class: StatusClass::BadRequest, - response, - }, - false, - &http(400, "body"), - ) - .unwrap(); - assert_eq!( - built.kind, - PublicKind::Status { - status_class: StatusClass::BadRequest, - response: expected, - } - ); - } - - #[rstest::rstest] - #[case::fixed(ApiStatus::Fixed(500), http(409, "body"), 500)] - #[case::original(ApiStatus::Original, http(409, "body"), 409)] - #[case::original_without_a_status( - ApiStatus::Original, - OriginalException::Response { message: "body".into() }, - 500 - )] - fn api_status_is_fixed_or_the_originals( - #[case] status: ApiStatus, - #[case] original: OriginalException, - #[case] expected: u16, - ) { - let built = apply_one( - Kind::Api { - status, - request_url: "https://api.test", - }, - false, - &original, - ) - .unwrap(); - assert_eq!( - built, - failure( - PublicKind::Api { - status: expected, - request_url: "https://api.test" - }, - "seen body", - "mistral" - ) - ); - } - - #[rstest::rstest] - #[case::with_debug(true, Some("\nModel: ocr-model"))] - #[case::without_debug(false, None)] - fn debug_rules_carry_the_extra_information( - #[case] debug: bool, - #[case] expected: Option<&str>, - ) { - let built = apply_one(Kind::Timeout(Some(504)), debug, &http(504, "body")).unwrap(); - assert_eq!( - built, - PublicFailure { - litellm_debug_info: expected.map(str::to_string), - ..failure( - PublicKind::Timeout { status: Some(504) }, - "seen body", - "mistral" - ) - } - ); + assert!(first_match(&ORDERED[..1], &mapping(Some(400), "second")).is_none()); } #[rstest::rstest] #[case::standalone_429_with_429_status("got 429 back", Some(429), true)] #[case::standalone_429_with_other_status("got 429 back", Some(400), false)] + #[case::standalone_429_with_unknown_status("got 429 back", None, true)] #[case::embedded_429("token4290", Some(429), false)] #[case::phrase_spaced("Rate Limit reached", None, true)] #[case::phrase_underscored("rate_limit", None, true)] diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs index 3d817fe9ae2..cb8924d6c2a 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/status.rs @@ -1,153 +1,49 @@ -use super::public::{PublicFailure, StatusClass}; -use super::rules::{ApiStatus, Kind, ResponseChoice, Rule, apply}; -use super::{DOCS_URL, Mapping}; +use super::public::PublicError; -const fn with_response(class: StatusClass) -> Kind { - Kind::Status { - class, - response: ResponseChoice::Provider, - } -} - -fn message(mapping: &Mapping<'_>) -> String { - format!("{} - {}", mapping.exception_provider, mapping.error_str) -} - -fn status(mapping: &Mapping<'_>) -> u16 { - mapping.original.status.unwrap_or_default() -} - -/// `_map_exception_by_status`, the fallback for a provider error no provider mapper claimed. -const RULES: &[Rule] = &[ - Rule { - when: |mapping| status(mapping) == 401, - kind: with_response(StatusClass::Authentication), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 403, - kind: with_response(StatusClass::PermissionDenied), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 404, - kind: with_response(StatusClass::NotFound), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 408, - kind: Kind::Timeout(None), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 429, - kind: with_response(StatusClass::RateLimit), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 500, - kind: with_response(StatusClass::InternalServer), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 502, - kind: with_response(StatusClass::BadGateway), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 503, - kind: with_response(StatusClass::ServiceUnavailable), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) == 504, - kind: Kind::Timeout(Some(504)), - message, - debug: true, - }, - Rule { - when: |mapping| status(mapping) < 500, - kind: with_response(StatusClass::BadRequest), - message, - debug: true, - }, - Rule { - when: |_| true, - kind: Kind::Api { - status: ApiStatus::Original, - request_url: DOCS_URL, - }, - message, - debug: true, - }, -]; - -/// Only a real provider status of 400 or more reaches the table; a status the HTTP handler -/// synthesized for a failure without a response does not. -pub(super) fn map(mapping: &Mapping<'_>) -> Option { - let status = mapping.original.status?; - if status < 400 || mapping.original.status_is_synthesized { - return None; - } - apply(RULES, mapping) +/// `_map_exception_by_status`, the one place a provider status picks a class. Statuses +/// below 400 are not failures the table claims. +pub(super) fn classify(status: u16) -> Option { + let error = match status { + ..400 => return None, + 401 => PublicError::Authentication, + 403 => PublicError::PermissionDenied, + 404 => PublicError::NotFound, + 408 | 504 => PublicError::Timeout { status }, + 429 => PublicError::RateLimit, + 500 => PublicError::InternalServer, + 502 => PublicError::BadGateway, + 503 => PublicError::ServiceUnavailable, + 400..500 => PublicError::BadRequest, + _ => PublicError::Api { status }, + }; + Some(error) } #[cfg(test)] mod tests { - use super::super::testing::{context, failure, http, upstream, with_debug}; - use super::super::{ExceptionFamily, OriginalException, PublicKind}; use super::*; - fn mapped(original: &OriginalException) -> Option { - let context = context("reducto", ExceptionFamily::Other); - map(&Mapping::new(&context, original)) - } - - fn classified(class: StatusClass, status_code: u16) -> PublicKind { - PublicKind::Status { - status_class: class, - response: upstream(status_code, "rejected"), - } - } - #[rstest::rstest] - #[case::authentication(401, classified(StatusClass::Authentication, 401))] - #[case::permission_denied(403, classified(StatusClass::PermissionDenied, 403))] - #[case::not_found(404, classified(StatusClass::NotFound, 404))] - #[case::request_timeout(408, PublicKind::Timeout { status: None })] - #[case::rate_limited(429, classified(StatusClass::RateLimit, 429))] - #[case::internal_server(500, classified(StatusClass::InternalServer, 500))] - #[case::bad_gateway(502, classified(StatusClass::BadGateway, 502))] - #[case::service_unavailable(503, classified(StatusClass::ServiceUnavailable, 503))] - #[case::gateway_timeout(504, PublicKind::Timeout { status: Some(504) })] - #[case::lowest_client_error(400, classified(StatusClass::BadRequest, 400))] - #[case::other_client_error(409, classified(StatusClass::BadRequest, 409))] - #[case::highest_client_error(499, classified(StatusClass::BadRequest, 499))] - #[case::other_server_error(501, PublicKind::Api { status: 501, request_url: DOCS_URL })] - fn every_mapped_status_and_the_fallback(#[case] status_code: u16, #[case] kind: PublicKind) { - assert_eq!( - mapped(&http(status_code, "rejected")), - Some(with_debug(failure( - kind, - "ReductoException - rejected", - "reducto" - ))) - ); - } - - #[rstest::rstest] - #[case::below_client_errors(http(399, "rejected"))] - #[case::synthesized(OriginalException::Connection { message: "refused".into() })] - #[case::no_status(OriginalException::Response { message: "bad body".into() })] - fn failures_the_table_does_not_claim(#[case] original: OriginalException) { - assert_eq!(mapped(&original), None); + #[case::below_client_errors(399, None)] + #[case::lowest_client_error(400, Some(PublicError::BadRequest))] + #[case::authentication(401, Some(PublicError::Authentication))] + #[case::permission_denied(403, Some(PublicError::PermissionDenied))] + #[case::not_found(404, Some(PublicError::NotFound))] + #[case::request_timeout(408, Some(PublicError::Timeout { status: 408 }))] + #[case::other_client_error(409, Some(PublicError::BadRequest))] + #[case::unprocessable(422, Some(PublicError::BadRequest))] + #[case::rate_limited(429, Some(PublicError::RateLimit))] + #[case::highest_client_error(499, Some(PublicError::BadRequest))] + #[case::internal_server(500, Some(PublicError::InternalServer))] + #[case::other_server_error(501, Some(PublicError::Api { status: 501 }))] + #[case::bad_gateway(502, Some(PublicError::BadGateway))] + #[case::service_unavailable(503, Some(PublicError::ServiceUnavailable))] + #[case::gateway_timeout(504, Some(PublicError::Timeout { status: 504 }))] + #[case::highest_server_error(599, Some(PublicError::Api { status: 599 }))] + fn every_mapped_status_and_the_fallback( + #[case] status: u16, + #[case] expected: Option, + ) { + assert_eq!(classify(status), expected); } } diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs index 0fa8c19c4f6..dab1adb2329 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/vertex_ai.rs @@ -1,60 +1,17 @@ -use super::public::{PublicFailure, PublicKind, ResponseArg, StatusClass, UpstreamResponse}; -use super::rules::{ - Kind, ResponseChoice, Rule, apply, body_error_code, contains_any, is_context_window_exceeded, -}; -use super::{Mapping, python_capitalize}; - -const VERTEX_URL: &str = "https://cloud.google.com/vertex-ai/"; -const VERTEX_URL_WITH_SPACE: &str = " https://cloud.google.com/vertex-ai/"; +use super::public::PublicError; +use super::rules::{Rule, body_error_code, contains_any, is_context_window_exceeded}; const QUOTA_MARKERS: &[&str] = &[ "429 Quota exceeded", "Quota exceeded for", "Resource exhausted", - "IndexError: list index out of range", "429 Unable to submit request because the service is temporarily out of capacity.", ]; -const fn stubbed(class: StatusClass, status: u16, url: &'static str) -> Kind { - Kind::Status { - class, - response: ResponseChoice::Stub { status, url }, - } -} - -const fn bare(class: StatusClass) -> Kind { - Kind::Status { - class, - response: ResponseChoice::Omitted, - } -} - -/// `{Provider}Exception{label} - {error_str}` with Python's `str.capitalize()`. -fn capitalized(mapping: &Mapping<'_>, label: &str) -> String { - format!( - "{}Exception{label} - {}", - python_capitalize(mapping.provider), - mapping.error_str - ) -} - -/// `litellm.{Class}: {provider}Exception - {error_str}` with the provider as given. -fn litellm_prefixed(mapping: &Mapping<'_>, class: &str) -> String { - format!( - "litellm.{class}: {}Exception - {}", - mapping.provider, mapping.error_str - ) -} - -fn status_is(mapping: &Mapping<'_>, status: u16) -> bool { - mapping.original.status == Some(status) -} - -/// `_map_vertex_exception`, in its branch order. A failure no rule claims falls through -/// to the status table. -const RULES: &[Rule] = &[ - Rule { - when: |mapping| { +/// The text branches of `_map_vertex_exception`, in its order. +pub(super) const RULES: &[Rule] = &[ + Rule::new( + |mapping| { contains_any( &mapping.error_str, &[ @@ -63,54 +20,32 @@ const RULES: &[Rule] = &[ ], ) }, - kind: stubbed(StatusClass::BadRequest, 400, VERTEX_URL_WITH_SPACE), - message: |mapping| litellm_prefixed(mapping, "BadRequestError"), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::BadRequest, + ), + Rule::new( + |mapping| { mapping .error_str .contains("400 Request payload size exceeds") + || is_context_window_exceeded(&mapping.error_str) }, - kind: bare(StatusClass::ContextWindowExceeded), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| is_context_window_exceeded(&mapping.error_str), - kind: bare(StatusClass::ContextWindowExceeded), - message: |mapping| format!("ContextWindowExceededError: {}", capitalized(mapping, "")), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::ContextWindowExceeded, + ), + Rule::new( + |mapping| { contains_any( &mapping.error_str, &["None Unknown Error.", "Content has no parts."], ) }, - kind: Kind::Status { - class: StatusClass::InternalServer, - response: ResponseChoice::InternalServerStub, - }, - message: |mapping| litellm_prefixed(mapping, "InternalServerError"), - debug: true, - }, - Rule { - when: |mapping| mapping.error_str.contains("API key not valid."), - kind: bare(StatusClass::Authentication), - message: |mapping| capitalized(mapping, ""), - debug: true, - }, - Rule { - when: |mapping| mapping.error_str.contains("403"), - kind: stubbed(StatusClass::BadRequest, 403, VERTEX_URL_WITH_SPACE), - message: |mapping| capitalized(mapping, " BadRequestError"), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::InternalServer, + ), + Rule::new( + |mapping| mapping.error_str.contains("API key not valid."), + PublicError::Authentication, + ), + Rule::new( + |mapping| { contains_any( &mapping.error_str, &[ @@ -119,456 +54,124 @@ const RULES: &[Rule] = &[ ], ) }, - kind: stubbed( - StatusClass::ContentPolicyViolation, - 400, - VERTEX_URL_WITH_SPACE, - ), - message: |mapping| capitalized(mapping, " ContentPolicyViolationError"), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::ContentPolicyViolation, + ), + Rule::new( + |mapping| { contains_any(&mapping.error_str, QUOTA_MARKERS) || (mapping - .original .status .is_some_and(|status| (500..600).contains(&status)) && body_error_code(&mapping.error_str) == Some(429)) }, - kind: stubbed(StatusClass::RateLimit, 429, VERTEX_URL_WITH_SPACE), - message: |mapping| litellm_prefixed(mapping, "RateLimitError"), - debug: true, - }, - Rule { - when: |mapping| { + PublicError::RateLimit, + ), + Rule::new( + |mapping| { contains_any( &mapping.error_str, &["500 Internal Server Error", "The model is overloaded."], ) }, - kind: bare(StatusClass::InternalServer), - message: |mapping| litellm_prefixed(mapping, "InternalServerError"), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, 400), - kind: stubbed(StatusClass::BadRequest, 400, VERTEX_URL), - message: |mapping| capitalized(mapping, " BadRequestError"), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, 401), - kind: bare(StatusClass::Authentication), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, 403), - kind: stubbed(StatusClass::PermissionDenied, 403, VERTEX_URL), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, 404), - kind: bare(StatusClass::NotFound), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, 408), - kind: Kind::Timeout(None), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, 429), - kind: stubbed(StatusClass::RateLimit, 429, VERTEX_URL_WITH_SPACE), - message: |mapping| format!("litellm.RateLimitError: {}", capitalized(mapping, "")), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, 500), - kind: Kind::Status { - class: StatusClass::InternalServer, - response: ResponseChoice::InternalServerStub, - }, - message: |mapping| capitalized(mapping, " InternalServerError"), - debug: true, - }, - Rule { - when: |mapping| status_is(mapping, 502), - kind: Kind::ApiConnection, - message: |mapping| capitalized(mapping, ""), - debug: false, - }, - Rule { - when: |mapping| status_is(mapping, 503), - kind: bare(StatusClass::ServiceUnavailable), - message: |mapping| capitalized(mapping, ""), - debug: false, - }, + PublicError::InternalServer, + ), ]; -pub(super) fn map(mapping: &Mapping<'_>) -> Option { - apply(RULES, mapping).map(|failure| keep_upstream_response(mapping, failure)) -} - -/// Deliberate divergence from `_map_vertex_exception`, which replaces the provider response -/// with a stub and so drops the upstream body and `retry-after`. The response keeps the -/// status the public class carries. -fn keep_upstream_response(mapping: &Mapping<'_>, failure: PublicFailure) -> PublicFailure { - let (PublicKind::Status { status_class, .. }, Some(upstream), false) = ( - &failure.kind, - &mapping.original.response, - mapping.original.status_is_synthesized, - ) else { - return failure; - }; - PublicFailure { - kind: PublicKind::Status { - status_class: *status_class, - response: Some(ResponseArg::Upstream(UpstreamResponse { - status: status_class.status_code(), - ..upstream.clone() - })), - }, - ..failure - } -} - #[cfg(test)] mod tests { - use super::super::testing::{context, failure, http, status, upstream, with_debug}; - use super::super::{ExceptionFamily, HttpStub, OriginalException}; + use super::super::rules::first_match; + use super::super::testing::mapping; use super::*; - fn mapped(original: &OriginalException) -> Option { - let context = context("vertex_ai", ExceptionFamily::VertexAi); - map(&Mapping::new(&context, original)) - } - - fn kept(class: StatusClass, body: &str) -> PublicKind { - status(class, upstream(class.status_code(), body)) + fn classified(status: Option, text: &str) -> Option { + first_match(RULES, &mapping(status, text)).map(|rule| rule.error) } #[rstest::rstest] #[case::api_not_enabled( - 400, "Vertex AI API has not been used in project x", - with_debug(failure( - kept( - StatusClass::BadRequest, - "Vertex AI API has not been used in project x" - ), - "litellm.BadRequestError: vertex_aiException - Vertex AI API has not been used in project x", - "vertex_ai", - )) - )] - #[case::project_not_found( - 400, - "Unable to find your project", - with_debug(failure( - kept(StatusClass::BadRequest, "Unable to find your project"), - "litellm.BadRequestError: vertex_aiException - Unable to find your project", - "vertex_ai", - )) + PublicError::BadRequest )] + #[case::project_not_found("Unable to find your project", PublicError::BadRequest)] #[case::payload_too_large( - 400, "400 Request payload size exceeds the limit", - failure( - kept( - StatusClass::ContextWindowExceeded, - "400 Request payload size exceeds the limit" - ), - "Vertex_aiException - 400 Request payload size exceeds the limit", - "vertex_ai", - ) + PublicError::ContextWindowExceeded )] #[case::context_window( - 500, "This model's maximum context length is 10", - with_debug(failure( - kept( - StatusClass::ContextWindowExceeded, - "This model's maximum context length is 10" - ), - "ContextWindowExceededError: Vertex_aiException - This model's maximum context length is 10", - "vertex_ai", - )) - )] - #[case::unknown_error( - 400, - "None Unknown Error.", - with_debug(failure( - kept(StatusClass::InternalServer, "None Unknown Error."), - "litellm.InternalServerError: vertex_aiException - None Unknown Error.", - "vertex_ai", - )) - )] - #[case::no_parts( - 400, - "Content has no parts.", - with_debug(failure( - kept(StatusClass::InternalServer, "Content has no parts."), - "litellm.InternalServerError: vertex_aiException - Content has no parts.", - "vertex_ai", - )) - )] - #[case::api_key_not_valid( - 400, - "API key not valid.", - with_debug(failure( - kept(StatusClass::Authentication, "API key not valid."), - "Vertex_aiException - API key not valid.", - "vertex_ai", - )) - )] - #[case::forbidden_text( - 400, - "got a 403", - with_debug(failure( - kept(StatusClass::BadRequest, "got a 403"), - "Vertex_aiException BadRequestError - got a 403", - "vertex_ai", - )) - )] - #[case::response_blocked( - 400, - "The response was blocked.", - with_debug(failure( - kept(StatusClass::ContentPolicyViolation, "The response was blocked."), - "Vertex_aiException ContentPolicyViolationError - The response was blocked.", - "vertex_ai", - )) + PublicError::ContextWindowExceeded )] + #[case::unknown_error("None Unknown Error.", PublicError::InternalServer)] + #[case::no_parts("Content has no parts.", PublicError::InternalServer)] + #[case::api_key_not_valid("API key not valid.", PublicError::Authentication)] + #[case::response_blocked("The response was blocked.", PublicError::ContentPolicyViolation)] #[case::output_blocked( - 400, "Output blocked by content filtering policy", - with_debug(failure( - kept( - StatusClass::ContentPolicyViolation, - "Output blocked by content filtering policy" - ), - "Vertex_aiException ContentPolicyViolationError - Output blocked by content filtering policy", - "vertex_ai", - )) + PublicError::ContentPolicyViolation )] - #[case::quota_marker( - 400, - "Quota exceeded for aiplatform", - with_debug(failure( - kept(StatusClass::RateLimit, "Quota exceeded for aiplatform"), - "litellm.RateLimitError: vertex_aiException - Quota exceeded for aiplatform", - "vertex_ai", - )) + #[case::quota_exceeded_429("429 Quota exceeded", PublicError::RateLimit)] + #[case::quota_exceeded_for("Quota exceeded for aiplatform", PublicError::RateLimit)] + #[case::resource_exhausted("Resource exhausted", PublicError::RateLimit)] + #[case::out_of_capacity( + "429 Unable to submit request because the service is temporarily out of capacity.", + PublicError::RateLimit )] - #[case::wrapped_429( - 503, - r#"{"error": {"code": "429"}}"#, - with_debug(failure( - kept(StatusClass::RateLimit, r#"{"error": {"code": "429"}}"#), - r#"litellm.RateLimitError: vertex_aiException - {"error": {"code": "429"}}"#, - "vertex_ai", - )) - )] - #[case::overloaded( - 400, - "The model is overloaded.", - with_debug(failure( - kept(StatusClass::InternalServer, "The model is overloaded."), - "litellm.InternalServerError: vertex_aiException - The model is overloaded.", - "vertex_ai", - )) - )] - #[case::internal_server_text( - 400, - "500 Internal Server Error", - with_debug(failure( - kept(StatusClass::InternalServer, "500 Internal Server Error"), - "litellm.InternalServerError: vertex_aiException - 500 Internal Server Error", - "vertex_ai", - )) - )] - fn each_text_rule_maps_by_the_body( - #[case] status_code: u16, - #[case] body: &str, - #[case] expected: PublicFailure, - ) { - assert_eq!(mapped(&http(status_code, body)), Some(expected)); + #[case::internal_server_text("500 Internal Server Error", PublicError::InternalServer)] + #[case::overloaded("The model is overloaded.", PublicError::InternalServer)] + fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(Some(400), text), Some(expected)); } #[rstest::rstest] - #[case::bad_request( - 400, - with_debug(failure( - kept(StatusClass::BadRequest, "rejected"), - "Vertex_aiException BadRequestError - rejected", - "vertex_ai" - )) - )] - #[case::authentication( - 401, - failure( - kept(StatusClass::Authentication, "rejected"), - "Vertex_aiException - rejected", - "vertex_ai" - ) - )] - #[case::permission_denied( - 403, - failure( - kept(StatusClass::PermissionDenied, "rejected"), - "Vertex_aiException - rejected", - "vertex_ai" - ) - )] - #[case::not_found( - 404, - failure( - kept(StatusClass::NotFound, "rejected"), - "Vertex_aiException - rejected", - "vertex_ai" - ) - )] - #[case::request_timeout(408, failure(PublicKind::Timeout { status: None }, "Vertex_aiException - rejected", "vertex_ai"))] - #[case::rate_limited( - 429, - with_debug(failure( - kept(StatusClass::RateLimit, "rejected"), - "litellm.RateLimitError: Vertex_aiException - rejected", - "vertex_ai" - )) - )] - #[case::internal_server( - 500, - with_debug(failure( - kept(StatusClass::InternalServer, "rejected"), - "Vertex_aiException InternalServerError - rejected", - "vertex_ai" - )) - )] - #[case::bad_gateway( - 502, - failure( - PublicKind::ApiConnection, - "Vertex_aiException - rejected", - "vertex_ai" - ) - )] - #[case::service_unavailable( - 503, - failure( - kept(StatusClass::ServiceUnavailable, "rejected"), - "Vertex_aiException - rejected", - "vertex_ai" - ) - )] - fn each_status_rule_maps_by_the_status( - #[case] status_code: u16, - #[case] expected: PublicFailure, + #[case::server_error_wrapping_a_429(Some(503), Some(PublicError::RateLimit))] + #[case::lowest_server_error(Some(500), Some(PublicError::RateLimit))] + #[case::highest_server_error(Some(599), Some(PublicError::RateLimit))] + #[case::client_error(Some(400), None)] + #[case::no_status(None, None)] + fn a_wrapped_429_is_a_rate_limit_only_behind_a_server_error( + #[case] status: Option, + #[case] expected: Option, ) { - assert_eq!(mapped(&http(status_code, "rejected")), Some(expected)); - } - - #[rstest::rstest] - #[case::unmapped_status(409)] - #[case::gateway_timeout(504)] - fn statuses_without_a_rule_fall_through(#[case] status_code: u16) { - assert_eq!(mapped(&http(status_code, "rejected")), None); - } - - #[rstest::rstest] - #[case::stub_without_an_upstream_response( - OriginalException::Response { message: "got a 403".into() }, - status(StatusClass::BadRequest, Some(ResponseArg::Stub(HttpStub { status: 403, method: "POST", url: VERTEX_URL_WITH_SPACE, content: None }))) - )] - #[case::stub_for_a_synthesized_status( - OriginalException::Connection { message: "got a 403".into() }, - status(StatusClass::BadRequest, Some(ResponseArg::Stub(HttpStub { status: 403, method: "POST", url: VERTEX_URL_WITH_SPACE, content: None }))) - )] - fn the_rule_response_stays_when_there_is_no_real_upstream_response( - #[case] original: OriginalException, - #[case] kind: PublicKind, - ) { - assert_eq!(mapped(&original).map(|failure| failure.kind), Some(kind)); - } - - #[test] - fn a_synthesized_500_keeps_the_internal_server_stub() { - let original = OriginalException::Connection { - message: "refused".into(), - }; assert_eq!( - mapped(&original), - Some(with_debug(failure( - status( - StatusClass::InternalServer, - Some(ResponseArg::Stub(HttpStub { - status: 500, - method: "completion", - url: "https://github.com/BerriAI/litellm", - content: Some("refused".into()), - })) - ), - "Vertex_aiException InternalServerError - refused", - "vertex_ai" - ))) + classified(status, r#"{"error": {"code": "429"}}"#), + expected ); } #[rstest::rstest] #[case::project_before_payload_size( "Unable to find your project 400 Request payload size exceeds", - StatusClass::BadRequest, - "litellm.BadRequestError: vertex_aiException - Unable to find your project 400 Request payload size exceeds", - true + PublicError::BadRequest )] - #[case::payload_size_before_context_window( - "400 Request payload size exceeds; This model's maximum context length is 10", - StatusClass::ContextWindowExceeded, - "Vertex_aiException - 400 Request payload size exceeds; This model's maximum context length is 10", - false + #[case::context_window_before_unknown_error( + "This model's maximum context length is 10 None Unknown Error.", + PublicError::ContextWindowExceeded )] - #[case::api_key_before_forbidden( - "API key not valid. 403", - StatusClass::Authentication, - "Vertex_aiException - API key not valid. 403", - true + #[case::unknown_error_before_api_key( + "Content has no parts. API key not valid.", + PublicError::InternalServer )] - #[case::forbidden_before_blocked( - "403 The response was blocked.", - StatusClass::BadRequest, - "Vertex_aiException BadRequestError - 403 The response was blocked.", - true + #[case::api_key_before_blocked( + "API key not valid. The response was blocked.", + PublicError::Authentication )] #[case::blocked_before_quota( "The response was blocked. Resource exhausted", - StatusClass::ContentPolicyViolation, - "Vertex_aiException ContentPolicyViolationError - The response was blocked. Resource exhausted", - true + PublicError::ContentPolicyViolation )] #[case::quota_before_overloaded( "Resource exhausted The model is overloaded.", - StatusClass::RateLimit, - "litellm.RateLimitError: vertex_aiException - Resource exhausted The model is overloaded.", - true + PublicError::RateLimit )] - fn the_earlier_rule_wins_when_two_apply( - #[case] body: &str, - #[case] class: StatusClass, - #[case] message: &str, - #[case] debug: bool, - ) { - let expected = failure(kept(class, body), message, "vertex_ai"); - assert_eq!( - mapped(&http(401, body)), - Some(if debug { - with_debug(expected) - } else { - expected - }) - ); + fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) { + assert_eq!(classified(Some(400), text), Some(expected)); + } + + #[rstest::rstest] + #[case::a_403_in_the_text("got a 403 from 4031 tokens")] + #[case::python_client_crash("IndexError: list index out of range")] + #[case::unmarked("rejected")] + fn text_without_a_marker_is_left_to_the_status_table(#[case] text: &str) { + assert_eq!(classified(Some(400), text), None); } } diff --git a/litellm-rust/crates/core-utils/src/lib.rs b/litellm-rust/crates/core-utils/src/lib.rs index 0c26aa50cc3..fcb232d8980 100644 --- a/litellm-rust/crates/core-utils/src/lib.rs +++ b/litellm-rust/crates/core-utils/src/lib.rs @@ -4,7 +4,6 @@ pub mod exception_mapping_utils; pub mod get_llm_provider_logic; pub mod params; pub mod prompt_templates; -pub mod python_repr; pub mod secret_redaction; pub mod serde_compat; pub mod url_utils; diff --git a/litellm-rust/crates/core-utils/src/python_repr.rs b/litellm-rust/crates/core-utils/src/python_repr.rs deleted file mode 100644 index 7ccfb826377..00000000000 --- a/litellm-rust/crates/core-utils/src/python_repr.rs +++ /dev/null @@ -1,93 +0,0 @@ -/// `repr()` of a Python `str`: single quotes unless the text holds a single quote and no -/// double quote, with backslashes, the chosen quote and control characters escaped. -pub fn python_str_repr(value: &str) -> String { - let quote = if value.contains('\'') && !value.contains('"') { - '"' - } else { - '\'' - }; - let escaped: String = value - .chars() - .map(|character| match character { - '\\' => "\\\\".to_string(), - '\t' => "\\t".to_string(), - '\n' => "\\n".to_string(), - '\r' => "\\r".to_string(), - character if character == quote => format!("\\{character}"), - character - if (character as u32) < 0x20 || (0x7f..0xa0).contains(&(character as u32)) => - { - format!("\\x{:02x}", character as u32) - } - character => character.to_string(), - }) - .collect(); - format!("{quote}{escaped}{quote}") -} - -/// `repr()` of the Python value a JSON value decodes to. -pub fn python_value_repr(value: &serde_json::Value) -> String { - use serde_json::Value; - match value { - Value::Null => "None".to_string(), - Value::Bool(true) => "True".to_string(), - Value::Bool(false) => "False".to_string(), - Value::Number(number) => number.to_string(), - Value::String(text) => python_str_repr(text), - Value::Array(items) => format!( - "[{}]", - items - .iter() - .map(python_value_repr) - .collect::>() - .join(", ") - ), - Value::Object(fields) => format!( - "{{{}}}", - fields - .iter() - .map(|(key, value)| format!( - "{}: {}", - python_str_repr(key), - python_value_repr(value) - )) - .collect::>() - .join(", ") - ), - } -} - -#[cfg(test)] -mod tests { - use serde_json::json; - - use super::{python_str_repr, python_value_repr}; - - #[rstest::rstest] - #[case::null(json!(null), "None")] - #[case::true_(json!(true), "True")] - #[case::false_(json!(false), "False")] - #[case::integer(json!(5), "5")] - #[case::float(json!(1.5), "1.5")] - #[case::string(json!("it's"), "\"it's\"")] - #[case::list(json!(["a", 1]), "['a', 1]")] - #[case::dict(json!({"format": "native"}), "{'format': 'native'}")] - #[case::empty_list(json!([]), "[]")] - fn value_repr_matches_python(#[case] value: serde_json::Value, #[case] expected: &str) { - assert_eq!(python_value_repr(&value), expected); - } - - #[rstest::rstest] - #[case::plain("native", "'native'")] - #[case::single_quote("it's", "\"it's\"")] - #[case::both_quotes("it's \"x\"", "'it\\'s \"x\"'")] - #[case::double_quote("say \"x\"", "'say \"x\"'")] - #[case::backslash("a\\b", "'a\\\\b'")] - #[case::whitespace("a\tb\nc\rd", "'a\\tb\\nc\\rd'")] - #[case::control("a\u{1}b\u{7f}c\u{85}", "'a\\x01b\\x7fc\\x85'")] - #[case::unicode("café", "'café'")] - #[case::empty("", "''")] - fn matches_python_repr(#[case] value: &str, #[case] expected: &str) { - assert_eq!(python_str_repr(value), expected); - } -} diff --git a/litellm-rust/crates/core-utils/src/secret_redaction.rs b/litellm-rust/crates/core-utils/src/secret_redaction.rs index 32922d7430c..e3caee4799a 100644 --- a/litellm-rust/crates/core-utils/src/secret_redaction.rs +++ b/litellm-rust/crates/core-utils/src/secret_redaction.rs @@ -1,5 +1,3 @@ -use std::sync::LazyLock; - use fancy_regex::Regex; pub const REDACTED: &str = "REDACTED"; @@ -51,21 +49,32 @@ fn secret_patterns(minimum_custom_key_length: usize) -> String { .join("|") } -static SECRET_RE: LazyLock = LazyLock::new(|| { - Regex::new(&format!( - "(?i){}", - secret_patterns(minimum_custom_key_length()) - )) - .expect("secret redaction patterns compile") -}); - -pub fn redact_string(value: &str) -> String { - SECRET_RE.replace_all(value, REDACTED).into_owned() +/// Python's `_ENABLE_SECRET_REDACTION` pattern set, compiled once per configuration. +#[derive(Clone, Debug)] +pub struct SecretRedactor { + pattern: Regex, } -pub fn secret_redaction_enabled() -> bool { - !std::env::var("LITELLM_DISABLE_REDACT_SECRETS") - .is_ok_and(|value| value.eq_ignore_ascii_case("true")) +impl SecretRedactor { + pub fn new(minimum_custom_key_length: usize) -> Self { + let pattern = Regex::new(&format!( + "(?i){}", + secret_patterns(minimum_custom_key_length) + )) + .expect("secret redaction patterns compile"); + Self { pattern } + } + + /// `None` when `LITELLM_DISABLE_REDACT_SECRETS` turns redaction off. + pub fn from_env() -> Option { + let disabled = std::env::var("LITELLM_DISABLE_REDACT_SECRETS") + .is_ok_and(|value| value.eq_ignore_ascii_case("true")); + (!disabled).then(|| Self::new(minimum_custom_key_length())) + } + + pub fn redact(&self, value: &str) -> String { + self.pattern.replace_all(value, REDACTED).into_owned() + } } #[cfg(test)] @@ -85,13 +94,16 @@ mod tests { #[case::password_needs_word_boundary("db_password=hunter2", "REDACTED")] #[case::plain_text_is_kept(r#"{"message": "rejected"}"#, r#"{"message": "rejected"}"#)] fn redacts_the_same_spans_as_the_python_patterns(#[case] input: &str, #[case] expected: &str) { - assert_eq!(redact_string(input), expected); + assert_eq!( + SecretRedactor::new(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH).redact(input), + expected + ); } #[test] fn sk_threshold_follows_the_minimum_custom_key_length() { - let patterns = Regex::new(&format!("(?i){}", secret_patterns(8))).unwrap(); - assert_eq!(patterns.replace_all("sk-abcde", REDACTED), REDACTED); - assert_eq!(patterns.replace_all("sk-abcd", REDACTED), "sk-abcd"); + let redactor = SecretRedactor::new(8); + assert_eq!(redactor.redact("sk-abcde"), REDACTED); + assert_eq!(redactor.redact("sk-abcd"), "sk-abcd"); } } diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json deleted file mode 100644 index ae629114589..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/api.json +++ /dev/null @@ -1,13 +0,0 @@ -{ - "kind": { - "type": "api", - "status": 409, - "request_url": "https://docs.litellm.ai/docs" - }, - "message": "MistralException - api", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json deleted file mode 100644 index ab400ed9e01..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/api_connection.json +++ /dev/null @@ -1,11 +0,0 @@ -{ - "kind": { - "type": "api_connection" - }, - "message": "MistralException - api_connection", - "model": "ocr-model", - "llm_provider": null, - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json deleted file mode 100644 index 057392575a1..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_authentication.json +++ /dev/null @@ -1,13 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "authentication", - "response": null - }, - "message": "MistralException - status_authentication", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json deleted file mode 100644 index abb1425f686..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_gateway.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "bad_gateway", - "response": { - "type": "upstream", - "status": 502, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_bad_gateway", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json deleted file mode 100644 index 171a994cd35..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_bad_request.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "bad_request", - "response": { - "type": "upstream", - "status": 400, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_bad_request", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json deleted file mode 100644 index ae9c1e145de..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_content_policy_violation.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "content_policy_violation", - "response": { - "type": "upstream", - "status": 400, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_content_policy_violation", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json deleted file mode 100644 index 61e1a56a622..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_context_window_exceeded.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "context_window_exceeded", - "response": { - "type": "upstream", - "status": 400, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_context_window_exceeded", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json deleted file mode 100644 index b3c5c51a785..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_internal_server.json +++ /dev/null @@ -1,19 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "internal_server", - "response": { - "type": "stub", - "status": 500, - "method": "completion", - "url": "https://github.com/BerriAI/litellm", - "content": "upstream text" - } - }, - "message": "MistralException - status_internal_server", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json deleted file mode 100644 index 26ffd872961..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_not_found.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "not_found", - "response": { - "type": "upstream", - "status": 404, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_not_found", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json deleted file mode 100644 index 42772f98830..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_permission_denied.json +++ /dev/null @@ -1,19 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "permission_denied", - "response": { - "type": "stub", - "status": 403, - "method": "POST", - "url": " https://cloud.google.com/vertex-ai/", - "content": null - } - }, - "message": "MistralException - status_permission_denied", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json deleted file mode 100644 index c9b88822ebd..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_rate_limit.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "rate_limit", - "response": { - "type": "upstream", - "status": 429, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_rate_limit", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json deleted file mode 100644 index 5af65202ef1..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_service_unavailable.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "service_unavailable", - "response": { - "type": "upstream", - "status": 503, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_service_unavailable", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json deleted file mode 100644 index a1b318ce03c..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/status_unsupported_params.json +++ /dev/null @@ -1,28 +0,0 @@ -{ - "kind": { - "type": "status", - "status_class": "unsupported_params", - "response": { - "type": "upstream", - "status": 400, - "body": "{\"message\": \"rejected\"}", - "headers": [ - [ - "retry-after", - "7" - ] - ] - } - }, - "message": "MistralException - status_unsupported_params", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": [ - [ - "retry-after", - "7" - ] - ], - "print_banner": false -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json deleted file mode 100644 index 8b215bdb7e6..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_with_status.json +++ /dev/null @@ -1,12 +0,0 @@ -{ - "kind": { - "type": "timeout", - "status": 504 - }, - "message": "MistralException - timeout_with_status", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": "\nModel: ocr-model", - "litellm_response_headers": null, - "print_banner": true -} diff --git a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json b/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json deleted file mode 100644 index 4a79169776e..00000000000 --- a/tests/test_litellm/rust_bridge/fixtures/public_failures/timeout_without_status.json +++ /dev/null @@ -1,12 +0,0 @@ -{ - "kind": { - "type": "timeout", - "status": null - }, - "message": "MistralException - timeout_without_status", - "model": "ocr-model", - "llm_provider": "mistral", - "litellm_debug_info": null, - "litellm_response_headers": null, - "print_banner": false -} From 6449632c7bc65015431abc08920737248b64b544 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:01:05 -0700 Subject: [PATCH 210/253] fix(proxy): persist only the router settings keys the request set --- litellm/proxy/proxy_server.py | 4 +++- .../proxy/proxy_server/test_routes_config.py | 23 +++++++++++++++++++ .../test_router_retry_policy_update.py | 2 +- 3 files changed, 27 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f43042942de..c7f91ad004e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -16982,7 +16982,9 @@ async def update_config( } ) typed_router_settings: Final[Mapping[str, JsonValue]] = ( - config_info.router_settings.model_dump(exclude_none=True) if config_info.router_settings is not None else {} + config_info.router_settings.model_dump(exclude_none=True, exclude_unset=True) + if config_info.router_settings is not None + else {} ) router_settings_updates: Final[Mapping[str, JsonValue]] = { **typed_router_settings, diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index e4d80c86f57..4234cdad23d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -197,6 +197,29 @@ def test_config_update_persists_only_the_general_settings_keys_the_request_set( assert persisted == {"alerting_threshold": 600} +def test_config_update_persists_only_the_router_settings_keys_the_request_set( + client, auth_as, mock_prisma, monkeypatch +): + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock()) + ps.proxy_config.router_settings.load_yaml({"model_group_alias": {"opus": "claude-opus-5"}}) + try: + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/update", json={"router_settings": {"retry_policy": {"TimeoutErrorRetries": 3}}} + ) + finally: + ps.proxy_config.router_settings.load_yaml({}) + + assert response.status_code == 200, response.text + persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"]) + assert persisted == {"retry_policy": {"TimeoutErrorRetries": 3}} + + def test_config_update_accepts_a_config_owned_success_callback_the_file_spells_in_mixed_case( client, auth_as, mock_prisma, monkeypatch ): diff --git a/tests/test_litellm/test_router_retry_policy_update.py b/tests/test_litellm/test_router_retry_policy_update.py index be568134763..0a3dcba325a 100644 --- a/tests/test_litellm/test_router_retry_policy_update.py +++ b/tests/test_litellm/test_router_retry_policy_update.py @@ -360,7 +360,7 @@ async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch): async def _apply_router_settings(*args, **kwargs): await proxy_server.proxy_config._add_router_settings_from_db_config( - config_data={}, llm_router=router, prisma_client=prisma_client + llm_router=router, prisma_client=prisma_client ) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) From 9d134413c9b6ad28b693d844a99ac1f167599f8d Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 14:00:47 -0700 Subject: [PATCH 211/253] test(rust): rename the standalone 429 test to match the rule --- .../crates/core-utils/src/exception_mapping_utils/openai.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs index b4f4496dbf0..d45078c415e 100644 --- a/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs +++ b/litellm-rust/crates/core-utils/src/exception_mapping_utils/openai.rs @@ -183,7 +183,7 @@ mod tests { } #[test] - fn a_standalone_429_counts_only_with_a_429_status() { + fn a_standalone_429_counts_with_a_429_status() { assert_eq!( classified(Some(429), "got 429 back"), Some(PublicError::RateLimit) From 3d3a46fc1fcd66549958807360ad4b182117a673 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:08:11 -0700 Subject: [PATCH 212/253] refactor: build the Converse beta list once and keep the OpenAPI snapshot as CI generates it --- .../adapters/transformation.py | 18 +++++++++------- .../bedrock/chat/converse_transformation.py | 21 ++++++++++--------- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- 3 files changed, 22 insertions(+), 19 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index d1047c8d86c..b213b5e30ba 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -198,6 +198,15 @@ def target_supports_mid_conversation_system(model: str | None, custom_llm_provid return supports_mid_conversation_system(model=model, custom_llm_provider=custom_llm_provider) +def _chat_tool_param(function_chunk: ChatCompletionToolParamFunctionChunk, tool: object) -> ChatCompletionToolParam: + eager_input_streaming: Final = eager_input_streaming_flag(tool) + if eager_input_streaming is None: + return ChatCompletionToolParam(type="function", function=function_chunk) + return ChatCompletionToolParam( + type="function", function=function_chunk, eager_input_streaming=eager_input_streaming + ) + + class AnthropicAdapter: def __init__(self) -> None: pass @@ -810,14 +819,7 @@ class LiteLLMAnthropicMessagesAdapter: for k, v in tool.items(): if k not in mapped_tool_params: # pass additional computer kwargs function_chunk.setdefault("parameters", {}).update({k: v}) - eager_input_streaming = eager_input_streaming_flag(tool) - tool_param = ( - ChatCompletionToolParam(type="function", function=function_chunk) - if eager_input_streaming is None - else ChatCompletionToolParam( - type="function", function=function_chunk, eager_input_streaming=eager_input_streaming - ) - ) + tool_param = _chat_tool_param(function_chunk, tool) self._add_cache_control_if_applicable(tool, tool_param, model) new_tools.append(tool_param) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 87e28054817..1a32fec45e3 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1518,9 +1518,6 @@ class AmazonConverseConfig(BaseConfig): """Process tools and collect anthropic_beta values.""" bedrock_tools: list[ToolBlock] = [] - # Collect anthropic_beta values from user headers - anthropic_beta_list: Final = list(get_anthropic_beta_from_headers(headers or {})) - # Separate pre-formatted Bedrock tools (e.g. systemTool from web_search_options) # from OpenAI-format tools that need transformation via _bedrock_tools_pt filtered_tools: Final = [] @@ -1540,6 +1537,17 @@ class AmazonConverseConfig(BaseConfig): continue filtered_tools.append(tool) + base_model: Final = BedrockModelInfo.get_base_model(model) + client_beta_list: Final = get_anthropic_beta_from_headers(headers or {}) + eager_beta: Final = ( + (ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER,) + if base_model.startswith("anthropic") + and AnthropicModelInfo().is_eager_input_streaming_used(filtered_tools) + and ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER not in client_beta_list + else () + ) + anthropic_beta_list: Final = [*client_beta_list, *eager_beta] + # Only separate tools if computer use tools are actually present if filtered_tools and self.is_computer_use_tool_used(filtered_tools, model): # Separate computer use tools from regular function tools @@ -1617,7 +1625,6 @@ class AmazonConverseConfig(BaseConfig): # Opus 4.5 gates ``output_config.effort`` behind a beta header; # Claude 4.6/4.7 accept it without one. - base_model: Final = BedrockModelInfo.get_base_model(model) if base_model.startswith("anthropic"): output_config: Final = additional_request_params.get("output_config") if ( @@ -1632,12 +1639,6 @@ class AmazonConverseConfig(BaseConfig): if ANTHROPIC_EFFORT_BETA_HEADER not in anthropic_beta_list: anthropic_beta_list.append(ANTHROPIC_EFFORT_BETA_HEADER) - if ( - AnthropicModelInfo().is_eager_input_streaming_used(filtered_tools) - and ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER not in anthropic_beta_list - ): - anthropic_beta_list.append(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER) - # Bedrock Converse: compact_20260112 edits only (+ beta header). AmazonConverseConfig._filter_context_management_for_bedrock_converse( additional_request_params, anthropic_beta_list diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 64daa695316..7a4c966a628 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19370,7 +19370,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " }, "500": { "content": { From 93d61abfa58e52d7f4b8154ce56420522be7eb48 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 21:08:14 +0000 Subject: [PATCH 213/253] fix(deepgram): refuse /listen sessions that have no streaming price A caller could pick a model with only a pre-recorded registry row, or no row at all, and the session would be billed at the pre-recorded rate or logged at zero cost, so budgets did not apply. The route now closes the WebSocket with 1008 before dialing Deepgram unless deepgram/streaming/ (or the -multilingual row for language=multi) is an exact registry hit, and the logging handler applies the same check so a registry change under a live session records the duration with no cost instead of a substitute rate Regression tests cover the route refusal, an operator-supplied streaming row for another model being accepted, and the handler never substituting the pre-recorded rate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 39 ++++++++----- .../llm_passthrough_endpoints.py | 21 +++++-- ...gram_listen_passthrough_logging_handler.py | 19 +++--- .../deepgram/test_deepgram_common_utils.py | 58 ++++++++++++++----- ...gram_listen_passthrough_logging_handler.py | 34 +++++------ .../test_deepgram_ws_passthrough_routes.py | 57 ++++++++++++++++++ 6 files changed, 165 insertions(+), 63 deletions(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index ede5b4f157d..676391dc744 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -6,8 +6,10 @@ from urllib.parse import parse_qs, urlparse import httpx +import litellm from litellm.constants import DEEPGRAM_DEFAULT_API_BASE, DEEPGRAM_LISTEN_DEFAULT_MODEL from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.types.utils import LlmProviders _WEBSOCKET_SCHEMES: Final = MappingProxyType({"https": "wss", "http": "ws", "wss": "wss", "ws": "ws"}) DEEPGRAM_LISTEN_CALLBACK_PARAMS: Final = frozenset({"callback", "callback_method"}) @@ -57,19 +59,30 @@ def _param_enabled(values: Sequence[str]) -> bool: return any(value.strip().lower() not in _DISABLED_PARAM_VALUES for value in values) -def deepgram_listen_base_pricing_models(upstream_url: str) -> tuple[str, ...]: - """Registry keys to try, in order, for the per-second base rate of a streaming session: the streaming entry for - the language mode Deepgram bills (multilingual when ``language=multi``), then the plain streaming entry, then - the pre-recorded entry for models that have no streaming price of their own.""" - model: Final = deepgram_listen_model(upstream_url) - params: Final = parse_qs(urlparse(upstream_url).query) - streaming: Final = f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{model}" - multilingual: Final = params.get("language", ("",))[-1].strip().lower() == DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE - return ( - (f"{streaming}{DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX}", streaming, model) - if multilingual - else (streaming, model) - ) +def deepgram_listen_pricing_model(upstream_url: str) -> str: + """Registry key, without the provider prefix, for the per-second base rate Deepgram bills a streaming session at: + the multilingual streaming entry when ``language=multi``, otherwise the model's own streaming entry. Pre-recorded + entries are never a substitute: Deepgram prices the two products differently.""" + streaming: Final = f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{deepgram_listen_model(upstream_url)}" + language: Final = parse_qs(urlparse(upstream_url).query).get("language", ("",))[-1] + if language.strip().lower() == DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE: + return f"{streaming}{DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX}" + return streaming + + +def deepgram_listen_registry_key(upstream_url: str) -> str: + return f"{LlmProviders.DEEPGRAM.value}/{deepgram_listen_pricing_model(upstream_url)}" + + +def deepgram_listen_is_priced(upstream_url: str) -> bool: + """Only an exact registry hit counts: the cost calculator resolves a missing ``streaming/`` row to the + pre-recorded ```` row, which is not the rate Deepgram bills a WebSocket session at.""" + registry_key: Final = deepgram_listen_registry_key(upstream_url) + try: + model_info: Final = litellm.get_model_info(model=registry_key, custom_llm_provider=LlmProviders.DEEPGRAM.value) + except Exception: + return False + return model_info["key"] == registry_key def deepgram_listen_addon_pricing_models(upstream_url: str) -> tuple[str, ...]: diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index ea90648e526..629784b3cee 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -38,6 +38,8 @@ from litellm.llms.azure.passthrough.transformation import foreign_azure_deployme from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.deepgram.common_utils import ( deepgram_listen_callback_params, + deepgram_listen_is_priced, + deepgram_listen_registry_key, deepgram_listen_requested_model, deepgram_listen_websocket_target, ) @@ -2897,6 +2899,9 @@ _DEEPGRAM_WS_MISSING_KEY_REASON: Final = ( "Required 'DEEPGRAM_API_KEY' in environment to make pass-through calls to Deepgram." ) _DEEPGRAM_WS_CALLBACK_REASON: Final = "Deepgram callback delivery is not supported through the proxy: remove {params}" +_DEEPGRAM_WS_UNPRICED_REASON: Final = ( + "No streaming price for '{registry_key}': add it to the model cost map to enable it" +) async def deepgram_listen_user_api_key_auth(websocket: WebSocket) -> UserAPIKeyAuth: @@ -2929,12 +2934,20 @@ async def deepgram_listen_websocket_route( ) return + target: Final = deepgram_listen_websocket_target( + api_base=get_secret_str("DEEPGRAM_API_BASE"), + query_string=websocket.url.query, + ) + if not deepgram_listen_is_priced(target): + await websocket.close( + code=1008, + reason=_DEEPGRAM_WS_UNPRICED_REASON.format(registry_key=deepgram_listen_registry_key(target)), + ) + return + await relay( websocket=websocket, - target=deepgram_listen_websocket_target( - api_base=get_secret_str("DEEPGRAM_API_BASE"), - query_string=websocket.url.query, - ), + target=target, custom_headers={ # mutable-ok: websocket_passthrough_request requires a plain dict of upstream headers "Authorization": f"Token {deepgram_api_key}" }, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py index c0953ae7b85..8386c154600 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/deepgram_listen_passthrough_logging_handler.py @@ -9,9 +9,11 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.deepgram.common_utils import ( deepgram_listen_addon_pricing_models, deepgram_listen_audio_seconds, - deepgram_listen_base_pricing_models, deepgram_listen_channel_count, + deepgram_listen_is_priced, deepgram_listen_model, + deepgram_listen_pricing_model, + deepgram_listen_registry_key, deepgram_listen_transcript, ) from litellm.proxy._types import PassThroughEndpointLoggingTypedDict @@ -34,19 +36,14 @@ def _registry_cost(response: TranscriptionResponse, pricing_model: str) -> float def _audio_cost(response: TranscriptionResponse, upstream_url: str) -> float | None: - base_cost: Final = next( - ( - cost - for pricing_model in deepgram_listen_base_pricing_models(upstream_url) - if (cost := _registry_cost(response, pricing_model)) is not None - ), - None, - ) - if base_cost is None: + if not deepgram_listen_is_priced(upstream_url): verbose_proxy_logger.warning( - "Deepgram listen passthrough: no pricing for model '%s'", deepgram_listen_model(upstream_url) + "Deepgram listen passthrough: no registry entry '%s'", deepgram_listen_registry_key(upstream_url) ) return None + base_cost: Final = _registry_cost(response, deepgram_listen_pricing_model(upstream_url)) + if base_cost is None: + return None addon_costs: Final = tuple( _registry_cost(response, pricing_model) for pricing_model in deepgram_listen_addon_pricing_models(upstream_url) ) diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index 8b61d5bffa0..a1fa8f26b70 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -8,10 +8,12 @@ import litellm from litellm.llms.deepgram.common_utils import ( deepgram_listen_addon_pricing_models, deepgram_listen_audio_seconds, - deepgram_listen_base_pricing_models, deepgram_listen_callback_params, deepgram_listen_channel_count, + deepgram_listen_is_priced, deepgram_listen_model, + deepgram_listen_pricing_model, + deepgram_listen_registry_key, deepgram_listen_requested_model, deepgram_listen_transcript, deepgram_listen_websocket_target, @@ -220,27 +222,51 @@ def test_requested_model_is_the_model_the_upstream_target_will_carry(query_strin @pytest.mark.parametrize( ("upstream_url", "expected"), [ - pytest.param(NOVA_3_URL, ("streaming/nova-3", "nova-3"), id="monolingual"), - pytest.param(f"{NOVA_3_URL}&language=en", ("streaming/nova-3", "nova-3"), id="explicit language"), - pytest.param( - f"{NOVA_3_URL}&language=multi", - ("streaming/nova-3-multilingual", "streaming/nova-3", "nova-3"), - id="multilingual", - ), - pytest.param( - f"{NOVA_3_URL}&language=MULTI", - ("streaming/nova-3-multilingual", "streaming/nova-3", "nova-3"), - id="multilingual any case", - ), + pytest.param(NOVA_3_URL, "streaming/nova-3", id="monolingual"), + pytest.param(f"{NOVA_3_URL}&language=en", "streaming/nova-3", id="explicit language"), + pytest.param(f"{NOVA_3_URL}&language=multi", "streaming/nova-3-multilingual", id="multilingual"), + pytest.param(f"{NOVA_3_URL}&language=MULTI", "streaming/nova-3-multilingual", id="multilingual any case"), pytest.param( "wss://api.deepgram.com/v1/listen?model=nova-2&language=multi", - ("streaming/nova-2-multilingual", "streaming/nova-2", "nova-2"), + "streaming/nova-2-multilingual", id="other model", ), + pytest.param("wss://api.deepgram.com/v1/listen?encoding=linear16", "streaming/nova-3", id="default model"), ], ) -def test_deepgram_listen_base_pricing_models(upstream_url: str, expected: tuple[str, ...]): - assert deepgram_listen_base_pricing_models(upstream_url) == expected +def test_deepgram_listen_pricing_model_is_the_streaming_entry_never_the_prerecorded_one( + upstream_url: str, expected: str +): + assert deepgram_listen_pricing_model(upstream_url) == expected + assert deepgram_listen_registry_key(upstream_url) == f"deepgram/{expected}" + + +NOVA_2_URL: Final = "wss://api.deepgram.com/v1/listen?model=nova-2" + + +@pytest.mark.usefixtures("local_model_cost_map") +@pytest.mark.parametrize( + ("upstream_url", "extra_rows", "expected"), + [ + pytest.param(NOVA_3_URL, (), True, id="streaming entry present"), + pytest.param(f"{NOVA_3_URL}&language=multi", (), True, id="multilingual entry present"), + pytest.param(NOVA_2_URL, (), False, id="only the pre-recorded entry"), + pytest.param(f"{NOVA_2_URL}&language=multi", ("deepgram/streaming/nova-2",), False, id="needs multilingual"), + pytest.param("wss://api.deepgram.com/v1/listen?model=nova-99-unmapped", (), False, id="nothing priced"), + pytest.param(NOVA_2_URL, ("deepgram/streaming/nova-2",), True, id="operator-supplied streaming entry"), + pytest.param(NOVA_2_URL, ("streaming/nova-2",), False, id="a row under another key is not the entry"), + ], +) +def test_deepgram_listen_is_priced( + monkeypatch: pytest.MonkeyPatch, upstream_url: str, extra_rows: tuple[str, ...], expected: bool +): + """The bundled map prices only nova-3 for streaming; nova-2 has a pre-recorded row, which must never count.""" + monkeypatch.delitem(litellm.model_cost, "deepgram/streaming/nova-2", raising=False) + assert "deepgram/nova-2" in litellm.model_cost + for row in extra_rows: + monkeypatch.setitem(litellm.model_cost, row, dict(litellm.model_cost["deepgram/streaming/nova-3"])) + + assert deepgram_listen_is_priced(upstream_url) is expected @pytest.mark.parametrize( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py index 2ae56709b11..ac742c0ab46 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_deepgram_listen_passthrough_logging_handler.py @@ -158,17 +158,25 @@ def test_handler_add_ons_scale_with_channels_like_the_base_rate(): assert stereo_redacted - stereo_plain == pytest.approx(_registry_cost("streaming/redact", 120.0)) -def test_handler_falls_back_to_the_prerecorded_rate_for_a_model_without_a_streaming_entry(): - assert "deepgram/streaming/nova-2" not in litellm.model_cost +@pytest.mark.parametrize( + "upstream_url", + [ + pytest.param("wss://api.deepgram.com/v1/listen?model=nova-2", id="only a pre-recorded entry"), + pytest.param("wss://api.deepgram.com/v1/listen?model=nova-99-not-in-registry", id="no entry at all"), + ], +) +def test_handler_never_substitutes_another_rate_for_a_missing_streaming_entry(monkeypatch, upstream_url): + """The route refuses these sessions up front; should the registry change under a live one, the spend row + keeps the duration and carries no cost, rather than the pre-recorded rate or any other stand-in.""" + monkeypatch.delitem(litellm.model_cost, "deepgram/streaming/nova-2", raising=False) + assert "deepgram/nova-2" in litellm.model_cost handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( - websocket_messages=(_metadata(60.0),), - logging_obj=_logging_obj(), - upstream_url="wss://api.deepgram.com/v1/listen?model=nova-2", + websocket_messages=(_metadata(60.0),), logging_obj=_logging_obj(), upstream_url=upstream_url ) - assert handler_result["kwargs"]["model"] == "nova-2" - assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("nova-2", 60.0)) + assert handler_result["kwargs"]["response_cost"] is None + assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 60.0 def test_handler_falls_back_to_results_frames_when_the_stream_ends_without_metadata(): @@ -222,18 +230,6 @@ def test_handler_bills_the_declared_channels_when_the_stream_dies_before_any_fra assert handler_result["kwargs"]["response_cost"] == pytest.approx(_registry_cost("streaming/nova-3", 30.0)) -def test_handler_keeps_the_spend_row_but_no_cost_for_an_unpriced_model(): - handler_result = DeepgramListenPassthroughLoggingHandler().deepgram_listen_passthrough_handler( - websocket_messages=(_metadata(12.5),), - logging_obj=_logging_obj(), - upstream_url="wss://api.deepgram.com/v1/listen?model=nova-99-not-in-registry", - ) - - assert handler_result["kwargs"]["model"] == "nova-99-not-in-registry" - assert handler_result["kwargs"]["response_cost"] is None - assert handler_result["result"]._hidden_params["audio_transcription_duration"] == 12.5 - - class _CapturingLogger(CustomLogger): def __init__(self) -> None: super().__init__() diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 85baaacf1e9..4eb183b14ce 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -30,6 +30,14 @@ GET_CREDENTIALS: Final = ( ) USER_API_KEY_AUTH: Final = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth" LISTEN_PATHS: Final = ("/deepgram/v1/listen", "/deepgram/listen") +NOVA_2_STREAMING_KEY: Final = "deepgram/streaming/nova-2" + +pytestmark: Final = pytest.mark.usefixtures("local_model_cost_map") + + +def _price_nova_2_streaming(monkeypatch: pytest.MonkeyPatch) -> None: + """An operator-supplied streaming row: the bundled map prices only nova-3 for streaming.""" + monkeypatch.setitem(litellm.model_cost, NOVA_2_STREAMING_KEY, dict(litellm.model_cost["deepgram/streaming/nova-3"])) class _FakeWebSocket: @@ -138,6 +146,7 @@ async def test_deepgram_listen_forwards_query_and_injects_only_provider_auth(pat @pytest.mark.asyncio async def test_deepgram_listen_keeps_caller_chosen_model(monkeypatch): monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + _price_nova_2_streaming(monkeypatch) websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2&language=en") with patch(GET_CREDENTIALS, return_value="dg-provider-key"): @@ -236,6 +245,53 @@ async def test_deepgram_listen_rejects_callback_delivery_that_would_go_unbilled( assert "dg-provider-key" not in websocket.closed[1] +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("query", "missing_key"), + [ + pytest.param("model=nova-2", "deepgram/streaming/nova-2", id="model with only a pre-recorded price"), + pytest.param("model=nova-99", "deepgram/streaming/nova-99", id="model unknown to the registry"), + pytest.param( + "model=nova-3&language=multi", + "deepgram/streaming/nova-3-multilingual", + id="multilingual session without its own price", + ), + ], +) +async def test_deepgram_listen_refuses_sessions_it_cannot_price(query, missing_key, monkeypatch): + """A session with no streaming price would be logged at zero (or at the pre-recorded rate), letting a caller run + up unmetered spend, so the proxy closes it before Deepgram is contacted and names the registry row to add.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + monkeypatch.delitem(litellm.model_cost, missing_key, raising=False) + assert "deepgram/nova-2" in litellm.model_cost + websocket = _FakeWebSocket("/deepgram/v1/listen", query) + + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(websocket) + + assert relay.calls == [] + assert websocket.closed is not None + assert websocket.closed[0] == 1008 + assert missing_key in websocket.closed[1] + assert "dg-provider-key" not in websocket.closed[1] + + +@pytest.mark.asyncio +async def test_deepgram_listen_relays_once_the_operator_prices_the_model(monkeypatch): + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2") + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + assert (await _serve(websocket)).calls == [] + + _price_nova_2_streaming(monkeypatch) + priced_websocket = _FakeWebSocket("/deepgram/v1/listen", "model=nova-2") + with patch(GET_CREDENTIALS, return_value="dg-provider-key"): + relay = await _serve(priced_websocket) + + assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2"] + assert priced_websocket.closed is None + + def _app_with_relay(relay: _FakeRelay) -> FastAPI: app = FastAPI() app.include_router(router) @@ -332,6 +388,7 @@ def test_deepgram_listen_authorizes_the_model_it_will_actually_send_upstream(que in its default: the real key auth path must see the same model the upstream target will carry.""" monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) monkeypatch.setattr(litellm, "max_budget", 0.0) + _price_nova_2_streaming(monkeypatch) cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"])) relay = _FakeRelay() client = TestClient(_app_with_relay(relay)) From f1b9642c4108e5bb3c9fbae858d1f39099f815cb Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 21:15:01 +0000 Subject: [PATCH 214/253] fix(proxy): classify Azure Speech short audio behind a prefixed api base Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../azure_speech_passthrough_logging_handler.py | 4 +++- .../test_azure_speech_passthrough_logging_handler.py | 8 ++++++++ 2 files changed, 11 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py index 8cd1b137de8..33d1815b3c4 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/azure_speech_passthrough_logging_handler.py @@ -8,6 +8,7 @@ import httpx from litellm._logging import verbose_proxy_logger from litellm.constants import ( AZURE_SPEECH_BATCH_MODEL, + AZURE_SPEECH_BATCH_PATH_PREFIX, AZURE_SPEECH_CUSTOM_LLM_PROVIDER, AZURE_SPEECH_FAST_TRANSCRIPTION_MODEL, AZURE_SPEECH_FAST_TRANSCRIPTION_PATH, @@ -30,7 +31,8 @@ from litellm.types.utils import StandardPassThroughResponseObject class AzureSpeechPassthroughLoggingHandler: @staticmethod def _is_short_audio_route(url_route: str) -> bool: - return urlparse(url_route).path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX) + path: Final = urlparse(url_route).path + return path.rfind(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX) > path.rfind(AZURE_SPEECH_BATCH_PATH_PREFIX) @staticmethod def _is_fast_transcription_route(url_route: str) -> bool: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py index 670f8e65823..a0d27e618f9 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_azure_speech_passthrough_logging_handler.py @@ -21,6 +21,11 @@ SHORT_AUDIO_URL = ( ) BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions" FAST_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/transcriptions:transcribe?api-version=2024-11-15" +PREFIXED_SHORT_AUDIO_URL = ( + "https://apim.example.com/speech-proxy/speech/recognition/conversation/cognitiveservices/v1?language=en-US" +) +PREFIXED_BATCH_URL = "https://apim.example.com/speech/speechtotext/v3.2/transcriptions" +PREFIXED_FAST_URL = "https://apim.example.com/speech/speechtotext/transcriptions:transcribe?api-version=2024-11-15" FAST_BODY = {"durationMilliseconds": 5061, "combinedPhrases": [{"text": "Hello world."}]} FAST_AUDIO_SECONDS = 5.061 TRANSCRIPT_BODY = { @@ -87,6 +92,9 @@ class TestAzureSpeechPassthroughHandler: (FAST_URL, "azure_speech/fast-transcription", FAST_AUDIO_SECONDS * PRICE_PER_SECOND), (BATCH_URL, "azure_speech/batch-transcription", 0.0), (f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription", 0.0), + (PREFIXED_SHORT_AUDIO_URL, "azure_speech/short-audio", TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND), + (PREFIXED_FAST_URL, "azure_speech/fast-transcription", FAST_AUDIO_SECONDS * PRICE_PER_SECOND), + (PREFIXED_BATCH_URL, "azure_speech/batch-transcription", 0.0), ], ) def test_records_model_provider_and_cost(self, url_route: str, expected_model: str, expected_cost: float): From 3449ae9d0d1d339146fa5ffb7318a62553860a46 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 21:15:32 +0000 Subject: [PATCH 215/253] fix(proxy): advance the daily global spend marker in one conditional upsert so overlapping runs cannot rewind it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../daily_global_spend_rollup.py | 51 +++++++--------- .../test_daily_global_spend_rollup.py | 61 ++++++++++++++++--- 2 files changed, 74 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py index d50feb6f16b..376b113ed02 100644 --- a/litellm/proxy/spend_tracking/daily_global_spend_rollup.py +++ b/litellm/proxy/spend_tracking/daily_global_spend_rollup.py @@ -23,7 +23,6 @@ from litellm.constants import ( DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS, DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, ) -from litellm.repositories.config_repository import ConfigRepository if TYPE_CHECKING: from litellm.caching.redis_cache import RedisCache @@ -82,6 +81,17 @@ _PENDING_DAYS_SQL: Final = ( 'AND ("date" > $2 OR "updated_at" >= $3::timestamp - INTERVAL \'1 hour\') ' 'ORDER BY "date"' ) +# Runs can overlap (Redis unreachable, lock expired on a long backfill), so the database keeps the +# later of the stored and the incoming day and scan time in one statement; GREATEST skips NULL. +_ADVANCE_MARKER_SQL: Final = ( + 'INSERT INTO "LiteLLM_Config" ("param_name", "param_value") ' + "VALUES ($1, jsonb_build_object('reconciled_through', $2::text, 'scanned_at', $3::text)) " + 'ON CONFLICT ("param_name") DO UPDATE SET "param_value" = jsonb_build_object(' + "'reconciled_through', GREATEST(\"LiteLLM_Config\".\"param_value\" ->> 'reconciled_through', " + "EXCLUDED.\"param_value\" ->> 'reconciled_through'), " + "'scanned_at', GREATEST(\"LiteLLM_Config\".\"param_value\" ->> 'scanned_at', " + "EXCLUDED.\"param_value\" ->> 'scanned_at'))" +) class ReconciledThrough(BaseModel): @@ -152,33 +162,20 @@ async def reconciled_through(prisma_client: "PrismaClient") -> str | None: return None if marker is None else marker.reconciled_through -async def _record_marker(prisma_client: "PrismaClient", marker: ReconciledThrough) -> None: +async def _advance_marker(prisma_client: "PrismaClient", days: tuple[str, ...], *, scanned_at: str | None) -> None: + """Move the stored marker to the last of ``days`` and to ``scanned_at`` where those are later + than what is stored, so a slower overlapping run can only add to a faster run's marker.""" from litellm.proxy.utils import invalidate_config_param - await ConfigRepository(prisma_client).set_param( - DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, marker.model_dump_json() + await prisma_client.db.execute_raw( + _ADVANCE_MARKER_SQL, + DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM, + max(days) if days else None, + scanned_at, ) await invalidate_config_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) -async def _stored_marker(prisma_client: "PrismaClient") -> ReconciledThrough | None: - """The marker as another pod may have just written it, bypassing this pod's config cache.""" - param: Final = await ConfigRepository(prisma_client).get_param(DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM) - return None if param is None else _marker_from_param_value(param.param_value) - - -async def _record_advanced(prisma_client: "PrismaClient", days: tuple[str, ...], *, scanned_at: str | None) -> None: - """Advance the stored marker by ``days``. Two runs can overlap (Redis unreachable, lock expired - on a long backfill), so the base is what is stored now, not the snapshot this run scanned from: - a slower run may then only add to the faster run's marker, never rewind it. Without a new scan - time the stored one is kept.""" - stored: Final = await _stored_marker(prisma_client) - kept_scanned_at: Final = None if stored is None else stored.scanned_at - await _record_marker( - prisma_client, _advanced(stored, days, scanned_at=scanned_at if scanned_at is not None else kept_scanned_at) - ) - - async def _db_now(prisma_client: "PrismaClient") -> _NowRow: rows: Final = await prisma_client.db.query_raw(_DB_NOW_SQL) return _NowRow.model_validate(rows[0]) @@ -221,16 +218,10 @@ async def run_daily_global_spend_reconcile(prisma_client: "PrismaClient") -> Rec marker: Final = await reconciled_through(prisma_client) return ReconcileResult(days_reconciled=done, reconciled_through=marker, failed_day=scan.days[len(done)]) if scan.marker is not None or done: - await _record_advanced(prisma_client, done, scanned_at=scan.scanned_at) + await _advance_marker(prisma_client, done, scanned_at=scan.scanned_at) return ReconcileResult(days_reconciled=done, reconciled_through=await reconciled_through(prisma_client)) -def _advanced(marker: ReconciledThrough | None, days: tuple[str, ...], *, scanned_at: str | None) -> ReconciledThrough: - """The marker after ``days`` were rewritten: a late old day never moves it back.""" - through: Final = max((marker.reconciled_through if marker is not None else "", *days)) - return ReconciledThrough(reconciled_through=through, scanned_at=scanned_at) - - async def _reconcile_until_failure(prisma_client: "PrismaClient", scan: _PendingScan) -> tuple[str, ...]: for index, day in enumerate(scan.days): if not await _reconcile_and_record(prisma_client, scan.days[: index + 1]): @@ -242,7 +233,7 @@ async def _reconcile_and_record(prisma_client: "PrismaClient", done_with_this: t day: Final = done_with_this[-1] try: await reconcile_day(prisma_client, day) - await _record_advanced(prisma_client, done_with_this, scanned_at=None) + await _advance_marker(prisma_client, done_with_this, scanned_at=None) except Exception as exc: # noqa: BLE001 # one bad day must not lose the days already done verbose_proxy_logger.exception("Daily global spend reconcile: day %s failed: %s", day, exc) return False diff --git a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py index 69f06fad081..3da587435ad 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_daily_global_spend_rollup.py @@ -1,5 +1,6 @@ """Tests for the LiteLLM_DailyGlobalSpend reconcile job (LIT-7818).""" +import json import pathlib import re from datetime import date @@ -14,6 +15,7 @@ from pytest_postgresql import factories from litellm.constants import DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM from litellm.proxy.db.daily_spend_bulk_upsert import DAILY_SPEND_TABLES, build_bulk_upsert, merge_by_conflict_key from litellm.proxy.spend_tracking.daily_global_spend_rollup import ( + _ADVANCE_MARKER_SQL, RECONCILE_DAY_SQL, read_marker, reconciled_through, @@ -36,13 +38,21 @@ class _FakeConfigTable: def __init__(self) -> None: self.rows: dict[str, object] = {} - async def upsert(self, *, where: dict[str, str], data: dict[str, dict[str, str]]) -> _FakeConfigRow: - self.rows[where["param_name"]] = data["update"]["param_value"] - return _FakeConfigRow(where["param_name"], data["update"]["param_value"]) + def advance(self, param_name: str, through: str | None, scanned_at: str | None) -> None: + """What ``_ADVANCE_MARKER_SQL`` does in Postgres: keep the later of stored and incoming per field.""" + stored = self.rows.get(param_name) + current: dict[str, str | None] = json.loads(stored) if isinstance(stored, str) else {} + self.rows[param_name] = json.dumps( + { + "reconciled_through": _greatest(current.get("reconciled_through"), through), + "scanned_at": _greatest(current.get("scanned_at"), scanned_at), + } + ) - async def find_unique(self, *, where: dict[str, str]) -> _FakeConfigRow | None: - stored = self.rows.get(where["param_name"]) - return None if stored is None else _FakeConfigRow(where["param_name"], stored) + +def _greatest(stored: str | None, incoming: str | None) -> str | None: + present = [value for value in (stored, incoming) if value is not None] + return max(present) if present else None class _FakeDb: @@ -67,9 +77,14 @@ class _FakeDb: {"date": d} for d, written in sorted(rows.items()) if d <= last and (d > marker or written >= scanned_at) ] - async def execute_raw(self, sql: str, *params: str) -> int: + async def execute_raw(self, sql: str, *params: str | None) -> int: + if sql == _ADVANCE_MARKER_SQL: + param_name, through, scanned_at = params + assert param_name is not None + self.litellm_config.advance(param_name, through, scanned_at) + return 1 (day,) = params - if day in self._prisma.failing_days: + if day is None or day in self._prisma.failing_days: raise RuntimeError(f"day {day} exploded") self._prisma.reconciled.append(day) landing = self._prisma.marker_landing_on_day.get(day) @@ -485,3 +500,33 @@ def test_reconcile_day_sql_makes_the_global_day_equal_the_per_key_sums(_rollup_p assert sum(int(r["total_response_time_ms"]) for r in global_rows) == 1600 # pyright: ignore[reportArgumentType] # dict_row values are untyped assert [(r["model"], r["model_group"]) for r in global_rows] == [("gpt-5", ""), ("gpt-5", "gpt-5")] assert untouched == [] + + +_CONFIG_DDL: Final = 'CREATE TABLE "LiteLLM_Config" (param_name TEXT PRIMARY KEY, param_value JSONB)' +_MARKER_SQL: Final = 'SELECT param_value FROM "LiteLLM_Config" WHERE param_name = %s' + + +def test_advance_marker_sql_only_ever_moves_the_stored_marker_forward(_rollup_postgresql: psycopg.Connection): + """Against real Postgres: the statement a slower overlapping run issues after the faster run + already stored a later marker leaves that marker alone, whether it carries an older scan time or + none at all, while a run that is further along moves both fields on.""" + conn: Final = _rollup_postgresql + conn.execute(_CONFIG_DDL) # pyright: ignore[reportArgumentType] # DDL literal + conn.commit() + param: Final = DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM + + def stored() -> object: + with conn.cursor(row_factory=dict_row) as cur: + row = cur.execute(_MARKER_SQL, (param,)).fetchone() + return None if row is None else row["param_value"] + + _execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-01", None)) + assert stored() == {"reconciled_through": "2026-09-01", "scanned_at": None} + + _execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-14", "2026-09-15 00:30:02.5")) + _execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-02", None)) + _execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-03", "2026-09-15 00:30:01.25")) + assert stored() == {"reconciled_through": "2026-09-14", "scanned_at": "2026-09-15 00:30:02.5"} + + _execute_dollar_sql(conn, _ADVANCE_MARKER_SQL, (param, "2026-09-15", "2026-09-16 00:30:00.75")) + assert stored() == {"reconciled_through": "2026-09-15", "scanned_at": "2026-09-16 00:30:00.75"} From 4a951847bbcb25a8c42eb526b0526604d1b3bb1a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:15:39 -0700 Subject: [PATCH 216/253] fix(responses): merge deployment litellm_params into native websocket response.create frames --- litellm/llms/custom_httpx/llm_http_handler.py | 3 + litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/responses/main.py | 25 ++- litellm/responses/streaming_iterator.py | 31 ++- .../types/responses/streaming_websocket.py | 13 ++ .../test_responses_websocket_all_providers.py | 210 ++++++++++++++++++ 6 files changed, 275 insertions(+), 9 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 98fe0014386..477d10a3cbd 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -144,6 +144,7 @@ from litellm.types.llms.openai import ( from litellm.types.realtime import RealtimeQueryParams from litellm.types.rerank import RerankResponse from litellm.types.responses.main import DeleteResponseResult +from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CallTypes, @@ -6588,6 +6589,7 @@ class BaseLLMHTTPHandler: litellm_metadata: dict[str, object] | None = None, custom_llm_provider: str | None = None, first_message: str | None = None, + request_defaults: ResponsesWebSocketRequestDefaults | None = None, **kwargs: Any, ): """ @@ -6742,6 +6744,7 @@ class BaseLLMHTTPHandler: output_guardrail_callbacks=_ws_output_guardrail_callbacks, quota_callbacks=_ws_quota_callbacks, authorized_model=model, + request_defaults=request_defaults, ) await streaming.bidirectional_forward() diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 8a8d08c6887..2aa2cf15ac1 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19394,7 +19394,7 @@ } } }, - "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" }, "500": { "content": { diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 6dc34bb93ef..6064fa66d91 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -8,7 +8,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias, cast import httpx -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter from typing_extensions import assert_never import litellm @@ -53,6 +53,7 @@ from litellm.litellm_core_utils.get_litellm_params import get_litellm_params from litellm.llms.openai.data_residency import infer_openai_data_residency from litellm.secret_managers.main import get_secret_str from litellm.types.responses.main import * +from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import all_litellm_params from litellm.utils import ( @@ -2261,6 +2262,27 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict: return metadata +_EXTRA_BODY_ADAPTER: Final = TypeAdapter(dict[str, object] | None) + + +def _build_responses_websocket_request_defaults(kwargs: Mapping[str, object]) -> ResponsesWebSocketRequestDefaults: + reasoning_effort: Final = kwargs.get("reasoning_effort") + mapped_reasoning: Final = ( + LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) + if kwargs.get("reasoning") is None and isinstance(reasoning_effort, str) + else None + ) + candidate_params: Final[dict[str, object]] = { + **kwargs, + **({"reasoning": mapped_reasoning} if mapped_reasoning is not None else {}), + } + fill_missing: Final = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(candidate_params) + return ResponsesWebSocketRequestDefaults( + fill_missing=MappingProxyType(dict(fill_missing)), + overrides=MappingProxyType(_EXTRA_BODY_ADAPTER.validate_python(kwargs.get("extra_body")) or {}), + ) + + @client async def _aresponses_websocket( model: str, @@ -2352,5 +2374,6 @@ async def _aresponses_websocket( user_api_key_dict=kwargs.get("user_api_key_dict"), litellm_metadata=_build_litellm_metadata_for_ws(kwargs), custom_llm_provider=_custom_llm_provider, + request_defaults=_build_responses_websocket_request_defaults(kwargs), **remaining_kwargs, ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 8d766cf1cd0..1c82df8c664 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -54,6 +54,7 @@ if TYPE_CHECKING: PresidioGuardrailCallback, ResponsesBackendWebSocket, ResponsesClientWebSocket, + ResponsesWebSocketRequestDefaults, ) from litellm.types.router import LiteLLM_Params @@ -1717,6 +1718,7 @@ class ResponsesWebSocketStreaming: output_guardrail_callbacks: list[PresidioGuardrailCallback] | None = None, quota_callbacks: Sequence[ProjectQuotaCallback] | None = None, authorized_model: str | None = None, + request_defaults: ResponsesWebSocketRequestDefaults | None = None, ): self.websocket = websocket self.backend_ws = backend_ws @@ -1732,6 +1734,7 @@ class ResponsesWebSocketStreaming: # Model name authorized at connection time; enforced on every # response.create frame to prevent deployment-substitution attacks. self.authorized_model: str | None = authorized_model + self.request_defaults: ResponsesWebSocketRequestDefaults | None = request_defaults def _should_store_event(self, event_obj: _MutableJsonObject) -> bool: return event_obj.get("type") in RESPONSES_WS_LOGGED_EVENT_TYPES @@ -1874,12 +1877,23 @@ class ResponsesWebSocketStreaming: modified = True return modified + def _with_request_defaults(self, msg_obj: dict[str, object]) -> dict[str, object]: + if self.request_defaults is None: + return msg_obj + nested: Final = msg_obj.get("response") + if _is_json_object(nested): + return {**msg_obj, "response": self.request_defaults.merged_into(nested)} + return self.request_defaults.merged_into(msg_obj) + async def _mask_response_create(self, message: str) -> str: """ - Enforce the authorized model and apply Presidio PII masking to a - ``response.create`` message before it is forwarded to the upstream - provider. + Merge deployment defaults, enforce the authorized model, and apply + Presidio PII masking to a ``response.create`` message before it is + forwarded to the upstream provider. + - Fills the deployment's ``litellm_params`` request defaults into the + frame the way the HTTP ``/v1/responses`` path does: client-set keys + win, ``extra_body`` entries override. - Overwrites any ``model`` field with the connection-authorized model to prevent deployment-substitution attacks (always applied). - Walks the ``input`` and ``instructions`` fields, calls ``check_pii`` @@ -1889,23 +1903,26 @@ class ResponsesWebSocketStreaming: Non-``response.create`` messages are returned unchanged. """ try: - msg_obj: Final = _load_json_object(message) + parsed: Final = _load_json_object(message) except (json.JSONDecodeError, TypeError): return message - if msg_obj.get("type") != "response.create": + if parsed.get("type") != "response.create": return message + msg_obj: Final = self._with_request_defaults(parsed) + defaults_applied: Final = msg_obj != parsed + # Always enforce the authorized model, even when PII masking is off. model_modified: Final = self._enforce_authorized_model(msg_obj) if not self.guardrail_callbacks: - return json.dumps(msg_obj) if model_modified else message + return json.dumps(msg_obj) if model_modified or defaults_applied else message if "metadata" not in self.request_data: self.request_data["metadata"] = {} - modified = model_modified + modified = model_modified or defaults_applied guardrail_cbs: Final[tuple[PresidioGuardrailCallback, ...]] = tuple(self.guardrail_callbacks) for cb in guardrail_cbs: presidio_config = cb.get_presidio_settings_from_request_data(self.request_data) diff --git a/litellm/types/responses/streaming_websocket.py b/litellm/types/responses/streaming_websocket.py index 2aa71647955..f369cbcebf8 100644 --- a/litellm/types/responses/streaming_websocket.py +++ b/litellm/types/responses/streaming_websocket.py @@ -1,5 +1,7 @@ from __future__ import annotations +from collections.abc import Mapping +from dataclasses import dataclass from typing import Protocol from litellm.types.guardrails import PresidioPerRequestConfig @@ -39,3 +41,14 @@ class PresidioGuardrailCallback(Protocol): presidio_config: PresidioPerRequestConfig | None, request_data: dict[str, object], ) -> str: ... + + +@dataclass(frozen=True, slots=True) +class ResponsesWebSocketRequestDefaults: + """Deployment-level request parameters merged into every ``response.create`` frame relayed over a native websocket.""" + + fill_missing: Mapping[str, object] + overrides: Mapping[str, object] + + def merged_into(self, request: Mapping[str, object]) -> dict[str, object]: + return {**self.fill_missing, **request, **self.overrides} diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index fe3c4a0640d..918c4a33d40 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1204,6 +1204,216 @@ class TestWebSocketProjectQuotaEnforcement: quota_callback.enforce_project_io_token_quota_for_frame.assert_awaited_once() +def _deployment_defaults(): + from types import MappingProxyType + + from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults + + return ResponsesWebSocketRequestDefaults( + fill_missing=MappingProxyType({"reasoning": {"effort": "high"}, "service_tier": "priority"}), + overrides=MappingProxyType({"provider_default": "configured"}), + ) + + +class TestNativeWebSocketDeploymentDefaults: + """The native relay merges deployment litellm_params into every response.create like HTTP does.""" + + def test_builder_maps_router_kwargs_like_the_http_path(self): + from litellm.responses.main import _build_responses_websocket_request_defaults + + defaults = _build_responses_websocket_request_defaults( + { + "model": "gpt-5-pro", + "reasoning_effort": "high", + "service_tier": "priority", + "extra_body": {"provider_default": "configured"}, + "temperature": None, + "timeout": 600, + "max_retries": 2, + "caching": False, + "custom_llm_provider": "openai", + "litellm_metadata": {"user_api_key": "hashed"}, + "user_api_key_dict": MagicMock(), + "litellm_logging_obj": MagicMock(), + "websocket": MagicMock(), + } + ) + + assert dict(defaults.fill_missing) == {"reasoning": {"effort": "high"}, "service_tier": "priority"} + assert dict(defaults.overrides) == {"provider_default": "configured"} + + def test_builder_keeps_explicit_reasoning_over_reasoning_effort(self): + from litellm.responses.main import _build_responses_websocket_request_defaults + + defaults = _build_responses_websocket_request_defaults( + {"model": "gpt-5-pro", "reasoning": {"effort": "low"}, "reasoning_effort": "high"} + ) + + assert dict(defaults.fill_missing) == {"reasoning": {"effort": "low"}} + assert dict(defaults.overrides) == {} + + @pytest.mark.asyncio + async def test_flat_frame_gets_defaults_client_keys_win_extra_body_overrides(self): + handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults()) + + forwarded = json.loads( + await handler._mask_response_create( + json.dumps( + { + "type": "response.create", + "model": "gpt-5-pro", + "input": "Say hello", + "service_tier": "default", + "provider_default": "client", + } + ) + ) + ) + + assert forwarded == { + "type": "response.create", + "model": "gpt-5-pro", + "input": "Say hello", + "service_tier": "default", + "provider_default": "configured", + "reasoning": {"effort": "high"}, + } + + @pytest.mark.asyncio + async def test_nested_response_frame_gets_defaults_inside_response(self): + handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults()) + + forwarded = json.loads( + await handler._mask_response_create( + json.dumps({"type": "response.create", "response": {"model": "gpt-5-pro", "input": "hi"}}) + ) + ) + + assert forwarded == { + "type": "response.create", + "response": { + "model": "gpt-5-pro", + "input": "hi", + "reasoning": {"effort": "high"}, + "service_tier": "priority", + "provider_default": "configured", + }, + } + + @pytest.mark.asyncio + async def test_frames_that_need_nothing_pass_through_untouched(self): + handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults()) + cancel_frame = json.dumps({"type": "response.cancel"}) + complete_frame = json.dumps( + { + "type": "response.create", + "model": "gpt-5-pro", + "input": "hi", + "reasoning": {"effort": "high"}, + "service_tier": "priority", + "provider_default": "configured", + } + ) + + assert await handler._mask_response_create(cancel_frame) is cancel_frame + assert await handler._mask_response_create(complete_frame) is complete_frame + + @pytest.mark.asyncio + async def test_handler_applies_defaults_to_the_first_frame_sent_upstream(self): + import asyncio + from unittest.mock import AsyncMock, patch + + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + + class FakeBackend: + def __init__(self): + self.sent = [] + + async def send(self, message): + self.sent.append(message) + + async def recv(self, decode=False): + raise RuntimeError("backend closed") + + async def close(self): + pass + + backend = FakeBackend() + + class FakeConnect: + def __init__(self, url, **kwargs): + pass + + async def __aenter__(self): + return backend + + async def __aexit__(self, *args): + pass + + mock_config = MagicMock(spec=OpenAIResponsesAPIConfig) + mock_config.supports_native_websocket.return_value = True + mock_config.model_in_websocket_url.return_value = True + mock_config.get_websocket_url.return_value = "wss://api.openai.com/v1/responses" + mock_config.validate_environment.return_value = {} + + mock_logging = MagicMock() + mock_logging.pre_call = MagicMock() + mock_logging.dispatch_success_handlers = AsyncMock() + + client_ws = MagicMock() + client_ws.receive_text = AsyncMock(side_effect=RuntimeError("client closed")) + client_ws.send_text = AsyncMock() + client_ws.close = AsyncMock() + + with patch("websockets.connect", FakeConnect): + await BaseLLMHTTPHandler().async_responses_websocket( + model="gpt-5-pro", + websocket=client_ws, + logging_obj=mock_logging, + responses_api_provider_config=mock_config, + api_key="sk-test", + first_message=json.dumps({"type": "response.create", "model": "gpt-5-pro", "input": "Say hello"}), + request_defaults=_deployment_defaults(), + ) + await asyncio.sleep(0) + + assert [json.loads(frame) for frame in backend.sent] == [ + { + "type": "response.create", + "model": "gpt-5-pro", + "input": "Say hello", + "reasoning": {"effort": "high"}, + "service_tier": "priority", + "provider_default": "configured", + } + ] + + @pytest.mark.asyncio + async def test_aresponses_websocket_builds_defaults_from_deployment_kwargs(self, monkeypatch): + import importlib + from unittest.mock import AsyncMock + + responses_main = importlib.import_module("litellm.responses.main") + + stub = MagicMock() + stub.async_responses_websocket = AsyncMock() + monkeypatch.setattr(responses_main, "base_llm_http_handler", stub) + + await responses_main._aresponses_websocket.__wrapped__( + model="openai/gpt-5-pro", + websocket=MagicMock(), + api_key="sk-test", + litellm_logging_obj=MagicMock(), + reasoning_effort="high", + service_tier="priority", + extra_body={"provider_default": "configured"}, + ) + + request_defaults = stub.async_responses_websocket.call_args.kwargs["request_defaults"] + assert dict(request_defaults.fill_missing) == {"reasoning": {"effort": "high"}, "service_tier": "priority"} + assert dict(request_defaults.overrides) == {"provider_default": "configured"} + + class TestNativeWebSocketGuardrails: @pytest.mark.asyncio async def test_response_create_injects_authorized_model(self): From 1f27c442b4146e61d77c904a5c5d9e5165602267 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 14:17:00 -0700 Subject: [PATCH 217/253] ci(duplicate-check): let Codex reach GitHub from its sandbox Every gh search in the first real runs failed with "error connecting to api.github.com", so the verdict was always null. The legacy sandbox_permissions key no longer grants network in read-only mode; the workspace-write sandbox has a network_access switch that does. Pin the CLI to the version the prompt was proven on --- .github/workflows/duplicate_issue_check.yml | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/.github/workflows/duplicate_issue_check.yml b/.github/workflows/duplicate_issue_check.yml index b12f894328e..83c6cc06de7 100644 --- a/.github/workflows/duplicate_issue_check.yml +++ b/.github/workflows/duplicate_issue_check.yml @@ -93,10 +93,11 @@ jobs: responses-api-endpoint: ${{ vars.LITELLM_API_BASE }}/v1/responses prompt-file: .github/prompts/duplicate-issue-check.md output-schema-file: .github/prompts/duplicate-issue-check.schema.json - sandbox: read-only - # read-only denies network, and the whole method is searching the tracker with gh - codex-args: '["-c", "sandbox_permissions=[\"network-full-access\"]"]' + sandbox: workspace-write + # The whole method is searching the tracker with gh, and network is only switchable in workspace-write + codex-args: '["-c", "sandbox_workspace_write.network_access=true"]' model: ${{ vars.DUPLICATE_CHECK_MODEL }} + codex-version: "0.154.0" # Issue authors have no write access and the action refuses them by default; the # prompt is fixed, the sandbox read-only, and the only token is read-only on a public repo allow-users: "*" From 1e8b8f7c33b83b35463e5efa9c11c33eb3929d54 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 14:21:22 -0700 Subject: [PATCH 218/253] ci(duplicate-check): describe the workspace-write sandbox accurately --- .github/workflows/duplicate_issue_check.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/duplicate_issue_check.yml b/.github/workflows/duplicate_issue_check.yml index 83c6cc06de7..91eee3ca383 100644 --- a/.github/workflows/duplicate_issue_check.yml +++ b/.github/workflows/duplicate_issue_check.yml @@ -98,8 +98,8 @@ jobs: codex-args: '["-c", "sandbox_workspace_write.network_access=true"]' model: ${{ vars.DUPLICATE_CHECK_MODEL }} codex-version: "0.154.0" - # Issue authors have no write access and the action refuses them by default; the - # prompt is fixed, the sandbox read-only, and the only token is read-only on a public repo + # Issue authors have no write access and the action refuses them by default; the prompt is + # fixed, writes stay inside the throwaway checkout, and the only token is read-only on a public repo allow-users: "*" - name: Summary From 9767878425ea43788925d567ce5499cc02dc44f8 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 21:27:03 +0000 Subject: [PATCH 219/253] test(ocr): replace per-provider OCR test classes with a declarative provider x auth x input matrix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/ocr_tests/base_ocr_unit_tests.py | 203 ----------- tests/ocr_tests/test_ocr_azure_ai.py | 29 -- .../test_ocr_azure_document_intelligence.py | 49 +-- tests/ocr_tests/test_ocr_matrix.py | 317 ++++++++++++++++++ tests/ocr_tests/test_ocr_mistral.py | 92 ----- tests/ocr_tests/test_ocr_vertex_ai.py | 127 +------ tests/ocr_tests/vertex_key.json | 13 - 7 files changed, 328 insertions(+), 502 deletions(-) delete mode 100644 tests/ocr_tests/base_ocr_unit_tests.py delete mode 100644 tests/ocr_tests/test_ocr_azure_ai.py create mode 100644 tests/ocr_tests/test_ocr_matrix.py delete mode 100644 tests/ocr_tests/test_ocr_mistral.py delete mode 100644 tests/ocr_tests/vertex_key.json diff --git a/tests/ocr_tests/base_ocr_unit_tests.py b/tests/ocr_tests/base_ocr_unit_tests.py deleted file mode 100644 index ae65efd952d..00000000000 --- a/tests/ocr_tests/base_ocr_unit_tests.py +++ /dev/null @@ -1,203 +0,0 @@ -""" -Base test class for OCR functionality across different providers. - -This follows the same pattern as BaseLLMChatTest in tests/llm_translation/base_llm_unit_tests.py -""" - -import pytest -import litellm -import os -from abc import ABC, abstractmethod - - -# Test resources -TEST_IMAGE_PATH = "test_image_edit.png" -# Tiny in-repo PDF served via jsdelivr (sha-pinned, immutable). The arxiv -# PDF previously used here was several MB — once base64-encoded into the -# Vertex OCR request it ballooned cassettes past 100 MB per test. Keep -# the URL stable across runs so cassettes don't churn. -TEST_PDF_URL = ( - "https://cdn.jsdelivr.net/gh/BerriAI/litellm" - "@d769e81c90d453240c61fc572cdb27fae06a89d0" - "/tests/llm_translation/fixtures/dummy.pdf" -) - - -class BaseOCRTest(ABC): - """ - Abstract base test class that enforces common OCR tests across all providers. - - Each provider-specific test class should inherit from this and implement - get_base_ocr_call_args() to return provider-specific configuration. - """ - - @abstractmethod - def get_base_ocr_call_args(self) -> dict: - """Must return the base OCR call args for the specific provider""" - pass - - @pytest.mark.parametrize("sync_mode", [True, False]) - @pytest.mark.asyncio - async def test_basic_ocr_with_url(self, sync_mode): - """ - Test basic OCR with a public URL. - """ - litellm._turn_on_debug() - base_ocr_call_args = self.get_base_ocr_call_args() - print("BASE OCR Call args=", base_ocr_call_args) - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - - try: - if sync_mode: - response = litellm.ocr( - document={"type": "document_url", "document_url": TEST_PDF_URL}, - **base_ocr_call_args, - ) - else: - response = await litellm.aocr( - document={"type": "document_url", "document_url": TEST_PDF_URL}, - **base_ocr_call_args, - ) - - print(f"\n{'='*80}") - print(f"Sync Mode: {sync_mode}") - print(f"Response type: {type(response)}") - print( - f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}" - ) - - # Check if response has expected OCR format - assert hasattr(response, "pages"), "Response should have 'pages' attribute" - assert hasattr(response, "model"), "Response should have 'model' attribute" - assert hasattr( - response, "object" - ), "Response should have 'object' attribute" - assert ( - response.object == "ocr" - ), f"Expected object='ocr', got '{response.object}'" - - # Validate pages structure - assert isinstance(response.pages, list), "pages should be a list" - assert len(response.pages) > 0, "Should have at least one page" - - # Check first page structure - first_page = response.pages[0] - assert hasattr(first_page, "index"), "Page should have 'index' attribute" - assert hasattr( - first_page, "markdown" - ), "Page should have 'markdown' attribute" - - # Extract text from all pages for validation - total_text = "\n\n".join( - page.markdown for page in response.pages if page.markdown - ) - print(f"Total pages: {len(response.pages)}") - print(f"Total extracted text length: {len(total_text)} characters") - print(f"First 200 chars: {total_text[:200]}") - print(f"Model: {response.model}") - if response.usage_info: - print(f"Pages processed: {response.usage_info.pages_processed}") - print(f"{'='*80}\n") - - assert len(total_text) > 0, "Should extract some text from the document" - - ######################################################### - # validate we get a response cost in hidden parameters - ######################################################### - hidden_params = response._hidden_params - assert isinstance( - hidden_params, dict - ), "Hidden parameters should be a dictionary" - - print("response usage_info:", response.usage_info) - - response_cost = hidden_params.get("response_cost") - assert ( - response_cost is not None - ), "Response cost should be in hidden parameters" - assert response_cost > 0, "Response cost should be greater than 0" - print("response_cost=", response_cost) - - except litellm.RateLimitError as e: - error_msg = str(e) - if "Quota exceeded" in error_msg or "RESOURCE_EXHAUSTED" in error_msg: - pytest.skip(f"Quota exceeded - {error_msg}") - else: - pytest.skip(f"Rate limit exceeded - {error_msg}") - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - except litellm.BadRequestError as e: - error_msg = str(e) - if ( - "URL_REJECTED" in error_msg - or "Cannot fetch content from the provided URL" in error_msg - ): - pytest.skip(f"URL rejected by provider - {error_msg}") - else: - pytest.fail(f"OCR call failed: {str(e)}") - except Exception as e: - pytest.fail(f"OCR call failed: {str(e)}") - - def test_ocr_response_structure(self): - """ - Test that the OCR response has the correct structure. - """ - litellm.set_verbose = True - base_ocr_call_args = self.get_base_ocr_call_args() - - try: - response = litellm.ocr( - document={"type": "document_url", "document_url": TEST_PDF_URL}, - **base_ocr_call_args, - ) - - # Validate response structure - assert hasattr(response, "pages"), "Response should have 'pages' attribute" - assert hasattr(response, "model"), "Response should have 'model' attribute" - assert hasattr( - response, "object" - ), "Response should have 'object' attribute" - assert hasattr( - response, "usage_info" - ), "Response should have 'usage_info' attribute" - - assert isinstance(response.pages, list), "pages should be a list" - assert len(response.pages) > 0, "Should have at least one page" - assert response.object == "ocr", "object should be 'ocr'" - - # Validate first page structure - first_page = response.pages[0] - assert hasattr(first_page, "index"), "Page should have 'index' attribute" - assert hasattr( - first_page, "markdown" - ), "Page should have 'markdown' attribute" - assert isinstance(first_page.markdown, str), "markdown should be a string" - - print(f"\nResponse structure validated:") - print(f" - object: {response.object}") - print(f" - model: {response.model}") - print(f" - pages: {len(response.pages)}") - if response.usage_info: - print(f" - pages_processed: {response.usage_info.pages_processed}") - print(f" - doc_size_bytes: {response.usage_info.doc_size_bytes}") - - except litellm.RateLimitError as e: - error_msg = str(e) - if "Quota exceeded" in error_msg or "RESOURCE_EXHAUSTED" in error_msg: - pytest.skip(f"Quota exceeded - {error_msg}") - else: - pytest.skip(f"Rate limit exceeded - {error_msg}") - except litellm.InternalServerError: - pytest.skip("Model is overloaded") - except litellm.BadRequestError as e: - error_msg = str(e) - if ( - "URL_REJECTED" in error_msg - or "Cannot fetch content from the provided URL" in error_msg - ): - pytest.skip(f"URL rejected by provider - {error_msg}") - else: - pytest.fail(f"OCR response structure test failed: {str(e)}") - except Exception as e: - pytest.fail(f"OCR response structure test failed: {str(e)}") diff --git a/tests/ocr_tests/test_ocr_azure_ai.py b/tests/ocr_tests/test_ocr_azure_ai.py deleted file mode 100644 index acb44958fd9..00000000000 --- a/tests/ocr_tests/test_ocr_azure_ai.py +++ /dev/null @@ -1,29 +0,0 @@ -""" -Test OCR functionality with Azure AI API. - -Note: Azure AI OCR automatically converts URLs to base64 data URIs since -the Azure AI endpoint doesn't have internet access. -""" - -import os -from base_ocr_unit_tests import BaseOCRTest - - -class TestAzureAIOCR(BaseOCRTest): - """ - Test class for Azure AI OCR functionality. - Inherits from BaseOCRTest and provides Azure AI-specific configuration. - - Note: For Azure AI, LiteLLM will automatically convert URLs to base64 data URIs before - sending to the API, since Azure AI OCR endpoint doesn't have internet access. - """ - - def get_base_ocr_call_args(self) -> dict: - """ - Return the base OCR call args for Azure AI. - """ - return { - "model": "azure_ai/mistral-document-ai-2512", - "api_key": os.getenv("AZURE_API_KEY"), - "api_base": os.getenv("AZURE_API_BASE"), - } diff --git a/tests/ocr_tests/test_ocr_azure_document_intelligence.py b/tests/ocr_tests/test_ocr_azure_document_intelligence.py index e6a2e5e5735..521f85ca8f7 100644 --- a/tests/ocr_tests/test_ocr_azure_document_intelligence.py +++ b/tests/ocr_tests/test_ocr_azure_document_intelligence.py @@ -1,53 +1,13 @@ -""" -Test OCR functionality with Azure Document Intelligence API. - -Azure Document Intelligence provides advanced document analysis capabilities -using the v4.0 (2024-11-30) API. -""" - -import os +"""Azure Document Intelligence request transformation: Mistral-shaped `pages` to Azure's query string.""" import pytest -from base_ocr_unit_tests import BaseOCRTest from litellm.constants import AZURE_DOCUMENT_INTELLIGENCE_API_VERSION from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, ) -class TestAzureDocumentIntelligenceOCR(BaseOCRTest): - """ - Test class for Azure Document Intelligence OCR functionality. - - Inherits from BaseOCRTest and provides Azure Document Intelligence-specific configuration. - - Tests the azure_ai/doc-intelligence/ provider route. - """ - - def get_base_ocr_call_args(self) -> dict: - """ - Return the base OCR call args for Azure Document Intelligence. - - Uses prebuilt-layout model which is closest to Mistral OCR format. - """ - # Check for required environment variables - api_key = os.environ.get("AZURE_DOCUMENT_INTELLIGENCE_API_KEY") - endpoint = os.environ.get("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - - if not api_key or not endpoint: - pytest.skip( - "AZURE_DOCUMENT_INTELLIGENCE_API_KEY and AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT " - "environment variables are required for Azure Document Intelligence tests" - ) - - return { - "model": "azure_ai/doc-intelligence/prebuilt-layout", - "api_key": api_key, - "api_base": endpoint, - } - - class TestAzureDocumentIntelligencePagesParam: """ Unit tests for the Mistral-compatible `pages` parameter translation to @@ -101,7 +61,7 @@ class TestAzureDocumentIntelligencePagesParam: cfg.map_ocr_params({"pages": [True, False]}, {}, "prebuilt-layout") def test_map_ocr_params_unsupported_type_raises(self, cfg): - with pytest.raises(ValueError, match='based, Mistral-style\\) or a string like'): + with pytest.raises(ValueError, match="based, Mistral-style\\) or a string like"): cfg.map_ocr_params({"pages": 5}, {}, "prebuilt-layout") def test_get_complete_url_appends_pages_query(self, cfg): @@ -110,9 +70,7 @@ class TestAzureDocumentIntelligencePagesParam: model="azure_ai/doc-intelligence/prebuilt-layout", optional_params={"pages": "1-3,5"}, ) - assert ( - f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url - ), url + assert f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url, url assert "pages=1-3,5" in url, url assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url @@ -168,4 +126,3 @@ class TestAzureDocumentIntelligencePagesParam: assert "pages=3,4,5,6,7,8,9" in url assert req.data == {"urlSource": "https://example.com/x.pdf"} - diff --git a/tests/ocr_tests/test_ocr_matrix.py b/tests/ocr_tests/test_ocr_matrix.py new file mode 100644 index 00000000000..13cffbbc9a1 --- /dev/null +++ b/tests/ocr_tests/test_ocr_matrix.py @@ -0,0 +1,317 @@ +"""Live provider x auth x input coverage for ``litellm.ocr`` / ``litellm.aocr``. + +Each ``Case`` is one hand-picked cell, not the full cross product: every provider +exercises each of its credential kinds in both ``explicit`` (kwargs) and ``env`` +(monkeypatched environment) mode at least once, and every input kind a provider +accepts is exercised at least once. Sync and async are spread across the cells. +Every cell also checks the success callback saw the same response and cost. +""" + +from __future__ import annotations + +import asyncio +import base64 +import io +import os +import re +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal + +import pytest + +import litellm +from litellm import Router +from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.base_llm.ocr.transformation import OCRResponse + +Document = Mapping[str, object] +AuthMode = Literal["explicit", "env"] +CallStyle = Literal["sync", "async"] + + +@dataclass(frozen=True, slots=True) +class LoggedCall: + payload: Mapping[str, object] + response: object + + +class RecordingLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() # pyright: ignore[reportUnknownMemberType] # CustomLogger.__init__ is untyped + self.calls: Final[list[LoggedCall]] = [] # mutable-ok: append-only sink the callback hooks write into + + def _record(self, kwargs: Mapping[str, object], response_obj: object) -> None: + payload: Final = _string_keyed(kwargs.get("standard_logging_object")) + self.calls.append(LoggedCall(payload, response_obj)) + + def log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self._record(kwargs, response_obj) + + async def async_log_success_event( + self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + self._record(kwargs, response_obj) + + async def wait_for_call(self, timeout: float = 10.0) -> LoggedCall: + deadline: Final = asyncio.get_running_loop().time() + timeout + while not self.calls: + assert asyncio.get_running_loop().time() < deadline, "success callback never fired" + await asyncio.sleep(0.05) + assert len(self.calls) == 1, self.calls + return self.calls[0] + + +@pytest.fixture +def logger(monkeypatch: pytest.MonkeyPatch) -> RecordingLogger: + recorder: Final = RecordingLogger() + monkeypatch.setattr(litellm, "callbacks", [recorder]) + for registry in ("success_callback", "_async_success_callback", "failure_callback", "_async_failure_callback"): + monkeypatch.setattr(litellm, registry, []) + return recorder + + +TESTS_DIR: Final = Path(__file__).resolve().parents[1] +PDF_PATH: Final = TESTS_DIR / "llm_translation" / "fixtures" / "dummy.pdf" +PNG_PATH: Final = TESTS_DIR / "image_gen_tests" / "test_image.png" +PINNED_CDN: Final = "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0" +PDF_URL: Final = f"{PINNED_CDN}/tests/llm_translation/fixtures/dummy.pdf" +PNG_URL: Final = f"{PINNED_CDN}/tests/image_gen_tests/test_image.png" +PDF_TEXT: Final = "Test PDF File" +PNG_TEXT: Final = "LiteLLM" + + +class _NamedReader(io.BytesIO): + def __init__(self, path: Path) -> None: + super().__init__(path.read_bytes()) + self.name: Final = path.name + + +def _data_uri(path: Path, mime: str) -> str: + return f"data:{mime};base64,{base64.b64encode(path.read_bytes()).decode()}" + + +@dataclass(frozen=True, slots=True) +class Input: + id: str + build: Callable[[], Document] + expected_text: str + + +PDF_BY_URL: Final = Input("pdf_url", lambda: {"type": "document_url", "document_url": PDF_URL}, PDF_TEXT) +PNG_BY_URL: Final = Input("image_url", lambda: {"type": "image_url", "image_url": PNG_URL}, PNG_TEXT) +PDF_DATA_URI: Final = Input( + "pdf_data_uri", + lambda: {"type": "document_url", "document_url": _data_uri(PDF_PATH, "application/pdf")}, + PDF_TEXT, +) +PNG_DATA_URI: Final = Input( + "image_data_uri", lambda: {"type": "image_url", "image_url": _data_uri(PNG_PATH, "image/png")}, PNG_TEXT +) +PDF_AS_PATH: Final = Input("pdf_path", lambda: {"type": "file", "file": PDF_PATH}, PDF_TEXT) +PDF_AS_BYTES: Final = Input( + "pdf_bytes", lambda: {"type": "file", "file": PDF_PATH.read_bytes(), "mime_type": "application/pdf"}, PDF_TEXT +) +PNG_AS_BYTES: Final = Input( + "image_bytes", lambda: {"type": "file", "file": PNG_PATH.read_bytes(), "mime_type": "image/png"}, PNG_TEXT +) +PNG_AS_FILE_OBJECT: Final = Input( + "image_file_object", lambda: {"type": "file", "file": _NamedReader(PNG_PATH)}, PNG_TEXT +) + + +@dataclass(frozen=True, slots=True) +class Secret: + """One credential value: the ``litellm.ocr`` kwarg it travels in, the env var litellm reads + when the kwarg is omitted, and the env var that holds the value in the test process.""" + + kwarg: str + env: str + source: str | None = None + + @property + def source_env(self) -> str: + return self.source or self.env + + +@dataclass(frozen=True, slots=True) +class Credential: + id: str + secrets: tuple[Secret, ...] + + +@dataclass(frozen=True, slots=True) +class Provider: + id: str + model: str + credentials: tuple[Credential, ...] + params: Mapping[str, str] = MappingProxyType({}) + + @property + def env_vars(self) -> frozenset[str]: + return frozenset(secret.env for credential in self.credentials for secret in credential.secrets) + + +MISTRAL_KEY: Final = Credential("api_key", (Secret("api_key", "MISTRAL_API_KEY"),)) +COHERE_KEY: Final = Credential("api_key", (Secret("api_key", "COHERE_API_KEY"),)) +REDUCTO_KEY: Final = Credential("api_key", (Secret("api_key", "REDUCTO_API_KEY"),)) + +AZURE_ENTRA_SECRETS: Final = ( + Secret("tenant_id", "AZURE_TENANT_ID", "AZURE_FOUNDRY_TENANT_ID"), + Secret("client_id", "AZURE_CLIENT_ID", "AZURE_FOUNDRY_ADMIN_CLIENT_ID"), + Secret("client_secret", "AZURE_CLIENT_SECRET", "AZURE_FOUNDRY_ADMIN_CLIENT_SECRET"), +) +AZURE_AI_BASE: Final = Secret("api_base", "AZURE_AI_API_BASE") +AZURE_AI_KEY: Final = Credential("api_key", (AZURE_AI_BASE, Secret("api_key", "AZURE_AI_API_KEY"))) +AZURE_AI_ENTRA: Final = Credential("entra", (AZURE_AI_BASE, *AZURE_ENTRA_SECRETS)) + +AZURE_DI_BASE: Final = Secret("api_base", "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") +AZURE_DI_KEY: Final = Credential("api_key", (AZURE_DI_BASE, Secret("api_key", "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"))) +AZURE_DI_ENTRA: Final = Credential("entra", (AZURE_DI_BASE, *AZURE_ENTRA_SECRETS)) + +VERTEX_SERVICE_ACCOUNT: Final = Credential( + "service_account", + (Secret("vertex_credentials", "VERTEXAI_CREDENTIALS"), Secret("vertex_project", "VERTEXAI_PROJECT")), +) + +MISTRAL: Final = Provider("mistral", "mistral/mistral-ocr-latest", (MISTRAL_KEY,)) +AZURE_AI_MISTRAL: Final = Provider( + "azure_ai_mistral", "azure_ai/mistral-document-ai-2512", (AZURE_AI_KEY, AZURE_AI_ENTRA) +) +AZURE_DOC_INTELLIGENCE: Final = Provider( + "azure_doc_intelligence", "azure_ai/doc-intelligence/prebuilt-layout", (AZURE_DI_KEY, AZURE_DI_ENTRA) +) +COHERE: Final = Provider("cohere", "cohere/parse-v5.0", (COHERE_KEY,)) +REDUCTO_V3: Final = Provider("reducto_v3", "reducto/parse-v3", (REDUCTO_KEY,)) +REDUCTO_LEGACY: Final = Provider("reducto_legacy", "reducto/parse-legacy", (REDUCTO_KEY,)) +VERTEX_MISTRAL: Final = Provider( + "vertex_mistral", + "vertex_ai/mistral-ocr-2505", + (VERTEX_SERVICE_ACCOUNT,), + MappingProxyType({"vertex_location": "us-central1"}), +) + + +@dataclass(frozen=True, slots=True) +class Case: + provider: Provider + credential: Credential + auth: AuthMode + document: Input + call: CallStyle + + @property + def id(self) -> str: + return f"{self.provider.id}-{self.credential.id}-{self.auth}-{self.document.id}-{self.call}" + + def bind_credentials(self, monkeypatch: pytest.MonkeyPatch) -> Mapping[str, str]: + """Clear every env var the provider could fall back to, then supply this case's values via kwargs or env.""" + values: Final = {secret: os.environ.get(secret.source_env) for secret in self.credential.secrets} + missing: Final = tuple(secret.source_env for secret, value in values.items() if not value) + if missing: + pytest.skip(f"{', '.join(missing)} not set") + for env_var in self.provider.env_vars: + monkeypatch.delenv(env_var, raising=False) + if self.auth == "explicit": + return {secret.kwarg: value for secret, value in values.items() if value} + for secret, value in values.items(): + monkeypatch.setenv(secret.env, value or "") + return {} + + async def run(self, credentials: Mapping[str, str]) -> OCRResponse: + kwargs: Final = {**self.provider.params, **credentials} + document: Final = self.document.build() + response: Final = ( + await litellm.aocr(model=self.provider.model, document=document, **kwargs) # pyright: ignore[reportUnknownMemberType] # @client erases the signature + if self.call == "async" + else litellm.ocr(model=self.provider.model, document=document, **kwargs) + ) + assert isinstance(response, OCRResponse) + return response + + +CASES: Final = ( + Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "sync"), + Case(MISTRAL, MISTRAL_KEY, "env", PNG_BY_URL, "async"), + Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_AS_PATH, "sync"), + Case(MISTRAL, MISTRAL_KEY, "explicit", PNG_AS_BYTES, "async"), + Case(MISTRAL, MISTRAL_KEY, "explicit", PNG_AS_FILE_OBJECT, "sync"), + Case(AZURE_AI_MISTRAL, AZURE_AI_KEY, "explicit", PDF_BY_URL, "sync"), + Case(AZURE_AI_MISTRAL, AZURE_AI_KEY, "env", PNG_BY_URL, "async"), + Case(AZURE_AI_MISTRAL, AZURE_AI_ENTRA, "explicit", PDF_AS_PATH, "sync"), + Case(AZURE_AI_MISTRAL, AZURE_AI_ENTRA, "env", PDF_DATA_URI, "async"), + Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_KEY, "explicit", PDF_BY_URL, "sync"), + Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_KEY, "env", PNG_AS_BYTES, "async"), + Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_ENTRA, "explicit", PNG_BY_URL, "async"), + Case(AZURE_DOC_INTELLIGENCE, AZURE_DI_ENTRA, "env", PDF_AS_PATH, "sync"), + Case(COHERE, COHERE_KEY, "explicit", PNG_BY_URL, "sync"), + Case(COHERE, COHERE_KEY, "env", PNG_DATA_URI, "async"), + Case(REDUCTO_V3, REDUCTO_KEY, "explicit", PDF_AS_PATH, "sync"), + Case(REDUCTO_V3, REDUCTO_KEY, "env", PNG_AS_BYTES, "async"), + Case(REDUCTO_V3, REDUCTO_KEY, "explicit", PDF_DATA_URI, "async"), + Case(REDUCTO_LEGACY, REDUCTO_KEY, "explicit", PDF_AS_BYTES, "sync"), + Case(VERTEX_MISTRAL, VERTEX_SERVICE_ACCOUNT, "explicit", PDF_BY_URL, "sync"), + Case(VERTEX_MISTRAL, VERTEX_SERVICE_ACCOUNT, "env", PNG_BY_URL, "async"), +) + + +def _response_cost(response: OCRResponse) -> float: + response_cost: Final[object] = response._hidden_params.get("response_cost") # pyright: ignore[reportPrivateUsage, reportUnknownMemberType, reportUnknownVariableType] # response_cost is only surfaced on _hidden_params + assert isinstance(response_cost, float) and response_cost > 0 + return response_cost + + +def _assert_ocr_response(response: OCRResponse, model: str, expected_text: str) -> None: + assert response.object == "ocr" + assert response.model == model.split("/", 1)[1] + assert [page.index for page in response.pages] == list(range(len(response.pages))) + text: Final = re.sub(r"\s+", " ", " ".join(page.markdown for page in response.pages)) + assert expected_text.lower() in text.lower(), text + assert response.usage_info is not None + assert response.usage_info.pages_processed == len(response.pages) + _response_cost(response) + + +def _string_keyed(value: object) -> Mapping[str, object]: + assert isinstance(value, Mapping), type(value) + items: Final = tuple(value.items()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType, reportUnknownArgumentType] # narrowed from object + return MappingProxyType({str(key): value for key, value in items}) # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # narrowed from object + + +def _assert_logged(logged: LoggedCall, response: OCRResponse, model: str, logged_model: str, call: CallStyle) -> None: + assert isinstance(logged.response, OCRResponse) + assert logged.response.pages == response.pages + assert logged.payload["status"] == "success" + assert logged.payload["call_type"] == ("aocr" if call == "async" else "ocr") + assert logged.payload["custom_llm_provider"] == model.split("/", 1)[0] + assert logged.payload["model"] == logged_model + assert logged.payload["response_cost"] == _response_cost(response) + + +@pytest.mark.parametrize("case", CASES, ids=[case.id for case in CASES]) +async def test_ocr(case: Case, monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None: + credentials: Final = case.bind_credentials(monkeypatch) + response: Final = await case.run(credentials) + _assert_ocr_response(response, case.provider.model, case.document.expected_text) + _assert_logged(await logger.wait_for_call(), response, case.provider.model, response.model, case.call) + + +async def test_router_aocr(monkeypatch: pytest.MonkeyPatch, logger: RecordingLogger) -> None: + case: Final = Case(MISTRAL, MISTRAL_KEY, "explicit", PDF_BY_URL, "async") + router: Final = Router( + model_list=[ + { + "model_name": "ocr-alias", + "litellm_params": {"model": MISTRAL.model, **case.bind_credentials(monkeypatch)}, + } + ] + ) + response: Final = await router.aocr(model="ocr-alias", document=PDF_BY_URL.build()) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # Router.aocr is untyped + assert isinstance(response, OCRResponse) + _assert_ocr_response(response, MISTRAL.model, PDF_TEXT) + _assert_logged(await logger.wait_for_call(), response, MISTRAL.model, MISTRAL.model, case.call) diff --git a/tests/ocr_tests/test_ocr_mistral.py b/tests/ocr_tests/test_ocr_mistral.py deleted file mode 100644 index cdc093620b2..00000000000 --- a/tests/ocr_tests/test_ocr_mistral.py +++ /dev/null @@ -1,92 +0,0 @@ -""" -Test OCR functionality with Mistral API. -""" - -import os -import sys -import pytest -import litellm -from litellm import Router -from base_ocr_unit_tests import BaseOCRTest, TEST_PDF_URL - - -class TestMistralOCR(BaseOCRTest): - """ - Test class for Mistral OCR functionality. - """ - - def get_base_ocr_call_args(self) -> dict: - """Return the base OCR call args for Mistral""" - return { - "model": "mistral/mistral-ocr-latest", - "api_key": os.getenv("MISTRAL_API_KEY"), - } - - -@pytest.mark.asyncio -async def test_router_aocr_with_mistral(): - """ - Test OCR with Router using Mistral OCR deployment. - """ - litellm.set_verbose = True - - # Create router with Mistral OCR deployment - router = Router( - model_list=[ - { - "model_name": "mistral-ocr", - "litellm_params": { - "model": "mistral/mistral-ocr-latest", - "api_key": os.getenv("MISTRAL_API_KEY"), - }, - } - ] - ) - - try: - # Call OCR through router - response = await router.aocr( - model="mistral-ocr", - document={"type": "document_url", "document_url": TEST_PDF_URL}, - ) - - print(f"\n{'='*80}") - print("Router OCR Test") - print(f"Response type: {type(response)}") - print( - f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}" - ) - - # Check if response has expected Mistral OCR format - assert hasattr(response, "pages"), "Response should have 'pages' attribute" - assert hasattr(response, "model"), "Response should have 'model' attribute" - assert hasattr(response, "object"), "Response should have 'object' attribute" - assert ( - response.object == "ocr" - ), f"Expected object='ocr', got '{response.object}'" - - # Validate pages structure - assert isinstance(response.pages, list), "pages should be a list" - assert len(response.pages) > 0, "Should have at least one page" - - # Check first page structure - first_page = response.pages[0] - assert hasattr(first_page, "index"), "Page should have 'index' attribute" - assert hasattr(first_page, "markdown"), "Page should have 'markdown' attribute" - - # Extract text from all pages for validation - total_text = "\n\n".join( - page.markdown for page in response.pages if page.markdown - ) - print(f"Total pages: {len(response.pages)}") - print(f"Total extracted text length: {len(total_text)} characters") - print(f"First 200 chars: {total_text[:200]}") - print(f"Model: {response.model}") - if response.usage_info: - print(f"Pages processed: {response.usage_info.pages_processed}") - print(f"{'='*80}\n") - - assert len(total_text) > 0, "Should extract some text from the document" - - except Exception as e: - pytest.fail(f"Router OCR call failed: {str(e)}") diff --git a/tests/ocr_tests/test_ocr_vertex_ai.py b/tests/ocr_tests/test_ocr_vertex_ai.py index 1842eb063a5..beddc9cd35e 100644 --- a/tests/ocr_tests/test_ocr_vertex_ai.py +++ b/tests/ocr_tests/test_ocr_vertex_ai.py @@ -1,117 +1,8 @@ -""" -Test OCR functionality with Vertex AI OCR APIs (Mistral and DeepSeek). +"""Vertex AI OCR config routing and DeepSeek request shaping (no network).""" -Note: Vertex AI OCR automatically converts URLs to base64 data URIs since -the Vertex AI endpoint doesn't have internet access. -""" - -import json -import os -import tempfile from typing import Final import pytest -from base_ocr_unit_tests import BaseOCRTest - - -def load_vertex_ai_credentials(): - """Load Vertex AI credentials for tests""" - # Define the path to the vertex_key.json file - print("loading vertex ai credentials") - filepath = os.path.dirname(os.path.abspath(__file__)) - vertex_key_path = filepath + "/vertex_key.json" - - # Read the existing content of the file or create an empty dictionary - try: - with open(vertex_key_path, "r") as file: - # Read the file content - print("Read vertexai file path") - content = file.read() - - # If the file is empty or not valid JSON, create an empty dictionary - if not content or not content.strip(): - service_account_key_data = {} - else: - # Attempt to load the existing JSON content - file.seek(0) - service_account_key_data = json.load(file) - except FileNotFoundError: - # If the file doesn't exist, create an empty dictionary - service_account_key_data = {} - - # Update the service_account_key_data with environment variables - private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") - private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") - private_key = private_key.replace("\\n", "\n") - service_account_key_data["private_key_id"] = private_key_id - service_account_key_data["private_key"] = private_key - - # Create a temporary file - with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file: - # Write the updated content to the temporary files - json.dump(service_account_key_data, temp_file, indent=2) - - # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS - os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) - - -class TestVertexAIMistralOCR(BaseOCRTest): - """ - Test class for Vertex AI Mistral OCR functionality. - Inherits from BaseOCRTest and provides Vertex AI-specific configuration. - - Note: For Vertex AI, LiteLLM will automatically convert URLs to base64 data URIs before - sending to the API, since Vertex AI OCR endpoint doesn't have internet access. - """ - - def setup_method(self): - if os.environ.get("LITELLM_RUN_LIVE_VERTEX_MISTRAL_OCR_TESTS") != "1": - pytest.skip("Live Vertex AI Mistral OCR E2E tests are opt-in") - if os.environ.get("CASSETTE_REDIS_URL"): - pytest.skip( - "Live Vertex AI Mistral OCR E2E tests cannot run under VCR replay" - ) - - def get_base_ocr_call_args(self) -> dict: - """ - Return the base OCR call args for Vertex AI Mistral OCR. - """ - load_vertex_ai_credentials() - return { - "model": "vertex_ai/mistral-ocr-2505", - "vertex_location": "us-central1", - } - - -class TestVertexAIDeepSeekOCR(BaseOCRTest): - """ - Test class for Vertex AI DeepSeek OCR functionality. - Inherits from BaseOCRTest and provides Vertex AI-specific configuration. - - Note: DeepSeek OCR uses the chat completion API format through the openapi endpoint. - Note: DeepSeek OCR does not support PDF URLs - only image URLs and base64 data. - """ - - def get_base_ocr_call_args(self) -> dict: - """ - Return the base OCR call args for Vertex AI DeepSeek OCR. - """ - load_vertex_ai_credentials() - return { - "model": "vertex_ai/deepseek-ocr-maas", - "vertex_location": "us-central1", - } - - # Skip PDF URL tests for DeepSeek OCR as it doesn't support PDF URLs - @pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs") - async def test_basic_ocr_with_url(self, sync_mode): - """Skip this test for DeepSeek OCR - PDF URLs not supported""" - pass - - @pytest.mark.skip(reason="DeepSeek OCR does not support PDF URLs") - def test_ocr_response_structure(self): - """Skip this test for DeepSeek OCR - PDF URLs not supported""" - pass def test_vertex_ai_ocr_routing(): @@ -126,21 +17,19 @@ def test_vertex_ai_ocr_routing(): # Test DeepSeek OCR routing deepseek_config = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas") - assert isinstance( - deepseek_config, VertexAIDeepSeekOCRConfig - ), "DeepSeek model should route to VertexAIDeepSeekOCRConfig" + assert isinstance(deepseek_config, VertexAIDeepSeekOCRConfig), ( + "DeepSeek model should route to VertexAIDeepSeekOCRConfig" + ) # Test Mistral OCR routing (should use default VertexAIOCRConfig) mistral_config = get_vertex_ai_ocr_config("vertex_ai/mistral-ocr-2505") - assert isinstance( - mistral_config, VertexAIOCRConfig - ), "Mistral model should route to VertexAIOCRConfig" + assert isinstance(mistral_config, VertexAIOCRConfig), "Mistral model should route to VertexAIOCRConfig" # Test other DeepSeek variants deepseek_variant = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas") - assert isinstance( - deepseek_variant, VertexAIDeepSeekOCRConfig - ), "DeepSeek variant should route to VertexAIDeepSeekOCRConfig" + assert isinstance(deepseek_variant, VertexAIDeepSeekOCRConfig), ( + "DeepSeek variant should route to VertexAIDeepSeekOCRConfig" + ) @pytest.mark.parametrize("model", ("deepseek-ocr-maas", "deepseek-ai/deepseek-ocr-maas")) diff --git a/tests/ocr_tests/vertex_key.json b/tests/ocr_tests/vertex_key.json deleted file mode 100644 index 800969fb305..00000000000 --- a/tests/ocr_tests/vertex_key.json +++ /dev/null @@ -1,13 +0,0 @@ -{ - "type": "service_account", - "project_id": "litellm-ci-cd", - "private_key_id": "", - "private_key": "", - "client_email": "test-litellm-ci-cd@litellm-ci-cd.iam.gserviceaccount.com", - "client_id": "116563532503305622785", - "auth_uri": "https://accounts.google.com/o/oauth2/auth", - "token_uri": "https://oauth2.googleapis.com/token", - "auth_provider_x509_cert_url": "https://www.googleapis.com/oauth2/v1/certs", - "client_x509_cert_url": "https://www.googleapis.com/robot/v1/metadata/x509/test-litellm-ci-cd%40litellm-ci-cd.iam.gserviceaccount.com", - "universe_domain": "googleapis.com" -} From 0b5b69ea3ae4aa2c8aeb5764240046c85713d5a0 Mon Sep 17 00:00:00 2001 From: yassin Date: Fri, 18 Sep 2026 21:29:29 +0000 Subject: [PATCH 220/253] fix(deepgram): forward only the first model and language values to /listen Authorization and pricing read the first model and language query value, but the raw query was forwarded, so Deepgram (which honours the last repeated value) could be sent a model the key was never allowed. Later duplicates of those two keys are now dropped before the upstream URL is built Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/deepgram/common_utils.py | 21 ++++++++--- .../deepgram/test_deepgram_common_utils.py | 35 +++++++++++++++++-- .../test_deepgram_ws_passthrough_routes.py | 30 ++++++++++++++++ 3 files changed, 78 insertions(+), 8 deletions(-) diff --git a/litellm/llms/deepgram/common_utils.py b/litellm/llms/deepgram/common_utils.py index 676391dc744..9b071ac8321 100644 --- a/litellm/llms/deepgram/common_utils.py +++ b/litellm/llms/deepgram/common_utils.py @@ -26,6 +26,7 @@ DEEPGRAM_LISTEN_ADDON_PRICING_PARAMS: Final = MappingProxyType( } ) _DISABLED_PARAM_VALUES: Final = frozenset({"", "false"}) +_SINGLE_VALUED_PARAMS: Final = frozenset({"model", "language"}) class DeepgramException(BaseLLMException): @@ -36,13 +37,23 @@ def deepgram_listen_requested_model(query_string: str) -> str: return httpx.QueryParams(query_string).get("model") or DEEPGRAM_LISTEN_DEFAULT_MODEL +def _first_occurrences(query_string: str) -> httpx.QueryParams: + """Authorization and pricing read the first ``model`` and ``language`` value; Deepgram must not see a second one.""" + items: Final = httpx.QueryParams(query_string).multi_items() + return httpx.QueryParams( + tuple( + (key, value) + for index, (key, value) in enumerate(items) + if key not in _SINGLE_VALUED_PARAMS or all(earlier != key for earlier, _ in items[:index]) + ) + ) + + def deepgram_listen_websocket_target(api_base: str | None, query_string: str) -> str: listen_url: Final = httpx.URL(f"{(api_base or DEEPGRAM_DEFAULT_API_BASE).rstrip('/')}/listen") websocket_url: Final = listen_url.copy_with(scheme=_WEBSOCKET_SCHEMES.get(listen_url.scheme, listen_url.scheme)) - params: Final = httpx.QueryParams(query_string) - query: Final = ( - query_string if params.get("model") else str(params.remove("model").add("model", DEEPGRAM_LISTEN_DEFAULT_MODEL)) - ) + params: Final = _first_occurrences(query_string) + query: Final = params if params.get("model") else params.remove("model").add("model", DEEPGRAM_LISTEN_DEFAULT_MODEL) return f"{websocket_url}?{query}" @@ -64,7 +75,7 @@ def deepgram_listen_pricing_model(upstream_url: str) -> str: the multilingual streaming entry when ``language=multi``, otherwise the model's own streaming entry. Pre-recorded entries are never a substitute: Deepgram prices the two products differently.""" streaming: Final = f"{DEEPGRAM_LISTEN_STREAMING_PRICING_PREFIX}{deepgram_listen_model(upstream_url)}" - language: Final = parse_qs(urlparse(upstream_url).query).get("language", ("",))[-1] + language: Final = parse_qs(urlparse(upstream_url).query).get("language", ("",))[0] if language.strip().lower() == DEEPGRAM_LISTEN_MULTILINGUAL_LANGUAGE: return f"{streaming}{DEEPGRAM_LISTEN_MULTILINGUAL_PRICING_SUFFIX}" return streaming diff --git a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py index a1fa8f26b70..530888b70c4 100644 --- a/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py +++ b/tests/test_litellm/llms/deepgram/test_deepgram_common_utils.py @@ -1,6 +1,7 @@ import math from collections.abc import Mapping, Sequence from typing import Final +from urllib.parse import parse_qs, urlparse import pytest @@ -76,6 +77,24 @@ def _metadata(duration: object, channels: object = 1) -> dict[str, object]: "wss://dg.internal/v1/listen?model=nova-3&keywords=a&keywords=b", id="repeated keys preserved", ), + pytest.param( + None, + "model=nova-2&encoding=linear16&model=nova-3", + "wss://api.deepgram.com/v1/listen?model=nova-2&encoding=linear16", + id="only the authorized first model reaches deepgram", + ), + pytest.param( + None, + "language=en&model=nova-3&language=multi", + "wss://api.deepgram.com/v1/listen?language=en&model=nova-3", + id="only the priced first language reaches deepgram", + ), + pytest.param( + None, + "model=&model=nova-2", + "wss://api.deepgram.com/v1/listen?model=nova-3", + id="blank first model is the default, later models dropped", + ), ], ) def test_deepgram_listen_websocket_target(api_base: str | None, query_string: str, expected: str): @@ -210,12 +229,22 @@ def test_deepgram_listen_model_comes_from_the_upstream_query(upstream_url: str, @pytest.mark.parametrize( "query_string", - ["model=nova-2&language=en", "language=en", "model=&language=en", "", "model=nova-3-medical"], + [ + "model=nova-2&language=en", + "language=en", + "model=&language=en", + "", + "model=nova-3-medical", + "model=nova-2&model=nova-3", + "model=&model=nova-3-medical", + ], ) -def test_requested_model_is_the_model_the_upstream_target_will_carry(query_string: str): +def test_requested_model_is_the_only_model_the_upstream_target_carries(query_string: str): """Authorization runs against ``deepgram_listen_requested_model``; the upstream URL is built separately, so the - two must always agree or a key could be authorized for one model and reach another.""" + two must always agree or a key could be authorized for one model and reach another. Deepgram reads the last + repeated ``model``, so the target must carry exactly one.""" target: Final = deepgram_listen_websocket_target(None, query_string) + assert parse_qs(urlparse(target).query)["model"] == [deepgram_listen_requested_model(query_string)] assert deepgram_listen_requested_model(query_string) == deepgram_listen_model(target) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py index 4eb183b14ce..44533f35c72 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_deepgram_ws_passthrough_routes.py @@ -421,6 +421,36 @@ def test_deepgram_listen_authorizes_the_model_it_will_actually_send_upstream(que assert relay.calls == [] +def test_deepgram_listen_strips_a_second_model_that_would_outrank_the_authorized_one(monkeypatch): + """Deepgram honours the last repeated ``model``; auth and pricing read the first. A key allowed only ``nova-2`` + must not smuggle ``nova-3`` past authorization behind an authorized first value.""" + monkeypatch.delenv("DEEPGRAM_API_BASE", raising=False) + monkeypatch.setattr(litellm, "max_budget", 0.0) + _price_nova_2_streaming(monkeypatch) + cache = asyncio.run(_cache_restricted_key("sk-only-nova-2", ["nova-2"])) + relay = _FakeRelay() + client = TestClient(_app_with_relay(relay)) + + with ( + patch(GET_CREDENTIALS, return_value="dg-provider-key"), + patch.multiple( # test-quality-ok: the real key auth path reads these proxy_server globals and has no injection seam + "litellm.proxy.proxy_server", + master_key="sk-master", + prisma_client=MagicMock(), + user_api_key_cache=cache, + llm_model_list=None, + llm_router=None, + ), + ): + with client.websocket_connect( + "/deepgram/v1/listen?model=nova-2&language=en&model=nova-3&language=multi", + headers={"Authorization": "Bearer sk-only-nova-2"}, + ): + pass + + assert [call.target for call in relay.calls] == ["wss://api.deepgram.com/v1/listen?model=nova-2&language=en"] + + def test_deepgram_listen_echoes_the_browser_subprotocol_that_carries_the_litellm_key(monkeypatch): """Browsers cannot set headers, so they send the key as a subprotocol and abort the handshake unless the server echoes that subprotocol back; the key itself must still stay off the upstream connection.""" From 595768b54b7ccfefdee9234c0bb70fbe219ee255 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 18 Sep 2026 14:33:41 -0700 Subject: [PATCH 221/253] fix(proxy): narrow project_id without a cast and fold the unbudgeted cases into the budget matrix test The lint job failed on one new typing.cast (LIT006) in the cost callback, and the test-quality gate behind it would have failed next on a test whose only assertion inspected a mock (TQ002). project_id is now narrowed with isinstance, and the zero and negative max_budget cases run through the existing parametrized budget test, which asserts the raised error or a clean admit with no alert --- .../proxy/hooks/proxy_track_cost_callback.py | 6 ++- .../proxy/auth/test_auth_checks.py | 45 ++++++------------- 2 files changed, 19 insertions(+), 32 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 5e525108ade..903255c7b6c 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -319,7 +319,11 @@ class _ProxyDBLogger(CustomLogger): user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None)) org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None)) - project_id: Final = cast(str | None, metadata.get("user_api_key_project_id", None)) + project_id: Final = ( + project_id_value + if isinstance(project_id_value := metadata.get("user_api_key_project_id"), str) + else None + ) key_alias: Final = cast(str | None, metadata.get("user_api_key_alias", None)) end_user_max_budget: Final = metadata.get("user_api_end_user_max_budget", None) sl_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index beadaa27831..85673df57ba 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7559,15 +7559,19 @@ def _project_with_budget(spend: float, max_budget: float): @pytest.mark.asyncio @pytest.mark.parametrize( - "counter_spend, db_spend, blocks", + "counter_spend, db_spend, max_budget, blocks", [ - pytest.param(5.0, 0.0, True, id="counter-at-budget-blocks-despite-stale-db-row"), - pytest.param(4.99, 0.0, False, id="counter-under-budget-admits"), - pytest.param(None, 5.0, True, id="no-counter-falls-back-to-persisted-spend"), - pytest.param(None, 0.0, False, id="no-counter-and-no-persisted-spend-admits"), + pytest.param(5.0, 0.0, 5.0, True, id="counter-at-budget-blocks-despite-stale-db-row"), + pytest.param(4.99, 0.0, 5.0, False, id="counter-under-budget-admits"), + pytest.param(None, 5.0, 5.0, True, id="no-counter-falls-back-to-persisted-spend"), + pytest.param(None, 0.0, 5.0, False, id="no-counter-and-no-persisted-spend-admits"), + pytest.param(12.5, 12.5, 0.0, False, id="zero-budget-is-unbudgeted"), + pytest.param(12.5, 12.5, -1.0, False, id="negative-budget-is-unbudgeted"), ], ) -async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, db_spend, blocks): +async def test_project_max_budget_check_blocks_only_when_live_spend_reaches_a_positive_budget( + counter_spend, db_spend, max_budget, blocks +): from litellm.caching.dual_cache import DualCache from litellm.proxy.auth.auth_checks import _project_max_budget_check @@ -7583,14 +7587,16 @@ async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, ): if not blocks: await _project_max_budget_check( - project_object=_project_with_budget(spend=db_spend, max_budget=5.0), + project_object=_project_with_budget(spend=db_spend, max_budget=max_budget), valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) + await asyncio.sleep(0) + proxy_logging_obj.budget_alerts.assert_not_awaited() return with pytest.raises(litellm.BudgetExceededError) as exc_info: await _project_max_budget_check( - project_object=_project_with_budget(spend=db_spend, max_budget=5.0), + project_object=_project_with_budget(spend=db_spend, max_budget=max_budget), valid_token=valid_token, proxy_logging_obj=proxy_logging_obj, ) @@ -7603,29 +7609,6 @@ async def test_project_max_budget_check_reads_live_spend_counter(counter_spend, assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "project_budget" -@pytest.mark.asyncio -@pytest.mark.parametrize("max_budget", [0.0, -1.0]) -async def test_project_max_budget_check_treats_non_positive_budget_as_unbudgeted(max_budget): - from litellm.caching.dual_cache import DualCache - from litellm.proxy.auth.auth_checks import _project_max_budget_check - - real_spend_counter_cache = DualCache() - real_spend_counter_cache.in_memory_cache.set_cache(key="spend:project:p-budget", value=12.5) - proxy_logging_obj = MagicMock() - proxy_logging_obj.budget_alerts = AsyncMock() - - with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock - "litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache - ): - await _project_max_budget_check( - project_object=_project_with_budget(spend=12.5, max_budget=max_budget), - valid_token=UserAPIKeyAuth(api_key="hashed-key", project_id="p-budget"), - proxy_logging_obj=proxy_logging_obj, - ) - - proxy_logging_obj.budget_alerts.assert_not_awaited() - - def test_is_user_proxy_admin_rejects_view_only_admin(): """This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an Admin Viewer answering True here would gain every write route. Read parity for From f72b7155acf0b17e61f10733fb897f1f5defb832 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 21:40:13 +0000 Subject: [PATCH 222/253] fix(ocr): map Rust upstream 401/403 to the public auth exceptions The httpx.Response built for a Rust upstream failure had no request attached, so constructing openai.AuthenticationError raised RuntimeError inside the exception mapper and every bad-key OCR call surfaced as APIConnectionError 500 instead of AuthenticationError 401 (the Python path already returned 401) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/rust_bridge/ocr/route_host.py | 10 +++++++--- tests/test_litellm/rust_bridge/ocr/test_route_host.py | 11 +++++++++++ 2 files changed, 18 insertions(+), 3 deletions(-) diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index 0bc7b383eea..277fdceb734 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -27,13 +27,17 @@ class UpstreamFailure(Exception): self.__cause__ = cause -def _upstream_failure(error: Exception) -> Exception: +def _upstream_failure(error: Exception, request: LiteLLMOcrRequest) -> Exception: try: status, body = _UPSTREAM_ARGS.validate_python(error.args) headers: Final = _UPSTREAM_HEADERS.validate_python(getattr(error, "headers", None)) except ValidationError: return error - return UpstreamFailure(httpx.Response(status, content=body.encode(), headers=headers), error) + http_request: Final = httpx.Request("POST", request.api_base or "https://docs.litellm.ai/docs") + return UpstreamFailure( + httpx.Response(status, content=body.encode(), headers=headers, request=http_request), + error, + ) def response(value: Mapping[str, object]) -> OCRResponse: @@ -57,7 +61,7 @@ def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: model=request.model.removeprefix(f"{request_provider}/"), llm_provider=request_provider, ) - original: Final = _upstream_failure(error) + original: Final = _upstream_failure(error, request) public_error: Final = failures.map_failure(original, request.model, request_provider, arguments(request)) if isinstance(original, UpstreamFailure) and public_error.__context__ is original: public_error.__context__ = error diff --git a/tests/test_litellm/rust_bridge/ocr/test_route_host.py b/tests/test_litellm/rust_bridge/ocr/test_route_host.py index a328579400c..699492e4424 100644 --- a/tests/test_litellm/rust_bridge/ocr/test_route_host.py +++ b/tests/test_litellm/rust_bridge/ocr/test_route_host.py @@ -59,6 +59,17 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N assert public_error.llm_provider == "mistral" +def test_map_failure_maps_upstream_401_to_authentication_error() -> None: + error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ()) + + public_error: Final = map_failure(error, REQUEST, "mistral") + + assert isinstance(public_error, litellm.AuthenticationError) + assert public_error.status_code == 401 + assert public_error.response.text == '{"message": "Unauthorized"}' + assert public_error.__context__ is error + + def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: error: Final = RuntimeError("bridge exploded") From cf05466a27ad9409d1edcd9697ddb826ca8f93da Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:54:18 -0700 Subject: [PATCH 223/253] fix(gemini): map every documented finishReason and reset per-candidate state A content-less candidate is now kept as a choice whenever it carries a finishReason, with the raw value on the choice's provider_specific_fields. NO_IMAGE, IMAGE_RECITATION, IMAGE_OTHER and ESCALATION map to content_filter; UNEXPECTED_TOOL_CALL and MISSING_THOUGHT_SIGNATURE map to stop. The /v1/responses bridge reports content_filter and refusal as incomplete with incomplete_details, and tool calls and reasoning no longer leak from one candidate into the next. --- litellm/litellm_core_utils/core_helpers.py | 5 ++ .../vertex_and_google_ai_studio_gemini.py | 29 ++++---- .../transformation.py | 33 ++++++--- .../litellm_core_utils/test_core_helpers.py | 6 ++ ...test_vertex_and_google_ai_studio_gemini.py | 68 +++++++++++++++++++ .../test_litellm_completion_responses.py | 2 +- 6 files changed, 117 insertions(+), 26 deletions(-) diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index 2615c58de2c..d29b1fc74ef 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -225,6 +225,11 @@ _FINISH_REASON_MAP: Final[dict[str, OpenAIChatCompletionFinishReason]] = { "TOO_MANY_TOOL_CALLS": "stop", "MALFORMED_RESPONSE": "stop", "NO_IMAGE": "content_filter", + "IMAGE_RECITATION": "content_filter", + "IMAGE_OTHER": "content_filter", + "ESCALATION": "content_filter", + "UNEXPECTED_TOOL_CALL": "stop", + "MISSING_THOUGHT_SIGNATURE": "stop", # Zhipu GLM "network_error": "stop", "sensitive": "content_filter", 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 61bac4793a8..a95b845718a 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 @@ -1348,6 +1348,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "TOO_MANY_TOOL_CALLS", "MALFORMED_RESPONSE", "NO_IMAGE", + "IMAGE_RECITATION", + "IMAGE_OTHER", + "ESCALATION", + "UNEXPECTED_TOOL_CALL", + "MISSING_THOUGHT_SIGNATURE", } ) @@ -2243,7 +2248,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): image_response: list[ImageURLListItem] | None = None chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} chat_completion_logprobs: ChoiceLogprobs | None = None - tools: list[ChatCompletionToolCallChunk] | None = [] + tools: list[ChatCompletionToolCallChunk] | None = None functions: ChatCompletionToolCallFunctionChunk | None = None thinking_blocks: list[ChatCompletionThinkingBlock] | None = None reasoning_content: str | None = None @@ -2358,11 +2363,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): tool_invocation_fields["server_side_tool_invocations"] = server_side_tool_invocations chat_completion_message["provider_specific_fields"] = tool_invocation_fields - if candidate.get("finishReason"): - finish_reason_fields = chat_completion_message.get("provider_specific_fields") or {} - finish_reason_fields["native_finish_reason"] = candidate.get("finishReason") - chat_completion_message["provider_specific_fields"] = finish_reason_fields - if isinstance(model_response, ModelResponseStream): choice = VertexGeminiConfig._create_streaming_choice( chat_completion_message=chat_completion_message, @@ -2375,15 +2375,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ) model_response.choices.append(choice) elif isinstance(model_response, ModelResponse): + native_finish_reason = candidate.get("finishReason") choice = litellm.Choices( finish_reason=VertexGeminiConfig._check_finish_reason( - chat_completion_message, candidate.get("finishReason") + chat_completion_message, native_finish_reason ), index=candidate.get("index", idx), message=chat_completion_message, logprobs=chat_completion_logprobs, enhancements=None, - provider_specific_fields=chat_completion_message.get("provider_specific_fields"), + provider_specific_fields=( + {"native_finish_reason": native_finish_reason} if native_finish_reason is not None else None + ), ) model_response.choices.append(choice) @@ -3181,12 +3184,10 @@ class ModelResponseIterator: self.has_seen_tool_calls = True break - # _process_candidates skips candidates without a "content" part, so a - # content-less chunk leaves choices empty and the downstream streaming - # handler hits IndexError on choices[0]. This covers the final chunk - # (finishReason, no content) and mid-stream metadata-only chunks - # (grounding/web-search/thought, no content and no finishReason — seen - # with web_search + reasoning) by emitting an empty-delta choice. + # _process_candidates skips candidates with neither "content" nor + # "finishReason", so a metadata-only chunk (grounding/web-search/thought, + # seen with web_search + reasoning) leaves choices empty and the downstream + # streaming handler hits IndexError on choices[0]. Emit an empty-delta choice. if not model_response.choices and _candidates: from litellm.types.utils import Delta, StreamingChoices diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 7792bc11405..961fffd3d42 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -61,6 +61,7 @@ from litellm.types.llms.openai import ( ChatCompletionToolParamFunctionChunk, ChatCompletionUserMessage, GenericChatCompletionMessage, + IncompleteDetails, InputTokensDetails, OpenAIChatCompletionTextObject, OpenAIMcpServerTool, @@ -2295,6 +2296,21 @@ class LiteLLMCompletionResponsesConfig: # Default to completed for unknown finish reasons return "completed" + @staticmethod + def _incomplete_details_for_finish_reason( + finish_reason: str | None, + existing: IncompleteDetails | None, + ) -> IncompleteDetails | None: + if existing is not None: + return existing + match finish_reason: + case "length": + return IncompleteDetails(reason="max_output_tokens") + case "content_filter" | "refusal": + return IncompleteDetails(reason="content_filter") + case _: + return None + @staticmethod def _tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str: """Bedrock Mantle returns a non-unique, index-based ``call_id`` (``call_0``, @@ -2411,17 +2427,10 @@ class LiteLLMCompletionResponsesConfig: if choices and len(choices) > 0: finish_reason = choices[0].finish_reason - status: Final[ResponsesAPIStatus] = ( - LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status(finish_reason) + incomplete_details: Final = LiteLLMCompletionResponsesConfig._incomplete_details_for_finish_reason( + finish_reason=finish_reason, + existing=getattr(chat_completion_response, "incomplete_details", None), ) - incomplete_details = getattr(chat_completion_response, "incomplete_details", None) - if incomplete_details is None and status == "incomplete": - from openai.types.responses.response import IncompleteDetails - - if finish_reason == "length": - incomplete_details = IncompleteDetails(reason="max_output_tokens") - elif finish_reason in ["content_filter", "refusal"]: - incomplete_details = IncompleteDetails(reason="content_filter") responses_api_response: Final[ResponsesAPIResponse] = ResponsesAPIResponse( id=chat_completion_response.id, @@ -2447,7 +2456,9 @@ class LiteLLMCompletionResponsesConfig: max_output_tokens=getattr(chat_completion_response, "max_output_tokens", None), previous_response_id=getattr(chat_completion_response, "previous_response_id", None), reasoning=None, - status=status, + status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status( + finish_reason + ), text={}, truncation=getattr(chat_completion_response, "truncation", None), usage=LiteLLMCompletionResponsesConfig._transform_chat_completion_usage_to_responses_usage( diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index e937be47441..b2ad13c205e 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -151,6 +151,12 @@ class TestMapFinishReasonGemini: ("IMAGE_PROHIBITED_CONTENT", "content_filter"), ("TOO_MANY_TOOL_CALLS", "stop"), ("MALFORMED_RESPONSE", "stop"), + ("NO_IMAGE", "content_filter"), + ("IMAGE_RECITATION", "content_filter"), + ("IMAGE_OTHER", "content_filter"), + ("ESCALATION", "content_filter"), + ("UNEXPECTED_TOOL_CALL", "stop"), + ("MISSING_THOUGHT_SIGNATURE", "stop"), ], ) def test_gemini_finish_reasons(self, gemini_reason, expected): 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 216d6f0db7e..4ee199bd9af 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 @@ -968,6 +968,12 @@ def test_finish_reason_unspecified_and_malformed_function_call(): # Test new Gemini finish reasons assert finish_reason_mappings["TOO_MANY_TOOL_CALLS"] == "stop" assert finish_reason_mappings["MALFORMED_RESPONSE"] == "stop" + assert finish_reason_mappings["NO_IMAGE"] == "content_filter" + assert finish_reason_mappings["IMAGE_RECITATION"] == "content_filter" + assert finish_reason_mappings["IMAGE_OTHER"] == "content_filter" + assert finish_reason_mappings["ESCALATION"] == "content_filter" + assert finish_reason_mappings["UNEXPECTED_TOOL_CALL"] == "stop" + assert finish_reason_mappings["MISSING_THOUGHT_SIGNATURE"] == "stop" def test_vertex_ai_usage_metadata_response_token_count(): @@ -6219,3 +6225,65 @@ def test_gemini_candidate_other_finish_reasons_no_content(): ) assert responses_length.status == "incomplete" assert responses_length.incomplete_details.reason == "max_output_tokens" + + +def test_gemini_candidate_with_finish_reason_no_content_streaming_chunk(): + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + chunk: Final = { + "candidates": [{"finishReason": "NO_IMAGE", "index": 0}], + "usageMetadata": {"promptTokenCount": 19, "candidatesTokenCount": 0, "totalTokenCount": 19}, + } + iterator: Final = ModelResponseIterator(streaming_response=[], sync_stream=True, logging_obj=MagicMock()) + + streaming_chunk: Final = iterator.chunk_parser(chunk) + + assert len(streaming_chunk.choices) == 1 + assert streaming_chunk.choices[0].finish_reason == "content_filter" + assert streaming_chunk.choices[0].delta.content is None + assert streaming_chunk.choices[0].delta.tool_calls is None + + +def test_gemini_multi_candidate_messages_do_not_share_state(): + config: Final = VertexGeminiConfig() + completion_response: Final = { + "candidates": [ + { + "content": { + "role": "model", + "parts": [ + {"text": "Let me check the weather.", "thought": True}, + {"functionCall": {"name": "get_weather", "args": {"city": "Paris"}}}, + ], + }, + "finishReason": "STOP", + "index": 0, + }, + { + "content": {"role": "model", "parts": [{"text": "It is sunny in Paris."}]}, + "finishReason": "STOP", + "index": 1, + }, + ], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20, "totalTokenCount": 30}, + } + + resp: Final = config._transform_google_generate_content_to_openai_model_response( + completion_response=completion_response, + model_response=ModelResponse(), + model="gemini-2.5-flash", + logging_obj=MagicMock(), + raw_response=MagicMock(headers={}), + ) + + assert len(resp.choices) == 2 + assert resp.choices[0].finish_reason == "tool_calls" + assert resp.choices[0].message.tool_calls[0].function.name == "get_weather" + assert resp.choices[0].message.reasoning_content == "Let me check the weather." + assert resp.choices[1].finish_reason == "stop" + assert resp.choices[1].message.content == "It is sunny in Paris." + assert resp.choices[1].message.tool_calls is None + assert getattr(resp.choices[1].message, "reasoning_content", None) is None + assert resp.choices[1].provider_specific_fields["native_finish_reason"] == "STOP" diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 849c997838d..52d06acfd64 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -4909,7 +4909,7 @@ class TestStreamingSnapshotItemIds: def test_transform_chat_completion_response_incomplete_details(): - from openai.types.responses.response import IncompleteDetails + from litellm.types.llms.openai import IncompleteDetails resp_length = ModelResponse( id="resp-length", From 3edbf60e9c1d3070ffe6e1e70bea39c688f39a21 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:58:34 -0700 Subject: [PATCH 224/253] fix(proxy): requeue the daily tag rollup on commit failure without the Redis buffer --- litellm/proxy/db/db_spend_update_writer.py | 20 +++++-------- .../proxy/db/test_db_spend_update_writer.py | 29 +++++++++++++++++++ 2 files changed, 37 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index a4f1713d61d..b2e6f9dc54d 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1305,7 +1305,7 @@ class DBSpendUpdateWriter: async def _flush_daily_spend_queue( self, queue: DailySpendUpdateQueue, - entity_type: Literal["user", "team", "org", "end_user", "agent"], + entity_type: Literal["user", "team", "org", "tag", "end_user", "agent"], commit: _DailySpendCommit[_DailySpendTransactionT], n_retry_times: int, prisma_client: PrismaClient, @@ -1447,19 +1447,15 @@ class DBSpendUpdateWriter: Commit only tag spend updates to database. This is called by a separate scheduler job at a longer interval. """ - daily_tag_spend_update_transactions: Final = cast( - dict[str, DailyTagSpendTransaction], - await self.daily_tag_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(), + await self._flush_daily_spend_queue( + queue=self.daily_tag_spend_update_queue, + entity_type="tag", + commit=DBSpendUpdateWriter.update_daily_tag_spend, + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, ) - if daily_tag_spend_update_transactions: - await DBSpendUpdateWriter.update_daily_tag_spend( - n_retry_times=n_retry_times, - prisma_client=prisma_client, - proxy_logging_obj=proxy_logging_obj, - daily_spend_transactions=daily_tag_spend_update_transactions, - ) - async def _commit_daily_tag_spend_to_db_with_redis( self, prisma_client: PrismaClient, diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 3feb117edb7..b27b838133b 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -2864,6 +2864,35 @@ async def test_failed_daily_spend_commit_requeues_the_rows_and_flushes_the_other assert db_writer.daily_spend_update_queue.update_queue.empty() +@pytest.mark.asyncio +async def test_failed_daily_tag_spend_commit_requeues_the_rows(): + """The tag rollup drains on its own scheduler job with the same no-Redis drop: + a failed LiteLLM_DailyTagSpend commit has to put the rows back for the next tick.""" + db_writer = DBSpendUpdateWriter() + tag_txn = {key: value for key, value in _daily_txn().items() if key != "user_id"} | {"tag": "tag-1"} + await db_writer.daily_tag_spend_update_queue.add_update({"tag-key": tag_txn}) + db = _DailySpendFakeDB(failing_table="LiteLLM_DailyTagSpend") + proxy_logging_obj = MagicMock() + proxy_logging_obj.failure_handler = AsyncMock() + + await db_writer._commit_daily_tag_spend_to_db( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + assert _daily_upserts(db, "LiteLLM_DailyTagSpend") == [] + assert not db_writer.daily_tag_spend_update_queue.update_queue.empty() + + db.failing_table = None + await db_writer._commit_daily_tag_spend_to_db( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + (tag_upsert,) = _daily_upserts(db, "LiteLLM_DailyTagSpend") + assert _row_values(tag_upsert, "tag") == ["tag-1"] + assert _row_values(tag_upsert, "spend") == [0.1] + assert db_writer.daily_tag_spend_update_queue.update_queue.empty() + + @pytest.mark.asyncio async def test_failed_window_spend_commit_from_redis_is_restored_to_redis(): """The Redis drain is destructive, so a failed window commit has to push From 55249c7128bca0c47641f561c30e96040e734d69 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 14:58:58 -0700 Subject: [PATCH 225/253] fix: set vertex gemma-4-26b-a4b-it-maas context window to 262144 --- litellm/model_prices_and_context_window_backup.json | 2 +- model_prices_and_context_window.json | 2 +- .../gemma/test_vertex_ai_gemma_global_endpoint.py | 12 +++++++++++- 3 files changed, 13 insertions(+), 3 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 377e30a23b7..ef4fce37f23 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -50721,7 +50721,7 @@ "vertex_ai/google/gemma-4-26b-a4b-it-maas": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-openai_models", - "max_input_tokens": 256000, + "max_input_tokens": 262144, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 377e30a23b7..ef4fce37f23 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -50721,7 +50721,7 @@ "vertex_ai/google/gemma-4-26b-a4b-it-maas": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-openai_models", - "max_input_tokens": 256000, + "max_input_tokens": 262144, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py index e9b58622a4b..9d08daa0a81 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py @@ -30,7 +30,7 @@ from litellm.types.llms.vertex_ai import VertexPartnerProvider _GEMMA_MODEL_COST_ENTRY = { "vertex_ai/google/gemma-4-26b-a4b-it-maas": { "litellm_provider": "vertex_ai-openai_models", - "max_input_tokens": 256000, + "max_input_tokens": 262144, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", @@ -180,6 +180,16 @@ class TestCreateVertexURLGemma: # --------------------------------------------------------------------------- +def test_gemma_maas_context_window_matches_google(monkeypatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + info = litellm.get_model_info("vertex_ai/google/gemma-4-26b-a4b-it-maas") + + assert info["max_input_tokens"] == 262144 + assert info["max_output_tokens"] == 128000 + + # --------------------------------------------------------------------------- # Integration tests: verify payloads reach the global OpenAI endpoint # From 19ef8e47a6dc3453b414b2efe556b9e2893c3e24 Mon Sep 17 00:00:00 2001 From: joshua Date: Fri, 18 Sep 2026 22:01:56 +0000 Subject: [PATCH 226/253] feat(ui): link MCP Servers page to the user's connected MCP servers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp-servers/_components/mcp_servers.test.tsx | 14 ++++++++++++++ .../mcp-servers/_components/mcp_servers.tsx | 11 +++++++++-- 2 files changed, 23 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 1394a923174..091b2f1403f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -62,6 +62,20 @@ describe("MCPServers", () => { expect(screen.getByText("MCP Servers")).toBeInTheDocument(); }); + it.each(["Admin", "Internal User"])("links a %s to their MCP connections page", async (userRole) => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + + render( + + + , + ); + + const myConnections = await screen.findByRole("link", { name: "My Connections" }); + expect(myConnections).toBeVisible(); + expect(myConnections).toHaveAttribute("href", "/ui/connect"); + }); + it("should render mocked MCP servers data in the table", async () => { // Mock MCP servers data const mockServers = [ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index aa0a031c55c..33b2f344bba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -1,7 +1,8 @@ import { isAdminRole, isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; -import { CircleHelp, Search } from "lucide-react"; +import { CircleHelp, Plug, Search } from "lucide-react"; +import Link from "next/link"; import { Badge } from "@/components/ui/badge"; -import { Button } from "@/components/ui/button"; +import { Button, buttonVariants } from "@/components/ui/button"; import { InputGroup, InputGroupAddon, InputGroupInput } from "@/components/ui/input-group"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; @@ -43,6 +44,8 @@ import MCPDiscovery from "./mcp_discovery"; import { ByokCredentialModal } from "@/components/mcp_tools/ByokCredentialModal"; import { getSecureItem } from "@/utils/secureStorage"; import { TOOLS_OAUTH_UI_STATE_KEY } from "@/hooks/mcpOAuthUtils"; +import { uiHref } from "@/utils/uiHref"; +import { cn } from "@/lib/cva.config"; import UserEnvVarsModal from "./UserEnvVarsModal"; import { listMCPUserEnvVarStatus } from "@/components/networking"; @@ -485,6 +488,10 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i

Configure and manage your MCP servers

+ + + My Connections + {isAdminRole(userRole) && ( <>

Configure and manage your MCP servers

-
+
My Connections From 417a88daedf6e5e4d2de62fab5a33f4ec5a6ae3f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:13:07 -0700 Subject: [PATCH 228/253] fix(responses): carry dict-valued reasoning_effort and keep the frame type on websocket defaults A deployment whose reasoning_effort is an object is copied through as reasoning the way the HTTP mapper does it instead of being dropped, and the relay re-asserts the response.create frame type after merging extra_body so a type key inside it can never replace it. The lazy OpenAPI snapshot goes back to main: the earlier regeneration came from a Python 3.14 interpreter dedenting docstrings, which CI on 3.12 rejects --- litellm/proxy/_lazy_openapi_snapshot.json | 2 +- litellm/responses/main.py | 22 +++++++----- litellm/responses/streaming_iterator.py | 2 +- .../test_responses_websocket_all_providers.py | 36 +++++++++++++++++++ 4 files changed, 51 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 2aa2cf15ac1..8a8d08c6887 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -19394,7 +19394,7 @@ } } }, - "description": "\nUnified rate-limit error.\n\nEvery rate-limit condition surfaced by litellm \u2014 whether it originated from\nan upstream LLM provider, a vendor batch endpoint, or one of litellm's own\nproxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\nmax-iterations, etc.) \u2014 is raised as an instance of this class.\n\nThe :attr:`category` attribute lets callers distinguish the source. See\n:class:`RateLimitErrorCategory` for the available values.\n" + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " }, "500": { "content": { diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 6064fa66d91..a5912bb42b1 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -2262,24 +2262,28 @@ def _build_litellm_metadata_for_ws(kwargs: dict) -> dict: return metadata -_EXTRA_BODY_ADAPTER: Final = TypeAdapter(dict[str, object] | None) +_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object] | None) + + +def _deployment_reasoning_default(kwargs: Mapping[str, object]) -> Reasoning | dict[str, object] | None: + if kwargs.get("reasoning") is not None: + return None + reasoning_effort: Final = kwargs.get("reasoning_effort") + if isinstance(reasoning_effort, str): + return LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) + return _JSON_OBJECT_ADAPTER.validate_python(reasoning_effort) if isinstance(reasoning_effort, Mapping) else None def _build_responses_websocket_request_defaults(kwargs: Mapping[str, object]) -> ResponsesWebSocketRequestDefaults: - reasoning_effort: Final = kwargs.get("reasoning_effort") - mapped_reasoning: Final = ( - LiteLLMResponsesTransformationHandler()._map_reasoning_effort(reasoning_effort) - if kwargs.get("reasoning") is None and isinstance(reasoning_effort, str) - else None - ) + default_reasoning: Final = _deployment_reasoning_default(kwargs) candidate_params: Final[dict[str, object]] = { **kwargs, - **({"reasoning": mapped_reasoning} if mapped_reasoning is not None else {}), + **({"reasoning": default_reasoning} if default_reasoning is not None else {}), } fill_missing: Final = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(candidate_params) return ResponsesWebSocketRequestDefaults( fill_missing=MappingProxyType(dict(fill_missing)), - overrides=MappingProxyType(_EXTRA_BODY_ADAPTER.validate_python(kwargs.get("extra_body")) or {}), + overrides=MappingProxyType(_JSON_OBJECT_ADAPTER.validate_python(kwargs.get("extra_body")) or {}), ) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 1c82df8c664..8f65552ec20 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -1883,7 +1883,7 @@ class ResponsesWebSocketStreaming: nested: Final = msg_obj.get("response") if _is_json_object(nested): return {**msg_obj, "response": self.request_defaults.merged_into(nested)} - return self.request_defaults.merged_into(msg_obj) + return {**self.request_defaults.merged_into(msg_obj), "type": msg_obj["type"]} async def _mask_response_create(self, message: str) -> str: """ diff --git a/tests/test_litellm/responses/test_responses_websocket_all_providers.py b/tests/test_litellm/responses/test_responses_websocket_all_providers.py index 918c4a33d40..b05904acdcb 100644 --- a/tests/test_litellm/responses/test_responses_websocket_all_providers.py +++ b/tests/test_litellm/responses/test_responses_websocket_all_providers.py @@ -1252,6 +1252,42 @@ class TestNativeWebSocketDeploymentDefaults: assert dict(defaults.fill_missing) == {"reasoning": {"effort": "low"}} assert dict(defaults.overrides) == {} + def test_builder_copies_dict_valued_reasoning_effort_like_the_http_path(self): + from litellm.responses.main import _build_responses_websocket_request_defaults + + defaults = _build_responses_websocket_request_defaults( + {"model": "gpt-5-pro", "reasoning_effort": {"effort": "xhigh", "summary": "auto"}} + ) + + assert dict(defaults.fill_missing) == {"reasoning": {"effort": "xhigh", "summary": "auto"}} + + @pytest.mark.asyncio + async def test_extra_body_type_key_never_replaces_the_frame_type(self): + from types import MappingProxyType + + from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults + + handler = _make_streaming( + authorized_model="gpt-5-pro", + request_defaults=ResponsesWebSocketRequestDefaults( + fill_missing=MappingProxyType({}), + overrides=MappingProxyType({"type": "session.update", "provider_default": "configured"}), + ), + ) + + forwarded = json.loads( + await handler._mask_response_create( + json.dumps({"type": "response.create", "model": "gpt-5-pro", "input": "hi"}) + ) + ) + + assert forwarded == { + "type": "response.create", + "model": "gpt-5-pro", + "input": "hi", + "provider_default": "configured", + } + @pytest.mark.asyncio async def test_flat_frame_gets_defaults_client_keys_win_extra_body_overrides(self): handler = _make_streaming(authorized_model="gpt-5-pro", request_defaults=_deployment_defaults()) From d057e82e6482ec5a75f7952db45cb5a7fe7aa973 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:13:51 -0700 Subject: [PATCH 229/253] test(proxy): assert stored login throttle limits never outrank the config file --- tests/test_litellm/proxy/test_proxy_server.py | 43 +++++++++---------- 1 file changed, 20 insertions(+), 23 deletions(-) diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 9309c318573..1b3a460c730 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -14060,32 +14060,29 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the @pytest.mark.asyncio -async def test_login_throttle_settings_are_not_hot_applied_from_the_database(): - """LIT-5285: a stored sign-in limit does not take effect on a live worker. - - _update_general_settings copies an allowlist of keys out of the DB row on every config - poll. Adding these to it would let a stored value outrank config.yaml without a restart, - so an operator locked out by a bad value could not fix it by editing YAML and restarting. - """ +async def test_login_throttle_limits_from_the_config_file_outrank_the_database(monkeypatch): import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ProxyConfig - original = dict(ps.general_settings) - try: - ps.general_settings.clear() - await ProxyConfig()._update_general_settings( - db_general_settings={ - "max_failed_login_attempts_per_source": 999, - "failed_login_window_seconds": 1, - "failed_login_block_seconds": 1, - } - ) - assert "max_failed_login_attempts_per_source" not in ps.general_settings - assert "failed_login_window_seconds" not in ps.general_settings - assert "failed_login_block_seconds" not in ps.general_settings - finally: - ps.general_settings.clear() - ps.general_settings.update(original) + monkeypatch.setattr( + ps, + "general_settings", + { + "max_failed_login_attempts_per_source": 10, + "failed_login_window_seconds": 60, + "failed_login_block_seconds": 300, + }, + ) + await ProxyConfig()._update_general_settings( + db_general_settings={ + "max_failed_login_attempts_per_source": 999, + "failed_login_window_seconds": 1, + "failed_login_block_seconds": 1, + } + ) + assert ps.general_settings.get("max_failed_login_attempts_per_source") == 10 + assert ps.general_settings.get("failed_login_window_seconds") == 60 + assert ps.general_settings.get("failed_login_block_seconds") == 300 @pytest.mark.asyncio From 1adbfbfbb11b534df9c4c75f5937c51ca2d2ec1b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:28:45 -0700 Subject: [PATCH 230/253] fix: strip eager_input_streaming for non-Claude providers next to input_examples --- litellm/main.py | 46 +++++++++++----------- tests/test_litellm/test_main.py | 70 +++++++++++++++++++++++++-------- 2 files changed, 76 insertions(+), 40 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 859e0df142c..ac8fa507728 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1146,37 +1146,35 @@ def responses_api_bridge_check( return model_info, model -def _should_allow_input_examples(custom_llm_provider: str | None, model: str) -> bool: +_ANTHROPIC_ONLY_TOOL_KEYS: Final = frozenset({"input_examples", "eager_input_streaming"}) + + +def _is_claude_tool_target(custom_llm_provider: str | None, model: str) -> bool: if custom_llm_provider == "anthropic": return True - if custom_llm_provider == "azure_ai" or custom_llm_provider == "bedrock" or custom_llm_provider == "vertex_ai": - return "claude" in model.lower() + model_lower: Final = model.lower() + if custom_llm_provider == "bedrock": + return "claude" in model_lower or ("arn:" in model_lower and ":bedrock:" in model_lower) + if custom_llm_provider == "azure_ai" or custom_llm_provider == "vertex_ai": + return "claude" in model_lower return False -def _drop_input_examples_from_tool(tool: dict) -> dict: - tool_copy: Final = tool.copy() - tool_copy.pop("input_examples", None) - function = tool_copy.get("function") - if isinstance(function, dict): - function = function.copy() - function.pop("input_examples", None) - tool_copy["function"] = function - return tool_copy +def _without_anthropic_only_tool_keys(tool: dict) -> dict: + kept: Final = {key: value for key, value in tool.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS} + function: Final = tool.get("function") + if not isinstance(function, dict): + return kept + return { + **kept, + "function": {key: value for key, value in function.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS}, + } -def _drop_input_examples_from_tools( - tools: list[dict] | None, -) -> list[dict] | None: +def _drop_anthropic_only_tool_keys(tools: list[dict] | None) -> list[dict] | None: if tools is None: return None - cleaned_tools: Final[list[dict]] = [] - for tool in tools: - if isinstance(tool, dict): - cleaned_tools.append(_drop_input_examples_from_tool(tool)) - else: - cleaned_tools.append(tool) - return cleaned_tools + return [_without_anthropic_only_tool_keys(tool) if isinstance(tool, dict) else tool for tool in tools] class _ProxyAuthHeadersProvider(Protocol): @@ -5360,8 +5358,8 @@ def completion( api_base=api_base, ) - if not _should_allow_input_examples(custom_llm_provider=custom_llm_provider, model=model): - tools = _drop_input_examples_from_tools(tools=tools) + if not _is_claude_tool_target(custom_llm_provider=custom_llm_provider, model=model): + tools = _drop_anthropic_only_tool_keys(tools=tools) if provider_specific_header is not None: headers.update( diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index cbbac3d247f..bd115c699d5 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -349,28 +349,66 @@ def test_bedrock_latency_optimized_inference(): assert json_data["performanceConfig"]["latency"] == "optimized" -def test_strip_input_examples_for_non_anthropic_providers(): +@pytest.mark.parametrize( + ("custom_llm_provider", "model", "expected"), + [ + ("anthropic", "claude-sonnet-5", True), + ("bedrock", "us.anthropic.claude-sonnet-5-20260501-v1:0", True), + ("bedrock", "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123", True), + ("bedrock", "us.amazon.nova-2-lite-v1:0", False), + ("vertex_ai", "claude-sonnet-5", True), + ("vertex_ai", "gemini-3.8-flash", False), + ("azure_ai", "claude-sonnet-4-6", True), + ("azure_ai", "gpt-5.6", False), + ("openai", "gpt-5.6", False), + ("gemini", "gemini-3.8-flash", False), + ], +) +def test_is_claude_tool_target(custom_llm_provider: str, model: str, expected: bool): + assert litellm_main._is_claude_tool_target(custom_llm_provider=custom_llm_provider, model=model) is expected + + +@pytest.mark.parametrize("key", ["input_examples", "eager_input_streaming"]) +def test_drop_anthropic_only_tool_keys_strips_tool_and_function_levels(key: str): tools = [ - { - "type": "function", - "name": "example_tool", - "input_examples": [{"foo": "bar"}], - "function": { - "name": "example_tool", - "input_examples": [{"foo": "bar"}], - }, - } + {"type": "function", "name": "example_tool", key: True, "function": {"name": "example_tool", key: True}}, + "opaque_tool", ] - assert not litellm_main._should_allow_input_examples( - custom_llm_provider="openai", model="gpt-4o-mini" + cleaned = litellm_main._drop_anthropic_only_tool_keys(tools=tools) + + assert cleaned == [ + {"type": "function", "name": "example_tool", "function": {"name": "example_tool"}}, + "opaque_tool", + ] + assert tools[0][key] is True + assert tools[0]["function"][key] is True + + +def test_completion_strips_eager_input_streaming_before_openai(respx_mock: respx.MockRouter, openai_api_response): + api_base: Final = "http://localhost:12346/v1" + mock_route: Final = respx_mock.post(url__regex=rf"{api_base}/chat/completions.*").mock( + return_value=httpx.Response(status_code=200, json=openai_api_response) ) - cleaned = litellm_main._drop_input_examples_from_tools(tools=tools) + litellm.completion( + model="openai/gpt-5.6", + messages=[{"role": "user", "content": "Write the file"}], + tools=[ + { + "type": "function", + "function": {"name": "write_file", "parameters": {"type": "object", "properties": {}}}, + "eager_input_streaming": True, + } + ], + api_base=api_base, + api_key="fake_openai_api_key", + ) - assert isinstance(cleaned, list) - assert "input_examples" not in cleaned[0] - assert "input_examples" not in cleaned[0]["function"] + assert mock_route.called + sent_tool: Final = json.loads(respx_mock.calls[0].request.content)["tools"][0] + assert "eager_input_streaming" not in sent_tool + assert sent_tool["function"]["name"] == "write_file" def test_custom_provider_with_extra_headers(): From 66c01cf35c14d06877f7560af7604a5c8e36154c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:39:02 -0700 Subject: [PATCH 231/253] refactor(responses): map finish reasons to incomplete_details through a lookup table The match statement in _incomplete_details_for_finish_reason tripped CodeQL's mixed explicit and implicit returns alert (code-scanning 12640). A module-level MappingProxyType keyed by finish reason gives the same three mappings with one explicit return path --- .../transformation.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 961fffd3d42..7337beba45c 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -112,6 +112,9 @@ ResponseTools: TypeAlias = Sequence[Mapping[str, object]] | None ChatToolParam: TypeAlias = ChatCompletionToolParam | OpenAIMcpServerTool NAMESPACE_DESCRIPTION_SEPARATOR: Final = "\n\n" NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: Final = frozenset({"function", "custom"}) +_INCOMPLETE_REASON_BY_FINISH_REASON: Final[Mapping[str, Literal["max_output_tokens", "content_filter"]]] = ( + MappingProxyType({"length": "max_output_tokens", "content_filter": "content_filter", "refusal": "content_filter"}) +) @dataclass(frozen=True, slots=True) @@ -2303,13 +2306,10 @@ class LiteLLMCompletionResponsesConfig: ) -> IncompleteDetails | None: if existing is not None: return existing - match finish_reason: - case "length": - return IncompleteDetails(reason="max_output_tokens") - case "content_filter" | "refusal": - return IncompleteDetails(reason="content_filter") - case _: - return None + if finish_reason is None: + return None + reason: Final = _INCOMPLETE_REASON_BY_FINISH_REASON.get(finish_reason) + return IncompleteDetails(reason=reason) if reason is not None else None @staticmethod def _tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str: From 44034c1d5e25c8f1888dfd3f13e847cc44860fab Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:43:07 -0700 Subject: [PATCH 232/253] test: source the gemma context window limits and isolate the cost map cache --- .../gemma/test_vertex_ai_gemma_global_endpoint.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py index 9d08daa0a81..85a2124ab02 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/gemma/test_vertex_ai_gemma_global_endpoint.py @@ -180,12 +180,11 @@ class TestCreateVertexURLGemma: # --------------------------------------------------------------------------- -def test_gemma_maas_context_window_matches_google(monkeypatch): - monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") - monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) - +def test_gemma_maas_context_window_matches_google(local_model_cost_map): info = litellm.get_model_info("vertex_ai/google/gemma-4-26b-a4b-it-maas") + # 262,144 context length and 128,000 maximum output per Google's model page, checked 2026-09-18: + # https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/maas/google/gemma-4-26b-a4b-it assert info["max_input_tokens"] == 262144 assert info["max_output_tokens"] == 128000 From b4bfd92a2a4ab11df53d886704b621b6e8cd2339 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 13:36:06 -0700 Subject: [PATCH 233/253] refactor(rust): route-neutral callback contract Every legacy callback call from callbacks-legacy now goes through one typed Python shim, litellm.rust_bridge.legacy_callbacks, the only Python module the crate reaches. Before, the crate called Logging methods, litellm.utils hooks, the logging worker, the executor and several litellm globals directly, and its tests retyped those signatures by hand, so an outdated fake could accept a call the real code rejects. python_contract.json lists each shim function's parameters: a Python test pins it to the real signatures and a Rust test pins it to the Rust enum. The lifecycle contract changes to match the Python @client wrapper: - the driver emits CallEvent::Started before begin, so every host sees one start time - RequestContext carries the route-resolved api_key, so legacy pre_call and post_call receive it, and post_call's additional_args match the Python OCR path - Passthrough and its re-aliasing are gone - async deployment hooks always run, and the "no callbacks" shortcut that skipped the logging payload is removed, as in the Python path The OCR api_key is a SecretValue from the wire request onward, so Debug output upstream of the callback contract cannot leak it. host-python's RouteHost now classifies native failures once through classify, and host ops return HostOpError. The OCR route host keeps main's public errors by sending both through the existing Python map_failure. --- litellm-rust/Cargo.lock | 3 + litellm-rust/crates/auth/src/secret.rs | 4 +- .../crates/callbacks-legacy/AGENTS.md | 14 +- .../crates/callbacks-legacy/Cargo.toml | 4 + .../callbacks-legacy/python_contract.json | 98 +++++ .../crates/callbacks-legacy/src/adapter.rs | 101 +++-- .../crates/callbacks-legacy/src/call.rs | 43 +-- .../crates/callbacks-legacy/src/callbacks.rs | 242 +++--------- .../callbacks-legacy/src/legacy_python.rs | 155 ++++++++ .../crates/callbacks-legacy/src/lib.rs | 5 +- .../crates/callbacks-legacy/src/logger.rs | 91 +---- .../callbacks-legacy/src/preparation.rs | 44 +-- .../crates/callbacks-legacy/tests/deferred.rs | 18 +- .../tests/deployment_hooks.rs | 18 +- .../crates/callbacks-legacy/tests/payload.rs | 113 +++--- .../crates/callbacks-legacy/tests/support.rs | 135 ++++--- .../crates/callbacks-legacy/tests/terminal.rs | 46 ++- litellm-rust/crates/callbacks/Cargo.toml | 1 + litellm-rust/crates/callbacks/src/event.rs | 77 +--- litellm-rust/crates/callbacks/src/run.rs | 36 +- litellm-rust/crates/core/src/ocr/handler.rs | 13 +- litellm-rust/crates/core/src/ocr/mod.rs | 4 +- litellm-rust/crates/core/src/ocr/prepare.rs | 4 +- .../crates/core/src/ocr/provider_config.rs | 46 ++- litellm-rust/crates/core/src/ocr/types.rs | 18 +- litellm-rust/crates/core/src/ocr/wire.rs | 4 +- .../tests/azure_document_intelligence_ocr.rs | 2 +- litellm-rust/crates/core/tests/ocr.rs | 22 +- .../crates/core/tests/ocr/document.rs | 152 ++++++++ .../crates/core/tests/ocr/passthrough.rs | 282 -------------- litellm-rust/crates/core/tests/ocr/support.rs | 10 +- litellm-rust/crates/host-python/AGENTS.md | 5 +- .../crates/host-python/src/adapter.rs | 55 ++- .../crates/host-python/src/argument.rs | 51 +++ litellm-rust/crates/host-python/src/driver.rs | 364 +++++++++++++----- litellm-rust/crates/host-python/src/lib.rs | 8 +- .../document_intelligence/transformation.rs | 21 +- .../llms/src/azure_ai/ocr/transformation.rs | 21 +- .../llms/src/base_llm/ocr/transformation.rs | 17 +- .../llms/src/cohere/ocr/transformation.rs | 6 +- .../llms/src/custom_httpx/llm_http_handler.rs | 32 +- .../llms/src/mistral/ocr/transformation.rs | 6 +- .../llms/src/reducto/ocr/transformation.rs | 8 +- .../llms/src/vertex_ai/ocr/transformation.rs | 5 +- litellm-rust/crates/python-bridge/AGENTS.md | 2 +- .../python-bridge/src/routes/ocr/host.rs | 61 +-- .../python-bridge/src/routes/ocr/project.rs | 10 +- litellm/litellm_core_utils/litellm_logging.py | 43 +-- litellm/rust_bridge/legacy_callbacks.py | 291 ++++++++++---- .../rust_bridge/test_legacy_callbacks.py | 19 +- tests/test_litellm_rust/ocr/test_lifecycle.py | 31 +- 51 files changed, 1564 insertions(+), 1297 deletions(-) create mode 100644 litellm-rust/crates/callbacks-legacy/python_contract.json create mode 100644 litellm-rust/crates/callbacks-legacy/src/legacy_python.rs create mode 100644 litellm-rust/crates/core/tests/ocr/document.rs delete mode 100644 litellm-rust/crates/core/tests/ocr/passthrough.rs create mode 100644 litellm-rust/crates/host-python/src/argument.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index c359ca19986..0265e0adbc2 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2031,6 +2031,7 @@ dependencies = [ name = "litellm-callbacks" version = "0.1.0" dependencies = [ + "litellm-auth", "rstest", "serde_json", "tokio", @@ -2040,11 +2041,13 @@ dependencies = [ name = "litellm-callbacks-legacy" version = "0.1.0" dependencies = [ + "litellm-auth", "litellm-callbacks", "litellm-host-python", "pyo3", "rstest", "serde_json", + "strum", ] [[package]] diff --git a/litellm-rust/crates/auth/src/secret.rs b/litellm-rust/crates/auth/src/secret.rs index 3ecb0a835ee..a07fe3eaad9 100644 --- a/litellm-rust/crates/auth/src/secret.rs +++ b/litellm-rust/crates/auth/src/secret.rs @@ -1,6 +1,8 @@ +use serde::Deserialize; use veil::Redact; -#[derive(Redact, Clone)] +#[derive(Redact, Clone, Deserialize)] +#[serde(transparent)] pub struct SecretValue(#[redact(with = "[REDACTED]")] String); impl SecretValue { diff --git a/litellm-rust/crates/callbacks-legacy/AGENTS.md b/litellm-rust/crates/callbacks-legacy/AGENTS.md index e4762d3037a..e184e2fb415 100644 --- a/litellm-rust/crates/callbacks-legacy/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy/AGENTS.md @@ -1,15 +1,17 @@ - Target invariants, not completion claims - Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits) - - The driver in `litellm-host-python`, the routes and core see one `CallbackAdapter`; they never learn which Python objects consume a call + - The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call +- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`) + - The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it + - Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython` - `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy -- `setup` decides once who owns the `Logging` instance and returns it as `CallSetup.bridge_owned`; `PythonLogger` carries it and nothing on the instance records it - - A logger the caller passed as `litellm_logging_obj` is caller-owned and observed in full, because the caller reads it after the call; the proxy is the live case - - A logger `function_setup` built for this call is bridge-owned, so each fan-out phase is skipped when `callbacks_needed` finds no registry, dynamic callback, `logger_fn` or debug switch for it; cost, timing and response metadata still run +- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does + - Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's - Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation - Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view - - Re-alias every `passthrough_fields` body key to the caller's object before `pre_call`; a keyword wins over the request attribute even when it is an explicit `None` + - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup` - Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only - - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-callbacks`, `litellm-host-python` and the bridge; the only facts that cross from the route are the prepared keyword view and `RequestContext.passthrough_fields` + - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-callbacks`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view - Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts - Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once diff --git a/litellm-rust/crates/callbacks-legacy/Cargo.toml b/litellm-rust/crates/callbacks-legacy/Cargo.toml index 96c9c9ed560..3cf9382f857 100644 --- a/litellm-rust/crates/callbacks-legacy/Cargo.toml +++ b/litellm-rust/crates/callbacks-legacy/Cargo.toml @@ -9,8 +9,12 @@ autotests = false [dependencies] litellm-callbacks.workspace = true litellm-host-python.workspace = true + pyo3.workspace = true +strum.workspace = true +serde_json.workspace = true [dev-dependencies] +litellm-auth.workspace = true rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/callbacks-legacy/python_contract.json b/litellm-rust/crates/callbacks-legacy/python_contract.json new file mode 100644 index 00000000000..840c0abfa45 --- /dev/null +++ b/litellm-rust/crates/callbacks-legacy/python_contract.json @@ -0,0 +1,98 @@ +{ + "setup": [ + "call_type", + "args", + "kwargs", + "start_time", + "asynchronous" + ], + "check_limits": [ + "kwargs" + ], + "finalize": [ + "response", + "logger", + "kwargs", + "start_time", + "end_time" + ], + "update_logging": [ + "logger", + "kwargs", + "model", + "optional_params", + "litellm_params", + "custom_llm_provider" + ], + "pre_call": [ + "logger", + "input", + "api_key", + "additional_args" + ], + "post_call": [ + "logger", + "original_response", + "api_key", + "additional_args" + ], + "defers_async_logging": [ + "logger" + ], + "defer_success": [ + "logger", + "pending" + ], + "sync_success_for_async_call": [ + "logger", + "response", + "start", + "end" + ], + "failure_handler": [ + "logger", + "error", + "start", + "end", + "asynchronous" + ], + "submit_success": [ + "logger", + "response", + "start", + "end" + ], + "async_success_handler": [ + "logger", + "response", + "start", + "end" + ], + "enqueue_logging": [ + "coroutine" + ], + "restore_context": [ + "logger" + ], + "custom_pricing_fields": [], + "is_internal_call": [], + "credential_list": [], + "warn_unknown_credential": [ + "name", + "loaded" + ], + "before_deployment_call": [ + "kwargs", + "call_type" + ], + "after_deployment_success": [ + "kwargs", + "response", + "call_type" + ], + "after_deployment_failure": [ + "kwargs", + "error", + "call_type" + ] +} diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index df346506094..7204207cd62 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -4,7 +4,7 @@ use litellm_callbacks::event::{CallEvent, FailureOrigin, RequestContext, Timing, WireRequest}; use litellm_host_python::{ - AdapterStep, CallbackAdapter, PublicValue, from_py, missing_state, to_py, + LifecycleStep, PublicValue, PythonLifecycle, from_py, missing_state, to_py, }; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -12,6 +12,7 @@ use pyo3::{ prelude::*, types::PyDict, }; +use serde_json::Value; use crate::{ DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger, @@ -43,7 +44,7 @@ pub struct LegacyLogging { response: Option>, error: Option>, body: Option>, - headers: Option>, + context: Option, asynchronous: bool, internal: bool, pending: Option, @@ -76,7 +77,7 @@ impl LegacyLogging { response: None, error: None, body: None, - headers: None, + context: None, asynchronous, internal: false, pending: None, @@ -85,8 +86,8 @@ impl LegacyLogging { /// Deployment hooks are awaited, and Python's synchronous `@client` wrapper never /// runs them. - fn deployment_hooks(&self, py: Python<'_>) -> PyResult { - Ok(self.asynchronous && DeploymentHooks::needed(py)?) + fn runs_deployment_hooks(&self) -> bool { + self.asynchronous } fn logger(&self) -> PyResult<&PythonLogger> { @@ -95,13 +96,13 @@ impl LegacyLogging { }) } - fn prepare(&mut self, py: Python<'_>) -> PyResult { + fn prepare(&mut self, py: Python<'_>) -> PyResult { let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind(); self.call.set_kwargs(prepared); - Ok(AdapterStep::Arguments(self.call.kwargs().clone_ref(py))) + Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py))) } - fn finalize(&mut self, py: Python<'_>) -> PyResult { + fn finalize(&mut self, py: Python<'_>) -> PyResult { finalize( py, &self.response, @@ -112,7 +113,7 @@ impl LegacyLogging { )?; self.response .as_ref() - .map(|response| AdapterStep::Response(response.clone_ref(py))) + .map(|response| LifecycleStep::Response(response.clone_ref(py))) .ok_or_else(missing_state) } @@ -145,9 +146,7 @@ impl LegacyLogging { .get_item("fallbacks")? .is_none_or(|value| value.is_none()) { - if !logger.callbacks_needed(py, "async_success")? { - logger.success_bookkeeping(py, &self.response, &self.start, &self.end, true)?; - } else if logger.defers_async_logging(py) { + if logger.defers_async_logging(py) { let pending = Py::new( py, PendingLogging { @@ -165,12 +164,12 @@ impl LegacyLogging { /// The sync failure handler, then the async one for async calls. Ordinary handler /// errors never replace the selected failure or suppress the other family; a /// cancellation does end the call. - fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult { + fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult { let (Some(logger), Some(error)) = (&self.logger, &self.error) else { - return Ok(AdapterStep::Done); + return Ok(LifecycleStep::Done); }; if self.asynchronous && self.internal { - return Ok(AdapterStep::Done); + return Ok(LifecycleStep::Done); } if let Err(failure) = logger.failure(py, error, &self.start, &self.end, false) && is_cancellation(py, &failure) @@ -178,27 +177,27 @@ impl LegacyLogging { return Err(failure); } if !self.asynchronous { - return Ok(AdapterStep::Done); + return Ok(LifecycleStep::Done); } match logger.failure(py, error, &self.start, &self.end, true) { Ok(Some(awaitable)) => { self.pending = Some(Pending::AsyncFailure); - Ok(AdapterStep::Await(awaitable)) + Ok(LifecycleStep::Await(awaitable)) } - Ok(None) => Ok(AdapterStep::Done), + Ok(None) => Ok(LifecycleStep::Done), Err(failure) if is_cancellation(py, &failure) => Err(failure), - Err(_) => Ok(AdapterStep::Done), + Err(_) => Ok(LifecycleStep::Done), } } } -impl CallbackAdapter for LegacyLogging { +impl PythonLifecycle for LegacyLogging { fn begin( &mut self, py: Python<'_>, arguments: Py, started_at: f64, - ) -> PyResult { + ) -> PyResult { self.call.set_kwargs(arguments); self.start = datetime(py, started_at)?; self.internal = is_internal_call(py)?; @@ -212,9 +211,9 @@ impl CallbackAdapter for LegacyLogging { )?; self.logger = Some(result.logger()?); self.call.set_kwargs(result.kwargs()?); - if self.deployment_hooks(py)? { + if self.runs_deployment_hooks() { self.pending = Some(Pending::DeploymentPreCall); - return Ok(AdapterStep::Await(DeploymentHooks::before_call( + return Ok(LifecycleStep::Await(DeploymentHooks::before_call( py, self.call.kwargs(), self.surface.call_type, @@ -228,18 +227,16 @@ impl CallbackAdapter for LegacyLogging { py: Python<'_>, wire: Box, context: &RequestContext, - ) -> PyResult { + ) -> PyResult { let logger = self.logger()?; logger.update_from_kwargs(py, self.call.kwargs(), &wire, context)?; - if !logger.callbacks_needed(py, "payload")? { - logger.record_api_call_start(py)?; - return Ok(AdapterStep::Wire(wire)); - } let body = to_py(py, &wire.body)? .into_bound(py) .cast_into::()?; - for name in context.passthrough_fields.iter() { - if let Some(value) = self.call.lookup(py, name)? { + for (name, sent) in wire.body.as_object().into_iter().flatten() { + if let Some(value) = self.call.lookup(py, name)? + && from_py::(&value).is_ok_and(|caller| caller == *sent) + { body.set_item(name, value)?; } } @@ -248,12 +245,11 @@ impl CallbackAdapter for LegacyLogging { headers.set_item(name, value)?; } self.body = Some(body.clone().unbind()); - self.headers = Some(headers.clone().unbind()); - let api_key = self.call.lookup(py, "api_key")?; + self.context = Some(context.clone()); self.logger()?.pre_call( py, self.surface.input_description, - api_key.as_ref(), + context.api_key.as_ref().map(|api_key| api_key.expose()), &body, &headers, &wire.url, @@ -262,7 +258,7 @@ impl CallbackAdapter for LegacyLogging { .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) .collect::>>()?; - Ok(AdapterStep::Wire(Box::new(WireRequest { + Ok(LifecycleStep::Wire(Box::new(WireRequest { body: from_py(&body)?, headers, ..*wire @@ -274,12 +270,12 @@ impl CallbackAdapter for LegacyLogging { py: Python<'_>, response: Py, timing: Timing, - ) -> PyResult { + ) -> PyResult { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response); - if self.deployment_hooks(py)? { + if self.runs_deployment_hooks() { self.pending = Some(Pending::DeploymentPostCall); - return Ok(AdapterStep::Await(DeploymentHooks::after_success( + return Ok(LifecycleStep::Await(DeploymentHooks::after_success( py, self.call.kwargs(), &self.response, @@ -294,31 +290,35 @@ impl CallbackAdapter for LegacyLogging { py: Python<'_>, event: &CallEvent, public: Option>, - ) -> PyResult { + ) -> PyResult { match (event, public) { + (CallEvent::Started { .. }, _) => Ok(LifecycleStep::Done), (CallEvent::ResponseReceived { raw }, _) => { - let logger = self.logger()?; - if logger.callbacks_needed(py, "payload")? { - logger.post_call(py, &raw.body, self.body.as_ref(), self.headers.as_ref())?; - } - Ok(AdapterStep::Done) + let api_key = self + .context + .as_ref() + .and_then(|context| context.api_key.as_ref()) + .map(|api_key| api_key.expose()); + self.logger()? + .post_call(py, &raw.body, api_key, self.body.as_ref())?; + Ok(LifecycleStep::Done) } (CallEvent::Succeeded { timing }, Some(PublicValue::Response(response))) => { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response.clone_ref(py)); self.dispatch_success(py)?; - Ok(AdapterStep::Done) + Ok(LifecycleStep::Done) } (CallEvent::Failed { timing, origin }, Some(PublicValue::Error(error))) => { self.end = Some(datetime(py, timing.end_time)?); self.error = Some(error.clone_ref(py).into_value(py)); if *origin == FailureOrigin::Call && self.logger.is_some() - && self.deployment_hooks(py)? + && self.runs_deployment_hooks() { let error = self.error.as_ref().ok_or_else(missing_state)?; self.pending = Some(Pending::DeploymentFailure); - return Ok(AdapterStep::Await(DeploymentHooks::after_failure( + return Ok(LifecycleStep::Await(DeploymentHooks::after_failure( py, self.call.kwargs(), error, @@ -331,7 +331,7 @@ impl CallbackAdapter for LegacyLogging { } } - fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult { + fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult { match self.pending.take().ok_or_else(missing_state)? { Pending::DeploymentPreCall => { self.call @@ -345,7 +345,7 @@ impl CallbackAdapter for LegacyLogging { Pending::DeploymentFailure => self.dispatch_failure(py), Pending::AsyncFailure => match result { Err(failure) if is_cancellation(py, &failure) => Err(failure), - _ => Ok(AdapterStep::Done), + _ => Ok(LifecycleStep::Done), }, } } @@ -357,7 +357,7 @@ impl CallbackAdapter for LegacyLogging { error.write_unraisable(py, None); } self.body = None; - self.headers = None; + self.context = None; } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -369,8 +369,7 @@ impl CallbackAdapter for LegacyLogging { visit.call(&self.end)?; visit.call(&self.response)?; visit.call(&self.error)?; - visit.call(&self.body)?; - visit.call(&self.headers) + visit.call(&self.body) } } diff --git a/litellm-rust/crates/callbacks-legacy/src/call.rs b/litellm-rust/crates/callbacks-legacy/src/call.rs index 59090ee8d60..bd1e2525d3d 100644 --- a/litellm-rust/crates/callbacks-legacy/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy/src/call.rs @@ -4,7 +4,7 @@ //! this crate holds them. use litellm_callbacks::{machine::Machine, route::Route}; -use litellm_host_python::{RouteHost, run_call}; +use litellm_host_python::{RouteHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -63,21 +63,6 @@ impl PublicCall { } } -/// The caller's own object for a public argument, as every legacy reader resolves it: the -/// keyword if given, even an explicit `None`, else the bound request's attribute. A route -/// host projecting from the prepared keyword view uses the same rule, so the callbacks -/// and the provider see one object per argument. -pub fn lookup<'py>( - kwargs: &Bound<'py, PyDict>, - request: &Bound<'py, PyAny>, - name: &str, -) -> PyResult>> { - if let Some(value) = kwargs.get_item(name)? { - return Ok(Some(value)); - } - request.getattr_opt(name) -} - /// Runs one native call under the legacy `Logging` contract: the route host projects from /// the keyword view the contract prepares, and the contract observes the call. pub fn run_legacy_call( @@ -121,32 +106,6 @@ mod tests { (call, locals) } - #[test] - fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() { - Python::initialize(); - Python::attach(|py| { - let (call, locals) = capture( - py, - c" -key = object() -document = {'type': 'document_url'} -class Request: - api_key = 'from-request' - api_base = 'from-request' - document = document -request = Request() -kwargs = {'api_key': key, 'api_base': None} -", - ); - let key = locals.get_item("key").unwrap().unwrap(); - let document = locals.get_item("document").unwrap().unwrap(); - assert!(call.lookup(py, "api_key").unwrap().unwrap().is(&key)); - assert!(call.lookup(py, "api_base").unwrap().unwrap().is_none()); - assert!(call.lookup(py, "document").unwrap().unwrap().is(&document)); - assert!(call.lookup(py, "model").unwrap().is_none()); - }); - } - #[test] fn capture_copies_the_keyword_dict_without_copying_its_values() { Python::initialize(); diff --git a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs index aa586013e75..9464f1d6612 100644 --- a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs +++ b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs @@ -6,11 +6,10 @@ use litellm_callbacks::event::{RequestContext, WireRequest}; use litellm_host_python::to_py; use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict}; +use crate::legacy_python::{Logging, Wrapper}; use crate::logger::PythonLogger; pub trait LegacyCallbacks { - fn callbacks_needed(&self, py: Python<'_>, phase: &str) -> PyResult; - /// `Logging.update_from_kwargs`: what the logger is told about the request it is /// about to see, with consumed credentials redacted. fn update_from_kwargs( @@ -21,26 +20,24 @@ pub trait LegacyCallbacks { context: &RequestContext, ) -> PyResult<()>; - fn record_api_call_start(&self, py: Python<'_>) -> PyResult<()>; - - /// `Logging.pre_call`, or its payload-free shortcut when no input callback listens. + /// `Logging.pre_call`. fn pre_call( &self, py: Python<'_>, input: &str, - api_key: Option<&Bound<'_, PyAny>>, + api_key: Option<&str>, body: &Bound<'_, PyDict>, headers: &Bound<'_, PyDict>, url: &str, ) -> PyResult<()>; - /// `Logging.post_call`, or its payload-free shortcut when no input callback listens. + /// `Logging.post_call`. fn post_call( &self, py: Python<'_>, original_response: &str, + api_key: Option<&str>, body: Option<&Py>, - headers: Option<&Py>, ) -> PyResult<()>; fn defers_async_logging(&self, py: Python<'_>) -> bool; @@ -82,16 +79,6 @@ pub trait LegacyCallbacks { } impl LegacyCallbacks for PythonLogger { - fn callbacks_needed(&self, py: Python<'_>, phase: &str) -> PyResult { - if !self.bridge_owned() { - return Ok(true); - } - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("callbacks_needed")? - .call1((self.object(py), phase))? - .extract() - } - fn update_from_kwargs( &self, py: Python<'_>, @@ -100,18 +87,13 @@ impl LegacyCallbacks for PythonLogger { context: &RequestContext, ) -> PyResult<()> { let secret_fields: Vec<&str> = context.secret_fields.iter().map(String::as_str).collect(); - let update = PyDict::new(py); - update.set_item("kwargs", redact(py, kwargs.bind(py), &secret_fields)?)?; - update.set_item("model", &context.model)?; - update.set_item( - "optional_params", - redact( - py, - &to_py(py, &context.optional_params)? - .into_bound(py) - .cast_into::()?, - &secret_fields, - )?, + let redacted_kwargs = redact(py, kwargs.bind(py), &secret_fields)?; + let optional_params = redact( + py, + &to_py(py, &context.optional_params)? + .into_bound(py) + .cast_into::()?, + &secret_fields, )?; let params = PyDict::new(py); params.set_item( @@ -131,15 +113,17 @@ impl LegacyCallbacks for PythonLogger { params.set_item(name, value)?; } } - update.set_item("litellm_params", params)?; - update.set_item("custom_llm_provider", &context.custom_llm_provider)?; - self.object(py) - .call_method("update_from_kwargs", (), Some(&update))?; - Ok(()) - } - - fn record_api_call_start(&self, py: Python<'_>) -> PyResult<()> { - self.object(py).call_method0("record_api_call_start_time")?; + Logging::Update.call( + py, + ( + self.object(py), + redacted_kwargs, + &context.model, + optional_params, + params, + &context.custom_llm_provider, + ), + )?; Ok(()) } @@ -147,7 +131,7 @@ impl LegacyCallbacks for PythonLogger { &self, py: Python<'_>, input: &str, - api_key: Option<&Bound<'_, PyAny>>, + api_key: Option<&str>, body: &Bound<'_, PyDict>, headers: &Bound<'_, PyDict>, url: &str, @@ -156,17 +140,7 @@ impl LegacyCallbacks for PythonLogger { additional.set_item("complete_input_dict", body)?; additional.set_item("headers", headers)?; additional.set_item("api_base", url)?; - let kwargs = PyDict::new(py); - kwargs.set_item("input", input)?; - kwargs.set_item("api_key", api_key)?; - kwargs.set_item("additional_args", &additional)?; - if self.callbacks_needed(py, "input")? { - self.object(py).call_method("pre_call", (), Some(&kwargs))?; - } else { - self.object(py) - .call_method("_pre_call", (), Some(&kwargs))?; - self.record_api_call_start(py)?; - } + Logging::PreCall.call(py, (self.object(py), input, api_key, &additional))?; Ok(()) } @@ -174,37 +148,28 @@ impl LegacyCallbacks for PythonLogger { &self, py: Python<'_>, original_response: &str, + api_key: Option<&str>, body: Option<&Py>, - headers: Option<&Py>, ) -> PyResult<()> { let additional = PyDict::new(py); additional.set_item("complete_input_dict", body)?; - additional.set_item("headers", headers)?; - if self.callbacks_needed(py, "input")? { - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", original_response)?; - kwargs.set_item("additional_args", &additional)?; - self.object(py) - .call_method("post_call", (), Some(&kwargs))?; - } else { - let response = py - .import("json")? - .call_method1("dumps", (original_response,))?; - self.object(py).call_method1( - "record_post_call", - (response, py.None(), py.None(), additional), - )?; - } + Logging::PostCall.call( + py, + (self.object(py), original_response, api_key, &additional), + )?; Ok(()) } + fn defers_async_logging(&self, py: Python<'_>) -> bool { - self.object(py) - .getattr("_defer_async_logging") - .is_ok_and(|value| value.is_truthy().unwrap_or(false)) + Logging::DefersAsync + .call(py, (self.object(py),)) + .and_then(|value| value.extract()) + .unwrap_or(false) } fn defer_success(&self, py: Python<'_>, pending: &Bound<'_, PyAny>) -> PyResult<()> { - self.object(py).setattr("_native_pending_logging", pending) + Logging::DeferSuccess.call(py, (self.object(py), pending))?; + Ok(()) } fn sync_success_for_async_call( @@ -214,13 +179,7 @@ impl LegacyCallbacks for PythonLogger { start: &Py, end: &Option>, ) -> PyResult<()> { - if !self.callbacks_needed(py, "sync_success_async")? { - return Ok(()); - } - self.object(py).call_method1( - "handle_sync_success_callbacks_for_async_calls", - (response, start, end), - )?; + Logging::SyncSuccessForAsyncCall.call(py, (self.object(py), response, start, end))?; Ok(()) } @@ -232,34 +191,11 @@ impl LegacyCallbacks for PythonLogger { end: &Option>, asynchronous: bool, ) -> PyResult>> { - if !self.callbacks_needed( - py, - if asynchronous { - "async_failure" - } else { - "sync_failure" - }, - )? { - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("failure_bookkeeping")? - .call1((self.object(py), error, start, end, asynchronous))?; - return Ok(None); - } - let trace = py - .import("traceback")? - .getattr("format_exception")? - .call1((error,))?; - let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?; - let value = self.object(py).call_method1( - if asynchronous { - "async_failure_handler" - } else { - "failure_handler" - }, - (error, trace, start, end), - )?; + let value = + Logging::FailureHandler.call(py, (self.object(py), error, start, end, asynchronous))?; Ok(asynchronous.then(|| value.unbind())) } + fn submit_success( &self, py: Python<'_>, @@ -267,22 +203,7 @@ impl LegacyCallbacks for PythonLogger { start: &Py, end: &Option>, ) -> PyResult<()> { - if !self.callbacks_needed(py, "sync_success")? { - return self.success_bookkeeping(py, response, start, end, false); - } - let context = py.import("contextvars")?.call_method0("copy_context")?; - py.import("litellm.litellm_core_utils.litellm_logging")? - .getattr("executor")? - .call_method1( - "submit", - ( - context.getattr("run")?, - self.object(py).getattr("success_handler")?, - response, - start, - end, - ), - )?; + Logging::SubmitSuccess.call(py, (self.object(py), response, start, end))?; Ok(()) } @@ -293,18 +214,9 @@ impl LegacyCallbacks for PythonLogger { start: &Py, end: &Option>, ) -> PyResult<()> { - if !self.callbacks_needed(py, "async_success")? { - return self.success_bookkeeping(py, response, start, end, true); - } - let context = py.import("contextvars")?.call_method0("copy_context")?; - let worker = py - .import("litellm.litellm_core_utils.logging_worker")? - .getattr("GLOBAL_LOGGING_WORKER")? - .getattr("ensure_initialized_and_enqueue")?; - let coroutine = self - .object(py) - .call_method1("async_success_handler", (response, start, end))?; - let enqueue = context.call_method1("run", (worker, &coroutine)); + let coroutine = + Logging::AsyncSuccessHandler.call(py, (self.object(py), response, start, end))?; + let enqueue = Logging::Enqueue.call(py, (&coroutine,)); if enqueue.is_err() && let Err(error) = coroutine.call_method0("close") { @@ -315,14 +227,7 @@ impl LegacyCallbacks for PythonLogger { } fn custom_pricing_fields(py: Python<'_>) -> PyResult> { - py.import("litellm.types.utils")? - .getattr("CustomPricingLiteLLMParams")? - .getattr("model_fields")? - .cast_into::()? - .keys() - .iter() - .map(|name| name.extract::()) - .collect() + Logging::CustomPricingFields.call(py, ())?.extract() } fn redact( @@ -347,58 +252,5 @@ fn redact( /// Proxy-internal calls skip the legacy success fan-out. pub fn is_internal_call(py: Python<'_>) -> PyResult { - py.import("litellm._internal_context")? - .getattr("is_internal_call")? - .call_method0("get")? - .extract() -} - -#[cfg(test)] -mod tests { - use pyo3::types::PyDict; - - use super::*; - - fn logger_whose_registries_need_no_input(py: Python<'_>, bridge_owned: bool) -> PythonLogger { - let locals = PyDict::new(py); - py.run( - c" -import sys -import types -for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.legacy_callbacks'): - sys.modules.setdefault(name, types.ModuleType(name)) -legacy = sys.modules['litellm.rust_bridge.legacy_callbacks'] -legacy.callbacks_needed = lambda logger, phase: logger.needed.get(phase, True) -class Logger: - needed = {'input': False} -logger = Logger() -", - Some(&locals), - Some(&locals), - ) - .unwrap(); - PythonLogger::new( - locals.get_item("logger").unwrap().unwrap().unbind(), - bridge_owned, - ) - } - - #[test] - fn a_caller_owned_logger_is_observed_in_full() { - Python::initialize(); - Python::attach(|py| { - let logger = logger_whose_registries_need_no_input(py, false); - assert!(logger.callbacks_needed(py, "input").unwrap()); - }); - } - - #[test] - fn a_bridge_owned_logger_is_elided_where_no_registry_needs_it() { - Python::initialize(); - Python::attach(|py| { - let logger = logger_whose_registries_need_no_input(py, true); - assert!(!logger.callbacks_needed(py, "input").unwrap()); - assert!(logger.callbacks_needed(py, "payload").unwrap()); - }); - } + Wrapper::IsInternalCall.call(py, ())?.extract() } diff --git a/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs b/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs new file mode 100644 index 00000000000..a924775070c --- /dev/null +++ b/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs @@ -0,0 +1,155 @@ +use pyo3::prelude::*; +use strum::{IntoStaticStr, VariantArray}; + +const MODULE: &str = "litellm.rust_bridge.legacy_callbacks"; + +/// Every litellm Python internal the native call still borrows, grouped by the subsystem it +/// belongs to. Rust drives the call; these exist only so behaviour that Python owns today +/// (span tracking, the standard logging payload, spend, callback fan-out) keeps working. +/// A group is deleted once Rust owns that subsystem, so this enum only shrinks. Calling a +/// user's own callback is not borrowing and does not belong here. +/// +/// `litellm/rust_bridge/legacy_callbacks.py` is the only Python module behind it, and +/// `python_contract.json` pins each function's parameters on both sides. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum LegacyPython { + Wrapper(Wrapper), + Logging(Logging), + DeploymentHooks(DeploymentHooks), +} + +/// The `@client` wrapper around the call: `function_setup`, limits, credentials, +/// response metadata and the correlation context. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum Wrapper { + #[strum(serialize = "setup")] + Setup, + #[strum(serialize = "check_limits")] + CheckLimits, + #[strum(serialize = "credential_list")] + CredentialList, + #[strum(serialize = "warn_unknown_credential")] + WarnUnknownCredential, + #[strum(serialize = "is_internal_call")] + IsInternalCall, + #[strum(serialize = "finalize")] + Finalize, + #[strum(serialize = "restore_context")] + RestoreContext, +} + +/// litellm's `Logging` object and the sync and async callback fan-out behind it. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum Logging { + #[strum(serialize = "custom_pricing_fields")] + CustomPricingFields, + #[strum(serialize = "update_logging")] + Update, + #[strum(serialize = "pre_call")] + PreCall, + #[strum(serialize = "post_call")] + PostCall, + #[strum(serialize = "defers_async_logging")] + DefersAsync, + #[strum(serialize = "defer_success")] + DeferSuccess, + #[strum(serialize = "sync_success_for_async_call")] + SyncSuccessForAsyncCall, + #[strum(serialize = "submit_success")] + SubmitSuccess, + #[strum(serialize = "async_success_handler")] + AsyncSuccessHandler, + #[strum(serialize = "enqueue_logging")] + Enqueue, + #[strum(serialize = "failure_handler")] + FailureHandler, +} + +/// The `litellm.utils` fan-outs that run every callback's deployment hook. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum DeploymentHooks { + #[strum(serialize = "before_deployment_call")] + BeforeDeploymentCall, + #[strum(serialize = "after_deployment_success")] + AfterDeploymentSuccess, + #[strum(serialize = "after_deployment_failure")] + AfterDeploymentFailure, +} + +impl LegacyPython { + fn name(self) -> &'static str { + match self { + Self::Wrapper(function) => function.into(), + Self::Logging(function) => function.into(), + Self::DeploymentHooks(function) => function.into(), + } + } + + pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + py.import(MODULE)?.getattr(self.name())?.call1(args) + } +} + +impl Wrapper { + pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + LegacyPython::Wrapper(self).call(py, args) + } +} + +impl Logging { + pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + LegacyPython::Logging(self).call(py, args) + } +} + +impl DeploymentHooks { + pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + LegacyPython::DeploymentHooks(self).call(py, args) + } +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeSet; + + use strum::VariantArray; + + use super::{DeploymentHooks, LegacyPython, Logging, Wrapper}; + use crate::test_support::PYTHON_CONTRACT; + + #[test] + fn every_borrowed_function_is_in_the_python_contract() { + let contract: serde_json::Map = + serde_json::from_str(PYTHON_CONTRACT).unwrap(); + let declared: BTreeSet<&str> = contract.keys().map(String::as_str).collect(); + let called: Vec<&str> = Wrapper::VARIANTS + .iter() + .map(|&function| LegacyPython::Wrapper(function)) + .chain( + Logging::VARIANTS + .iter() + .map(|&function| LegacyPython::Logging(function)), + ) + .chain( + DeploymentHooks::VARIANTS + .iter() + .map(|&function| LegacyPython::DeploymentHooks(function)), + ) + .map(LegacyPython::name) + .collect(); + assert_eq!(called.len(), declared.len(), "a function is borrowed twice"); + assert_eq!(called.into_iter().collect::>(), declared); + } +} diff --git a/litellm-rust/crates/callbacks-legacy/src/lib.rs b/litellm-rust/crates/callbacks-legacy/src/lib.rs index 06783ac255d..42ffd545e2b 100644 --- a/litellm-rust/crates/callbacks-legacy/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy/src/lib.rs @@ -2,7 +2,7 @@ //! sync and async callback registries it fans out to, the deployment hooks, the deferred //! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name //! inheritance, budget and retry-count limits). All of it sits behind one -//! [`CallbackAdapter`](litellm_host_python::CallbackAdapter), so the driver, the routes and +//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and //! core never learn which Python object is on the other end. //! //! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`] @@ -13,6 +13,7 @@ mod adapter; mod call; mod callbacks; mod deferred; +mod legacy_python; mod logger; mod preparation; #[cfg(test)] @@ -21,7 +22,7 @@ mod test_support; pub(crate) use adapter::LegacyLogging; pub use adapter::LegacySurface; -pub use call::{PublicCall, lookup, run_legacy_call}; +pub use call::{PublicCall, run_legacy_call}; pub(crate) use callbacks::{LegacyCallbacks, is_internal_call}; pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup}; pub(crate) use preparation::prepare; diff --git a/litellm-rust/crates/callbacks-legacy/src/logger.rs b/litellm-rust/crates/callbacks-legacy/src/logger.rs index a0e525000b8..061941f05b9 100644 --- a/litellm-rust/crates/callbacks-legacy/src/logger.rs +++ b/litellm-rust/crates/callbacks-legacy/src/logger.rs @@ -5,34 +5,25 @@ use pyo3::{ types::{PyDict, PyTuple}, }; -/// The `Logging` instance one call fans out through, and who owns it. A logger the caller -/// handed in is observed in full, because the caller reads it after the call; one this -/// crate built through `function_setup` is elided wherever no registry needs it. +use crate::legacy_python::{self, Wrapper}; + +/// The `Logging` instance one call fans out through. pub struct PythonLogger { object: Py, - bridge_owned: bool, } impl PythonLogger { - pub(crate) fn new(object: Py, bridge_owned: bool) -> Self { - Self { - object, - bridge_owned, - } + pub(crate) fn new(object: Py) -> Self { + Self { object } } pub(crate) fn object<'py>(&self, py: Python<'py>) -> &Bound<'py, PyAny> { self.object.bind(py) } - pub(crate) fn bridge_owned(&self) -> bool { - self.bridge_owned - } - pub fn clone_ref(&self, py: Python<'_>) -> Self { Self { object: self.object.clone_ref(py), - bridge_owned: self.bridge_owned, } } @@ -40,34 +31,17 @@ impl PythonLogger { visit.call(&self.object) } - pub fn success_bookkeeping( - &self, - py: Python<'_>, - response: &Option>, - start: &Py, - end: &Option>, - asynchronous: bool, - ) -> PyResult<()> { - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("success_bookkeeping")? - .call1((self.object(py), response, start, end, asynchronous))?; - Ok(()) - } - pub fn restore_context(&self, py: Python<'_>) -> PyResult<()> { - py.import("litellm.utils")? - .getattr("_restore_correlation_context_if_supported")? - .call1((self.object(py),))?; + Wrapper::RestoreContext.call(py, (self.object(py),))?; Ok(()) } } -/// A bare Python object was not obtained from `setup`, so it is caller-owned. impl FromPyObject<'_, '_> for PythonLogger { type Error = PyErr; fn extract(object: Borrowed<'_, '_, PyAny>) -> PyResult { - Ok(Self::new(object.to_owned().unbind(), false)) + Ok(Self::new(object.to_owned().unbind())) } } @@ -75,9 +49,7 @@ pub struct SetupResult<'py>(Bound<'py, PyAny>); impl SetupResult<'_> { pub fn logger(&self) -> PyResult { - let object = self.0.getattr("logger")?.unbind(); - let bridge_owned = self.0.getattr("bridge_owned")?.extract()?; - Ok(PythonLogger::new(object, bridge_owned)) + Ok(PythonLogger::new(self.0.getattr("logger")?.unbind())) } pub fn kwargs(&self) -> PyResult> { @@ -93,9 +65,8 @@ pub fn setup<'py>( start: &Py, asynchronous: bool, ) -> PyResult> { - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("setup")? - .call1((call_type, args, kwargs, start, asynchronous)) + Wrapper::Setup + .call(py, (call_type, args, kwargs, start, asynchronous)) .map(SetupResult) } @@ -107,30 +78,20 @@ pub fn finalize( start: &Py, end: &Option>, ) -> PyResult<()> { - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("finalize")? - .call1((response, logger.object(py), kwargs, start, end))?; + Wrapper::Finalize.call(py, (response, logger.object(py), kwargs, start, end))?; Ok(()) } pub struct DeploymentHooks; impl DeploymentHooks { - pub fn needed(py: Python<'_>) -> PyResult { - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("deployment_callbacks_needed")? - .call0()? - .extract() - } - pub fn before_call( py: Python<'_>, kwargs: &Py, call_type: &str, ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_pre_call_deployment_hook")? - .call1((kwargs, call_type)) + legacy_python::DeploymentHooks::BeforeDeploymentCall + .call(py, (kwargs, call_type)) .map(Bound::unbind) } @@ -140,9 +101,8 @@ impl DeploymentHooks { response: &Option>, call_type: &str, ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_post_call_success_deployment_hook")? - .call1((kwargs, response, call_type)) + legacy_python::DeploymentHooks::AfterDeploymentSuccess + .call(py, (kwargs, response, call_type)) .map(Bound::unbind) } @@ -152,9 +112,8 @@ impl DeploymentHooks { error: &Py, call_type: &str, ) -> PyResult> { - py.import("litellm.utils")? - .getattr("async_post_call_failure_deployment_hook")? - .call1((kwargs, error, call_type)) + legacy_python::DeploymentHooks::AfterDeploymentFailure + .call(py, (kwargs, error, call_type)) .map(Bound::unbind) } } @@ -185,10 +144,6 @@ class Setup: reads.append('logger') return logger @property - def bridge_owned(self): - reads.append('bridge_owned') - return True - @property def kwargs(self): reads.append('kwargs') return [] @@ -206,7 +161,6 @@ result = Setup() .object(py) .is(locals.get_item("logger").unwrap().unwrap()) ); - assert!(logger.bridge_owned()); assert!( result .kwargs() @@ -220,17 +174,8 @@ result = Setup() .unwrap() .extract::>() .unwrap(), - ["logger", "bridge_owned", "kwargs"] + ["logger", "kwargs"] ); }); } - - #[test] - fn a_logger_extracted_from_a_bare_object_is_caller_owned() { - Python::initialize(); - Python::attach(|py| { - let logger: PythonLogger = py.None().into_bound(py).extract().unwrap(); - assert!(!logger.bridge_owned()); - }); - } } diff --git a/litellm-rust/crates/callbacks-legacy/src/preparation.rs b/litellm-rust/crates/callbacks-legacy/src/preparation.rs index 981b1702f2e..fa1ff9acd4d 100644 --- a/litellm-rust/crates/callbacks-legacy/src/preparation.rs +++ b/litellm-rust/crates/callbacks-legacy/src/preparation.rs @@ -3,6 +3,8 @@ use pyo3::{ types::{PyDict, PyList}, }; +use crate::legacy_python::Wrapper; + struct CredentialEntry<'py>(Bound<'py, PyAny>); impl<'py> CredentialEntry<'py> { @@ -22,18 +24,19 @@ pub fn prepare<'py>( ) -> PyResult> { let arguments = kwargs.copy()?; arguments.set_item("litellm_logging_obj", logger.object(py))?; - let litellm = py.import("litellm")?; - inherit_credentials(py, &litellm, &arguments)?; - py.import("litellm.rust_bridge.legacy_callbacks")? - .getattr("check_limits")? - .call1((&arguments,))?; + inherit_credentials(py, &arguments, || { + Ok(Wrapper::CredentialList + .call(py, ())? + .cast_into::()?) + })?; + Wrapper::CheckLimits.call(py, (&arguments,))?; Ok(arguments) } -fn inherit_credentials( - py: Python<'_>, - litellm: &Bound<'_, PyModule>, - arguments: &Bound<'_, PyDict>, +fn inherit_credentials<'py>( + py: Python<'py>, + arguments: &Bound<'py, PyDict>, + credential_list: impl FnOnce() -> PyResult>, ) -> PyResult<()> { let Some(requested) = arguments .get_item("litellm_credential_name")? @@ -45,16 +48,13 @@ fn inherit_credentials( return Ok(()); } let requested: String = requested.extract()?; - let credentials = litellm.getattr("credential_list")?.cast_into::()?; + let credentials = credential_list()?; let names = credentials .iter() .map(|credential| CredentialEntry(credential).name()) .collect::>>()?; let Some(index) = names.iter().position(|name| *name == requested) else { - py.import("litellm._logging")?.getattr("verbose_logger")?.call_method1( - "warning", - ("litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", requested, names.len()), - )?; + Wrapper::WarnUnknownCredential.call(py, (requested, names.len()))?; return Ok(()); }; let selected = CredentialEntry(credentials.get_item(index)?); @@ -80,19 +80,19 @@ mod tests { } fn inherit(py: Python<'_>, locals: &Bound<'_, PyDict>) -> PyResult<()> { - let litellm = PyModule::new(py, "credential_host")?; - litellm.setattr( - "credential_list", - locals.get_item("credentials").unwrap().unwrap(), - )?; inherit_credentials( py, - &litellm, &locals .get_item("arguments") .unwrap() .unwrap() .cast_into::()?, + || { + Ok(locals + .get_item("credentials")? + .unwrap() + .cast_into::()?) + }, ) } @@ -304,11 +304,11 @@ arguments = {'litellm_credential_name': 'ocr-test'} fn falsy_credential_names_return_before_loading_credentials() { Python::initialize(); Python::attach(|py| { - let litellm = PyModule::new(py, "credential_host").unwrap(); for name in [py.None(), py.eval(c"''", None, None).unwrap().unbind()] { let arguments = PyDict::new(py); arguments.set_item("litellm_credential_name", name).unwrap(); - inherit_credentials(py, &litellm, &arguments).unwrap(); + inherit_credentials(py, &arguments, || panic!("credentials must not be loaded")) + .unwrap(); } }); } diff --git a/litellm-rust/crates/callbacks-legacy/tests/deferred.rs b/litellm-rust/crates/callbacks-legacy/tests/deferred.rs index 3daea8840d8..289ea1b2e7f 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/deferred.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/deferred.rs @@ -16,7 +16,7 @@ fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> { py, PendingLogging { pending: Some(PendingSuccess { - logger: PythonLogger::new(local(&locals, "logger").unbind(), true), + logger: PythonLogger::new(local(&locals, "logger").unbind()), response: Some(local(&locals, "response").unbind()), start: py.None(), end: Some(py.None()), @@ -79,22 +79,6 @@ assert logger.calls == [], logger.calls }); } -#[test] -fn a_release_after_the_async_callbacks_went_away_only_keeps_the_books() { - Python::initialize(); - Python::attach(|py| { - let locals = defer(py, c"logger.needed = {'async_success': False}"); - run( - py, - &locals, - c" -pending.release(True) -assert logger.calls == [('success_bookkeeping', True)], logger.calls -", - ); - }); -} - #[rstest] #[case::ordinary_error(c"RuntimeError('queue full')", false)] #[case::cancellation(c"asyncio.CancelledError()", true)] diff --git a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs index 3ceda4441a7..ea3510de17e 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs @@ -1,7 +1,7 @@ use std::ffi::CStr; use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing}; -use litellm_host_python::{AdapterStep, CallbackAdapter, PublicValue}; +use litellm_host_python::{LifecycleStep, PublicValue, PythonLifecycle}; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; use pyo3::types::PyDict; @@ -24,7 +24,7 @@ fn begin<'py>( py: Python<'py>, locals: &Bound<'py, PyDict>, asynchronous: bool, -) -> (LegacyLogging, AdapterStep) { +) -> (LegacyLogging, LifecycleStep) { let mut logging = legacy_call(py, locals, asynchronous); let kwargs = local(locals, "kwargs") .cast_into::() @@ -34,15 +34,15 @@ fn begin<'py>( (logging, step) } -fn arguments<'py>(py: Python<'py>, step: AdapterStep) -> Bound<'py, PyDict> { - let AdapterStep::Arguments(arguments) = step else { +fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> { + let LifecycleStep::Arguments(arguments) = step else { panic!("expected the prepared arguments"); }; arguments.into_bound(py) } -fn awaits_deployment_hook(step: &AdapterStep) -> bool { - matches!(step, AdapterStep::Await(_)) +fn awaits_deployment_hook(step: &LifecycleStep) -> bool { + matches!(step, LifecycleStep::Await(_)) } #[rstest] @@ -121,7 +121,7 @@ logger.hooks = {'pre': lambda kwargs: kwargs} let step = logging .resume(py, Ok(local(&locals, "replacement").unbind())) .unwrap(); - let AdapterStep::Response(returned) = step else { + let LifecycleStep::Response(returned) = step else { panic!("expected the finalized response"); }; assert!(returned.bind(py).is(local(&locals, "replacement"))); @@ -195,7 +195,7 @@ fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelle }; assert!(matches!( logging.resume(py, hook_result).unwrap(), - AdapterStep::Await(_) + LifecycleStep::Await(_) )); run( py, @@ -237,7 +237,7 @@ kwargs = {'logger': logger} .unwrap() .unbind(); let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step { - AdapterStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())), + LifecycleStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())), step => Ok(step), }); let error = result.err().unwrap(); diff --git a/litellm-rust/crates/callbacks-legacy/tests/payload.rs b/litellm-rust/crates/callbacks-legacy/tests/payload.rs index 480bedf8548..68c0a2b1e15 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/payload.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/payload.rs @@ -1,7 +1,8 @@ use std::ffi::CStr; -use litellm_callbacks::event::{CallEvent, Passthrough, RawResponse, RequestContext, WireRequest}; -use litellm_host_python::{AdapterStep, CallbackAdapter}; +use litellm_auth::SecretValue; +use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest}; +use litellm_host_python::{LifecycleStep, PythonLifecycle}; use pyo3::prelude::*; use rstest::rstest; use serde_json::{Value, json}; @@ -23,20 +24,12 @@ class PayloadLogger(StubLogger): def pre_call(self, input, api_key, additional_args): self.record('pre_call', None) self.pre = additional_args + self.pre_api_key = api_key on_pre_call(additional_args) - def _pre_call(self, input, api_key, additional_args): - self.record('_pre_call', None) - - def record_api_call_start_time(self): - self.record('record_api_call_start_time', None) - - def post_call(self, original_response, additional_args): + def post_call(self, original_response, api_key, additional_args): self.record('post_call', None) - self.post = (original_response, additional_args) - - def record_post_call(self, response, *rest): - self.record('record_post_call', response) + self.post = (original_response, api_key, additional_args) request = Request() kwargs = {} @@ -52,16 +45,16 @@ fn document(source: &str) -> Value { json!({"type": "document_url", "document_url": source}) } -fn before_send(script: &CStr, caller: Value, body: Value) -> WireRequest { - before_send_with_secrets(script, caller, body, &[]) +fn before_send(script: &CStr, body: Value) -> WireRequest { + before_send_with_secrets(script, json!({}), body, &[]) } -/// Runs `before_send` over `body` for a caller whose route-side view is `caller`, with the -/// Python objects `script` binds, then delivers the provider's raw response the way the +/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with +/// the Python objects `script` binds, then delivers the provider's raw response the way the /// driver does and runs the script's `check()`. fn before_send_with_secrets( script: &CStr, - caller: Value, + optional_params: Value, body: Value, secret_fields: &[&str], ) -> WireRequest { @@ -70,15 +63,15 @@ fn before_send_with_secrets( let locals = namespace(py, PAYLOAD_LOGGER); run(py, &locals, script); let mut logging = LegacyLogging { - logger: Some(PythonLogger::new(local(&locals, "logger").unbind(), true)), + logger: Some(PythonLogger::new(local(&locals, "logger").unbind())), ..legacy_call(py, &locals, false) }; let context = RequestContext { model: "model".into(), custom_llm_provider: "provider".into(), - optional_params: caller.clone(), - passthrough_fields: Passthrough::unchanged(caller.as_object().unwrap(), &body), + optional_params, secret_fields: secret_fields.iter().map(|name| name.to_string()).collect(), + api_key: Some(SecretValue::new("route-key")), }; let wire = WireRequest { url: "https://provider.invalid/ocr".into(), @@ -93,10 +86,10 @@ fn before_send_with_secrets( }; assert!(matches!( logging.emit(py, &raw, None).unwrap(), - AdapterStep::Done + LifecycleStep::Done )); run(py, &locals, c"check()"); - let AdapterStep::Wire(wire) = step else { + let LifecycleStep::Wire(wire) = step else { panic!("before_send did not hand back the wire request"); }; *wire @@ -129,11 +122,7 @@ def check(): ")] fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) { let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]}); - let wire = before_send( - script, - json!({"document": document(DOCUMENT), "pages": [0]}), - body.clone(), - ); + let wire = before_send(script, body.clone()); assert_eq!(wire.body, body); } @@ -149,7 +138,6 @@ def check(): assert document['document_url'] == 'data:application/pdf;base64,ZWRpdGVk' ", json!({"document": document(DOCUMENT)}), - json!({"document": document(DOCUMENT)}), ); assert_eq!(wire.body["document"], document(EDITED)); } @@ -168,7 +156,6 @@ def check(): assert observed == [False], observed assert document == {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'} ", - json!({"document": document("https://example.invalid/scan.pdf")}), json!({"document": document(DOCUMENT)}), ); assert_eq!( @@ -177,6 +164,23 @@ def check(): ); } +#[test] +fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() { + let body = json!({"pages": [0]}); + let wire = before_send( + c" +opaque = object() +kwargs = {'pages': opaque} +observed = [] +on_pre_call = lambda args: observed.append(args['complete_input_dict']['pages']) +def check(): + assert observed == [[0]], observed +", + body.clone(), + ); + assert_eq!(wire.body, body); +} + #[rstest] #[case::body( c" @@ -192,7 +196,7 @@ def on_pre_call(args): )] fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) { let body = json!({"document": document(DOCUMENT)}); - let wire = before_send(script, json!({}), body.clone()); + let wire = before_send(script, body.clone()); assert_eq!(wire.body, body); assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]); } @@ -205,7 +209,6 @@ def on_pre_call(args): args['headers']['x-callback'] = 'edited' ", json!({}), - json!({}), ); assert_eq!( wire.headers, @@ -288,7 +291,7 @@ def on_pre_call(args): )] fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) { let body = json!({"document": document(DOCUMENT)}); - let wire = before_send(script, json!({"document": document(DOCUMENT)}), body); + let wire = before_send(script, body); assert_eq!(wire.body, expected); } @@ -302,7 +305,6 @@ def on_pre_call(args): retained['x-retained'] = 'sent' ", json!({}), - json!({}), ); assert_eq!( wire.headers, @@ -314,52 +316,33 @@ def on_pre_call(args): } #[test] -fn post_call_receives_the_raw_response_and_the_payload_dicts_pre_call_saw() { +fn post_call_receives_the_raw_response_the_route_key_and_the_body_pre_call_saw() { before_send( c" def check(): - original_response, additional_args = logger.post + original_response, api_key, additional_args = logger.post assert original_response == 'raw response', original_response + assert api_key == logger.pre_api_key == 'route-key', (api_key, logger.pre_api_key) + assert additional_args == {'complete_input_dict': logger.pre['complete_input_dict']}, additional_args assert additional_args['complete_input_dict'] is logger.pre['complete_input_dict'] - assert additional_args['headers'] is logger.pre['headers'] ", - json!({}), json!({"document": document(DOCUMENT)}), ); } -#[rstest] -#[case::every_phase_listens(c"{}", &["pre_call", "post_call"])] -#[case::no_input_callback( - c"{'input': False}", - &["_pre_call", "record_api_call_start_time", "record_post_call"] -)] -#[case::no_payload_consumer(c"{'payload': False}", &["record_api_call_start_time"])] -fn payload_callbacks_run_only_for_the_phases_someone_listens_to( - #[case] needed: &CStr, - #[case] expected_calls: &[&str], -) { - let script = std::ffi::CString::new(format!( - " -logger.needed = {needed} +#[test] +fn every_request_runs_the_full_pre_call_and_post_call() { + let wire = before_send( + c" def on_pre_call(args): args['complete_input_dict']['include_image_base64'] = True def check(): - assert logger.names() == {expected_calls:?}, logger.calls + assert logger.names() == ['pre_call', 'post_call'], logger.calls ", - needed = needed.to_str().unwrap(), - expected_calls = expected_calls, - )) - .unwrap(); - let body = json!({"document": document(DOCUMENT)}); - let wire = before_send(&script, json!({}), body.clone()); - let edited = json!({"document": document(DOCUMENT), "include_image_base64": true}); + json!({"document": document(DOCUMENT)}), + ); assert_eq!( wire.body, - if expected_calls.contains(&"pre_call") { - edited - } else { - body - } + json!({"document": document(DOCUMENT), "include_image_base64": true}) ); } diff --git a/litellm-rust/crates/callbacks-legacy/tests/support.rs b/litellm-rust/crates/callbacks-legacy/tests/support.rs index 1663e11963e..444655ea77b 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/support.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/support.rs @@ -5,64 +5,89 @@ use pyo3::types::{PyDict, PyTuple}; use crate::{LegacyLogging, LegacySurface, PublicCall}; -/// Stand-ins for every litellm function the legacy contract calls. Tests share one -/// interpreter and run concurrently, so each stub is installed idempotently and forwards to -/// the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`). +/// The parameters of every `legacy_callbacks` function, as the real module declares them. +/// `tests/test_litellm/rust_bridge/test_legacy_callbacks.py` pins this file to the Python +/// signatures, and [`namespace`] binds every fake call against it. +pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json"); + +/// Stand-ins for `legacy_callbacks`, the only Python module the crate calls. Tests +/// share one interpreter and run concurrently, so each fake is installed idempotently and +/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`). +/// Every fake is bound against the contract first, so a call the real module would reject +/// fails here too. const STUBS: &CStr = c" import contextvars +import inspect +import json import sys +import traceback import types -for name in ( - 'litellm', - 'litellm.utils', - 'litellm.types', - 'litellm.types.utils', - 'litellm._internal_context', - 'litellm.litellm_core_utils', - 'litellm.litellm_core_utils.logging_worker', - 'litellm.litellm_core_utils.litellm_logging', - 'litellm.rust_bridge', - 'litellm.rust_bridge.legacy_callbacks', -): +for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.legacy_callbacks'): sys.modules.setdefault(name, types.ModuleType(name)) legacy = sys.modules['litellm.rust_bridge.legacy_callbacks'] -legacy.setup = lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace( - logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'], - kwargs=kwargs, - bridge_owned=True, -) -legacy.deployment_callbacks_needed = lambda: True -legacy.check_limits = lambda arguments: arguments['logger'].check_limits(arguments) -legacy.callbacks_needed = lambda logger, phase: logger.needed.get(phase, True) -legacy.success_bookkeeping = lambda logger, response, start, end, asynchronous: logger.record( - 'success_bookkeeping', asynchronous -) -legacy.failure_bookkeeping = lambda logger, error, start, end, asynchronous: logger.record( - 'failure_bookkeeping', asynchronous -) -legacy.finalize = lambda response, logger, kwargs, start, end: logger.record('finalize', response) +CONTRACT = json.loads(python_contract) -utils = sys.modules['litellm.utils'] -utils.async_pre_call_deployment_hook = lambda kwargs, call_type: kwargs['logger'].hook( - 'pre', kwargs, call_type -) -utils.async_post_call_success_deployment_hook = lambda kwargs, response, call_type: kwargs[ - 'logger' -].hook('success', response, call_type) -utils.async_post_call_failure_deployment_hook = lambda kwargs, error, call_type: kwargs[ - 'logger' -].hook('failure', error, call_type) -utils._restore_correlation_context_if_supported = lambda logger: logger.record('restore', None) -internal = sys.modules['litellm._internal_context'] -if not hasattr(internal, 'is_internal_call'): - internal.is_internal_call = contextvars.ContextVar('is_internal_call', default=False) +def contracted(name, fake): + signature = inspect.Signature( + [inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]] + ) -sys.modules['litellm.types.utils'].CustomPricingLiteLLMParams = type( - 'CustomPricingLiteLLMParams', (), {'model_fields': {'ocr_cost_per_page': None}} -) + def checked(*args, **kwargs): + signature.bind(*args, **kwargs) + return fake(*args, **kwargs) + + return checked + + +if not hasattr(legacy, 'is_internal'): + legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False) + +FAKES = { + 'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace( + logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'], + kwargs=kwargs, + ), + 'check_limits': lambda arguments: arguments['logger'].check_limits(arguments), + 'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response), + 'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider=provider, + ), + 'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args), + 'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call( + original_response, api_key, additional_args + ), + 'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)), + 'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending), + 'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls( + response, start, end + ), + 'failure_handler': lambda logger, error, start, end, asynchronous: ( + logger.async_failure_handler if asynchronous else logger.failure_handler + )(error, ''.join(traceback.format_exception(error)), start, end), + 'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)), + 'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end), + 'enqueue_logging': lambda coroutine: coroutine.enqueue(), + 'restore_context': lambda logger: logger.record('restore', None), + 'custom_pricing_fields': lambda: ('ocr_cost_per_page',), + 'is_internal_call': lambda: legacy.is_internal.get(), + 'credential_list': lambda: [], + 'warn_unknown_credential': lambda name, loaded: None, + 'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type), + 'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook( + 'success', response, call_type + ), + 'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type), +} +assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys()) +for name, fake in FAKES.items(): + setattr(legacy, name, contracted(name, fake)) unraisable = sys.modules.setdefault( @@ -77,20 +102,6 @@ def unraisable_from(owner): return [error for source, error in unraisable.events if source is owner] -class Worker: - def ensure_initialized_and_enqueue(self, coroutine): - return coroutine.enqueue() - - -class Executor: - def submit(self, run, handler, *args): - handler.__self__.record('submit', args) - - -sys.modules['litellm.litellm_core_utils.logging_worker'].GLOBAL_LOGGING_WORKER = Worker() -sys.modules['litellm.litellm_core_utils.litellm_logging'].executor = Executor() - - class StubCoroutine: def __init__(self, logger): self.logger = logger @@ -106,7 +117,6 @@ class StubCoroutine: class StubLogger: def __init__(self): self.calls = [] - self.needed = {} self.hooks = {} self.on_enqueue = lambda coroutine: None @@ -147,6 +157,7 @@ logger = StubLogger() /// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it. pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> { let locals = PyDict::new(py); + locals.set_item("python_contract", PYTHON_CONTRACT).unwrap(); py.run(STUBS, Some(&locals), Some(&locals)).unwrap(); py.run(script, Some(&locals), Some(&locals)).unwrap(); locals diff --git a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs index 9b9d29108f6..3094d7b88d2 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs @@ -1,7 +1,7 @@ use std::ffi::CStr; use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing}; -use litellm_host_python::{AdapterStep, CallbackAdapter, PublicValue}; +use litellm_host_python::{LifecycleStep, PublicValue, PythonLifecycle}; use pyo3::exceptions::PyRuntimeError; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; @@ -19,12 +19,16 @@ const TIMING: Timing = Timing { fn logged(py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool) -> LegacyLogging { LegacyLogging { - logger: Some(PythonLogger::new(local(locals, "logger").unbind(), true)), + logger: Some(PythonLogger::new(local(locals, "logger").unbind())), ..legacy_call(py, locals, asynchronous) } } -fn succeed(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> AdapterStep { +fn succeed( + py: Python<'_>, + locals: &Bound<'_, PyDict>, + logging: &mut LegacyLogging, +) -> LifecycleStep { let response = local(locals, "response").unbind(); logging .emit( @@ -35,7 +39,7 @@ fn succeed(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLoggi .unwrap() } -fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> AdapterStep { +fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> LifecycleStep { let failure = PyErr::from_value(local(locals, "failure")); logging .emit( @@ -51,20 +55,14 @@ fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) #[rstest] #[case::sync_listened(false, c"", &["submit"])] -#[case::sync_unlistened(false, c"logger.needed = {'sync_success': False}", &["success_bookkeeping"])] #[case::async_listened( true, c"", &["async_success_handler", "enqueued", "sync_success_for_async_call"] )] -#[case::async_unlistened( - true, - c"logger.needed = {'async_success': False, 'sync_success_async': False}", - &["success_bookkeeping"] -)] #[case::async_deferred(true, c"logger._defer_async_logging = True", &["sync_success_for_async_call"])] #[case::async_with_fallbacks(true, c"kwargs = {'fallbacks': ['other']}", &["sync_success_for_async_call"])] -fn success_reaches_only_the_callbacks_that_listen( +fn success_reaches_the_logging_handlers( #[case] asynchronous: bool, #[case] script: &CStr, #[case] expected: &[&str], @@ -76,7 +74,7 @@ fn success_reaches_only_the_callbacks_that_listen( let mut logging = logged(py, &locals, asynchronous); assert!(matches!( succeed(py, &locals, &mut logging), - AdapterStep::Done + LifecycleStep::Done )); let names: Vec = local(&locals, "logger") .call_method0("names") @@ -109,7 +107,10 @@ fn internal_calls_skip_failure_callbacks_only_when_asynchronous( internal: true, ..logged(py, &locals, asynchronous) }; - assert!(matches!(fail(py, &locals, &mut logging), AdapterStep::Done)); + assert!(matches!( + fail(py, &locals, &mut logging), + LifecycleStep::Done + )); let names: Vec = local(&locals, "logger") .call_method0("names") .unwrap() @@ -157,7 +158,7 @@ logger = FailingLogger() let mut logging = logged(py, &locals, true); assert!(matches!( succeed(py, &locals, &mut logging), - AdapterStep::Done + LifecycleStep::Done )); assert!( logging @@ -173,14 +174,8 @@ logger = FailingLogger() #[rstest] #[case::sync_listened(false, c"", &["failure_handler"])] -#[case::sync_unlistened(false, c"logger.needed = {'sync_failure': False}", &["failure_bookkeeping"])] #[case::async_listened(true, c"", &["failure_handler", "async_failure_handler"])] -#[case::async_unlistened( - true, - c"logger.needed = {'sync_failure': False, 'async_failure': False}", - &["failure_bookkeeping", "failure_bookkeeping"] -)] -fn failure_reaches_only_the_callbacks_that_listen( +fn failure_reaches_the_logging_handlers( #[case] asynchronous: bool, #[case] script: &CStr, #[case] expected: &[&str], @@ -192,7 +187,10 @@ fn failure_reaches_only_the_callbacks_that_listen( let mut logging = logged(py, &locals, asynchronous); let step = fail(py, &locals, &mut logging); let awaits_async_handler = expected.contains(&"async_failure_handler"); - assert_eq!(matches!(step, AdapterStep::Await(_)), awaits_async_handler); + assert_eq!( + matches!(step, LifecycleStep::Await(_)), + awaits_async_handler + ); let names: Vec = local(&locals, "logger") .call_method0("names") .unwrap() @@ -227,7 +225,7 @@ logger = FailingLogger() let mut logging = logged(py, &locals, true); assert!(matches!( fail(py, &locals, &mut logging), - AdapterStep::Await(_) + LifecycleStep::Await(_) )); assert!( logging @@ -265,7 +263,7 @@ fn the_async_failure_handler_ends_the_call_unless_it_was_cancelled( }; let expected = result.as_ref().err().map(|error| error.value(py).clone()); match logging.resume(py, result) { - Ok(step) => assert!(done && matches!(step, AdapterStep::Done)), + Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)), Err(propagated) => { assert!(!done); assert!(propagated.value(py).is(expected.unwrap())); diff --git a/litellm-rust/crates/callbacks/Cargo.toml b/litellm-rust/crates/callbacks/Cargo.toml index 4b966271478..a68ebc26a8d 100644 --- a/litellm-rust/crates/callbacks/Cargo.toml +++ b/litellm-rust/crates/callbacks/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-auth.workspace = true serde_json.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/callbacks/src/event.rs b/litellm-rust/crates/callbacks/src/event.rs index e6f88fd9709..3bf12e553b9 100644 --- a/litellm-rust/crates/callbacks/src/event.rs +++ b/litellm-rust/crates/callbacks/src/event.rs @@ -1,6 +1,6 @@ use std::time::{SystemTime, UNIX_EPOCH}; -use serde_json::{Map, Value}; +use serde_json::Value; /// Seconds since the Unix epoch, on one clock for every host. pub fn epoch_seconds() -> f64 { @@ -33,34 +33,10 @@ pub struct RequestContext { pub custom_llm_provider: String, /// The route's parameters before the provider transformation. pub optional_params: Value, - pub passthrough_fields: Passthrough, /// Optional-param names that carry credentials and must be redacted when logged. pub secret_fields: Vec, -} - -/// Body keys whose values are the caller's inputs, unchanged by the route. The only way to -/// build one is to compare the two, so a route cannot name a key it rewrote. -#[derive(Clone, Debug, Default, PartialEq, Eq)] -pub struct Passthrough(Vec); - -impl Passthrough { - pub fn unchanged(caller: &Map, body: &Value) -> Self { - Self( - caller - .iter() - .filter(|(name, value)| body.get(name.as_str()) == Some(*value)) - .map(|(name, _)| name.clone()) - .collect(), - ) - } - - pub fn iter(&self) -> impl Iterator { - self.0.iter().map(String::as_str) - } - - pub fn contains(&self, name: &str) -> bool { - self.0.iter().any(|field| field == name) - } + /// The credential the route resolved for the provider call. + pub api_key: Option, } #[derive(Clone, Debug, PartialEq, Eq)] @@ -78,6 +54,9 @@ pub enum FailureOrigin { #[derive(Clone, Debug, PartialEq)] pub enum CallEvent { + Started { + start_time: f64, + }, ResponseReceived { raw: RawResponse, }, @@ -89,47 +68,3 @@ pub enum CallEvent { origin: FailureOrigin, }, } - -#[cfg(test)] -mod tests { - use rstest::rstest; - use serde_json::json; - - use super::*; - - #[rstest] - #[case::unchanged_scalar(json!({"pages": [0]}), json!({"pages": [0]}), &["pages"])] - #[case::unchanged_explicit_null(json!({"pages": null}), json!({"pages": null}), &["pages"])] - #[case::unchanged_nested_object( - json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}}), - json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}, "model": "m"}), - &["document"] - )] - #[case::rewritten_value( - json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}}), - json!({"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}}), - &[] - )] - #[case::dropped_nested_field( - json!({"document": {"type": "image_url", "image_url": "https://a/b.png", "document_name": "b.png"}}), - json!({"document": {"type": "image_url", "image_url": "https://a/b.png"}}), - &[] - )] - #[case::added_nested_field( - json!({"document": {"type": "image_url", "image_url": "https://a/b.png"}}), - json!({"document": {"type": "image_url", "image_url": "https://a/b.png", "detail": "high"}}), - &[] - )] - #[case::reordered_array(json!({"pages": [0, 1]}), json!({"pages": [1, 0]}), &[])] - #[case::consumed_by_the_route(json!({"api_key": "k", "pages": [0]}), json!({"pages": [0]}), &["pages"])] - #[case::added_by_the_route(json!({}), json!({"model": "m"}), &[])] - #[case::non_object_body(json!({"pages": [0]}), json!([{"pages": [0]}]), &[])] - fn passthrough_is_exactly_the_callers_unchanged_keys( - #[case] caller: Value, - #[case] body: Value, - #[case] expected: &[&str], - ) { - let passthrough = Passthrough::unchanged(caller.as_object().unwrap(), &body); - assert_eq!(passthrough.iter().collect::>(), expected); - } -} diff --git a/litellm-rust/crates/callbacks/src/run.rs b/litellm-rust/crates/callbacks/src/run.rs index 57bf134f345..5accaa2e25e 100644 --- a/litellm-rust/crates/callbacks/src/run.rs +++ b/litellm-rust/crates/callbacks/src/run.rs @@ -11,6 +11,7 @@ where H: Host, { let start_time = epoch_seconds(); + let _ = host.emit(&CallEvent::Started { start_time }).await; let mut result = None; let outcome = loop { let step = match machine.resume(result.take()).await { @@ -102,6 +103,7 @@ mod tests { async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { self.seen.lock().unwrap().push(match event { + CallEvent::Started { .. } => "started".into(), CallEvent::Succeeded { .. } => "succeeded".into(), CallEvent::Failed { .. } => "failed".into(), other => format!("{other:?}"), @@ -124,7 +126,7 @@ mod tests { assert_eq!(outcome, Ok(())); assert_eq!( *host.seen.lock().unwrap(), - ["route:project", "route:send", "succeeded"] + ["started", "route:project", "route:send", "succeeded"] ); } @@ -133,7 +135,7 @@ mod tests { let host = Recording::default(); let outcome = run(scripted(&[], Err("boom")), &host).await; assert_eq!(outcome, Err("boom")); - assert_eq!(*host.seen.lock().unwrap(), ["failed"]); + assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]); let host = Recording { fail: Some("send"), @@ -143,7 +145,35 @@ mod tests { assert_eq!(outcome, Err("host failed")); assert_eq!( *host.seen.lock().unwrap(), - ["route:project", "route:send", "failed"] + ["started", "route:project", "route:send", "failed"] ); } + + struct StartTimes(Mutex>); + + impl Host for StartTimes { + async fn route(&self, _: &'static str) -> Result<(), &'static str> { + Ok(()) + } + + async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> { + if let CallEvent::Started { start_time } + | CallEvent::Succeeded { + timing: Timing { start_time, .. }, + } = event + { + self.0.lock().unwrap().push(*start_time); + } + Err("observer failed") + } + } + + #[tokio::test] + async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() { + let host = StartTimes(Mutex::default()); + assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(())); + let times = host.0.lock().unwrap(); + assert_eq!(times.len(), 2); + assert_eq!(times[0], times[1]); + } } diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 33cb8a8d32a..a6af190eb91 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,5 +1,6 @@ use futures_util::future::BoxFuture; -use litellm_callbacks::event::{CallEvent, Passthrough, RawResponse, RequestContext, WireRequest}; +use litellm_auth::SecretValue; +use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest}; use litellm_llms::{ base_llm::ocr::{ error::Error, @@ -36,6 +37,7 @@ pub(crate) struct OcrCallHooks { custom_llm_provider: &'static str, optional_params: Value, secret_fields: Vec, + api_key: Option, } impl OcrCallHooks { @@ -51,22 +53,19 @@ impl OcrCallHooks { .filter(|name| is_secret_param(name)) .cloned() .collect(), + api_key: request.connection.api_key.clone(), } } } impl CallHooks for OcrCallHooks { - fn before_send( - &self, - wire: WireRequest, - passthrough_fields: Passthrough, - ) -> BoxFuture<'_, Result> { + fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { let context = RequestContext { model: self.model.clone(), custom_llm_provider: self.custom_llm_provider.into(), optional_params: self.optional_params.clone(), - passthrough_fields, secret_fields: self.secret_fields.clone(), + api_key: self.api_key.clone(), }; Box::pin(self.host.before_send(wire, context)) } diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index e7f77acc3f8..c977f721a70 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -21,8 +21,8 @@ mod cohere_tests; #[path = "../../tests/deepseek_ocr.rs"] mod deepseek_tests; #[cfg(test)] -#[path = "../../tests/ocr/passthrough.rs"] -mod passthrough_tests; +#[path = "../../tests/ocr/document.rs"] +mod document_tests; #[cfg(test)] #[path = "../../tests/reducto_ocr.rs"] mod reducto_tests; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 24c3f43e2b4..8ac038290b7 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -1,4 +1,4 @@ -use litellm_auth::{InputSource, Sourced}; +use litellm_auth::{InputSource, SecretValue, Sourced}; use litellm_llms::base_llm::ocr::transformation::{ OcrConnection, OcrCredentialInputs, PreparedOcrRequest, credential_env, }; @@ -22,7 +22,7 @@ pub(crate) fn prepare_request( .config .get_api_key_env_var() .and_then(credential_env) - .map(|value| Sourced::new(value, InputSource::Environment)) + .map(|value| Sourced::new(SecretValue::new(value), InputSource::Environment)) }) }); let dynamic_api_base = credentials.dynamic_api_base.or_else(|| { diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index d12b8cfee95..14b34ea4564 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -277,12 +277,18 @@ mod tests { #[test] fn connection_resolution_preserves_dynamic_precedence_and_input_sources() { let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs { - api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)), + api_key: Some(Sourced::new( + litellm_auth::SecretValue::new("explicit-key"), + InputSource::Deployment, + )), api_base: Some(Sourced::new( "https://explicit.test".into(), InputSource::Deployment, )), - dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)), + dynamic_api_key: Some(Sourced::new( + litellm_auth::SecretValue::new("dynamic-key"), + InputSource::Environment, + )), dynamic_api_base: Some(Sourced::new( "https://dynamic.test".into(), InputSource::Request, @@ -292,7 +298,7 @@ mod tests { connection .api_key .as_ref() - .map(|value| value.value().as_str()), + .map(|value| value.value().expose()), Some("dynamic-key") ); assert_eq!( @@ -318,22 +324,31 @@ mod tests { fn empty_or_missing_dynamic_credentials_preserve_explicit_values( #[case] dynamic_value: Option<&str>, ) { - let dynamic = + let dynamic_key = dynamic_value.map(|value| { + Sourced::new( + litellm_auth::SecretValue::new(value), + InputSource::Environment, + ) + }); + let dynamic_base = dynamic_value.map(|value| Sourced::new(value.into(), InputSource::Environment)); let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs { - api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)), + api_key: Some(Sourced::new( + litellm_auth::SecretValue::new("explicit-key"), + InputSource::Deployment, + )), api_base: Some(Sourced::new( "https://explicit.test".into(), InputSource::Deployment, )), - dynamic_api_key: dynamic.clone(), - dynamic_api_base: dynamic, + dynamic_api_key: dynamic_key, + dynamic_api_base: dynamic_base, }); assert_eq!( connection .api_key .as_ref() - .map(|value| value.value().as_str()), + .map(|value| value.value().expose()), Some("explicit-key") ); assert_eq!( @@ -356,11 +371,18 @@ mod tests { ) { let connection = OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params( OcrCredentialInputs { - api_key: explicit_key - .map(|value| Sourced::new(value.into(), InputSource::Deployment)), + api_key: explicit_key.map(|value| { + Sourced::new( + litellm_auth::SecretValue::new(value), + InputSource::Deployment, + ) + }), api_base: explicit_base .map(|value| Sourced::new(value.into(), InputSource::Deployment)), - dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)), + dynamic_api_key: Some(Sourced::new( + litellm_auth::SecretValue::new("dynamic-key"), + InputSource::Environment, + )), dynamic_api_base: Some(Sourced::new( "https://dynamic.test".into(), InputSource::Deployment, @@ -371,7 +393,7 @@ mod tests { connection .api_key .as_ref() - .map(|value| value.value().as_str()), + .map(|value| value.value().expose()), explicit_key.map(|_| "dynamic-key") ); assert_eq!( diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 75202ed52a5..6316088dec8 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -1,7 +1,7 @@ use std::{collections::BTreeMap, path::PathBuf, time::Duration}; use bytes::Bytes; -use litellm_auth::{InputSource, TokenProviderHandle}; +use litellm_auth::{InputSource, SecretValue, TokenProviderHandle}; use litellm_core_utils::call_arguments::CallArguments; use litellm_llms::base_llm::ocr::{ error::Error, @@ -56,7 +56,7 @@ pub struct OcrFileContent { /// credentials, and per-field provenance in `input_sources`. #[derive(Clone, Debug, Default)] pub struct OcrConnectionInputs { - pub api_key: Option, + pub api_key: Option, pub api_base: Option, pub extra_headers: Map, pub timeout: Option, @@ -237,6 +237,16 @@ mod tests { .unwrap() } + #[test] + fn connection_inputs_debug_hides_the_api_key() { + let inputs = OcrConnectionInputs { + api_key: Some(SecretValue::new("caller-api-key")), + ..OcrConnectionInputs::default() + }; + + assert!(!format!("{inputs:?}").contains("caller-api-key")); + } + #[test] fn from_inputs_applies_connection_overrides_with_field_sources() { let request = LiteLLMOcrRequest::from_inputs( @@ -245,7 +255,7 @@ mod tests { None, Default::default(), OcrConnectionInputs { - api_key: Some(" key ".into()), + api_key: Some(SecretValue::new(" key ")), api_base: Some("".into()), extra_headers: json!({"x-a": "1"}).as_object().unwrap().clone(), timeout: Some(Duration::from_secs(7)), @@ -259,7 +269,7 @@ mod tests { .unwrap(); let api_key = request.credentials.api_key.as_ref().unwrap(); - assert_eq!(api_key.clone().into_value(), "key"); + assert_eq!(api_key.value().expose(), "key"); assert_eq!(api_key.source(), InputSource::Request); assert!(request.credentials.api_base.is_none()); assert_eq!( diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index 29345e38885..b9c60f57e3c 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -1,6 +1,6 @@ use std::{collections::BTreeMap, time::Duration}; -use litellm_auth::InputSource; +use litellm_auth::{InputSource, SecretValue}; use litellm_llms::base_llm::ocr::{ error::Error, transformation::{OcrDocument, decode_request_value}, @@ -44,7 +44,7 @@ pub fn consumed_optional_param_names( pub struct OcrWireRequest { pub model: String, pub document: D, - pub api_key: Option, + pub api_key: Option, pub api_base: Option, pub custom_llm_provider: Option, pub extra_headers: Option>, diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index 01a4e5efb3b..1fc4d6c2b9e 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -69,7 +69,7 @@ async fn rejects_invalid_pages_features_and_format( let result = decode_request(OcrWireRequest { model: "azure_ai/doc-intelligence/prebuilt-read".into(), document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), - api_key: Some("key".into()), + api_key: Some(litellm_auth::SecretValue::new("key")), api_base: Some(base), custom_llm_provider: None, extra_headers: None, diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 1f591d74d5d..ca2e14a7f0d 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -81,7 +81,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() { let request = OcrWireRequest { model: "mistral/model".into(), document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some("key".into()), + api_key: Some(litellm_auth::SecretValue::new("key")), api_base: None, custom_llm_provider: None, extra_headers: None, @@ -97,7 +97,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() { decode_request(OcrWireRequest { model: "model".into(), document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some("key".into()), + api_key: Some(litellm_auth::SecretValue::new("key")), api_base: None, custom_llm_provider: Some("unknown".into()), extra_headers: None, @@ -194,6 +194,7 @@ async fn facade_uses_the_injected_http_client() { fn event_name(event: &CallEvent) -> &'static str { match event { + CallEvent::Started { .. } => "started", CallEvent::ResponseReceived { .. } => "response", CallEvent::Succeeded { .. } => "success", CallEvent::Failed { .. } => "failure", @@ -235,7 +236,7 @@ async fn lifecycle_sends_headers_returned_by_the_before_send_operation() { } #[tokio::test] -async fn before_send_context_names_passthrough_fields_and_secrets() { +async fn before_send_context_names_the_route_and_its_secrets() { let (base, _, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; let observed = Arc::new(Mutex::new(None)); let captured = observed.clone(); @@ -254,8 +255,6 @@ async fn before_send_context_names_passthrough_fields_and_secrets() { assert_eq!(context.custom_llm_provider, "mistral"); assert_eq!(context.model, "model"); assert_eq!(wire.body["pages"], json!([0])); - assert!(context.passthrough_fields.contains("pages")); - assert!(context.passthrough_fields.contains("document")); assert!(context.secret_fields.is_empty()); assert_eq!(context.optional_params["req_format"], "native"); @@ -279,7 +278,6 @@ async fn before_send_context_names_passthrough_fields_and_secrets() { perform_ocr_with(host).await.unwrap(); server.await.unwrap(); let context = observed.lock().unwrap().take().unwrap(); - assert!(!context.passthrough_fields.contains("document")); assert_eq!(context.secret_fields, ["client_secret"]); } @@ -296,7 +294,7 @@ async fn lifecycle_orders_hooks_and_emits_one_success() { server.await.unwrap(); assert_eq!( *events.lock().unwrap(), - ["before_send", "response", "success"] + ["started", "before_send", "response", "success"] ); assert_eq!(seen.lock().unwrap().len(), 1); } @@ -311,7 +309,10 @@ async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() { ); let error = perform_ocr_with(host).await.unwrap_err(); assert!(matches!(error, OcrError::InvalidRequest(message) if message == "blocked")); - assert_eq!(*events.lock().unwrap(), ["before_send", "failure"]); + assert_eq!( + *events.lock().unwrap(), + ["started", "before_send", "failure"] + ); } #[tokio::test] @@ -330,7 +331,10 @@ async fn upstream_failure_emits_one_terminal_failure() { ); assert!(perform_ocr_with(host).await.is_err()); server.await.unwrap(); - assert_eq!(*events.lock().unwrap(), ["before_send", "failure"]); + assert_eq!( + *events.lock().unwrap(), + ["started", "before_send", "failure"] + ); assert_eq!(seen.lock().unwrap().len(), 1); } diff --git a/litellm-rust/crates/core/tests/ocr/document.rs b/litellm-rust/crates/core/tests/ocr/document.rs new file mode 100644 index 00000000000..5e10ce3ad2a --- /dev/null +++ b/litellm-rust/crates/core/tests/ocr/document.rs @@ -0,0 +1,152 @@ +use litellm_callbacks::event::WireRequest; +use litellm_llms::base_llm::ocr::error::Error; +use rstest::rstest; +use serde_json::{Value, json}; + +use super::test_support::{ + MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body, + wire_request_with_document, +}; +use crate::ocr::route::LocalOcrHost; + +#[derive(Clone, Copy, Debug)] +enum Route { + Mistral, + AzureAi, + VertexMistral, + AzureCohereParse, + Cohere, +} + +impl Route { + fn model(self) -> &'static str { + match self { + Self::Mistral => "mistral/model", + Self::AzureAi => "azure_ai/model", + Self::VertexMistral => "vertex_ai/mistral-ocr-maas", + Self::AzureCohereParse => "azure_ai/cohere-parse", + Self::Cohere => "cohere/model", + } + } + + fn document_type(self) -> &'static str { + match self { + Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url", + Self::AzureCohereParse | Self::Cohere => "image_url", + } + } + + fn options(self) -> Value { + match self { + Self::Mistral | Self::AzureAi => json!({"pages": [0]}), + Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}), + Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}), + } + } +} + +/// What the host does to the wire request in `before_send`. +#[derive(Clone, Copy, Debug)] +enum Host { + Detached, + ReplacesDocument, +} + +const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ="; + +impl Host { + fn before_send(self, wire: WireRequest) -> WireRequest { + let Value::Object(fields) = wire.body else { + return wire; + }; + let body = fields + .into_iter() + .map(|(name, value)| match self { + Self::Detached => (name, value), + Self::ReplacesDocument if name == "document" => { + let document_type = value["type"].clone(); + let key = document_type.as_str().unwrap_or_default().to_string(); + (name, json!({"type": document_type, key: REPLACED_DOCUMENT})) + } + Self::ReplacesDocument => (name, value), + }) + .collect(); + WireRequest { + body: Value::Object(body), + ..wire + } + } +} + +struct Sent { + result: Result<(), Error>, + provider_body: Option, +} + +async fn send(route: Route, host: Host, document_base: &str) -> Sent { + let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await; + let document_type = route.document_type(); + let document = + json!({"type": document_type, document_type: format!("{document_base}/scan.png")}); + let request = wire_request_with_document(route.model(), &base, document, route.options()); + let local = + LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire))); + let result = perform_ocr_with(local).await.map(|_| ()); + match result { + Ok(()) => provider.await.unwrap(), + Err(_) => provider.abort(), + } + let provider_body = seen + .lock() + .unwrap() + .first() + .map(|request| request_body(request)); + Sent { + result, + provider_body, + } +} + +fn served_document_uri() -> String { + use base64::Engine; + format!( + "data:image/png;base64,{}", + base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT) + ) +} + +#[rstest] +#[case::azure_ai(Route::AzureAi)] +#[case::vertex_mistral(Route::VertexMistral)] +#[case::azure_cohere_parse(Route::AzureCohereParse)] +#[tokio::test] +async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) { + let (document_base, _documents) = document_server().await; + let sent = send(route, Host::Detached, &document_base).await; + sent.result.unwrap(); + assert_eq!( + sent.provider_body.unwrap()["document"][route.document_type()], + json!(served_document_uri()) + ); +} + +#[rstest] +#[tokio::test] +async fn document_replaced_by_the_host_reaches_the_provider( + #[values( + Route::Mistral, + Route::AzureAi, + Route::VertexMistral, + Route::AzureCohereParse, + Route::Cohere + )] + route: Route, +) { + let (document_base, _documents) = document_server().await; + let sent = send(route, Host::ReplacesDocument, &document_base).await; + sent.result.unwrap(); + assert_eq!( + sent.provider_body.unwrap()["document"][route.document_type()], + json!(REPLACED_DOCUMENT) + ); +} diff --git a/litellm-rust/crates/core/tests/ocr/passthrough.rs b/litellm-rust/crates/core/tests/ocr/passthrough.rs deleted file mode 100644 index 0273b48664d..00000000000 --- a/litellm-rust/crates/core/tests/ocr/passthrough.rs +++ /dev/null @@ -1,282 +0,0 @@ -use std::{ - collections::BTreeSet, - sync::{Arc, Mutex}, -}; - -use litellm_callbacks::event::{RequestContext, WireRequest}; -use litellm_llms::base_llm::ocr::error::Error; -use rstest::rstest; -use rstest_reuse::{self, apply, template}; -use serde_json::{Map, Value, json}; - -use super::test_support::{ - MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body, - wire_request_with_document, -}; -use crate::ocr::route::LocalOcrHost; - -#[derive(Clone, Copy, Debug)] -enum Route { - Mistral, - AzureAi, - VertexMistral, - AzureCohereParse, - Cohere, -} - -impl Route { - fn model(self) -> &'static str { - match self { - Self::Mistral => "mistral/model", - Self::AzureAi => "azure_ai/model", - Self::VertexMistral => "vertex_ai/mistral-ocr-maas", - Self::AzureCohereParse => "azure_ai/cohere-parse", - Self::Cohere => "cohere/model", - } - } - - fn document_type(self) -> &'static str { - match self { - Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url", - Self::AzureCohereParse | Self::Cohere => "image_url", - } - } - - fn options(self) -> Value { - match self { - Self::Mistral | Self::AzureAi => json!({"pages": [0]}), - Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}), - Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}), - } - } -} - -#[derive(Clone, Copy, Debug)] -enum Source { - Inline, - Remote, - RemoteWithExtraField, -} - -/// What the host does to the wire request in `before_send`. -#[derive(Clone, Copy, Debug)] -enum Host { - Detached, - /// What `litellm-callbacks-legacy` does before `pre_call`: every passthrough body key - /// is replaced by the caller's own value. - Realiasing, - ReplacesDocument, -} - -const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ="; - -impl Host { - fn before_send( - self, - caller: &Map, - wire: WireRequest, - context: &RequestContext, - ) -> WireRequest { - let Value::Object(fields) = wire.body else { - return wire; - }; - let body = fields - .into_iter() - .map(|(name, value)| match self { - Self::Detached => (name, value), - Self::Realiasing => { - let aliased = context - .passthrough_fields - .contains(&name) - .then(|| caller.get(&name).cloned()) - .flatten() - .unwrap_or(value); - (name, aliased) - } - Self::ReplacesDocument if name == "document" => { - let document_type = value["type"].clone(); - let key = document_type.as_str().unwrap_or_default().to_string(); - (name, json!({"type": document_type, key: REPLACED_DOCUMENT})) - } - Self::ReplacesDocument => (name, value), - }) - .collect(); - WireRequest { - body: Value::Object(body), - ..wire - } - } -} - -struct Sent { - caller: Map, - result: Result<(), Error>, - before_send: Option<(WireRequest, RequestContext)>, - provider_body: Option, -} - -fn caller_document(route: Route, source: Source, document_base: &str) -> Value { - let document_type = route.document_type(); - let remote = format!("{document_base}/scan.png"); - match source { - Source::Inline => { - json!({"type": document_type, document_type: "data:image/png;base64,YWJj"}) - } - Source::Remote => json!({"type": document_type, document_type: remote}), - Source::RemoteWithExtraField => { - json!({"type": document_type, document_type: remote, "document_name": "scan.png"}) - } - } -} - -async fn send(route: Route, source: Source, host: Host, document_base: &str) -> Sent { - let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await; - let document = caller_document(route, source, document_base); - let caller: Map = route - .options() - .as_object() - .unwrap() - .clone() - .into_iter() - .chain([("document".to_string(), document.clone())]) - .collect(); - let observed = Arc::new(Mutex::new(None)); - let captured = observed.clone(); - let host_caller = caller.clone(); - let request = wire_request_with_document(route.model(), &base, document, route.options()); - let local = LocalOcrHost::new(request).with_before_send(move |wire, context| { - *captured.lock().unwrap() = Some((wire.clone(), context.clone())); - Ok(host.before_send(&host_caller, wire, context)) - }); - let result = perform_ocr_with(local).await.map(|_| ()); - match result { - Ok(()) => provider.await.unwrap(), - Err(_) => provider.abort(), - } - let provider_body = seen - .lock() - .unwrap() - .first() - .map(|request| request_body(request)); - let before_send = observed.lock().unwrap().take(); - Sent { - caller, - result, - before_send, - provider_body, - } -} - -fn served_document_uri() -> String { - use base64::Engine; - format!( - "data:image/png;base64,{}", - base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT) - ) -} - -#[template] -#[rstest] -fn every_route_and_source( - #[values( - Route::Mistral, - Route::AzureAi, - Route::VertexMistral, - Route::AzureCohereParse, - Route::Cohere - )] - route: Route, - #[values(Source::Inline, Source::Remote, Source::RemoteWithExtraField)] source: Source, -) { -} - -#[template] -#[rstest] -fn every_route( - #[values( - Route::Mistral, - Route::AzureAi, - Route::VertexMistral, - Route::AzureCohereParse, - Route::Cohere - )] - route: Route, -) { -} - -#[template] -#[rstest] -#[case::azure_ai(Route::AzureAi)] -#[case::vertex_mistral(Route::VertexMistral)] -#[case::azure_cohere_parse(Route::AzureCohereParse)] -fn inlining_routes(#[case] route: Route) {} - -#[apply(every_route_and_source)] -#[tokio::test] -async fn passthrough_fields_are_exactly_the_caller_values_sent_unchanged( - route: Route, - source: Source, -) { - let (document_base, _documents) = document_server().await; - let sent = send(route, source, Host::Detached, &document_base).await; - sent.result.unwrap(); - let (wire, context) = sent.before_send.unwrap(); - let passthrough: BTreeSet<&str> = context.passthrough_fields.iter().collect(); - let unchanged: BTreeSet<&str> = sent - .caller - .iter() - .filter(|(name, value)| wire.body.get(name.as_str()) == Some(*value)) - .map(|(name, _)| name.as_str()) - .collect(); - assert_eq!( - passthrough, - unchanged, - "body: {:#}\ncaller: {:#}", - wire.body, - Value::Object(sent.caller.clone()) - ); -} - -#[apply(every_route_and_source)] -#[tokio::test] -async fn realiasing_leaves_the_provider_request_unchanged(route: Route, source: Source) { - let (document_base, _documents) = document_server().await; - let detached = send(route, source, Host::Detached, &document_base).await; - let realiased = send(route, source, Host::Realiasing, &document_base).await; - detached.result.unwrap(); - realiased.result.unwrap(); - assert_eq!(realiased.provider_body, detached.provider_body); -} - -#[apply(inlining_routes)] -#[tokio::test] -async fn inlining_routes_send_the_downloaded_document( - route: Route, - #[values(Host::Detached, Host::Realiasing)] host: Host, -) { - let (document_base, _documents) = document_server().await; - let sent = send(route, Source::Remote, host, &document_base).await; - sent.result.unwrap(); - assert_eq!( - sent.provider_body.unwrap()["document"][route.document_type()], - json!(served_document_uri()) - ); -} - -#[apply(every_route)] -#[tokio::test] -async fn document_replaced_by_the_host_reaches_the_provider(route: Route) { - let (document_base, _documents) = document_server().await; - let sent = send( - route, - Source::Remote, - Host::ReplacesDocument, - &document_base, - ) - .await; - sent.result.unwrap(); - assert_eq!( - sent.provider_body.unwrap()["document"][route.document_type()], - json!(REPLACED_DOCUMENT) - ); -} diff --git a/litellm-rust/crates/core/tests/ocr/support.rs b/litellm-rust/crates/core/tests/ocr/support.rs index f3adf27cfa6..44313d5f552 100644 --- a/litellm-rust/crates/core/tests/ocr/support.rs +++ b/litellm-rust/crates/core/tests/ocr/support.rs @@ -1,7 +1,7 @@ use std::sync::{Arc, Mutex}; use futures_util::future::BoxFuture; -use litellm_callbacks::event::{Passthrough, WireRequest}; +use litellm_callbacks::event::WireRequest; use litellm_llms::{ base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}, custom_httpx::llm_http_handler::{CallHooks, OcrClient}, @@ -23,11 +23,7 @@ use crate::ocr::{ pub(crate) struct NoHooks; impl CallHooks for NoHooks { - fn before_send( - &self, - wire: WireRequest, - _passthrough_fields: Passthrough, - ) -> BoxFuture<'_, Result> { + fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result> { Box::pin(async move { Ok(wire) }) } @@ -70,7 +66,7 @@ pub(crate) fn wire_request_with_document( decode_request(OcrWireRequest { model: model.into(), document, - api_key: Some("test-key".into()), + api_key: Some(litellm_auth::SecretValue::new("test-key")), api_base: Some(base.into()), custom_llm_provider: None, extra_headers: None, diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index a3fdd2340b3..903ccb06c77 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,9 +1,10 @@ - Target invariants; implementation and runtime validation may lag these rules -- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `CallbackAdapter`/`RouteHost` traits +- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits - No LiteLLM domain dependencies beyond `litellm-callbacks`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) - - A failure that surfaces inside the call, including a host op the call asked for, is mapped through the route's `map_failure`; a failure in `begin` or `after_success` is raised as is + - A native failure, including one a host op returns as `HostOpError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is + - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs - Prefer `Bound<'py, T>` for attached operations/results, `Py` for retention; binding/unbinding does not copy payloads - Use `pythonize` for selected Serde data, never a JSON-text round trip; share conversion with `Pythonized` diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index f1bc3142a25..4aa7a2163ca 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -11,7 +11,7 @@ pub fn missing_state() -> PyErr { /// What an adapter step produced: either the value the driver asked for, or a Python /// awaitable the driver hands back to the caller's task before asking again. -pub enum AdapterStep { +pub enum LifecycleStep { Await(Py), Arguments(Py), Wire(Box), @@ -28,55 +28,74 @@ pub enum PublicValue<'a> { /// One consumer of a call's lifecycle on the Python side. The driver calls the steps in /// order: `begin` before the machine starts, `before_send` and `emit` while it runs, /// `after_success` and one terminal `emit` after it completes. Whenever a step returns -/// [`AdapterStep::Await`], the driver awaits it in the caller's task and continues the +/// [`LifecycleStep::Await`], the driver awaits it in the caller's task and continues the /// same step through `resume`. /// /// A step that fails with an ordinary exception fails the call with that exception, /// except on a terminal event, where the adapter is expected to report and swallow its /// own errors. An exception that is not a `PyException`, such as a cancellation, ends /// the call without further dispatch. -pub trait CallbackAdapter: Send + Sync { +pub trait PythonLifecycle: Send + Sync { fn begin( &mut self, py: Python<'_>, arguments: Py, started_at: f64, - ) -> PyResult; + ) -> PyResult; fn before_send( &mut self, py: Python<'_>, wire: Box, context: &RequestContext, - ) -> PyResult; + ) -> PyResult; fn after_success( &mut self, py: Python<'_>, response: Py, timing: Timing, - ) -> PyResult; + ) -> PyResult; fn emit( &mut self, py: Python<'_>, event: &CallEvent, public: Option>, - ) -> PyResult; + ) -> PyResult; - fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult; + fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult; fn close(&mut self, py: Python<'_>); fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -/// The Python side of one route: answers the route's own operations, builds the public -/// response and maps failures to public exceptions. -pub trait RouteHost: Send + Sync { - type Route: Route; +/// Why a route operation the host answered did not produce a result: the route's own code +/// rejected it, which the route classifies like any other native failure, or Python code +/// raised, which reaches the caller as it was raised. +#[derive(Debug)] +pub enum HostOpError { + Native(E), + Python(PyErr), +} - /// `arguments` is the keyword view the callback adapter's `begin` produced, not the +impl From for HostOpError { + fn from(error: PyErr) -> Self { + Self::Python(error) + } +} + +/// The Python side of one route: answers the route's own operations, builds the public +/// response and classifies native failures into public exceptions. +pub trait RouteHost: Send + Sync { + type Route: Route; + + /// The public exception a native failure maps to, kept as a value until the driver + /// raises it. + type Failure: Into; + + /// `arguments` is the keyword view the lifecycle's `begin` produced, not the /// caller's own dict. A route host that projects from it inherits whatever that /// adapter rewrote. fn invoke( @@ -84,7 +103,7 @@ pub trait RouteHost: Send + Sync { py: Python<'_>, arguments: &Bound<'_, PyDict>, op: ::Op, - ) -> PyResult<::OpResult>; + ) -> Result<::OpResult, HostOpError<::Error>>; fn complete( &mut self, @@ -92,12 +111,14 @@ pub trait RouteHost: Send + Sync { response: ::Response, ) -> PyResult>; - fn native_error(error: ::Error) -> PyErr; + fn classify( + &self, + py: Python<'_>, + error: ::Error, + ) -> PyResult; fn host_error(error: &PyErr) -> ::Error; - fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult; - fn close(&mut self, py: Python<'_>); fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; diff --git a/litellm-rust/crates/host-python/src/argument.rs b/litellm-rust/crates/host-python/src/argument.rs new file mode 100644 index 00000000000..34e07cdfbd5 --- /dev/null +++ b/litellm-rust/crates/host-python/src/argument.rs @@ -0,0 +1,51 @@ +use pyo3::{prelude::*, types::PyDict}; + +/// The caller's own object for a public argument: the keyword if given, even an explicit +/// `None`, else the bound request's attribute. Every reader of a public Python call uses +/// this rule, so the callbacks and the provider see one object per argument. +pub fn lookup<'py>( + kwargs: &Bound<'py, PyDict>, + request: &Bound<'py, PyAny>, + name: &str, +) -> PyResult>> { + if let Some(value) = kwargs.get_item(name)? { + return Ok(Some(value)); + } + request.getattr_opt(name) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() { + crate::initialize_python(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +key = object() +document = {'type': 'document_url'} +class Request: + api_key = 'from-request' + api_base = 'from-request' + document = document +request = Request() +kwargs = {'api_key': key, 'api_base': None} +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let item = |name: &str| locals.get_item(name).unwrap().unwrap(); + let kwargs = item("kwargs").cast_into::().unwrap(); + let request = item("request"); + let find = |name: &str| lookup(&kwargs, &request, name).unwrap(); + assert!(find("api_key").unwrap().is(item("key"))); + assert!(find("api_base").unwrap().is_none()); + assert!(find("document").unwrap().is(item("document"))); + assert!(find("model").is_none()); + }); + } +} diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 8bda13b44d0..7895d803cb5 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -12,7 +12,9 @@ use pyo3::prelude::*; use pyo3::types::PyDict; use tokio::sync::Mutex; -use crate::adapter::{AdapterStep, CallbackAdapter, PublicValue, RouteHost, missing_state}; +use crate::adapter::{ + HostOpError, LifecycleStep, PublicValue, PythonLifecycle, RouteHost, missing_state, +}; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; @@ -43,6 +45,7 @@ enum Stage { #[derive(Clone, Copy)] enum Expect { + Started, Arguments, Wire, Emitted, @@ -66,7 +69,7 @@ where M: Machine> + 'static, { route: H, - adapter: Box, + adapter: Box, machine: Option>>>, arguments: Option>, started_at: f64, @@ -84,7 +87,7 @@ pub fn run_call( py: Python<'_>, machine: M, route: H, - adapter: Box, + adapter: Box, arguments: Py, asynchronous: bool, ) -> PyResult> @@ -146,9 +149,11 @@ where match (self.pending.take(), result) { (None, None) => { self.started_at = epoch_seconds(); - let arguments = self.arguments.take().ok_or_else(missing_state)?; - match self.adapter.begin(py, arguments, self.started_at) { - Ok(step) => self.on_adapter(py, step, Expect::Arguments), + let started = CallEvent::Started { + start_time: self.started_at, + }; + match self.adapter.emit(py, &started, None) { + Ok(step) => self.on_adapter(py, step, Expect::Started), Err(error) => self.adapter_failed(py, error), } } @@ -170,27 +175,28 @@ where fn on_adapter( &mut self, py: Python<'_>, - step: AdapterStep, + step: LifecycleStep, expect: Expect, ) -> PyResult { match (expect, step) { - (_, AdapterStep::Await(awaitable)) => { + (_, LifecycleStep::Await(awaitable)) => { self.pending = Some(Pending::Adapter(expect)); Ok(ExecutionStep::Await(awaitable)) } - (Expect::Arguments, AdapterStep::Arguments(arguments)) => { + (Expect::Started, LifecycleStep::Done) => self.begin(py), + (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) } - (Expect::Wire, AdapterStep::Wire(wire)) => { + (Expect::Wire, LifecycleStep::Wire(wire)) => { self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire)))) } - (Expect::Emitted, AdapterStep::Done) => { + (Expect::Emitted, LifecycleStep::Done) => { self.resume_machine(py, Some(Ok(HostResult::Emitted))) } - (Expect::Response, AdapterStep::Response(response)) => self.succeeded(py, response), - (Expect::Terminal, AdapterStep::Done) => match &self.stage { + (Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response), + (Expect::Terminal, LifecycleStep::Done) => match &self.stage { Stage::Succeeded(response) => Ok(ExecutionStep::Return(response.clone_ref(py))), Stage::Failed(error) => Err(PyErr::from_value(error.bind(py).clone().into_any())), _ => Err(missing_state()), @@ -199,6 +205,14 @@ where } } + fn begin(&mut self, py: Python<'_>) -> PyResult { + let arguments = self.arguments.take().ok_or_else(missing_state)?; + match self.adapter.begin(py, arguments, self.started_at) { + Ok(step) => self.on_adapter(py, step, Expect::Arguments), + Err(error) => self.adapter_failed(py, error), + } + } + fn adapter_failed(&mut self, py: Python<'_>, error: PyErr) -> PyResult { match self.stage { Stage::Begin | Stage::AfterSuccess => self.failure(py, error, FailureOrigin::Host), @@ -248,14 +262,20 @@ where let answer = match op { HostOp::Route(op) => { let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; - self.route - .invoke(py, arguments.bind(py), op) - .map(HostResult::Route) + match self.route.invoke(py, arguments.bind(py), op) { + Ok(result) => Ok(HostResult::Route(result)), + Err(HostOpError::Native(error)) => { + return self + .resume_core(py, Some(Err(HostFailure::Error(error)))) + .map(Next::Continue); + } + Err(HostOpError::Python(error)) => Err(error), + } } HostOp::BeforeSend { wire, context } => { match self.adapter.before_send(py, wire, &context) { - Ok(AdapterStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)), - Ok(AdapterStep::Await(awaitable)) => { + Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)), + Ok(LifecycleStep::Await(awaitable)) => { self.pending = Some(Pending::Adapter(Expect::Wire)); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } @@ -264,8 +284,8 @@ where } } HostOp::Emit(event) => match self.adapter.emit(py, &event, None) { - Ok(AdapterStep::Done) => Ok(HostResult::Emitted), - Ok(AdapterStep::Await(awaitable)) => { + Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), + Ok(LifecycleStep::Await(awaitable)) => { self.pending = Some(Pending::Adapter(Expect::Emitted)); return Ok(Next::Return(ExecutionStep::Await(awaitable))); } @@ -360,11 +380,29 @@ where self.ended_at.get_or_insert_with(epoch_seconds); let error = match self.interrupted.take() { Some(retained) => PyErr::from_value(retained.into_bound(py).into_any()), - None => H::native_error(error), + None => self.classified(py, error), }; self.failure(py, error, FailureOrigin::Call) } + /// The route's public exception for a native failure. When classification itself + /// fails, that failure is raised with the native error's text as its `__context__`. + fn classified(&self, py: Python<'_>, error: ErrorOf) -> PyErr { + let native = error.to_string(); + let classifier_error = match self.route.classify(py, error) { + Ok(failure) => return failure.into(), + Err(classifier_error) => classifier_error, + }; + let attached = classifier_error.value(py).setattr( + "__context__", + PyRuntimeError::new_err(native).into_value(py), + ); + match attached { + Ok(()) => classifier_error, + Err(error) => error, + } + } + fn succeeded(&mut self, py: Python<'_>, response: Py) -> PyResult { let event = CallEvent::Succeeded { timing: self.timing(), @@ -386,18 +424,14 @@ where if is_cancellation(py, &error) { return Err(error); } - let public = match origin { - FailureOrigin::Call => self.route.map_failure(py, &error).unwrap_or(error), - FailureOrigin::Host => error, - }; let event = CallEvent::Failed { timing: self.timing(), origin, }; let step = self .adapter - .emit(py, &event, Some(PublicValue::Error(&public)))?; - self.stage = Stage::Failed(public.into_value(py)); + .emit(py, &event, Some(PublicValue::Error(&error)))?; + self.stage = Stage::Failed(error.into_value(py)); self.on_adapter(py, step, Expect::Terminal) } @@ -489,6 +523,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri #[derive(Clone, Debug, PartialEq, Eq)] struct Error(String); + impl std::fmt::Display for Error { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(&self.0) + } + } + struct Synthetic; impl Route for Synthetic { @@ -518,8 +558,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri model: "model".into(), custom_llm_provider: "provider".into(), optional_params: serde_json::json!({}), - passthrough_fields: Default::default(), secret_fields: Vec::new(), + api_key: None, } } @@ -566,25 +606,46 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } + #[derive(Clone, Copy)] + enum OpScript { + Answer, + RaisePython, + RejectNatively, + } + struct SyntheticHost { log: Log, - fail_op: bool, + op: OpScript, + classifier_fails: bool, + } + + /// The fake route's public exception, kept as a value so a test sees what `classify` + /// produced before the driver raises it. + #[derive(Debug, PartialEq, Eq)] + struct Classified(String); + + impl From for PyErr { + fn from(classified: Classified) -> Self { + PyValueError::new_err(format!("classified: {}", classified.0)) + } } impl RouteHost for SyntheticHost { type Route = Synthetic; + type Failure = Classified; fn invoke( &mut self, _: Python<'_>, arguments: &Bound<'_, PyDict>, op: &'static str, - ) -> PyResult { + ) -> Result> { self.log.push(format!("route:{op}")); - if self.fail_op { - return Err(PyValueError::new_err("op failed")); + match self.op { + OpScript::Answer => Ok(format!("{op}:{}", arguments.len())), + OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), + OpScript::RejectNatively => Err(HostOpError::Native(Error("op rejected".into()))), } - Ok(format!("{op}:{}", arguments.len())) } fn complete(&mut self, py: Python<'_>, response: String) -> PyResult> { @@ -594,22 +655,18 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri .unbind()) } - fn native_error(error: Error) -> PyErr { - PyValueError::new_err(error.0) + fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + self.log.push(format!("classify:{error}")); + if self.classifier_fails { + return Err(pyo3::exceptions::PyTypeError::new_err("classifier failed")); + } + Ok(Classified(error.0)) } fn host_error(error: &PyErr) -> Error { Error(error.to_string()) } - fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult { - self.log.push("map_failure"); - Ok(PyValueError::new_err(format!( - "mapped: {}", - error.value(py) - ))) - } - fn close(&mut self, _: Python<'_>) { self.log.push("route.close"); } @@ -632,13 +689,18 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri script: AdapterScript, } - impl CallbackAdapter for SyntheticAdapter { - fn begin(&mut self, _: Python<'_>, arguments: Py, _: f64) -> PyResult { + impl PythonLifecycle for SyntheticAdapter { + fn begin( + &mut self, + _: Python<'_>, + arguments: Py, + _: f64, + ) -> PyResult { self.log.push("begin"); if matches!(self.script, AdapterScript::FailBegin) { return Err(PyValueError::new_err("begin failed")); } - Ok(AdapterStep::Arguments(arguments)) + Ok(LifecycleStep::Arguments(arguments)) } fn before_send( @@ -646,9 +708,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri _: Python<'_>, wire: Box, _: &RequestContext, - ) -> PyResult { + ) -> PyResult { self.log.push("before_send"); - Ok(AdapterStep::Wire(Box::new(WireRequest { + Ok(LifecycleStep::Wire(Box::new(WireRequest { url: "rewritten".into(), ..*wire }))) @@ -659,17 +721,17 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py: Python<'_>, response: Py, _: Timing, - ) -> PyResult { + ) -> PyResult { self.log.push("after_success"); match self.script { - AdapterScript::ReplaceResponse => Ok(AdapterStep::Response( + AdapterScript::ReplaceResponse => Ok(LifecycleStep::Response( "replaced".into_pyobject(py)?.into_any().unbind(), )), AdapterScript::FailAfterSuccess => { Err(PyValueError::new_err("after_success failed")) } AdapterScript::Plain | AdapterScript::FailBegin => { - Ok(AdapterStep::Response(response)) + Ok(LifecycleStep::Response(response)) } } } @@ -679,8 +741,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py: Python<'_>, event: &CallEvent, public: Option>, - ) -> PyResult { + ) -> PyResult { self.log.push(match (event, public) { + (CallEvent::Started { .. }, None) => "started".into(), (CallEvent::ResponseReceived { raw }, None) => format!("response:{}", raw.body), (CallEvent::Succeeded { .. }, Some(PublicValue::Response(value))) => { format!("succeeded:{}", value.bind(py)) @@ -690,10 +753,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } _ => "unexpected".into(), }); - Ok(AdapterStep::Done) + Ok(LifecycleStep::Done) } - fn resume(&mut self, _: Python<'_>, _: PyResult>) -> PyResult { + fn resume(&mut self, _: Python<'_>, _: PyResult>) -> PyResult { Err(missing_state()) } @@ -709,15 +772,31 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn run_scripted( py: Python<'_>, machine: ScriptedMachine, - fail_op: bool, + op: OpScript, script: AdapterScript, asynchronous: bool, ) -> (PyResult>, Vec) { - let log = Log::default(); - let route = SyntheticHost { - log: Log(log.0.clone()), - fail_op, - }; + run_hosted( + py, + machine, + SyntheticHost { + log: Log::default(), + op, + classifier_fails: false, + }, + script, + asynchronous, + ) + } + + fn run_hosted( + py: Python<'_>, + machine: ScriptedMachine, + route: SyntheticHost, + script: AdapterScript, + asynchronous: bool, + ) -> (PyResult>, Vec) { + let log = Log(route.log.0.clone()); let adapter = SyntheticAdapter { log: Log(log.0.clone()), script, @@ -777,7 +856,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_scripted( py, success_machine(), - false, + OpScript::Answer, AdapterScript::Plain, asynchronous, ); @@ -785,6 +864,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri assert_eq!( log, [ + "started", "begin", "route:project", "before_send", @@ -800,28 +880,75 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } + fn failing_machine() -> ScriptedMachine { + ScriptedMachine { + ops: vec![HostOp::Route("project")], + outcome: Some(Err(Error("provider exploded".into()))), + answers: Vec::new(), + } + } + #[test] - fn machine_failures_are_mapped_and_dispatched_once_as_call_failures() { + fn a_native_failure_is_classified_once_and_reported_classified() { let _guard = PYTHON_GLOBALS .lock() .unwrap_or_else(|error| error.into_inner()); crate::initialize_python(); Python::attach(|py| { - let machine = ScriptedMachine { - ops: vec![HostOp::Route("project")], - outcome: Some(Err(Error("provider exploded".into()))), - answers: Vec::new(), - }; - let (result, log) = run_scripted(py, machine, false, AdapterScript::Plain, false); - let error = result.unwrap_err(); - assert_eq!(error.value(py).to_string(), "mapped: provider exploded"); + install_lifecycle_module(py); + for asynchronous in [false, true] { + let (result, log) = run_scripted( + py, + failing_machine(), + OpScript::Answer, + AdapterScript::Plain, + asynchronous, + ); + let error = result.unwrap_err(); + assert!(error.is_instance_of::(py)); + assert_eq!(error.value(py).to_string(), "classified: provider exploded"); + assert_eq!( + log, + [ + "started", + "begin", + "route:project", + "classify:provider exploded", + "failed:Call:classified: provider exploded", + "adapter.close", + "route.close", + ] + ); + } + }); + } + + #[test] + fn a_native_rejection_from_a_host_operation_is_classified_once() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + let (result, log) = run_scripted( + py, + success_machine(), + OpScript::RejectNatively, + AdapterScript::Plain, + false, + ); + assert_eq!( + result.unwrap_err().value(py).to_string(), + "classified: op rejected" + ); assert_eq!( log, [ + "started", "begin", "route:project", - "map_failure", - "failed:Call:mapped: provider exploded", + "classify:op rejected", + "failed:Call:classified: op rejected", "adapter.close", "route.close", ] @@ -830,18 +957,72 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } #[test] - fn host_operation_failures_interrupt_the_call_and_keep_the_python_exception() { + fn a_python_exception_from_a_host_operation_is_reported_as_raised() { let _guard = PYTHON_GLOBALS .lock() .unwrap_or_else(|error| error.into_inner()); crate::initialize_python(); Python::attach(|py| { - let (result, log) = - run_scripted(py, success_machine(), true, AdapterScript::Plain, false); + let (result, log) = run_scripted( + py, + success_machine(), + OpScript::RaisePython, + AdapterScript::Plain, + false, + ); let error = result.unwrap_err(); - assert_eq!(error.value(py).to_string(), "mapped: op failed"); - assert!(!log.contains(&"before_send".to_string())); - assert!(log.contains(&"failed:Call:mapped: op failed".to_string())); + assert!(error.is_instance_of::(py)); + assert_eq!(error.value(py).to_string(), "op failed"); + assert_eq!( + log, + [ + "started", + "begin", + "route:project", + "failed:Call:op failed", + "adapter.close", + "route.close", + ] + ); + }); + } + + #[test] + fn a_failing_classifier_surfaces_with_the_native_error_as_context() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + let (result, log) = run_hosted( + py, + failing_machine(), + SyntheticHost { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: true, + }, + AdapterScript::Plain, + false, + ); + let error = result.unwrap_err(); + assert!(error.is_instance_of::(py)); + assert_eq!(error.value(py).to_string(), "classifier failed"); + let context = error.value(py).getattr("__context__").unwrap(); + assert!(context.is_instance_of::()); + assert_eq!(context.str().unwrap().to_string(), "provider exploded"); + assert_eq!( + log, + [ + "started", + "begin", + "route:project", + "classify:provider exploded", + "failed:Call:classifier failed", + "adapter.close", + "route.close", + ] + ); }); } @@ -855,7 +1036,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_scripted( py, success_machine(), - false, + OpScript::Answer, AdapterScript::FailBegin, false, ); @@ -864,6 +1045,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri assert_eq!( log, [ + "started", "begin", "failed:Host:begin failed", "adapter.close", @@ -885,7 +1067,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_scripted( py, success_machine(), - false, + OpScript::Answer, AdapterScript::ReplaceResponse, asynchronous, ); @@ -908,7 +1090,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri let (result, log) = run_scripted( py, success_machine(), - false, + OpScript::Answer, AdapterScript::FailAfterSuccess, asynchronous, ); @@ -938,12 +1120,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri struct Cancelling(Log); impl RouteHost for Cancelling { type Route = Synthetic; + type Failure = Classified; fn invoke( &mut self, py: Python<'_>, _: &Bound<'_, PyDict>, _: &'static str, - ) -> PyResult { + ) -> Result> { self.0.push("route"); Err(PyErr::from_value( py.import("asyncio") @@ -952,21 +1135,19 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri .unwrap() .call0() .unwrap(), - )) + ) + .into()) } fn complete(&mut self, _: Python<'_>, _: String) -> PyResult> { Err(missing_state()) } - fn native_error(error: Error) -> PyErr { - PyValueError::new_err(error.0) + fn classify(&self, _: Python<'_>, error: Error) -> PyResult { + self.0.push("classify"); + Ok(Classified(error.0)) } fn host_error(error: &PyErr) -> Error { Error(error.to_string()) } - fn map_failure(&self, _: Python<'_>, _: &PyErr) -> PyResult { - self.0.push("map_failure"); - Err(missing_state()) - } fn close(&mut self, _: Python<'_>) {} fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { Ok(()) @@ -988,7 +1169,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ) .unwrap_err(); assert!(!error.is_instance_of::(py)); - assert_eq!(log.entries(), ["begin", "route", "adapter.close"]); + assert_eq!( + log.entries(), + ["started", "begin", "route", "adapter.close"] + ); }); } diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index bb0b5b1c3b1..8889b1513db 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,9 +1,10 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and //! asyncio glue, and the driver that runs a native [`Machine`](litellm_callbacks::machine::Machine) -//! against a Python route host and a callback adapter. Everything here is Python-specific by +//! against a Python route host and a Python lifecycle. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. mod adapter; +mod argument; mod callable; mod driver; mod execution; @@ -11,7 +12,10 @@ mod gil; mod handle; mod marshal; -pub use adapter::{AdapterStep, CallbackAdapter, PublicValue, RouteHost, missing_state}; +pub use adapter::{ + HostOpError, LifecycleStep, PublicValue, PythonLifecycle, RouteHost, missing_state, +}; +pub use argument::lookup; pub use callable::wrap_failure; pub use driver::run_call; pub use execution::{poll_async_value, run_async, run_async_value, run_sync, run_sync_value}; diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs index 87945bf8785..2e398d0287e 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs @@ -150,7 +150,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig { api_key: inputs.api_key.and_then(|key| { inputs .dynamic_api_key - .filter(|value| !value.value().is_empty()) + .filter(|value| !value.value().expose().is_empty()) .or(Some(key)) }), api_base: inputs.api_base.and_then(|base| { @@ -592,12 +592,17 @@ impl AzureDocumentIntelligenceOcrConfig { )?; return Ok(connection.extra_headers.clone()); } - let key = nonblank(connection.api_key.clone()) - .map(|value| Sourced::new(value, connection.api_key_source)) - .or_else(|| { - nonblank(self.get_api_key_env_var().and_then(env_lookup)) - .map(|value| Sourced::new(value, InputSource::Environment)) - }); + let key = nonblank( + connection + .api_key + .as_ref() + .map(|key| key.expose().to_string()), + ) + .map(|value| Sourced::new(value, connection.api_key_source)) + .or_else(|| { + nonblank(self.get_api_key_env_var().and_then(env_lookup)) + .map(|value| Sourced::new(value, InputSource::Environment)) + }); if let Some(key) = key { super::super::common_utils::validate_destination(connection, key.source())?; return Ok( @@ -796,7 +801,7 @@ mod tests { #[tokio::test] async fn request_endpoint_accepts_request_owned_key() { let connection = OcrConnection { - api_key: Some("request-key".into()), + api_key: Some(litellm_auth::SecretValue::new("request-key")), api_key_source: InputSource::Request, api_base: Some("https://request.example".into()), api_base_source: InputSource::Request, diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs index 2012f740173..7ef051e8986 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/transformation.rs @@ -142,12 +142,17 @@ impl AzureAiOcrConfig { super::common_utils::validate_destination(connection, connection.extra_headers_source)?; return Ok(connection.extra_headers.clone()); } - let key = nonblank(connection.api_key.clone()) - .map(|value| Sourced::new(value, connection.api_key_source)) - .or_else(|| { - nonblank(self.get_api_key_env_var().and_then(env_lookup)) - .map(|value| Sourced::new(value, InputSource::Environment)) - }); + let key = nonblank( + connection + .api_key + .as_ref() + .map(|key| key.expose().to_string()), + ) + .map(|value| Sourced::new(value, connection.api_key_source)) + .or_else(|| { + nonblank(self.get_api_key_env_var().and_then(env_lookup)) + .map(|value| Sourced::new(value, InputSource::Environment)) + }); if let Some(key) = key { super::common_utils::validate_destination(connection, key.source())?; return Ok(bearer_headers(connection, key.value())); @@ -196,7 +201,7 @@ mod tests { #[fixture] fn connection() -> OcrConnection { OcrConnection { - api_key: Some("request-key".into()), + api_key: Some(litellm_auth::SecretValue::new("request-key")), api_base: Some("https://example.com".into()), ..Default::default() } @@ -288,7 +293,7 @@ mod tests { #[tokio::test] async fn request_endpoint_accepts_request_owned_key() { let connection = OcrConnection { - api_key: Some("request-key".into()), + api_key: Some(litellm_auth::SecretValue::new("request-key")), api_key_source: InputSource::Request, api_base: Some("https://request.example".into()), api_base_source: InputSource::Request, diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs index e6fe5d9556d..8321dcfb4ce 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs @@ -1,6 +1,6 @@ use std::{collections::BTreeMap, future::Future, time::Duration}; -use litellm_auth::{InputSource, Sourced, TokenProviderHandle}; +use litellm_auth::{InputSource, SecretValue, Sourced, TokenProviderHandle}; use litellm_core_utils::{ call_arguments::CallArguments, serde_compat::{FiniteF64, LaxI64}, @@ -90,21 +90,22 @@ pub enum OcrResponseFormat { #[derive(Clone, Default)] pub struct OcrCredentialInputs { - pub api_key: Option>, - pub dynamic_api_key: Option>, + pub api_key: Option>, + pub dynamic_api_key: Option>, pub api_base: Option>, pub dynamic_api_base: Option>, } impl OcrCredentialInputs { pub fn new( - api_key: Option, + api_key: Option, api_key_source: InputSource, api_base: Option, api_base_source: InputSource, ) -> Self { Self { - api_key: nonblank(api_key).map(|value| Sourced::new(value, api_key_source)), + api_key: nonblank(api_key.as_ref().map(|key| key.expose().to_string())) + .map(|value| Sourced::new(SecretValue::new(value), api_key_source)), dynamic_api_key: None, api_base: nonblank(api_base).map(|value| Sourced::new(value, api_base_source)), dynamic_api_base: None, @@ -159,7 +160,7 @@ fn nonblank(value: Option) -> Option { #[derive(Clone)] pub struct OcrConnection { - pub api_key: Option, + pub api_key: Option, pub api_key_source: InputSource, pub api_base: Option, pub api_base_source: InputSource, @@ -209,7 +210,7 @@ impl Default for OcrConnection { #[derive(Clone, Default)] pub struct ResolvedOcrCredentials { - pub api_key: Option>, + pub api_key: Option>, pub api_base: Option>, } @@ -428,7 +429,7 @@ pub trait BaseOcrConfig: Send + Sync + Sized + 'static { ResolvedOcrCredentials { api_key: inputs .dynamic_api_key - .filter(|value| !value.value().is_empty()) + .filter(|value| !value.value().expose().is_empty()) .or(inputs.api_key), api_base: inputs .dynamic_api_base diff --git a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs index f353c22d8c4..2528c967f41 100644 --- a/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/cohere/ocr/transformation.rs @@ -179,8 +179,8 @@ impl CohereParseConfig { } let key = connection .api_key - .as_deref() - .map(str::trim) + .as_ref() + .map(|key| key.expose().trim()) .filter(|key| !key.is_empty()) .map(str::to_string) .or_else(|| { @@ -718,7 +718,7 @@ mod tests { assert!(matches!( CohereParseConfig.resolve_headers( &OcrConnection { - api_key: Some(" ".into()), + api_key: Some(litellm_auth::SecretValue::new(" ")), ..Default::default() }, &|_| None, diff --git a/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs b/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs index e93ddee3c50..7635bdd3d04 100644 --- a/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs +++ b/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs @@ -3,9 +3,9 @@ use std::{sync::OnceLock, time::Duration}; use bytes::{Bytes, BytesMut}; use futures_util::future::BoxFuture; use litellm_auth_gcp::VertexAuth; -use litellm_callbacks::event::{Passthrough, WireRequest}; +use litellm_callbacks::event::WireRequest; use serde::{Serialize, de::DeserializeOwned}; -use serde_json::{Map, Value}; +use serde_json::Value; use crate::{ base_llm::ocr::{ @@ -26,11 +26,7 @@ use crate::{ /// The route's view of one call, handed to provider code that has to reach the /// caller's hooks mid-flight (guardrails on the outgoing body, raw response events). pub trait CallHooks: Send + Sync { - fn before_send( - &self, - wire: WireRequest, - passthrough_fields: Passthrough, - ) -> BoxFuture<'_, Result>; + fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result>; fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), E>>; } @@ -231,9 +227,8 @@ pub async fn transform_request_body( config.get_supported_ocr_params(&request.model), )?; config.validate_request_body(&composed)?; - let passthrough_fields = Passthrough::unchanged(&caller_inputs(request)?, &composed); let changed = hooks - .before_send(wire_request(url, headers, composed), passthrough_fields) + .before_send(wire_request(url, headers, composed)) .await?; if !changed.body.is_object() { return Err(Error::RequestField { @@ -252,21 +247,6 @@ fn wire_request(url: &str, headers: &[(String, String)], body: Value) -> WireReq } } -fn caller_inputs(request: &PreparedOcrRequest) -> Result, Error> { - let document = request - .caller_document - .then(|| serde_json::to_value(&request.document)) - .transpose() - .map_err(|_| Error::RequestField { - path: "document".into(), - })?; - let params: Map = request.optional_params.clone().into(); - Ok(params - .into_iter() - .chain(document.map(|document| ("document".to_string(), document))) - .collect()) -} - pub fn build_http_request( client: &OcrClient, request: &PreparedOcrRequest, @@ -294,9 +274,7 @@ pub async fn guardrail_document( let body = serde_json::to_value(&request.document).map_err(|_| Error::RequestField { path: "document".into(), })?; - let changed = hooks - .before_send(wire_request(url, headers, body), Passthrough::default()) - .await?; + let changed = hooks.before_send(wire_request(url, headers, body)).await?; let document = decode_request_value(changed.body, "guardrail.document")?; Ok((document, changed.headers)) } diff --git a/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs b/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs index c2038d0552d..9028f09c5ab 100644 --- a/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/mistral/ocr/transformation.rs @@ -135,8 +135,8 @@ impl MistralOcrConfig { } let api_key = connection .api_key - .as_deref() - .map(str::trim) + .as_ref() + .map(|key| key.expose().trim()) .filter(|key| !key.is_empty()) .map(str::to_string) .or_else(|| { @@ -212,7 +212,7 @@ mod tests { #[default(vec![])] extra_headers: Vec<(String, String)>, ) -> OcrConnection { OcrConnection { - api_key: api_key.map(str::to_string), + api_key: api_key.map(litellm_auth::SecretValue::new), extra_headers, ..OcrConnection::default() } diff --git a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs index ca2bae9c3bb..ec876fafb8f 100644 --- a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs @@ -442,8 +442,8 @@ fn resolve_headers( } let api_key = connection .api_key - .as_deref() - .map(str::trim) + .as_ref() + .map(|key| key.expose().trim()) .filter(|key| !key.is_empty()) .map(str::to_string) .or_else(|| { @@ -629,7 +629,7 @@ mod tests { #[test] fn explicit_key_precedes_environment_key() { let connection = OcrConnection { - api_key: Some("passed-key".into()), + api_key: Some(litellm_auth::SecretValue::new("passed-key")), ..Default::default() }; let headers = resolve_headers(&connection, &|_| Some("env-key".into())).unwrap(); @@ -639,7 +639,7 @@ mod tests { #[test] fn blank_explicit_key_uses_environment_key() { let connection = OcrConnection { - api_key: Some(" ".into()), + api_key: Some(litellm_auth::SecretValue::new(" ")), ..Default::default() }; let headers = resolve_headers(&connection, &|_| Some(" env-key ".into())).unwrap(); diff --git a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs index ea0bcf3d08c..c2cb23d0010 100644 --- a/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/vertex_ai/ocr/transformation.rs @@ -134,7 +134,10 @@ impl VertexAiOcrConfig { .vertex_auth() .validate_environment( connection.extra_headers.clone(), - connection.api_key.as_deref(), + connection + .api_key + .as_ref() + .map(litellm_auth::SecretValue::expose), config, &credential_env, ) diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index 9932594e2f5..5dccfb4aca8 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -1,7 +1,7 @@ - Target invariants, not completion claims; these supersede older conflicting bridge guidance - Keep this crate the product-specific PyO3 consumer of `litellm-host-python` - Own registration, input projection, the route host and the caller callables it answers operations with (file readers, token providers), public response/error construction and the per-call composition of machine, route host and callback contract - - Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, `passthrough_fields` re-aliasing) lives in `litellm-callbacks-legacy` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy + - Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, re-aliasing unchanged body keys) lives in `litellm-callbacks-legacy` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy - Value-oriented execution, sync waiting, nested-runtime checks, signal polling and panic containment live in `litellm-host-python`; native async work uses `pyo3-async-runtimes`, Serde output uses `Pythonized` - Core owns typed native state, the route machine, provider preparation/I/O and normalization; the host driver owns terminal events; the legacy adapter in `litellm-callbacks-legacy` owns `Logging` dispatch policy - Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 9dc891a91d6..77212e3d38e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,9 +1,9 @@ use litellm_auth::ResolvedCredential; use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult}; -use litellm_host_python::{RouteHost, missing_state, to_py}; +use litellm_host_python::{HostOpError, RouteHost, missing_state, to_py}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ - exceptions::PyBaseException, + exceptions::{PyBaseException, PyException}, gc::{PyTraverseError, PyVisit}, prelude::*, types::PyDict, @@ -57,12 +57,8 @@ impl OcrRouteHost { .ok_or_else(missing_state)? .acquire(py) } -} -impl RouteHost for OcrRouteHost { - type Route = Ocr; - - fn invoke( + fn answer( &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, @@ -88,6 +84,40 @@ impl RouteHost for OcrRouteHost { } } + fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr { + if !error.is_instance_of::(py) { + return error; + } + let provider = match &self.data { + OcrHostData::Projected(handles) => handles.provider, + _ => "", + }; + let mapped = py + .import("litellm.rust_bridge.ocr.route_host") + .and_then(|module| module.getattr("map_failure")) + .and_then(|map| map.call1((error.value(py), self.request.bind(py), provider))) + .and_then(|mapped| mapped.extract::>().map_err(PyErr::from)); + match mapped { + Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()), + Err(_) => error, + } + } +} + +impl RouteHost for OcrRouteHost { + type Route = Ocr; + type Failure = PyErr; + + fn invoke( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + op: OcrOp, + ) -> Result> { + self.answer(py, arguments, op) + .map_err(|error| HostOpError::Python(self.map_failure(py, error))) + } + fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { py.import("litellm.rust_bridge.ocr.route_host")? .getattr("response")? @@ -95,27 +125,14 @@ impl RouteHost for OcrRouteHost { .map(Bound::unbind) } - fn native_error(error: Error) -> PyErr { - ocr_error_to_pyerr(error) + fn classify(&self, py: Python<'_>, error: Error) -> PyResult { + Ok(self.map_failure(py, ocr_error_to_pyerr(error))) } fn host_error(error: &PyErr) -> Error { Error::InvalidRequest(error.to_string()) } - fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult { - let provider = match &self.data { - OcrHostData::Projected(handles) => handles.provider, - _ => "", - }; - let mapped: Py = py - .import("litellm.rust_bridge.ocr.route_host")? - .getattr("map_failure")? - .call1((error.value(py), self.request.bind(py), provider))? - .extract()?; - Ok(PyErr::from_value(mapped.into_bound(py).into_any())) - } - fn close(&mut self, _: Python<'_>) { self.data = OcrHostData::Released; } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 7ffa129f85c..5dd2aa804b8 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -1,3 +1,4 @@ +use litellm_auth::SecretValue; use litellm_core::ocr::{ types::{LiteLLMOcrRequest, OcrDocumentInput}, wire::{OcrWireRequest, consumed_optional_params, decode_document, decode_request_input}, @@ -31,7 +32,7 @@ struct OcrArguments<'a, 'py> { impl<'py> OcrArguments<'_, 'py> { fn lookup(&self, name: &str) -> PyResult> { - litellm_callbacks_legacy::lookup(self.kwargs, self.request, name)? + litellm_host_python::lookup(self.kwargs, self.request, name)? .ok_or_else(|| PyValueError::new_err(format!("missing argument: {name}"))) } @@ -47,8 +48,11 @@ impl<'py> OcrArguments<'_, 'py> { self.lookup("document") } - fn api_key(&self) -> PyResult> { - self.lookup("api_key")?.extract() + fn api_key(&self) -> PyResult> { + Ok(self + .lookup("api_key")? + .extract::>()? + .map(SecretValue::new)) } fn api_base(&self) -> PyResult> { diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 4a9a65b1485..7248c2f3590 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1265,11 +1265,6 @@ class Logging(LiteLLMLoggingBaseClass): additional_args.get("api_base", "") ) - def record_api_call_start_time(self) -> None: - self.model_call_details["api_call_start_time"] = datetime.datetime.now() - if self.model_call_details.get("first_api_call_start_time") is None: - self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"] - def pre_call(self, input, api_key, model=None, additional_args={}): # Log the exact input to the LLM API try: @@ -1334,7 +1329,15 @@ class Logging(LiteLLMLoggingBaseClass): "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", e ) - self.record_api_call_start_time() + self.model_call_details["api_call_start_time"] = datetime.datetime.now() + # Set-once first provider-handoff instant. api_call_start_time + # is overwritten on every retry, so it can't measure one-time + # preprocessing; pinning the first attempt excludes retry loops + # + backoff. Logging object only — must NOT go into + # litellm_params["metadata"] (caller request metadata, typed + # Dict[str, str], echoed downstream; a datetime breaks it). + if self.model_call_details.get("first_api_call_start_time") is None: + self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"] # Input Integration Logging -> If you want to log the fact that an attempt to call the model was made callbacks: Final = litellm.input_callback + (self.dynamic_input_callbacks or []) for callback in callbacks: @@ -1468,21 +1471,16 @@ class Logging(LiteLLMLoggingBaseClass): """ return _get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers) - def record_post_call( - self, original_response: object, input: object, api_key: object, additional_args: dict[str, object] - ) -> None: - self.model_call_details["input"] = input - self.model_call_details["api_key"] = api_key - self.model_call_details["original_response"] = original_response - self.model_call_details["additional_args"] = additional_args - self.model_call_details["log_event_type"] = "post_api_call" - def post_call(self, original_response, input=None, api_key=None, additional_args={}): # Log the exact result from the LLM API, for streaming - log the type of response received if isinstance(original_response, dict): original_response = json.dumps(original_response, default=str) try: - self.record_post_call(original_response, input, api_key, additional_args) + self.model_call_details["input"] = input + self.model_call_details["api_key"] = api_key + self.model_call_details["original_response"] = original_response + self.model_call_details["additional_args"] = additional_args + self.model_call_details["log_event_type"] = "post_api_call" attr: Literal["warning", "debug"] if self.litellm_request_debug: @@ -2177,7 +2175,6 @@ class Logging(LiteLLMLoggingBaseClass): logging_result, start_time, end_time, - build_logging_payload: bool = True, ): """Resolve hidden params, compute response cost, and emit the standard logging payload.""" hidden_params: Final = getattr(logging_result, "_hidden_params", {}) @@ -2202,9 +2199,6 @@ class Logging(LiteLLMLoggingBaseClass): else: self.model_call_details["response_cost"] = self._response_cost_calculator(result=logging_result) - if not build_logging_payload: - return - self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload( logging_result, start_time, end_time ) @@ -2266,7 +2260,6 @@ class Logging(LiteLLMLoggingBaseClass): end_time=None, cache_hit=None, standard_logging_object: StandardLoggingPayload | None = None, - build_logging_payload: bool = True, ): try: if start_time is None: @@ -2304,7 +2297,6 @@ class Logging(LiteLLMLoggingBaseClass): logging_result=logging_result, start_time=start_time, end_time=end_time, - build_logging_payload=build_logging_payload, ) elif standard_logging_object is not None: self.model_call_details["standard_logging_object"] = standard_logging_object @@ -3328,9 +3320,7 @@ class Logging(LiteLLMLoggingBaseClass): except Exception as e: verbose_logger.debug("Error in _handle_callback_failure: %s", e) - def _failure_handler_helper_fn( - self, exception, traceback_exception, start_time=None, end_time=None, build_logging_payload: bool = True - ): + def _failure_handler_helper_fn(self, exception, traceback_exception, start_time=None, end_time=None): if start_time is None: start_time = self.start_time if end_time is None: @@ -3365,9 +3355,6 @@ class Logging(LiteLLMLoggingBaseClass): metadata: Final = self.model_call_details["litellm_params"].get("metadata", {}) or {} metadata.update(exception.headers) - if not build_logging_payload: - return start_time, end_time - ## STANDARDIZED LOGGING PAYLOAD self.model_call_details["standard_logging_object"] = get_standard_logging_object_payload( diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index e05d9368fa8..406dc55cfea 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -6,24 +6,22 @@ registries it fans out to. It expires with that contract. from __future__ import annotations +import contextvars import datetime -import os +import traceback import uuid -from collections.abc import Mapping +from collections.abc import Awaitable, Coroutine, Mapping from dataclasses import dataclass from typing import ( TYPE_CHECKING, Final, - Literal, Protocol, - TypeAlias, cast, # noqa: TID251 # bounded compatibility calls into legacy Python integrations ) -from typing_extensions import assert_never - if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.types.utils import CredentialItem class MetadataUpdater(Protocol): @@ -42,7 +40,6 @@ class MetadataUpdater(Protocol): class CallSetup: logger: Logging kwargs: dict[str, object] - bridge_owned: bool def setup( @@ -61,9 +58,9 @@ def setup( } supplied: Final = arguments.get("litellm_logging_obj") if isinstance(supplied, Logging): - return CallSetup(supplied, arguments, bridge_owned=False) + return CallSetup(supplied, arguments) logger, prepared = function_setup(call_type, Rules(), start_time, *args, is_async_call=asynchronous, **arguments) - return CallSetup(logger, prepared, bridge_owned=True) + return CallSetup(logger, prepared) def check_limits(kwargs: Mapping[str, object]) -> None: @@ -93,87 +90,219 @@ def finalize( update(response, logger, model if isinstance(model, str) else None, kwargs, start_time, end_time) -def deployment_callbacks_needed() -> bool: - import litellm - from litellm.integrations.custom_logger import CustomLogger +class LoggingSurface(Protocol): + def update_from_kwargs( + self, + kwargs: dict[str, object], + litellm_params: dict[str, object] | None = None, + optional_params: dict[str, object] | None = None, + model: str | None = None, + user: str | None = None, + **additional_params: object, + ) -> None: ... - return any(isinstance(callback, CustomLogger) for callback in litellm.callbacks) + def pre_call( + self, input: object, api_key: object, model: object = None, additional_args: dict[str, object] = ... + ) -> object: ... + + def post_call( + self, + original_response: object, + input: object = None, + api_key: object = None, + additional_args: dict[str, object] = ..., + ) -> object: ... + + def handle_sync_success_callbacks_for_async_calls( + self, result: object, start_time: datetime.datetime, end_time: datetime.datetime, cache_hit: object = None + ) -> None: ... + + def failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime | None = None, + end_time: datetime.datetime | None = None, + ) -> None: ... + + def async_failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime | None = None, + end_time: datetime.datetime | None = None, + ) -> Coroutine[object, object, None]: ... + + def success_handler( + self, + result: object = None, + start_time: datetime.datetime | None = None, + end_time: datetime.datetime | None = None, + cache_hit: bool | None = None, + **kwargs: object, + ) -> None: ... + + def async_success_handler( + self, + result: object = None, + start_time: datetime.datetime | None = None, + end_time: datetime.datetime | None = None, + cache_hit: bool | None = None, + **kwargs: object, + ) -> Coroutine[object, object, None]: ... -Phase: TypeAlias = Literal[ - "input", "sync_success", "sync_success_async", "async_success", "sync_failure", "async_failure", "payload" -] +if TYPE_CHECKING: + _LOGGING_CONFORMS: type[LoggingSurface] = Logging -def callbacks_needed(logger: Logging, phase: Phase) -> bool: - import litellm - from litellm._logging import ( - _is_debugging_on, # pyright: ignore[reportPrivateUsage] # use the same debug gate as Logging +class LoggingWorker(Protocol): + def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: ... + + +class DeploymentHook(Protocol): + def __call__(self, kwargs: dict[str, object], call_type: str) -> Awaitable[object]: ... + + +class DeploymentSuccessHook(Protocol): + def __call__(self, request_data: dict[str, object], response: object, call_type: object) -> Awaitable[object]: ... + + +class DeploymentFailureHook(Protocol): + def __call__(self, request_data: Mapping[str, object], exception: Exception, call_type: str) -> Awaitable[None]: ... + + +def update_logging( + logger: LoggingSurface, + kwargs: dict[str, object], + model: str, + optional_params: dict[str, object], + litellm_params: dict[str, object], + custom_llm_provider: str, +) -> None: + logger.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params=litellm_params, + custom_llm_provider=custom_llm_provider, ) - if ( - _is_debugging_on() - or getattr(logger, "litellm_request_debug", False) - or os.getenv("LITELLM_PRINT_STANDARD_LOGGING_PAYLOAD") - ): - return True - input_needed: Final = bool( - litellm.input_callback - or litellm._async_input_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor - or logger.dynamic_input_callbacks - or callable(getattr(logger, "logger_fn", None)) - or logger.log_raw_request_response - or litellm.log_raw_request_response + +def pre_call(logger: LoggingSurface, input: str, api_key: str | None, additional_args: dict[str, object]) -> None: + logger.pre_call(input=input, api_key=api_key, additional_args=additional_args) + + +def post_call( + logger: LoggingSurface, original_response: str, api_key: str | None, additional_args: dict[str, object] +) -> None: + logger.post_call(original_response=original_response, api_key=api_key, additional_args=additional_args) + + +def defers_async_logging(logger: LoggingSurface) -> bool: + return bool(getattr(logger, "_defer_async_logging", False)) + + +def defer_success(logger: LoggingSurface, pending: object) -> None: + setattr(logger, "_native_pending_logging", pending) + + +def sync_success_for_async_call( + logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime +) -> None: + logger.handle_sync_success_callbacks_for_async_calls(result=response, start_time=start, end_time=end) + + +def failure_handler( + logger: LoggingSurface, error: Exception, start: datetime.datetime, end: datetime.datetime, asynchronous: bool +) -> Coroutine[object, object, None] | None: + trace: Final = "".join(traceback.format_exception(error)) + if asynchronous: + return logger.async_failure_handler(error, trace, start, end) + logger.failure_handler(error, trace, start, end) + return None + + +def submit_success(logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime) -> None: + from litellm.litellm_core_utils.litellm_logging import executor + + executor.submit(contextvars.copy_context().run, logger.success_handler, response, start, end) + + +def async_success_handler( + logger: LoggingSurface, response: object, start: datetime.datetime, end: datetime.datetime +) -> Coroutine[object, object, None]: + return logger.async_success_handler(response, start, end) + + +def enqueue_logging(coroutine: Coroutine[object, object, None]) -> None: + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + worker: Final = cast( # cast-ok: bounded adapter for the untyped logging worker + LoggingWorker, GLOBAL_LOGGING_WORKER ) - match phase: - case "input": - return input_needed - case "sync_success": - return bool(litellm.success_callback or logger.dynamic_success_callbacks) - case "sync_success_async": - return bool( - (litellm.success_callback or logger.dynamic_success_callbacks) - and logger._should_run_sync_callbacks_for_async_calls() # pyright: ignore[reportPrivateUsage] # preserve async call filtering of sync callbacks - ) - case "async_success": - return bool(litellm._async_success_callback or logger.dynamic_async_success_callbacks) # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor - case "sync_failure": - return bool(litellm.failure_callback or logger.dynamic_failure_callbacks) - case "async_failure": - return bool(litellm._async_failure_callback or logger.dynamic_async_failure_callbacks) # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor - case "payload": - return bool( - input_needed - or litellm.success_callback - or litellm.failure_callback - or litellm._async_success_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor - or litellm._async_failure_callback # pyright: ignore[reportPrivateUsage] # live async registries have no public accessor - or logger.dynamic_success_callbacks - or logger.dynamic_async_success_callbacks - or logger.dynamic_failure_callbacks - or logger.dynamic_async_failure_callbacks - ) - case _: - assert_never(phase) + contextvars.copy_context().run(worker.ensure_initialized_and_enqueue, coroutine) -def success_bookkeeping( - logger: Logging, response: object, start: datetime.datetime, end: datetime.datetime, asynchronous: bool -) -> None: - phase: Final = "async_success" if asynchronous else "sync_success" - if logger.should_run_logging(phase): - logger._success_handler_helper_fn( # pyright: ignore[reportPrivateUsage] # retain success bookkeeping without constructing a callback payload - result=response, start_time=start, end_time=end, build_logging_payload=False - ) - logger.has_run_logging(phase) +def restore_context(logger: LoggingSurface) -> None: + from litellm.utils import ( + _restore_correlation_context_if_supported, # pyright: ignore[reportPrivateUsage] # the @client wrapper restores the same correlation context + ) + + _restore_correlation_context_if_supported(logger) -def failure_bookkeeping( - logger: Logging, error: BaseException, start: datetime.datetime, end: datetime.datetime, asynchronous: bool -) -> None: - phase: Final = "async_failure" if asynchronous else "sync_failure" - if logger.should_run_logging(phase): - logger._failure_handler_helper_fn( # pyright: ignore[reportPrivateUsage] # retain failure accounting without formatting an unused traceback or payload - error, "", start, end, build_logging_payload=False - ) - logger.has_run_logging(phase) +def custom_pricing_fields() -> tuple[str, ...]: + from litellm.types.utils import CustomPricingLiteLLMParams + + return tuple(CustomPricingLiteLLMParams.model_fields) + + +def is_internal_call() -> bool: + from litellm._internal_context import is_internal_call as internal + + return internal.get() + + +def credential_list() -> list[CredentialItem]: + import litellm + + return litellm.credential_list + + +def warn_unknown_credential(name: str, loaded: int) -> None: + from litellm._logging import verbose_logger + + verbose_logger.warning( + "litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", + name, + loaded, + ) + + +def before_deployment_call(kwargs: dict[str, object], call_type: str) -> Awaitable[object]: + from litellm import utils + + hook: Final = cast( # cast-ok: bounded adapter for the untyped deployment hook + DeploymentHook, utils.async_pre_call_deployment_hook + ) + return hook(kwargs, call_type) + + +def after_deployment_success(kwargs: dict[str, object], response: object, call_type: str) -> Awaitable[object]: + from litellm import utils + from litellm.types.utils import CallTypes + + hook: Final = cast( # cast-ok: bounded adapter for the untyped deployment hook + DeploymentSuccessHook, utils.async_post_call_success_deployment_hook + ) + return hook(kwargs, response, CallTypes(call_type)) + + +def after_deployment_failure(kwargs: dict[str, object], error: Exception, call_type: str) -> Awaitable[None]: + from litellm import utils + + hook: Final = cast( # cast-ok: bounded adapter for the untyped deployment hook + DeploymentFailureHook, utils.async_post_call_failure_deployment_hook + ) + return hook(kwargs, error, call_type) diff --git a/tests/test_litellm/rust_bridge/test_legacy_callbacks.py b/tests/test_litellm/rust_bridge/test_legacy_callbacks.py index a4474c85230..a0906c7c5be 100644 --- a/tests/test_litellm/rust_bridge/test_legacy_callbacks.py +++ b/tests/test_litellm/rust_bridge/test_legacy_callbacks.py @@ -1,12 +1,16 @@ import datetime +import inspect from collections.abc import Mapping +from pathlib import Path from types import MappingProxyType from typing import Final import pytest +from pydantic import TypeAdapter import litellm from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.rust_bridge import legacy_callbacks as legacy from litellm.rust_bridge.legacy_callbacks import check_limits, setup _OCR_KWARGS: Final = MappingProxyType( @@ -56,13 +60,12 @@ def _supplied_logger() -> Logging: ) -def test_setup_adopts_a_supplied_logger_as_caller_owned() -> None: +def test_setup_reuses_a_supplied_logger() -> None: supplied: Final = _supplied_logger() result: Final = setup( "aocr", (), {**_OCR_KWARGS, "litellm_logging_obj": supplied}, datetime.datetime.now(), asynchronous=True ) assert result.logger is supplied - assert result.bridge_owned is False @pytest.mark.parametrize( @@ -73,7 +76,15 @@ def test_setup_adopts_a_supplied_logger_as_caller_owned() -> None: ], ids=["ocr", "embedding"], ) -def test_setup_owns_every_logger_it_builds(call_type: str, kwargs: Mapping[str, object]) -> None: +def test_setup_builds_a_logger_when_none_is_supplied(call_type: str, kwargs: Mapping[str, object]) -> None: result: Final = setup(call_type, (), kwargs, datetime.datetime.now(), asynchronous=True) - assert result.bridge_owned is True assert result.logger.litellm_call_id == result.kwargs["litellm_call_id"] + + +CONTRACT_PATH: Final = Path(__file__).parents[3] / "litellm-rust/crates/callbacks-legacy/python_contract.json" + + +def test_the_rust_contract_matches_the_shim_signatures() -> None: + contract: Final = TypeAdapter(dict[str, list[str]]).validate_json(CONTRACT_PATH.read_text()) + + assert contract == {name: list(inspect.signature(getattr(legacy, name)).parameters) for name in contract} diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 085ea4a14c0..5fca927bea3 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -4,7 +4,7 @@ import gc import json import threading import weakref -from collections.abc import Coroutine +from collections.abc import Awaitable, Callable, Coroutine from contextvars import ContextVar from typing import Final @@ -400,7 +400,7 @@ async def test_response_limit_is_enforced_at_the_public_boundary(ocr_server: Rec @pytest.mark.asyncio @pytest.mark.parametrize("failure", [False, True]) -async def test_empty_callbacks_keep_bookkeeping_without_optional_dispatch( +async def test_empty_callbacks_run_deployment_hooks_and_defer_like_the_python_client_wrapper( ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, failure: bool, @@ -414,8 +414,12 @@ async def test_empty_callbacks_keep_bookkeeping_without_optional_dispatch( submissions = 0 enqueues = 0 - def deployment(self, *args: object, **kwargs: object) -> None: - self.deployments += 1 + def counting(self, hook: Callable[..., Awaitable[object]]) -> Callable[..., Awaitable[object]]: + async def counted(*args: object, **kwargs: object) -> object: + self.deployments += 1 + return await hook(*args, **kwargs) + + return counted def submit(self, *args: object, **kwargs: object) -> None: self.submissions += 1 @@ -430,7 +434,7 @@ async def test_empty_callbacks_keep_bookkeeping_without_optional_dispatch( "async_post_call_success_deployment_hook", "async_post_call_failure_deployment_hook", ): - monkeypatch.setattr(utils, name, probe.deployment) + monkeypatch.setattr(utils, name, probe.counting(getattr(utils, name))) monkeypatch.setattr(litellm_logging, "executor", probe) monkeypatch.setattr(logging_worker, "GLOBAL_LOGGING_WORKER", probe) if failure: @@ -447,17 +451,16 @@ async def test_empty_callbacks_keep_bookkeeping_without_optional_dispatch( assert response._hidden_params["response_cost"] is not None assert response._hidden_params["_response_ms"] > 0 assert trace_id_var.get() == "callback-free-parent" - assert probe.deployments == probe.submissions == probe.enqueues == 0 + assert probe.deployments == 2 + assert probe.submissions == probe.enqueues == 0 assert len(created_loggers) == 1 logger: Final = created_loggers[0] - assert not hasattr(logger, "_native_pending_logging") - assert logger.model_call_details["first_api_call_start_time"] <= logger.model_call_details["end_time"] - assert "standard_logging_object" not in logger.model_call_details - assert ( - "original_response" not in logger.model_call_details or logger.model_call_details["original_response"] is None - ) - assert "complete_input_dict" not in logger.model_call_details.get("additional_args", {}) - assert logger.model_call_details["response_cost"] == (0 if failure else response._hidden_params["response_cost"]) + if failure: + assert logger.model_call_details["first_api_call_start_time"] <= logger.model_call_details["end_time"] + assert logger.model_call_details["response_cost"] == 0 + else: + assert getattr(logger, "_native_pending_logging", None) is not None + assert "end_time" not in logger.model_call_details @pytest.mark.asyncio From a737e3430a979b78c4d84ceac6d7605aca729d79 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 14:41:50 -0700 Subject: [PATCH 234/253] test(rust): property and parametrized tests for legacy callback contracts The payload boundary of callbacks-legacy gets a model-based proptest: for any JSON body, caller keywords and callback edit, keywords the route sends unchanged reach pre_call as the caller's own objects, and the wire is the body pre_call received as the callback left it. A parametrized test pins that a keyword the bridge never reads keeps its identity through setup, the deployment hook, check_limits and prepare. Behaviour owned by the real Logging object is pinned end to end in the OCR tests: a hypothesis version of the body property over HTTP, sync hooks seeing no running event loop, retained payloads staying intact after the call, success callbacks sharing one standard logging payload, and state stashed before a blocking deployment hook raises reaching both failure callback families --- litellm-rust/Cargo.toml | 1 + .../crates/callbacks-legacy/Cargo.toml | 1 + .../tests/deployment_hooks.rs | 37 ++++ .../crates/callbacks-legacy/tests/payload.rs | 175 +++++++++++++++++- tests/test_litellm_rust/ocr/test_callbacks.py | 163 +++++++++++++++- 5 files changed, 370 insertions(+), 7 deletions(-) diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index ffdbf64bb49..32f925b8d8b 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -26,6 +26,7 @@ litellm-token-counter = { path = "crates/token-counter" } litellm-host-python = { path = "crates/host-python" } bytes = "1" +proptest = "1.7.0" pyo3 = "0.29.2" pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] } pythonize = "0.29.0" diff --git a/litellm-rust/crates/callbacks-legacy/Cargo.toml b/litellm-rust/crates/callbacks-legacy/Cargo.toml index 3cf9382f857..efacde051ed 100644 --- a/litellm-rust/crates/callbacks-legacy/Cargo.toml +++ b/litellm-rust/crates/callbacks-legacy/Cargo.toml @@ -16,5 +16,6 @@ serde_json.workspace = true [dev-dependencies] litellm-auth.workspace = true +proptest.workspace = true rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs index ea3510de17e..7bad09c7890 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs @@ -97,6 +97,43 @@ assert checked is prepared }); } +#[rstest] +#[case::synchronous(false)] +#[case::asynchronous(true)] +fn a_keyword_the_bridge_never_reads_reaches_every_reader_as_the_callers_object( + #[case] asynchronous: bool, +) { + Python::initialize(); + Python::attach(|py| { + let locals = namespace( + py, + c" +opaque = object() +hooked = [] +logger.hooks = {'pre': lambda kwargs: hooked.append(kwargs['vendor_extension']) or kwargs} +kwargs = {'logger': logger, 'vendor_extension': opaque} +", + ); + let (mut logging, step) = begin(py, &locals, asynchronous); + let step = match step { + LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(), + step => step, + }; + locals.set_item("prepared", arguments(py, step)).unwrap(); + locals.set_item("asynchronous", asynchronous).unwrap(); + run( + py, + &locals, + c" +assert prepared['vendor_extension'] is opaque +[checked] = [value for name, value in logger.calls if name == 'check_limits'] +assert checked['vendor_extension'] is opaque +assert hooked == ([opaque] if asynchronous else []), hooked +", + ); + }); +} + #[test] fn response_returned_by_the_post_call_hook_is_finalized_and_returned() { Python::initialize(); diff --git a/litellm-rust/crates/callbacks-legacy/tests/payload.rs b/litellm-rust/crates/callbacks-legacy/tests/payload.rs index 68c0a2b1e15..5c49383bf37 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/payload.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/payload.rs @@ -2,10 +2,11 @@ use std::ffi::CStr; use litellm_auth::SecretValue; use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest}; -use litellm_host_python::{LifecycleStep, PythonLifecycle}; +use litellm_host_python::{LifecycleStep, PythonLifecycle, to_py}; +use proptest::prelude::*; use pyo3::prelude::*; use rstest::rstest; -use serde_json::{Value, json}; +use serde_json::{Map, Value, json}; use super::LegacyLogging; use crate::PythonLogger; @@ -57,10 +58,24 @@ fn before_send_with_secrets( optional_params: Value, body: Value, secret_fields: &[&str], +) -> WireRequest { + before_send_bound(&[], script, optional_params, body, secret_fields) +} + +/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs. +fn before_send_bound( + bindings: &[(&str, &Value)], + script: &CStr, + optional_params: Value, + body: Value, + secret_fields: &[&str], ) -> WireRequest { Python::initialize(); Python::attach(|py| { let locals = namespace(py, PAYLOAD_LOGGER); + for &(name, value) in bindings { + locals.set_item(name, to_py(py, value).unwrap()).unwrap(); + } run(py, &locals, script); let mut logging = LegacyLogging { logger: Some(PythonLogger::new(local(&locals, "logger").unbind())), @@ -346,3 +361,159 @@ def check(): json!({"document": document(DOCUMENT), "include_image_base64": true}) ); } + +/// What one pre-call callback does to the payload it is handed. +#[derive(Clone, Debug)] +enum Edit { + Nothing, + Set(String, Value), + Remove(String), + Rebind(Value), + RebindThenSetRetained(String, Value), +} + +impl Edit { + fn script(&self) -> Value { + match self { + Self::Nothing => json!({"kind": "nothing"}), + Self::Set(key, value) => json!({"kind": "set", "key": key, "value": value}), + Self::Remove(key) => json!({"kind": "remove", "key": key}), + Self::Rebind(value) => json!({"kind": "rebind", "value": value}), + Self::RebindThenSetRetained(key, value) => { + json!({"kind": "rebind_then_set_retained", "key": key, "value": value}) + } + } + } + + /// The legacy contract: the provider is sent the body object `pre_call` received, as + /// the callback left it. Rebinding the envelope's key points the envelope elsewhere and + /// leaves that object alone. + fn sent(&self, body: &Map) -> Value { + let mut sent = body.clone(); + match self { + Self::Nothing | Self::Rebind(_) => {} + Self::Set(key, value) | Self::RebindThenSetRetained(key, value) => { + sent.insert(key.clone(), value.clone()); + } + Self::Remove(key) => { + sent.remove(key); + } + } + Value::Object(sent) + } +} + +/// How the caller's keyword for a body key relates to what the route sends under it. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Caller { + PassedUnchanged, + RewrittenByTheRoute, + NotPassed, +} + +const MODEL: &CStr = c" +aliased = {} +def on_pre_call(args): + body = args['complete_input_dict'] + aliased.update({name: body[name] is kwargs[name] for name in unchanged}) + kind = edit['kind'] + if kind == 'set': + body[edit['key']] = edit['value'] + elif kind == 'remove': + body.pop(edit['key'], None) + elif kind == 'rebind': + args['complete_input_dict'] = edit['value'] + elif kind == 'rebind_then_set_retained': + args['complete_input_dict'] = {} + body[edit['key']] = edit['value'] +def check(): + assert aliased == {name: True for name in unchanged}, aliased + assert logger.names() == ['pre_call', 'post_call'], logger.calls +"; + +fn json_value() -> impl Strategy { + let leaf = prop_oneof![ + Just(Value::Null), + any::().prop_map(Value::from), + any::().prop_map(Value::from), + any::() + .prop_filter("JSON has no NaN or infinity", |number| number.is_finite()) + .prop_map(Value::from), + ".{0,8}".prop_map(Value::from), + ]; + leaf.prop_recursive(3, 24, 4, |inner| { + prop_oneof![ + prop::collection::vec(inner.clone(), 0..4).prop_map(Value::from), + prop::collection::btree_map(key(), inner, 0..4) + .prop_map(|fields| Value::Object(fields.into_iter().collect())), + ] + }) +} + +fn key() -> impl Strategy { + "[a-z]{1,6}" +} + +fn caller() -> impl Strategy { + prop_oneof![ + Just(Caller::PassedUnchanged), + Just(Caller::RewrittenByTheRoute), + Just(Caller::NotPassed), + ] +} + +fn edit() -> impl Strategy { + prop_oneof![ + Just(Edit::Nothing), + (key(), json_value()).prop_map(|(key, value)| Edit::Set(key, value)), + key().prop_map(Edit::Remove), + json_value().prop_map(Edit::Rebind), + (key(), json_value()).prop_map(|(key, value)| Edit::RebindThenSetRetained(key, value)), + ] +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(128))] + + /// For any body, any caller keywords and any callback edit: every keyword the route + /// sends unchanged reaches `pre_call` as the caller's own object, and the provider is + /// sent exactly what the model says, so a callback that edits nothing changes nothing. + #[test] + fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it( + fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5), + edit in edit(), + ) { + let body: Map = fields + .iter() + .map(|(name, (value, _))| (name.clone(), value.clone())) + .collect(); + let kwargs: Map = fields + .iter() + .filter_map(|(name, (value, caller))| match caller { + Caller::PassedUnchanged => Some((name.clone(), value.clone())), + Caller::RewrittenByTheRoute => Some((name.clone(), json!([value]))), + Caller::NotPassed => None, + }) + .collect(); + let unchanged: Value = fields + .iter() + .filter(|(_, (_, caller))| *caller == Caller::PassedUnchanged) + .map(|(name, _)| Value::from(name.clone())) + .collect(); + + let wire = before_send_bound( + &[ + ("kwargs", &Value::Object(kwargs)), + ("unchanged", &unchanged), + ("edit", &edit.script()), + ], + MODEL, + json!({}), + Value::Object(body.clone()), + &[], + ); + + prop_assert_eq!(wire.body, edit.sent(&body)); + prop_assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]); + } +} diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 27cdcc4d997..7cc19b8c090 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -1,24 +1,28 @@ import asyncio import copy +import gc import queue import threading from typing import Final import pytest +from hypothesis import HealthCheck, given, settings +from hypothesis import strategies as st import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse -from tests.test_litellm_rust.support.callback_recorder import RecordingLogger +from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import ( OCR_DOCUMENT, OCR_RESPONSE, + call_native, call_native_aocr, call_native_ocr, request_body, request_headers, ) -from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec pytestmark = pytest.mark.requires_rust_extension @@ -123,9 +127,7 @@ async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_ "callbacks": [Retain(), Edit()], } response: Final = ( - await call_native_aocr(ocr_server, **arguments) - if asynchronous - else call_native_ocr(ocr_server, **arguments) + await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments) ) assert aliases == [True] @@ -291,6 +293,157 @@ def test_native_ocr_dispatches_each_callback_phase_once_when_logger_is_registere assert "log_failure_event" not in recorder.names +JSON_SCALARS: Final = ( + st.none() + | st.booleans() + | st.integers(min_value=-(2**63), max_value=2**63 - 1) + | st.floats(allow_nan=False, allow_infinity=False) + | st.text(max_size=8) +) +JSON_VALUES: Final = st.recursive( + JSON_SCALARS, + lambda children: st.lists(children, max_size=3) | st.dictionaries(st.text(max_size=6), children, max_size=3), + max_leaves=8, +) + + +LATEST_EDITS: Final[list[dict[str, object]]] = [] + + +class ApplyLatestEdits(CustomLogger): + """Registrations can outlive one hypothesis example, so every instance applies the current example's edits.""" + + def __init__(self, latest: list[dict[str, object]]) -> None: + super().__init__() + self.latest = latest + + def log_pre_api_call(self, model, messages, kwargs): + request_body(kwargs).update(copy.deepcopy(self.latest[-1])) + + +@settings(max_examples=25, deadline=None, suppress_health_check=[HealthCheck.function_scoped_fixture]) +@given(edits=st.dictionaries(st.from_regex(r"x_[a-z]{1,6}", fullmatch=True), JSON_VALUES, max_size=3)) +def test_native_ocr_provider_receives_the_body_exactly_as_pre_call_callbacks_left_it( + ocr_server: RecordingServer, edits: dict[str, object] +) -> None: + LATEST_EDITS.append(edits) + + call_native_ocr_with_callbacks(ocr_server, [ApplyLatestEdits(LATEST_EDITS)]) + + assert ocr_server.requests[-1].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT, **edits} + + +@pytest.mark.parametrize("hook", ["log_pre_api_call", "logging_hook", "log_success_event"]) +def test_native_ocr_sync_hooks_see_no_running_event_loop(ocr_server: RecordingServer, hook: str) -> None: + recorder: Final = RecordingLogger() + + call_native_ocr_with_callbacks(ocr_server, [recorder]) + + [event] = recorder.wait_for(hook) + assert event.loop is None + assert (event.thread is threading.current_thread()) == (hook == "log_pre_api_call") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +async def test_native_ocr_payload_a_callback_retains_outlives_the_call_intact( + ocr_server: RecordingServer, asynchronous: bool +) -> None: + retained: Final = [] + + class Retain(CustomLogger): + def log_pre_api_call(self, model, messages, kwargs): + retained.append((kwargs, request_body(kwargs), request_headers(kwargs))) + + await call_native(ocr_server, asynchronous, callbacks=[Retain()]) + await drain_logging() + gc.collect() + + [(details, body, headers)] = retained + assert body == ocr_server.requests[0].body + assert headers + assert all(ocr_server.requests[0].headers[name] == value for name, value in headers.items()) + assert details["additional_args"]["complete_input_dict"] is body + assert details["additional_args"]["headers"] is headers + + +@pytest.mark.asyncio +@pytest.mark.parametrize("family", ["sync", "async"]) +async def test_native_ocr_success_callbacks_share_one_logging_payload(ocr_server: RecordingServer, family: str) -> None: + queued: Final = [] + finished: Final = threading.Event() + + def queue_payload(kwargs: dict[str, object]) -> None: + queued.append(kwargs["standard_logging_object"]) + + def strip_payload(kwargs: dict[str, object]) -> None: + payload: Final = kwargs["standard_logging_object"] + assert isinstance(payload, dict) + payload["stripped-by-a-later-callback"] = True + finished.set() + + class QueuePayload(CustomLogger): + if family == "sync": + + def log_success_event(self, kwargs, response_obj, start_time, end_time): + queue_payload(kwargs) + + else: + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + queue_payload(kwargs) + + class StripPayload(CustomLogger): + if family == "sync": + + def log_success_event(self, kwargs, response_obj, start_time, end_time): + strip_payload(kwargs) + + else: + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + strip_payload(kwargs) + + await call_native(ocr_server, family == "async", callbacks=[QueuePayload(), StripPayload()]) + await drain_logging() + + assert await asyncio.to_thread(finished.wait, 10) + assert [payload["stripped-by-a-later-callback"] for payload in queued] == [True] + + +@pytest.mark.asyncio +async def test_native_aocr_state_stashed_before_a_blocking_hook_raises_reaches_failure_callbacks( + ocr_server: RecordingServer, +) -> None: + token: Final = object() + observed: Final = [] + + class Blocked(Exception): + pass + + class Block(CustomLogger): + async def async_post_call_success_deployment_hook(self, request_data, response, call_type): + request_data["litellm_logging_obj"].model_call_details["blocked-by"] = token + raise Blocked("blocked after the provider answered") + + def log_success_event(self, kwargs, response_obj, start_time, end_time): + observed.append(("success", None, None)) + + def log_failure_event(self, kwargs, response_obj, start_time, end_time): + observed.append(("sync", kwargs.get("blocked-by"), kwargs["exception"])) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + observed.append(("async", kwargs.get("blocked-by"), kwargs["exception"])) + + litellm.callbacks.append(Block()) + + with pytest.raises(Blocked) as raised: + await call_native_aocr(ocr_server) + await drain_logging() + + assert observed == [("sync", token, raised.value), ("async", token, raised.value)] + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context( From 30f7b8442bd8b7258ad4bc5a41a9b85546d7a83c Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 14:42:59 -0700 Subject: [PATCH 235/253] chore(rust): lock proptest --- litellm-rust/Cargo.lock | 62 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0265e0adbc2..bf58a81a6c5 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2044,6 +2044,7 @@ dependencies = [ "litellm-auth", "litellm-callbacks", "litellm-host-python", + "proptest", "pyo3", "rstest", "serde_json", @@ -2601,6 +2602,25 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "proptest" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744" +dependencies = [ + "bit-set", + "bit-vec", + "bitflags", + "num-traits", + "rand 0.9.5", + "rand_chacha 0.9.0", + "rand_xorshift", + "regex-syntax", + "rusty-fork", + "tempfile", + "unarray", +] + [[package]] name = "pyo3" version = "0.29.2" @@ -2682,6 +2702,12 @@ dependencies = [ "serde", ] +[[package]] +name = "quick-error" +version = "1.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0" + [[package]] name = "quinn" version = "0.11.11" @@ -2845,6 +2871,15 @@ dependencies = [ "rand_core 0.10.1", ] +[[package]] +name = "rand_xorshift" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "513962919efc330f829edb2535844d1b912b0fbe2ca165d613e4e8788bb05a5a" +dependencies = [ + "rand_core 0.9.5", +] + [[package]] name = "rayon" version = "1.12.0" @@ -3244,6 +3279,18 @@ version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" +[[package]] +name = "rusty-fork" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc6bf79ff24e648f6da1f8d1f011e9cac26491b619e6b9280f2b47f1774e6ee2" +dependencies = [ + "fnv", + "quick-error", + "tempfile", + "wait-timeout", +] + [[package]] name = "ryu" version = "1.0.23" @@ -4099,6 +4146,12 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "unarray" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eaea85b334db583fe3274d12b4cd1880032beab409c0d774be044d4480ab9a94" + [[package]] name = "unicase" version = "2.9.0" @@ -4212,6 +4265,15 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5c3082ca00d5a5ef149bb8b555a72ae84c9c59f7250f013ac822ac2e49b19c64" +[[package]] +name = "wait-timeout" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ac3b126d3914f9849036f826e054cbabdc8519970b8998ddaf3b5bd3c65f11" +dependencies = [ + "libc", +] + [[package]] name = "walkdir" version = "2.5.0" From ab27ee7efcbe982991095dbfbf5bd01662f8e69c Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 14:44:19 -0700 Subject: [PATCH 236/253] test(rust): fix request count and header case in new OCR callback tests --- tests/test_litellm_rust/ocr/test_callbacks.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 7cc19b8c090..45b99d19d90 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -326,6 +326,7 @@ class ApplyLatestEdits(CustomLogger): def test_native_ocr_provider_receives_the_body_exactly_as_pre_call_callbacks_left_it( ocr_server: RecordingServer, edits: dict[str, object] ) -> None: + ocr_server.expected_requests = None LATEST_EDITS.append(edits) call_native_ocr_with_callbacks(ocr_server, [ApplyLatestEdits(LATEST_EDITS)]) @@ -362,7 +363,7 @@ async def test_native_ocr_payload_a_callback_retains_outlives_the_call_intact( [(details, body, headers)] = retained assert body == ocr_server.requests[0].body assert headers - assert all(ocr_server.requests[0].headers[name] == value for name, value in headers.items()) + assert all(ocr_server.requests[0].headers[name.lower()] == value for name, value in headers.items()) assert details["additional_args"]["complete_input_dict"] is body assert details["additional_args"]["headers"] is headers From dd2e6c17bc19b839ff0a5e8500dfe40a7612f430 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 21:56:13 +0000 Subject: [PATCH 237/253] fix(rust): preserve OCR callback headers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/callbacks-legacy/src/adapter.rs | 12 ++++++++++-- .../crates/callbacks-legacy/src/callbacks.rs | 3 +++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index 7204207cd62..53aa6ce9d2e 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -44,6 +44,7 @@ pub struct LegacyLogging { response: Option>, error: Option>, body: Option>, + headers: Option>, context: Option, asynchronous: bool, internal: bool, @@ -77,6 +78,7 @@ impl LegacyLogging { response: None, error: None, body: None, + headers: None, context: None, asynchronous, internal: false, @@ -245,6 +247,7 @@ impl PythonLifecycle for LegacyLogging { headers.set_item(name, value)?; } self.body = Some(body.clone().unbind()); + self.headers = Some(headers.clone().unbind()); self.context = Some(context.clone()); self.logger()?.pre_call( py, @@ -299,8 +302,13 @@ impl PythonLifecycle for LegacyLogging { .as_ref() .and_then(|context| context.api_key.as_ref()) .map(|api_key| api_key.expose()); - self.logger()? - .post_call(py, &raw.body, api_key, self.body.as_ref())?; + self.logger()?.post_call( + py, + &raw.body, + api_key, + self.body.as_ref(), + self.headers.as_ref(), + )?; Ok(LifecycleStep::Done) } (CallEvent::Succeeded { timing }, Some(PublicValue::Response(response))) => { diff --git a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs index 9464f1d6612..9fcfe98368e 100644 --- a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs +++ b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs @@ -38,6 +38,7 @@ pub trait LegacyCallbacks { original_response: &str, api_key: Option<&str>, body: Option<&Py>, + headers: Option<&Py>, ) -> PyResult<()>; fn defers_async_logging(&self, py: Python<'_>) -> bool; @@ -150,9 +151,11 @@ impl LegacyCallbacks for PythonLogger { original_response: &str, api_key: Option<&str>, body: Option<&Py>, + headers: Option<&Py>, ) -> PyResult<()> { let additional = PyDict::new(py); additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; Logging::PostCall.call( py, (self.object(py), original_response, api_key, &additional), From 4627ec4ea8f4ab6b53388f079bab54075d05a0ea Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 21:58:14 +0000 Subject: [PATCH 238/253] test(rust): cover post-call header identity Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/callbacks-legacy/tests/payload.rs | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/litellm-rust/crates/callbacks-legacy/tests/payload.rs b/litellm-rust/crates/callbacks-legacy/tests/payload.rs index 5c49383bf37..43128bc38ea 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/payload.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/payload.rs @@ -331,15 +331,19 @@ def on_pre_call(args): } #[test] -fn post_call_receives_the_raw_response_the_route_key_and_the_body_pre_call_saw() { +fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() { before_send( c" def check(): original_response, api_key, additional_args = logger.post assert original_response == 'raw response', original_response assert api_key == logger.pre_api_key == 'route-key', (api_key, logger.pre_api_key) - assert additional_args == {'complete_input_dict': logger.pre['complete_input_dict']}, additional_args + assert additional_args == { + 'complete_input_dict': logger.pre['complete_input_dict'], + 'headers': logger.pre['headers'], + }, additional_args assert additional_args['complete_input_dict'] is logger.pre['complete_input_dict'] + assert additional_args['headers'] is logger.pre['headers'] ", json!({"document": document(DOCUMENT)}), ); From 1a52bae77950ddab49f1d53a0a511fdaf1d7b570 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:32:28 -0700 Subject: [PATCH 239/253] add streaming message --- .../callbacks-legacy/python_contract.json | 17 ++ .../crates/callbacks-legacy/src/adapter.rs | 90 +++++++- .../callbacks-legacy/src/legacy_python.rs | 30 ++- .../crates/callbacks-legacy/tests/support.rs | 5 + litellm-rust/crates/callbacks/src/event.rs | 4 + litellm-rust/crates/callbacks/src/host.rs | 21 ++ litellm-rust/crates/callbacks/src/route.rs | 5 + litellm-rust/crates/callbacks/src/run.rs | 4 + litellm-rust/crates/core/src/machine/mod.rs | 17 +- .../crates/core/src/messages/handler.rs | 128 +++++------- litellm-rust/crates/core/src/messages/mod.rs | 36 +++- .../crates/core/src/messages/prepare.rs | 11 +- .../crates/core/src/messages/route.rs | 197 ++++++++++++++++++ litellm-rust/crates/core/src/ocr/route.rs | 2 + litellm-rust/crates/core/tests/ocr.rs | 2 + .../crates/host-python/src/adapter.rs | 8 + litellm-rust/crates/host-python/src/driver.rs | 77 ++++++- litellm-rust/crates/host-python/src/handle.rs | 13 ++ .../crates/python-bridge/src/marshal.rs | 16 -- .../python-bridge/src/routes/messages.rs | 88 -------- .../python-bridge/src/routes/messages/host.rs | 186 +++++++++++++++++ .../python-bridge/src/routes/messages/mod.rs | 64 ++++++ .../crates/python-bridge/src/routes/mod.rs | 31 --- .../python-bridge/src/routes/ocr/host.rs | 4 + litellm/rust_bridge/_native.pyi | 28 +-- litellm/rust_bridge/catalog.py | 1 + litellm/rust_bridge/failures.py | 43 ++++ litellm/rust_bridge/legacy_callbacks.py | 80 +++++++ litellm/rust_bridge/lifecycle.py | 137 ++++++++++-- litellm/rust_bridge/messages/entrypoints.py | 10 +- litellm/rust_bridge/messages/route_host.py | 2 +- litellm/rust_bridge/ocr/route_host.py | 40 +--- .../rust_bridge/native_route_wheel_test.py | 31 +-- .../test_litellm/rust_bridge/test_bindings.py | 4 +- .../test_litellm/rust_bridge/test_catalog.py | 4 + .../test_litellm/rust_bridge/test_runtime.py | 1 - tests/test_litellm_rust/messages/__init__.py | 0 .../messages/test_callbacks.py | 175 ++++++++++++++++ .../support/recording_server.py | 16 +- tests/test_litellm_rust/support/requests.py | 35 ++++ 40 files changed, 1319 insertions(+), 344 deletions(-) create mode 100644 litellm-rust/crates/core/src/messages/route.rs delete mode 100644 litellm-rust/crates/python-bridge/src/routes/messages.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/messages/host.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/messages/mod.rs create mode 100644 tests/test_litellm_rust/messages/__init__.py create mode 100644 tests/test_litellm_rust/messages/test_callbacks.py diff --git a/litellm-rust/crates/callbacks-legacy/python_contract.json b/litellm-rust/crates/callbacks-legacy/python_contract.json index 840c0abfa45..a09bdc711a3 100644 --- a/litellm-rust/crates/callbacks-legacy/python_contract.json +++ b/litellm-rust/crates/callbacks-legacy/python_contract.json @@ -94,5 +94,22 @@ "kwargs", "error", "call_type" + ], + "stream_opened": [ + "logger" + ], + "stream_success": [ + "logger", + "request_body", + "chunks", + "start", + "end", + "first_chunk" + ], + "stream_failure": [ + "logger", + "request_body", + "chunks", + "error" ] } diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index 53aa6ce9d2e..9d28db92add 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -2,7 +2,9 @@ //! raises is answered with the same `Logging` calls, in the same order, as the Python //! `@client` path makes them. -use litellm_callbacks::event::{CallEvent, FailureOrigin, RequestContext, Timing, WireRequest}; +use litellm_callbacks::event::{ + CallEvent, FailureOrigin, RequestContext, Timing, WireRequest, epoch_seconds, +}; use litellm_host_python::{ LifecycleStep, PublicValue, PythonLifecycle, from_py, missing_state, to_py, }; @@ -10,14 +12,16 @@ use pyo3::{ exceptions::{PyBaseException, PyException}, gc::{PyTraverseError, PyVisit}, prelude::*, - types::PyDict, + types::{PyDict, PyList}, }; use serde_json::Value; use crate::{ DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger, deferred::{PendingLogging, PendingSuccess}, - finalize, is_internal_call, prepare, setup, + finalize, is_internal_call, + legacy_python::Streaming, + prepare, setup, }; /// What the legacy contract needs to know about the route it is logging. @@ -28,6 +32,12 @@ pub struct LegacySurface { pub input_description: &'static str, } +/// What the Messages stream iterator keeps for its end-of-stream billing. +struct DeliveredStream { + chunks: Py, + first_chunk: Option>, +} + enum Pending { DeploymentPreCall, DeploymentPostCall, @@ -46,6 +56,7 @@ pub struct LegacyLogging { body: Option>, headers: Option>, context: Option, + stream: Option, asynchronous: bool, internal: bool, pending: Option, @@ -80,6 +91,7 @@ impl LegacyLogging { body: None, headers: None, context: None, + stream: None, asynchronous, internal: false, pending: None, @@ -163,6 +175,49 @@ impl LegacyLogging { logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end) } + fn stream_success(&self, py: Python<'_>, stream: &DeliveredStream) -> PyResult<()> { + let logger = self.logger()?; + let billed = Streaming::Success.call( + py, + ( + logger.object(py), + &self.body, + &stream.chunks, + &self.start, + &self.end, + &stream.first_chunk, + ), + ); + match billed { + Err(error) if error.is_instance_of::(py) => { + error.write_unraisable(py, Some(logger.object(py))); + Ok(()) + } + result => result.map(|_| ()), + } + } + + /// A failure after the stream reached the caller bills the delivered chunks as + /// partial usage. The sync path has no loop to schedule that on, so it falls back to + /// the plain failure handler. + fn stream_failure(&mut self, py: Python<'_>) -> PyResult { + let (Some(logger), Some(error), Some(stream)) = (&self.logger, &self.error, &self.stream) + else { + return Ok(LifecycleStep::Done); + }; + if !self.asynchronous { + return self.dispatch_failure(py); + } + match Streaming::Failure.call(py, (logger.object(py), &self.body, &stream.chunks, error)) { + Ok(awaitable) => { + self.pending = Some(Pending::AsyncFailure); + Ok(LifecycleStep::Await(awaitable.unbind())) + } + Err(failure) if is_cancellation(py, &failure) => Err(failure), + Err(_) => Ok(LifecycleStep::Done), + } + } + /// The sync failure handler, then the async one for async calls. Ordinary handler /// errors never replace the selected failure or suppress the other family; a /// cancellation does end the call. @@ -296,6 +351,22 @@ impl PythonLifecycle for LegacyLogging { ) -> PyResult { match (event, public) { (CallEvent::Started { .. }, _) => Ok(LifecycleStep::Done), + (CallEvent::Opened, _) => { + Streaming::Opened.call(py, (self.logger()?.object(py),))?; + self.stream = Some(DeliveredStream { + chunks: PyList::empty(py).unbind(), + first_chunk: None, + }); + Ok(LifecycleStep::Done) + } + (CallEvent::Delivered, Some(PublicValue::Chunk(chunk))) => { + let stream = self.stream.as_mut().ok_or_else(missing_state)?; + if stream.first_chunk.is_none() { + stream.first_chunk = Some(datetime(py, epoch_seconds())?); + } + stream.chunks.bind(py).append(chunk)?; + Ok(LifecycleStep::Done) + } (CallEvent::ResponseReceived { raw }, _) => { let api_key = self .context @@ -314,12 +385,18 @@ impl PythonLifecycle for LegacyLogging { (CallEvent::Succeeded { timing }, Some(PublicValue::Response(response))) => { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response.clone_ref(py)); - self.dispatch_success(py)?; + match &self.stream { + Some(stream) => self.stream_success(py, stream)?, + None => self.dispatch_success(py)?, + } Ok(LifecycleStep::Done) } (CallEvent::Failed { timing, origin }, Some(PublicValue::Error(error))) => { self.end = Some(datetime(py, timing.end_time)?); self.error = Some(error.clone_ref(py).into_value(py)); + if self.stream.is_some() { + return self.stream_failure(py); + } if *origin == FailureOrigin::Call && self.logger.is_some() && self.runs_deployment_hooks() @@ -366,6 +443,7 @@ impl PythonLifecycle for LegacyLogging { } self.body = None; self.context = None; + self.stream = None; } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { @@ -377,6 +455,10 @@ impl PythonLifecycle for LegacyLogging { visit.call(&self.end)?; visit.call(&self.response)?; visit.call(&self.error)?; + if let Some(stream) = &self.stream { + visit.call(&stream.chunks)?; + visit.call(&stream.first_chunk)?; + } visit.call(&self.body) } } diff --git a/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs b/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs index a924775070c..7f5c77c1735 100644 --- a/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs +++ b/litellm-rust/crates/callbacks-legacy/src/legacy_python.rs @@ -16,6 +16,7 @@ pub(crate) enum LegacyPython { Wrapper(Wrapper), Logging(Logging), DeploymentHooks(DeploymentHooks), + Streaming(Streaming), } /// The `@client` wrapper around the call: `function_setup`, limits, credentials, @@ -76,12 +77,25 @@ pub(crate) enum DeploymentHooks { AfterDeploymentFailure, } +/// The Messages stream iterator's logging: the stream flag, the end-of-stream billing +/// from the delivered chunks, and the partial-usage failure path. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum Streaming { + #[strum(serialize = "stream_opened")] + Opened, + #[strum(serialize = "stream_success")] + Success, + #[strum(serialize = "stream_failure")] + Failure, +} + impl LegacyPython { fn name(self) -> &'static str { match self { Self::Wrapper(function) => function.into(), Self::Logging(function) => function.into(), Self::DeploymentHooks(function) => function.into(), + Self::Streaming(function) => function.into(), } } @@ -111,6 +125,15 @@ impl Logging { } } +impl Streaming { + pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + LegacyPython::Streaming(self).call(py, args) + } +} + impl DeploymentHooks { pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> where @@ -126,7 +149,7 @@ mod tests { use strum::VariantArray; - use super::{DeploymentHooks, LegacyPython, Logging, Wrapper}; + use super::{DeploymentHooks, LegacyPython, Logging, Streaming, Wrapper}; use crate::test_support::PYTHON_CONTRACT; #[test] @@ -147,6 +170,11 @@ mod tests { .iter() .map(|&function| LegacyPython::DeploymentHooks(function)), ) + .chain( + Streaming::VARIANTS + .iter() + .map(|&function| LegacyPython::Streaming(function)), + ) .map(LegacyPython::name) .collect(); assert_eq!(called.len(), declared.len(), "a function is borrowed twice"); diff --git a/litellm-rust/crates/callbacks-legacy/tests/support.rs b/litellm-rust/crates/callbacks-legacy/tests/support.rs index 444655ea77b..42ca184eb16 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/support.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/support.rs @@ -84,6 +84,11 @@ FAKES = { 'success', response, call_type ), 'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type), + 'stream_opened': lambda logger: logger.record('stream_opened', None), + 'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record( + 'stream_success', list(chunks) + ), + 'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error), } assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys()) for name, fake in FAKES.items(): diff --git a/litellm-rust/crates/callbacks/src/event.rs b/litellm-rust/crates/callbacks/src/event.rs index 3bf12e553b9..d19e973b812 100644 --- a/litellm-rust/crates/callbacks/src/event.rs +++ b/litellm-rust/crates/callbacks/src/event.rs @@ -60,6 +60,10 @@ pub enum CallEvent { ResponseReceived { raw: RawResponse, }, + /// The call streams and its stream was handed to the caller. + Opened, + /// One chunk of an open stream reached the caller. + Delivered, Succeeded { timing: Timing, }, diff --git a/litellm-rust/crates/callbacks/src/host.rs b/litellm-rust/crates/callbacks/src/host.rs index 2392718a18d..eef3e1da8d5 100644 --- a/litellm-rust/crates/callbacks/src/host.rs +++ b/litellm-rust/crates/callbacks/src/host.rs @@ -11,12 +11,25 @@ pub enum HostOp { context: Box, }, Emit(CallEvent), + /// The response streams: the host hands the caller a stream and answers once the + /// caller asks for the first chunk or goes away. + Open(R::StreamHead), + /// The next chunk of an open stream, answered once the caller asks for the one after. + Deliver(R::Chunk), } pub enum HostResult { Route(R::OpResult), BeforeSend(Box), Emitted, + Demand(Demand), +} + +/// Whether the caller of a streamed call still reads it. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum Demand { + More, + Detached, } /// A host answer that is either available now or arrives once the host's own @@ -42,4 +55,12 @@ pub trait Host: Send + Sync { fn emit(&self, _event: &CallEvent) -> impl Future> + Send { async { Ok(()) } } + + fn open(&self, _head: R::StreamHead) -> impl Future> + Send { + async { Ok(Demand::More) } + } + + fn deliver(&self, _chunk: R::Chunk) -> impl Future> + Send { + async { Ok(Demand::More) } + } } diff --git a/litellm-rust/crates/callbacks/src/route.rs b/litellm-rust/crates/callbacks/src/route.rs index 97738c8da8b..8ab2b125760 100644 --- a/litellm-rust/crates/callbacks/src/route.rs +++ b/litellm-rust/crates/callbacks/src/route.rs @@ -6,4 +6,9 @@ pub trait Route: Send + Sync + 'static { type Error: Clone + Send + Sync + 'static; type Op: Send + 'static; type OpResult: Send + 'static; + /// One piece of a streamed response, handed to the caller as it arrives. A route + /// that never streams uses `Infallible`. + type Chunk: Send + 'static; + /// What the route knows once a streamed response starts, before its first chunk. + type StreamHead: Send + 'static; } diff --git a/litellm-rust/crates/callbacks/src/run.rs b/litellm-rust/crates/callbacks/src/run.rs index 5accaa2e25e..705c5504c01 100644 --- a/litellm-rust/crates/callbacks/src/run.rs +++ b/litellm-rust/crates/callbacks/src/run.rs @@ -26,6 +26,8 @@ where .await .map(|wire| HostResult::BeforeSend(Box::new(wire))), HostOp::Emit(event) => host.emit(&event).await.map(|()| HostResult::Emitted), + HostOp::Open(head) => host.open(head).await.map(HostResult::Demand), + HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand), }; match answer { Ok(answer) => result = Some(answer), @@ -61,6 +63,8 @@ mod tests { type Error = &'static str; type Op = &'static str; type OpResult = (); + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } struct Scripted { diff --git a/litellm-rust/crates/core/src/machine/mod.rs b/litellm-rust/crates/core/src/machine/mod.rs index 279a2d65c97..929f0a423c4 100644 --- a/litellm-rust/crates/core/src/machine/mod.rs +++ b/litellm-rust/crates/core/src/machine/mod.rs @@ -9,7 +9,7 @@ use std::{future::Future, pin::Pin}; pub use auth::{HostTokenProvider, TokenRoute}; use litellm_callbacks::{ event::{CallEvent, RequestContext, WireRequest}, - host::{HostOp, HostResult}, + host::{Demand, HostOp, HostResult}, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, route::Route, }; @@ -88,6 +88,21 @@ where _ => Err(MachineFault::Mismatch.into()), } } + + pub async fn open(&self, head: R::StreamHead) -> Result { + self.demand(HostOp::Open(head)).await + } + + pub async fn deliver(&self, chunk: R::Chunk) -> Result { + self.demand(HostOp::Deliver(chunk)).await + } + + async fn demand(&self, op: HostOp) -> Result { + match self.invoke(op).await? { + HostResult::Demand(demand) => Ok(demand), + _ => Err(MachineFault::Mismatch.into()), + } + } } enum Execution { diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index b95402b1a7a..22e2c398ff7 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,88 +1,54 @@ -use litellm_llms::custom_httpx::http_handler::http_request; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use std::time::Duration; -use super::{ - Error, client::http_client, common_utils::truncate_error_body, - prepare::prepare_provider_request, +use litellm_llms::{ + base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, + custom_httpx::{http_handler::http_request, transport::Error as TransportError}, }; -use crate::{constants::ANTHROPIC_MESSAGES_PROVIDER, messages::types::MessagesRequest}; +use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use serde_json::Value; -pub(super) async fn execute_messages_provider_call( - request: MessagesRequest<'_>, +use super::{Error, client::http_client, common_utils::truncate_error_body}; + +pub(super) fn network(error: reqwest::Error) -> Error { + Error::Transport(TransportError::Network(error.to_string())) +} + +pub(super) async fn send( + url: &str, + headers: &[(String, String)], + body: &Value, + timeout: Option, +) -> Result { + let builder = headers.iter().fold( + http_client().post(url).json(body), + |builder, (key, value)| builder.header(key, value), + ); + let builder = match timeout { + Some(duration) => builder.timeout(duration), + None => builder, + }; + http_request(builder).await.map_err(network) +} + +pub(super) async fn provider_error(response: reqwest::Response) -> Error { + let status = response.status().as_u16(); + match response.text().await { + Ok(text) => Error::Transport(TransportError::Http { + status, + body: truncate_error_body(&text), + }), + Err(error) => network(error), + } +} + +pub(super) fn decode_response( + config: &dyn BaseAnthropicMessagesConfig, + model: &str, + text: &str, ) -> Result { - let request = prepare_provider_request(request)?; - let mut request_builder = http_client().post(&request.url).json(&request.body); - for (key, value) in &request.upstream_headers { - request_builder = request_builder.header(key, value); - } - if let Some(duration) = request.timeout { - request_builder = request_builder.timeout(duration); - } - - let response = http_request(request_builder).await.map_err(|err| { - Error::Transport(litellm_llms::custom_httpx::transport::Error::Network( - err.to_string(), - )) - })?; - - let status = response.status(); - let text = response.text().await.map_err(|err| { - Error::Transport(litellm_llms::custom_httpx::transport::Error::Network( - err.to_string(), - )) - })?; - - if !status.is_success() { - return Err(Error::Transport( - litellm_llms::custom_httpx::transport::Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }, - )); - } - - let response = serde_json::from_str(&text) + let response = serde_json::from_str(text) .map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?; - request - .config - .transform_anthropic_messages_response(&request.model, response) + config + .transform_anthropic_messages_response(model, response) .map_err(Error::from) } - -pub(super) async fn execute_messages_provider_stream( - request: MessagesRequest<'_>, -) -> Result { - let request = prepare_provider_request(request)?; - if request.provider != ANTHROPIC_MESSAGES_PROVIDER { - return Err(Error::Unsupported("streaming messages for this provider")); - } - - let mut request_builder = http_client().post(&request.url).json(&request.body); - for (key, value) in &request.upstream_headers { - request_builder = request_builder.header(key, value); - } - if let Some(duration) = request.timeout { - request_builder = request_builder.timeout(duration); - } - - let response = http_request(request_builder).await.map_err(|err| { - Error::Transport(litellm_llms::custom_httpx::transport::Error::Network( - err.to_string(), - )) - })?; - let status = response.status(); - if !status.is_success() { - let text = response.text().await.map_err(|err| { - Error::Transport(litellm_llms::custom_httpx::transport::Error::Network( - err.to_string(), - )) - })?; - return Err(Error::Transport( - litellm_llms::custom_httpx::transport::Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }, - )); - } - Ok(response) -} diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index c3d7bea48ff..e36c6668efe 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,11 +1,8 @@ //! The Anthropic Messages call, the Rust equivalent of Python's //! `litellm.messages()`. //! -//! [`messages`] is the top-level entrypoint: give it a model, a body, and -//! credentials, and it resolves the provider, transforms the request, calls the -//! provider, and returns a typed non-streaming response. [`messages_stream`] -//! is the streaming variant; it hands the raw upstream response back so a host -//! can splice the event stream to its own caller. +//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs +//! it in process for a caller that already holds the request and wants the message. mod error; pub mod types; @@ -14,17 +11,34 @@ mod client; mod common_utils; mod handler; mod prepare; -use handler::{execute_messages_provider_call, execute_messages_provider_stream}; +pub mod route; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}; +use serde_json::Value; use crate::messages::types::MessagesRequest; pub async fn messages(request: MessagesRequest<'_>) -> Result { - execute_messages_provider_call(request).await -} - -pub async fn messages_stream(request: MessagesRequest<'_>) -> Result { - execute_messages_provider_stream(request).await + let Value::Object(body) = request.body else { + return Err(Error::InvalidRequest( + "messages body must be an object".into(), + )); + }; + let call = MessagesCall { + model: request.model.into(), + body, + api_key: request.api_key.map(Into::into), + api_base: request.api_base.map(Into::into), + custom_llm_provider: request.custom_llm_provider.map(Into::into), + extra_headers: request.extra_headers, + timeout: request.timeout, + }; + match litellm_callbacks::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? { + MessagesOutput::Message(message) => Ok(message), + MessagesOutput::Streamed => Err(Error::Unsupported( + "streamed responses need a streaming host", + )), + } } #[cfg(test)] diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index 8b676803871..850f9108869 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -2,6 +2,7 @@ use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_l use litellm_llms::base_llm::anthropic_messages::transformation::{ BaseAnthropicMessagesConfig, MessagesAuthStrategy, }; +use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; use serde_json::{Map, Value}; use super::{ @@ -37,10 +38,14 @@ pub(super) fn prepare_provider_request( let headers = validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?; - let typed_request = serde_json::from_value(request.body).map_err(|err| { - Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) + let typed_request: AnthropicMessagesRequest = + serde_json::from_value(request.body).map_err(|err| { + Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) + })?; + let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest { + model: model.clone(), + ..typed_request })?; - let transformed = config.transform_anthropic_messages_request(typed_request)?; let body = serde_json::to_value(transformed).map_err(|err| { Error::InvalidRequest(format!( "failed to serialize Anthropic messages request: {err}" diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs new file mode 100644 index 00000000000..a680b5ce1ec --- /dev/null +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -0,0 +1,197 @@ +use std::{sync::Mutex, time::Duration}; + +use bytes::Bytes; +use litellm_auth::SecretValue; +use litellm_callbacks::{ + event::{CallEvent, RawResponse, RequestContext, WireRequest}, + host::{Demand, Host}, + route::Route, +}; +use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; +use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use serde_json::{Map, Value}; + +use super::{ + Error, + common_utils::messages_provider_config, + handler::{decode_response, network, provider_error, send}, + prepare::prepare_provider_request, + types::MessagesRequest, +}; +use crate::{ + constants::ANTHROPIC_MESSAGES_PROVIDER, + machine::{HostChannel, MachineFault, RouteMachine}, +}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum MessagesOp { + ProjectRequest, +} + +pub enum MessagesOpResult { + Request(Box), +} + +/// The caller's request as the host projects it. +pub struct MessagesCall { + pub model: String, + pub body: Map, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub timeout: Option, +} + +impl MessagesCall { + fn streams(&self) -> bool { + self.body.get("stream").and_then(Value::as_bool) == Some(true) + } +} + +pub enum MessagesOutput { + Message(AnthropicMessagesResponse), + /// Every chunk already reached the host through `Deliver`. + Streamed, +} + +pub struct Messages; + +impl Route for Messages { + type Response = MessagesOutput; + type Error = Error; + type Op = MessagesOp; + type OpResult = MessagesOpResult; + type Chunk = Bytes; + type StreamHead = (); +} + +impl From for Error { + fn from(fault: MachineFault) -> Self { + Self::InvalidRequest(match fault { + MachineFault::Abandoned => "messages host driver was abandoned".into(), + MachineFault::Protocol(message) => format!("messages {message}"), + MachineFault::Mismatch => "invalid messages host operation result".into(), + }) + } +} + +pub type MessagesHost = HostChannel; +pub type MessagesMachine = RouteMachine; + +/// Whether this route serves the request, decided before any callback runs so a host +/// can still run its own path. +pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool { + let provider = get_custom_llm_provider(model, custom_llm_provider) + .map(|resolved| resolved.custom_llm_provider) + .or(custom_llm_provider); + match provider { + Some(ANTHROPIC_MESSAGES_PROVIDER) => true, + Some(provider) => !stream && messages_provider_config(provider).is_some(), + None => false, + } +} + +/// The in-process host for a request already in hand. It answers projection once and +/// observes nothing. +pub struct LocalMessagesHost { + call: Mutex>, +} + +impl LocalMessagesHost { + pub fn new(call: MessagesCall) -> Self { + Self { + call: Mutex::new(Some(call)), + } + } +} + +impl Host for LocalMessagesHost { + async fn route(&self, op: MessagesOp) -> Result { + match op { + MessagesOp::ProjectRequest => self + .call + .lock() + .unwrap_or_else(|error| error.into_inner()) + .take() + .map(|call| MessagesOpResult::Request(Box::new(call))) + .ok_or_else(|| { + Error::InvalidRequest("messages request was already projected".into()) + }), + } + } +} + +pub fn messages_machine() -> MessagesMachine { + RouteMachine::new(|host| Box::pin(execute(host))) +} + +async fn execute(host: MessagesHost) -> Result { + let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; + let stream = call.streams(); + let request = prepare_provider_request(MessagesRequest { + model: &call.model, + body: Value::Object(call.body.clone()), + api_key: call.api_key.as_deref(), + api_base: call.api_base.as_deref(), + custom_llm_provider: call.custom_llm_provider.as_deref(), + extra_headers: call.extra_headers.clone(), + timeout: call.timeout, + })?; + if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER { + return Err(Error::Unsupported("streaming messages for this provider")); + } + let context = RequestContext { + model: request.model.clone(), + custom_llm_provider: request.provider.clone(), + optional_params: Value::Object( + call.body + .iter() + .filter(|(name, _)| !matches!(name.as_str(), "model" | "messages")) + .map(|(name, value)| (name.clone(), value.clone())) + .collect(), + ), + secret_fields: Vec::new(), + api_key: call.api_key.clone().map(SecretValue::new), + }; + let wire = host + .before_send( + WireRequest { + url: request.url, + headers: request.upstream_headers, + body: request.body, + }, + context, + ) + .await?; + let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?; + if !response.status().is_success() { + return Err(provider_error(response).await); + } + if stream { + return relay(&host, response).await; + } + let text = response.text().await.map_err(network)?; + host.emit(CallEvent::ResponseReceived { + raw: RawResponse { body: text.clone() }, + }) + .await?; + decode_response(request.config, &request.model, &text).map(MessagesOutput::Message) +} + +/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading +/// ends the upstream read, and the call completes with what it delivered. +async fn relay( + host: &MessagesHost, + mut response: reqwest::Response, +) -> Result { + if host.open(()).await? == Demand::Detached { + return Ok(MessagesOutput::Streamed); + } + while let Some(chunk) = response.chunk().await.map_err(network)? { + if host.deliver(chunk).await? == Demand::Detached { + break; + } + } + Ok(MessagesOutput::Streamed) +} diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index ac4237651da..e6ee45c64a8 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -39,6 +39,8 @@ impl Route for Ocr { type Error = Error; type Op = OcrOp; type OpResult = OcrOpResult; + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } impl TokenRoute for Ocr { diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index ca2e14a7f0d..779d2037bb3 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -198,6 +198,8 @@ fn event_name(event: &CallEvent) -> &'static str { CallEvent::ResponseReceived { .. } => "response", CallEvent::Succeeded { .. } => "success", CallEvent::Failed { .. } => "failure", + CallEvent::Opened => "opened", + CallEvent::Delivered => "delivered", } } diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 4aa7a2163ca..795977438fd 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -23,6 +23,7 @@ pub enum LifecycleStep { pub enum PublicValue<'a> { Response(&'a Py), Error(&'a PyErr), + Chunk(&'a Py), } /// One consumer of a call's lifecycle on the Python side. The driver calls the steps in @@ -111,6 +112,13 @@ pub trait RouteHost: Send + Sync { response: ::Response, ) -> PyResult>; + /// One streamed chunk as the caller receives it. + fn chunk( + &mut self, + py: Python<'_>, + chunk: ::Chunk, + ) -> PyResult>; + fn classify( &self, py: Python<'_>, diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 7895d803cb5..c59ba1925dc 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -3,7 +3,7 @@ use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; -use litellm_callbacks::host::{HostOp, HostResult, HostStep}; +use litellm_callbacks::host::{Demand, HostOp, HostResult, HostStep}; use litellm_callbacks::machine::{HostFailure, Machine, MachineStep}; use litellm_callbacks::route::Route; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; @@ -38,6 +38,7 @@ struct MachineState { enum Stage { Begin, Call, + Streaming, AfterSuccess, Succeeded(Py), Failed(Py), @@ -56,6 +57,8 @@ enum Expect { enum Pending { Native, Adapter(Expect), + /// The stream handed to the caller waits for its next read or its close. + Consumer, } enum Next { @@ -121,7 +124,14 @@ where } match driver.resume(None)? { ExecutionStep::Return(value) => Ok(value), - ExecutionStep::Await(_) => Err(PyRuntimeError::new_err("sync call suspended")), + ExecutionStep::Open => py + .import("litellm.rust_bridge.lifecycle")? + .getattr("SyncStream")? + .call1((Py::new(py, Execution::suspended(driver))?,)) + .map(Bound::unbind), + ExecutionStep::Await(_) | ExecutionStep::Yield(_) => { + Err(PyRuntimeError::new_err("sync call suspended")) + } } } @@ -162,6 +172,14 @@ where self.run_steps(py, HostStep::Ready(result)) } (Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error), + (Some(Pending::Consumer), Some(read)) => { + let demand = if read.is_ok() { + Demand::More + } else { + Demand::Detached + }; + self.resume_machine(py, Some(Ok(HostResult::Demand(demand)))) + } (Some(Pending::Adapter(expect)), Some(result)) => { match self.adapter.resume(py, result) { Ok(step) => self.on_adapter(py, step, expect), @@ -216,7 +234,7 @@ where fn adapter_failed(&mut self, py: Python<'_>, error: PyErr) -> PyResult { match self.stage { Stage::Begin | Stage::AfterSuccess => self.failure(py, error, FailureOrigin::Host), - Stage::Call => self.interrupt(py, error), + Stage::Call | Stage::Streaming => self.interrupt(py, error), Stage::Succeeded(_) | Stage::Failed(_) => Err(error), } } @@ -283,6 +301,8 @@ where Err(error) => Err(error), } } + HostOp::Open(_) => return self.opened(py).map(Next::Return), + HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return), HostOp::Emit(event) => match self.adapter.emit(py, &event, None) { Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), Ok(LifecycleStep::Await(awaitable)) => { @@ -299,6 +319,40 @@ where } } + fn opened(&mut self, py: Python<'_>) -> PyResult { + self.stage = Stage::Streaming; + match self.adapter.emit(py, &CallEvent::Opened, None) { + Ok(LifecycleStep::Done) => { + self.pending = Some(Pending::Consumer); + Ok(ExecutionStep::Open) + } + Ok(_) => Err(missing_state()), + Err(error) => self.interrupt(py, error), + } + } + + fn delivered( + &mut self, + py: Python<'_>, + chunk: as Route>::Chunk, + ) -> PyResult { + let chunk = match self.route.chunk(py, chunk) { + Ok(chunk) => chunk, + Err(error) => return self.interrupt(py, error), + }; + let observed = + self.adapter + .emit(py, &CallEvent::Delivered, Some(PublicValue::Chunk(&chunk))); + match observed { + Ok(LifecycleStep::Done) => { + self.pending = Some(Pending::Consumer); + Ok(ExecutionStep::Yield(chunk)) + } + Ok(_) => Err(missing_state()), + Err(error) => self.interrupt(py, error), + } + } + fn interrupt(&mut self, py: Python<'_>, error: PyErr) -> PyResult { let cancelled = is_cancellation(py, &error); let native = H::host_error(&error); @@ -369,6 +423,9 @@ where Ok(public) => public, Err(error) => return self.failure(py, error, FailureOrigin::Call), }; + if let Stage::Streaming = self.stage { + return self.succeeded(py, public); + } self.stage = Stage::AfterSuccess; match self.adapter.after_success(py, public, self.timing()) { Ok(step) => self.on_adapter(py, step, Expect::Response), @@ -536,6 +593,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri type Error = Error; type Op = &'static str; type OpResult = String; + type Chunk = std::convert::Infallible; + type StreamHead = std::convert::Infallible; } /// Yields the scripted ops in order, then completes or fails as scripted. @@ -574,6 +633,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri HostResult::Route(value) => value, HostResult::BeforeSend(wire) => wire.url, HostResult::Emitted => "emitted".into(), + HostResult::Demand(demand) => format!("{demand:?}"), }); } if !self.ops.is_empty() { @@ -648,6 +708,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } + fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { + match chunk {} + } + fn complete(&mut self, py: Python<'_>, response: String) -> PyResult> { self.log.push("complete"); Ok(pyo3::types::PyString::new(py, &response) @@ -1138,6 +1202,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri ) .into()) } + fn chunk( + &mut self, + _: Python<'_>, + chunk: std::convert::Infallible, + ) -> PyResult> { + match chunk {} + } fn complete(&mut self, _: Python<'_>, _: String) -> PyResult> { Err(missing_state()) } diff --git a/litellm-rust/crates/host-python/src/handle.rs b/litellm-rust/crates/host-python/src/handle.rs index d8cd6c92130..10abbadbda5 100644 --- a/litellm-rust/crates/host-python/src/handle.rs +++ b/litellm-rust/crates/host-python/src/handle.rs @@ -8,6 +8,10 @@ use pyo3::prelude::*; pub enum ExecutionStep { Return(Py), Await(Py), + /// The call streams: the caller gets a stream over this execution, which stays + /// suspended until the stream asks for a chunk. + Open, + Yield(Py), } pub trait ExecutionBody: Send + Sync { @@ -34,6 +38,13 @@ impl Execution { } } + /// An execution already started elsewhere and now waiting for its next input. + pub fn suspended(body: impl ExecutionBody + 'static) -> Self { + Self { + state: ExecutionState::Suspended(Box::new(body)), + } + } + fn advance( slf: &Bound<'_, Self>, py: Python<'_>, @@ -64,6 +75,8 @@ impl Execution { let step = body.resume(result)?; let (tag, value, suspended) = match step { ExecutionStep::Await(value) => ("Await", value, true), + ExecutionStep::Open => ("Open", py.None(), true), + ExecutionStep::Yield(value) => ("Yield", value, true), ExecutionStep::Return(value) => ("Complete", value, false), }; let step = py diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index ea4077b102f..2aba51cc4ff 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -18,10 +18,6 @@ pub(crate) struct RouteOptions { pub(crate) timeout: Option, } -pub(crate) fn body_argument(value: &Bound<'_, PyAny>) -> PyResult> { - required_object("body", from_py_argument(value)?) -} - pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult> { match from_py_argument(value)? { Value::Array(values) => Ok(values), @@ -192,18 +188,6 @@ mod tests { json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) ); - let body = py - .eval( - c"{'model': 'claude', 'metadata': {'user': '1'}}", - None, - None, - ) - .unwrap(); - assert_eq!( - Value::Object(body_argument(&body).unwrap()), - json!({"model": "claude", "metadata": {"user": "1"}}) - ); - let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap(); assert_eq!( optional_params_argument(¶ms).unwrap(), diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs deleted file mode 100644 index daec931c92e..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ /dev/null @@ -1,88 +0,0 @@ -use litellm_core::messages::{Error, messages as run_messages, types::MessagesRequest}; -use litellm_host_python::{run_async, run_sync}; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; -use pyo3::prelude::*; -use serde_json::{Map, Value}; - -use crate::{ - errors::messages_error_to_pyerr, - marshal::{RouteOptions, body_argument, extra_headers_argument, optional_timeout}, -}; - -async fn execute( - body: Map, - options: RouteOptions, -) -> Result { - let RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout, - } = options; - run_messages(MessagesRequest { - model: &model, - body: Value::Object(body), - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) - .await -} - -#[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] -pub(crate) fn messages( - py: Python<'_>, - model: String, - #[pyo3(from_py_with = body_argument)] body: Map, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - run_sync(py, execute(body, options), messages_error_to_pyerr) -} - -#[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] -#[expect( - clippy::too_many_arguments, - reason = "one parameter per Python keyword" -)] -pub(crate) fn amessages<'py>( - py: Python<'py>, - model: String, - #[pyo3(from_py_with = body_argument)] body: Map, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option>, - timeout_seconds: Option, -) -> PyResult> { - let options = RouteOptions { - model, - api_key, - api_base, - custom_llm_provider, - extra_headers, - timeout: optional_timeout(timeout_seconds), - }; - run_async(py, execute(body, options), messages_error_to_pyerr) -} diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs new file mode 100644 index 00000000000..0c87aea1ca2 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -0,0 +1,186 @@ +use bytes::Bytes; +use litellm_core::messages::{ + Error, + route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput}, +}; +use litellm_host_python::{HostOpError, RouteHost, from_py, lookup, to_py}; +use litellm_llms::custom_httpx::transport::Error as TransportError; +use pyo3::{ + exceptions::{PyException, PyValueError}, + gc::{PyTraverseError, PyVisit}, + prelude::*, + types::{PyBytes, PyDict}, +}; +use serde_json::{Map, Value}; + +use crate::{ + errors::{RustUpstreamError, messages_error_to_pyerr}, + marshal::{optional_timeout, python_timeout_seconds}, +}; + +/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`, +/// as `AnthropicMessagesRequestOptionalParams` declares them. +const BODY_FIELDS: [&str; 20] = [ + "max_tokens", + "metadata", + "stop_sequences", + "stream", + "system", + "temperature", + "thinking", + "tool_choice", + "tools", + "top_k", + "inference_geo", + "top_p", + "mcp_servers", + "context_management", + "container", + "output_format", + "speed", + "output_config", + "cache_control", + "reasoning_effort", +]; + +/// The Python side of the Messages route: projects the prepared arguments and builds the +/// public response, chunks and exceptions. +pub(super) struct MessagesRouteHost { + request: Py, +} + +impl MessagesRouteHost { + pub(super) fn new(request: Py) -> Self { + Self { request } + } + + fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult { + let request = self.request.bind(py); + let argument = |name: &str| -> PyResult>> { + Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none())) + }; + let string = |name: &str| -> PyResult> { + argument(name)?.map(|value| value.extract()).transpose() + }; + let model = string("model")?.ok_or_else(|| PyValueError::new_err("model is required"))?; + let messages = + argument("messages")?.ok_or_else(|| PyValueError::new_err("messages is required"))?; + let fields = BODY_FIELDS + .iter() + .filter_map(|name| match argument(name) { + Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))), + Ok(None) => None, + Err(error) => Some(Err(error)), + }) + .collect::>>()?; + let body = [ + ("model".to_string(), Value::String(model.clone())), + ("messages".to_string(), from_py(&messages)?), + ] + .into_iter() + .chain(fields) + .collect::>(); + let timeout = argument("timeout")? + .map(|value| python_timeout_seconds(py, value.unbind())) + .transpose()? + .flatten(); + Ok(MessagesCall { + model, + body, + api_key: string("api_key")?, + api_base: string("api_base")?, + custom_llm_provider: string("custom_llm_provider")?, + extra_headers: argument("extra_headers")? + .map(|value| from_py(&value)) + .transpose()?, + timeout: optional_timeout(timeout), + }) + } + + fn provider(&self, py: Python<'_>) -> String { + self.request + .bind(py) + .getattr("custom_llm_provider") + .and_then(|value| value.extract::>()) + .ok() + .flatten() + .unwrap_or_else(|| "anthropic".into()) + } + + fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr { + if !error.is_instance_of::(py) { + return error; + } + let mapped = py + .import("litellm.rust_bridge.messages.route_host") + .and_then(|module| module.getattr("map_failure")) + .and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py)))) + .and_then(|mapped| { + mapped + .extract::>() + .map_err(PyErr::from) + }); + match mapped { + Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()), + Err(_) => error, + } + } +} + +impl RouteHost for MessagesRouteHost { + type Route = Messages; + type Failure = PyErr; + + fn invoke( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + op: MessagesOp, + ) -> Result> { + match op { + MessagesOp::ProjectRequest => self + .project(py, arguments) + .map(|call| MessagesOpResult::Request(Box::new(call))) + .map_err(|error| HostOpError::Python(self.map_failure(py, error))), + } + } + + fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { + match response { + MessagesOutput::Message(message) => py + .import("litellm.rust_bridge.messages.route_host")? + .getattr("response")? + .call1((to_py(py, &message)?,)) + .map(Bound::unbind), + MessagesOutput::Streamed => Ok(py.None()), + } + } + + fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult> { + Ok(PyBytes::new(py, &chunk).into_any().unbind()) + } + + fn classify(&self, py: Python<'_>, error: Error) -> PyResult { + let native = match error { + Error::Transport(TransportError::Http { status, body }) => { + let error = RustUpstreamError::new_err((status, body)); + error + .value(py) + .setattr("headers", Vec::<(String, String)>::new())?; + error + } + other => messages_error_to_pyerr(other), + }; + Ok(self.map_failure(py, native)) + } + + fn host_error(error: &PyErr) -> Error { + Error::InvalidRequest(error.to_string()) + } + + fn close(&mut self, _: Python<'_>) {} + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.request) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs new file mode 100644 index 00000000000..b4259aca5ba --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -0,0 +1,64 @@ +mod host; + +use host::MessagesRouteHost; +use litellm_callbacks_legacy::{LegacySurface, PublicCall, run_legacy_call}; +use litellm_core::messages::route::{messages_machine, supports}; +use pyo3::{ + prelude::*, + types::{PyDict, PyTuple}, +}; + +use crate::errors::RustBridgeDeclined; + +const SURFACE: LegacySurface = LegacySurface { + call_type: "anthropic_messages", + input_description: "Messages", +}; + +fn run_messages( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, + asynchronous: bool, +) -> PyResult> { + let model: String = request.getattr("model")?.extract()?; + let provider: Option = request.getattr("custom_llm_provider")?.extract()?; + let stream = request + .getattr("stream")? + .extract::>()? + .unwrap_or(false); + if !supports(&model, provider.as_deref(), stream) { + return Err(RustBridgeDeclined::new_err( + "the Rust Messages route does not serve this provider", + )); + } + run_legacy_call( + py, + SURFACE, + PublicCall::capture(&request, &args, &kwargs)?, + messages_machine(), + MessagesRouteHost::new(request.unbind()), + asynchronous, + ) +} + +#[pyfunction] +pub(crate) fn messages( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, +) -> PyResult> { + run_messages(py, request, args, kwargs, false) +} + +#[pyfunction] +pub(crate) fn amessages( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, +) -> PyResult> { + run_messages(py, request, args, kwargs, true) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index f59e32a28e2..2d6b849a6b1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -22,11 +22,6 @@ mod tests { "atranscription", "(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", ), - ( - "messages", - "amessages", - "(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)", - ), ( "chat_completions", "achat_completions", @@ -113,25 +108,6 @@ value = Broken() ); assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string()); - let invalid_body = PyList::empty(py); - let sync_messages_error = module - .getattr("messages") - .and_then(|function| function.call1(("model", &invalid_body))) - .expect_err("sync Messages should reject a non-dict body"); - let async_messages_error = module - .getattr("amessages") - .and_then(|function| function.call1(("model", &invalid_body))) - .expect_err("async Messages should reject a non-dict body"); - - assert_eq!( - sync_messages_error.to_string(), - "ValueError: body must be a dict" - ); - assert_eq!( - async_messages_error.to_string(), - sync_messages_error.to_string() - ); - let invalid_headers = PyList::empty(py); let kwargs = PyDict::new(py); kwargs @@ -193,13 +169,6 @@ value = Broken() headers_kwargs .set_item("extra_headers", &invalid) .expect("kwargs should accept extra_headers"); - let invalid_body = PyList::empty(py); - let error = module - .getattr("messages") - .and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs))) - .expect_err("body should be validated before headers"); - assert_eq!(error.to_string(), "ValueError: body must be a dict"); - let invalid_payload = PyModule::new(py, "invalid_payload").expect("invalid payload should be created"); let error = module diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 77212e3d38e..19211671fd7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -125,6 +125,10 @@ impl RouteHost for OcrRouteHost { .map(Bound::unbind) } + fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult> { + match chunk {} + } + fn classify(&self, py: Python<'_>, error: Error) -> PyResult { Ok(self.map_failure(py, ocr_error_to_pyerr(error))) } diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 488e278cca7..9f959c056de 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -1,9 +1,11 @@ from asyncio import Future -from collections.abc import Coroutine, Mapping, Sequence +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence from typing import Never, final from litellm.llms.base_llm.ocr.transformation import OCRResponse +from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest +from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse class RustBridgeDeclined(Exception): ... class RustUpstreamError(Exception): ... @@ -39,23 +41,15 @@ def atranscription( timeout_seconds: float | None = None, ) -> Future[dict[str, object]]: ... def messages( - model: str, - body: Mapping[str, object], - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, -) -> dict[str, object]: ... + request: LiteLLMMessagesRequest, + args: tuple[object, ...], + kwargs: dict[str, object], +) -> AnthropicMessagesResponse | Iterator[bytes]: ... def amessages( - model: str, - body: Mapping[str, object], - api_key: str | None = None, - api_base: str | None = None, - custom_llm_provider: str | None = None, - extra_headers: Mapping[str, object] | None = None, - timeout_seconds: float | None = None, -) -> Future[dict[str, object]]: ... + request: LiteLLMMessagesRequest, + args: tuple[object, ...], + kwargs: dict[str, object], +) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ... def chat_completions_decline( model: str, messages: Sequence[object], diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 9efbbfa2e9e..d843a874fe3 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -59,6 +59,7 @@ Rules: TypeAlias = tuple[Rule, ...] RULES: Final[Rules] = ( Rule(Route.OCR, Rollout.RUST_OPT_OUT), + Rule(Route.MESSAGES, Rollout.RUST_OPT_IN), Rule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})), ) diff --git a/litellm/rust_bridge/failures.py b/litellm/rust_bridge/failures.py index b714341fe43..80805b7ff69 100644 --- a/litellm/rust_bridge/failures.py +++ b/litellm/rust_bridge/failures.py @@ -5,8 +5,37 @@ from __future__ import annotations from collections.abc import Mapping from typing import Final, Protocol, cast # noqa: TID251 # adapts the public exception mapper +import httpx +import openai +from pydantic import TypeAdapter, ValidationError + import litellm +_UPSTREAM_ARGS: Final = TypeAdapter(tuple[int, str]) +_UPSTREAM_HEADERS: Final = TypeAdapter(list[tuple[str, str]]) + + +class UpstreamFailure(Exception): + def __init__(self, response: httpx.Response, cause: Exception) -> None: + super().__init__(str(cause)) + self.message: Final = str(cause) + self.response: Final = response + self.status_code: Final = response.status_code + self.__cause__ = cause + + +def _upstream_failure(error: Exception, api_base: str | None) -> Exception: + try: + status, body = _UPSTREAM_ARGS.validate_python(error.args) + headers: Final = _UPSTREAM_HEADERS.validate_python(getattr(error, "headers", None)) + except ValidationError: + return error + http_request: Final = httpx.Request("POST", api_base or "https://docs.litellm.ai/docs") + return UpstreamFailure( + httpx.Response(status, content=body.encode(), headers=headers, request=http_request), + error, + ) + class ExceptionMapper(Protocol): def __call__( @@ -35,3 +64,17 @@ def map_failure(error: Exception, model: str, request_provider: str, kwargs: Map except Exception as public_error: public_error.__context__ = error return public_error + + +def map_native_failure( + error: Exception, model: str, request_provider: str, kwargs: Mapping[str, object], api_base: str | None = None +) -> Exception: + """`map_failure`, reading a native `(status, body)` provider failure as the HTTP response it was.""" + original: Final = _upstream_failure(error, api_base) + public_error: Final = map_failure(original, model, request_provider, kwargs) + if isinstance(original, UpstreamFailure) and public_error.__context__ is original: + public_error.__context__ = error + if isinstance(public_error, openai.APIStatusError): + public_error.response = original.response + public_error.status_code = original.status_code + return public_error diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index 406dc55cfea..fd0fac1799a 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -6,6 +6,7 @@ registries it fans out to. It expires with that contract. from __future__ import annotations +import asyncio import contextvars import datetime import traceback @@ -160,6 +161,21 @@ class LoggingWorker(Protocol): def ensure_initialized_and_enqueue(self, async_coroutine: Coroutine[object, object, None]) -> None: ... +class StreamingLogBuilder(Protocol): + def __call__( + self, + *, + litellm_logging_obj: Logging, + passthrough_success_handler_obj: object, + url_route: str, + request_body: dict[str, object], + endpoint_type: object, + start_time: datetime.datetime, + raw_bytes: list[bytes], + end_time: datetime.datetime, + ) -> Coroutine[object, object, None]: ... + + class DeploymentHook(Protocol): def __call__(self, kwargs: dict[str, object], call_type: str) -> Awaitable[object]: ... @@ -306,3 +322,67 @@ def after_deployment_failure(kwargs: dict[str, object], error: Exception, call_t DeploymentFailureHook, utils.async_post_call_failure_deployment_hook ) return hook(kwargs, error, call_type) + + +def stream_opened(logger: Logging) -> None: + logger.stream = True + logger.model_call_details["stream"] = True + + +def stream_success( + logger: Logging, + request_body: dict[str, object], + chunks: list[bytes], + start: datetime.datetime, + end: datetime.datetime, + first_chunk: datetime.datetime | None, +) -> None: + from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, + ) + from litellm.proxy.pass_through_endpoints.streaming_handler import PassThroughStreamingHandler + from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType + + if first_chunk is not None: + logger.completion_start_time = first_chunk + logger.model_call_details["completion_start_time"] = first_chunk + build: Final = cast( # cast-ok: bounded adapter for the untyped pass-through logging builder + StreamingLogBuilder, + PassThroughStreamingHandler._route_streaming_logging_to_handler, # pyright: ignore[reportPrivateUsage] # the Messages stream iterator bills through the same builder + ) + coroutine: Final = build( + litellm_logging_obj=logger, + passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, + url_route="/v1/messages", + request_body=request_body, + endpoint_type=EndpointType.ANTHROPIC, + start_time=start, + raw_bytes=chunks, + end_time=end, + ) + if getattr(logger, "_on_deferred_stream_complete", None) is not None: + logger._deferred_stream_complete_args = (coroutine,) # pyright: ignore[reportAttributeAccessIssue] # the proxy's deferred stream release reads this slot + return + try: + asyncio.get_running_loop() + except RuntimeError: + from litellm.litellm_core_utils.litellm_logging import executor + + executor.submit(contextvars.copy_context().run, asyncio.run, coroutine) + return + enqueue_logging(coroutine) + + +def stream_failure( + logger: Logging, request_body: dict[str, object], chunks: list[bytes], error: Exception +) -> Coroutine[object, object, None]: + from litellm.proxy.pass_through_endpoints.streaming_handler import PassThroughStreamingHandler + from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType + + return PassThroughStreamingHandler.schedule_stream_failure_logging( + litellm_logging_obj=logger, + endpoint_type=EndpointType.ANTHROPIC, + request_body=request_body, + raw_bytes=chunks, + exception=error, + ) diff --git a/litellm/rust_bridge/lifecycle.py b/litellm/rust_bridge/lifecycle.py index d903021b6f3..4096d386964 100644 --- a/litellm/rust_bridge/lifecycle.py +++ b/litellm/rust_bridge/lifecycle.py @@ -1,8 +1,8 @@ from __future__ import annotations -from collections.abc import Awaitable +from collections.abc import AsyncIterator, Awaitable, Iterator from dataclasses import dataclass -from typing import Protocol +from typing import Final, Protocol @dataclass(frozen=True, slots=True) @@ -15,28 +15,133 @@ class Complete: value: object +@dataclass(frozen=True, slots=True) +class Open: + value: None + + +@dataclass(frozen=True, slots=True) +class Yield: + value: object + + +Settled = Complete | Open | Yield +Step = Await | Settled + + class Execution(Protocol): - def start(self) -> Await | Complete: ... + def start(self) -> Step: ... - def resume_value(self, value: object) -> Await | Complete: ... + def resume_value(self, value: object) -> Step: ... - def resume_error(self, error: BaseException) -> Await | Complete: ... + def resume_error(self, error: BaseException) -> Step: ... def close(self) -> None: ... +class StreamClosed(Exception): + """Tells a streaming execution that its caller stopped reading.""" + + +async def _settle(execution: Execution, step: Step) -> Settled: + while isinstance(step, Await): + try: + value = await step.awaitable # rebind-ok: each selected await produces the next protocol input + except GeneratorExit: + raise + except BaseException as error: + step = execution.resume_error(error) # rebind-ok: advance the execution protocol + else: + step = execution.resume_value(value) # rebind-ok: advance the execution protocol + return step + + +def _settled(step: Step) -> Settled: + if isinstance(step, Await): + raise RuntimeError("sync call suspended") + return step + + async def drive(execution: Execution) -> object: + handed_off = False # rebind-ok: set once the execution belongs to the returned stream try: - step = execution.start() # rebind-ok: the execution protocol advances after each selected await - while isinstance(step, Await): - try: - value = await step.awaitable # rebind-ok: each selected await produces the next protocol input - except GeneratorExit: - raise - except BaseException as error: - step = execution.resume_error(error) # rebind-ok: advance the execution protocol - else: - step = execution.resume_value(value) # rebind-ok: advance the execution protocol + step: Final = await _settle(execution, execution.start()) + if isinstance(step, Open): + handed_off = True + return Stream(execution) return step.value finally: - execution.close() + if not handed_off: + execution.close() + + +class Stream(AsyncIterator[object]): + """A streamed native call: each read resumes the execution until its next chunk.""" + + def __init__(self, execution: Execution) -> None: + self._execution: Final = execution + self._done = False + + def __aiter__(self) -> Stream: + return self + + async def __anext__(self) -> object: + if self._done: + raise StopAsyncIteration + try: + step: Final = await _settle(self._execution, self._execution.resume_value(None)) + except BaseException: + self._finish() + raise + if isinstance(step, Yield): + return step.value + self._finish() + raise StopAsyncIteration + + async def aclose(self) -> None: + if self._done: + return + try: + await _settle(self._execution, self._execution.resume_error(StreamClosed())) + finally: + self._finish() + + def _finish(self) -> None: + self._done = True + self._execution.close() + + +class SyncStream(Iterator[object]): + """The sync form of `Stream`; its execution never suspends on an awaitable.""" + + def __init__(self, execution: Execution) -> None: + self._execution: Final = execution + self._done = False + + def __iter__(self) -> SyncStream: + return self + + def __next__(self) -> object: + if self._done: + raise StopIteration + try: + step: Final = _settled(self._execution.resume_value(None)) + except BaseException: + self._finish() + raise + if isinstance(step, Yield): + return step.value + self._finish() + raise StopIteration + + def close(self) -> None: + if self._done: + return + try: + _settled(self._execution.resume_error(StreamClosed())) + finally: + self._finish() + + def _finish(self) -> None: + self._done = True + self._execution.close() diff --git a/litellm/rust_bridge/messages/entrypoints.py b/litellm/rust_bridge/messages/entrypoints.py index 46565bfd46a..d25c906c4c1 100644 --- a/litellm/rust_bridge/messages/entrypoints.py +++ b/litellm/rust_bridge/messages/entrypoints.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping, Sequence +from collections.abc import AsyncIterator, Awaitable, Iterator, Mapping, Sequence from dataclasses import dataclass from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables @@ -26,7 +26,7 @@ class NativeMessages(Protocol): request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object], - ) -> AnthropicMessagesResponse: ... + ) -> AnthropicMessagesResponse | Iterator[bytes]: ... class NativeAmessages(Protocol): @@ -35,7 +35,7 @@ class NativeAmessages(Protocol): request: LiteLLMMessagesRequest, args: tuple[object, ...], kwargs: Mapping[str, object], - ) -> Awaitable[AnthropicMessagesResponse]: ... + ) -> Awaitable[AnthropicMessagesResponse | AsyncIterator[bytes]]: ... def _messages_binding(value: object) -> NativeMessages | None: @@ -50,5 +50,5 @@ def _amessages_binding(value: object) -> NativeAmessages | None: return cast("NativeAmessages", value) # cast-ok: callable validated at the native binding boundary -NATIVE_MESSAGES: Final = NativeBinding("anthropic_messages_handler", validate=_messages_binding) -NATIVE_AMESSAGES: Final = NativeBinding("anthropic_messages", validate=_amessages_binding) +NATIVE_MESSAGES: Final = NativeBinding("messages", validate=_messages_binding) +NATIVE_AMESSAGES: Final = NativeBinding("amessages", validate=_amessages_binding) diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index 1aff6c7f75d..beef0f81eca 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -20,4 +20,4 @@ def arguments(request: LiteLLMMessagesRequest) -> Mapping[str, object]: def map_failure(error: Exception, request: LiteLLMMessagesRequest, request_provider: str) -> Exception: - return failures.map_failure(error, request.model, request_provider, arguments(request)) + return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base) diff --git a/litellm/rust_bridge/ocr/route_host.py b/litellm/rust_bridge/ocr/route_host.py index 277fdceb734..bfbd5c11d4e 100644 --- a/litellm/rust_bridge/ocr/route_host.py +++ b/litellm/rust_bridge/ocr/route_host.py @@ -4,40 +4,17 @@ from collections.abc import Mapping from types import MappingProxyType from typing import Final -import httpx -import openai -from pydantic import TypeAdapter, ValidationError +from pydantic import TypeAdapter import litellm from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge import failures +from litellm.rust_bridge.failures import UpstreamFailure from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest +__all__ = ("UpstreamFailure", "arguments", "map_failure", "response") + _RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object]) -_UPSTREAM_ARGS: Final = TypeAdapter(tuple[int, str]) -_UPSTREAM_HEADERS: Final = TypeAdapter(list[tuple[str, str]]) - - -class UpstreamFailure(Exception): - def __init__(self, response: httpx.Response, cause: Exception) -> None: - super().__init__(str(cause)) - self.message: Final = str(cause) - self.response: Final = response - self.status_code: Final = response.status_code - self.__cause__ = cause - - -def _upstream_failure(error: Exception, request: LiteLLMOcrRequest) -> Exception: - try: - status, body = _UPSTREAM_ARGS.validate_python(error.args) - headers: Final = _UPSTREAM_HEADERS.validate_python(getattr(error, "headers", None)) - except ValidationError: - return error - http_request: Final = httpx.Request("POST", request.api_base or "https://docs.litellm.ai/docs") - return UpstreamFailure( - httpx.Response(status, content=body.encode(), headers=headers, request=http_request), - error, - ) def response(value: Mapping[str, object]) -> OCRResponse: @@ -61,11 +38,4 @@ def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: model=request.model.removeprefix(f"{request_provider}/"), llm_provider=request_provider, ) - original: Final = _upstream_failure(error, request) - public_error: Final = failures.map_failure(original, request.model, request_provider, arguments(request)) - if isinstance(original, UpstreamFailure) and public_error.__context__ is original: - public_error.__context__ = error - if isinstance(public_error, openai.APIStatusError): - public_error.response = original.response - public_error.status_code = original.status_code - return public_error + return failures.map_native_failure(error, request.model, request_provider, arguments(request), request.api_base) diff --git a/tests/test_litellm/rust_bridge/native_route_wheel_test.py b/tests/test_litellm/rust_bridge/native_route_wheel_test.py index 4fa4c0b95ec..0b442f1f269 100644 --- a/tests/test_litellm/rust_bridge/native_route_wheel_test.py +++ b/tests/test_litellm/rust_bridge/native_route_wheel_test.py @@ -73,7 +73,7 @@ def assert_native_request( headers: HTTPMessage, body: object, ) -> None: - if route not in {"transcription", "messages", "chat_completions"}: + if route not in {"transcription", "chat_completions"}: raise AssertionError(f"unexpected route marker: {route!r}") if outcome not in {"success", "429", "hang"}: raise AssertionError(f"unexpected outcome marker: {outcome!r}") @@ -89,10 +89,6 @@ def assert_native_request( assert path == "/v1/messages" assert headers.get("x-api-key") == "sk-native" assert body["model"] == "claude-sonnet-4-5" - if route == "messages": - assert body["max_tokens"] == 16 - assert body["messages"][0]["content"] == "hello-from-messages" - return assert body["max_tokens"] == 17 assert body["messages"][0]["content"] == [{"type": "text", "text": "hello-from-chat"}] @@ -132,17 +128,6 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: "language": "en", }, } - if route == "messages": - return common | { - "model": "claude-sonnet-4-5", - "body": { - "model": "claude-sonnet-4-5", - "max_tokens": 16, - "messages": [{"role": "user", "content": "hello-from-messages"}], - }, - "api_key": "sk-native", - "custom_llm_provider": "anthropic", - } if route == "chat_completions": return common | { "model": "anthropic/claude-sonnet-4-5", @@ -165,8 +150,6 @@ def assert_success(route: str, response: object) -> None: def success_value(route: str, response: dict[object, object]) -> object: if route == "transcription": return response["text"] - if route == "messages": - return response["content"][0]["text"] return response["choices"][0]["message"]["content"] @@ -181,7 +164,7 @@ def assert_rate_limit(native: object, route: str, error: BaseException) -> None: def exercise_sync(native: object, api_base: str) -> None: - for route in ("transcription", "messages", "chat_completions"): + for route in ("transcription", "chat_completions"): function: Final = getattr(native, route) assert_success(route, function(**route_kwargs(route, api_base, "success"))) try: @@ -193,7 +176,7 @@ def exercise_sync(native: object, api_base: str) -> None: async def exercise_async(native: object, api_base: str) -> None: - for route in ("transcription", "messages", "chat_completions"): + for route in ("transcription", "chat_completions"): function: Final = getattr(native, f"a{route}") assert_success(route, await function(**route_kwargs(route, api_base, "success"))) try: @@ -206,11 +189,11 @@ async def exercise_async(native: object, api_base: str) -> None: async def exercise_async_concurrency(native: object, api_base: str) -> None: responses: Final = await asyncio.wait_for( - asyncio.gather(*(native.amessages(**route_kwargs("messages", api_base, "success")) for _ in range(32))), + asyncio.gather(*(native.achat_completions(**route_kwargs("chat_completions", api_base, "success")) for _ in range(32))), timeout=15, ) for response in responses: - assert_success("messages", response) + assert_success("chat_completions", response) def exercise_routes(native_path: Path, api_base: str) -> object: @@ -223,8 +206,8 @@ def exercise_routes(native_path: Path, api_base: str) -> object: def exercise_signal(native: object, api_base: str) -> int: try: - native.messages( - **route_kwargs("messages", api_base, "hang"), + native.chat_completions( + **route_kwargs("chat_completions", api_base, "hang"), ) except KeyboardInterrupt: sys.stdout.write("KeyboardInterrupt\n") diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index 72390b79141..b882a1bb8c2 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -43,8 +43,8 @@ def test_binding_validates_native_attribute( ROUTE_BINDINGS: Final = ( ("completion", chat_completions.NATIVE_COMPLETION), ("acompletion", chat_completions.NATIVE_ACOMPLETION), - ("anthropic_messages_handler", messages.NATIVE_MESSAGES), - ("anthropic_messages", messages.NATIVE_AMESSAGES), + ("messages", messages.NATIVE_MESSAGES), + ("amessages", messages.NATIVE_AMESSAGES), ("responses", responses.NATIVE_RESPONSES), ("aresponses", responses.NATIVE_ARESPONSES), ("ocr", ocr.NATIVE_OCR), diff --git a/tests/test_litellm/rust_bridge/test_catalog.py b/tests/test_litellm/rust_bridge/test_catalog.py index 2c737b0160e..e9fdbf859f4 100644 --- a/tests/test_litellm/rust_bridge/test_catalog.py +++ b/tests/test_litellm/rust_bridge/test_catalog.py @@ -40,6 +40,10 @@ def test_shipped_decisions( enabled: Final = environment == "1" if environment is not None else process is not False assert catalog.rollout(context) is Rollout.RUST_OPT_OUT assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if enabled else Decision.PYTHON) + elif route is Route.MESSAGES: + enabled: Final = environment == "1" if environment is not None else process is True + assert catalog.rollout(context) is Rollout.RUST_OPT_IN + assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if enabled else Decision.PYTHON) elif route is Route.TRANSCRIPTION and provider == "bedrock": assert catalog.rollout(context) is Rollout.RUST_REQUIRED assert catalog.decision(context) is Decision.RUST_REQUIRED diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index ade0ae549fb..fa6c0b30413 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -157,7 +157,6 @@ def test_context_outside_rule_stays_on_python() -> None: ( Context(Route.CHAT_COMPLETIONS, provider="anthropic"), Context(Route.CHAT_COMPLETIONS, provider="bedrock"), - Context(Route.MESSAGES, provider="anthropic"), Context(Route.RESPONSES, provider="openai"), Context(Route.TRANSCRIPTION, provider="openai"), ), diff --git a/tests/test_litellm_rust/messages/__init__.py b/tests/test_litellm_rust/messages/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm_rust/messages/test_callbacks.py b/tests/test_litellm_rust/messages/test_callbacks.py new file mode 100644 index 00000000000..b55bc47d640 --- /dev/null +++ b/tests/test_litellm_rust/messages/test_callbacks.py @@ -0,0 +1,175 @@ +from collections.abc import AsyncIterator, Iterator +from typing import Final + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger +from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec +from tests.test_litellm_rust.support.requests import ( + MESSAGES, + MESSAGES_EVENTS, + MESSAGES_MODEL, + MESSAGES_RESPONSE, + request_body, +) + +pytestmark = pytest.mark.requires_rust_extension + +STREAM: Final = ResponseSpec(body=None, events=MESSAGES_EVENTS) + + +@pytest.fixture +def messages_server(recording_server: RecordingServer) -> RecordingServer: + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + return recording_server + + +def arguments(server: RecordingServer, **kwargs: object) -> dict[str, object]: + return { + "model": MESSAGES_MODEL, + "messages": [dict(message) for message in MESSAGES], + "max_tokens": 64, + "api_key": "test-key", + "api_base": server.base_url, + **kwargs, + } + + +def assert_served_natively(server: RecordingServer) -> None: + assert len(server.requests) == 1 + assert not server.requests[0].headers.get("user-agent", "").startswith("python-httpx") + + +@pytest.mark.asyncio +async def test_native_messages_callbacks_see_the_provider_request_and_the_public_response( + messages_server: RecordingServer, +) -> None: + recorder: Final = RecordingLogger() + + response: Final = await litellm.anthropic.messages.acreate( + **arguments(messages_server, callbacks=[recorder], litellm_call_id="messages-success") + ) + + assert_served_natively(messages_server) + assert response["content"] == MESSAGES_RESPONSE["content"] + sent: Final = messages_server.requests[0] + assert sent.path == "/v1/messages" + assert sent.body == {"model": "claude-sonnet-5", "messages": list(MESSAGES), "max_tokens": 64, "stream": False} + pre_call: Final = recorder.wait_for("log_pre_api_call") + assert request_body(pre_call[0].kwargs) == sent.body + success: Final = await recorder.wait_for_async("async_log_success_event") + assert len(success) == 1 + assert success[0].call_type == "anthropic_messages" + assert success[0].kwargs["litellm_call_id"] == "messages-success" + assert success[0].response.choices[0].message.content == "Hello from native Messages" + + +@pytest.mark.asyncio +async def test_native_messages_pre_call_body_edit_reaches_the_provider(messages_server: RecordingServer) -> None: + class Edit(CustomLogger): + def log_pre_api_call(self, model, messages, kwargs): + request_body(kwargs)["temperature"] = 0.25 + + await litellm.anthropic.messages.acreate(**arguments(messages_server, callbacks=[Edit()])) + + assert messages_server.requests[0].body["temperature"] == 0.25 + + +@pytest.mark.asyncio +async def test_native_messages_provider_error_reaches_caller_and_failure_callbacks_as_one_public_error( + messages_server: RecordingServer, +) -> None: + messages_server.enqueue( + ResponseSpec(body={"type": "error", "error": {"type": "invalid_request_error", "message": "bad"}}, status=400) + ) + observed: Final = [] + + class Observe(CustomLogger): + def log_failure_event(self, kwargs, response_obj, start_time, end_time): + observed.append(("sync", kwargs["exception"])) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + observed.append(("async", kwargs["exception"])) + + with pytest.raises(litellm.BadRequestError) as raised: + await litellm.anthropic.messages.acreate(**arguments(messages_server, callbacks=[Observe()])) + + assert_served_natively(messages_server) + assert [phase for phase, _ in observed] == ["sync", "async"] + assert all(error is raised.value for _, error in observed) + + +def sse_payload() -> bytes: + return b"".join(STREAM.payloads()) + + +@pytest.mark.asyncio +async def test_native_messages_stream_relays_provider_events_and_logs_success_once_after_the_last_chunk( + messages_server: RecordingServer, +) -> None: + messages_server.enqueue(STREAM) + recorder: Final = RecordingLogger() + + stream: Final = await litellm.anthropic.messages.acreate( + **arguments(messages_server, stream=True, callbacks=[recorder]) + ) + assert isinstance(stream, AsyncIterator) + first: Final = await anext(stream) + await drain_logging() + assert "async_log_success_event" not in recorder.names + rest: Final = [chunk async for chunk in stream] + + assert first + b"".join(rest) == sse_payload() + assert_served_natively(messages_server) + assert messages_server.requests[0].body["stream"] is True + success: Final = await recorder.wait_for_async("async_log_success_event") + assert len(success) == 1 + assert success[0].kwargs["stream"] is True + assert success[0].kwargs["completion_start_time"] is not None + assert "log_failure_event" not in recorder.names + + +@pytest.mark.asyncio +async def test_native_messages_stream_closed_early_logs_success_once_for_what_was_delivered( + messages_server: RecordingServer, +) -> None: + messages_server.enqueue(STREAM) + recorder: Final = RecordingLogger() + + stream: Final = await litellm.anthropic.messages.acreate( + **arguments(messages_server, stream=True, callbacks=[recorder]) + ) + assert isinstance(stream, AsyncIterator) + await anext(stream) + await stream.aclose() + + success: Final = await recorder.wait_for_async("async_log_success_event") + assert len(success) == 1 + with pytest.raises(StopAsyncIteration): + await anext(stream) + + +def test_native_sync_messages_stream_relays_provider_events_and_logs_success_once( + messages_server: RecordingServer, +) -> None: + messages_server.enqueue(STREAM) + recorder: Final = RecordingLogger() + + stream: Final = litellm.anthropic.messages.create(**arguments(messages_server, stream=True, callbacks=[recorder])) + assert isinstance(stream, Iterator) + + assert b"".join(stream) == sse_payload() + assert_served_natively(messages_server) + assert len(recorder.wait_for("async_log_success_event")) == 1 + + +def test_native_sync_messages_returns_the_provider_message(messages_server: RecordingServer) -> None: + recorder: Final = RecordingLogger() + + response: Final = litellm.anthropic.messages.create(**arguments(messages_server, callbacks=[recorder])) + + assert_served_natively(messages_server) + assert response["content"] == MESSAGES_RESPONSE["content"] + assert len(recorder.wait_for("log_success_event")) == 1 diff --git a/tests/test_litellm_rust/support/recording_server.py b/tests/test_litellm_rust/support/recording_server.py index 228ed2cc454..3eea47751d3 100644 --- a/tests/test_litellm_rust/support/recording_server.py +++ b/tests/test_litellm_rust/support/recording_server.py @@ -25,6 +25,12 @@ class ResponseSpec: status: int = 200 headers: dict[str, str] = field(default_factory=dict) delay: float = 0 + events: tuple[tuple[str, object], ...] = () + + def payloads(self) -> tuple[bytes, ...]: + if not self.events: + return (json.dumps(self.body).encode(),) + return tuple(f"event: {event}\ndata: {json.dumps(data)}\n\n".encode() for event, data in self.events) @dataclass @@ -73,15 +79,17 @@ def recording_service() -> Iterator[RecordingServer]: response: Final = responses.pop(0) if responses else copy.deepcopy(recording_server.default_response) if response.delay: time.sleep(response.delay) - payload: Final = json.dumps(response.body).encode() + payloads: Final = response.payloads() self.send_response(response.status) - self.send_header("Content-Type", "application/json") - self.send_header("Content-Length", str(len(payload))) + self.send_header("Content-Type", "text/event-stream" if response.events else "application/json") + self.send_header("Content-Length", str(sum(len(payload) for payload in payloads))) for name, value in response.headers.items(): self.send_header(name, value) self.end_headers() try: - self.wfile.write(payload) + for payload in payloads: + self.wfile.write(payload) + self.wfile.flush() except (BrokenPipeError, ConnectionResetError): pass diff --git a/tests/test_litellm_rust/support/requests.py b/tests/test_litellm_rust/support/requests.py index b60cf5eac02..c9cf81b83ca 100644 --- a/tests/test_litellm_rust/support/requests.py +++ b/tests/test_litellm_rust/support/requests.py @@ -12,6 +12,41 @@ OCR_RESPONSE: Final = { "usage_info": {"pages_processed": 1, "doc_size_bytes": 3}, } +MESSAGES_MODEL: Final = "anthropic/claude-sonnet-5" +MESSAGES: Final = ({"role": "user", "content": "Hello"},) +MESSAGES_RESPONSE: Final = { + "id": "msg_native", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5", + "content": [{"type": "text", "text": "Hello from native Messages"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 4}, +} +MESSAGES_EVENTS: Final = ( + ("message_start", {"type": "message_start", "message": {**MESSAGES_RESPONSE, "content": [], "stop_reason": None}}), + ("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + ( + "content_block_delta", + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": "Hello from native Messages"}, + }, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 4}, + }, + ), + ("message_stop", {"type": "message_stop"}), +) + def ocr_arguments(server: RecordingServer, **kwargs: object) -> dict[str, object]: return { From 3a5b7c12ef119474a596cd58f59807cce5854fb5 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:35:20 -0700 Subject: [PATCH 240/253] refactor(rust): separate machine events from the Python lifecycle's events --- .../crates/callbacks-legacy/src/adapter.rs | 57 +++++++-------- .../tests/deployment_hooks.rs | 11 ++- .../crates/callbacks-legacy/tests/payload.rs | 8 +-- .../crates/callbacks-legacy/tests/terminal.rs | 14 ++-- litellm-rust/crates/callbacks/src/event.rs | 16 +++-- litellm-rust/crates/callbacks/src/host.rs | 4 +- litellm-rust/crates/callbacks/src/run.rs | 5 +- litellm-rust/crates/core/src/machine/mod.rs | 4 +- .../crates/core/src/messages/route.rs | 4 +- litellm-rust/crates/core/src/ocr/handler.rs | 4 +- .../tests/azure_document_intelligence_ocr.rs | 8 +-- litellm-rust/crates/core/tests/ocr.rs | 9 ++- litellm-rust/crates/core/tests/reducto_ocr.rs | 8 +-- .../crates/host-python/src/adapter.rs | 33 ++++++--- litellm-rust/crates/host-python/src/driver.rs | 69 ++++++++++--------- litellm-rust/crates/host-python/src/lib.rs | 2 +- 16 files changed, 142 insertions(+), 114 deletions(-) diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index 9d28db92add..a67da2188fb 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -3,10 +3,10 @@ //! `@client` path makes them. use litellm_callbacks::event::{ - CallEvent, FailureOrigin, RequestContext, Timing, WireRequest, epoch_seconds, + FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds, }; use litellm_host_python::{ - LifecycleStep, PublicValue, PythonLifecycle, from_py, missing_state, to_py, + LifecycleEvent, LifecycleStep, PythonLifecycle, from_py, missing_state, to_py, }; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -346,28 +346,11 @@ impl PythonLifecycle for LegacyLogging { fn emit( &mut self, py: Python<'_>, - event: &CallEvent, - public: Option>, + event: LifecycleEvent<'_>, ) -> PyResult { - match (event, public) { - (CallEvent::Started { .. }, _) => Ok(LifecycleStep::Done), - (CallEvent::Opened, _) => { - Streaming::Opened.call(py, (self.logger()?.object(py),))?; - self.stream = Some(DeliveredStream { - chunks: PyList::empty(py).unbind(), - first_chunk: None, - }); - Ok(LifecycleStep::Done) - } - (CallEvent::Delivered, Some(PublicValue::Chunk(chunk))) => { - let stream = self.stream.as_mut().ok_or_else(missing_state)?; - if stream.first_chunk.is_none() { - stream.first_chunk = Some(datetime(py, epoch_seconds())?); - } - stream.chunks.bind(py).append(chunk)?; - Ok(LifecycleStep::Done) - } - (CallEvent::ResponseReceived { raw }, _) => { + match event { + LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done), + LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { let api_key = self .context .as_ref() @@ -382,7 +365,7 @@ impl PythonLifecycle for LegacyLogging { )?; Ok(LifecycleStep::Done) } - (CallEvent::Succeeded { timing }, Some(PublicValue::Response(response))) => { + LifecycleEvent::Succeeded { timing, response } => { self.end = Some(datetime(py, timing.end_time)?); self.response = Some(response.clone_ref(py)); match &self.stream { @@ -391,13 +374,17 @@ impl PythonLifecycle for LegacyLogging { } Ok(LifecycleStep::Done) } - (CallEvent::Failed { timing, origin }, Some(PublicValue::Error(error))) => { + LifecycleEvent::Failed { + timing, + origin, + error, + } => { self.end = Some(datetime(py, timing.end_time)?); self.error = Some(error.clone_ref(py).into_value(py)); if self.stream.is_some() { return self.stream_failure(py); } - if *origin == FailureOrigin::Call + if origin == FailureOrigin::Call && self.logger.is_some() && self.runs_deployment_hooks() { @@ -412,10 +399,26 @@ impl PythonLifecycle for LegacyLogging { } self.dispatch_failure(py) } - _ => Err(missing_state()), } } + fn opened(&mut self, py: Python<'_>) -> PyResult<()> { + Streaming::Opened.call(py, (self.logger()?.object(py),))?; + self.stream = Some(DeliveredStream { + chunks: PyList::empty(py).unbind(), + first_chunk: None, + }); + Ok(()) + } + + fn delivered(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()> { + let stream = self.stream.as_mut().ok_or_else(missing_state)?; + if stream.first_chunk.is_none() { + stream.first_chunk = Some(datetime(py, epoch_seconds())?); + } + stream.chunks.bind(py).append(chunk) + } + fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult { match self.pending.take().ok_or_else(missing_state)? { Pending::DeploymentPreCall => { diff --git a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs index 7bad09c7890..f4602739548 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs @@ -1,7 +1,7 @@ use std::ffi::CStr; -use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing}; -use litellm_host_python::{LifecycleStep, PublicValue, PythonLifecycle}; +use litellm_callbacks::event::{FailureOrigin, Timing}; +use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; use pyo3::types::PyDict; @@ -217,13 +217,12 @@ fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelle .resume(py, Ok(local(&locals, "kwargs").unbind())) .unwrap(); let failure = PyErr::from_value(local(&locals, "failure")); - let failed = CallEvent::Failed { + let failed = LifecycleEvent::Failed { timing: TIMING, origin: FailureOrigin::Call, + error: &failure, }; - let step = logging - .emit(py, &failed, Some(PublicValue::Error(&failure))) - .unwrap(); + let step = logging.emit(py, failed).unwrap(); assert!(awaits_deployment_hook(&step)); let hook_result = if cancelled { Err(CancelledError::new_err("cancelled")) diff --git a/litellm-rust/crates/callbacks-legacy/tests/payload.rs b/litellm-rust/crates/callbacks-legacy/tests/payload.rs index 43128bc38ea..67ad4ab8a2a 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/payload.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/payload.rs @@ -1,8 +1,8 @@ use std::ffi::CStr; use litellm_auth::SecretValue; -use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest}; -use litellm_host_python::{LifecycleStep, PythonLifecycle, to_py}; +use litellm_callbacks::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; +use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py}; use proptest::prelude::*; use pyo3::prelude::*; use rstest::rstest; @@ -94,13 +94,13 @@ fn before_send_bound( body, }; let step = logging.before_send(py, Box::new(wire), &context).unwrap(); - let raw = CallEvent::ResponseReceived { + let raw = MachineEvent::ResponseReceived { raw: RawResponse { body: "raw response".into(), }, }; assert!(matches!( - logging.emit(py, &raw, None).unwrap(), + logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(), LifecycleStep::Done )); run(py, &locals, c"check()"); diff --git a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs index 3094d7b88d2..5688f70f387 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs @@ -1,7 +1,7 @@ use std::ffi::CStr; -use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing}; -use litellm_host_python::{LifecycleStep, PublicValue, PythonLifecycle}; +use litellm_callbacks::event::{FailureOrigin, Timing}; +use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; use pyo3::exceptions::PyRuntimeError; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; @@ -33,8 +33,10 @@ fn succeed( logging .emit( py, - &CallEvent::Succeeded { timing: TIMING }, - Some(PublicValue::Response(&response)), + LifecycleEvent::Succeeded { + timing: TIMING, + response: &response, + }, ) .unwrap() } @@ -44,11 +46,11 @@ fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) logging .emit( py, - &CallEvent::Failed { + LifecycleEvent::Failed { timing: TIMING, origin: FailureOrigin::Host, + error: &failure, }, - Some(PublicValue::Error(&failure)), ) .unwrap() } diff --git a/litellm-rust/crates/callbacks/src/event.rs b/litellm-rust/crates/callbacks/src/event.rs index d19e973b812..182dab657d3 100644 --- a/litellm-rust/crates/callbacks/src/event.rs +++ b/litellm-rust/crates/callbacks/src/event.rs @@ -52,18 +52,20 @@ pub enum FailureOrigin { Host, } +/// What a machine reports while it runs. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum MachineEvent { + ResponseReceived { raw: RawResponse }, +} + +/// What an in-process host observes: the machine's own events between the driver's +/// start and terminal ones. #[derive(Clone, Debug, PartialEq)] pub enum CallEvent { Started { start_time: f64, }, - ResponseReceived { - raw: RawResponse, - }, - /// The call streams and its stream was handed to the caller. - Opened, - /// One chunk of an open stream reached the caller. - Delivered, + Machine(MachineEvent), Succeeded { timing: Timing, }, diff --git a/litellm-rust/crates/callbacks/src/host.rs b/litellm-rust/crates/callbacks/src/host.rs index eef3e1da8d5..aba35185a18 100644 --- a/litellm-rust/crates/callbacks/src/host.rs +++ b/litellm-rust/crates/callbacks/src/host.rs @@ -1,6 +1,6 @@ use std::future::Future; -use crate::event::{CallEvent, RequestContext, WireRequest}; +use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest}; use crate::route::Route; /// One suspension point of a native call, performed by the host. @@ -10,7 +10,7 @@ pub enum HostOp { wire: Box, context: Box, }, - Emit(CallEvent), + Emit(MachineEvent), /// The response streams: the host hands the caller a stream and answers once the /// caller asks for the first chunk or goes away. Open(R::StreamHead), diff --git a/litellm-rust/crates/callbacks/src/run.rs b/litellm-rust/crates/callbacks/src/run.rs index 705c5504c01..6a0c08fba68 100644 --- a/litellm-rust/crates/callbacks/src/run.rs +++ b/litellm-rust/crates/callbacks/src/run.rs @@ -25,7 +25,10 @@ where .before_send(*wire, &context) .await .map(|wire| HostResult::BeforeSend(Box::new(wire))), - HostOp::Emit(event) => host.emit(&event).await.map(|()| HostResult::Emitted), + HostOp::Emit(event) => host + .emit(&CallEvent::Machine(event)) + .await + .map(|()| HostResult::Emitted), HostOp::Open(head) => host.open(head).await.map(HostResult::Demand), HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand), }; diff --git a/litellm-rust/crates/core/src/machine/mod.rs b/litellm-rust/crates/core/src/machine/mod.rs index 929f0a423c4..d6db488159e 100644 --- a/litellm-rust/crates/core/src/machine/mod.rs +++ b/litellm-rust/crates/core/src/machine/mod.rs @@ -8,7 +8,7 @@ use std::{future::Future, pin::Pin}; pub use auth::{HostTokenProvider, TokenRoute}; use litellm_callbacks::{ - event::{CallEvent, RequestContext, WireRequest}, + event::{MachineEvent, RequestContext, WireRequest}, host::{Demand, HostOp, HostResult}, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, route::Route, @@ -82,7 +82,7 @@ where } } - pub async fn emit(&self, event: CallEvent) -> Result<(), R::Error> { + pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> { match self.invoke(HostOp::Emit(event)).await? { HostResult::Emitted => Ok(()), _ => Err(MachineFault::Mismatch.into()), diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index a680b5ce1ec..2fd2f79907a 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -3,7 +3,7 @@ use std::{sync::Mutex, time::Duration}; use bytes::Bytes; use litellm_auth::SecretValue; use litellm_callbacks::{ - event::{CallEvent, RawResponse, RequestContext, WireRequest}, + event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, route::Route, }; @@ -172,7 +172,7 @@ async fn execute(host: MessagesHost) -> Result { return relay(&host, response).await; } let text = response.text().await.map_err(network)?; - host.emit(CallEvent::ResponseReceived { + host.emit(MachineEvent::ResponseReceived { raw: RawResponse { body: text.clone() }, }) .await?; diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index a6af190eb91..aff1eded2cc 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,6 +1,6 @@ use futures_util::future::BoxFuture; use litellm_auth::SecretValue; -use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest}; +use litellm_callbacks::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; use litellm_llms::{ base_llm::ocr::{ error::Error, @@ -71,7 +71,7 @@ impl CallHooks for OcrCallHooks { } fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> { - Box::pin(self.host.emit(CallEvent::ResponseReceived { + Box::pin(self.host.emit(MachineEvent::ResponseReceived { raw: RawResponse { body: String::from_utf8_lossy(body).into_owned(), }, diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index 1fc4d6c2b9e..62544931dfc 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::CallEvent; +use litellm_callbacks::event::{CallEvent, MachineEvent}; use litellm_llms::base_llm::ocr::error::Error; use rstest::rstest; use serde_json::{Value, json}; @@ -263,7 +263,7 @@ async fn accepted_response_emits_response_received_before_polling() { json!({}), )) .with_observer(move |event| { - let CallEvent::ResponseReceived { raw } = event else { + let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event else { return; }; match request_count.lock().unwrap().len() { @@ -466,7 +466,7 @@ async fn model_id_is_encoded_and_dot_segments_are_rejected() { mod transformation { use std::sync::{Arc, Mutex}; - use litellm_callbacks::event::CallEvent; + use litellm_callbacks::event::{CallEvent, MachineEvent}; use litellm_llms::base_llm::ocr::transformation::OcrDocument; use serde_json::{Value, json}; @@ -646,7 +646,7 @@ mod transformation { json!({}), )) .with_observer(move |event| { - if let CallEvent::ResponseReceived { raw } = event { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { observed .lock() .unwrap() diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 779d2037bb3..a6001951361 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -1,7 +1,7 @@ use std::sync::{Arc, Mutex}; use litellm_callbacks::{ - event::{CallEvent, WireRequest}, + event::{CallEvent, MachineEvent, WireRequest}, host::{Host, HostOp, HostResult}, machine::{HostFailure, Machine, MachineStep}, }; @@ -195,11 +195,9 @@ async fn facade_uses_the_injected_http_client() { fn event_name(event: &CallEvent) -> &'static str { match event { CallEvent::Started { .. } => "started", - CallEvent::ResponseReceived { .. } => "response", + CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response", CallEvent::Succeeded { .. } => "success", CallEvent::Failed { .. } => "failure", - CallEvent::Opened => "opened", - CallEvent::Delivered => "delivered", } } @@ -377,6 +375,7 @@ async fn drive_until( intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire))) } HostOp::Emit(event) => { + let event = CallEvent::Machine(event); ops.push(event_name(&event)); host.emit(&event) .await @@ -420,7 +419,7 @@ async fn invalid_provider_response_emits_response_received_before_normalization_ let observed = responses_received.clone(); let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))).with_observer( move |event| { - if let CallEvent::ResponseReceived { raw } = event { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { observed.lock().unwrap().push(raw.body.clone()); } }, diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs index 59891b16e90..8c037889a17 100644 --- a/litellm-rust/crates/core/tests/reducto_ocr.rs +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::{CallEvent, WireRequest}; +use litellm_callbacks::event::{CallEvent, MachineEvent, WireRequest}; use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument}; use rstest::rstest; use serde_json::{Value, json}; @@ -139,7 +139,7 @@ async fn response_received_stays_after_reducto_upload_and_parse() { let request_count = seen.clone(); let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))).with_observer( move |event| { - if let CallEvent::ResponseReceived { raw } = event { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { assert_eq!(request_count.lock().unwrap().len(), 2); assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#); } @@ -351,7 +351,7 @@ async fn guardrail_rewrites_document_before_upload() { } mod transformation { - use litellm_callbacks::event::{CallEvent, WireRequest}; + use litellm_callbacks::event::{CallEvent, MachineEvent, WireRequest}; use litellm_llms::{ base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext}, reducto::ocr::transformation::*, @@ -506,7 +506,7 @@ mod transformation { let request_count = seen.clone(); let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))) .with_observer(move |event| { - if let CallEvent::ResponseReceived { raw } = event { + if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event { assert_eq!(request_count.lock().unwrap().len(), 2); assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#); } diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 795977438fd..f946dbc763a 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::{CallEvent, RequestContext, Timing, WireRequest}; +use litellm_callbacks::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; use litellm_callbacks::route::Route; use pyo3::exceptions::PyRuntimeError; use pyo3::gc::{PyTraverseError, PyVisit}; @@ -19,11 +19,22 @@ pub enum LifecycleStep { Done, } -/// The host-typed value the driver attaches to a terminal event. -pub enum PublicValue<'a> { - Response(&'a Py), - Error(&'a PyErr), - Chunk(&'a Py), +/// What a lifecycle observes: the driver's start, the machine's own events, and one +/// terminal event carrying the public value the caller receives. +pub enum LifecycleEvent<'a> { + Started { + start_time: f64, + }, + Machine(&'a MachineEvent), + Succeeded { + timing: Timing, + response: &'a Py, + }, + Failed { + timing: Timing, + origin: FailureOrigin, + error: &'a PyErr, + }, } /// One consumer of a call's lifecycle on the Python side. The driver calls the steps in @@ -61,10 +72,16 @@ pub trait PythonLifecycle: Send + Sync { fn emit( &mut self, py: Python<'_>, - event: &CallEvent, - public: Option>, + event: LifecycleEvent<'_>, ) -> PyResult; + /// The call streams and its stream was handed to the caller. The caller is not + /// inside an await here, so this step and `delivered` cannot suspend. + fn opened(&mut self, py: Python<'_>) -> PyResult<()>; + + /// One chunk of an open stream is about to reach the caller. + fn delivered(&mut self, py: Python<'_>, chunk: &Py) -> PyResult<()>; + fn resume(&mut self, py: Python<'_>, result: PyResult>) -> PyResult; fn close(&mut self, py: Python<'_>); diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index c59ba1925dc..25f78d7013e 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; -use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing, epoch_seconds}; +use litellm_callbacks::event::{FailureOrigin, Timing, epoch_seconds}; use litellm_callbacks::host::{Demand, HostOp, HostResult, HostStep}; use litellm_callbacks::machine::{HostFailure, Machine, MachineStep}; use litellm_callbacks::route::Route; @@ -13,7 +13,7 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - HostOpError, LifecycleStep, PublicValue, PythonLifecycle, RouteHost, missing_state, + HostOpError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; @@ -159,10 +159,10 @@ where match (self.pending.take(), result) { (None, None) => { self.started_at = epoch_seconds(); - let started = CallEvent::Started { + let started = LifecycleEvent::Started { start_time: self.started_at, }; - match self.adapter.emit(py, &started, None) { + match self.adapter.emit(py, started) { Ok(step) => self.on_adapter(py, step, Expect::Started), Err(error) => self.adapter_failed(py, error), } @@ -303,7 +303,7 @@ where } HostOp::Open(_) => return self.opened(py).map(Next::Return), HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return), - HostOp::Emit(event) => match self.adapter.emit(py, &event, None) { + HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) { Ok(LifecycleStep::Done) => Ok(HostResult::Emitted), Ok(LifecycleStep::Await(awaitable)) => { self.pending = Some(Pending::Adapter(Expect::Emitted)); @@ -321,12 +321,11 @@ where fn opened(&mut self, py: Python<'_>) -> PyResult { self.stage = Stage::Streaming; - match self.adapter.emit(py, &CallEvent::Opened, None) { - Ok(LifecycleStep::Done) => { + match self.adapter.opened(py) { + Ok(()) => { self.pending = Some(Pending::Consumer); Ok(ExecutionStep::Open) } - Ok(_) => Err(missing_state()), Err(error) => self.interrupt(py, error), } } @@ -340,15 +339,11 @@ where Ok(chunk) => chunk, Err(error) => return self.interrupt(py, error), }; - let observed = - self.adapter - .emit(py, &CallEvent::Delivered, Some(PublicValue::Chunk(&chunk))); - match observed { - Ok(LifecycleStep::Done) => { + match self.adapter.delivered(py, &chunk) { + Ok(()) => { self.pending = Some(Pending::Consumer); Ok(ExecutionStep::Yield(chunk)) } - Ok(_) => Err(missing_state()), Err(error) => self.interrupt(py, error), } } @@ -461,12 +456,11 @@ where } fn succeeded(&mut self, py: Python<'_>, response: Py) -> PyResult { - let event = CallEvent::Succeeded { + let event = LifecycleEvent::Succeeded { timing: self.timing(), + response: &response, }; - let step = self - .adapter - .emit(py, &event, Some(PublicValue::Response(&response)))?; + let step = self.adapter.emit(py, event)?; self.stage = Stage::Succeeded(response); self.on_adapter(py, step, Expect::Terminal) } @@ -481,13 +475,12 @@ where if is_cancellation(py, &error) { return Err(error); } - let event = CallEvent::Failed { + let event = LifecycleEvent::Failed { timing: self.timing(), origin, + error: &error, }; - let step = self - .adapter - .emit(py, &event, Some(PublicValue::Error(&error)))?; + let step = self.adapter.emit(py, event)?; self.stage = Stage::Failed(error.into_value(py)); self.on_adapter(py, step, Expect::Terminal) } @@ -541,7 +534,7 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_callbacks::event::{RequestContext, WireRequest}; + use litellm_callbacks::event::{MachineEvent, RequestContext, WireRequest}; use litellm_callbacks::machine::{Interrupted, Step}; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; @@ -803,23 +796,33 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri fn emit( &mut self, py: Python<'_>, - event: &CallEvent, - public: Option>, + event: LifecycleEvent<'_>, ) -> PyResult { - self.log.push(match (event, public) { - (CallEvent::Started { .. }, None) => "started".into(), - (CallEvent::ResponseReceived { raw }, None) => format!("response:{}", raw.body), - (CallEvent::Succeeded { .. }, Some(PublicValue::Response(value))) => { - format!("succeeded:{}", value.bind(py)) + self.log.push(match event { + LifecycleEvent::Started { .. } => "started".into(), + LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { + format!("response:{}", raw.body) } - (CallEvent::Failed { origin, .. }, Some(PublicValue::Error(error))) => { + LifecycleEvent::Succeeded { response, .. } => { + format!("succeeded:{}", response.bind(py)) + } + LifecycleEvent::Failed { origin, error, .. } => { format!("failed:{origin:?}:{}", error.value(py)) } - _ => "unexpected".into(), }); Ok(LifecycleStep::Done) } + fn opened(&mut self, _: Python<'_>) -> PyResult<()> { + self.log.push("opened"); + Ok(()) + } + + fn delivered(&mut self, _: Python<'_>, _: &Py) -> PyResult<()> { + self.log.push("delivered"); + Ok(()) + } + fn resume(&mut self, _: Python<'_>, _: PyResult>) -> PyResult { Err(missing_state()) } @@ -899,7 +902,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri wire: Box::new(wire()), context: Box::new(context()), }, - HostOp::Emit(CallEvent::ResponseReceived { + HostOp::Emit(MachineEvent::ResponseReceived { raw: litellm_callbacks::event::RawResponse { body: "raw".into() }, }), ], diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 8889b1513db..3738a9b1c3c 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -13,7 +13,7 @@ mod handle; mod marshal; pub use adapter::{ - HostOpError, LifecycleStep, PublicValue, PythonLifecycle, RouteHost, missing_state, + HostOpError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; From c2679757b3ad93bd9f8def74021e78be42a025c7 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:36:42 -0700 Subject: [PATCH 241/253] refactor(rust): put stream billing on the legacy surface --- .../callbacks-legacy/python_contract.json | 3 ++ .../crates/callbacks-legacy/src/adapter.rs | 37 +++++++++++++++---- .../crates/callbacks-legacy/src/lib.rs | 2 +- .../crates/callbacks-legacy/tests/support.rs | 1 + .../crates/host-python/src/adapter.rs | 6 +-- litellm-rust/crates/host-python/src/driver.rs | 6 +-- .../python-bridge/src/routes/messages/mod.rs | 6 ++- .../python-bridge/src/routes/ocr/mod.rs | 1 + litellm/rust_bridge/legacy_callbacks.py | 14 +++++-- 9 files changed, 53 insertions(+), 23 deletions(-) diff --git a/litellm-rust/crates/callbacks-legacy/python_contract.json b/litellm-rust/crates/callbacks-legacy/python_contract.json index a09bdc711a3..8a7f3b98f47 100644 --- a/litellm-rust/crates/callbacks-legacy/python_contract.json +++ b/litellm-rust/crates/callbacks-legacy/python_contract.json @@ -100,6 +100,8 @@ ], "stream_success": [ "logger", + "url_route", + "endpoint_type", "request_body", "chunks", "start", @@ -108,6 +110,7 @@ ], "stream_failure": [ "logger", + "endpoint_type", "request_body", "chunks", "error" diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index a67da2188fb..4e9167100e0 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -30,6 +30,16 @@ pub struct LegacySurface { pub call_type: &'static str, /// What `Logging.pre_call` is told the input was. pub input_description: &'static str, + /// How a streamed response is billed; `None` for a route that never streams. + pub stream: Option, +} + +/// The pass-through billing a streamed response goes through once its chunks are in. +#[derive(Clone, Copy, Debug)] +pub struct PassThroughStream { + pub url_route: &'static str, + /// A value of Python's `EndpointType`. + pub endpoint_type: &'static str, } /// What the Messages stream iterator keeps for its end-of-stream billing. @@ -177,10 +187,13 @@ impl LegacyLogging { fn stream_success(&self, py: Python<'_>, stream: &DeliveredStream) -> PyResult<()> { let logger = self.logger()?; + let billing = self.surface.stream.ok_or_else(missing_state)?; let billed = Streaming::Success.call( py, ( logger.object(py), + billing.url_route, + billing.endpoint_type, &self.body, &stream.chunks, &self.start, @@ -201,14 +214,25 @@ impl LegacyLogging { /// partial usage. The sync path has no loop to schedule that on, so it falls back to /// the plain failure handler. fn stream_failure(&mut self, py: Python<'_>) -> PyResult { - let (Some(logger), Some(error), Some(stream)) = (&self.logger, &self.error, &self.stream) + let (Some(logger), Some(error), Some(stream), Some(billing)) = + (&self.logger, &self.error, &self.stream, self.surface.stream) else { return Ok(LifecycleStep::Done); }; if !self.asynchronous { return self.dispatch_failure(py); } - match Streaming::Failure.call(py, (logger.object(py), &self.body, &stream.chunks, error)) { + let scheduled = Streaming::Failure.call( + py, + ( + logger.object(py), + billing.endpoint_type, + &self.body, + &stream.chunks, + error, + ), + ); + match scheduled { Ok(awaitable) => { self.pending = Some(Pending::AsyncFailure); Ok(LifecycleStep::Await(awaitable.unbind())) @@ -343,11 +367,7 @@ impl PythonLifecycle for LegacyLogging { self.finalize(py) } - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult { + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { match event { LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done), LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { @@ -403,6 +423,9 @@ impl PythonLifecycle for LegacyLogging { } fn opened(&mut self, py: Python<'_>) -> PyResult<()> { + if self.surface.stream.is_none() { + return Err(missing_state()); + } Streaming::Opened.call(py, (self.logger()?.object(py),))?; self.stream = Some(DeliveredStream { chunks: PyList::empty(py).unbind(), diff --git a/litellm-rust/crates/callbacks-legacy/src/lib.rs b/litellm-rust/crates/callbacks-legacy/src/lib.rs index 42ffd545e2b..eaa1a8b714e 100644 --- a/litellm-rust/crates/callbacks-legacy/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy/src/lib.rs @@ -21,7 +21,7 @@ mod preparation; mod test_support; pub(crate) use adapter::LegacyLogging; -pub use adapter::LegacySurface; +pub use adapter::{LegacySurface, PassThroughStream}; pub use call::{PublicCall, run_legacy_call}; pub(crate) use callbacks::{LegacyCallbacks, is_internal_call}; pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup}; diff --git a/litellm-rust/crates/callbacks-legacy/tests/support.rs b/litellm-rust/crates/callbacks-legacy/tests/support.rs index 42ca184eb16..d3cc32e301f 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/support.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/support.rs @@ -197,6 +197,7 @@ pub(crate) fn legacy_call( LegacySurface { call_type: "test", input_description: "test input", + stream: None, }, call, asynchronous, diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index f946dbc763a..e70e4a8c58f 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -69,11 +69,7 @@ pub trait PythonLifecycle: Send + Sync { timing: Timing, ) -> PyResult; - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult; + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult; /// The call streams and its stream was handed to the caller. The caller is not /// inside an await here, so this step and `delivered` cannot suspend. diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 25f78d7013e..d6b31b27314 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -793,11 +793,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri } } - fn emit( - &mut self, - py: Python<'_>, - event: LifecycleEvent<'_>, - ) -> PyResult { + fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult { self.log.push(match event { LifecycleEvent::Started { .. } => "started".into(), LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => { diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index b4259aca5ba..8c42315ac59 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -1,7 +1,7 @@ mod host; use host::MessagesRouteHost; -use litellm_callbacks_legacy::{LegacySurface, PublicCall, run_legacy_call}; +use litellm_callbacks_legacy::{LegacySurface, PassThroughStream, PublicCall, run_legacy_call}; use litellm_core::messages::route::{messages_machine, supports}; use pyo3::{ prelude::*, @@ -13,6 +13,10 @@ use crate::errors::RustBridgeDeclined; const SURFACE: LegacySurface = LegacySurface { call_type: "anthropic_messages", input_description: "Messages", + stream: Some(PassThroughStream { + url_route: "/v1/messages", + endpoint_type: "anthropic", + }), }; fn run_messages( diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b5bb941708d..8afa1e2a906 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -15,6 +15,7 @@ use pyo3::{ const SURFACE: LegacySurface = LegacySurface { call_type: "ocr", input_description: "OCR document processing", + stream: None, }; const ASYNC_SURFACE: LegacySurface = LegacySurface { diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index fd0fac1799a..bac40442ce5 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -331,6 +331,8 @@ def stream_opened(logger: Logging) -> None: def stream_success( logger: Logging, + url_route: str, + endpoint_type: str, request_body: dict[str, object], chunks: list[bytes], start: datetime.datetime, @@ -353,9 +355,9 @@ def stream_success( coroutine: Final = build( litellm_logging_obj=logger, passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ, - url_route="/v1/messages", + url_route=url_route, request_body=request_body, - endpoint_type=EndpointType.ANTHROPIC, + endpoint_type=EndpointType(endpoint_type), start_time=start, raw_bytes=chunks, end_time=end, @@ -374,14 +376,18 @@ def stream_success( def stream_failure( - logger: Logging, request_body: dict[str, object], chunks: list[bytes], error: Exception + logger: Logging, + endpoint_type: str, + request_body: dict[str, object], + chunks: list[bytes], + error: Exception, ) -> Coroutine[object, object, None]: from litellm.proxy.pass_through_endpoints.streaming_handler import PassThroughStreamingHandler from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType return PassThroughStreamingHandler.schedule_stream_failure_logging( litellm_logging_obj=logger, - endpoint_type=EndpointType.ANTHROPIC, + endpoint_type=EndpointType(endpoint_type), request_body=request_body, raw_bytes=chunks, exception=error, From 19ffb584eb6b9c28118f96fb47642d1fca5442b3 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:37:05 -0700 Subject: [PATCH 242/253] refactor(rust): rename litellm-callbacks to litellm-host and HostOpError to InvokeError --- litellm-rust/Cargo.lock | 28 +++++++++---------- litellm-rust/Cargo.toml | 2 +- .../crates/callbacks-legacy/AGENTS.md | 2 +- .../crates/callbacks-legacy/Cargo.toml | 2 +- .../crates/callbacks-legacy/src/adapter.rs | 2 +- .../crates/callbacks-legacy/src/call.rs | 2 +- .../crates/callbacks-legacy/src/callbacks.rs | 2 +- .../tests/deployment_hooks.rs | 2 +- .../crates/callbacks-legacy/tests/payload.rs | 2 +- .../crates/callbacks-legacy/tests/terminal.rs | 2 +- litellm-rust/crates/core/Cargo.toml | 2 +- litellm-rust/crates/core/src/machine/auth.rs | 2 +- litellm-rust/crates/core/src/machine/mod.rs | 2 +- litellm-rust/crates/core/src/messages/mod.rs | 2 +- .../crates/core/src/messages/route.rs | 4 +-- litellm-rust/crates/core/src/ocr/client.rs | 2 +- litellm-rust/crates/core/src/ocr/handler.rs | 2 +- litellm-rust/crates/core/src/ocr/route.rs | 4 +-- .../tests/azure_document_intelligence_ocr.rs | 4 +-- litellm-rust/crates/core/tests/ocr.rs | 6 ++-- .../crates/core/tests/ocr/document.rs | 2 +- litellm-rust/crates/core/tests/ocr/support.rs | 4 +-- litellm-rust/crates/core/tests/reducto_ocr.rs | 4 +-- litellm-rust/crates/host-python/AGENTS.md | 4 +-- litellm-rust/crates/host-python/Cargo.toml | 2 +- .../crates/host-python/src/adapter.rs | 10 +++---- litellm-rust/crates/host-python/src/driver.rs | 26 ++++++++--------- litellm-rust/crates/host-python/src/lib.rs | 4 +-- .../crates/{callbacks => host}/Cargo.toml | 2 +- .../crates/{callbacks => host}/src/event.rs | 0 .../crates/{callbacks => host}/src/host.rs | 0 .../crates/{callbacks => host}/src/lib.rs | 0 .../crates/{callbacks => host}/src/machine.rs | 0 .../crates/{callbacks => host}/src/route.rs | 0 .../crates/{callbacks => host}/src/run.rs | 0 litellm-rust/crates/llms/Cargo.toml | 2 +- .../llms/src/custom_httpx/llm_http_handler.rs | 2 +- .../python-bridge/src/routes/messages/host.rs | 6 ++-- .../python-bridge/src/routes/ocr/host.rs | 6 ++-- 39 files changed, 75 insertions(+), 75 deletions(-) rename litellm-rust/crates/{callbacks => host}/Cargo.toml (91%) rename litellm-rust/crates/{callbacks => host}/src/event.rs (100%) rename litellm-rust/crates/{callbacks => host}/src/host.rs (100%) rename litellm-rust/crates/{callbacks => host}/src/lib.rs (100%) rename litellm-rust/crates/{callbacks => host}/src/machine.rs (100%) rename litellm-rust/crates/{callbacks => host}/src/route.rs (100%) rename litellm-rust/crates/{callbacks => host}/src/run.rs (100%) diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index bf58a81a6c5..5a8a613204f 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2027,22 +2027,12 @@ dependencies = [ "tokio", ] -[[package]] -name = "litellm-callbacks" -version = "0.1.0" -dependencies = [ - "litellm-auth", - "rstest", - "serde_json", - "tokio", -] - [[package]] name = "litellm-callbacks-legacy" version = "0.1.0" dependencies = [ "litellm-auth", - "litellm-callbacks", + "litellm-host", "litellm-host-python", "proptest", "pyo3", @@ -2060,8 +2050,8 @@ dependencies = [ "futures-util", "litellm-auth", "litellm-auth-aws", - "litellm-callbacks", "litellm-core-utils", + "litellm-host", "litellm-llms", "litellm-types", "mime_guess", @@ -2114,12 +2104,22 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-host" +version = "0.1.0" +dependencies = [ + "litellm-auth", + "rstest", + "serde_json", + "tokio", +] + [[package]] name = "litellm-host-python" version = "0.1.0" dependencies = [ "futures-util", - "litellm-callbacks", + "litellm-host", "pyo3", "pyo3-async-runtimes", "pythonize", @@ -2143,9 +2143,9 @@ dependencies = [ "litellm-auth-aws", "litellm-auth-azure", "litellm-auth-gcp", - "litellm-callbacks", "litellm-core-utils", "litellm-framing", + "litellm-host", "litellm-types", "reqwest 0.12.28", "rstest", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 32f925b8d8b..de6eacc62ee 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -10,7 +10,7 @@ repository = "https://github.com/BerriAI/litellm" [workspace.dependencies] litellm-core = { path = "crates/core" } -litellm-callbacks = { path = "crates/callbacks" } +litellm-host = { path = "crates/host" } litellm-callbacks-legacy = { path = "crates/callbacks-legacy" } litellm-framing = { path = "crates/framer" } litellm-auth = { path = "crates/auth" } diff --git a/litellm-rust/crates/callbacks-legacy/AGENTS.md b/litellm-rust/crates/callbacks-legacy/AGENTS.md index e184e2fb415..8b2e1c15f6e 100644 --- a/litellm-rust/crates/callbacks-legacy/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy/AGENTS.md @@ -11,7 +11,7 @@ - Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup` - Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only - - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-callbacks`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view + - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view - Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts - Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once diff --git a/litellm-rust/crates/callbacks-legacy/Cargo.toml b/litellm-rust/crates/callbacks-legacy/Cargo.toml index efacde051ed..023c13d912b 100644 --- a/litellm-rust/crates/callbacks-legacy/Cargo.toml +++ b/litellm-rust/crates/callbacks-legacy/Cargo.toml @@ -7,7 +7,7 @@ repository.workspace = true autotests = false [dependencies] -litellm-callbacks.workspace = true +litellm-host.workspace = true litellm-host-python.workspace = true pyo3.workspace = true diff --git a/litellm-rust/crates/callbacks-legacy/src/adapter.rs b/litellm-rust/crates/callbacks-legacy/src/adapter.rs index 4e9167100e0..6c013cd1ea5 100644 --- a/litellm-rust/crates/callbacks-legacy/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy/src/adapter.rs @@ -2,7 +2,7 @@ //! raises is answered with the same `Logging` calls, in the same order, as the Python //! `@client` path makes them. -use litellm_callbacks::event::{ +use litellm_host::event::{ FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds, }; use litellm_host_python::{ diff --git a/litellm-rust/crates/callbacks-legacy/src/call.rs b/litellm-rust/crates/callbacks-legacy/src/call.rs index bd1e2525d3d..b37790f60a8 100644 --- a/litellm-rust/crates/callbacks-legacy/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy/src/call.rs @@ -3,7 +3,7 @@ //! lifetime. No other callback host has that obligation, which is why nothing outside //! this crate holds them. -use litellm_callbacks::{machine::Machine, route::Route}; +use litellm_host::{machine::Machine, route::Route}; use litellm_host_python::{RouteHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, diff --git a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs index 9fcfe98368e..5f04224e6d7 100644 --- a/litellm-rust/crates/callbacks-legacy/src/callbacks.rs +++ b/litellm-rust/crates/callbacks-legacy/src/callbacks.rs @@ -2,7 +2,7 @@ //! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls //! duplication. All of it expires with the legacy callback contract. -use litellm_callbacks::event::{RequestContext, WireRequest}; +use litellm_host::event::{RequestContext, WireRequest}; use litellm_host_python::to_py; use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict}; diff --git a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs index f4602739548..52c5e47f83f 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/deployment_hooks.rs @@ -1,6 +1,6 @@ use std::ffi::CStr; -use litellm_callbacks::event::{FailureOrigin, Timing}; +use litellm_host::event::{FailureOrigin, Timing}; use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; use pyo3::exceptions::asyncio::CancelledError; use pyo3::prelude::*; diff --git a/litellm-rust/crates/callbacks-legacy/tests/payload.rs b/litellm-rust/crates/callbacks-legacy/tests/payload.rs index 67ad4ab8a2a..5459b36af27 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/payload.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/payload.rs @@ -1,7 +1,7 @@ use std::ffi::CStr; use litellm_auth::SecretValue; -use litellm_callbacks::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; +use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py}; use proptest::prelude::*; use pyo3::prelude::*; diff --git a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs index 5688f70f387..f68209233f2 100644 --- a/litellm-rust/crates/callbacks-legacy/tests/terminal.rs +++ b/litellm-rust/crates/callbacks-legacy/tests/terminal.rs @@ -1,6 +1,6 @@ use std::ffi::CStr; -use litellm_callbacks::event::{FailureOrigin, Timing}; +use litellm_host::event::{FailureOrigin, Timing}; use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle}; use pyo3::exceptions::PyRuntimeError; use pyo3::exceptions::asyncio::CancelledError; diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index db6cfc4b340..3995a235778 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -9,7 +9,7 @@ autotests = false [dependencies] litellm-types.workspace = true litellm-core-utils.workspace = true -litellm-callbacks.workspace = true +litellm-host.workspace = true bytes.workspace = true futures-util.workspace = true base64.workspace = true diff --git a/litellm-rust/crates/core/src/machine/auth.rs b/litellm-rust/crates/core/src/machine/auth.rs index 6a3e4daf6ee..cf91458ae43 100644 --- a/litellm-rust/crates/core/src/machine/auth.rs +++ b/litellm-rust/crates/core/src/machine/auth.rs @@ -1,7 +1,7 @@ use std::sync::Arc; use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -use litellm_callbacks::route::Route; +use litellm_host::route::Route; use super::{HostChannel, MachineFault}; diff --git a/litellm-rust/crates/core/src/machine/mod.rs b/litellm-rust/crates/core/src/machine/mod.rs index d6db488159e..ffcefc663ab 100644 --- a/litellm-rust/crates/core/src/machine/mod.rs +++ b/litellm-rust/crates/core/src/machine/mod.rs @@ -7,7 +7,7 @@ mod auth; use std::{future::Future, pin::Pin}; pub use auth::{HostTokenProvider, TokenRoute}; -use litellm_callbacks::{ +use litellm_host::{ event::{MachineEvent, RequestContext, WireRequest}, host::{Demand, HostOp, HostResult}, machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index e36c6668efe..c07a83c7e43 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -33,7 +33,7 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result Ok(message), MessagesOutput::Streamed => Err(Error::Unsupported( "streamed responses need a streaming host", diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 2fd2f79907a..549f848049f 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -2,12 +2,12 @@ use std::{sync::Mutex, time::Duration}; use bytes::Bytes; use litellm_auth::SecretValue; -use litellm_callbacks::{ +use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; +use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, route::Route, }; -use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 03782d91f24..c05622932b1 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -12,7 +12,7 @@ pub async fn perform( client: &OcrClient, request: LiteLLMOcrRequest, ) -> Result { - litellm_callbacks::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await + litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await } pub async fn ocr(request: LiteLLMOcrRequest) -> Result { diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index aff1eded2cc..bbf9cfa0e02 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,6 +1,6 @@ use futures_util::future::BoxFuture; use litellm_auth::SecretValue; -use litellm_callbacks::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; +use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest}; use litellm_llms::{ base_llm::ocr::{ error::Error, diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index e6ee45c64a8..d3711c34fa7 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -1,7 +1,7 @@ use std::sync::{Arc, Mutex}; use litellm_auth::ResolvedCredential; -use litellm_callbacks::{ +use litellm_host::{ event::{CallEvent, RequestContext, WireRequest}, route::Route, }; @@ -175,7 +175,7 @@ impl LocalOcrHost { } } -impl litellm_callbacks::host::Host for LocalOcrHost { +impl litellm_host::host::Host for LocalOcrHost { async fn route(&self, op: OcrOp) -> Result { match op { OcrOp::ProjectRequest => self diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index 62544931dfc..3cbe6fe3159 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::{CallEvent, MachineEvent}; +use litellm_host::event::{CallEvent, MachineEvent}; use litellm_llms::base_llm::ocr::error::Error; use rstest::rstest; use serde_json::{Value, json}; @@ -466,7 +466,7 @@ async fn model_id_is_encoded_and_dot_segments_are_rejected() { mod transformation { use std::sync::{Arc, Mutex}; - use litellm_callbacks::event::{CallEvent, MachineEvent}; + use litellm_host::event::{CallEvent, MachineEvent}; use litellm_llms::base_llm::ocr::transformation::OcrDocument; use serde_json::{Value, json}; diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index a6001951361..41a650945bc 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -1,6 +1,6 @@ use std::sync::{Arc, Mutex}; -use litellm_callbacks::{ +use litellm_host::{ event::{CallEvent, MachineEvent, WireRequest}, host::{Host, HostOp, HostResult}, machine::{HostFailure, Machine, MachineStep}, @@ -820,7 +820,7 @@ impl Host for CallerTokenHost { async fn before_send( &self, wire: WireRequest, - _: &litellm_callbacks::event::RequestContext, + _: &litellm_host::event::RequestContext, ) -> Result { let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization"); let authorization = wire @@ -855,7 +855,7 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_ trace: Mutex::new(Vec::new()), }; - litellm_callbacks::run::run(ocr_machine(ocr_client()), &host) + litellm_host::run::run(ocr_machine(ocr_client()), &host) .await .unwrap(); server.await.unwrap(); diff --git a/litellm-rust/crates/core/tests/ocr/document.rs b/litellm-rust/crates/core/tests/ocr/document.rs index 5e10ce3ad2a..855548dc6bf 100644 --- a/litellm-rust/crates/core/tests/ocr/document.rs +++ b/litellm-rust/crates/core/tests/ocr/document.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::WireRequest; +use litellm_host::event::WireRequest; use litellm_llms::base_llm::ocr::error::Error; use rstest::rstest; use serde_json::{Value, json}; diff --git a/litellm-rust/crates/core/tests/ocr/support.rs b/litellm-rust/crates/core/tests/ocr/support.rs index 44313d5f552..b368a754656 100644 --- a/litellm-rust/crates/core/tests/ocr/support.rs +++ b/litellm-rust/crates/core/tests/ocr/support.rs @@ -1,7 +1,7 @@ use std::sync::{Arc, Mutex}; use futures_util::future::BoxFuture; -use litellm_callbacks::event::WireRequest; +use litellm_host::event::WireRequest; use litellm_llms::{ base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}, custom_httpx::llm_http_handler::{CallHooks, OcrClient}, @@ -45,7 +45,7 @@ pub(crate) async fn perform_ocr(request: LiteLLMOcrRequest) -> Result Result { - litellm_callbacks::run::run(ocr_machine(ocr_client()), &host).await + litellm_host::run::run(ocr_machine(ocr_client()), &host).await } pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest { diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs index 8c037889a17..83e7754122b 100644 --- a/litellm-rust/crates/core/tests/reducto_ocr.rs +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -1,4 +1,4 @@ -use litellm_callbacks::event::{CallEvent, MachineEvent, WireRequest}; +use litellm_host::event::{CallEvent, MachineEvent, WireRequest}; use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument}; use rstest::rstest; use serde_json::{Value, json}; @@ -351,7 +351,7 @@ async fn guardrail_rewrites_document_before_upload() { } mod transformation { - use litellm_callbacks::event::{CallEvent, MachineEvent, WireRequest}; + use litellm_host::event::{CallEvent, MachineEvent, WireRequest}; use litellm_llms::{ base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext}, reducto::ocr::transformation::*, diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 903ccb06c77..5aca13eeb18 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -1,9 +1,9 @@ - Target invariants; implementation and runtime validation may lag these rules - Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits - - No LiteLLM domain dependencies beyond `litellm-callbacks`: no route types, no `Logging` policy, no public API registration, no cdylib build features + - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) - - A native failure, including one a host op returns as `HostOpError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is + - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs - Prefer `Bound<'py, T>` for attached operations/results, `Py` for retention; binding/unbinding does not copy payloads diff --git a/litellm-rust/crates/host-python/Cargo.toml b/litellm-rust/crates/host-python/Cargo.toml index ae0cebada59..e2c83fe1081 100644 --- a/litellm-rust/crates/host-python/Cargo.toml +++ b/litellm-rust/crates/host-python/Cargo.toml @@ -7,7 +7,7 @@ repository.workspace = true [dependencies] futures-util.workspace = true -litellm-callbacks.workspace = true +litellm-host.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true pythonize.workspace = true diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index e70e4a8c58f..3a4cb49be4d 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -1,5 +1,5 @@ -use litellm_callbacks::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; -use litellm_callbacks::route::Route; +use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest}; +use litellm_host::route::Route; use pyo3::exceptions::PyRuntimeError; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -89,12 +89,12 @@ pub trait PythonLifecycle: Send + Sync { /// rejected it, which the route classifies like any other native failure, or Python code /// raised, which reaches the caller as it was raised. #[derive(Debug)] -pub enum HostOpError { +pub enum InvokeError { Native(E), Python(PyErr), } -impl From for HostOpError { +impl From for InvokeError { fn from(error: PyErr) -> Self { Self::Python(error) } @@ -117,7 +117,7 @@ pub trait RouteHost: Send + Sync { py: Python<'_>, arguments: &Bound<'_, PyDict>, op: ::Op, - ) -> Result<::OpResult, HostOpError<::Error>>; + ) -> Result<::OpResult, InvokeError<::Error>>; fn complete( &mut self, diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index d6b31b27314..392d36e10f4 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -2,10 +2,10 @@ use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; -use litellm_callbacks::event::{FailureOrigin, Timing, epoch_seconds}; -use litellm_callbacks::host::{Demand, HostOp, HostResult, HostStep}; -use litellm_callbacks::machine::{HostFailure, Machine, MachineStep}; -use litellm_callbacks::route::Route; +use litellm_host::event::{FailureOrigin, Timing, epoch_seconds}; +use litellm_host::host::{Demand, HostOp, HostResult, HostStep}; +use litellm_host::machine::{HostFailure, Machine, MachineStep}; +use litellm_host::route::Route; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -13,7 +13,7 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - HostOpError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; @@ -282,12 +282,12 @@ where let arguments = self.arguments.as_ref().ok_or_else(missing_state)?; match self.route.invoke(py, arguments.bind(py), op) { Ok(result) => Ok(HostResult::Route(result)), - Err(HostOpError::Native(error)) => { + Err(InvokeError::Native(error)) => { return self .resume_core(py, Some(Err(HostFailure::Error(error)))) .map(Next::Continue); } - Err(HostOpError::Python(error)) => Err(error), + Err(InvokeError::Python(error)) => Err(error), } } HostOp::BeforeSend { wire, context } => { @@ -534,8 +534,8 @@ where mod tests { use std::sync::{Arc, Mutex}; - use litellm_callbacks::event::{MachineEvent, RequestContext, WireRequest}; - use litellm_callbacks::machine::{Interrupted, Step}; + use litellm_host::event::{MachineEvent, RequestContext, WireRequest}; + use litellm_host::machine::{Interrupted, Step}; use pyo3::exceptions::{PyBaseException, PyValueError}; use pyo3::types::PyDict; @@ -692,12 +692,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri _: Python<'_>, arguments: &Bound<'_, PyDict>, op: &'static str, - ) -> Result> { + ) -> Result> { self.log.push(format!("route:{op}")); match self.op { OpScript::Answer => Ok(format!("{op}:{}", arguments.len())), OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()), - OpScript::RejectNatively => Err(HostOpError::Native(Error("op rejected".into()))), + OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))), } } @@ -899,7 +899,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri context: Box::new(context()), }, HostOp::Emit(MachineEvent::ResponseReceived { - raw: litellm_callbacks::event::RawResponse { body: "raw".into() }, + raw: litellm_host::event::RawResponse { body: "raw".into() }, }), ], outcome: Some(Ok("done".into())), @@ -1189,7 +1189,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri py: Python<'_>, _: &Bound<'_, PyDict>, _: &'static str, - ) -> Result> { + ) -> Result> { self.0.push("route"); Err(PyErr::from_value( py.import("asyncio") diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 3738a9b1c3c..583a4eb91b6 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -1,5 +1,5 @@ //! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and -//! asyncio glue, and the driver that runs a native [`Machine`](litellm_callbacks::machine::Machine) +//! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine) //! against a Python route host and a Python lifecycle. Everything here is Python-specific by //! construction; another host language gets its own crate of the same shape. @@ -13,7 +13,7 @@ mod handle; mod marshal; pub use adapter::{ - HostOpError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; diff --git a/litellm-rust/crates/callbacks/Cargo.toml b/litellm-rust/crates/host/Cargo.toml similarity index 91% rename from litellm-rust/crates/callbacks/Cargo.toml rename to litellm-rust/crates/host/Cargo.toml index a68ebc26a8d..ebabe5ccbdc 100644 --- a/litellm-rust/crates/callbacks/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -1,5 +1,5 @@ [package] -name = "litellm-callbacks" +name = "litellm-host" version = "0.1.0" edition.workspace = true license.workspace = true diff --git a/litellm-rust/crates/callbacks/src/event.rs b/litellm-rust/crates/host/src/event.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/event.rs rename to litellm-rust/crates/host/src/event.rs diff --git a/litellm-rust/crates/callbacks/src/host.rs b/litellm-rust/crates/host/src/host.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/host.rs rename to litellm-rust/crates/host/src/host.rs diff --git a/litellm-rust/crates/callbacks/src/lib.rs b/litellm-rust/crates/host/src/lib.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/lib.rs rename to litellm-rust/crates/host/src/lib.rs diff --git a/litellm-rust/crates/callbacks/src/machine.rs b/litellm-rust/crates/host/src/machine.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/machine.rs rename to litellm-rust/crates/host/src/machine.rs diff --git a/litellm-rust/crates/callbacks/src/route.rs b/litellm-rust/crates/host/src/route.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/route.rs rename to litellm-rust/crates/host/src/route.rs diff --git a/litellm-rust/crates/callbacks/src/run.rs b/litellm-rust/crates/host/src/run.rs similarity index 100% rename from litellm-rust/crates/callbacks/src/run.rs rename to litellm-rust/crates/host/src/run.rs diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index 4ca6c7cb2a5..d295e4407ba 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -15,7 +15,7 @@ litellm-auth.workspace = true litellm-auth-aws.workspace = true litellm-auth-azure.workspace = true litellm-auth-gcp.workspace = true -litellm-callbacks.workspace = true +litellm-host.workspace = true litellm-framing.workspace = true base64.workspace = true bytes.workspace = true diff --git a/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs b/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs index 7635bdd3d04..fdd568d83fd 100644 --- a/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs +++ b/litellm-rust/crates/llms/src/custom_httpx/llm_http_handler.rs @@ -3,7 +3,7 @@ use std::{sync::OnceLock, time::Duration}; use bytes::{Bytes, BytesMut}; use futures_util::future::BoxFuture; use litellm_auth_gcp::VertexAuth; -use litellm_callbacks::event::WireRequest; +use litellm_host::event::WireRequest; use serde::{Serialize, de::DeserializeOwned}; use serde_json::Value; diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 0c87aea1ca2..b590018bc4e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -3,7 +3,7 @@ use litellm_core::messages::{ Error, route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput}, }; -use litellm_host_python::{HostOpError, RouteHost, from_py, lookup, to_py}; +use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py}; use litellm_llms::custom_httpx::transport::Error as TransportError; use pyo3::{ exceptions::{PyException, PyValueError}, @@ -136,12 +136,12 @@ impl RouteHost for MessagesRouteHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, op: MessagesOp, - ) -> Result> { + ) -> Result> { match op { MessagesOp::ProjectRequest => self .project(py, arguments) .map(|call| MessagesOpResult::Request(Box::new(call))) - .map_err(|error| HostOpError::Python(self.map_failure(py, error))), + .map_err(|error| InvokeError::Python(self.map_failure(py, error))), } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 19211671fd7..77c8d5d6641 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,6 +1,6 @@ use litellm_auth::ResolvedCredential; use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult}; -use litellm_host_python::{HostOpError, RouteHost, missing_state, to_py}; +use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py}; use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse}; use pyo3::{ exceptions::{PyBaseException, PyException}, @@ -113,9 +113,9 @@ impl RouteHost for OcrRouteHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, op: OcrOp, - ) -> Result> { + ) -> Result> { self.answer(py, arguments, op) - .map_err(|error| HostOpError::Python(self.map_failure(py, error))) + .map_err(|error| InvokeError::Python(self.map_failure(py, error))) } fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult> { From b7686d78b675a020bb15123c571438f647674431 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:40:02 -0700 Subject: [PATCH 243/253] refactor(rust): move RouteMachine out of core into litellm-host next to the Machine trait core/src/machine was route-neutral runtime code sitting among route surfaces, and the workspace had two modules named machine. It now lives in litellm-host beside the contract it implements, so core only holds routes. The OCR conversion from MachineFault moves to litellm-llms because the orphan rule no longer allows it in core --- litellm-rust/crates/core/src/lib.rs | 1 - litellm-rust/crates/core/src/messages/route.rs | 6 ++---- litellm-rust/crates/core/src/ocr/route.rs | 16 ++-------------- litellm-rust/crates/host/Cargo.toml | 2 +- litellm-rust/crates/host/src/lib.rs | 2 +- .../crates/{core => host}/src/machine/auth.rs | 5 ++--- .../host/src/{machine.rs => machine/mod.rs} | 6 ++++++ .../mod.rs => host/src/machine/route_machine.rs} | 10 ++++------ .../crates/llms/src/base_llm/ocr/error.rs | 11 +++++++++++ 9 files changed, 29 insertions(+), 30 deletions(-) rename litellm-rust/crates/{core => host}/src/machine/auth.rs (98%) rename litellm-rust/crates/host/src/{machine.rs => machine/mod.rs} (91%) rename litellm-rust/crates/{core/src/machine/mod.rs => host/src/machine/route_machine.rs} (97%) diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 58aef6cd629..e3e2fb48721 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -2,7 +2,6 @@ pub mod audio_transcription; pub mod chat_completions; pub mod constants; pub mod error; -pub mod machine; pub mod messages; pub mod ocr; pub mod responses; diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 549f848049f..3d607ba67f6 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -6,6 +6,7 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_host::{ event::{MachineEvent, RawResponse, RequestContext, WireRequest}, host::{Demand, Host}, + machine::{HostChannel, MachineFault, RouteMachine}, route::Route, }; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; @@ -18,10 +19,7 @@ use super::{ prepare::prepare_provider_request, types::MessagesRequest, }; -use crate::{ - constants::ANTHROPIC_MESSAGES_PROVIDER, - machine::{HostChannel, MachineFault, RouteMachine}, -}; +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum MessagesOp { diff --git a/litellm-rust/crates/core/src/ocr/route.rs b/litellm-rust/crates/core/src/ocr/route.rs index d3711c34fa7..bfc8c5ca965 100644 --- a/litellm-rust/crates/core/src/ocr/route.rs +++ b/litellm-rust/crates/core/src/ocr/route.rs @@ -3,6 +3,7 @@ use std::sync::{Arc, Mutex}; use litellm_auth::ResolvedCredential; use litellm_host::{ event::{CallEvent, RequestContext, WireRequest}, + machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute}, route::Route, }; use litellm_llms::{ @@ -11,10 +12,7 @@ use litellm_llms::{ }; use super::handler::perform_ocr_request; -use crate::{ - machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute}, - ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest}, -}; +use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum OcrOp { @@ -56,16 +54,6 @@ impl TokenRoute for Ocr { } } -impl From for Error { - fn from(fault: MachineFault) -> Self { - Self::InvalidRequest(match fault { - MachineFault::Abandoned => "OCR host driver was abandoned".into(), - MachineFault::Protocol(message) => format!("OCR {message}"), - MachineFault::Mismatch => "invalid OCR host operation result".into(), - }) - } -} - pub type OcrHost = HostChannel; pub type OcrMachine = RouteMachine; diff --git a/litellm-rust/crates/host/Cargo.toml b/litellm-rust/crates/host/Cargo.toml index ebabe5ccbdc..0c7c46192b5 100644 --- a/litellm-rust/crates/host/Cargo.toml +++ b/litellm-rust/crates/host/Cargo.toml @@ -8,7 +8,7 @@ repository.workspace = true [dependencies] litellm-auth.workspace = true serde_json.workspace = true +tokio = { workspace = true, features = ["sync"] } [dev-dependencies] rstest.workspace = true -tokio = { workspace = true, features = ["macros"] } diff --git a/litellm-rust/crates/host/src/lib.rs b/litellm-rust/crates/host/src/lib.rs index 41b0983f0ce..65479c2380f 100644 --- a/litellm-rust/crates/host/src/lib.rs +++ b/litellm-rust/crates/host/src/lib.rs @@ -1,7 +1,7 @@ //! The contract between a native call and the host runtime that drives it. //! //! A host is whatever sits on the far side of the language boundary: CPython today, -//! another runtime later. Core implements [`machine::Machine`] per route and never learns +//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns //! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers //! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent. diff --git a/litellm-rust/crates/core/src/machine/auth.rs b/litellm-rust/crates/host/src/machine/auth.rs similarity index 98% rename from litellm-rust/crates/core/src/machine/auth.rs rename to litellm-rust/crates/host/src/machine/auth.rs index cf91458ae43..ba7e242e766 100644 --- a/litellm-rust/crates/core/src/machine/auth.rs +++ b/litellm-rust/crates/host/src/machine/auth.rs @@ -1,9 +1,8 @@ use std::sync::Arc; -use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -use litellm_host::route::Route; - use super::{HostChannel, MachineFault}; +use crate::route::Route; +use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; /// A route whose host can mint credentials on the call's behalf. pub trait TokenRoute: Route { diff --git a/litellm-rust/crates/host/src/machine.rs b/litellm-rust/crates/host/src/machine/mod.rs similarity index 91% rename from litellm-rust/crates/host/src/machine.rs rename to litellm-rust/crates/host/src/machine/mod.rs index 2942913f095..2c26db61582 100644 --- a/litellm-rust/crates/host/src/machine.rs +++ b/litellm-rust/crates/host/src/machine/mod.rs @@ -1,6 +1,12 @@ +mod auth; +mod route_machine; + use std::future::Future; use std::pin::Pin; +pub use auth::{HostTokenProvider, TokenRoute}; +pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine}; + use crate::host::{HostOp, HostResult}; use crate::route::Route; diff --git a/litellm-rust/crates/core/src/machine/mod.rs b/litellm-rust/crates/host/src/machine/route_machine.rs similarity index 97% rename from litellm-rust/crates/core/src/machine/mod.rs rename to litellm-rust/crates/host/src/machine/route_machine.rs index ffcefc663ab..38a0b8bc16a 100644 --- a/litellm-rust/crates/core/src/machine/mod.rs +++ b/litellm-rust/crates/host/src/machine/route_machine.rs @@ -2,18 +2,16 @@ //! place, and turns the host operations that future requests into [`Machine`] steps. No //! task is spawned; dropping the machine drops the in-flight call. -mod auth; - use std::{future::Future, pin::Pin}; -pub use auth::{HostTokenProvider, TokenRoute}; -use litellm_host::{ +use tokio::sync::{mpsc, oneshot}; + +use super::{HostFailure, Interrupted, Machine, MachineStep, Step}; +use crate::{ event::{MachineEvent, RequestContext, WireRequest}, host::{Demand, HostOp, HostResult}, - machine::{HostFailure, Interrupted, Machine, MachineStep, Step}, route::Route, }; -use tokio::sync::{mpsc, oneshot}; /// The machine's own failures, distinct from anything the provider call reports. #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs index c3f481d7d44..3061a9fe2b2 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/error.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/error.rs @@ -102,6 +102,17 @@ pub enum Error { Headers(#[from] crate::custom_httpx::http_handler::HeaderError), } +impl From for Error { + fn from(fault: litellm_host::machine::MachineFault) -> Self { + use litellm_host::machine::MachineFault; + Self::InvalidRequest(match fault { + MachineFault::Abandoned => "OCR host driver was abandoned".into(), + MachineFault::Protocol(message) => format!("OCR {message}"), + MachineFault::Mismatch => "invalid OCR host operation result".into(), + }) + } +} + impl From for Error { fn from(error: litellm_core_utils::call_arguments::ArgumentError) -> Self { Self::RequestField { From 41873e2bc796bce93b74ec9431ae01b403ed46a3 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 22:41:17 +0000 Subject: [PATCH 244/253] fix(rust): box messages response output Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/core/src/messages/mod.rs | 2 +- litellm-rust/crates/core/src/messages/route.rs | 5 +++-- .../crates/python-bridge/src/routes/messages/host.rs | 2 +- 3 files changed, 5 insertions(+), 4 deletions(-) diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index c07a83c7e43..289f79109dd 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -34,7 +34,7 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result Ok(message), + MessagesOutput::Message(message) => Ok(*message), MessagesOutput::Streamed => Err(Error::Unsupported( "streamed responses need a streaming host", )), diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 3d607ba67f6..838b56fcb4b 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -48,7 +48,7 @@ impl MessagesCall { } pub enum MessagesOutput { - Message(AnthropicMessagesResponse), + Message(Box), /// Every chunk already reached the host through `Deliver`. Streamed, } @@ -174,7 +174,8 @@ async fn execute(host: MessagesHost) -> Result { raw: RawResponse { body: text.clone() }, }) .await?; - decode_response(request.config, &request.model, &text).map(MessagesOutput::Message) + decode_response(request.config, &request.model, &text) + .map(|message| MessagesOutput::Message(Box::new(message))) } /// Hands each upstream chunk to the caller as it arrives. A caller that stops reading diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index b590018bc4e..c1b3f59df58 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -150,7 +150,7 @@ impl RouteHost for MessagesRouteHost { MessagesOutput::Message(message) => py .import("litellm.rust_bridge.messages.route_host")? .getattr("response")? - .call1((to_py(py, &message)?,)) + .call1((to_py(py, message.as_ref())?,)) .map(Bound::unbind), MessagesOutput::Streamed => Ok(py.None()), } From ba6b22cf56ead7aba23d530b342da6d342230581 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 18 Sep 2026 15:57:50 -0700 Subject: [PATCH 245/253] test(rust): isolate callback registries per hypothesis example Replace the module-level LATEST_EDITS list with per-example callback registry isolation, and import litellm names with from-imports in the legacy callback shim so the module uses one import style. Co-Authored-By: Claude Opus 5 --- litellm/rust_bridge/legacy_callbacks.py | 18 +++--- tests/test_litellm_rust/conftest.py | 63 +++---------------- tests/test_litellm_rust/ocr/test_callbacks.py | 20 +++--- tests/test_litellm_rust/support/isolation.py | 58 +++++++++++++++++ 4 files changed, 88 insertions(+), 71 deletions(-) create mode 100644 tests/test_litellm_rust/support/isolation.py diff --git a/litellm/rust_bridge/legacy_callbacks.py b/litellm/rust_bridge/legacy_callbacks.py index bac40442ce5..30aa1d97bfc 100644 --- a/litellm/rust_bridge/legacy_callbacks.py +++ b/litellm/rust_bridge/legacy_callbacks.py @@ -65,13 +65,17 @@ def setup( def check_limits(kwargs: Mapping[str, object]) -> None: - import litellm + from litellm import ( + BudgetExceededError, + _current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor + max_budget, + num_retries_per_request, + ) from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit - current_cost: Final = litellm._current_cost # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor - if litellm.max_budget and current_cost > litellm.max_budget: - raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget) - if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request): + if max_budget and _current_cost > max_budget: + raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget) + if max_retries_per_request_hit(kwargs, num_retries_per_request): raise RuntimeError("Max retries per request hit!") @@ -281,9 +285,9 @@ def is_internal_call() -> bool: def credential_list() -> list[CredentialItem]: - import litellm + from litellm import credential_list as credentials - return litellm.credential_list + return credentials def warn_unknown_credential(name: str, loaded: int) -> None: diff --git a/tests/test_litellm_rust/conftest.py b/tests/test_litellm_rust/conftest.py index 4387ea2e2fd..1b6fcfa00db 100644 --- a/tests/test_litellm_rust/conftest.py +++ b/tests/test_litellm_rust/conftest.py @@ -1,10 +1,9 @@ import asyncio import os -from collections.abc import AsyncIterator, Generator, Iterator +from collections.abc import AsyncIterator, Generator from concurrent.futures import ThreadPoolExecutor -from contextlib import ExitStack, contextmanager -from types import ModuleType -from typing import Final, cast +from contextlib import ExitStack +from typing import Final import pytest import pytest_asyncio @@ -18,62 +17,20 @@ from litellm.rust_bridge.configuration import ( # pyright: ignore[reportPrivate _parse_env_bool, ) from tests.test_litellm_rust.support.callback_recorder import drain_logging +from tests.test_litellm_rust.support.isolation import isolated_callback_registries, rebound from tests.test_litellm_rust.support.recording_server import RecordingServer, recording_service -CALLBACK_ATTRIBUTES: Final = ( - "callbacks", - "input_callback", - "success_callback", - "failure_callback", - "_async_input_callback", - "_async_success_callback", - "_async_failure_callback", -) - - -def _list_attribute(container: ModuleType, attribute: str) -> list[object]: - value: Final = getattr(container, attribute) - if not isinstance(value, list): - raise AssertionError(f"{container.__name__}.{attribute} is not a list") - return cast(list[object], value) - - -@contextmanager -def _isolated_list(container: ModuleType, attribute: str) -> Iterator[None]: - source: Final = _list_attribute(container, attribute) - original: Final = list(source) - source.clear() # mutable-ok: test isolation mutates global registries by design - try: - yield - finally: - source.clear() - source.extend(original) - setattr(container, attribute, source) - - -@contextmanager -def _rebound(container: object, attribute: str, value: object) -> Iterator[None]: - original: Final[object] = getattr(container, attribute) - setattr(container, attribute, value) - try: - yield - finally: - setattr(container, attribute, original) - @pytest_asyncio.fixture(autouse=True, loop_scope="function") async def isolate_ocr_test_state() -> AsyncIterator[None]: with ExitStack() as stack: - for attribute in CALLBACK_ATTRIBUTES: - stack.enter_context(_isolated_list(litellm, attribute)) - stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor - stack.enter_context(_rebound(utils, "callback_list", [])) # rebind-ok: isolate legacy callback registry - stack.enter_context(_rebound(litellm, "cache", None)) # test-quality-ok: isolate process-global cache - stack.enter_context(_rebound(_CONFIGURATION, "override", None)) + stack.enter_context(isolated_callback_registries()) + stack.enter_context(rebound(litellm, "cache", None)) # test-quality-ok: isolate process-global cache + stack.enter_context(rebound(_CONFIGURATION, "override", None)) executor: Final = ThreadPoolExecutor(thread_name_prefix="rust-ocr-test-logging") - stack.enter_context(_rebound(litellm_logging, "executor", executor)) - stack.enter_context(_rebound(utils, "executor", executor)) - stack.enter_context(_rebound(thread_pool_executor, "executor", executor)) + stack.enter_context(rebound(litellm_logging, "executor", executor)) + stack.enter_context(rebound(utils, "executor", executor)) + stack.enter_context(rebound(thread_pool_executor, "executor", executor)) try: yield finally: diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index 45b99d19d90..ac4a1a11a80 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -3,6 +3,8 @@ import copy import gc import queue import threading +from collections.abc import Mapping +from types import MappingProxyType from typing import Final import pytest @@ -13,6 +15,7 @@ import litellm from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging +from tests.test_litellm_rust.support.isolation import isolated_callback_registries from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import ( OCR_DOCUMENT, @@ -307,18 +310,13 @@ JSON_VALUES: Final = st.recursive( ) -LATEST_EDITS: Final[list[dict[str, object]]] = [] - - -class ApplyLatestEdits(CustomLogger): - """Registrations can outlive one hypothesis example, so every instance applies the current example's edits.""" - - def __init__(self, latest: list[dict[str, object]]) -> None: +class ApplyEdits(CustomLogger): + def __init__(self, edits: Mapping[str, object]) -> None: super().__init__() - self.latest = latest + self.edits: Final = edits def log_pre_api_call(self, model, messages, kwargs): - request_body(kwargs).update(copy.deepcopy(self.latest[-1])) + request_body(kwargs).update(copy.deepcopy(dict(self.edits))) @settings(max_examples=25, deadline=None, suppress_health_check=[HealthCheck.function_scoped_fixture]) @@ -327,9 +325,9 @@ def test_native_ocr_provider_receives_the_body_exactly_as_pre_call_callbacks_lef ocr_server: RecordingServer, edits: dict[str, object] ) -> None: ocr_server.expected_requests = None - LATEST_EDITS.append(edits) - call_native_ocr_with_callbacks(ocr_server, [ApplyLatestEdits(LATEST_EDITS)]) + with isolated_callback_registries(): + call_native_ocr_with_callbacks(ocr_server, [ApplyEdits(MappingProxyType(edits))]) assert ocr_server.requests[-1].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT, **edits} diff --git a/tests/test_litellm_rust/support/isolation.py b/tests/test_litellm_rust/support/isolation.py new file mode 100644 index 00000000000..f98ce4843a8 --- /dev/null +++ b/tests/test_litellm_rust/support/isolation.py @@ -0,0 +1,58 @@ +from collections.abc import Generator +from contextlib import ExitStack, contextmanager +from types import ModuleType +from typing import Final, cast + +import litellm +from litellm import utils +from litellm.litellm_core_utils import litellm_logging + +CALLBACK_ATTRIBUTES: Final = ( + "callbacks", + "input_callback", + "success_callback", + "failure_callback", + "_async_input_callback", + "_async_success_callback", + "_async_failure_callback", +) + + +def _list_attribute(container: ModuleType, attribute: str) -> list[object]: + value: Final = getattr(container, attribute) + if not isinstance(value, list): + raise AssertionError(f"{container.__name__}.{attribute} is not a list") + return cast(list[object], value) + + +@contextmanager +def _isolated_list(container: ModuleType, attribute: str) -> Generator[None]: + source: Final = _list_attribute(container, attribute) + original: Final = list(source) + source.clear() # mutable-ok: test isolation mutates global registries by design + try: + yield + finally: + source.clear() + source.extend(original) + setattr(container, attribute, source) + + +@contextmanager +def rebound(container: object, attribute: str, value: object) -> Generator[None]: + original: Final[object] = getattr(container, attribute) + setattr(container, attribute, value) + try: + yield + finally: + setattr(container, attribute, original) + + +@contextmanager +def isolated_callback_registries() -> Generator[None]: + with ExitStack() as stack: + for attribute in CALLBACK_ATTRIBUTES: + stack.enter_context(_isolated_list(litellm, attribute)) + stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor + stack.enter_context(rebound(utils, "callback_list", [])) # rebind-ok: isolate legacy callback registry + yield From b3cf45e9f232c094e2f2b2a6bf59f609464bf742 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 15:58:24 -0700 Subject: [PATCH 246/253] fix(proxy): drop daily spend batches that cannot be re-sent safely instead of requeueing them --- litellm/proxy/db/db_spend_update_writer.py | 24 ++++++++- litellm/proxy/db/exception_handler.py | 18 +++++++ .../proxy/db/test_db_spend_update_writer.py | 51 ++++++++++++++++++- .../proxy/db/test_exception_handler.py | 22 ++++++++ 4 files changed, 112 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index b2e6f9dc54d..e9967fb0d67 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -31,6 +31,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.litellm_logging import coerce_model_access_groups from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy._types import ( + DB_CONNECTION_ERROR_TYPES, DB_RETRY_SAFE_ERROR_TYPES, BaseDailySpendTransaction, DailyAgentSpendTransaction, @@ -64,6 +65,7 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import ( WindowSpendTransaction, WindowSpendUpdateQueue, ) +from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING from litellm.proxy.spend_tracking.compression_savings import ( extract_compression_saved_tokens, @@ -157,6 +159,16 @@ class _DailySpendCommit(Protocol[_DailySpendTransactionT]): ) -> None: ... +_DATA_REJECTED_SQLSTATE_CLASSES: Final = frozenset({"22", "23"}) + + +def _daily_spend_commit_failure_is_requeue_safe(e: Exception) -> bool: + if isinstance(e, DB_CONNECTION_ERROR_TYPES): + return isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) + sqlstate: Final = PrismaDBExceptionHandler.postgres_sqlstate(e) + return sqlstate is None or sqlstate[:2] not in _DATA_REJECTED_SQLSTATE_CLASSES + + def _timed_request_duration_ms( payload: dict | SpendLogsPayload, request_status: Literal["success", "failure"], @@ -1319,7 +1331,17 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions), ) - except Exception as e: # noqa: BLE001 # the uncommitted rows go back on the queue; the other tables must still flush + except Exception as e: # noqa: BLE001 # whatever failed here, the other tables must still flush + if not _daily_spend_commit_failure_is_requeue_safe(e): + spend_log_error( + "Spend tracking - dropped %d daily %s spend rows: the failed commit may have applied " + "or the database refused the data, so re-sending it is not safe. Error: %s", + len(transactions), + entity_type, + str(e), + exc=e, + ) + return spend_log_error( "Spend tracking - failed to commit daily %s spend updates. " "Re-queued %d rows for retry on next tick. Error: %s", diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index 2cee5128c66..460bf5db3b1 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -1,6 +1,8 @@ from collections.abc import Awaitable, Callable, Iterator from typing import Any, Final, TypeVar +from pydantic import TypeAdapter, ValidationError + from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( DB_CONNECTION_ERROR_TYPES, @@ -17,6 +19,8 @@ _TRANSIENT_DB_UNAVAILABLE_MESSAGE: Final = ( "Service Unavailable, the authentication database is temporarily unreachable. Please retry shortly." ) +_DATABASE_ERROR_META: Final = TypeAdapter(dict[str, object]) + def _exception_chain(e: BaseException) -> Iterator[BaseException]: current = e # rebind-ok: advances one link per iteration of the bounded walk @@ -221,6 +225,20 @@ class PrismaDBExceptionHandler: or "write conflict or a deadlock" in error_message ) + @staticmethod + def postgres_sqlstate(e: Exception) -> str | None: + """The SQLSTATE Postgres attached to a failed statement, as prisma surfaces it, or None.""" + import prisma + + if not isinstance(e, _exception_types(prisma.errors.DataError)): + return None + try: + meta: Final = _DATABASE_ERROR_META.validate_python(getattr(e, "meta", None)) + except ValidationError: + return None + code: Final = meta.get("code") + return code if isinstance(code, str) else None + @staticmethod def is_read_only_transaction_error(e: Exception) -> bool: """True iff ``e`` is Postgres SQLSTATE 25006 surfaced through prisma: the diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index b27b838133b..155bca656d5 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -11,7 +11,9 @@ from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, call, patch +import httpx import pytest +from prisma.errors import RawQueryError from redis.exceptions import DataError import litellm @@ -2812,14 +2814,15 @@ async def test_failed_window_spend_commit_requeues_the_increments_and_continues_ class _DailySpendFakeDB(_WindowSpendFakeDB): """Records the daily rollup upserts it is handed and fails the ones aimed at one table.""" - def __init__(self, failing_table: str | None) -> None: + def __init__(self, failing_table: str | None, failure: Exception | None = None) -> None: super().__init__() self.failing_table = failing_table + self.failure = failure self.execute_raw_calls: list[Statement] = [] async def execute_raw(self, query: str, *args: object) -> int: if self.failing_table is not None and self.failing_table in query: - raise Exception("connection reset") + raise self.failure if self.failure is not None else Exception("connection reset") self.execute_raw_calls.append((query, args)) return len(args) @@ -2828,6 +2831,50 @@ def _daily_upserts(db: _DailySpendFakeDB, table: str) -> list[Statement]: return [statement for statement in db.execute_raw_calls if table in statement[0]] +def _postgres_rejection(sqlstate: str) -> RawQueryError: + return RawQueryError( + data={"user_facing_error": {"error_code": "P2010", "meta": {"code": sqlstate, "message": "db error"}}} + ) + + +@pytest.mark.parametrize( + ("failure", "lands_on_the_next_tick"), + [ + pytest.param(httpx.ReadTimeout("no reply"), False, id="reply lost after the statement was sent"), + pytest.param(httpx.ConnectError("refused"), True, id="statement never reached the database"), + pytest.param(_postgres_rejection("22021"), False, id="postgres refused the data itself"), + pytest.param(_postgres_rejection("23502"), False, id="postgres refused a constraint violation"), + pytest.param(_postgres_rejection("42P01"), True, id="table missing"), + pytest.param(_postgres_rejection("57014"), True, id="statement cancelled"), + ], +) +@pytest.mark.asyncio +async def test_failed_daily_spend_commit_is_requeued_only_when_the_rows_are_provably_uncommitted( + failure: Exception, lands_on_the_next_tick: bool +): + """A lost reply means the statement may already have applied, and re-sending it stacks a + second increment into the same transaction (LIT-4823); a row Postgres refuses would fail + every tick forever. Both are dropped loudly. Every other failure left nothing committed, + so its rows go back on the queue and land on the next tick.""" + db_writer = DBSpendUpdateWriter() + await db_writer.daily_spend_update_queue.add_update({"user-key": _daily_txn(user_id="user-1")}) + db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend", failure=failure) + db_writer._flush_tool_discovery_queue = AsyncMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.failure_handler = AsyncMock() + + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + db.failing_table = None + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + assert len(_daily_upserts(db, "LiteLLM_DailyUserSpend")) == (1 if lands_on_the_next_tick else 0) + assert db_writer.daily_spend_update_queue.update_queue.empty() + + @pytest.mark.asyncio async def test_failed_daily_spend_commit_requeues_the_rows_and_flushes_the_other_tables(): """With the Redis buffer off, a daily batch that failed to commit was discarded along diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 3f009137a1c..26ac1ea65ad 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -665,6 +665,28 @@ def test_is_deadlock_error_excludes_non_deadlocks(error): assert PrismaDBExceptionHandler.is_deadlock_error(error) is False +@pytest.mark.parametrize( + ("error", "sqlstate"), + [ + ( + RawQueryError( + data={"user_facing_error": {"error_code": "P2010", "meta": {"code": "22021", "message": "m"}}} + ), + "22021", + ), + (RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"message": "m"}}}), None), + (RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}), None), + (prisma_errors.DataError(data={"user_facing_error": {"meta": None}}), None), + (PrismaError("db error"), None), + (httpx.ReadTimeout("no reply"), None), + ], +) +def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error, sqlstate): + """Only a prisma data error carrying Postgres's own error code yields a SQLSTATE; a + codeless or malformed payload, an engine-level error, and a transport error yield None.""" + assert PrismaDBExceptionHandler.postgres_sqlstate(error) == sqlstate + + READ_ONLY_CONNECTOR_ERROR: Final = ( "Error occurred during query execution:\nConnectorError(ConnectorError { user_facing_error: None, " 'kind: QueryError(PostgresError { code: "25006", message: "cannot execute UPDATE in a read-only transaction", ' From 139445179a8a2fd1804cbeba1154e77b4e044fb6 Mon Sep 17 00:00:00 2001 From: ryan Date: Fri, 18 Sep 2026 22:59:32 +0000 Subject: [PATCH 247/253] ci: remove the dead Agent Shin triage workflows and scripts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/_agent_shin_actions.py | 50 - .github/scripts/agent_shin_shared.py | 211 -- .github/scripts/close_low_quality_prs.py | 573 ----- .github/scripts/triage-requirements.txt | 282 --- .github/scripts/triage_with_llm.py | 1797 -------------- .github/workflows/close_low_quality_prs.yml | 92 - .../create_daily_oss_agent_shin_branch.yml | 28 - .github/workflows/triage_reconsider.yml | 172 -- .../test_github_close_low_quality_prs.py | 856 ------- tests/test_litellm/test_github_review_gate.py | 524 ---- .../test_github_triage_with_llm.py | 2134 ----------------- .../test_github_triage_workflows.py | 264 -- 12 files changed, 6983 deletions(-) delete mode 100644 .github/scripts/_agent_shin_actions.py delete mode 100644 .github/scripts/agent_shin_shared.py delete mode 100644 .github/scripts/close_low_quality_prs.py delete mode 100644 .github/scripts/triage-requirements.txt delete mode 100644 .github/scripts/triage_with_llm.py delete mode 100644 .github/workflows/close_low_quality_prs.yml delete mode 100644 .github/workflows/create_daily_oss_agent_shin_branch.yml delete mode 100644 .github/workflows/triage_reconsider.yml delete mode 100644 tests/test_litellm/test_github_close_low_quality_prs.py delete mode 100644 tests/test_litellm/test_github_review_gate.py delete mode 100644 tests/test_litellm/test_github_triage_with_llm.py delete mode 100644 tests/test_litellm/test_github_triage_workflows.py diff --git a/.github/scripts/_agent_shin_actions.py b/.github/scripts/_agent_shin_actions.py deleted file mode 100644 index b3d1ff055b3..00000000000 --- a/.github/scripts/_agent_shin_actions.py +++ /dev/null @@ -1,50 +0,0 @@ -"""Dry-run wrapper(s) around Agent Shin GitHub mutations. - -The rollout scripts currently need only one mutation wrapped, so this module -exposes a single ``maybe_post_comment`` helper. It takes a ``dry_run: bool`` -keyword argument and the body is intentionally trivial: - - if dry_run: - print(...) # log what we would do, return - return - real_mutation(...) # otherwise, actually do it - -That shape means a dry-run preview differs from the real run in exactly one -line per side effect: the call site. So when you `python3 script.py` locally -without ``--close``, you can be confident the actions printed are the ones the -GitHub Action would have performed (modulo ordering on retry/error paths, -which are deliberately simple). Any further mutation a rollout script needs -should get the same ``maybe_*`` treatment instead of calling the raw -``triage_with_llm`` mutation directly. - -Importing from this module pulls in the real mutation from ``triage_with_llm`` -— call sites in the rollout scripts should NEVER import ``post_comment`` -directly; that would skip the dry-run gate and is the bug class this module -exists to prevent. -""" - -from __future__ import annotations - -import sys -import textwrap - -# Import the module itself rather than the bare names so monkeypatching -# `triage_with_llm.post_comment` (or any of the other mutations) in tests is -# reflected here — `from triage_with_llm import post_comment` would bind the -# original function to a local name and bypass the patch, defeating the whole -# point of these wrappers. -import triage_with_llm - - -def _log(line: str) -> None: - """Print a single dry-run line to stdout (one log statement per side effect).""" - print(line, file=sys.stdout, flush=True) - - -def maybe_post_comment(repo: str, number: int, body: str, *, dry_run: bool) -> None: - """Post a comment on ``repo#number`` — or, in dry-run, log what we would post.""" - if dry_run: - _log(f"[DRY RUN] comment {repo}#{number}:") - _log(textwrap.indent(body, " ")) - return - triage_with_llm.post_comment(repo, number, body) diff --git a/.github/scripts/agent_shin_shared.py b/.github/scripts/agent_shin_shared.py deleted file mode 100644 index 8f3dc3c2322..00000000000 --- a/.github/scripts/agent_shin_shared.py +++ /dev/null @@ -1,211 +0,0 @@ -"""Constants and helpers shared by Agent Shin's triage scripts. - -Both `triage_with_llm.py` (the LLM-judge entrypoint) and -`close_low_quality_prs.py` (the daily Greptile-score sweep) need to -agree on the same notions of: - - * What counts as a Greptile-authored review comment - (``GREPTILE_BOT_LOGINS``) and how to extract a confidence score from - its body (``SCORE_PATTERN`` / :func:`extract_greptile_score`). - * How long the 2-hour grace window is (``GRACE_PERIOD_SECONDS``) and - the HTML marker stamped into a grace-warning comment so the *other* - script can see "Agent Shin already warned" and behave accordingly - (``GRACE_COMMENT_MARKER``). - * Who Agent Shin is on GitHub (``AGENT_SHIN_DEFAULT_BOT_LOGIN``). - * How GitHub-style ISO-8601 timestamps round-trip into timezone-aware - :class:`datetime.datetime` (:func:`parse_iso8601`). - -Keeping these in one module means a future change (new Greptile output -format, a longer grace window, a new allowlisted account) is a single edit -instead of two — the original split version had to call out in comments -that the two copies "must stay in sync" precisely because nothing -enforced it. -""" - -from __future__ import annotations - -import datetime as dt -import json -import os -import re -import subprocess -from typing import Iterable - -GREPTILE_BOT_LOGINS = frozenset({"greptile-apps", "greptile-apps[bot]"}) - -SCORE_PATTERN = re.compile( - r"confidence\s*score\s*[:\-]?\s*(\d+)\s*/\s*5", - re.IGNORECASE, -) - -GRACE_COMMENT_MARKER = "" - -# Hidden HTML marker stamped on every Agent Shin auto-close comment (the LLM -# judge's grace/review-gate close and the daily Greptile sweep's close). -# `was_closed_by_agent_shin` requires this marker — not just the closing actor — -# before `@agent-shin reconsider` may reopen, because the `github-actions[bot]` -# identity is shared with every other workflow in the repo and is not unique to -# Agent Shin. Both close paths must stamp it or the reconsider path silently -# rejects the contributor. -AGENT_SHIN_CLOSE_MARKER = "" - -# 2 hours between the grace warning and the auto-close. Short enough to -# dogfood the "fix it before it closes" loop in one sitting; bump back up -# (e.g. 86400 for a day) for the public rollout. -GRACE_PERIOD_SECONDS = 7200 - -AGENT_SHIN_DEFAULT_BOT_LOGIN = "github-actions[bot]" - - -def _logins(*names: str) -> frozenset[str]: - """Build a login set normalized for case-insensitive membership checks. - - Callers compare via ``login.lower() in ``, so the stored values - must be lowercase. Normalizing here lets the literals keep each - account's canonical GitHub casing (e.g. ``SwiftWinds``) for - readability without breaking the lookup. - """ - return frozenset(name.lower() for name in names) - - -# Dogfood rollout gate. While this set is non-empty, Agent Shin acts ONLY on -# PRs/issues authored by these logins and skips everyone else. For an -# allowlisted author the usual internal/external classification is bypassed, so -# an internal account (e.g. a maintainer's own work login) still gets triaged -# while the bot is being tested on a small set of accounts. Empty the set to -# lift the restriction and restore full triage for the public rollout. Logins -# are compared case-insensitively. -ALLOWLIST_LOGINS = _logins("mateo-berri", "SwiftWinds") - -# `gh {pr,issue} list` has no "fetch everything" flag — `--limit` is the only -# control and it defaults to 30. Pass a ceiling far above any realistic open -# backlog (low thousands today) so gh paginates the API until the queue is -# exhausted rather than silently truncating. The bulk sweeps MUST see the whole -# backlog: gh lists newest-first, so a low cap drops the *oldest* PRs/issues — -# exactly the stale ones a low-quality sweep is meant to catch. -GH_LIST_ALL_LIMIT = 100_000 - - -def extract_greptile_score(comments: Iterable[dict]) -> tuple[int, dict] | None: - """Return (score, comment) for the most recent Greptile-authored comment - that contains a "Confidence Score: X/5". Returns None if no such comment. - - "Most recent" is determined by the comment's `updated_at` (falling back to - `created_at`), so re-reviews override earlier passes. - """ - candidates: list[tuple[str, int, dict]] = [] - for comment in comments: - user = (comment.get("user") or {}).get("login", "") - if user not in GREPTILE_BOT_LOGINS: - continue - body = comment.get("body") or "" - match = SCORE_PATTERN.search(body) - if not match: - continue - score = int(match.group(1)) - timestamp = comment.get("updated_at") or comment.get("created_at") or "" - candidates.append((timestamp, score, comment)) - - if not candidates: - return None - - candidates.sort(key=lambda triple: triple[0]) - _, score, comment = candidates[-1] - return score, comment - - -def parse_iso8601(value: str) -> dt.datetime: - """Parse a GitHub ISO-8601 timestamp into a timezone-aware datetime.""" - return dt.datetime.fromisoformat(value.replace("Z", "+00:00")) - - -def gh(*args: str) -> str: - """Run a `gh` CLI command and return stdout. Raises on non-zero exit. - - Shared by both Agent Shin entrypoints so a future change here - (timeout handling, logging, retry on transient failures) only needs - to be made once. - """ - result = subprocess.run( - ["gh", *args], - capture_output=True, - text=True, - check=True, - ) - return result.stdout - - -def list_open_items(kind: str, *, repo: str | None, fields: str) -> list[dict]: - """Return EVERY open PR (``kind="pr"``) or issue (``kind="issue"``) in ``repo``. - - Wraps ``gh {pr,issue} list`` with ``--limit GH_LIST_ALL_LIMIT`` so the full - backlog is fetched instead of the default 30 (or any other arbitrary cap). - Both bulk sweeps — the daily Greptile closer and the one-shot rollout - heads-up — rely on this seeing the whole queue, including the oldest items. - - ``fields`` is the comma-separated ``--json`` field list the caller needs - (e.g. ``"number"`` for the rollout, the full set for the closer). - """ - if kind not in ("pr", "issue"): - raise ValueError(f"kind must be 'pr' or 'issue', got {kind!r}") - repo_args = ["--repo", repo] if repo else [] - raw = gh( - kind, - "list", - "--state", - "open", - "--limit", - str(GH_LIST_ALL_LIMIT), - "--json", - fields, - *repo_args, - ) - return json.loads(raw) - - -def seconds_since_latest_marker_comment( - comments: Iterable[dict], - *, - marker: str, - bot_login: str | None = None, - now: dt.datetime | None = None, -) -> float | None: - """Return seconds since the bot's most recent comment containing ``marker``. - - Filters comments by author so a contributor who quotes the HTML - marker (e.g. via GitHub's "Quote reply" feature, which preserves - HTML comments in the raw markdown of the quoted text) is not - mistaken for a bot warning — that would silently reset cooldown - timers and suppress legitimate notifications. - - ``bot_login`` defaults to the `AGENT_SHIN_BOT_LOGIN` env override or - ``AGENT_SHIN_DEFAULT_BOT_LOGIN`` so callers normally don't need to - pass it. ``now`` is injectable for tests / callers (like the daily - sweep) that want every age calculation pinned to one snapshot. - """ - expected_login = ( - bot_login - or os.environ.get("AGENT_SHIN_BOT_LOGIN") - or AGENT_SHIN_DEFAULT_BOT_LOGIN - ).lower() - latest: dt.datetime | None = None - for comment in comments: - author = ((comment.get("user") or {}).get("login") or "").lower() - if author != expected_login: - continue - body = comment.get("body") or "" - if marker not in body: - continue - created = comment.get("created_at") - if not created: - continue - try: - ts = parse_iso8601(created) - except ValueError: - continue - if latest is None or ts > latest: - latest = ts - if latest is None: - return None - reference = now if now is not None else dt.datetime.now(dt.timezone.utc) - return (reference - latest).total_seconds() diff --git a/.github/scripts/close_low_quality_prs.py b/.github/scripts/close_low_quality_prs.py deleted file mode 100644 index 7b9bbb579e3..00000000000 --- a/.github/scripts/close_low_quality_prs.py +++ /dev/null @@ -1,573 +0,0 @@ -#!/usr/bin/env python3 -""" -Auto-close low-quality pull requests. - -Closes open PRs (including drafts, regardless of age) that satisfy ALL of: - 1. Have a Greptile (`greptile-apps`) review comment whose latest - "Confidence Score: X/5" is below the configured threshold (default: 4). - 2. Are authored by an external OSS contributor (internal BerriAI - contributors are exempt). - 3. Do not carry an opt-out label (default: "do not close"). - -`--min-age-days` is retained as an opt-in safety net for one-off backfill -runs (default: 0). The team's intent is that the count of open PRs equals -the count of PRs internal collaborators need to action on, so neither age -nor draft status acts as a free pass. - -For each match, the script posts an explanatory comment and closes the PR. -Because OSS contributors *cannot* reopen a PR closed by the bot/maintainer -(GitHub limitation), the close-comment instructs them to push their fixes -and **open a fresh PR**, or to comment `@agent-shin reconsider` on the -closed PR to have the LLM judge re-evaluate (and reopen on pass). - -Requires the `gh` CLI to be authenticated. - -Usage examples: - # Dry run (default) - prints what would be closed - python3 close_low_quality_prs.py - - # Actually close matching PRs - python3 close_low_quality_prs.py --close - - # Restrict to PRs at least N days old (one-off backfill safety net) - python3 close_low_quality_prs.py --min-age-days 7 --min-score 4 --close -""" - -from __future__ import annotations - -import argparse -import datetime as dt -import json -import os -import subprocess -import sys -from typing import Iterable - -# Add this script's directory to `sys.path` so the sibling -# `agent_shin_shared` module is importable when the script is invoked -# directly (e.g. `python3 .github/scripts/close_low_quality_prs.py ...`). -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - -from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above - AGENT_SHIN_CLOSE_MARKER, - ALLOWLIST_LOGINS, - GRACE_COMMENT_MARKER, - GRACE_PERIOD_SECONDS, - GREPTILE_BOT_LOGINS, - SCORE_PATTERN, - extract_greptile_score, - gh, - list_open_items, - parse_iso8601, - seconds_since_latest_marker_comment, -) - -# `GREPTILE_BOT_LOGINS` and `SCORE_PATTERN` (Greptile's GitHub App login -# variants and the "Confidence Score: X/5" regex) are imported from -# `agent_shin_shared` so the LLM judge in `triage_with_llm.py` and this -# daily Greptile sweep read the score through the same set of logins -# and the same regex. - -# `author_association` values for internal BerriAI contributors who should be -# exempt from auto-triage. -INTERNAL_AUTHOR_ASSOCIATIONS = frozenset({"OWNER", "MEMBER", "COLLABORATOR"}) - -# Default labels that exempt a PR from auto-close. Defined at module scope (not -# as a mutable argparse default) so that `--optout-label foo` REPLACES the -# defaults instead of appending to them — the argparse `action="append"` + -# `default=[...]` combination silently mutates the shared default list. -DEFAULT_OPTOUT_LABELS = ("do not close", "keep open", "wip") - -# `GRACE_COMMENT_MARKER` (HTML marker appended to grace-period warning -# comments — used by either script to recognize that a warning was -# already posted) and `GRACE_PERIOD_SECONDS` (length of the grace -# period between the warning and the actual auto-close, 2 hours) are -# imported from `agent_shin_shared` so the Agent Shin LLM judge and -# this daily Greptile sweep agree on the same marker and duration. - - -def fetch_open_prs(repo: str | None) -> list[dict]: - """Fetch all open PRs (number, createdAt, isDraft, labels, author). - - Includes drafts: `gh pr list --state open` returns both ready-for-review - and draft PRs by default. This is the desired behavior — drafts are not - a free pass; the internal-collaborator open-PR queue should reflect every - PR that needs human attention regardless of draft status. - """ - fields = "number,title,createdAt,isDraft,labels,author,url" - return list_open_items("pr", repo=repo, fields=fields) - - -def fetch_pr_author_association(pr_number: int, repo: str | None) -> str: - """Return the GitHub `author_association` for a PR, uppercase. - - Values: OWNER, MEMBER, COLLABORATOR, CONTRIBUTOR, FIRST_TIME_CONTRIBUTOR, - FIRST_TIMER, MANNEQUIN, NONE. Returns "" on lookup failure. - """ - endpoint = ( - f"repos/{repo}/pulls/{pr_number}" - if repo - else f"repos/{{owner}}/{{repo}}/pulls/{pr_number}" - ) - try: - data = json.loads(gh("api", endpoint)) - except subprocess.CalledProcessError: - return "" - return (data.get("author_association") or "").upper() - - -def is_external_pr_author(pr: dict, repo: str | None) -> bool: - """Return True if the PR author is an external OSS contributor. - - Internal = `OWNER` / `MEMBER` / `COLLABORATOR` association, or a bot login. - """ - login = ((pr.get("author") or {}).get("login") or "").lower() - if login.endswith("[bot]") or login in {"dependabot", "github-actions"}: - return False - association = fetch_pr_author_association(pr["number"], repo) - # Fail-safe: if the API lookup failed (empty string), treat the author as - # internal so we don't auto-close their PR. Auto-close is destructive, so - # an unknown association should never make a PR eligible for closing. - if not association or association in INTERNAL_AUTHOR_ASSOCIATIONS: - return False - return True - - -def fetch_pr_comments(pr_number: int, repo: str | None) -> list[dict]: - """Fetch issue-level comments on a PR (where Greptile posts its summary).""" - endpoint = ( - f"repos/{repo}/issues/{pr_number}/comments?per_page=100" - if repo - else f"repos/{{owner}}/{{repo}}/issues/{pr_number}/comments?per_page=100" - ) - raw = gh("api", "--paginate", endpoint) - comments: list[dict] = [] - for line in raw.strip().splitlines(): - line = line.strip() - if not line: - continue - try: - parsed = json.loads(line) - except json.JSONDecodeError: - # A malformed line should not blow up the whole sweep. Skip and - # carry on so the remaining PRs in this run still get evaluated. - continue - if isinstance(parsed, list): - comments.extend(parsed) - else: - comments.append(parsed) - return comments - - -def has_optout_label(pr: dict, optout_labels: set[str]) -> bool: - labels = {label.get("name", "").lower() for label in pr.get("labels", [])} - return bool(labels & {lbl.lower() for lbl in optout_labels}) - - -def seconds_since_last_grace_warning( - comments: Iterable[dict], - *, - bot_login: str | None = None, - now: dt.datetime | None = None, -) -> float | None: - """Return seconds since the bot's most recent grace-period warning, or - None if no such warning has ever been posted on this PR. - - Thin wrapper over - `agent_shin_shared.seconds_since_latest_marker_comment` — the - centralized helper handles the bot-author filter, marker match, - timestamp parsing, and `now` injection. Keeping this wrapper - preserves the closer's "already-fetched comments + injectable now" - interface so callers (and tests) don't need to change. - """ - return seconds_since_latest_marker_comment( - comments, - marker=GRACE_COMMENT_MARKER, - bot_login=bot_login, - now=now, - ) - - -def format_grace_warning_comment(score: int, threshold: int) -> str: - """Comment posted on the FIRST low-Greptile-score detection — gives - the contributor a 2-hour grace window before the auto-close fires on - the next daily cron run. - - Mirrors `format_grace_warning_pr_comment` in - `triage_with_llm.py` in spirit (2-hour grace + escape hatches), but - framed around Greptile's confidence score instead of the LLM judge's - rubric since the close trigger here is the Greptile signal. - """ - return ( - "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " - "repository.\n" - "\n" - "Heads up: Greptile's most recent review scored this PR " - f"**{score}/5**, below our merge bar of **{threshold}/5**.\n" - "\n" - "If the score isn't lifted in the next **2 hours**, I'll auto-close this PR. That's " - "**not** us saying the change isn't worthwhile. We want the open-PR list to mirror " - "what a maintainer can act on *right now*, so contributors like you don't get lost in " - "a backlog. Take your time; everything below still works after the close.\n" - "\n" - "**During the grace period:** push fixes that address Greptile's feedback, then comment " - "`@greptileai` to request a fresh review. If " - f"the new score is **{threshold}/5 or higher**, the PR stays open and no further " - "action is needed on your side.\n" - "\n" - "**If the PR does get auto-closed in 2 hours, you still have an easy recovery path:**\n" - "\n" - "- Comment `@greptileai` to request a fresh review. **This still works even after " - f"the PR is closed**, and a score of {threshold}/5 or higher is one of the signals " - "that lifts the PR back into the review queue. A low Greptile score isn't a blocker.\n" - "- Comment `@agent-shin reconsider` after pushing fixes; I'll re-run the rubric and " - "reopen the PR if both gates (description rubric + Greptile score) now pass.\n" - "\n" - f"{GRACE_COMMENT_MARKER}" - ) - - -def post_grace_warning( - pr: dict, - score: int, - threshold: int, - repo: str | None, - dry_run: bool, -) -> None: - """Post the 2-hour grace-period warning comment on `pr`. - - The warning carries `GRACE_COMMENT_MARKER` so subsequent runs can - detect that the contributor has already been told about the - pending close. Does NOT close the PR — the close happens on the - next eligible run after `GRACE_PERIOD_SECONDS` elapses (handled - by `close_pr`). - """ - pr_number = pr["number"] - repo_args = ["--repo", repo] if repo else [] - - if dry_run: - print( - f" [DRY RUN] Would post grace warning to PR #{pr_number} " - f"(greptile={score}/5): {pr['title']}" - ) - return - - comment_body = format_grace_warning_comment(score, threshold) - gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args) - print(f" Posted grace warning on PR #{pr_number} (greptile={score}/5)") - - -def format_close_comment(score: int, threshold: int) -> str: - """Comment posted when a low-Greptile-score PR is auto-closed. - - Carries `AGENT_SHIN_CLOSE_MARKER` so the `@agent-shin reconsider` path - (guarded by `was_closed_by_agent_shin`) recognizes this as an Agent Shin - close and is allowed to reopen the PR once it passes again; without the - marker that recovery path the comment advertises silently rejects the - contributor. - """ - score_sentence = ( - f"Greptile's most recent review scored this PR **{score}/5**, below " - f"our merge bar of **{threshold}/5**, and the 2-hour grace period since " - "the warning has elapsed.\n\n" - ) - return ( - f"Closing as part of automated PR triage.\n\n" - f"{score_sentence}" - "We close low-confidence PRs aggressively to keep the review queue " - "manageable for maintainers and contributors alike. **This is not a " - "rejection of the idea.** To bring this back:\n\n" - "1. Push the fixes that address Greptile's feedback (continue using " - "your existing branch is fine).\n" - "2. **Open a new PR** with the updated branch. Greptile will review " - "it again, and if it scores " - f"**{threshold}/5 or higher** a maintainer will take another look.\n\n" - "_Why open a new PR instead of reopening this one?_ GitHub does not " - "let external contributors reopen a PR that was closed by a bot or " - "maintainer, so a fresh PR is the most reliable path forward. If you " - "would prefer this exact PR re-evaluated, comment " - "`@agent-shin reconsider` once you've pushed the fixes; Agent Shin " - "will re-run triage and reopen this PR if it now meets the bar. " - "You can also comment `@greptileai` to request a fresh Greptile " - "review; that works **even after the PR is closed**.\n\n" - "Thanks for contributing to LiteLLM. We know auto-closures can sting; " - "the goal is to keep the project healthy, not to dismiss your work." - f"\n\n{AGENT_SHIN_CLOSE_MARKER}" - ) - - -def close_pr( - pr: dict, - score: int, - threshold: int, - age_days: int, - repo: str | None, - dry_run: bool, - label: str | None, -) -> None: - """Post the explanatory comment and close the PR.""" - pr_number = pr["number"] - repo_args = ["--repo", repo] if repo else [] - - if dry_run: - print( - f" [DRY RUN] Would close PR #{pr_number} " - f"(age={age_days}d, greptile={score}/5): {pr['title']}" - ) - return - - comment_body = format_close_comment(score, threshold) - gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args) - - if label: - try: - gh("pr", "edit", str(pr_number), "--add-label", label, *repo_args) - except subprocess.CalledProcessError as exc: - stderr = (exc.stderr or "").strip() - print(f" warn: failed to add label '{label}' to #{pr_number}: {stderr}") - - gh("pr", "close", str(pr_number), *repo_args) - print(f" Closed PR #{pr_number} (greptile={score}/5, age={age_days}d)") - - -def evaluate_pr( - pr: dict, - now: dt.datetime, - min_age_days: int, - min_score: int, - repo: str | None, - optout_labels: set[str], - allowlist: frozenset[str] = ALLOWLIST_LOGINS, -) -> tuple[str, int | None, int | None]: - """Decide what to do with `pr` on this triage run. - - Returns (action, score_or_none, age_days_or_none) where action is one of: - "skip-too-young", "skip-optout-label", "skip-not-allowlisted", - "skip-internal", "skip-no-greptile-score", "skip-score-ok", - "warn-grace", "skip-in-grace-period", or "close". - - Drafts are NOT skipped — the goal is "open PR count == PRs internal - collaborators need to action on", and a draft that Greptile scored <4/5 - is still in that queue. Authors can opt out via the `wip` label (see - `DEFAULT_OPTOUT_LABELS`) if they need to keep a long-lived draft open. - - Grace-period semantics: the first time a PR fails the rubric, the - action is `warn-grace` — the caller should post a warning comment but - NOT close the PR. On a subsequent run, if the warning is still less - than `GRACE_PERIOD_SECONDS` old AND the PR still fails, the action is - `skip-in-grace-period`. Once the warning ages out and the rubric is - still failing, the action is `close`. - """ - if has_optout_label(pr, optout_labels): - return ("skip-optout-label", None, None) - - created = parse_iso8601(pr["createdAt"]) - age_days = (now - created).days - # `min_age_days` defaults to 0 (close as soon as Greptile scores low). - # Set a positive value via --min-age-days for one-off backfill runs that - # want to skip very-young PRs. - if min_age_days > 0 and age_days < min_age_days: - return ("skip-too-young", None, age_days) - - # While the allowlist is active it is the sole author gate: only those - # logins are acted on and the external-only restriction is bypassed for - # them. Otherwise auto-close only external OSS contributors — internal - # contributors (BerriAI org members) handle their own backlog. - login = ((pr.get("author") or {}).get("login") or "").lower() - if allowlist: - if login not in allowlist: - return ("skip-not-allowlisted", None, age_days) - elif not is_external_pr_author(pr, repo): - return ("skip-internal", None, age_days) - - comments = fetch_pr_comments(pr["number"], repo) - extraction = extract_greptile_score(comments) - if extraction is None: - return ("skip-no-greptile-score", None, age_days) - - score, _ = extraction - if score >= min_score: - return ("skip-score-ok", score, age_days) - - grace_age = seconds_since_last_grace_warning(comments, now=now) - if grace_age is None: - return ("warn-grace", score, age_days) - if grace_age < GRACE_PERIOD_SECONDS: - return ("skip-in-grace-period", score, age_days) - - return ("close", score, age_days) - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument( - "--repo", - type=str, - default=None, - help="Repository (owner/repo). Auto-detected if omitted.", - ) - parser.add_argument( - "--min-age-days", - type=int, - default=0, - help=( - "Minimum age (in days) before a PR is eligible. Default 0 = " - "close as soon as Greptile flags it. Set a positive value for " - "one-off backfill runs that want to spare very-young PRs." - ), - ) - parser.add_argument( - "--min-score", - type=int, - default=4, - choices=range(1, 6), - help="Greptile score below which a PR is closed (default: 4 -> closes <4/5).", - ) - parser.add_argument( - "--optout-label", - action="append", - default=None, - help=( - "Label(s) that exempt a PR from auto-close. Repeat to add more. " - "Case-insensitive. When omitted, defaults to " - f"{list(DEFAULT_OPTOUT_LABELS)!r}; passing this flag REPLACES the " - "defaults (argparse `append` with a mutable default would append " - "instead, which we explicitly avoid)." - ), - ) - parser.add_argument( - "--close-label", - type=str, - default=None, - help=( - "Optional label to add to PRs that get auto-closed " - "(e.g. 'auto-closed-low-quality'). Must already exist on the repo." - ), - ) - parser.add_argument( - "--close", - action="store_true", - help="Actually close matching PRs (default is dry-run).", - ) - parser.add_argument( - "--limit", - type=int, - default=None, - help="Maximum number of PRs to close in one run (safety net).", - ) - args = parser.parse_args() - - dry_run = not args.close - if dry_run: - print("=== DRY RUN MODE (pass --close to actually close PRs) ===\n") - - print("Fetching open PRs...") - prs = fetch_open_prs(args.repo) - print(f"Found {len(prs)} open PRs.\n") - - now = dt.datetime.now(dt.timezone.utc) - optout_labels = set(args.optout_label or DEFAULT_OPTOUT_LABELS) - - closed = 0 - summary = { - "close": 0, - "warn-grace": 0, - "skip-in-grace-period": 0, - "skip-too-young": 0, - "skip-optout-label": 0, - "skip-not-allowlisted": 0, - "skip-internal": 0, - "skip-no-greptile-score": 0, - "skip-score-ok": 0, - } - - # `warned` tracks grace-warning comments posted in this run so the - # `--limit` safety net bounds *all* destructive write actions, not - # just closures. Without this cap, a backlog of PRs failing the - # threshold simultaneously could flood contributors with comments. - warned = 0 - for pr in sorted(prs, key=lambda p: p["createdAt"]): - try: - action, score, age_days = evaluate_pr( - pr, - now, - args.min_age_days, - args.min_score, - args.repo, - optout_labels, - ) - summary[action] = summary.get(action, 0) + 1 - - if action == "warn-grace": - assert score is not None - print( - f"#{pr['number']}: \"{pr['title']}\" " - f"(age={age_days}d, greptile={score}/5) -> warn-grace" - ) - post_grace_warning( - pr, - score=score, - threshold=args.min_score, - repo=args.repo, - dry_run=dry_run, - ) - if not dry_run: - warned += 1 - if args.limit is not None and (warned + closed) >= args.limit: - print( - f"\nReached --limit={args.limit} " - f"(closed={closed}, warned={warned}); stopping." - ) - break - continue - - if action != "close": - continue - - assert score is not None and age_days is not None - print( - f"#{pr['number']}: \"{pr['title']}\" " - f"(age={age_days}d, greptile={score}/5) -> close" - ) - close_pr( - pr, - score=score, - threshold=args.min_score, - age_days=age_days, - repo=args.repo, - dry_run=dry_run, - label=args.close_label, - ) - - if not dry_run: - closed += 1 - if args.limit is not None and (warned + closed) >= args.limit: - print( - f"\nReached --limit={args.limit} " - f"(closed={closed}, warned={warned}); stopping." - ) - break - except Exception as exc: # noqa: BLE001 - per-PR errors don't abort the sweep - summary["error"] = summary.get("error", 0) + 1 - print( - f"!! PR #{pr.get('number')}: {exc}", - file=sys.stderr, - ) - continue - - print("\n=== Summary ===") - for key, value in summary.items(): - print(f" {key:28s} {value}") - if dry_run: - print(f"\nTotal would close: {summary['close']}") - else: - print(f"\nTotal closed: {closed}") - print( - f"Total {'would warn (grace)' if dry_run else 'warned (grace)'}: " - f"{summary['warn-grace']}" - ) - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/.github/scripts/triage-requirements.txt b/.github/scripts/triage-requirements.txt deleted file mode 100644 index a18f05fbb95..00000000000 --- a/.github/scripts/triage-requirements.txt +++ /dev/null @@ -1,282 +0,0 @@ -# Hash-pinned dependency set for the Agent Shin triage scripts. -# Installed in privileged triage workflows, so every package is pinned to an -# exact version with SHA-256 hashes and installed with pip --require-hashes. -# -# Regenerate after bumping openai: -# echo 'openai==' \ -# | uv pip compile - --generate-hashes --python-version 3.12 \ -# --no-annotate --no-header -o .github/scripts/triage-requirements.txt - -annotated-types==0.7.0 \ - --hash=sha256:1f02e8b43a8fbbc3f3e0d4f0f4bfc8131bcb4eebe8849b8e5c773f3a1c582a53 \ - --hash=sha256:aff07c09a53a08bc8cfccb9c85b05f1aa9a2a6f23728d790723543408344ce89 -anyio==4.14.0 \ - --hash=sha256:b47c1f9ccf73e67021df785332508f99379c68fa7d0684e8e3492cb1d4b23f89 \ - --hash=sha256:dd9b7a2a9799ed6552fde617b2c5df02b7fdd7d88392fc48101e51bae46164d9 -certifi==2026.6.17 \ - --hash=sha256:024c88eeec92ca068db80f02b8b07c9cef7b9fe261d1d535abfd5abd6f6af432 \ - --hash=sha256:2227dcbaafe0d2f59279d1762ddddc37783ed4354594f194ffc31d20f41fc3db -distro==1.9.0 \ - --hash=sha256:2fa77c6fd8940f116ee1d6b94a2f90b13b5ea8d019b98bc8bafdcabcdd9bdbed \ - --hash=sha256:7bffd925d65168f85027d8da9af6bddab658135b840670a223589bc0c8ef02b2 -h11==0.16.0 \ - --hash=sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1 \ - --hash=sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86 -httpcore==1.0.9 \ - --hash=sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55 \ - --hash=sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8 -httpx==0.28.1 \ - --hash=sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc \ - --hash=sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad -idna==3.18 \ - --hash=sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2 \ - --hash=sha256:ffb385a7e039654cef1ab9ef32c6fafe283c0c0467bba1d9029738ce4a14a848 -jiter==0.15.0 \ - --hash=sha256:01a8222cf05ab1128e239421156c207949808acaaea2bdfd33130ae666786e86 \ - --hash=sha256:032396229564bca02440396bd327710719f724f5e7b7e9f7a8eb3faa4a2c2281 \ - --hash=sha256:04b400bbf8c9efb03d9bdd976475c919c1d85593b04b9fff7ae234065daf87ae \ - --hash=sha256:05906b93d72f03339e6bb7cf8dc10ebda64a0266126eed6beba79e20abcf5fd4 \ - --hash=sha256:066f8f33f18b2419cd8213b2436fa7fbc9c499f315971cfa3ce1f9820c001b1b \ - --hash=sha256:0ab068bce62a45aa3e7367eceaffb5dde60b7eb853be8dece45132e3d0ff4879 \ - --hash=sha256:0be6f5ad41a809f303f416d17cec92a7a725902fb9b4f3de3d19362ac0ef8554 \ - --hash=sha256:0e90a1c315a0226ec822d973817967f9223b7701546c8c2a7913e7ab0926294d \ - --hash=sha256:0f862193b8696249d22ec433e85fd2ab0ad9596bc3e45e6c0bc55e8aeba97be2 \ - --hash=sha256:1303d4d68a9b051ea90502402063ecf3807da00ad2affa19ca1ae3b90b3c5f67 \ - --hash=sha256:144f8e72cb53dab146347b91cceac01f5481237f2b93b4a339a1ee8f8878b67c \ - --hash=sha256:182226cbc930c9fab81bc2e41a4da672f89539906dadb05e75670ac07b94f71f \ - --hash=sha256:1c11465f97e2abf45a014b83b730222f8f1c5335e802c7055a67d50de6f1f4e3 \ - --hash=sha256:1c15024a3d892223b18f597c86d59387249dc396590844ce6b9f6131d1093bae \ - --hash=sha256:1d54fb5b31dea401a41af3f8a7d2512e9b6a6a005491e6166c7e4ffab9639a9c \ - --hash=sha256:25ffbe229aa8cd98c28879d8aa1a6e34ae77992ab984a65fba800859dab16269 \ - --hash=sha256:2a77aadd57cac1682e4401a72724d2796d89a4ba129b1a5812aa94ee480826eb \ - --hash=sha256:2ae901f3a55bfafdde31d289590fa25e3245735a2b1e8c7cc15871710a002871 \ - --hash=sha256:2b0074e2f56eb2dacca1689760fd2852a068f85a0547a157b82cb4cafeb6768b \ - --hash=sha256:2c8aea7781d2a372227871de4e1a1332aa96f5a89fd76c5e835dafdbad102887 \ - --hash=sha256:2c9cb907439d20bd0c7d7565ca01ee52234203208433749bae5b516907526928 \ - --hash=sha256:2fb6a5d26af81fc0f00f9360a891e05cf755e149bba391c4d563adc54812973d \ - --hash=sha256:2fd73e3da91a0a722d67165e849ce2cdc10de0e0d48738c142be8c6c5f310f4c \ - --hash=sha256:30ce1a5d16b5641dc935d50ef775af6a0871e3d14ab05d6fc54dff371b78e558 \ - --hash=sha256:30ce785d2adb8e32c3f7741442370a74834ec4c01f3c48f0750227a0b4ef27d6 \ - --hash=sha256:30f2218e6a9e5c18bc10fe6d41ac189c442c88eacf11bad9f28ef95a9bef00e6 \ - --hash=sha256:351a341c2105aa430b7047e30f1bf7975f6313b00165d3fc07be2edaf741f279 \ - --hash=sha256:37a10c377ce3a4a85f4a67f28b7afe093154cde77eaf248a72e856aa08b4d865 \ - --hash=sha256:392b8ab019e5502d08aff85c6272209c24bc2cbe706ea82a56368f524236614a \ - --hash=sha256:3e4540b8e74e4268811ac05db226a6a128ff572e7e0ce3f1163b693cadb184cd \ - --hash=sha256:40b2c7e92c44a84d748d21706c68dc6ff8161d80b59c99d774721a0d2317d7c7 \ - --hash=sha256:411fa4dfa5a7ae3d11491027ffb9beadec3996010a986862db70d91abba1c750 \ - --hash=sha256:4251acc80e2b7c9b7b8823456ea0fceeb0734dac2df7636d3c711b38476b5a76 \ - --hash=sha256:42bfb257930800cf43e7c62c832402c704ab60797c992faf88d20e903eac8f32 \ - --hash=sha256:4363818355dbc70ae1a8e9eaba9de350d93ede4ff6992b8f8eb8cbb6e5122d42 \ - --hash=sha256:4ab395feec8d249ec4044e228e98a7033f043426a265df439dc3698823f0a4e4 \ - --hash=sha256:50164d7610c00e7cd913a873fce30b6beeebf4b37e53983e33f22de4c900f6b8 \ - --hash=sha256:50e51156192722a9c58db112837d3f8ef96fb3c5ecc14e95f409134b08b158ec \ - --hash=sha256:510c8b3c17a0ed9ac69850c0438dada3c9b82d9c4d589fcb62002a5a9cf3a866 \ - --hash=sha256:5157de9f76eb4bc5ea74a1219366a25f945ad305641d74e04f59c54087091aa9 \ - --hash=sha256:54d5d6090cdc1b7c9e780dfb04949a990adb1e301a2fc0bbcee7de4638d33f9a \ - --hash=sha256:553fcac2ef2cb990877f9fc0833b8b629a3e6a5670b6b5fd58219b41a653ddc4 \ - --hash=sha256:5607e6013ed7e6b0ec9661e467b7ffde0aa7ab36833a04850f26fcf88ed4845b \ - --hash=sha256:5d6a60072b44c3c2b797a7ddcbcbbf2b34ea3cfd4721580fbfd2a09d9d9b84ba \ - --hash=sha256:5f30bae8bc1c2d613e28e5af3e8cceb09b742f1c8a8a5f839fb67afaffc03b61 \ - --hash=sha256:62ebd14e47e9aed9df4472afcb2663668ce4d74891cd54f86bf6e44029d6dc89 \ - --hash=sha256:631f13a3d04e97d4e083993b10f4b99530e3a10d953e2eb5e196b7dc7f812ce0 \ - --hash=sha256:6550fa135c7deb8ead6af49ed7ff648532ea8334a1447fe34a36315ef79c5c29 \ - --hash=sha256:66b1880df2d01e206e8339769d1c7c1753bcb653efd6289e203f6f24ebada0c0 \ - --hash=sha256:6eac374c5c975709b69c10f09afd199df74150172156ad10c8d4fd785b7da995 \ - --hash=sha256:71683c38c825452999b5717fcae07ea708e8c93003e808be4319c1b02e3d176e \ - --hash=sha256:7553333dd0930c104a5a0db8df72bf7219fe663d731383b576bb6ed6351c984d \ - --hash=sha256:75e8a04e91432dde9f1838373cf93d23726c79d3e908d319acf0e796f85592e7 \ - --hash=sha256:773b6eb282ce11ee19f05f6b2d4404fa308e5bbd353b0b80a0262caad6db2cd7 \ - --hash=sha256:774f93f65031856bf14ad9f59bdcab8b8cad501e5ceabd51ba3525f76937a25b \ - --hash=sha256:7c468136b8bd6bb18c8786e4236a1fa27362f24cb23450ba0cb204ab379b8e6f \ - --hash=sha256:7ce8902f939970048b233087082e7bb829db29375811c7ad50687b8624c6fd08 \ - --hash=sha256:7d3d6683288c11cbab50e865f2e2f13950179aa45410e30b2cfbd3fb7b0177bf \ - --hash=sha256:7f6163c0f10b055245f814dcc59f4818da60dfe72f3e72ab89fc24b6bd5e9c52 \ - --hash=sha256:8020c99ec13a7db2b6f96cbe82ef4721c88b426a4892f27478044af0284615ef \ - --hash=sha256:813dfbb17d65328bf86e5f0905dd277ba2265d3ca20556e86c0c7035b7182e5a \ - --hash=sha256:860a74063284a2ae9bfedd694f299cc2c68e2696c5f3d440cc9d18bb81b9dd04 \ - --hash=sha256:8c9004af7c8d67cce7f1aae1026fb55607f4aa600710d08ede3a3ce4aeefe7e0 \ - --hash=sha256:8d2c0c44d569ce0f2850f5c926f8caeb5f245fbc84475aeb36efccc2103e6dbd \ - --hash=sha256:8f7e9bc0f1135039b22ee6eab588d42df1ce55842b30740a352885eb267bd941 \ - --hash=sha256:90c5db5527c221249a876160663ab891ace358c17f7b9c93ec1478b7f0550e5c \ - --hash=sha256:9100ddbec09741cc66feb0fc6773f8bdbd0e3c345689368f260082ff85dcc0cd \ - --hash=sha256:913d02d29c9606643418d9ccfc3b72492ab25a6bf7889934e09a3490f8d3438b \ - --hash=sha256:980c256edb05b78a111b99c4de3b1d32e31634b867fd1fc2cf726e7b7bba9854 \ - --hash=sha256:9f924585cdacf631cd382b657966847bb537bf9ed0a6f9b991da5f05a631480f \ - --hash=sha256:a254e10b593624d230c365b6d616b22ca0ad65e63a16e6631c2b3466022e6ba8 \ - --hash=sha256:a2a438005b6f22d0273413484d6094d7c2c5d10ec1b3a3bf128e0d1d3ba53258 \ - --hash=sha256:a97261f1fccb8e50ecd2890a96e46efdc3f57c80a197324c6777827231eca712 \ - --hash=sha256:ab596fa3837e91e7e6a31b5f639988bfc6a35d1f915ac3932d946062219d588f \ - --hash=sha256:abbf258599526ad0326fe51e252e24f2bd6f24f1852681b4b78feda3808f1d18 \ - --hash=sha256:ac0d9ddea4350974be7a221fc25895f251a8fee748c889bdced2141c0fec1a49 \ - --hash=sha256:acf4ee4d1fc55917239fe72972fb292dd773055d05eb040d36f4326e02cc2c0e \ - --hash=sha256:ae1b0d82ac2d987f9ea512b1c9adfcc71a28de3dea3a6039b54d76cffda9901e \ - --hash=sha256:b15741f501469009ae0ae90b7147958a664a7dede40aa7ff174a8a4645f546d0 \ - --hash=sha256:b15d3ec9b0449c40e85319bdb4caa8b77ab526e74f5532ed94bec15e2f66822c \ - --hash=sha256:b3b3b775e33d3bfaec9899edc526ae97b0da0bf9d071a46124ba419149a414f8 \ - --hash=sha256:b6c0ffae686c39bf3737be60793783267628783ea42545632c10b291105aee45 \ - --hash=sha256:c210f8b35dc6f30aafd4b4365ca89b9d1189f21ab49b8e68fa6322a847aef138 \ - --hash=sha256:c2f6bb8b5216ab9e7873bc08b5d7bef2b8abbb578a3069bf1cd14a45d71d771d \ - --hash=sha256:c60e71b6d10cfc284c9bf36bd885e8d44c46f688ce50aa91b5edd90181dea687 \ - --hash=sha256:c6694a173ecabc12eb60efbc0b474464ead1951ff65cd8b1e72100715c64512b \ - --hash=sha256:c77496cb10bd7549690fbbab3e5ec05857b83e49276f4a9423a766ddd2afcd4c \ - --hash=sha256:c84c1b7be454b0c16f8499b4ebfbfd82ea5cca6527cceefcbbc06a7557b5ed2e \ - --hash=sha256:cc0bc345cf2df9d1c00ac443f50d543c1ccfa8b0422cb85b1ab70d681c0b255b \ - --hash=sha256:ceb8fc27d38793f9c97149be8302720c5b22e5c195a37bf2c45dc36c4600a512 \ - --hash=sha256:cf4bd113a69c0a740e27cb962ce10630c36d2b8f59d759a651b955ee9d18a823 \ - --hash=sha256:d1aa62e277fc1cbd80e6deacae6f4d983b41b3d7728e0645c5d741a6149bba45 \ - --hash=sha256:d1e7b1776f0797956c509e123d0952d10d293a9492dea9f288ab9570ec01d1a5 \ - --hash=sha256:d636d5095155afd364247f65070fab7beda13498d7ff4de331046e704ab9657f \ - --hash=sha256:d726e3ceeb337191324b49de298142f27c3ad10886341555d1d5315b5f252c6a \ - --hash=sha256:d72d8af5c1013656a8870c866660627d1a75bc185814ee022c8533caa1de88ae \ - --hash=sha256:d8d2955167274e15d79a7a020afdd9b39c990eb80b2d89fca695d92dcfdd38ec \ - --hash=sha256:d92a5cd21fdb083931d546c207aa29633787c5dc5b02daab2d32b843f88a2c53 \ - --hash=sha256:e58585a58209d72691ce2d62a9147445f5a87beb0bde97fde284c96ae392a3d1 \ - --hash=sha256:e7196e56f1cd69af1dbb07dff02dcfb260a50b45a82d409d92a06fedb32473b5 \ - --hash=sha256:eda3071db3346334beae1360b46da4606da57bf3528c167b3c38533afaf9f2c5 \ - --hash=sha256:edebcf7d1f601199084bb6e844d7dc67e03e04f6ac786b0332d616635c4ff7a4 \ - --hash=sha256:ef1fd24d9413f6209e00d3d5a453e67acfe004a25cc6c8e8484faed4311ab9e8 \ - --hash=sha256:f0b271b462769543716f92d3a4f90527df6ef5ed05ee95ec4137f513e21e1b77 \ - --hash=sha256:f18f85e4218d1b40f000f42a92239a7a61a902cd42c65e6c360dbd17dcb20894 \ - --hash=sha256:f1e1754960f38ec40613a07e5e372df67acb3b890fb383b6fb3de3e49ddbf3c7 \ - --hash=sha256:f2143ab06181d2b029eedcb6af3cebe95f11bbac62441781860f98ee9330a6a6 \ - --hash=sha256:f3d37768fce7f88dd2a8c6091f2325dea27d30d30d5c6e7a1c0f0af77723b708 \ - --hash=sha256:fa248c9eb220197d363f688818dac2fd4b2f0cd7d843ca7105d652034823427d -openai==2.33.0 \ - --hash=sha256:03ac37d70e8c9e3a8124214e3afa785e2cbc12e627fbd98177a086ef2fd87ad5 \ - --hash=sha256:f850c435e2a4685bba3295bd54912dd26315d9c1b7733068186134d6e0599f9a -pydantic==2.13.4 \ - --hash=sha256:45a282cde31d808236fd7ea9d919b128653c8b38b393d1c4ab335c62924d9aba \ - --hash=sha256:c40756b57adaa8b1efeeced5c196f3f3b7c435f90e84ea7f443901bec8099ef6 -pydantic-core==2.46.4 \ - --hash=sha256:00c603d540afdd6b80eb39f078f33ebd46211f02f33e34a32d9f053bba711de0 \ - --hash=sha256:0186750b482eefa11d7f435892b09c5c606193ef3375bcf94aa00ae6bfb66262 \ - --hash=sha256:041bde0a48fd37cf71cab1c9d56d3e8625a3793fef1f7dd232b3ff37e978ecda \ - --hash=sha256:0c563b08bca408dc7f65f700633d8442fffb2421fc47b8101377e9fd65051ff0 \ - --hash=sha256:0cbe8b01f948de4286c74cdd6c667aceb38f5c1e26f0693b3983d9d74887c65e \ - --hash=sha256:0ce40cd7b21210e99342afafbd4d0f76d784eb5b1d60f3bdc566be4983c6c73b \ - --hash=sha256:0e96592440881c74a213e5ad528e2b24d3d4f940de2766bed9010ab1d9e51594 \ - --hash=sha256:10e17cbb10a330363733efc4d7c4d0dd827ac0909b8f6a6542298fed1ea62f29 \ - --hash=sha256:133878133d271ade3d41d1bfb2a45ec38dbdbda40bc065921c6b04e4630127e2 \ - --hash=sha256:14d4edf427bdcf950a8a02d7cb44a08614388dd6e1bdcbf4f67504fa7887da9c \ - --hash=sha256:14f4c5d6db102bd796a627bbb3a17b4cf4574b9ae861d8b7c9a9661c6dd3362d \ - --hash=sha256:17299feefe090f2caa5b8e37222bb5f663e4935a8bfa6931d4102e5df1a9f398 \ - --hash=sha256:184c081504d17f1c1066e430e117142b2c77d9448a97f7b65c6ac9fd9aee238d \ - --hash=sha256:18e5ceec2ab67e6d5f1a9085e5a24c9c4e2ac4545730bfe668680bca05e555f3 \ - --hash=sha256:19e51f073cd3df251856a8a4189fbdf1de4012c3ebacfb1884f94f1eb406079f \ - --hash=sha256:1a7dd0b3ee80d90150e3495a3a13ac34dbcbfd4f012996a6a1d8900e91b5c0fb \ - --hash=sha256:1d8ba486450b14f3b1d63bc521d410ec7565e52f887b9fb671791886436a42f7 \ - --hash=sha256:2108ba5c1c1eca18030634489dc544844144ee36357f2f9f780b93e7ddbb44b5 \ - --hash=sha256:228ee9bae8bef5b1e97ec58302f80357c37199e0d0a99174e138d28e6957b9d9 \ - --hash=sha256:23ace664830ee0bfe014a0c7bc248b1f7f25ed7ad103852c317624a1083af462 \ - --hash=sha256:2412e734dcb48da14d4e4006b82b46b74f2518b8a26ee7e58c6844a6cd6d03c4 \ - --hash=sha256:29c61fc04a3d840155ff08e475a04809278972fe6aef51e2720554e96367e34b \ - --hash=sha256:2f84c03c8607173d16b5a854ec68a2f9079ae03237a54fb506d13af47e1d018d \ - --hash=sha256:3009f12e4e90b7f88b4f9adb1b0c4a3d58fe7820f3238c190047209d148026df \ - --hash=sha256:3245406455a5d98187ec35530fd772b1d799b26667980872c8d4614991e2c4a2 \ - --hash=sha256:3447661d99f75a3683a4cf5c87da72f2161964611864dbbeac7fbb118bb4bfc0 \ - --hash=sha256:372429a130e469c9cd698925ce5fc50940b7a1336b0d82038e63d5bbc4edc519 \ - --hash=sha256:395aebd9183f9d112f569aeb5b2214d1a10a33bec8456447f7fbdfa51d38d4cd \ - --hash=sha256:3a233125ac121aa3ffba9a2b59edfc4a985a76092dc8279586ab4b71390875e7 \ - --hash=sha256:3be77f45df024d789a672ae34f8b06fb346c4f9f46ea714956660ea4862e89ac \ - --hash=sha256:3bf92c5d0e00fefaab325a4d27828fe6b6e2a21848686b5b60d2d9eeb09d76c6 \ - --hash=sha256:3ecbc122d18468d06ca279dc26a8c2e2d5acb10943bb35e36ae92096dc3b5565 \ - --hash=sha256:3fb702cd90b0446a3a1c5e470bfa0dd23c0233b676a9099ddcc964fa6ca13898 \ - --hash=sha256:428e04521a40150c85216fc8b85e8d39fece235a9cf5e383761238c7fa9b96fb \ - --hash=sha256:432c179df7874eeb73307aad2df0755e1ae0efa61ff0ea89b93e194411ae3928 \ - --hash=sha256:4a05d69cba51d852c5c3e92758653245a50c0b646ced0cf05bd793ed592839d6 \ - --hash=sha256:4c63ebc82684aa89d9a3bcbd13d515b3be44250dc68dd3bd81526c1cb31286c3 \ - --hash=sha256:4fc73cb559bdb54b1134a706a2802a4cddd27a0633f5abb7e53056268751ac6a \ - --hash=sha256:4fcbe087dbc2068af7eda3aa87634eba216dbda64d1ae73c8684b621d33f6596 \ - --hash=sha256:56cb4851bcaf3d117eddcef4fe66afd750a50274b0da8e22be256d10e5611987 \ - --hash=sha256:5855698a4856556d86e8e6cd8434bc3ac0314ee8e12089ae0e143f64c6256e4e \ - --hash=sha256:5a4330cdbc57162e4b3aa303f588ba752257694c9c9be3e7ebb11b4aca659b5d \ - --hash=sha256:5b712b53160b79a5850310b912a5ef8e57e56947c8ad690c227f5c9d7e561712 \ - --hash=sha256:5d5902252db0d3cedf8d4a1bc68f70eeb430f7e4c7104c8c476753519b423008 \ - --hash=sha256:617d7e2ca7dcb8c5cf6bcb8c59b8832c94b36196bbf1cbd1bfb56ed341905edd \ - --hash=sha256:62f875393d7f270851f20523dd2e29f082bcc82292d66db2b64ea71f64b6e1c1 \ - --hash=sha256:633147d34cf4550417f12e2b1a0383973bdf5cdfde212cb09e9a581cf10820be \ - --hash=sha256:66ce7632c22d837c95301830e111ad0128a32b8207533b60896a96c4915192ea \ - --hash=sha256:6b3ace8194b0e5204818c92802dcdca7fc6d88aabbb799d7c795540d9cd6d292 \ - --hash=sha256:6f2eeda33a839975441c86a4119e1383c50b47faf0cbb5176985565c6bb02c33 \ - --hash=sha256:7027560ee92211647d0d34e3f7cd6f50da56399d26a9c8ad0da286d3869a53f3 \ - --hash=sha256:7283d57845ecf5a163403eb0702dfc220cc4fbdd18919cb5ccea4f95ee1cdab4 \ - --hash=sha256:7a5f930472650a82629163023e630d160863fce524c616f4e5186e5de9d9a49b \ - --hash=sha256:7bfb192b3f4b9e8a89b6277b6ce787564f62cfd272055f6e685726b111dc7826 \ - --hash=sha256:811ff8e9c313ab425368bcbb36e5c4ebd7108c2bbf4e4089cfbb0b01eff63fac \ - --hash=sha256:8233f2947cf85404441fd7e0085f53b10c93e0ee78611099b5c7237e36aacbf7 \ - --hash=sha256:82cf5301172168103724d49a1444d3378cb20cdee30b116a1bd6031236298a5d \ - --hash=sha256:8358a950c8909158e3df31538a7e4edc2d7265a7c54b47f0864d9e5bae9dcebf \ - --hash=sha256:85bb3611ff1802f3ee7fdd7dbff26b56f343fb432d57a4728fdd49b6ef35e2f4 \ - --hash=sha256:86e1a4418c6cd97d60c95c71164158eaf7324fae7b0923264016baa993eba6fc \ - --hash=sha256:8b9bab013d1c7a79d3501ff86d0bc9c31bf587db4551677b96bec07df78c6b15 \ - --hash=sha256:8c5dac79fa1614d1e06ca695109c6105923bd9c7d1d6c918d4e637b7e6b32fd3 \ - --hash=sha256:8d0820e8192167f80d88d64038e609c31452eeca865b4e1d9950a27a4609b00b \ - --hash=sha256:8daafc69c93ee8a0204506a3b6b30f586ef54028f52aeeeb5c4cfc5184fd5914 \ - --hash=sha256:9037063db01f09b09e237c282b6792bd4da634b5402c4e7f0c61effed7701a04 \ - --hash=sha256:905a0ed8ea6f2d61c1738835f99b699348d7857379083e5fc497fa0c967a407c \ - --hash=sha256:90884113d8b48f760e9587002789ddd741e76ab9f89518cd1e43b1f1a52ec44b \ - --hash=sha256:91a06d2e259ecfbd8c901d70c3c507900458498142b3026a296b7de4d1322cc9 \ - --hash=sha256:926c9541b14b12b1681dca8a0b75feb510b06c6341b70a8e500c2fdcff837cce \ - --hash=sha256:9401557acd873c3a7f3eb9383edef8ac4968f9510e340f4808d427e75667e7b4 \ - --hash=sha256:9551187363ffc0de2a00b2e47c25aeaeb1020b69b668762966df15fc5659dd5a \ - --hash=sha256:962ccbab7b642487b1d8b7df90ef677e03134cf1fd8880bf698649b22a69371f \ - --hash=sha256:97e7cf2be5c77b7d1a9713a05605d49460d02c6078d38d8bef3cbe323c548424 \ - --hash=sha256:9aa768456404a8bf48a4406685ac2bec8e72b62c69313734fa3b73cf33b3a894 \ - --hash=sha256:9bc519fbf2b7578398853d815009ae5e4d4603d12f4e3f91da8c06852d3da3e9 \ - --hash=sha256:9d56801be94b86a9da183e5f3766e6310752b99ff647e38b09a9500d88e46e76 \ - --hash=sha256:9f444c499b3eefd3a92e348059471ea0c3a6e303d9c1cec09fa748fd9f895201 \ - --hash=sha256:9fa8ae11da9e2b3126c6426f147e0fba88d96d65921799bb30c6abd1cb2c97fb \ - --hash=sha256:a0f62d0a58f4e7da165457e995725421e0064f2255d8eccebc49f41bbc23b109 \ - --hash=sha256:a396dcc17e5a0b164dbe026896245a4fa9ff402edca1dff0be3d53a517f74de4 \ - --hash=sha256:aaa2a54443eff1950ba5ddc6b6ccda0d9c84a364276a62f969bdf2a390650848 \ - --hash=sha256:ad785e92e6dc634c21555edc8bd6b64957ab844541bcb96a1366c202951ae526 \ - --hash=sha256:af8244b2bef6aaad6d92cda81372de7f8c8d36c9f0c3ea36e827c60e7d9467a0 \ - --hash=sha256:b078afbc25f3a1436c7a1d2cd3e322497ee99615ba97c563566fdf46aff1ee01 \ - --hash=sha256:b2f69dec1725e79a012d920df1707de5caf7ed5e08f3be4435e25803efc47458 \ - --hash=sha256:b8458003118a712e66286df6a707db01c52c0f52f7db8e4a38f0da1d3b94fc4e \ - --hash=sha256:bb63e0198ca18aad131c089b9204c23079c3afa95487e561f4c522d519e55aba \ - --hash=sha256:bfec22eab3c8cc2ceec0248aec886624116dc079afa027ecc8ad4a7e62010f8a \ - --hash=sha256:c1747f85cee84c26985853c6f3d9bd3e75da5212912443fa111c113b9c246f39 \ - --hash=sha256:c1b3f518abeca3aa13c712fd202306e145abf59a18b094a6bafb2d2bbf59192c \ - --hash=sha256:c50f2528cf200c5eed56faf3f4e22fcd5f38c157a8b78576e6ba3168ec35f000 \ - --hash=sha256:c68fcd102d71ea85c5b2dfac3f4f8476eff42a9e078fd5faefff6d145063536b \ - --hash=sha256:c7a7bd4e39e8e4c12c39cd480356842b6a8a06e41b23a55a5e3e191718838ddf \ - --hash=sha256:c94f0688e7b8d0a67abf40e57a7eaaecd17cc9586706a31b76c031f63df052b4 \ - --hash=sha256:cbaf13819775b7f769bf4a1f066cb6df7a28d4480081a589828ef190226881cd \ - --hash=sha256:cd2213145bcc2ba85884d0ac63d222fece9209678f77b9b4d76f054c561adb28 \ - --hash=sha256:ce5c1d2a8b27468f433ca974829c44060b8097eedc39933e3c206a90ee49c4a9 \ - --hash=sha256:d396ec2b979760aaf3218e76c24e65bd0aca24983298653b3a9d7a45f9e47b30 \ - --hash=sha256:d51026d73fcfd93610abc7b27789c26b313920fcfb20e27462d74a7f8b06e983 \ - --hash=sha256:d80ee3d731373b24cebbc10d689ca4ee1875caf0d5703a245db18efd4dd37fc1 \ - --hash=sha256:d995260fdf4e1db774581b4900e0f832abe3c7c84996726bbc161b19c8f29e76 \ - --hash=sha256:da4b951fe36dc7c3a1ccb4e3cd1747c3542b8c9ceede8fc86cae054e764485f5 \ - --hash=sha256:daa27d92c36f24388fe3ad306b174781c747627f134452e4f128ea00ce1fe8c4 \ - --hash=sha256:db06ffe51636ffe9ca531fe9023dd64bdd794be8754cb5df57c5498ae5b518a7 \ - --hash=sha256:e0d65b8c354be7fb5f720c3caa8bc940bc2d20ce749c8e06135f07f8ed95dd7c \ - --hash=sha256:e68b7a074f65a2fd746c52a7ce6142ab7006074ac269ace0c25cd8ba171f8066 \ - --hash=sha256:e739fee756ba1010f8bcccb534252e85a35fe45ae92c295a06059ce58b74ccd3 \ - --hash=sha256:e846ae7835bf0703ae43f534ab79a867146dadd59dc9ca5c8b53d5c8f7c9ef02 \ - --hash=sha256:e9c26f834c65f5752f3f06cb08cb86a913ceb7274d0db6e267808a708b46bc89 \ - --hash=sha256:ea793e075b70290d89d8142074262885d3f7da19634845135751bd6344f73b50 \ - --hash=sha256:f027324c56cd5406ca49c124b0db10e56c69064fec039acc571c29020cc87c76 \ - --hash=sha256:f13a646d65d09fbf1bc6b3a9635d30095c8e7e5cc419ff35ecc563c5fd04cd49 \ - --hash=sha256:f47286a97f0bc9b8859519809077b91b2cefe4ae47fcbf5e466a009c1c5d742b \ - --hash=sha256:f747929cf940cddb5b3668a390056ddd5ba2e5010615ea2dcf4f9c4f3ab8791d \ - --hash=sha256:f99626688942fb746e545232e7726926f3be91b5975f8b55327665fafda991c7 \ - --hash=sha256:f9fa868638bf362d3d138ea55829cefb3d5f4b0d7f142234382a15e2485dbec4 \ - --hash=sha256:fbdb89b3e1c94a30cc5edfce477c6e6a5dc4d8f84665b455c27582f211a1c72c \ - --hash=sha256:fc010ab034c8c7452522748bf937df58020d256ccae0874463d1f4d01758af8e \ - --hash=sha256:fc3e9034a63de20e15e8ade85358bc6efc614008cab72898b4b4952bea0509ff \ - --hash=sha256:fd8b3d9fd264be37976686c7f65cd52a83f5e84f4bfd2adf9c1d469676bbb6ae -sniffio==1.3.1 \ - --hash=sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2 \ - --hash=sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc -tqdm==4.68.3 \ - --hash=sha256:00dfa48452b6b6cfae3dd9885636c23d3422d1ec97c66d96818cbd5e0821d482 \ - --hash=sha256:39832cc2def2789a6f29df83f172db7416cea70052c0907a57801c5f2fdccb03 -typing-extensions==4.15.0 \ - --hash=sha256:0cea48d173cc12fa28ecabc3b837ea3cf6f38c6d1136f85cbaaf598984861466 \ - --hash=sha256:f0fa19c6845758ab08074a0cfa8b7aecb71c999ca73d62883bc25cc018c4e548 -typing-inspection==0.4.2 \ - --hash=sha256:4ed1cacbdc298c220f1bd249ed5287caa16f34d44ef4e9c3d0cbad5b521545e7 \ - --hash=sha256:ba561c48a67c5958007083d386c3295464928b01faa735ab8547c5692e87f464 diff --git a/.github/scripts/triage_with_llm.py b/.github/scripts/triage_with_llm.py deleted file mode 100644 index e23a012425a..00000000000 --- a/.github/scripts/triage_with_llm.py +++ /dev/null @@ -1,1797 +0,0 @@ -#!/usr/bin/env python3 -""" -Agent Shin — LLM-as-judge triage for external OSS pull requests and issues. - -Evaluates a single PR or issue against the contribution rubric and, when the -LLM judge marks it as failing, posts an explanatory comment + closes the -PR/issue. Re-triggers on `reopened` so contributors can iterate back in by -filling in the missing pieces and reopening. - -Internal BerriAI contributors (`author_association` in {OWNER, MEMBER, -COLLABORATOR}) and bot accounts are skipped entirely. - -Usage: - triage_with_llm.py --repo owner/repo --pr 1234 - triage_with_llm.py --repo owner/repo --issue 5678 - triage_with_llm.py --repo owner/repo --pr 1234 --close # actually close - triage_with_llm.py --repo owner/repo --pr 1234 --print-prompt # show prompt - -Defaults are SAFE: without `--close` the script writes a verdict to stdout (and, -when running in GitHub Actions, to $GITHUB_STEP_SUMMARY) but takes no GitHub -write actions. - -Environment: - GH_TOKEN / GITHUB_TOKEN - for `gh` CLI auth (auto-set in Actions) - OPENAI_API_KEY - required when --close is passed - OPENAI_BASE_URL - optional (route to any OpenAI-compatible API) - TRIAGE_MODEL - optional model override (default: gpt-5.4-mini) -""" - -from __future__ import annotations - -import argparse -import datetime as dt -import json -import os -import re -import subprocess -import sys -import textwrap -import urllib.parse -from typing import Any, Iterable - -# Add this script's directory to `sys.path` so the sibling -# `agent_shin_shared` module is importable when the script is invoked -# directly (e.g. `python3 .github/scripts/triage_with_llm.py ...`) and -# also when the tests load this script via -# `importlib.util.spec_from_file_location`. -sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) - -from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above - AGENT_SHIN_CLOSE_MARKER, - AGENT_SHIN_DEFAULT_BOT_LOGIN, - ALLOWLIST_LOGINS, - GRACE_COMMENT_MARKER, - GRACE_PERIOD_SECONDS, - GREPTILE_BOT_LOGINS, - SCORE_PATTERN, - extract_greptile_score, - gh, - parse_iso8601, - seconds_since_latest_marker_comment, -) - -DEFAULT_MODEL = "gpt-5.4-mini" - -INTERNAL_ASSOCIATIONS = frozenset({"OWNER", "MEMBER", "COLLABORATOR"}) - -# `AGENT_SHIN_DEFAULT_BOT_LOGIN` is imported from `agent_shin_shared`. -# When the workflow uses the default `secrets.GITHUB_TOKEN`, the -# closure / reopen event's `actor.login` is `github-actions[bot]`. The -# env override `AGENT_SHIN_BOT_LOGIN` exists for local debugging and for -# repos that wire Agent Shin to a PAT. - -# HTML marker appended to every reconsider verdict comment. We grep for this -# on subsequent reconsider triggers to enforce a short cooldown so that -# repeated `@agent-shin reconsider` comments don't burn CI/LLM budget. -# Using a unique HTML comment keeps the marker invisible to humans while -# being trivially greppable from a comments-list API response. -RECONSIDER_COMMENT_MARKER = "" - -# Minimum gap between two reconsider verdicts on the same PR/issue. Set to -# 10 minutes — long enough that a contributor can't trivially spam the -# trigger, short enough that a genuine "I just pushed a fix and reupdated -# the body" iteration loop isn't punished. -RECONSIDER_RATE_LIMIT_SECONDS = 600 - -# `GRACE_COMMENT_MARKER` (HTML marker on the grace-period warning comment -# posted on the first low-quality detection — used on subsequent triage -# runs to detect that a warning was already posted and measure how long -# ago it was posted) and `GRACE_PERIOD_SECONDS` (length of the grace -# period between the warning and the actual auto-close, 2 hours) are -# imported from `agent_shin_shared` so the daily Greptile sweep and the -# LLM judge agree on the same marker and duration. - -# --- Review-gate ("ready for review" label lifecycle) configuration ---------- -# The review gate keeps a single label in sync with whether a PR currently -# clears BOTH quality bars: the LLM rubric (clear problem + expected/actual + -# QA proof, or a linked issue) AND Greptile's most recent confidence score. -READY_FOR_REVIEW_LABEL = "ready for review" -DEFAULT_GRACE_DAYS = 1 # 24h before an un-passing, un-tagged PR is auto-closed -DEFAULT_MIN_GREPTILE_SCORE = 4 # Greptile < 4/5 counts as "not passing" - -# Hidden HTML-comment markers stamped into review-gate comments. They never -# render in the GitHub UI but let the gate detect its own prior actions so it -# (a) posts the within-grace "what's missing" notice at most once and (b) can -# tell a first-time pass ("ready for review") from a recovery after a -# regression ("all clear again"). -READY_MARKER = "" -REGRESSED_MARKER = "" -WITHIN_GRACE_MARKER = "" - -# `GREPTILE_BOT_LOGINS` (Greptile's GitHub App login variants — -# `greptile-apps[bot]` in REST API comments, `greptile-apps` in -# `gh pr view --json` output) and `SCORE_PATTERN` (regex matching lines -# like `Confidence Score: 3/5`) are imported from `agent_shin_shared` -# so the daily sweep and the review gate read the score through the -# same set of logins / patterns. - -# `AGENT_SHIN_CLOSE_MARKER` is imported from `agent_shin_shared` so this LLM -# judge and the daily Greptile sweep stamp the same marker on their close -# comments — `was_closed_by_agent_shin` keys the reconsider reopen path off it. - -# Model families that require `reasoning_effort` to be set, and that reject -# `temperature != 1` unless `reasoning_effort` is "none". For these models we -# pass `reasoning_effort="none"` so a `temperature=0` deterministic judgment -# is still accepted. See litellm/llms/openai/chat/gpt_5_transformation.py for -# the full set of constraints LiteLLM applies to these models. -GPT5_FAMILY_PREFIX = "gpt-5" - -# Regexes for picking off "obvious passes" without burning LLM tokens. -# -# Keep this list to GitHub's documented PR-closing keywords only -# (https://docs.github.com/issues/tracking-your-work-with-issues/linking-a-pull-request-to-an-issue). -# Casual mentions like "see #1234" or "ref #1234" are intentionally NOT -# auto-passed — they should fall through to the LLM judge, which has the -# stricter rubric "a bare issue number without a closing keyword counts only -# if it's clearly the related issue (not a passing mention)". -LINKED_ISSUE_PATTERN = re.compile( - r"\b(?:fixes|fix|fixed|closes|close|closed|resolves|resolve|resolved)\s+" - r"(?:#\d+|https?://github\.com/[\w.-]+/[\w.-]+/issues/\d+)", - re.IGNORECASE, -) -HTML_COMMENT_PATTERN = re.compile(r"", re.DOTALL) - - -# --------------------------------------------------------------------------- -# gh helpers -# -# `gh` is imported from `agent_shin_shared` so a future change (timeout, -# logging, retry) only needs to be made once. - - -def fetch_pr(repo: str, number: int) -> dict: - """Return the full GitHub REST representation of a PR.""" - return json.loads(gh("api", f"repos/{repo}/pulls/{number}")) - - -def fetch_issue(repo: str, number: int) -> dict: - """Return the full GitHub REST representation of an issue.""" - return json.loads(gh("api", f"repos/{repo}/issues/{number}")) - - -def post_comment(repo: str, number: int, body: str) -> None: - """Post an issue-style comment (works for both issues and PRs).""" - gh( - "api", - f"repos/{repo}/issues/{number}/comments", - "-X", - "POST", - "-f", - f"body={body}", - ) - - -def close_pr(repo: str, number: int) -> None: - """Close a pull request (state=closed).""" - gh( - "api", - f"repos/{repo}/pulls/{number}", - "-X", - "PATCH", - "-f", - "state=closed", - ) - - -def reopen_pr(repo: str, number: int) -> None: - """Reopen a previously-closed pull request (state=open). - - Used by the `@agent-shin reconsider` comment-trigger flow: the bot has - write access via GH_TOKEN, so it can reopen on the contributor's behalf - even though GitHub doesn't let the OSS author do it themselves. - """ - gh( - "api", - f"repos/{repo}/pulls/{number}", - "-X", - "PATCH", - "-f", - "state=open", - ) - - -def close_issue(repo: str, number: int, *, not_planned: bool = True) -> None: - """Close an issue, marking state_reason=not_planned by default.""" - args = [ - "api", - f"repos/{repo}/issues/{number}", - "-X", - "PATCH", - "-f", - "state=closed", - ] - if not_planned: - args.extend(["-f", "state_reason=not_planned"]) - gh(*args) - - -def reopen_issue(repo: str, number: int) -> None: - """Reopen a previously-closed issue (state=open, state_reason=reopened).""" - gh( - "api", - f"repos/{repo}/issues/{number}", - "-X", - "PATCH", - "-f", - "state=open", - "-f", - "state_reason=reopened", - ) - - -def add_label(repo: str, number: int, label: str) -> None: - """Add a label to a PR/issue (GitHub creates the label if it's missing).""" - gh( - "api", - f"repos/{repo}/issues/{number}/labels", - "-X", - "POST", - "-f", - f"labels[]={label}", - ) - - -def remove_label(repo: str, number: int, label: str) -> None: - """Remove a label from a PR/issue. A missing label (404) is not an error.""" - encoded = urllib.parse.quote(label, safe="") - try: - gh( - "api", - f"repos/{repo}/issues/{number}/labels/{encoded}", - "-X", - "DELETE", - ) - except subprocess.CalledProcessError as exc: - stderr = (exc.stderr or "").lower() - if "404" in stderr or "not found" in stderr: - return - raise - - -def _iter_paginated_json(*api_args: str) -> Any: - """Yield JSON objects from `gh api --paginate ... -q '.[]'`. - - `gh api --paginate` on a JSON-array endpoint concatenates pages into - one stream; `-q '.[]'` flattens that stream into newline-delimited - objects (jq-style). This keeps memory bounded for chatty endpoints - like issue events/comments on long-lived PRs. - """ - raw = gh("api", "--paginate", *api_args, "-q", ".[]") - for line in raw.splitlines(): - line = line.strip() - if not line: - continue - try: - yield json.loads(line) - except json.JSONDecodeError: - # A malformed line should not blow up the whole guard. Skip and - # carry on — at worst the guard fail-closes (returns False / - # None) and the caller treats it as "unknown". - continue - - -def fetch_last_close_event( - repo: str, number: int -) -> tuple[str | None, dt.datetime | None]: - """Return the actor login and timestamp of the most recent `closed` event. - - Either field may be None: actor when the events API returns nothing - (unusual for a closed item, but possible on transient errors), and - timestamp when the event lacks `created_at` or the value can't be - parsed. `was_closed_by_agent_shin` fail-closes on either. - """ - actor: str | None = None - closed_at: dt.datetime | None = None - for event in _iter_paginated_json(f"repos/{repo}/issues/{number}/events"): - if event.get("event") != "closed": - continue - actor = (event.get("actor") or {}).get("login") - created = event.get("created_at") - if not created: - closed_at = None - continue - try: - closed_at = parse_iso8601(created) - except ValueError: - closed_at = None - return actor, closed_at - - -# How much older than the latest `closed` event the Agent Shin marker -# comment is allowed to be while still counting as "this close was Agent -# Shin's". Agent Shin posts the close comment immediately before closing, -# so the marker timestamp is normally at most a few seconds before the -# close event; the buffer just absorbs clock skew between the comments -# API and the events API. -AGENT_SHIN_CLOSE_MARKER_SKEW_SECONDS = 300 - - -def was_closed_by_agent_shin( - repo: str, number: int, *, bot_login: str | None = None -) -> bool: - """Return True iff Agent Shin itself most-recently closed this PR/issue. - - This is the guard that stops `@agent-shin reconsider` from reopening an - item Agent Shin did not close — a maintainer closing for non-rubric - reasons (security, duplicate, design rejection), or a different workflow - (stale/duplicate sweeps) closing under the shared `github-actions[bot]` - identity. Three independent signals must all hold, because that identity - is not unique to Agent Shin and a marker comment from a prior - closed/reopened cycle would otherwise vouch for an unrelated close: - - 1. The most recent `closed` event's actor is the bot identity. - 2. Agent Shin left one of its auto-close comments, detected via - `AGENT_SHIN_CLOSE_MARKER`. The actor check alone can't tell an - Agent Shin close from any other `github-actions[bot]` close. - 3. That marker comment was posted at (or just before) the latest - close event, not on a previous close in an - Agent-Shin-close -> reconsider-reopen -> other-bot-reclose cycle. - - The check is intentionally fail-closed: any uncertainty about who closed - the item is treated as "not Agent Shin" so the destructive reopen path - stays gated. - """ - expected = ( - bot_login - or os.environ.get("AGENT_SHIN_BOT_LOGIN") - or AGENT_SHIN_DEFAULT_BOT_LOGIN - ).lower() - actor, closed_at = fetch_last_close_event(repo, number) - if not actor or actor.lower() != expected or closed_at is None: - return False - marker_seconds = seconds_since_last_agent_shin_close( - repo, number, bot_login=bot_login - ) - if marker_seconds is None: - return False - close_age_seconds = (dt.datetime.now(dt.timezone.utc) - closed_at).total_seconds() - return marker_seconds <= close_age_seconds + AGENT_SHIN_CLOSE_MARKER_SKEW_SECONDS - - -def _seconds_since_latest_marker_comment( - repo: str, - number: int, - *, - marker: str, - bot_login: str | None = None, -) -> float | None: - """Return seconds since the bot's most recent comment with ``marker``. - - Fetches comments via `_iter_paginated_json` and delegates the - iteration / author-filter / timestamp logic to - `agent_shin_shared.seconds_since_latest_marker_comment` so the daily - Greptile sweep and the LLM judge use one source of truth for the - "bot already posted X" detection. The wall-clock `now` is resolved - against this module's `dt` so tests that freeze time via - `monkeypatch.setattr(triage_module, "dt", ...)` still apply. - """ - return seconds_since_latest_marker_comment( - _iter_paginated_json(f"repos/{repo}/issues/{number}/comments"), - marker=marker, - bot_login=bot_login, - now=dt.datetime.now(dt.timezone.utc), - ) - - -def seconds_since_last_reconsider_verdict( - repo: str, number: int, *, bot_login: str | None = None -) -> float | None: - """Return seconds since the bot's most recent reconsider verdict comment. - - Detects comments by matching the HTML marker `RECONSIDER_COMMENT_MARKER` - appended by `format_reopen_comment` and - `format_reconsider_still_failing_comment`. Returns None when the bot - has never posted a reconsider verdict on this PR/issue (or when the - only matching comments are missing a `created_at` timestamp, which - shouldn't happen on a real GitHub response). - """ - return _seconds_since_latest_marker_comment( - repo, number, marker=RECONSIDER_COMMENT_MARKER, bot_login=bot_login - ) - - -def seconds_since_last_grace_warning( - repo: str, number: int, *, bot_login: str | None = None -) -> float | None: - """Return seconds since the bot's most recent grace-period warning. - - Detects warning comments by matching the HTML marker - `GRACE_COMMENT_MARKER` appended by `format_grace_warning_pr_comment` - and `format_grace_warning_issue_comment`. Returns None when no - grace warning has ever been posted on this PR/issue — that's the - "first low-quality detection" signal that drives the warning path. - """ - return _seconds_since_latest_marker_comment( - repo, number, marker=GRACE_COMMENT_MARKER, bot_login=bot_login - ) - - -def seconds_since_last_agent_shin_close( - repo: str, number: int, *, bot_login: str | None = None -) -> float | None: - """Return seconds since Agent Shin's most recent auto-close comment. - - Detects close comments by matching `AGENT_SHIN_CLOSE_MARKER` (stamped by - `format_pr_close_comment` / `format_issue_close_comment`). Returns None - when Agent Shin has never closed this PR/issue — the signal - `was_closed_by_agent_shin` uses to keep the reconsider reopen path gated - against closures performed by other workflows sharing the bot identity. - """ - return _seconds_since_latest_marker_comment( - repo, number, marker=AGENT_SHIN_CLOSE_MARKER, bot_login=bot_login - ) - - -# --------------------------------------------------------------------------- -# Author classification - - -def is_internal_contributor(item: dict) -> bool: - """Return True if the PR/issue author should be exempted from triage. - - Fail-safe: if `author_association` is missing or empty (which should never - happen on a successful GitHub REST response but is possible on schema - changes or partial responses), treat the author as INTERNAL so the - destructive close path never fires on an unknown contributor. This matches - the sibling `is_external_pr_author` in `close_low_quality_prs.py`. - """ - login = ((item.get("user") or {}).get("login") or "").lower() - if login.endswith("[bot]") or login in {"dependabot", "github-actions"}: - return True - association = (item.get("author_association") or "").upper() - if not association or association in INTERNAL_ASSOCIATIONS: - return True - return False - - -# --------------------------------------------------------------------------- -# Greptile score + age helpers (`extract_greptile_score`, `parse_iso8601`) -# live in `agent_shin_shared` — they're imported at the top of this module -# so both `triage_with_llm.py` and `close_low_quality_prs.py` share a -# single source of truth for the Confidence-Score regex and ISO-8601 -# parsing. - - -# --------------------------------------------------------------------------- -# Prompt construction - - -def strip_html_comments(text: str) -> str: - """Remove HTML comments — template placeholder text shouldn't fool the judge.""" - return HTML_COMMENT_PATTERN.sub("", text or "") - - -def has_linked_issue(text: str) -> bool: - """Heuristic: does this body link to an open issue (Fixes #123 etc.)?""" - return bool(LINKED_ISSUE_PATTERN.search(strip_html_comments(text or ""))) - - -def build_pr_prompt(*, title: str, body: str) -> str: - cleaned_body = strip_html_comments(body or "").strip() or "(empty)" - # Dedent the static template *before* interpolating dynamic fields so that - # multi-line bodies (whose 2nd+ lines start at column 0) don't defeat the - # common-indent computation in textwrap.dedent. - template = textwrap.dedent(""" - You are "Agent Shin", the OSS triage bot for the LiteLLM open-source - repository (BerriAI/litellm). Decide whether this external pull request - meets the project's contribution standards. - - A PR PASSES triage only if BOTH (1) AND (2) are satisfied. A linked - issue alone is NOT enough — it covers context, not proof. - - (1) CONTEXT — the PR provides AT LEAST ONE of: - (a) A link to a related GitHub issue. Acceptable forms: - "Fixes #1234", "Closes #1234", "Resolves #1234", - "Refs https://github.com/BerriAI/litellm/issues/1234". A - bare "#1234" without a closing keyword counts only if it - is clearly the related issue (not a passing mention). - (b) A clear problem description in the body (what bug or - missing feature this addresses, beyond the title) AND - expected vs. actual behavior (or, for features, "what's - possible now vs. with this PR"). - - (2) END-TO-END QA PROOF: the PR body contains AT LEAST ONE of: - (a) A screen recording / video showing the behavior before - and after the change (the bug reproducing, then the fix - working). For a brand-new feature with no meaningful - "before", a recording of it working end-to-end is fine. - (b) A screenshot (or before/after screenshots) showing the - fix or feature working. - (c) Specific commands that were actually run (curl, python, - a CLI invocation, etc.) PAIRED WITH their real - output, demonstrating the change works end-to-end against - the real system. Commands whose external dependencies - (LLM provider, DB, network) are mocked or stubbed do NOT - satisfy (2c); they are not end-to-end. - - `has_qa_proof` must be set to `true` only when (2a), (2b), - or a non-mocked (2c) is actually present in the body. If the - only "proof" is mocked tests, `has_qa_proof` is `false` and - the verdict is "fail". - - The following do NOT count as QA proof: - - Generic claims like "I tested it", "works locally", "all - tests pass", or a checked "I added tests" checkbox with no - output shown. - - A description of what tests exist or were added, without - their actual output in the PR body. - - `pytest` (or any test runner) executed against the - repository's own unit tests. Those mock the LLM provider, - DB, and network, so they are NOT end-to-end and never - satisfy (2), no matter how much passing output is pasted. - - A linked issue. The linked issue is context (1a), never - proof (2). - - FAIL the PR if EITHER (1) or (2) is missing. Do not bias toward PASS: - if QA proof is absent, the verdict is "fail" even when the rest of - the PR is well-written. - - Respond with a single JSON object, no prose: - - {{ - "verdict": "pass" | "fail", - "linked_issue": boolean, - "has_problem_description": boolean, - "has_expected_vs_actual": boolean, - "has_qa_proof": boolean, - "qa_proof_type": "video" | "screenshot" | "commands_with_output" | "none", - "missing": ["plain-english strings naming what is missing"], - "explanation": "1-2 sentence reasoning for the team to skim" - }} - - --- - PR title: {title} - - PR body: - --- - {cleaned_body} - --- - """).strip() - return template.format(title=title, cleaned_body=cleaned_body) - - -def build_issue_prompt(*, title: str, body: str) -> str: - cleaned_body = strip_html_comments(body or "").strip() or "(empty)" - # Dedent the static template *before* interpolating dynamic fields so that - # multi-line bodies (whose 2nd+ lines start at column 0) don't defeat the - # common-indent computation in textwrap.dedent. - template = textwrap.dedent(""" - You are "Agent Shin", the OSS triage bot for the LiteLLM open-source - repository (BerriAI/litellm). Decide whether this GitHub issue meets - the project's reporting standards. - - For a BUG REPORT the issue PASSES triage only when it contains BOTH: - (1) END-TO-END EVIDENCE OF THE BUG (the "before"; set - `has_repro=true` only when this is present): AT LEAST ONE of: - (a) A screen recording / video of the bug happening. - (b) A screenshot of the bug. - (c) The exact command(s) actually run (curl, python, a CLI - invocation, etc.) PAIRED WITH their real output, traceback, - or logs showing the failure against the real system. - Commands whose external dependencies (LLM provider, DB, - network) are mocked or stubbed do NOT count. - Prose-only "steps to reproduce" with no run output, video, or - screenshot do NOT satisfy (1). An unfilled template scaffold - (bare headings such as "Version or commit:" with nothing under - them, empty numbered lists) counts as absent, not as evidence. - (2) Expected vs. actual behavior (`has_expected_vs_actual`). - - FAIL the bug report if either (1) or (2) is missing. Do not bias - toward PASS: if the bug isn't demonstrated end-to-end, the verdict is - "fail" even when the report is well-written. - - For a FEATURE REQUEST the issue PASSES triage only when it contains - ALL of: - - A clear description of the proposed feature (what should LiteLLM do - that it does not today). - - Motivation / use case with a concrete example (config, API call, - UI flow, or scenario showing what's blocked today). - - END-TO-END EVIDENCE OF THE DEAD-END (set - `has_dead_end_evidence=true` only when this is present): a video, - a screenshot, or the exact command(s) actually run paired with - their real output, showing the point where the flow stops today. - Mocked or stubbed dependencies do NOT count, and an unfilled - template scaffold (bare headings, empty numbered lists) counts as - absent. - - For an issue that is neither a bug report nor a feature request (a - question, support request, or discussion), PASS as long as it has a - clear, specific ask and is not empty or template placeholder text. - - Respond with a single JSON object, no prose: - - {{ - "verdict": "pass" | "fail", - "kind": "bug" | "feature" | "other", - "has_repro": boolean, - "has_expected_vs_actual": boolean, - "has_motivation_example": boolean, - "has_dead_end_evidence": boolean, - "missing": ["plain-english strings naming what is missing"], - "explanation": "1-2 sentence reasoning for the team to skim" - }} - - --- - Issue title: {title} - - Issue body: - --- - {cleaned_body} - --- - """).strip() - return template.format(title=title, cleaned_body=cleaned_body) - - -# --------------------------------------------------------------------------- -# LLM call + verdict parsing - - -def call_llm_judge( - prompt: str, *, model: str, api_key: str, base_url: str | None -) -> str: - """Call an OpenAI-compatible chat completions endpoint. Returns raw text.""" - # Import inside the function so unit tests that monkey-patch this never - # need the openai package installed. - from openai import OpenAI - - client = ( - OpenAI(api_key=api_key, base_url=base_url) - if base_url - else OpenAI(api_key=api_key) - ) - kwargs: dict[str, Any] = { - "model": model, - "messages": [{"role": "user", "content": prompt}], - "temperature": 0, - "response_format": {"type": "json_object"}, - } - # gpt-5.x reasoning models reject `temperature != 1` unless - # `reasoning_effort` is explicitly "none". Set it via `extra_body` so this - # works across openai SDK versions regardless of whether the SDK natively - # types `reasoning_effort` as a top-level chat-completions param yet. - if model.lower().startswith(GPT5_FAMILY_PREFIX): - kwargs["extra_body"] = {"reasoning_effort": "none"} - response = client.chat.completions.create(**kwargs) - return response.choices[0].message.content or "" - - -def parse_verdict(raw: str) -> dict: - """Parse the LLM's JSON response. Tolerates ```json fences and stray text.""" - if not raw: - raise ValueError("empty LLM response") - text = raw.strip() - if text.startswith("```"): - text = re.sub(r"^```(?:json)?\s*", "", text) - text = re.sub(r"\s*```$", "", text) - try: - return json.loads(text) - except json.JSONDecodeError: - match = re.search(r"\{.*\}", text, re.DOTALL) - if not match: - raise ValueError(f"could not extract JSON from LLM response: {raw[:200]}") - return json.loads(match.group(0)) - - -# --------------------------------------------------------------------------- -# Comment composition - - -def _format_missing(missing: list[str]) -> str: - if not missing: - return "- (see explanation below)" - return "\n".join(f"- {m}" for m in missing) - - -# Rubric items the judge can mark present. The first element of each tuple is -# the verdict-JSON boolean field, the second is the human-readable label we -# render in the "what you got right" section of close / grace-warning comments. -_PR_PRESENT_LABELS: tuple[tuple[str, str], ...] = ( - ("linked_issue", "Linked a related GitHub issue"), - ("has_problem_description", "Clear problem description"), - ("has_expected_vs_actual", "Expected vs. actual behavior"), - ("has_qa_proof", "End-to-end QA proof"), -) - -# Issue rubric labels grouped by `kind`. The judge sets `kind` to one of -# {"bug", "feature", "other"}; when "other" we render both groups so we don't -# silently drop a present-flag the judge actually set to True. -_ISSUE_BUG_LABELS: tuple[tuple[str, str], ...] = ( - ( - "has_repro", - "End-to-end evidence of the bug (video, screenshot, or command + real output)", - ), - ("has_expected_vs_actual", "Expected vs. actual behavior"), -) -_ISSUE_FEATURE_LABELS: tuple[tuple[str, str], ...] = ( - ("has_motivation_example", "Motivation and concrete example"), - ( - "has_dead_end_evidence", - "End-to-end evidence of the dead-end (video, screenshot, or command + real output)", - ), -) - - -def _format_present_for_pr(verdict: dict) -> list[str]: - """Human-readable rubric items the judge confirmed are present on a PR. - - Drives the "what you got right" section in close / grace-warning comments. - The user gave explicit feedback: contributors should see what they nailed - *before* the list of gaps, so the comment doesn't read as pure rejection. - """ - return [label for field, label in _PR_PRESENT_LABELS if verdict.get(field)] - - -def _format_present_for_issue(verdict: dict) -> list[str]: - """Human-readable rubric items the judge confirmed are present on an issue. - - Branches on the judge's `kind` field. For `"other"` (or missing kind) we - render the union so a present-flag isn't dropped just because the judge - couldn't classify the issue cleanly. - """ - kind = (verdict.get("kind") or "").lower() - groups: list[tuple[tuple[str, str], ...]] = [] - if kind in ("bug", "other", ""): - groups.append(_ISSUE_BUG_LABELS) - if kind in ("feature", "other", ""): - groups.append(_ISSUE_FEATURE_LABELS) - out: list[str] = [] - for group in groups: - for field, label in group: - if verdict.get(field) and label not in out: - out.append(label) - return out - - -def _format_present_block(items: list[str]) -> str: - """Render the optional "what you got right" block. Empty string when the - judge didn't confirm anything as present — better to omit the section - entirely than to show "What you got right: (nothing)". - """ - if not items: - return "" - bullets = "\n".join(f"- ✅ {item}" for item in items) - return f"**What you got right:**\n\n{bullets}\n\n" - - -def format_pr_close_comment(verdict: dict) -> str: - missing_lines = _format_missing(verdict.get("missing") or []) - present_block = _format_present_block(_format_present_for_pr(verdict)) - explanation = verdict.get("explanation") or "" - return ( - "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " - "repository. " - "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" - "\n" - "I read the description against our " - "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md). " - "Here's how it lined up:\n" - "\n" - f"{present_block}" - "**What's still missing:**\n" - "\n" - f"{missing_lines}\n" - "\n" - f"> {explanation}\n" - "\n" - "**Closing this PR isn't a rejection of the change.** We want the open-PR list to " - "mirror what a maintainer can act on *right now*, so contributors don't get lost in a " - 'backlog. A closed PR is a soft "park this for later"; your work is still here, ' - "the diff is still here, and getting it reopened is one comment away. Take your time.\n" - "\n" - "**To bring this PR back:**\n" - "\n" - "- Update the description with the missing pieces, then comment `@agent-shin reconsider` " - "on this PR. I'll re-evaluate and reopen if it now passes.\n" - "- Or **Open a new PR** with the same fix and the updated description. GitHub doesn't " - "always let external contributors reopen a bot-closed PR, so a fresh PR is the most " - "reliable path back into the review queue.\n" - "- If Greptile's most recent score on this PR was below 4/5, comment `@greptileai` to " - "request a fresh review; that **still works even after the PR is closed**, and a " - "stronger score is one of the signals that lifts the PR back into the queue. A low " - "Greptile score isn't a blocker.\n" - "\n" - '**What "end-to-end QA proof" means**, since it\'s the most common gap: at least one ' - "of a short before/after screen recording / video (the bug reproducing, then the fix " - "working; for a brand-new feature, a recording of it working end-to-end), a screenshot " - "(or before/after screenshots) of it working, or the exact commands you ran paired " - "with their **real output** against the real system. Running `pytest` on the repo's " - "unit tests doesn't count; those mock the LLM provider, DB, and network, so they " - "aren't end-to-end. Output from a real, no-mocks integration run is what we look " - "for. A linked issue alone isn't enough either: it covers context, not proof. See " - "[the full rubric](https://docs.litellm.ai/blog/agent-shin-triage#the-rubric-for-pull-requests).\n" - "\n" - "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" - "\n" - "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, comment " - "`@agent-shin reconsider` or ping a maintainer; they'll override me.)_" - f"\n\n{AGENT_SHIN_CLOSE_MARKER}" - ) - - -def format_issue_close_comment(verdict: dict) -> str: - missing_lines = _format_missing(verdict.get("missing") or []) - present_block = _format_present_block(_format_present_for_issue(verdict)) - explanation = verdict.get("explanation") or "" - return ( - "🚅 Hi, thanks for filing this! I'm **Agent Shin**, the automated triage bot for this " - "repository. " - "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" - "\n" - "I read the issue against our reporting checklist. Here's how it lined up:\n" - "\n" - f"{present_block}" - "**What's still missing:**\n" - "\n" - f"{missing_lines}\n" - "\n" - f"> {explanation}\n" - "\n" - "**Closing this isn't us saying the bug isn't real or the request isn't useful.** We " - "want the open-issue list to mirror what a maintainer can act on *right now*, so " - "reports like yours don't get buried in a backlog. A closed issue is a soft \"park " - 'this for later"; your report is still here, and getting it reopened is one comment ' - "away. Take your time.\n" - "\n" - "**To bring this issue back:**\n" - "\n" - "1. Edit the issue description to add the missing pieces:\n" - " - For **bug reports**: end-to-end evidence of the bug (a screen recording / " - "video, a screenshot, or the exact commands you ran with their real output / " - "traceback) plus expected vs. actual behavior. Written steps with no run output, " - "video, or screenshot don't count, and mocked or stubbed runs don't count.\n" - " - For **feature requests**: a concrete description of what should change, a " - "use case and example (config / API call / UI flow), plus end-to-end evidence of " - "the dead-end (a video, a screenshot, or the exact commands you ran with their " - "real output showing where the flow stops today). Mocked or stubbed runs don't " - "count.\n" - "2. Comment `@agent-shin reconsider`. I'll re-run triage and reopen the issue if it " - "now meets the bar. (GitHub doesn't let external authors reopen an issue a maintainer " - "or bot closed, so the comment-based reconsider is the reliable path.)\n" - "\n" - "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" - "\n" - "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, comment " - "`@agent-shin reconsider` or ping a maintainer; they'll override me.)_" - f"\n\n{AGENT_SHIN_CLOSE_MARKER}" - ) - - -def format_grace_warning_pr_comment(verdict: dict) -> str: - """Comment posted on the FIRST low-quality detection — gives the - contributor a 2-hour grace window to fix the PR before the next - triage run actually closes it. - - This is the "before-close" warning. On the second triage run, if the - grace marker is older than `GRACE_PERIOD_SECONDS` AND the PR still - fails the rubric, the close path runs (which posts - `format_pr_close_comment` and closes the PR). - """ - missing_lines = _format_missing(verdict.get("missing") or []) - present_block = _format_present_block(_format_present_for_pr(verdict)) - explanation = verdict.get("explanation") or "" - return ( - "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " - "repository. " - "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" - "\n" - "I read the description against our " - "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md). " - "Here's how it lined up:\n" - "\n" - f"{present_block}" - "**What's still missing:**\n" - "\n" - f"{missing_lines}\n" - "\n" - f"> {explanation}\n" - "\n" - "If the description isn't updated in the next **2 hours**, I'll auto-close this PR. " - "That's **not** us saying we don't care about the change; we want the open-PR list to " - "mirror what a maintainer can act on *right now*, so contributors don't get lost in a " - 'backlog. A closed PR is a soft "park this for later," not a rejection. Take your ' - "time; everything below still works after the close.\n" - "\n" - "**During the grace period:** just update the PR description with the missing pieces. " - "No need to ping me; I'll re-check on the next sweep and skip the auto-close if it " - "now passes. See " - "[what counts as QA proof](https://docs.litellm.ai/blog/agent-shin-triage#the-rubric-for-pull-requests) " - "for the full rubric (a linked issue alone isn't enough; it covers context, not proof).\n" - "\n" - "**If the PR does get auto-closed in 2 hours, you still have easy recovery paths:**\n" - "\n" - "- Comment `@agent-shin reconsider` after updating the description. I'll re-evaluate " - "and reopen the PR if it now passes.\n" - "- Comment `@greptileai` to request a fresh Greptile review; that **still works even " - "after the PR is closed**, and a stronger score is one of the signals that lifts the " - "PR back into the queue. So a low Greptile score isn't a blocker either.\n" - "\n" - "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" - "\n" - "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, ping a " - "maintainer; they'll override me.)_\n" - "\n" - f"{GRACE_COMMENT_MARKER}" - ) - - -def format_grace_warning_issue_comment(verdict: dict) -> str: - """Issue analogue of `format_grace_warning_pr_comment`.""" - missing_lines = _format_missing(verdict.get("missing") or []) - present_block = _format_present_block(_format_present_for_issue(verdict)) - explanation = verdict.get("explanation") or "" - return ( - "🚅 Hi, thanks for filing this! I'm **Agent Shin**, the automated triage bot for this " - "repository. " - "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" - "\n" - "I read the issue against our reporting checklist. Here's how it lined up:\n" - "\n" - f"{present_block}" - "**What's still missing:**\n" - "\n" - f"{missing_lines}\n" - "\n" - f"> {explanation}\n" - "\n" - "If the issue isn't updated in the next **2 hours**, I'll auto-close it. That's **not** us " - "saying the bug isn't real or the request isn't useful; we want the open-issue list " - "to mirror what a maintainer can act on *right now*, so reports like yours don't get " - 'buried in a backlog. A closed issue is a soft "park this for later," not a ' - "rejection. Take your time; reopening is one comment away.\n" - "\n" - "**During the grace period:** just edit the issue description with the missing " - "pieces. No need to ping me; I'll re-check on the next sweep and skip the auto-close " - "if it now passes.\n" - "\n" - "Missing pieces, depending on what this is:\n" - "\n" - "- For **bug reports**: end-to-end evidence of the bug (a screen recording / video, a " - "screenshot, or the exact commands you ran with their real output / traceback) plus " - "expected vs. actual behavior. Written steps with no run output don't count, and " - "mocked or stubbed runs don't count.\n" - "- For **feature requests**: a concrete description of what should change, a use " - "case and example (config / API call / UI flow), plus end-to-end evidence of the " - "dead-end (a video, a screenshot, or the exact commands you ran with their real " - "output showing where the flow stops today). Mocked or stubbed runs don't count.\n" - "\n" - "**If the issue does get auto-closed in 2 hours**, comment `@agent-shin reconsider` " - "and I'll re-evaluate. If it now meets the bar, I'll reopen the issue.\n" - "\n" - "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" - "\n" - "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, ping a " - "maintainer; they'll override me.)_\n" - "\n" - f"{GRACE_COMMENT_MARKER}" - ) - - -# --------------------------------------------------------------------------- -# Step-summary helpers - - -def write_step_summary(content: str) -> None: - """When running inside GitHub Actions, append to the step summary file.""" - path = os.environ.get("GITHUB_STEP_SUMMARY") - if not path: - return - try: - with open(path, "a", encoding="utf-8") as handle: - handle.write(content) - if not content.endswith("\n"): - handle.write("\n") - except OSError as exc: - print(f"warn: failed to write step summary: {exc}", file=sys.stderr) - - -# --------------------------------------------------------------------------- -# Core orchestration - - -def format_reopen_comment(kind: str) -> str: - """Comment posted when Agent Shin reopens after a successful reconsider.""" - noun = "PR" if kind == "pr" else "issue" - # The trailing HTML marker is used by `seconds_since_last_reconsider_verdict` - # to enforce a cooldown between repeated `@agent-shin reconsider` triggers. - # Keep the marker on its own line so it doesn't disturb the rendered text. - return ( - f"♻️ **Re-evaluated and reopened.** Thanks for updating the {noun}!\n" - "\n" - "Agent Shin re-ran triage on the latest description and it now meets " - "the bar. A maintainer will take another look soon; please don't " - f"close this {noun} again unless asked to.\n" - "\n" - "_(If a maintainer ends up closing this for non-rubric reasons, that " - "decision stands; comment `@agent-shin reconsider` again only if you " - "have substantively new information.)_\n" - "\n" - f"{RECONSIDER_COMMENT_MARKER}" - ) - - -def format_reconsider_still_failing_comment(kind: str, verdict: dict) -> str: - """Comment posted when reconsider re-runs triage but the verdict is still fail.""" - missing_lines = _format_missing(verdict.get("missing") or []) - explanation = verdict.get("explanation") or "" - noun = "PR" if kind == "pr" else "issue" - # The trailing HTML marker is used by `seconds_since_last_reconsider_verdict` - # to enforce a cooldown between repeated `@agent-shin reconsider` triggers. - return ( - f"⏸️ **Re-evaluated; this {noun} still doesn't meet the rubric.**\n" - "\n" - "Agent Shin re-ran triage on the current description but is still " - "missing:\n" - "\n" - f"{missing_lines}\n" - "\n" - f"> {explanation}\n" - "\n" - "Update the description with the missing pieces and comment " - "`@agent-shin reconsider` again, or ping a maintainer if you think " - "I got this wrong.\n" - "\n" - "_(I'm an LLM and I'm not infallible.)_\n" - "\n" - f"{RECONSIDER_COMMENT_MARKER}" - ) - - -# --------------------------------------------------------------------------- -# Review gate — "ready for review" label lifecycle - -_UNSET = object() - - -def _combine_missing( - verdict: dict, greptile_score: int | None, min_score: int -) -> list[str]: - """Merge the LLM rubric's `missing` list with a Greptile-score shortfall.""" - missing = list(verdict.get("missing") or []) - if greptile_score is not None and greptile_score < min_score: - missing.insert( - 0, - f"Greptile's most recent review scored this PR {greptile_score}/5 " - f"(below the {min_score}/5 bar)", - ) - return missing or ["(see explanation below)"] - - -def _has_marker( - comments: Iterable[dict], marker: str, *, bot_login: str | None = None -) -> bool: - """Return True iff the bot itself posted a comment containing ``marker``. - - Filters by author so a contributor who quotes the marker (e.g. via - GitHub's "Quote reply" feature, which preserves HTML comments in - raw markdown) is not mistaken for a bot action — that would - silently suppress notifications or change which "recovered" wording - is selected. Matches the author-filter pattern used by the sibling - `_seconds_since_latest_marker_comment` helper. - """ - expected_login = ( - bot_login - or os.environ.get("AGENT_SHIN_BOT_LOGIN") - or AGENT_SHIN_DEFAULT_BOT_LOGIN - ).lower() - for comment in comments: - author = ((comment.get("user") or {}).get("login") or "").lower() - if author != expected_login: - continue - if marker in (comment.get("body") or ""): - return True - return False - - -def format_ready_for_review_comment( - verdict: dict, - greptile_score: int | None, - min_greptile_score: int = DEFAULT_MIN_GREPTILE_SCORE, -) -> str: - """Posted the first time a PR clears the bar (label added).""" - score_line = ( - f" Greptile scored it **{greptile_score}/5**." - if greptile_score is not None - else "" - ) - explanation = verdict.get("explanation") or "" - return ( - "✅ **Triage passed, tagging `ready for review`.**\n" - "\n" - "Agent Shin checked this PR against the " - "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md) " - "and it clears the bar (a linked issue, or a clear problem description " - f"+ expected vs. actual + QA proof).{score_line}\n" - "\n" - f"> {explanation}\n" - "\n" - "A maintainer will take it from here. If a later re-check finds the PR " - f"has regressed (Greptile drops below {min_greptile_score}/5, " - "the QA proof is removed, etc.) I'll pull the tag and comment with " - "what's missing; fix it and the tag comes back automatically.\n" - f"{READY_MARKER}" - ) - - -def format_all_clear_comment(verdict: dict, greptile_score: int | None) -> str: - """Posted when a PR recovers after a regression (label re-added).""" - score_line = ( - f" Greptile is back to **{greptile_score}/5**." - if greptile_score is not None - else "" - ) - explanation = verdict.get("explanation") or "" - return ( - "✅ **All clear again, re-adding `ready for review`.**\n" - "\n" - "Thanks for addressing the earlier feedback. On re-check this PR meets " - f"the contribution bar once more.{score_line}\n" - "\n" - f"> {explanation}\n" - "\n" - "A maintainer will take another look.\n" - f"{READY_MARKER}" - ) - - -def format_regression_comment( - missing: list[str], explanation: str, grace_days: int -) -> str: - """Posted when a previously-tagged PR regresses (label removed, PR stays open). - - Discloses the same ``grace_days`` deadline the state machine enforces: - once that window elapses with the PR still failing, the close path fires. - Hiding the deadline behind a bare "stays open" would surprise contributors - with an auto-close they were never warned about. - """ - window = "24 hours" if grace_days == 1 else f"{grace_days} days" - return ( - "⚠️ **Removing the `ready for review` tag.**\n" - "\n" - "On a re-check this PR no longer meets the contribution bar. What's " - "missing now:\n" - "\n" - f"{_format_missing(missing)}\n" - "\n" - f"> {explanation}\n" - "\n" - f"The PR stays open for ~{window}; address the points above and Agent " - 'Shin will post an "all clear" comment and re-add the tag ' - "automatically. If the points still aren't addressed after that " - "window, the PR is auto-closed; that's not a rejection, and you can " - "comment `@agent-shin reconsider` to have it re-evaluated and reopened " - "once it passes.\n" - f"{REGRESSED_MARKER}" - ) - - -def format_within_grace_comment( - missing: list[str], explanation: str, grace_days: int -) -> str: - """Posted once while a failing PR is still inside its grace window.""" - window = "24 hours" if grace_days == 1 else f"{grace_days} days" - return ( - "🚅 Hi, thanks for the PR! This is **Agent Shin**, the automated triage " - "bot. This PR doesn't quite meet the contribution bar yet:\n" - "\n" - f"{_format_missing(missing)}\n" - "\n" - f"> {explanation}\n" - "\n" - f"You have ~{window} from when this PR was opened to add the missing " - "pieces; just update the description and I'll re-check on the next " - "sweep. Once it passes I'll tag it `ready for review`. If it does get " - "auto-closed, that's not a rejection; comment `@agent-shin reconsider` " - "and I'll re-evaluate and reopen if it now passes.\n" - f"{WITHIN_GRACE_MARKER}" - ) - - -def review_gate( - *, - repo: str, - number: int, - close: bool, - model: str, - judge: Any = None, - greptile_score: Any = _UNSET, - comments: Any = _UNSET, - now: dt.datetime | None = None, - grace_days: int = DEFAULT_GRACE_DAYS, - min_greptile_score: int = DEFAULT_MIN_GREPTILE_SCORE, - label: str = READY_FOR_REVIEW_LABEL, - allowlist: frozenset[str] = ALLOWLIST_LOGINS, -) -> dict: - """Reconcile the `ready for review` label with a PR's current quality. - - A PR is *passing* when it clears BOTH gates: the LLM rubric (linked issue, - or problem description + expected/actual + QA proof) AND Greptile's most - recent confidence score (>= ``min_greptile_score``; absence of a score is - not held against the PR). The gate then drives a small state machine, using - the label itself as the persisted state so comments fire only on - transitions (never on every scheduled run): - - passing, untagged -> add label + "ready for review" / "all clear" - passing, tagged -> noop-passing - not passing, tagged -> remove label + regression comment (stays open) - not passing, untagged, old -> close + comment (past the grace window) - not passing, untagged, new -> one-time "what's missing" notice (within grace) - - ``close`` gates every destructive side effect: with ``close=False`` the - function returns a ``would-*`` preview and touches nothing, mirroring the - dry-run contract of :func:`triage`. ``judge``/``greptile_score``/ - ``comments``/``now`` are injectable for tests; in production they are - resolved from the OpenAI judge, the PR's Greptile comment, the live comment - list, and the wall clock respectively. - """ - item = fetch_pr(repo, number) - - title = item.get("title") or "" - body = item.get("body") or "" - login = (item.get("user") or {}).get("login") or "" - association = item.get("author_association") or "" - state = item.get("state") or "" - # GitHub label names are case-insensitive; compare lowercased so a repo - # that already has e.g. "Ready for Review" is recognized as the same - # label as our READY_FOR_REVIEW_LABEL constant ("ready for review"). - labels_now = {(lbl.get("name") or "").lower() for lbl in (item.get("labels") or [])} - label_key = label.lower() - created_raw = item.get("created_at") or "" - - base_result = { - "kind": "pr", - "number": number, - "title": title, - "author": login, - "author_association": association, - "state": state, - "labeled": label_key in labels_now, - "review_gate": True, - } - - if state != "open": - return {**base_result, "action": "skip-not-open"} - - if allowlist: - if login.lower() not in allowlist: - return {**base_result, "action": "skip-not-allowlisted"} - elif is_internal_contributor(item): - return {**base_result, "action": "skip-internal-author"} - - # Resolve the comment list once — used for both the Greptile score and the - # marker-based dedup below. - if comments is _UNSET: - comments = list(_iter_paginated_json(f"repos/{repo}/issues/{number}/comments")) - - # --- rubric verdict: linked-issue short-circuit, else the LLM judge ------- - if has_linked_issue(body): - verdict = { - "verdict": "pass", - "linked_issue": True, - "missing": [], - "explanation": "Linked-issue regex matched; LLM was not called.", - } - rubric_pass = True - else: - prompt = build_pr_prompt(title=title, body=body) - if judge is None: - api_key = os.environ.get("OPENAI_API_KEY") - if not api_key: - return {**base_result, "action": "skip-no-llm-key"} - base_url = os.environ.get("OPENAI_BASE_URL") or None - - def judge(p: str) -> str: - return call_llm_judge( - p, model=model, api_key=api_key, base_url=base_url - ) - - try: - verdict = parse_verdict(judge(prompt)) - except Exception as exc: # noqa: BLE001 - judge errors must never act - return {**base_result, "action": "skip-llm-error", "error": str(exc)} - rubric_pass = (verdict.get("verdict") or "").lower() == "pass" - - # --- Greptile score ------------------------------------------------------- - if greptile_score is _UNSET: - extraction = extract_greptile_score(comments) - greptile_score = extraction[0] if extraction else None - greptile_ok = greptile_score is None or greptile_score >= min_greptile_score - passing = rubric_pass and greptile_ok - - # --- age ------------------------------------------------------------------ - age_days = None - if created_raw: - reference = now or dt.datetime.now(dt.timezone.utc) - age_days = (reference - parse_iso8601(created_raw)).days - - label_present = label_key in labels_now - explanation = verdict.get("explanation") or "" - # When the rubric short-circuited to pass (linked-issue regex) but - # Greptile dragged the PR below the bar, the synthetic verdict's - # explanation ("LLM was not called") would mislead a contributor reading - # the regression / close comment. Surface the real reason instead. - if rubric_pass and not greptile_ok: - explanation = ( - f"Greptile's most recent review scored this PR " - f"{greptile_score}/5 (below the {min_greptile_score}/5 bar)." - ) - verdict = {**verdict, "explanation": explanation} - base_result = { - **base_result, - "verdict": verdict, - "greptile_score": greptile_score, - "passing": passing, - "age_days": age_days, - } - - if passing: - if label_present: - return {**base_result, "action": "noop-passing"} - recovered = _has_marker(comments, REGRESSED_MARKER) - comment = ( - format_all_clear_comment(verdict, greptile_score) - if recovered - else format_ready_for_review_comment( - verdict, greptile_score, min_greptile_score - ) - ) - if not close: - return {**base_result, "action": "would-label-ready", "comment": comment} - post_comment(repo, number, comment) - add_label(repo, number, label) - return {**base_result, "action": "labeled-ready", "comment": comment} - - missing = _combine_missing(verdict, greptile_score, min_greptile_score) - - if label_present: - comment = format_regression_comment(missing, explanation, grace_days) - if not close: - return {**base_result, "action": "would-remove-label", "comment": comment} - remove_label(repo, number, label) - post_comment(repo, number, comment) - return {**base_result, "action": "label-removed-regressed", "comment": comment} - - # Not passing and not tagged. If the PR was previously tagged and then - # regressed (we removed the label and posted REGRESSED_MARKER), honor the - # "PR stays open — fix it and the tag comes back" promise from - # `format_regression_comment` and skip the close path. Without this guard, - # any PR older than `grace_days` would be closed on the next evaluation, - # giving the contributor no realistic window to address the regression. - # - # The promise has a deliberate expiration: once `grace_days` have elapsed - # since the regression notice, fall through to the close path so a PR that - # was abandoned post-regression doesn't sit open forever. - if _has_marker(comments, REGRESSED_MARKER): - reference = now or dt.datetime.now(dt.timezone.utc) - seconds_since_regression = seconds_since_latest_marker_comment( - comments, marker=REGRESSED_MARKER, now=reference - ) - grace_seconds = grace_days * 86400 - if seconds_since_regression is None or seconds_since_regression < grace_seconds: - return {**base_result, "action": "regressed-already-notified"} - - # Not passing and not tagged: close if past the grace window, else notify once. - if age_days is not None and age_days >= grace_days: - comment = format_pr_close_comment({**verdict, "missing": missing}) - if not close: - return {**base_result, "action": "would-close", "comment": comment} - post_comment(repo, number, comment) - close_pr(repo, number) - return {**base_result, "action": "closed", "comment": comment} - - if _has_marker(comments, WITHIN_GRACE_MARKER): - return {**base_result, "action": "within-grace-already-notified"} - comment = format_within_grace_comment(missing, explanation, grace_days) - if not close: - return { - **base_result, - "action": "would-notify-within-grace", - "comment": comment, - } - post_comment(repo, number, comment) - return {**base_result, "action": "within-grace-notified", "comment": comment} - - -def triage( - *, - repo: str, - kind: str, - number: int, - close: bool, - model: str, - judge: Any = None, - print_prompt: bool = False, - reconsider: bool = False, - allowlist: frozenset[str] = ALLOWLIST_LOGINS, -) -> dict: - """Triage a single PR or issue. Returns a result dict for logging/tests. - - `judge` is an optional callable `(prompt) -> str` for tests / dry-run with - a stub. In production, leave it None and the script uses `call_llm_judge`. - - When `reconsider=True`, the closed-state guard is skipped and a - fail-but-no-comment is replaced with a "still failing" comment + leave - closed; a pass triggers `reopen_pr`/`reopen_issue` plus a reopen comment. - Reconsider mode is intended for the `@agent-shin reconsider` comment - trigger. Like regular triage, `close=False` keeps reconsider in dry-run - (returns `would-reopen` / `would-reconsider-still-failing` so a local - operator can preview without write side effects); the workflow only - passes `--close` when `AGENT_SHIN_ENABLED=true`. - - Reconsider mode adds two extra safety guards on top of the regular - triage skip-internal-author check: - - 1. **Bot-closed guard.** Only reopens if the most recent close was - performed by the bot identity (default `github-actions[bot]`). - This stops a contributor from using `@agent-shin reconsider` to - override a maintainer's close for non-rubric reasons. - 2. **Rate-limit guard.** If the bot has already posted a reconsider - verdict on this PR/issue within `RECONSIDER_RATE_LIMIT_SECONDS`, - skip — repeated triggers from the same contributor shouldn't burn - CI minutes or LLM budget. - """ - fetcher = {"pr": fetch_pr, "issue": fetch_issue}[kind] - item = fetcher(repo, number) - - title = item.get("title") or "" - body = item.get("body") or "" - login = (item.get("user") or {}).get("login") or "" - association = item.get("author_association") or "" - state = item.get("state") or "" - - base_result = { - "kind": kind, - "number": number, - "title": title, - "author": login, - "author_association": association, - "state": state, - "reconsider": reconsider, - } - - # Reconsider only makes sense on a closed PR/issue. A "reconsider on an - # open PR" is a no-op (the regular triage flow already evaluates open - # PRs); return a clear skip so the workflow can short-circuit. - if reconsider: - if state != "closed": - return {**base_result, "action": "skip-not-closed"} - else: - if state != "open": - return {**base_result, "action": "skip-not-open"} - - if allowlist: - if login.lower() not in allowlist: - return {**base_result, "action": "skip-not-allowlisted"} - elif is_internal_contributor(item): - return {**base_result, "action": "skip-internal-author"} - - # Reconsider-only guards — these run BEFORE the LLM call so a - # maintainer-closed PR / rate-limited trigger never spends LLM budget. - if reconsider: - if not was_closed_by_agent_shin(repo, number): - return {**base_result, "action": "skip-not-bot-closed"} - age = seconds_since_last_reconsider_verdict(repo, number) - if age is not None and age < RECONSIDER_RATE_LIMIT_SECONDS: - return { - **base_result, - "action": "skip-rate-limited", - "rate_limit_age_seconds": age, - "rate_limit_window_seconds": RECONSIDER_RATE_LIMIT_SECONDS, - } - - if kind == "pr": - # Short-circuit: if body very clearly links a related issue, just pass. - if has_linked_issue(body): - base = { - **base_result, - "action": "pass-linked-issue", - "verdict": { - "verdict": "pass", - "linked_issue": True, - "explanation": "Linked-issue regex matched; LLM was not called.", - }, - } - if reconsider: - # Pass-on-reconsider -> reopen the PR with a friendly comment. - reopen_body = format_reopen_comment(kind) - if not close: - return { - **base, - "action": "would-reopen", - "comment": reopen_body, - } - post_comment(repo, number, reopen_body) - reopen_pr(repo, number) - return { - **base, - "action": "reopened", - "comment": reopen_body, - } - return base - prompt = build_pr_prompt(title=title, body=body) - else: - prompt = build_issue_prompt(title=title, body=body) - - if print_prompt: - return {**base_result, "action": "print-prompt", "prompt": prompt} - - if judge is None: - api_key = os.environ.get("OPENAI_API_KEY") - if not api_key: - # No key configured — never take a destructive action. Report skip. - return { - **base_result, - "action": "skip-no-llm-key", - "prompt_preview": prompt[:200], - } - base_url = os.environ.get("OPENAI_BASE_URL") or None - - def judge(p: str) -> str: - return call_llm_judge(p, model=model, api_key=api_key, base_url=base_url) - - try: - raw = judge(prompt) - verdict = parse_verdict(raw) - except Exception as exc: # noqa: BLE001 - judge errors must never close PRs - return {**base_result, "action": "skip-llm-error", "error": str(exc)} - - decision = (verdict.get("verdict") or "").lower() - - if reconsider: - # Reconsider: an explicit `pass` -> reopen + post reopen comment; - # anything else (fail, missing/malformed verdict, typo) -> leave - # closed + post a "still failing" comment so the contributor can - # iterate again. Reopen is destructive, so a flaky/empty verdict - # must not satisfy the gate. - # In dry-run (`close=False`) we return `would-*` actions instead - # of touching GitHub state, mirroring the regular triage flow's - # `would-close`. This lets a local operator preview the outcome - # of `python triage_with_llm.py --reconsider --pr N` without - # risking accidental comments or reopens. - if decision == "pass": - reopen_body = format_reopen_comment(kind) - if not close: - return { - **base_result, - "action": "would-reopen", - "verdict": verdict, - "comment": reopen_body, - } - post_comment(repo, number, reopen_body) - if kind == "pr": - reopen_pr(repo, number) - else: - reopen_issue(repo, number) - return { - **base_result, - "action": "reopened", - "verdict": verdict, - "comment": reopen_body, - } - still_failing = format_reconsider_still_failing_comment(kind, verdict) - if not close: - return { - **base_result, - "action": "would-reconsider-still-failing", - "verdict": verdict, - "comment": still_failing, - } - post_comment(repo, number, still_failing) - return { - **base_result, - "action": "reconsider-still-failing", - "verdict": verdict, - "comment": still_failing, - } - - if decision != "fail": - return {**base_result, "action": "pass-llm", "verdict": verdict} - - # Grace-period flow: on the first low-quality detection, post a warning - # comment instead of closing immediately. On a subsequent triage run - # (manual re-trigger, or the daily `close_low_quality_prs.py` cron - # finding the same PR in its own pass), if `GRACE_PERIOD_SECONDS` has - # elapsed since the warning AND the PR still fails the rubric, close. - grace_age = seconds_since_last_grace_warning(repo, number) - if grace_age is None: - warning_body = ( - format_grace_warning_pr_comment(verdict) - if kind == "pr" - else format_grace_warning_issue_comment(verdict) - ) - if not close: - return { - **base_result, - "action": "would-warn-grace", - "verdict": verdict, - "comment": warning_body, - } - post_comment(repo, number, warning_body) - return { - **base_result, - "action": "warned-grace", - "verdict": verdict, - "comment": warning_body, - } - if grace_age < GRACE_PERIOD_SECONDS: - return { - **base_result, - "action": "skip-in-grace-period", - "verdict": verdict, - "grace_age_seconds": grace_age, - "grace_period_seconds": GRACE_PERIOD_SECONDS, - } - - # The grace window has elapsed. `--close` still gates the destructive - # write so a dry-run preview never posts or closes — the workflow only - # passes `--close` when `AGENT_SHIN_ENABLED=true`, which keeps the bot - # inert by default. - if not close: - return {**base_result, "action": "would-close", "verdict": verdict} - - comment_body = ( - format_pr_close_comment(verdict) - if kind == "pr" - else format_issue_close_comment(verdict) - ) - post_comment(repo, number, comment_body) - if kind == "pr": - close_pr(repo, number) - else: - close_issue(repo, number) - - return { - **base_result, - "action": "closed", - "verdict": verdict, - "comment": comment_body, - } - - -# --------------------------------------------------------------------------- -# CLI - - -def render_summary(result: dict) -> str: - """Render a human-readable summary block (used for stdout + step summary).""" - lines = ["## Agent Shin verdict", ""] - lines.append( - f"- **{result['kind'].upper()} #{result['number']}**: {result.get('title', '')}" - ) - lines.append( - f"- **Author**: `{result.get('author', '')}` ({result.get('author_association', '')})" - ) - lines.append(f"- **State**: {result.get('state', '')}") - lines.append(f"- **Action**: `{result['action']}`") - verdict = result.get("verdict") - if verdict: - lines.append("") - lines.append("```json") - lines.append(json.dumps(verdict, indent=2)) - lines.append("```") - error = result.get("error") - if error: - lines.append("") - lines.append(f"_LLM error: {error}_") - comment = result.get("comment") - if comment: - lines.append("") - lines.append("### Posted comment:") - lines.append("") - lines.append("> " + comment.replace("\n", "\n> ")) - return "\n".join(lines) - - -def main() -> int: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--repo", required=True, help="Repository (owner/repo).") - target = parser.add_mutually_exclusive_group(required=True) - target.add_argument("--pr", type=int, help="Pull request number to triage.") - target.add_argument("--issue", type=int, help="Issue number to triage.") - parser.add_argument( - "--close", - action="store_true", - help="Actually post comment + close on fail (default: dry run).", - ) - parser.add_argument( - "--model", - # `os.environ.get("TRIAGE_MODEL", DEFAULT_MODEL)` would return "" when - # GitHub Actions exposes an unset repo variable as an empty-string env - # var, silently bypassing DEFAULT_MODEL and causing every call to fail - # as `skip-llm-error`. The `or` guard collapses empty -> default. - default=os.environ.get("TRIAGE_MODEL") or DEFAULT_MODEL, - help=f"OpenAI-compatible model name (default: {DEFAULT_MODEL}).", - ) - parser.add_argument( - "--print-prompt", - action="store_true", - help="Print the prompt that would be sent to the judge and exit.", - ) - parser.add_argument( - "--reconsider", - action="store_true", - help=( - "Re-run triage on a CLOSED PR/issue and reopen it on pass. " - "Used by the `@agent-shin reconsider` comment-trigger workflow. " - "Only invoke this from a workflow that has already gated on " - "AGENT_SHIN_ENABLED=true and verified the commenter is the " - "PR/issue author or an internal collaborator." - ), - ) - parser.add_argument( - "--review-gate", - action="store_true", - help=( - "Reconcile the `ready for review` label for an OPEN PR: tag on " - "pass, remove the tag + comment on regression, close after the " - "grace window if it never passed. PR-only." - ), - ) - parser.add_argument( - "--grace-days", - type=int, - default=DEFAULT_GRACE_DAYS, - help=( - "Review-gate only: hours/24 a failing, un-tagged PR may stay open " - f"before auto-close (default: {DEFAULT_GRACE_DAYS} = 24h)." - ), - ) - parser.add_argument( - "--min-greptile-score", - type=int, - default=DEFAULT_MIN_GREPTILE_SCORE, - choices=range(1, 6), - help=( - "Review-gate only: Greptile score below which a PR counts as not " - f"passing (default: {DEFAULT_MIN_GREPTILE_SCORE} -> <4/5 regresses)." - ), - ) - args = parser.parse_args() - - kind = "pr" if args.pr is not None else "issue" - number = args.pr if args.pr is not None else args.issue - - if args.review_gate: - if kind != "pr": - parser.error("--review-gate applies to pull requests only (use --pr).") - result = review_gate( - repo=args.repo, - number=number, - close=args.close, - model=args.model, - grace_days=args.grace_days, - min_greptile_score=args.min_greptile_score, - ) - else: - result = triage( - repo=args.repo, - kind=kind, - number=number, - close=args.close, - model=args.model, - print_prompt=args.print_prompt, - reconsider=args.reconsider, - ) - - if result.get("action") == "print-prompt": - print(result["prompt"]) - return 0 - - summary = render_summary(result) - print(summary) - write_step_summary(summary + "\n") - return 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/.github/workflows/close_low_quality_prs.yml b/.github/workflows/close_low_quality_prs.yml deleted file mode 100644 index 2401be84000..00000000000 --- a/.github/workflows/close_low_quality_prs.yml +++ /dev/null @@ -1,92 +0,0 @@ -name: Close Low-Quality PRs - -# Auto-close any open PR (including drafts, regardless of age) authored by an -# external OSS contributor that Greptile reviewed with a confidence score -# below 4/5. Closures are explained in a comment that tells the contributor -# to push fixes and open a fresh PR (since OSS authors cannot reopen a PR -# closed by a bot/maintainer) or comment `@agent-shin reconsider` to have -# Agent Shin re-evaluate. -# -# Manual one-off run: -# gh workflow run "Close Low-Quality PRs" -f close=true -# -# Dry-run preview (no PRs are touched): -# gh workflow run "Close Low-Quality PRs" -f close=false - -on: - schedule: - # Daily at 09:00 UTC. Pairs well with the stale-issue workflow at midnight. - - cron: "0 9 * * *" - workflow_dispatch: - inputs: - close: - description: "Actually close matching PRs (false = dry run)." - required: false - default: "false" - type: choice - options: - - "true" - - "false" - min_age_days: - description: "Minimum PR age in days (default 0 = no age filter)." - required: false - default: "0" - min_score: - description: "Greptile score below which a PR is closed (1-5)." - required: false - default: "4" - limit: - description: "Maximum number of PRs to close in a single run." - required: false - default: "25" - -permissions: - contents: read - pull-requests: write - issues: write - -jobs: - close-low-quality-prs: - if: github.repository == 'BerriAI/litellm' - runs-on: ubuntu-latest - steps: - - name: Checkout triage script - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - sparse-checkout: .github/scripts - persist-credentials: false - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Run low-quality PR closer - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - # Scheduled runs are ALWAYS dry-run, even when AGENT_SHIN_ENABLED is - # "true", so the team can QA the closer's verdicts in step summaries - # before any contributor sees a PR closed. Real closures only happen - # on manual workflow_dispatch with close=true (and the variable set). - CLOSE_FLAG: ${{ github.event.inputs.close || 'false' }} - AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} - MIN_AGE_DAYS: ${{ github.event.inputs.min_age_days || '0' }} - MIN_SCORE: ${{ github.event.inputs.min_score || '4' }} - LIMIT: ${{ github.event.inputs.limit || '25' }} - run: | - set -euo pipefail - ARGS=( - --repo "${{ github.repository }}" - --min-age-days "${MIN_AGE_DAYS}" - --min-score "${MIN_SCORE}" - --limit "${LIMIT}" - ) - if [ "${AGENT_SHIN_ENABLED:-false}" != "true" ]; then - echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> forcing dry-run regardless of close input." - elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${CLOSE_FLAG}" = "true" ]; then - ARGS+=(--close) - echo "::notice::Running in close-on-fail mode." - else - echo "::notice::AGENT_SHIN_ENABLED is true but this trigger is dry-run (scheduled event or close=false)." - fi - python3 .github/scripts/close_low_quality_prs.py "${ARGS[@]}" diff --git a/.github/workflows/create_daily_oss_agent_shin_branch.yml b/.github/workflows/create_daily_oss_agent_shin_branch.yml deleted file mode 100644 index 9baf9f142f6..00000000000 --- a/.github/workflows/create_daily_oss_agent_shin_branch.yml +++ /dev/null @@ -1,28 +0,0 @@ -name: Create Daily oss-agent-shin Branch - -on: - schedule: - - cron: "0 0 * * *" # Runs every day at midnight UTC - workflow_dispatch: # Allow manual trigger - -jobs: - create-oss-agent-shin-branch: - if: github.repository == 'BerriAI/litellm' - runs-on: ubuntu-latest - permissions: - contents: write - - steps: - - name: Create daily oss-agent-shin branch - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - run: | - BRANCH_NAME="litellm_oss_agent_shin_$(date +'%m_%d_%Y')" - echo "Creating branch: $BRANCH_NAME" - if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then - echo "Branch $BRANCH_NAME already exists. Skipping creation." - exit 0 - fi - MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha') - gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent - echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA" diff --git a/.github/workflows/triage_reconsider.yml b/.github/workflows/triage_reconsider.yml deleted file mode 100644 index f35f681d09a..00000000000 --- a/.github/workflows/triage_reconsider.yml +++ /dev/null @@ -1,172 +0,0 @@ -name: Agent Shin — reconsider - -# Comment-trigger workflow: when the PR/issue author (or an internal -# collaborator) comments `@agent-shin reconsider` on a CLOSED PR/issue, -# Agent Shin re-runs LLM-judge triage on the current title+body and: -# -# - on PASS: posts a "re-evaluated and reopened" comment + reopens. -# - on FAIL: posts a "still missing X" comment and leaves it closed, -# so the contributor can iterate again. -# -# This exists because GitHub does NOT let an external (non-write-access) -# OSS contributor reopen a PR/issue closed by a bot or maintainer. Without -# this comment trigger, a contributor whose PR Agent Shin auto-closed -# would have no path back into the review queue except opening a fresh PR -# (which loses the original PR's history). The bot, on the other hand, -# has write access via GH_TOKEN and can reopen on their behalf. -# -# DRY-RUN BY DEFAULT — gated on `vars.AGENT_SHIN_ENABLED == 'true'` just -# like the other Agent Shin workflows. The workflow also gates on the -# commenter being either the PR/issue author or an internal collaborator -# (OWNER/MEMBER/COLLABORATOR) so random commenters cannot DOS the LLM -# judge or force a reopen. - -on: - issue_comment: - types: [created] - -permissions: - contents: read - issues: write - pull-requests: write - -jobs: - reconsider: - if: | - github.repository == 'BerriAI/litellm' - && contains(github.event.comment.body, '@agent-shin reconsider') - runs-on: ubuntu-latest - steps: - - name: Authorize commenter - # Only the PR/issue author OR an internal collaborator may trigger - # a reconsider. Outside random commenters could otherwise spam the - # phrase to burn LLM budget or, if a fail-open bug were ever - # introduced, force a reopen on someone else's behalf. - # - # We expose the authorization decision as a step output and gate - # every subsequent (potentially destructive) step on it. A `run:` - # step with `exit 0` would NOT stop the job — only `if:` gating - # on a known-true output is safe here. - id: auth - env: - COMMENTER: ${{ github.event.comment.user.login }} - AUTHOR: ${{ github.event.issue.user.login }} - ASSOCIATION: ${{ github.event.comment.author_association }} - run: | - set -euo pipefail - if [ "${COMMENTER}" = "${AUTHOR}" ]; then - echo "::notice::Authorized: commenter is the PR/issue author." - echo "authorized=true" >> "$GITHUB_OUTPUT" - exit 0 - fi - case "${ASSOCIATION}" in - OWNER|MEMBER|COLLABORATOR) - echo "::notice::Authorized: commenter is an internal collaborator (${ASSOCIATION})." - echo "authorized=true" >> "$GITHUB_OUTPUT" - ;; - *) - echo "::notice::Commenter '${COMMENTER}' (${ASSOCIATION}) is not authorized to trigger reconsider; skipping subsequent steps." - echo "authorized=false" >> "$GITHUB_OUTPUT" - ;; - esac - - - name: React 👀 to acknowledge the reconsider - # Add an eyes reaction to the triggering comment the moment we accept - # it, so the contributor gets instant feedback that the bot saw their - # `@agent-shin reconsider` before the slower triage steps run. Gated on - # AGENT_SHIN_ENABLED so dry-run leaves no visible trace. Best-effort: - # a reactions API hiccup must never fail the actual reconsider. - if: steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true' - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - COMMENT_ID: ${{ github.event.comment.id }} - run: | - set -euo pipefail - gh api --method POST \ - -H "Accept: application/vnd.github+json" \ - "repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \ - -f content=eyes \ - || echo "::warning::failed to add 👀 reaction (non-fatal)" - - - name: Checkout triage script - if: steps.auth.outputs.authorized == 'true' - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - sparse-checkout: .github/scripts - persist-credentials: false - - - name: Set up Python - if: steps.auth.outputs.authorized == 'true' - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Install LLM client - if: steps.auth.outputs.authorized == 'true' - run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt - - - name: Run Agent Shin reconsider - if: steps.auth.outputs.authorized == 'true' - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - # Only expose the LLM key when the bot is enabled, so a PR/issue - # author can't force paid LLM calls by spamming `@agent-shin - # reconsider` while the bot is still in dry-run. The Python script - # calls the LLM whenever this var is set (regardless of `--close`); - # stripping `--close` doesn't suppress the API call, only the - # destructive side effects. Mirror the gating used by every other - # Agent Shin workflow (triage_pr_with_llm.yml, review_gate.yml, ...). - OPENAI_API_KEY: ${{ vars.AGENT_SHIN_ENABLED == 'true' && secrets.OPENAI_API_KEY || '' }} - OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} - TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} - AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} - # `issue_comment` events fire for both issues and PR comments. - # `issue.pull_request` is set iff this is a PR comment, so we use - # its presence to decide whether to invoke `--pr N` or `--issue N`. - IS_PR: ${{ github.event.issue.pull_request != null }} - NUMBER: ${{ github.event.issue.number }} - run: | - set -euo pipefail - if [ "${IS_PR}" = "true" ]; then - ARGS=(--repo "${{ github.repository }}" --pr "${NUMBER}" --reconsider) - else - ARGS=(--repo "${{ github.repository }}" --issue "${NUMBER}" --reconsider) - fi - # Reconsider's destructive actions (post comment + reopen) are - # gated on `--close`, mirroring the regular triage workflows. - # When AGENT_SHIN_ENABLED is not the EXACT string "true", we - # still run the script so its verdict + would-X action lands in - # the step summary for QA — but without `--close`, the script - # returns `would-reopen` / `would-reconsider-still-failing` - # instead of touching GitHub state. - # - # Use the positive `= "true"` gate (not `!= "true" -> exit`) so - # the workflow guardrails in - # tests/test_litellm/test_github_triage_workflows.py see the - # canonical fail-safe enable pattern. Unknown values like - # "True", "yes", "1", or typos fall through to the dry-run - # branch, which is the safe default. - if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then - ARGS+=(--close) - echo "::notice::Agent Shin reconsider ENABLED — running real triage (close=true)." - else - echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> reconsider stays in dry-run (no comment, no reopen)." - fi - python3 .github/scripts/triage_with_llm.py "${ARGS[@]}" - - - name: React 👍 when the reconsider finishes - # Once the reconsider run has completed successfully, add a thumbs-up so - # the contributor sees the bot is done (the 👀 stays, signalling - # seen -> handled). `success()` keeps this from firing if the run - # errored, and the AGENT_SHIN_ENABLED gate keeps dry-run inert. - if: success() && steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true' - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - COMMENT_ID: ${{ github.event.comment.id }} - run: | - set -euo pipefail - gh api --method POST \ - -H "Accept: application/vnd.github+json" \ - "repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \ - -f content=+1 \ - || echo "::warning::failed to add 👍 reaction (non-fatal)" diff --git a/tests/test_litellm/test_github_close_low_quality_prs.py b/tests/test_litellm/test_github_close_low_quality_prs.py deleted file mode 100644 index 2a891ca72f5..00000000000 --- a/tests/test_litellm/test_github_close_low_quality_prs.py +++ /dev/null @@ -1,856 +0,0 @@ -"""Unit tests for `.github/scripts/close_low_quality_prs.py`. - -These exercise the pure logic (score extraction and per-PR evaluation) without -hitting GitHub. Network/CLI calls are stubbed via monkeypatch. -""" - -from __future__ import annotations - -import datetime as dt -import importlib.util -import sys -from pathlib import Path - -import pytest - -SCRIPT_PATH = ( - Path(__file__).resolve().parents[2] - / ".github" - / "scripts" - / "close_low_quality_prs.py" -) - - -@pytest.fixture(scope="module") -def closer_module(): - """Load the script as a module via its file path (it lives outside the package).""" - spec = importlib.util.spec_from_file_location("close_low_quality_prs", SCRIPT_PATH) - assert spec and spec.loader, f"Could not load spec for {SCRIPT_PATH}" - module = importlib.util.module_from_spec(spec) - sys.modules["close_low_quality_prs"] = module - spec.loader.exec_module(module) - return module - - -def _greptile_comment( - body: str, - updated_at: str = "2026-05-10T00:00:00Z", - login: str = "greptile-apps[bot]", -) -> dict: - return { - "user": {"login": login}, - "body": body, - "created_at": updated_at, - "updated_at": updated_at, - } - - -class TestExtractGreptileScore: - def test_should_extract_score_from_html_header(self, closer_module): - comments = [ - _greptile_comment("

Confidence Score: 3/5

\nSome body text.") - ] - result = closer_module.extract_greptile_score(comments) - assert result is not None - score, _ = result - assert score == 3 - - def test_should_accept_both_greptile_login_variants(self, closer_module): - # REST API form ("greptile-apps[bot]") and GraphQL form ("greptile-apps") - for login in ("greptile-apps", "greptile-apps[bot]"): - comments = [ - _greptile_comment("

Confidence Score: 2/5

", login=login) - ] - result = closer_module.extract_greptile_score(comments) - assert result is not None, f"failed to detect score for login={login}" - score, _ = result - assert score == 2 - - def test_should_extract_score_from_plain_text(self, closer_module): - comments = [_greptile_comment("Confidence Score: 5/5 — looks good!")] - result = closer_module.extract_greptile_score(comments) - assert result is not None - score, _ = result - assert score == 5 - - def test_should_tolerate_whitespace_and_case(self, closer_module): - comments = [_greptile_comment("**confidence score : 2 / 5**")] - result = closer_module.extract_greptile_score(comments) - assert result is not None - score, _ = result - assert score == 2 - - def test_should_pick_most_recent_comment_when_rereview_happens(self, closer_module): - comments = [ - _greptile_comment( - "Confidence Score: 2/5", updated_at="2026-05-01T00:00:00Z" - ), - _greptile_comment( - "Confidence Score: 5/5", updated_at="2026-05-12T00:00:00Z" - ), - ] - result = closer_module.extract_greptile_score(comments) - assert result is not None - score, _ = result - assert score == 5 - - def test_should_ignore_non_greptile_authors(self, closer_module): - comments = [ - { - "user": {"login": "some-human"}, - "body": "Confidence Score: 1/5", - "created_at": "2026-05-12T00:00:00Z", - "updated_at": "2026-05-12T00:00:00Z", - } - ] - assert closer_module.extract_greptile_score(comments) is None - - def test_should_return_none_when_no_score_present(self, closer_module): - comments = [_greptile_comment("Greptile summary without a score.")] - assert closer_module.extract_greptile_score(comments) is None - - def test_should_return_none_for_empty_comments(self, closer_module): - assert closer_module.extract_greptile_score([]) is None - - -class TestEvaluatePr: - @pytest.fixture(autouse=True) - def _now(self): - return dt.datetime(2026, 5, 17, tzinfo=dt.timezone.utc) - - def _make_pr( - self, - *, - number: int = 1, - created_days_ago: int = 10, - is_draft: bool = False, - labels: list[str] | None = None, - author_login: str = "mateo-berri", - ) -> dict: - created = dt.datetime(2026, 5, 17, tzinfo=dt.timezone.utc) - dt.timedelta( - days=created_days_ago - ) - return { - "number": number, - "title": f"PR #{number}", - "createdAt": created.isoformat().replace("+00:00", "Z"), - "isDraft": is_draft, - "labels": [{"name": lbl} for lbl in (labels or [])], - "author": {"login": author_login}, - "url": f"https://example.com/pr/{number}", - } - - @pytest.fixture(autouse=True) - def _external_author(self, closer_module, monkeypatch): - """Treat every test PR as external unless overridden.""" - monkeypatch.setattr( - closer_module, "is_external_pr_author", lambda pr, repo: True - ) - - def test_should_warn_drafts_when_score_low_first_time( - self, closer_module, _now, monkeypatch - ): - # Drafts are NOT a free pass — the open-PR queue should reflect any - # PR that needs human attention regardless of draft status. Authors - # who need a long-lived draft can use the `wip` opt-out label. - # First run: warn the contributor (1-day grace), don't close yet. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 2/5")], - ) - action, score, age = closer_module.evaluate_pr( - self._make_pr(is_draft=True, created_days_ago=0), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "warn-grace" - assert score == 2 and age == 0 - - def test_should_warn_brand_new_pr_when_min_age_zero( - self, closer_module, _now, monkeypatch - ): - # `min_age_days=0` means no age filter — a freshly-opened PR is - # eligible the moment Greptile scores it below threshold. The - # first detection still goes through the warn-grace step rather - # than closing immediately, giving the contributor 2 hours to - # respond before the next run actually closes the PR. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 1/5")], - ) - action, score, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=0), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "warn-grace" - assert score == 1 and age == 0 - - def test_should_skip_optout_label_case_insensitive( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: pytest.fail("should not fetch comments for opt-outs"), - ) - action, _, _ = closer_module.evaluate_pr( - self._make_pr(labels=["WIP"]), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels={"wip"}, - ) - assert action == "skip-optout-label" - - def test_should_skip_too_young_when_min_age_set( - self, closer_module, _now, monkeypatch - ): - # The min-age-days flag is now opt-in (default 0). When a maintainer - # explicitly passes a positive value (e.g. for a backfill run that - # wants to spare brand-new PRs), the skip-too-young path still works. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: pytest.fail("should not fetch comments for young PRs"), - ) - action, _, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=2), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-too-young" - assert age == 2 - - def test_should_not_skip_when_min_age_is_zero( - self, closer_module, _now, monkeypatch - ): - # With the new default min_age_days=0, even a 0-day-old PR is - # evaluated. This test pins that behavior so future refactors don't - # silently restore an age filter. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 5/5")], - ) - action, score, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=0), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-score-ok" - assert score == 5 and age == 0 - - def test_should_skip_when_greptile_has_not_reviewed( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr(closer_module, "fetch_pr_comments", lambda *a, **kw: []) - action, score, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=10), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-no-greptile-score" - assert score is None and age == 10 - - def test_should_skip_when_score_meets_threshold( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 4/5")], - ) - action, score, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=10), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-score-ok" - assert score == 4 and age == 10 - - def test_should_warn_when_old_and_low_score_no_prior_warning( - self, closer_module, _now, monkeypatch - ): - # Even an old PR that still has no grace warning gets one on the - # first eligible run — the daily cron is the natural cadence, so - # an existing-but-never-warned PR enters the grace flow normally. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 3/5")], - ) - action, score, age = closer_module.evaluate_pr( - self._make_pr(created_days_ago=10), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "warn-grace" - assert score == 3 and age == 10 - - def test_should_close_when_grace_warning_aged_out_and_score_still_low( - self, closer_module, _now, monkeypatch - ): - # Day-1 the closer posted a warning. Day-2 the PR still scores <4 - # AND the warning is older than `GRACE_PERIOD_SECONDS`, so the - # action flips to `close`. This is the "grace expired" path. - old_warning = { - "user": {"login": "github-actions[bot]"}, - "body": ( - "you have 2 hours to fix this\n\n" + closer_module.GRACE_COMMENT_MARKER - ), - "created_at": ( - _now - dt.timedelta(seconds=closer_module.GRACE_PERIOD_SECONDS + 60) - ) - .isoformat() - .replace("+00:00", "Z"), - "updated_at": "2026-05-15T00:00:00Z", - } - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [ - _greptile_comment( - "

Confidence Score: 1/5

", - updated_at="2026-05-15T00:00:00Z", - ), - old_warning, - ], - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(created_days_ago=14), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "close" - assert score == 1 - - def test_should_skip_when_grace_warning_within_window( - self, closer_module, _now, monkeypatch - ): - # Within the 2-hour grace window the closer must NOT close the - # PR even if the score is still low. The warning is only an hour - # old; give the contributor time to push fixes before destruction. - recent_warning = { - "user": {"login": "github-actions[bot]"}, - "body": "warning text\n\n" + closer_module.GRACE_COMMENT_MARKER, - "created_at": (_now - dt.timedelta(hours=1)) - .isoformat() - .replace("+00:00", "Z"), - } - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [ - _greptile_comment("Confidence Score: 2/5"), - recent_warning, - ], - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(created_days_ago=10), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-in-grace-period" - assert score == 2 - - def test_should_warn_grace_for_swiftwinds_not_close_immediately( - self, closer_module, _now, monkeypatch - ): - # Regression: SwiftWinds (the dogfood account) used to be in a - # now-removed `IMMEDIATE_CLOSE_LOGINS` bypass that closed on first - # detection. It must now follow the SAME grace path as every other - # external author: warn first, close only after the window elapses. - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 1/5")], - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(created_days_ago=0, author_login="SwiftWinds"), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "warn-grace" - assert score == 1 - - def test_should_skip_internal_authors(self, closer_module, _now, monkeypatch): - # Override the fixture for this one test. - monkeypatch.setattr( - closer_module, "is_external_pr_author", lambda pr, repo: False - ) - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: pytest.fail("should not fetch comments for internal"), - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(created_days_ago=14, author_login="krrishdholakia"), - now=_now, - min_age_days=7, - min_score=4, - repo=None, - optout_labels=set(), - allowlist=frozenset(), - ) - assert action == "skip-internal" - assert score is None - - -class TestMainOptoutLabelDefault: - """`--optout-label` must REPLACE the canonical defaults, not append.""" - - def _patch_no_op(self, closer_module, monkeypatch): - monkeypatch.setattr(closer_module, "fetch_open_prs", lambda repo: []) - # `optout_labels` is captured indirectly via evaluate_pr; sniff the - # set passed in by stubbing evaluate_pr. - captured: dict = {} - - def fake_evaluate(pr, now, min_age_days, min_score, repo, optout_labels): - captured["optout_labels"] = set(optout_labels) - return ("skip-internal", None, None) - - monkeypatch.setattr(closer_module, "evaluate_pr", fake_evaluate) - return captured - - def test_should_use_canonical_defaults_when_flag_omitted( - self, closer_module, monkeypatch - ): - captured = self._patch_no_op(closer_module, monkeypatch) - # No PRs -> capture won't fire; instead inject one synthetic PR via - # fetch_open_prs so evaluate_pr is invoked at least once. - monkeypatch.setattr( - closer_module, - "fetch_open_prs", - lambda repo: [ - { - "number": 1, - "title": "p", - "createdAt": "2026-05-10T00:00:00Z", - "isDraft": True, - "labels": [], - "author": {"login": "x"}, - } - ], - ) - monkeypatch.setattr(sys, "argv", ["close_low_quality_prs.py"]) - rc = closer_module.main() - assert rc == 0 - assert captured["optout_labels"] == set(closer_module.DEFAULT_OPTOUT_LABELS) - - def test_should_replace_defaults_when_flag_provided( - self, closer_module, monkeypatch - ): - captured = self._patch_no_op(closer_module, monkeypatch) - monkeypatch.setattr( - closer_module, - "fetch_open_prs", - lambda repo: [ - { - "number": 1, - "title": "p", - "createdAt": "2026-05-10T00:00:00Z", - "isDraft": True, - "labels": [], - "author": {"login": "x"}, - } - ], - ) - monkeypatch.setattr( - sys, - "argv", - [ - "close_low_quality_prs.py", - "--optout-label", - "hold", - "--optout-label", - "needs-discussion", - ], - ) - rc = closer_module.main() - assert rc == 0 - # Crucially, none of the canonical defaults leak in. - assert captured["optout_labels"] == {"hold", "needs-discussion"} - for default in closer_module.DEFAULT_OPTOUT_LABELS: - assert default not in captured["optout_labels"], default - - -class TestSecondsSinceLastGraceWarning: - """Grace-period detection: only counts comments by the bot identity - that contain the shared `GRACE_COMMENT_MARKER`.""" - - def _make_marker_comment( - self, - closer_module, - *, - login: str = "github-actions[bot]", - created_at: str = "2026-05-16T00:00:00Z", - include_marker: bool = True, - ) -> dict: - body = "warning text" - if include_marker: - body += "\n\n" + closer_module.GRACE_COMMENT_MARKER - return { - "user": {"login": login}, - "body": body, - "created_at": created_at, - } - - def test_should_return_none_when_no_marker_comment(self, closer_module): - comments = [ - { - "user": {"login": "github-actions[bot]"}, - "body": "Some other bot comment", - "created_at": "2026-05-16T00:00:00Z", - } - ] - assert closer_module.seconds_since_last_grace_warning(comments) is None - - def test_should_return_none_for_empty(self, closer_module): - assert closer_module.seconds_since_last_grace_warning([]) is None - - def test_should_ignore_non_bot_comments_with_marker(self, closer_module): - # If a curious user quotes the marker in a comment, we must NOT - # treat it as a bot warning. The grace timer would then never fire. - comments = [ - self._make_marker_comment(closer_module, login="random-user"), - ] - assert closer_module.seconds_since_last_grace_warning(comments) is None - - def test_should_pick_latest_marker_comment(self, closer_module): - # When multiple grace warnings exist (e.g. a re-open cycle), use - # the most recent one to compute the age. - comments = [ - self._make_marker_comment(closer_module, created_at="2026-05-15T00:00:00Z"), - self._make_marker_comment(closer_module, created_at="2026-05-16T23:00:00Z"), - ] - now = dt.datetime(2026, 5, 17, 0, 0, 0, tzinfo=dt.timezone.utc) - age = closer_module.seconds_since_last_grace_warning(comments, now=now) - # 1h = 3600s - assert age == 3600.0 - - -class TestGraceWarningCommentText: - """Pin the user-facing language in the grace warning comment so the - grace-window and `@greptileai still works after close` promises - don't get accidentally dropped in a future refactor. - """ - - def test_should_state_grace_window(self, closer_module): - body = closer_module.format_grace_warning_comment(score=2, threshold=4) - # The user's PR explicitly said "specify in the comment" — pin - # that the grace window appears in the comment. - assert "2 hours" in body - - def test_should_mention_agent_shin_reconsider(self, closer_module): - body = closer_module.format_grace_warning_comment(score=2, threshold=4) - assert "@agent-shin reconsider" in body - - def test_should_promise_greptileai_works_after_close(self, closer_module): - body = closer_module.format_grace_warning_comment(score=2, threshold=4) - assert "@greptileai" in body - assert "even after the PR is closed" in body - - def test_should_carry_grace_marker(self, closer_module): - # The marker is what `seconds_since_last_grace_warning` greps for - # to detect a prior warning — dropping it would silently break - # the cooldown. - body = closer_module.format_grace_warning_comment(score=2, threshold=4) - assert closer_module.GRACE_COMMENT_MARKER in body - - def test_close_comment_should_mention_greptileai_post_close(self, closer_module): - # The close comment should ALSO point at the @greptileai post-close - # re-review path so contributors see the same options whether they - # read the warning or only catch the close comment. - body = closer_module.format_close_comment(score=2, threshold=4) - assert "@greptileai" in body - assert "even after the PR is closed" in body - - def test_close_comment_should_advertise_reconsider(self, closer_module): - body = closer_module.format_close_comment(score=2, threshold=4) - assert "@agent-shin reconsider" in body - - def test_close_comment_should_carry_agent_shin_close_marker(self, closer_module): - # The close comment advertises `@agent-shin reconsider`, and the - # reconsider reopen guard (`was_closed_by_agent_shin`) only treats a - # PR as Agent-Shin-closed when the close comment carries this marker. - # Dropping it silently breaks the advertised recovery path for every - # PR closed by this daily sweep. - body = closer_module.format_close_comment(score=2, threshold=4) - assert closer_module.AGENT_SHIN_CLOSE_MARKER in body - - def test_close_comment_should_state_score_and_threshold(self, closer_module): - body = closer_module.format_close_comment(score=1, threshold=4) - assert "1/5" in body - assert "4/5" in body - - -class TestHasOptoutLabel: - def test_should_match_label_case_insensitively(self, closer_module): - pr = {"labels": [{"name": "Do Not Close"}, {"name": "bug"}]} - assert closer_module.has_optout_label(pr, {"do not close"}) is True - - def test_should_return_false_when_no_match(self, closer_module): - pr = {"labels": [{"name": "bug"}, {"name": "enhancement"}]} - assert closer_module.has_optout_label(pr, {"wip", "keep open"}) is False - - def test_should_handle_missing_labels(self, closer_module): - assert closer_module.has_optout_label({}, {"wip"}) is False - - -class TestListOpenItemsNoCap: - """The bulk sweeps must fetch the ENTIRE open backlog. - - Regression guard for the old hard-coded ``--limit 1000``: gh lists - newest-first, so a low cap silently dropped the *oldest* PRs/issues — - exactly the stale ones a low-quality sweep exists to catch. - """ - - @staticmethod - def _shared(closer_module): - # `closer_module` loading puts `.github/scripts` on sys.path and - # imports agent_shin_shared, so it's already in sys.modules. - import agent_shin_shared - - return agent_shin_shared - - def _capture_gh_args(self, closer_module, monkeypatch, *, returns="[]"): - shared = self._shared(closer_module) - captured: dict = {} - - def fake_gh(*args): - captured["args"] = args - return returns - - # `list_open_items` looks up `gh` in agent_shin_shared's namespace. - monkeypatch.setattr(shared, "gh", fake_gh) - return shared, captured - - def test_list_open_items_passes_no_cap_limit_not_1000( - self, closer_module, monkeypatch - ): - shared, captured = self._capture_gh_args(closer_module, monkeypatch) - shared.list_open_items("pr", repo="o/r", fields="number,title") - args = captured["args"] - assert "--limit" in args - limit_value = args[args.index("--limit") + 1] - assert limit_value == str(shared.GH_LIST_ALL_LIMIT) - assert limit_value != "1000" - # A meaningful ceiling: comfortably above any realistic open backlog. - assert shared.GH_LIST_ALL_LIMIT >= 100_000 - - def test_list_open_items_uses_dedicated_command_state_and_fields( - self, closer_module, monkeypatch - ): - shared, captured = self._capture_gh_args(closer_module, monkeypatch) - shared.list_open_items("issue", repo="o/r", fields="number") - args = captured["args"] - assert args[0] == "issue" and args[1] == "list" - assert args[args.index("--state") + 1] == "open" - assert args[args.index("--json") + 1] == "number" - assert tuple(args[-2:]) == ("--repo", "o/r") - - def test_list_open_items_omits_repo_when_none(self, closer_module, monkeypatch): - shared, captured = self._capture_gh_args(closer_module, monkeypatch) - shared.list_open_items("pr", repo=None, fields="number") - assert "--repo" not in captured["args"] - - def test_list_open_items_parses_json_array(self, closer_module, monkeypatch): - shared, _ = self._capture_gh_args( - closer_module, monkeypatch, returns='[{"number": 1}, {"number": 2}]' - ) - items = shared.list_open_items("pr", repo=None, fields="number") - assert [i["number"] for i in items] == [1, 2] - - def test_list_open_items_rejects_unknown_kind(self, closer_module): - shared = self._shared(closer_module) - with pytest.raises(ValueError, match="kind must be 'pr' or 'issue', got 'both"): - shared.list_open_items("both", repo="o/r", fields="number") - - def test_fetch_open_prs_delegates_with_no_cap(self, closer_module, monkeypatch): - shared, captured = self._capture_gh_args(closer_module, monkeypatch) - closer_module.fetch_open_prs("o/r") - args = captured["args"] - assert args[0] == "pr" - assert args[args.index("--limit") + 1] == str(shared.GH_LIST_ALL_LIMIT) - # Still requests every field downstream evaluate_pr / labels logic needs. - assert "createdAt" in args[args.index("--json") + 1] - - -class TestEvaluatePrAllowlist: - """While the dogfood allowlist is active `evaluate_pr` only acts on the - named accounts and bypasses the external-only restriction for them. - Emptying it restores the internal-author skip.""" - - @pytest.fixture(autouse=True) - def _now(self): - return dt.datetime(2026, 5, 17, tzinfo=dt.timezone.utc) - - def _make_pr(self, *, author_login: str, created_days_ago: int = 10) -> dict: - created = dt.datetime(2026, 5, 17, tzinfo=dt.timezone.utc) - dt.timedelta( - days=created_days_ago - ) - return { - "number": 1, - "title": "PR #1", - "createdAt": created.isoformat().replace("+00:00", "Z"), - "isDraft": False, - "labels": [], - "author": {"login": author_login}, - "url": "https://example.com/pr/1", - } - - def test_should_skip_author_not_on_allowlist( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: pytest.fail("must not fetch comments for non-allowlisted"), - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(author_login="random-oss-dev"), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "skip-not-allowlisted" - assert score is None - - def test_should_act_on_allowlisted_internal_author( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr( - closer_module, "is_external_pr_author", lambda pr, repo: False - ) - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [_greptile_comment("Confidence Score: 2/5")], - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(author_login="mateo-berri", created_days_ago=0), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - ) - assert action == "warn-grace" - assert score == 2 - - def test_empty_allowlist_restores_internal_skip( - self, closer_module, _now, monkeypatch - ): - monkeypatch.setattr( - closer_module, "is_external_pr_author", lambda pr, repo: False - ) - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: pytest.fail("must not fetch comments for internal"), - ) - action, score, _ = closer_module.evaluate_pr( - self._make_pr(author_login="krrishdholakia"), - now=_now, - min_age_days=0, - min_score=4, - repo=None, - optout_labels=set(), - allowlist=frozenset(), - ) - assert action == "skip-internal" - - def test_allowlist_constant_is_the_two_dogfood_accounts(self, closer_module): - assert closer_module.ALLOWLIST_LOGINS == frozenset( - {"mateo-berri", "swiftwinds"} - ) - - -class TestDryRunGateOnClose: - """Regression: the daily sweep is dry-run unless `--close` is passed - (the workflow only adds it when `AGENT_SHIN_ENABLED=true`). A closeable - PR (low score, grace window elapsed) must be DETECTED and reported as - "would close", but the dry run must never make a real GitHub mutation, - so merging Agent Shin stays inert by default.""" - - def _closeable_pr(self) -> dict: - return { - "number": 7, - "title": "thin PR", - "createdAt": "2026-05-10T00:00:00Z", - "isDraft": False, - "labels": [], - "author": {"login": "SwiftWinds"}, - "url": "https://example.com/pr/7", - } - - def test_dry_run_sweep_detects_but_does_not_close( - self, closer_module, monkeypatch, capsys - ): - aged_out_warning = { - "user": {"login": "github-actions[bot]"}, - "body": "warned\n\n" + closer_module.GRACE_COMMENT_MARKER, - # Far enough in the past that it's aged out regardless of - # GRACE_PERIOD_SECONDS, since main() pins `now` to real time. - "created_at": "2020-01-01T00:00:00Z", - } - monkeypatch.setattr( - closer_module, "fetch_open_prs", lambda repo: [self._closeable_pr()] - ) - monkeypatch.setattr( - closer_module, - "fetch_pr_comments", - lambda *a, **kw: [ - _greptile_comment("Confidence Score: 1/5"), - aged_out_warning, - ], - ) - # Any real GitHub mutation during a dry run is the bug under test. - monkeypatch.setattr( - closer_module, - "gh", - lambda *a, **kw: pytest.fail(f"dry run must not call gh: {a}"), - ) - monkeypatch.setattr(sys, "argv", ["close_low_quality_prs.py"]) - - rc = closer_module.main() - - assert rc == 0 - # The PR is detected as closeable, just not acted on. - assert "Total would close: 1" in capsys.readouterr().out diff --git a/tests/test_litellm/test_github_review_gate.py b/tests/test_litellm/test_github_review_gate.py deleted file mode 100644 index 001fa8f43f5..00000000000 --- a/tests/test_litellm/test_github_review_gate.py +++ /dev/null @@ -1,524 +0,0 @@ -"""Unit tests for the `ready for review` label lifecycle (Agent Shin review gate). - -Exercises `triage_with_llm.review_gate`, the state machine that keeps the -`ready for review` label in sync with whether a PR clears both the LLM rubric -and Greptile's confidence score: - - * pass (untagged) -> add label + "ready for review" comment - * pass (untagged, recovered) -> add label + "all clear again" comment - * pass (already tagged) -> noop - * regress (tagged) -> remove label + "what's missing" comment, stays open - * fail (untagged, within 24h)-> one-time "what's missing" notice - * fail (untagged, >24h) -> close + comment - * dry run (close=False) -> would-* previews, no side effects -""" - -from __future__ import annotations - -import datetime as dt -import importlib.util -import sys -from pathlib import Path - -import pytest - -SCRIPT_PATH = ( - Path(__file__).resolve().parents[2] / ".github" / "scripts" / "triage_with_llm.py" -) - -NOW = dt.datetime(2026, 5, 24, 12, 0, 0, tzinfo=dt.timezone.utc) -JUST_NOW = "2026-05-24T11:00:00Z" # 1h old -> within 24h grace -TWO_DAYS_AGO = "2026-05-22T11:00:00Z" # >24h old -> past grace - - -@pytest.fixture(scope="module") -def triage_module(): - spec = importlib.util.spec_from_file_location("triage_with_llm", SCRIPT_PATH) - assert spec and spec.loader - module = importlib.util.module_from_spec(spec) - sys.modules["triage_with_llm"] = module - spec.loader.exec_module(module) - return module - - -class _Recorder: - """Captures every gh mutation review_gate could fire, and fails loudly - on the ones a given scenario forbids.""" - - def __init__(self, triage_module, monkeypatch): - self.comments: list[str] = [] - self.added: list[str] = [] - self.removed: list[str] = [] - self.closed: list[int] = [] - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: self.comments.append(body), - ) - monkeypatch.setattr( - triage_module, - "add_label", - lambda repo, n, label: self.added.append(label), - ) - monkeypatch.setattr( - triage_module, - "remove_label", - lambda repo, n, label: self.removed.append(label), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda repo, n: self.closed.append(n), - ) - - -def _make_pr(**overrides): - base = { - "number": 7, - "title": "feat: do a thing", - "body": "some body without a linked issue or QA proof", - "state": "open", - "author_association": "NONE", - "user": {"login": "mateo-berri"}, - "labels": [], - "created_at": JUST_NOW, - } - base.update(overrides) - return base - - -def _pass(prompt): - return '{"verdict": "pass", "missing": [], "explanation": "looks good"}' - - -def _fail(prompt): - return ( - '{"verdict": "fail", "missing": ["QA proof", "expected vs. actual"],' - ' "explanation": "thin description"}' - ) - - -def _gate(triage_module, **kwargs): - """Call review_gate with safe defaults for the injectable hooks.""" - params = dict( - repo="o/r", - number=7, - close=True, - model="m", - judge=_pass, - greptile_score=None, - comments=[], - now=NOW, - ) - params.update(kwargs) - return triage_module.review_gate(**params) - - -class TestReviewGatePass: - def test_pass_untagged_adds_label_and_ready_comment( - self, triage_module, monkeypatch - ): - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: _make_pr()) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_pass, greptile_score=5) - - assert result["action"] == "labeled-ready" - assert rec.added == [triage_module.READY_FOR_REVIEW_LABEL] - assert rec.removed == [] and rec.closed == [] - assert len(rec.comments) == 1 - assert "ready for review" in rec.comments[0].lower() - assert triage_module.READY_MARKER in rec.comments[0] - assert "5/5" in rec.comments[0] - - def test_pass_already_tagged_is_noop(self, triage_module, monkeypatch): - pr = _make_pr(labels=[{"name": "ready for review"}]) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_pass, greptile_score=5) - - assert result["action"] == "noop-passing" - assert rec.added == [] and rec.removed == [] and rec.comments == [] - - def test_pass_after_prior_regression_uses_all_clear_wording( - self, triage_module, monkeypatch - ): - # A regression marker in history -> this is a recovery, not a first pass. - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: _make_pr()) - rec = _Recorder(triage_module, monkeypatch) - prior = [ - { - "user": {"login": "github-actions[bot]"}, - "body": triage_module.REGRESSED_MARKER, - } - ] - - result = _gate(triage_module, judge=_pass, greptile_score=5, comments=prior) - - assert result["action"] == "labeled-ready" - assert "all clear" in rec.comments[0].lower() - - def test_linked_issue_passes_without_calling_judge( - self, triage_module, monkeypatch - ): - pr = _make_pr(body="Fixes #4321\n\nbody") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate( - triage_module, - judge=lambda p: pytest.fail("LLM must not be called for linked issue"), - greptile_score=5, - ) - assert result["action"] == "labeled-ready" - assert rec.added == [triage_module.READY_FOR_REVIEW_LABEL] - - -class TestReviewGateRegression: - def test_regression_removes_label_and_keeps_pr_open( - self, triage_module, monkeypatch - ): - pr = _make_pr(labels=[{"name": "ready for review"}]) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_fail, greptile_score=5) - - assert result["action"] == "label-removed-regressed" - assert rec.removed == [triage_module.READY_FOR_REVIEW_LABEL] - assert rec.closed == [] # regression NEVER closes the PR - assert triage_module.REGRESSED_MARKER in rec.comments[0] - assert "QA proof" in rec.comments[0] - # The state machine closes a still-failing PR `grace_days` after this - # notice (default 24h); the comment must disclose that deadline rather - # than implying the PR stays open indefinitely. - assert "24 hours" in rec.comments[0] - assert "auto-closed" in rec.comments[0] - - def test_regression_comment_discloses_grace_deadline(self, triage_module): - one_day = triage_module.format_regression_comment( - ["QA proof"], "needs work", grace_days=1 - ) - assert "24 hours" in one_day - assert "auto-closed" in one_day - - three_days = triage_module.format_regression_comment( - ["QA proof"], "needs work", grace_days=3 - ) - assert "3 days" in three_days - assert "auto-closed" in three_days - - def test_greptile_drop_alone_triggers_regression(self, triage_module, monkeypatch): - # Rubric still passes, but Greptile fell to 2/5 -> not passing. - pr = _make_pr(labels=[{"name": "ready for review"}]) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_pass, greptile_score=2) - - assert result["action"] == "label-removed-regressed" - assert rec.removed == [triage_module.READY_FOR_REVIEW_LABEL] - assert "2/5" in rec.comments[0] - - def test_greptile_score_read_from_comments_when_not_injected( - self, triage_module, monkeypatch - ): - pr = _make_pr(labels=[{"name": "ready for review"}]) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - greptile = [ - { - "user": {"login": "greptile-apps[bot]"}, - "body": "Confidence Score: 2/5", - "created_at": "2026-05-24T10:00:00Z", - } - ] - - result = _gate( - triage_module, - judge=_pass, - greptile_score=triage_module._UNSET, - comments=greptile, - ) - assert result["action"] == "label-removed-regressed" - assert "2/5" in rec.comments[0] - - -class TestReviewGateGraceAndClose: - def test_within_grace_posts_one_time_notice(self, triage_module, monkeypatch): - monkeypatch.setattr( - triage_module, "fetch_pr", lambda repo, n: _make_pr(created_at=JUST_NOW) - ) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_fail, greptile_score=None) - - assert result["action"] == "within-grace-notified" - assert rec.closed == [] and rec.added == [] and rec.removed == [] - assert triage_module.WITHIN_GRACE_MARKER in rec.comments[0] - assert "QA proof" in rec.comments[0] - - def test_within_grace_does_not_double_notify(self, triage_module, monkeypatch): - monkeypatch.setattr( - triage_module, "fetch_pr", lambda repo, n: _make_pr(created_at=JUST_NOW) - ) - rec = _Recorder(triage_module, monkeypatch) - prior = [ - { - "user": {"login": "github-actions[bot]"}, - "body": triage_module.WITHIN_GRACE_MARKER, - } - ] - - result = _gate(triage_module, judge=_fail, greptile_score=None, comments=prior) - - assert result["action"] == "within-grace-already-notified" - assert rec.comments == [] - - def test_past_grace_closes_with_comment(self, triage_module, monkeypatch): - monkeypatch.setattr( - triage_module, - "fetch_pr", - lambda repo, n: _make_pr(created_at=TWO_DAYS_AGO), - ) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, judge=_fail, greptile_score=None) - - assert result["action"] == "closed" - assert rec.closed == [7] - assert len(rec.comments) == 1 - # The close comment must carry the reconsider provenance marker so - # `was_closed_by_agent_shin` can later recognize this as an Agent Shin - # close (and not some other workflow's `github-actions[bot]` close). - assert triage_module.AGENT_SHIN_CLOSE_MARKER in rec.comments[0] - - def test_recent_regression_marker_blocks_close(self, triage_module, monkeypatch): - """A failing PR with a fresh regression notice must NOT be closed — - the contributor needs a window to address the regression.""" - monkeypatch.setattr( - triage_module, - "fetch_pr", - lambda repo, n: _make_pr(created_at=TWO_DAYS_AGO), - ) - rec = _Recorder(triage_module, monkeypatch) - prior = [ - { - "user": {"login": "github-actions[bot]"}, - "body": triage_module.REGRESSED_MARKER, - # Posted just an hour before NOW -> well inside grace_days. - "created_at": "2026-05-24T11:00:00Z", - } - ] - - result = _gate(triage_module, judge=_fail, greptile_score=None, comments=prior) - - assert result["action"] == "regressed-already-notified" - assert rec.closed == [] and rec.comments == [] - - def test_stale_regression_marker_allows_close(self, triage_module, monkeypatch): - """Once grace_days have elapsed since the regression notice, the - review gate must let the close path fire — otherwise PRs that were - regressed and then abandoned stay open forever.""" - monkeypatch.setattr( - triage_module, - "fetch_pr", - lambda repo, n: _make_pr(created_at=TWO_DAYS_AGO), - ) - rec = _Recorder(triage_module, monkeypatch) - prior = [ - { - "user": {"login": "github-actions[bot]"}, - "body": triage_module.REGRESSED_MARKER, - # Posted 30 days before NOW -> well past the default 1-day grace. - "created_at": "2026-04-24T11:00:00Z", - } - ] - - result = _gate(triage_module, judge=_fail, greptile_score=None, comments=prior) - - assert result["action"] == "closed" - assert rec.closed == [7] - assert len(rec.comments) == 1 - - def test_linked_issue_with_greptile_fail_uses_greptile_explanation( - self, triage_module, monkeypatch - ): - """When the rubric short-circuits to pass (linked-issue regex) but - Greptile dragged the PR under the bar, the close comment's - explanation must describe the Greptile shortfall, not the - misleading "LLM was not called" rubric placeholder.""" - pr = _make_pr(body="Fixes #4321\n\nbody", created_at=TWO_DAYS_AGO) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate( - triage_module, - judge=lambda p: pytest.fail("LLM must not be called for linked issue"), - greptile_score=2, - ) - - assert result["action"] == "closed" - assert len(rec.comments) == 1 - body = rec.comments[0] - assert "LLM was not called" not in body - assert "Greptile" in body and "2/5" in body - - -class TestReviewGateDryRun: - @pytest.mark.parametrize( - "scenario,labels,judge,score,created,expected", - [ - ("pass", [], _pass, 5, JUST_NOW, "would-label-ready"), - ( - "regress", - [{"name": "ready for review"}], - _fail, - 5, - JUST_NOW, - "would-remove-label", - ), - ("within-grace", [], _fail, None, JUST_NOW, "would-notify-within-grace"), - ("past-grace", [], _fail, None, TWO_DAYS_AGO, "would-close"), - ], - ) - def test_dry_run_previews_without_side_effects( - self, - triage_module, - monkeypatch, - scenario, - labels, - judge, - score, - created, - expected, - ): - pr = _make_pr(labels=labels, created_at=created) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - - result = _gate(triage_module, close=False, judge=judge, greptile_score=score) - - assert result["action"] == expected - # Dry run touches nothing. - assert rec.added == [] and rec.removed == [] and rec.closed == [] - assert rec.comments == [] - assert "comment" in result # preview body still surfaced - - -class TestReviewGateGuards: - def test_skips_internal_author(self, triage_module, monkeypatch): - pr = _make_pr(author_association="MEMBER", user={"login": "krrish"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = _gate( - triage_module, - judge=lambda p: pytest.fail("no LLM for internal"), - allowlist=frozenset(), - ) - assert result["action"] == "skip-internal-author" - - def test_skips_closed_pr(self, triage_module, monkeypatch): - pr = _make_pr(state="closed") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = _gate(triage_module, judge=lambda p: pytest.fail("no LLM for closed")) - assert result["action"] == "skip-not-open" - - def test_llm_error_is_non_destructive(self, triage_module, monkeypatch): - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: _make_pr()) - rec = _Recorder(triage_module, monkeypatch) - - def boom(prompt): - raise RuntimeError("api down") - - result = _gate(triage_module, judge=boom, greptile_score=None) - - assert result["action"] == "skip-llm-error" - assert rec.closed == [] and rec.added == [] and rec.removed == [] - - def test_full_recovery_cycle(self, triage_module, monkeypatch): - """pass -> regress -> recover, threading labels/comments like GitHub would.""" - state = {"labels": [], "comments": []} - - def fake_fetch(repo, n): - return _make_pr(labels=list(state["labels"]), created_at=JUST_NOW) - - monkeypatch.setattr(triage_module, "fetch_pr", fake_fetch) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: state["comments"].append( - {"user": {"login": "github-actions[bot]"}, "body": body} - ), - ) - monkeypatch.setattr( - triage_module, - "add_label", - lambda repo, n, label: state["labels"].append({"name": label}), - ) - monkeypatch.setattr( - triage_module, - "remove_label", - lambda repo, n, label: state["labels"].clear(), - ) - monkeypatch.setattr( - triage_module, "close_pr", lambda repo, n: pytest.fail("must not close") - ) - - # 1) passes -> tagged - r1 = _gate( - triage_module, judge=_pass, greptile_score=5, comments=state["comments"] - ) - assert r1["action"] == "labeled-ready" - assert any(lbl["name"] == "ready for review" for lbl in state["labels"]) - - # 2) regresses -> tag removed, comment posted, PR still open - r2 = _gate( - triage_module, judge=_fail, greptile_score=2, comments=state["comments"] - ) - assert r2["action"] == "label-removed-regressed" - assert state["labels"] == [] - - # 3) fixed again -> "all clear" + tag back - r3 = _gate( - triage_module, judge=_pass, greptile_score=5, comments=state["comments"] - ) - assert r3["action"] == "labeled-ready" - assert any(lbl["name"] == "ready for review" for lbl in state["labels"]) - assert "all clear" in state["comments"][-1]["body"].lower() - - -class TestReviewGateAllowlist: - """While the dogfood allowlist is active it is the sole author gate: - only the named accounts pass, and for them the internal-author exemption - is bypassed. Emptying it restores the normal internal-author skip.""" - - def test_should_skip_author_not_on_allowlist(self, triage_module, monkeypatch): - pr = _make_pr(user={"login": "random-oss-dev"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - result = _gate( - triage_module, judge=lambda p: pytest.fail("no LLM for non-allowlisted") - ) - assert result["action"] == "skip-not-allowlisted" - assert rec.added == [] and rec.comments == [] and rec.closed == [] - - def test_should_act_on_allowlisted_internal_author( - self, triage_module, monkeypatch - ): - pr = _make_pr(author_association="MEMBER", user={"login": "mateo-berri"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - rec = _Recorder(triage_module, monkeypatch) - result = _gate(triage_module, judge=_pass, greptile_score=5) - assert result["action"] == "labeled-ready" - assert rec.added == [triage_module.READY_FOR_REVIEW_LABEL] - - def test_empty_allowlist_restores_internal_skip(self, triage_module, monkeypatch): - pr = _make_pr(author_association="MEMBER", user={"login": "krrish"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = _gate( - triage_module, - judge=lambda p: pytest.fail("no LLM for internal"), - allowlist=frozenset(), - ) - assert result["action"] == "skip-internal-author" diff --git a/tests/test_litellm/test_github_triage_with_llm.py b/tests/test_litellm/test_github_triage_with_llm.py deleted file mode 100644 index ddffb978b48..00000000000 --- a/tests/test_litellm/test_github_triage_with_llm.py +++ /dev/null @@ -1,2134 +0,0 @@ -"""Unit tests for `.github/scripts/triage_with_llm.py` (Agent Shin).""" - -from __future__ import annotations - -import importlib.util -import json -import sys -from pathlib import Path - -import pytest - -SCRIPT_PATH = ( - Path(__file__).resolve().parents[2] / ".github" / "scripts" / "triage_with_llm.py" -) - - -@pytest.fixture(scope="module") -def triage_module(): - spec = importlib.util.spec_from_file_location("triage_with_llm", SCRIPT_PATH) - assert spec and spec.loader - module = importlib.util.module_from_spec(spec) - sys.modules["triage_with_llm"] = module - spec.loader.exec_module(module) - return module - - -class TestIsInternalContributor: - @pytest.mark.parametrize("association", ["OWNER", "MEMBER", "COLLABORATOR"]) - def test_should_mark_org_associations_as_internal(self, triage_module, association): - item = { - "author_association": association, - "user": {"login": "krrishdholakia"}, - } - assert triage_module.is_internal_contributor(item) is True - - @pytest.mark.parametrize( - "association", - ["CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR", "FIRST_TIMER", "NONE"], - ) - def test_should_mark_outside_associations_as_external( - self, triage_module, association - ): - item = { - "author_association": association, - "user": {"login": "random-oss-dev"}, - } - assert triage_module.is_internal_contributor(item) is False - - @pytest.mark.parametrize( - "item", - [ - {"author_association": "", "user": {"login": "random-oss-dev"}}, - {"user": {"login": "random-oss-dev"}}, # association field absent - ], - ) - def test_should_fail_safe_when_author_association_is_missing( - self, triage_module, item - ): - # Fail-safe: an empty/missing association must never make a PR - # eligible for the destructive close path. Treat as internal (skip). - assert triage_module.is_internal_contributor(item) is True - - @pytest.mark.parametrize( - "login", - ["dependabot[bot]", "greptile-apps[bot]", "dependabot", "github-actions"], - ) - def test_should_skip_bot_accounts_regardless_of_association( - self, triage_module, login - ): - item = {"author_association": "NONE", "user": {"login": login}} - assert triage_module.is_internal_contributor(item) is True - - -class TestHasLinkedIssue: - @pytest.mark.parametrize( - "body", - [ - "Fixes #1234", - "closes #1", - "Resolves #99", - "fix #42 — this addresses the regression", - "Closes https://github.com/BerriAI/litellm/issues/27000", - "Resolved https://github.com/BerriAI/litellm/issues/27001", - ], - ) - def test_should_detect_common_link_phrases(self, triage_module, body): - assert triage_module.has_linked_issue(body) is True - - @pytest.mark.parametrize( - "body", - [ - "", - "Some change", - # Casual mentions must NOT auto-pass — they should fall through to - # the LLM judge so the stricter "not a passing mention" rule fires. - "See #1234", - "see #1234 for context", - "ref #1234", - "Refs https://github.com/BerriAI/litellm/issues/27000", - "this addresses #1234", - ], - ) - def test_should_not_auto_pass_casual_mentions(self, triage_module, body): - assert triage_module.has_linked_issue(body) is False - - def test_should_not_detect_when_only_html_comment_template(self, triage_module): - body = "" - assert triage_module.has_linked_issue(body) is False - - -class TestStripHtmlComments: - def test_should_remove_single_line_comments(self, triage_module): - text = "before after" - assert "placeholder" not in triage_module.strip_html_comments(text) - - def test_should_remove_multiline_comments(self, triage_module): - text = "kept\n\nkept2" - cleaned = triage_module.strip_html_comments(text) - assert "Fixes #1" not in cleaned - assert "kept" in cleaned and "kept2" in cleaned - - def test_should_handle_none(self, triage_module): - assert triage_module.strip_html_comments(None) == "" - - -class TestCloseCommentText: - """Pin the user-facing language in close comments so changes are intentional.""" - - def test_pr_close_comment_should_recommend_new_pr_primarily(self, triage_module): - body = triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": ["QA proof"], "explanation": "thin"} - ) - # Primary path: open a new PR (because OSS authors can't reopen a - # bot-closed PR). Secondary path: `@agent-shin reconsider`. - assert "Open a new PR" in body - assert "@agent-shin reconsider" in body - # Old advice that no longer works for OSS contributors must NOT - # appear (they can't reopen a PR closed by a bot/maintainer). - assert "Reopen the PR" not in body - - def test_reopen_comment_should_carry_reconsider_marker(self, triage_module): - # The marker is what the rate-limit guard greps for to detect a - # prior reconsider verdict on the same PR. If the marker ever - # gets dropped from this comment, the cooldown silently breaks - # and a contributor can spam `@agent-shin reconsider` to burn - # LLM budget. - body = triage_module.format_reopen_comment("pr") - assert triage_module.RECONSIDER_COMMENT_MARKER in body - - def test_still_failing_comment_should_carry_reconsider_marker(self, triage_module): - body = triage_module.format_reconsider_still_failing_comment( - "pr", - {"verdict": "fail", "missing": ["QA proof"], "explanation": "thin"}, - ) - assert triage_module.RECONSIDER_COMMENT_MARKER in body - - def test_pr_close_comment_should_not_promise_automatic_reopen_on_open( - self, triage_module - ): - # The previous comment said "I'll re-evaluate automatically" — that - # only worked because the author could reopen, which they often - # can't. The new wording must point them at the comment trigger or - # a new PR instead. - body = triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "I'll re-evaluate automatically" not in body - - def test_issue_close_comment_should_use_reconsider_trigger(self, triage_module): - # OSS authors have read access, which only lets them reopen issues - # they closed themselves; they CANNOT reopen an issue a maintainer or - # bot closed. So the recovery path is `@agent-shin reconsider` (the - # bot reopens), exactly like the PR path. If this regresses to "reopen - # it yourself", contributors hit a dead end on bot-closed issues. - body = triage_module.format_issue_close_comment( - {"verdict": "fail", "missing": ["repro"], "explanation": "thin"} - ) - assert "@agent-shin reconsider" in body - - def test_pr_close_comment_should_link_blog_explainer(self, triage_module): - # The blog post is the canonical public explanation of what the bot - # checks and why. Every action-required bot comment must link to it - # so contributors landing on a bot-closed PR can self-serve context - # without pinging a maintainer. - body = triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "https://docs.litellm.ai/blog/agent-shin-triage" in body - - def test_issue_close_comment_should_link_blog_explainer(self, triage_module): - body = triage_module.format_issue_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "https://docs.litellm.ai/blog/agent-shin-triage" in body - - def test_pr_close_comment_should_flag_mocked_tests_as_insufficient_proof( - self, triage_module - ): - # The PR rubric was tightened to require end-to-end QA proof and - # explicitly exclude mocked-dependency unit tests. The user-facing - # close comment must say so — otherwise contributors will keep - # re-submitting "pytest passed (mocks)" runs and getting closed - # again with no explanation of why. - body = triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "end-to-end qa proof" in body.lower() - assert "mock" in body.lower() - - def test_issue_recovery_comments_should_name_feature_dead_end_evidence( - self, triage_module - ): - # The feature-request pass bar demands end-to-end evidence of the - # dead-end, so the close and grace-warning recovery bullets must ask - # for it too — otherwise a requester follows those exact instructions - # (description + use case only) and fails `reconsider` again with no - # hint of what else was needed. - verdict = {"verdict": "fail", "missing": [], "explanation": ""} - for body in ( - triage_module.format_issue_close_comment(verdict), - triage_module.format_grace_warning_issue_comment(verdict), - ): - normalized = " ".join(body.split()) - assert "end-to-end evidence of the dead-end" in normalized - assert "showing where the flow stops today" in normalized - - def test_all_agent_shin_comments_should_use_bullet_train_emoji(self, triage_module): - # The bullet train (🚅) is Agent Shin's symbol, matching the LiteLLM - # logo; the previous wave (👋) was generic and didn't match the bot's - # identity. Every action-required comment the bot can post must use the - # bullet train so the contributor recognizes who's writing without - # reading the signoff. - verdict = {"verdict": "fail", "missing": [], "explanation": ""} - comments = { - "pr_close": triage_module.format_pr_close_comment(verdict), - "issue_close": triage_module.format_issue_close_comment(verdict), - "pr_grace": triage_module.format_grace_warning_pr_comment(verdict), - "issue_grace": triage_module.format_grace_warning_issue_comment(verdict), - "within_grace": triage_module.format_within_grace_comment( - [], "", grace_days=1 - ), - } - for name, body in comments.items(): - assert "🚅" in body, f"{name} comment is missing the bullet train emoji" - assert "👋" not in body, f"{name} comment still uses the old wave emoji" - - def test_pr_close_comment_should_show_what_pr_got_right(self, triage_module): - # The user explicitly asked for a "things you got right" section so - # the comment doesn't read as pure rejection. When the judge confirms - # a field is present (e.g. linked_issue), the bullet for it MUST - # appear in the close comment. - body = triage_module.format_pr_close_comment( - { - "verdict": "fail", - "linked_issue": True, - "has_problem_description": True, - "has_expected_vs_actual": False, - "has_qa_proof": False, - "missing": ["QA proof"], - "explanation": "no proof", - } - ) - assert "What you got right" in body - # The two present fields surface as ✅ bullets; the two absent - # fields do not get a ✅ bullet (the QA-proof rubric block still - # mentions the concept, but only the affirmed fields get checkmarks). - assert "- ✅ Linked a related GitHub issue" in body - assert "- ✅ Clear problem description" in body - assert "- ✅ Expected vs. actual behavior" not in body - assert "- ✅ End-to-end QA proof" not in body - - def test_pr_close_comment_should_omit_present_section_when_nothing_present( - self, triage_module - ): - # If the judge says nothing is present (every flag False), the - # "what you got right" block is skipped entirely — better to omit - # than to render "What you got right: (nothing)". - body = triage_module.format_pr_close_comment( - { - "verdict": "fail", - "linked_issue": False, - "has_problem_description": False, - "has_expected_vs_actual": False, - "has_qa_proof": False, - "missing": [], - "explanation": "", - } - ) - assert "What you got right" not in body - - def test_issue_close_comment_should_show_what_issue_got_right(self, triage_module): - # `has_expected_vs_actual` is present, the end-to-end bug evidence is - # not: the "what you got right" block must surface the former and omit - # the latter (no "✅ (nothing)"-style noise for absent items). - body = triage_module.format_issue_close_comment( - { - "verdict": "fail", - "kind": "bug", - "has_repro": False, - "has_expected_vs_actual": True, - "missing": ["end-to-end evidence of the bug"], - "explanation": "no repro shown", - } - ) - assert "What you got right" in body - assert "Expected vs. actual behavior" in body - assert "- ✅ End-to-end evidence of the bug" not in body - - def test_issue_close_comment_should_credit_feature_dead_end_evidence( - self, triage_module - ): - # A feature requester who pasted their dead-end run but skipped the - # motivation must see the evidence credited and only the motivation - # listed as a gap — without a dedicated verdict field the praise - # block could never acknowledge the work they did do. - body = triage_module.format_issue_close_comment( - { - "verdict": "fail", - "kind": "feature", - "has_motivation_example": False, - "has_dead_end_evidence": True, - "missing": ["motivation / use case"], - "explanation": "no use case given", - } - ) - assert "What you got right" in body - assert "- ✅ End-to-end evidence of the dead-end" in body - assert "- ✅ Motivation and concrete example" not in body - - def test_close_comments_should_use_softer_park_for_later_framing( - self, triage_module - ): - # User feedback: the messaging shouldn't feel like punishment. The - # comment must explicitly frame close as a "park this for later," not - # a rejection, and ground that in the queue-hygiene reason. - for body in ( - triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - triage_module.format_issue_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - ): - assert "park this for later" in body - assert ( - "not a rejection" in body - or "isn't a rejection" in body - or ("isn't us saying" in body) - ) - - def test_only_close_comments_carry_the_agent_shin_close_marker(self, triage_module): - # The reconsider reopen guard keys off AGENT_SHIN_CLOSE_MARKER to tell - # an Agent Shin close from a same-identity close by another workflow. - # That only works if the marker is stamped on the close comments and - # NOT on the grace warnings (which don't close anything). - verdict = {"verdict": "fail", "missing": [], "explanation": ""} - marker = triage_module.AGENT_SHIN_CLOSE_MARKER - assert marker in triage_module.format_pr_close_comment(verdict) - assert marker in triage_module.format_issue_close_comment(verdict) - assert marker not in triage_module.format_grace_warning_pr_comment(verdict) - assert marker not in triage_module.format_grace_warning_issue_comment(verdict) - - -class TestWasClosedByAgentShin: - """Bot-closed guard: only Agent Shin's own closures are reopen candidates.""" - - @staticmethod - def _stub_close_event( - triage_module, - monkeypatch, - *, - actor: str | None, - closed_at: object = "now", - ): - """Stub the most recent `closed` event used by the guard. - - `actor` is the login that closed the item. `closed_at` defaults - to "now" so the marker comment (stubbed at 42s ago) reads as - recent enough relative to the close; tests can pass a concrete - ``datetime`` to simulate older closes (e.g. the stale-marker - regression scenario). - """ - import datetime as real_dt - - if closed_at == "now": - closed_at = real_dt.datetime.now(real_dt.timezone.utc) - monkeypatch.setattr( - triage_module, - "fetch_last_close_event", - lambda repo, n: (actor, closed_at), - ) - - @staticmethod - def _stub_close_marker_present( - triage_module, monkeypatch, *, present: bool, age_seconds: float = 42.0 - ): - """Stub the Agent Shin close-comment marker lookup. - - `was_closed_by_agent_shin` requires the closing actor AND a - recent Agent Shin close comment; these tests pin the latter so - they exercise the actor half in isolation. - """ - monkeypatch.setattr( - triage_module, - "seconds_since_last_agent_shin_close", - lambda *a, **kw: age_seconds if present else None, - ) - - def test_should_return_true_when_bot_closed_and_close_comment_present( - self, triage_module, monkeypatch - ): - self._stub_close_event(triage_module, monkeypatch, actor="github-actions[bot]") - self._stub_close_marker_present(triage_module, monkeypatch, present=True) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is True - - def test_should_return_false_when_bot_closed_but_no_agent_shin_comment( - self, triage_module, monkeypatch - ): - # The `github-actions[bot]` identity is shared across workflows. A - # stale/duplicate sweep closing under that identity must NOT let - # @agent-shin reconsider reopen the item: without an Agent Shin close - # comment the guard fails closed. - self._stub_close_event(triage_module, monkeypatch, actor="github-actions[bot]") - self._stub_close_marker_present(triage_module, monkeypatch, present=False) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - def test_should_return_false_when_last_close_actor_is_maintainer( - self, triage_module, monkeypatch - ): - # A maintainer closed it (e.g. duplicate, security, design). The - # bot must refuse to reopen on @agent-shin reconsider even if an - # earlier Agent Shin close comment is still on the thread. - self._stub_close_event(triage_module, monkeypatch, actor="krrishdholakia") - self._stub_close_marker_present(triage_module, monkeypatch, present=True) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - def test_should_fail_closed_when_no_close_event(self, triage_module, monkeypatch): - # If the events API returns nothing (network blip, repo permission - # quirk), the guard must fail-closed: refuse to reopen rather than - # assume the bot did it. - self._stub_close_event(triage_module, monkeypatch, actor=None, closed_at=None) - self._stub_close_marker_present(triage_module, monkeypatch, present=True) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - def test_should_fail_closed_when_close_event_has_no_timestamp( - self, triage_module, monkeypatch - ): - # Without a usable close timestamp the guard cannot prove the - # marker comment belongs to the latest close; fail-closed. - self._stub_close_event( - triage_module, monkeypatch, actor="github-actions[bot]", closed_at=None - ) - self._stub_close_marker_present(triage_module, monkeypatch, present=True) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - def test_should_return_false_when_marker_predates_latest_close( - self, triage_module, monkeypatch - ): - # Regression for the stale-marker bug: Agent Shin closed once - # (marker stamped), reconsider reopened, and a different workflow - # later closed under the same bot identity without stamping the - # marker. The old marker is still on the thread but does NOT - # belong to the latest close, so reconsider must not reopen. - import datetime as real_dt - - now = real_dt.datetime.now(real_dt.timezone.utc) - # Latest close happened a minute ago. - self._stub_close_event( - triage_module, - monkeypatch, - actor="github-actions[bot]", - closed_at=now - real_dt.timedelta(seconds=60), - ) - # The most recent Agent Shin marker is from an hour ago (a prior - # closed/reopened cycle), which is well outside the skew window. - self._stub_close_marker_present( - triage_module, monkeypatch, present=True, age_seconds=3600.0 - ) - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - def test_should_respect_bot_login_override_via_env( - self, triage_module, monkeypatch - ): - # Operators wiring Agent Shin to a PAT (instead of GITHUB_TOKEN) - # can override the expected bot login via env. The guard must - # respect the override so non-default deployments still work. - monkeypatch.setenv("AGENT_SHIN_BOT_LOGIN", "my-bot") - self._stub_close_marker_present(triage_module, monkeypatch, present=True) - self._stub_close_event(triage_module, monkeypatch, actor="my-bot") - assert triage_module.was_closed_by_agent_shin("o/r", 1) is True - # Default "github-actions[bot]" should NOT match when env is set. - self._stub_close_event(triage_module, monkeypatch, actor="github-actions[bot]") - assert triage_module.was_closed_by_agent_shin("o/r", 1) is False - - -class TestSecondsSinceLastAgentShinClose: - """Close-provenance lookup: detects the bot's own auto-close marker.""" - - def _make_comment(self, *, login: str, body: str) -> dict: - return { - "user": {"login": login}, - "body": body, - "created_at": "2026-05-18T05:00:00Z", - } - - def test_should_return_none_when_bot_never_closed(self, triage_module, monkeypatch): - # Comments exist, but none is an Agent Shin close — e.g. only a grace - # warning, or a close by another workflow with no Agent Shin comment. - comments = [ - self._make_comment(login="outside-dev", body="any update?"), - self._make_comment( - login="github-actions[bot]", - body=triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_agent_shin_close("o/r", 1) is None - - def test_should_detect_bot_close_comment(self, triage_module, monkeypatch): - comments = [ - self._make_comment( - login="github-actions[bot]", - body=triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_agent_shin_close("o/r", 1) is not None - - def test_should_ignore_non_bot_comment_quoting_marker( - self, triage_module, monkeypatch - ): - # A contributor quoting the hidden marker (GitHub "Quote reply" - # preserves HTML comments) must not be mistaken for a bot close. - comments = [ - self._make_comment( - login="curious-user", - body=f"what is this? {triage_module.AGENT_SHIN_CLOSE_MARKER}", - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_agent_shin_close("o/r", 1) is None - - -class TestSecondsSinceLastReconsiderVerdict: - """Rate-limit guard: detects the bot's own reconsider verdict marker.""" - - def _make_comment( - self, *, login: str, body: str, created_at: str | None = "2026-05-18T05:00:00Z" - ) -> dict: - comment: dict = {"user": {"login": login}, "body": body} - if created_at is not None: - comment["created_at"] = created_at - return comment - - def test_should_return_none_when_no_bot_reconsider_comments( - self, triage_module, monkeypatch - ): - # An issue with chatter from other users but no bot reconsider - # verdict must not be rate-limited. - comments = [ - self._make_comment(login="outside-dev", body="ping?"), - self._make_comment( - login="github-actions[bot]", body="some other bot message" - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_reconsider_verdict("o/r", 1) is None - - def test_should_pick_latest_bot_reconsider_marker(self, triage_module, monkeypatch): - # When multiple reconsider verdicts exist, return the AGE of the - # most recent one. Using a frozen reference helps pin the math. - comments = [ - self._make_comment( - login="github-actions[bot]", - body="old verdict " + triage_module.RECONSIDER_COMMENT_MARKER, - created_at="2026-05-18T04:00:00Z", - ), - self._make_comment( - login="github-actions[bot]", - body="newer verdict " + triage_module.RECONSIDER_COMMENT_MARKER, - created_at="2026-05-18T04:55:00Z", - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - - # Freeze "now" via a tiny shim on the module's `dt` import. - import datetime as real_dt - - class FrozenDateTime(real_dt.datetime): - @classmethod - def now(cls, tz=None): - return real_dt.datetime(2026, 5, 18, 5, 0, 0, tzinfo=tz) - - frozen_module = type(triage_module.dt)("datetime") - frozen_module.datetime = FrozenDateTime - frozen_module.timezone = real_dt.timezone - monkeypatch.setattr(triage_module, "dt", frozen_module) - - age = triage_module.seconds_since_last_reconsider_verdict("o/r", 1) - # newer verdict is 5 minutes (300 seconds) before "now" - assert age == 300.0 - - def test_should_ignore_non_bot_comments_with_marker( - self, triage_module, monkeypatch - ): - # A user comment that happens to quote the marker (e.g. in - # a "what does this hidden marker do?" question) must NOT count. - # The rate-limit guard only trusts comments authored by the bot. - comments = [ - self._make_comment( - login="curious-user", - body=f"Saw this marker: {triage_module.RECONSIDER_COMMENT_MARKER}", - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_reconsider_verdict("o/r", 1) is None - - def test_should_ignore_bot_comments_without_marker( - self, triage_module, monkeypatch - ): - # The bot posts other things too (Agent Shin close comments, - # CI status, etc.) — only the reconsider-verdict marker should - # arm the cooldown. - comments = [ - self._make_comment( - login="github-actions[bot]", - body="Agent Shin closed this PR (no marker)", - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_reconsider_verdict("o/r", 1) is None - - -class TestParseVerdict: - def test_should_parse_plain_json(self, triage_module): - raw = '{"verdict": "pass", "missing": []}' - assert triage_module.parse_verdict(raw)["verdict"] == "pass" - - def test_should_strip_markdown_fence(self, triage_module): - raw = '```json\n{"verdict": "fail", "missing": ["foo"]}\n```' - result = triage_module.parse_verdict(raw) - assert result["verdict"] == "fail" - assert result["missing"] == ["foo"] - - def test_should_extract_embedded_json_from_prose(self, triage_module): - raw = 'Here you go: {"verdict": "pass", "missing": []}\nThanks.' - assert triage_module.parse_verdict(raw)["verdict"] == "pass" - - def test_should_raise_for_unparseable_text(self, triage_module): - with pytest.raises(ValueError, match='could not extract JSON from LLM response: not even close to'): - triage_module.parse_verdict("not even close to json") - - def test_should_raise_for_empty(self, triage_module): - with pytest.raises(ValueError, match='empty LLM response'): - triage_module.parse_verdict("") - - -class TestBuildPrompts: - def test_should_include_pr_title_and_body(self, triage_module): - prompt = triage_module.build_pr_prompt( - title="Add foo", body=" Real body" - ) - assert "Add foo" in prompt - assert "Real body" in prompt - assert "comment" not in prompt # HTML comments are stripped - - def test_should_show_empty_marker_for_empty_pr_body(self, triage_module): - prompt = triage_module.build_pr_prompt(title="t", body="") - assert "(empty)" in prompt - - def test_should_include_issue_title_and_body(self, triage_module): - prompt = triage_module.build_issue_prompt(title="Bug", body="repro here") - assert "Bug" in prompt - assert "repro here" in prompt - - def test_issue_bug_rubric_requires_end_to_end_evidence_and_drops_pass_bias( - self, triage_module - ): - # The bug bar was tightened: a report needs the "before" half shown - # end-to-end (video / screenshot / real command output), prose-only - # repro steps no longer pass, and the old "bias toward PASS" leniency - # is gone. If any of these regress, the judge silently goes soft on - # undemonstrated bug reports again. - prompt = triage_module.build_issue_prompt(title="t", body="x") - normalized = " ".join(prompt.split()) - assert "Bias toward PASS when the issue has structure" not in normalized - assert "END-TO-END EVIDENCE OF THE BUG" in normalized - assert "Do not bias toward PASS" in normalized - # The three accepted forms of the "before" demonstration must be named. - assert "screen recording / video" in normalized - assert "screenshot of the bug" in normalized - assert "mocked or stubbed" in normalized - # Prose-only steps are explicitly insufficient now. - assert "steps to reproduce" in normalized - # An unedited issue-form scaffold must not read as evidence: the proof - # field ships with visible headings, so the judge has to be told that - # bare headings with nothing under them count as absent. - assert "unfilled template scaffold" in normalized - assert "counts as absent, not as evidence" in normalized - - def test_issue_feature_rubric_requires_evidence_of_the_dead_end( - self, triage_module - ): - # The feature form asks the requester to walk the ideal flow against a - # live proxy and paste output up to the step that dead-ends, so the - # judge has to demand that evidence, and must not accept an unedited - # scaffold of bare headings as if it were a real attempt. - prompt = triage_module.build_issue_prompt(title="t", body="x") - normalized = " ".join(prompt.split()) - assert "END-TO-END EVIDENCE OF THE DEAD-END" in normalized - assert "showing the point where the flow stops today" in normalized - assert "unfilled template scaffold" in normalized - # The evidence has its own verdict field so feature requesters who - # provided it get credited in "What you got right", exactly like - # `has_repro` credits bug evidence. - assert "`has_dead_end_evidence=true` only when this is present" in normalized - assert '"has_dead_end_evidence": boolean' in normalized - - def test_should_not_crash_when_pr_body_contains_curly_braces(self, triage_module): - """User-supplied content with `{` / `}` must NOT be re-parsed by - `str.format()`. `format` only scans the template literal for - replacement fields; values being substituted in are inserted as - plain strings, so a body like `{"foo": "bar"}` or `{unmatched` - cannot blow up the script. Pinning this here so a future - "improvement" to the templating doesn't reintroduce a crash on - every PR that quotes JSON. - """ - for body in ( - 'Here is some JSON: {"foo": "bar", "n": 1}', - "Half a brace { left dangling, and a stray }", - "Format-spec-looking thing: {0}, {name:>10}, {!r}", - "Nested {a: {b: c}} braces", - ): - pr_prompt = triage_module.build_pr_prompt(title="t", body=body) - issue_prompt = triage_module.build_issue_prompt(title="t", body=body) - assert body in pr_prompt - assert body in issue_prompt - - def test_should_not_crash_when_pr_title_contains_curly_braces(self, triage_module): - title = "Fix bug in {0:>10} format-spec handling" - pr_prompt = triage_module.build_pr_prompt(title=title, body="x") - issue_prompt = triage_module.build_issue_prompt(title=title, body="x") - assert title in pr_prompt - assert title in issue_prompt - - def test_should_preserve_template_indentation_with_multiline_body( - self, triage_module - ): - """`textwrap.dedent` runs on the static template *before* user - content is interpolated, so a multi-line body (whose 2nd+ lines - start at column 0) cannot defeat the common-indent computation - and leave 8-space indentation on every template line. Pin the - dedented shape so the rendered prompt stays consistent for the - LLM judge. - """ - body = "first line\nsecond line at column 0\nthird line at column 0" - for builder in ( - triage_module.build_pr_prompt, - triage_module.build_issue_prompt, - ): - prompt = builder(title="t", body=body) - # Template lines should NOT carry the 8 leading spaces from - # the source-file indentation of the triple-quoted string. - assert " You are " not in prompt - assert 'You are "Agent Shin"' in prompt - assert body in prompt - - -class TestMainModelDefault: - """`--model` falls back to DEFAULT_MODEL even when TRIAGE_MODEL is empty.""" - - def _stub_triage(self, triage_module, monkeypatch): - captured: dict = {} - - def fake_triage(**kwargs): - captured.update(kwargs) - return { - "kind": kwargs["kind"], - "number": kwargs["number"], - "title": "", - "author": "x", - "author_association": "NONE", - "state": "open", - "action": "skip-no-llm-key", - } - - monkeypatch.setattr(triage_module, "triage", fake_triage) - return captured - - def test_should_fall_back_to_default_when_triage_model_env_empty( - self, triage_module, monkeypatch - ): - captured = self._stub_triage(triage_module, monkeypatch) - monkeypatch.setenv("TRIAGE_MODEL", "") - monkeypatch.setattr( - sys, - "argv", - ["triage_with_llm.py", "--repo", "o/r", "--pr", "1"], - ) - rc = triage_module.main() - assert rc == 0 - assert captured["model"] == triage_module.DEFAULT_MODEL - - def test_should_respect_explicit_triage_model_env(self, triage_module, monkeypatch): - captured = self._stub_triage(triage_module, monkeypatch) - monkeypatch.setenv("TRIAGE_MODEL", "gpt-4o-mini") - monkeypatch.setattr( - sys, - "argv", - ["triage_with_llm.py", "--repo", "o/r", "--pr", "1"], - ) - rc = triage_module.main() - assert rc == 0 - assert captured["model"] == "gpt-4o-mini" - - -class TestCallLlmJudge: - """call_llm_judge sets gpt-5 specific kwargs correctly.""" - - def _stub_openai(self, monkeypatch, captured: dict): - """Install a fake `openai.OpenAI` client into sys.modules. - - The fake client records the kwargs passed to chat.completions.create - and returns a minimal response object whose .choices[0].message.content - is "ok". - """ - import types - - class FakeMessage: - content = '{"verdict": "pass"}' - - class FakeChoice: - message = FakeMessage() - - class FakeResponse: - choices = [FakeChoice()] - - class FakeCompletions: - def create(self, **kwargs): - captured.update(kwargs) - return FakeResponse() - - class FakeChat: - completions = FakeCompletions() - - class FakeClient: - def __init__(self, api_key, base_url=None): - captured["__client_kwargs__"] = { - "api_key": api_key, - "base_url": base_url, - } - self.chat = FakeChat() - - fake_module = types.ModuleType("openai") - fake_module.OpenAI = FakeClient - monkeypatch.setitem(sys.modules, "openai", fake_module) - - def test_should_set_reasoning_effort_none_for_gpt5_family( - self, triage_module, monkeypatch - ): - captured: dict = {} - self._stub_openai(monkeypatch, captured) - triage_module.call_llm_judge( - "prompt", model="gpt-5.4-mini", api_key="sk-test", base_url=None - ) - assert captured["model"] == "gpt-5.4-mini" - assert captured["temperature"] == 0 - assert captured["extra_body"] == {"reasoning_effort": "none"} - - def test_should_set_reasoning_effort_for_capitalized_or_dated_gpt5( - self, triage_module, monkeypatch - ): - for model in ("GPT-5.4-mini", "gpt-5.4-mini-2026-03-17", "gpt-5"): - captured: dict = {} - self._stub_openai(monkeypatch, captured) - triage_module.call_llm_judge( - "prompt", model=model, api_key="sk-test", base_url=None - ) - assert captured["extra_body"] == {"reasoning_effort": "none"}, model - - def test_should_omit_reasoning_effort_for_non_gpt5( - self, triage_module, monkeypatch - ): - captured: dict = {} - self._stub_openai(monkeypatch, captured) - triage_module.call_llm_judge( - "prompt", model="gpt-4o-mini", api_key="sk-test", base_url=None - ) - assert "extra_body" not in captured - - def test_should_pass_base_url_when_provided(self, triage_module, monkeypatch): - captured: dict = {} - self._stub_openai(monkeypatch, captured) - triage_module.call_llm_judge( - "p", - model="gpt-5.4-mini", - api_key="sk-test", - base_url="https://proxy.example.com/v1", - ) - assert ( - captured["__client_kwargs__"]["base_url"] == "https://proxy.example.com/v1" - ) - - -class TestTriageOrchestration: - """End-to-end-ish tests that mock both gh fetchers and the LLM.""" - - def _make_pr(self, **overrides): - base = { - "number": 1, - "title": "PR title", - "body": "PR body", - "state": "open", - "author_association": "NONE", - "user": {"login": "mateo-berri"}, - } - base.update(overrides) - return base - - def test_should_skip_internal_author(self, triage_module, monkeypatch): - pr = self._make_pr( - author_association="MEMBER", user={"login": "krrishdholakia"} - ) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - - def boom(*a, **kw): - pytest.fail("LLM should not be called for internal authors") - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=boom, - allowlist=frozenset(), - ) - assert result["action"] == "skip-internal-author" - - def test_should_skip_closed_pr(self, triage_module, monkeypatch): - pr = self._make_pr(state="closed") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("should not run on closed PRs"), - ) - assert result["action"] == "skip-not-open" - - def test_should_short_circuit_on_linked_issue(self, triage_module, monkeypatch): - pr = self._make_pr(body="Fixes #1234\n\nFoo bar") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM should not be called"), - ) - assert result["action"] == "pass-linked-issue" - assert result["verdict"]["verdict"] == "pass" - - def test_should_not_short_circuit_on_casual_mention( - self, triage_module, monkeypatch - ): - # "See #1234" is a passing mention, not a closing keyword. The LLM - # must get a chance to apply the stricter rubric. With no prior - # grace warning, the first failing verdict triggers the warning - # path (`would-warn-grace` in dry-run). - pr = self._make_pr(body="See #1234 for context. No QA proof here.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_no_warning(triage_module, monkeypatch) - called = {"judge": False} - - def judge(prompt): - called["judge"] = True - return json.dumps( - {"verdict": "fail", "missing": ["QA proof"], "explanation": "thin."} - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=False, - model="m", - judge=judge, - ) - assert called["judge"] is True - assert result["action"] == "would-warn-grace" - - def test_should_return_pass_llm_when_judge_passes(self, triage_module, monkeypatch): - pr = self._make_pr(body="Long body, no linked issue.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - captured = {} - - def judge(prompt): - captured["prompt"] = prompt - return json.dumps({"verdict": "pass", "missing": [], "explanation": "ok"}) - - result = triage_module.triage( - repo="o/r", kind="pr", number=1, close=True, model="m", judge=judge - ) - assert result["action"] == "pass-llm" - assert "Long body" in captured["prompt"] - - def test_should_return_would_close_in_dry_run_after_grace_aged_out( - self, triage_module, monkeypatch - ): - # When the grace warning has already aged out (>= GRACE_PERIOD_SECONDS) - # AND the rubric still fails, the dry-run preview returns - # `would-close` so a step-summary writer can render the close - # comment without touching GitHub state. - pr = self._make_pr(body="just a sentence.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_aged_out(triage_module, monkeypatch) - - def fake_post(*a, **kw): - pytest.fail("should not post comments in dry-run") - - def fake_close(*a, **kw): - pytest.fail("should not close in dry-run") - - monkeypatch.setattr(triage_module, "post_comment", fake_post) - monkeypatch.setattr(triage_module, "close_pr", fake_close) - - verdict = { - "verdict": "fail", - "missing": ["problem description", "QA proof"], - "explanation": "Body is one sentence.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=False, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "would-close" - assert result["verdict"]["missing"] == ["problem description", "QA proof"] - - def test_should_post_comment_and_close_after_grace_window( - self, triage_module, monkeypatch - ): - # The "real close" path: --close passed AND the grace warning has - # aged out AND the rubric still fails. The bot posts the close - # comment and closes the PR. - pr = self._make_pr(body="just a sentence.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_aged_out(triage_module, monkeypatch) - posted = {} - closed = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"repo": repo, "n": n, "body": body}), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda repo, n: closed.update({"repo": repo, "n": n}), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Body too thin.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "closed" - assert posted["n"] == 42 and closed["n"] == 42 - assert "Agent Shin" in posted["body"] - assert "QA proof" in posted["body"] - - def test_should_skip_on_llm_error_in_close_mode(self, triage_module, monkeypatch): - pr = self._make_pr(body="something.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not comment on LLM error"), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail("must not close on LLM error"), - ) - - def broken_judge(prompt): - raise RuntimeError("upstream 500") - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=broken_judge, - ) - assert result["action"] == "skip-llm-error" - assert "upstream 500" in result["error"] - - def test_should_skip_open_pr_in_reconsider_mode(self, triage_module, monkeypatch): - # Reconsider only makes sense on a CLOSED PR — running it on an open - # one is a no-op (the regular triage flow already evaluated it). - pr = self._make_pr(state="open") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=False, - model="m", - judge=lambda p: pytest.fail("should not run on open PR in reconsider"), - reconsider=True, - ) - assert result["action"] == "skip-not-closed" - - @staticmethod - def _stub_reconsider_guards(triage_module, monkeypatch): - """Default reconsider-guard stubs: pretend bot closed + no cooldown. - - The new safety guards (`was_closed_by_agent_shin`, - `seconds_since_last_reconsider_verdict`) hit the GitHub API in - production. Tests that exercise the reconsider happy path stub - them to "yes the bot closed it, no recent reconsider comment" - so the test stays focused on its actual assertion. - """ - monkeypatch.setattr( - triage_module, "was_closed_by_agent_shin", lambda *a, **kw: True - ) - monkeypatch.setattr( - triage_module, - "seconds_since_last_reconsider_verdict", - lambda *a, **kw: None, - ) - - @staticmethod - def _stub_grace_aged_out(triage_module, monkeypatch): - """Pretend the grace warning has aged out. - - For tests that exercise the post-grace close path. Set the age - to twice the grace window so a future tweak to - `GRACE_PERIOD_SECONDS` doesn't accidentally make the stub fall - back inside the window. - """ - monkeypatch.setattr( - triage_module, - "seconds_since_last_grace_warning", - lambda *a, **kw: triage_module.GRACE_PERIOD_SECONDS * 2, - ) - - @staticmethod - def _stub_grace_no_warning(triage_module, monkeypatch): - """Pretend no grace warning has been posted yet (first detection).""" - monkeypatch.setattr( - triage_module, - "seconds_since_last_grace_warning", - lambda *a, **kw: None, - ) - - def test_should_reopen_on_reconsider_pass(self, triage_module, monkeypatch): - # Reconsider on a closed PR with a passing verdict -> reopen + post a - # friendly "re-evaluated" comment. close=True is the production path - # (the workflow only adds --close when AGENT_SHIN_ENABLED=true). - pr = self._make_pr( - state="closed", body="Updated body with QA proof + screenshots." - ) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - posted = {} - reopened = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"n": n, "body": body}), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda repo, n: reopened.update({"n": n}), - ) - # close_pr / close_issue MUST NOT fire in reconsider mode. - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail("must not close on reconsider pass"), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p: json.dumps( - {"verdict": "pass", "missing": [], "explanation": "ok now"} - ), - reconsider=True, - ) - assert result["action"] == "reopened" - assert reopened["n"] == 42 - assert posted["n"] == 42 - assert "reopened" in posted["body"].lower() - - def test_should_dry_run_reconsider_pass_when_close_false( - self, triage_module, monkeypatch - ): - # Reconsider must honor `close=False` (dry-run) just like the - # regular triage flow. A local invocation of - # `python triage_with_llm.py --reconsider --pr N` (no --close) - # must NOT post a comment or reopen the PR — it should return - # `would-reopen` so the operator can preview the outcome. - pr = self._make_pr( - state="closed", body="Updated body with QA proof + screenshots." - ) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not post comment in dry-run reconsider"), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen PR in dry-run reconsider"), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=False, - model="m", - judge=lambda p: json.dumps( - {"verdict": "pass", "missing": [], "explanation": "ok now"} - ), - reconsider=True, - ) - assert result["action"] == "would-reopen" - # The previewed comment body is still returned so a step-summary - # writer can render exactly what would have been posted. - assert "reopened" in result["comment"].lower() - - def test_should_post_still_failing_on_reconsider_fail( - self, triage_module, monkeypatch - ): - pr = self._make_pr(state="closed", body="still empty") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - posted = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"n": n, "body": body}), - ) - # Neither reopen nor close should fire when reconsider verdict is fail. - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen on fail"), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail("must not close again on reconsider fail"), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Still no QA proof.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - reconsider=True, - ) - assert result["action"] == "reconsider-still-failing" - assert posted["n"] == 42 - assert "QA proof" in posted["body"] - - def test_should_not_reopen_on_reconsider_with_ambiguous_verdict( - self, triage_module, monkeypatch - ): - # Regression: only an explicit `pass` verdict reopens. Missing, - # empty, or unexpected verdict strings ("failed", "", garbage) - # must fall through to the still-failing branch rather than - # reopen a PR the rubric did not actually clear. - pr = self._make_pr(state="closed", body="still empty") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - posted = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"body": body}), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen on ambiguous verdict"), - ) - - for ambiguous in ("", "failed", "needs-info", "unknown"): - posted.clear() - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p, v=ambiguous: json.dumps( - {"verdict": v, "missing": [], "explanation": "weird"} - ), - reconsider=True, - ) - assert result["action"] == "reconsider-still-failing", ambiguous - assert "body" in posted, ambiguous - - def test_should_dry_run_reconsider_fail_when_close_false( - self, triage_module, monkeypatch - ): - # Mirror dry-run behavior for the FAIL branch — `close=False` - # must NOT post the "still failing" comment. - pr = self._make_pr(state="closed", body="still empty") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail( - "must not post still-failing comment in dry-run" - ), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Still no QA proof.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=False, - model="m", - judge=lambda p: json.dumps(verdict), - reconsider=True, - ) - assert result["action"] == "would-reconsider-still-failing" - assert "QA proof" in result["comment"] - - def test_should_reopen_on_reconsider_with_linked_issue_short_circuit( - self, triage_module, monkeypatch - ): - # The linked-issue short-circuit also has to honor reconsider mode: - # if the contributor edited the body to add `Fixes #1234`, the regex - # path should reopen the PR without calling the LLM. - pr = self._make_pr(state="closed", body="Fixes #1234\n\nAddresses the bug.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - posted = {} - reopened = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"body": body}), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda repo, n: reopened.update({"n": n}), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=55, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM must not run when linked-issue matches"), - reconsider=True, - ) - assert result["action"] == "reopened" - assert reopened["n"] == 55 - assert "reopened" in posted["body"].lower() - - def test_should_dry_run_reconsider_with_linked_issue_when_close_false( - self, triage_module, monkeypatch - ): - # Linked-issue short-circuit must ALSO honor dry-run. - pr = self._make_pr(state="closed", body="Fixes #1234\n\nAddresses the bug.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_reconsider_guards(triage_module, monkeypatch) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not post in dry-run"), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen in dry-run"), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=55, - close=False, - model="m", - judge=lambda p: pytest.fail("LLM must not run when linked-issue matches"), - reconsider=True, - ) - assert result["action"] == "would-reopen" - - def test_should_skip_internal_in_reconsider_mode(self, triage_module, monkeypatch): - # Internal authors are exempt from triage in both regular and - # reconsider mode — Agent Shin should never reopen one of their PRs - # automatically, in case a maintainer closed it intentionally. - pr = self._make_pr( - state="closed", - author_association="MEMBER", - user={"login": "krrishdholakia"}, - ) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen for internal author"), - ) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=False, - model="m", - judge=lambda p: pytest.fail("LLM must not run for internal author"), - reconsider=True, - allowlist=frozenset(), - ) - assert result["action"] == "skip-internal-author" - - def test_should_skip_reconsider_when_not_bot_closed( - self, triage_module, monkeypatch - ): - # SECURITY: `@agent-shin reconsider` must NOT reopen a PR/issue - # that a MAINTAINER closed for non-rubric reasons (e.g. duplicate, - # design rejection, security report). Only PRs closed by the bot - # itself should ever be candidates for the reconsider reopen path. - pr = self._make_pr(state="closed", body="something.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, "was_closed_by_agent_shin", lambda *a, **kw: False - ) - # Even though there's no rate-limit conflict, the bot-closed guard - # alone is sufficient to block. The LLM judge must never run on a - # maintainer-closed PR. - monkeypatch.setattr( - triage_module, - "seconds_since_last_reconsider_verdict", - lambda *a, **kw: None, - ) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not comment on maintainer-closed PR"), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen maintainer-closed PR"), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM must not run before bot-closed guard"), - reconsider=True, - ) - assert result["action"] == "skip-not-bot-closed" - - def test_should_rate_limit_repeated_reconsider_triggers( - self, triage_module, monkeypatch - ): - # COST CONTROL: each `@agent-shin reconsider` event burns CI - # minutes + an OpenAI API call. If the bot already posted a - # reconsider verdict within the cooldown window - # (RECONSIDER_RATE_LIMIT_SECONDS), refuse to run again. This - # bounds the damage from a contributor spamming the trigger. - pr = self._make_pr(state="closed", body="something with new edits.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, "was_closed_by_agent_shin", lambda *a, **kw: True - ) - # Pretend the bot posted a reconsider verdict 1 second ago. - monkeypatch.setattr( - triage_module, - "seconds_since_last_reconsider_verdict", - lambda *a, **kw: 1.0, - ) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not comment during cooldown"), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda *a, **kw: pytest.fail("must not reopen during cooldown"), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM must not run during cooldown"), - reconsider=True, - ) - assert result["action"] == "skip-rate-limited" - assert result["rate_limit_age_seconds"] == 1.0 - assert ( - result["rate_limit_window_seconds"] - == triage_module.RECONSIDER_RATE_LIMIT_SECONDS - ) - - def test_should_allow_reconsider_after_cooldown_window( - self, triage_module, monkeypatch - ): - # The cooldown is a window, not a one-shot lock — once - # RECONSIDER_RATE_LIMIT_SECONDS has elapsed since the last bot - # verdict, a fresh `@agent-shin reconsider` is allowed through. - pr = self._make_pr(state="closed", body="updated with screenshots now.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, "was_closed_by_agent_shin", lambda *a, **kw: True - ) - # Last reconsider was 1 hour ago — well outside the 10-min window. - monkeypatch.setattr( - triage_module, - "seconds_since_last_reconsider_verdict", - lambda *a, **kw: 3600.0, - ) - posted = {} - reopened = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"n": n, "body": body}), - ) - monkeypatch.setattr( - triage_module, - "reopen_pr", - lambda repo, n: reopened.update({"n": n}), - ) - - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: json.dumps( - {"verdict": "pass", "missing": [], "explanation": "ok"} - ), - reconsider=True, - ) - assert result["action"] == "reopened" - assert reopened["n"] == 1 - - def test_should_reopen_issue_on_reconsider_pass(self, triage_module, monkeypatch): - issue = { - "number": 7, - "title": "Bug: now with repro", - "body": "## Repro\n```bash\ncurl ...\n```\n\nExpected X, got Y.", - "state": "closed", - "author_association": "NONE", - "user": {"login": "mateo-berri"}, - } - monkeypatch.setattr(triage_module, "fetch_issue", lambda repo, n: issue) - self._stub_reconsider_guards(triage_module, monkeypatch) - posted = {} - reopened = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"body": body}), - ) - monkeypatch.setattr( - triage_module, - "reopen_issue", - lambda repo, n: reopened.update({"n": n}), - ) - - result = triage_module.triage( - repo="o/r", - kind="issue", - number=7, - close=True, - model="m", - judge=lambda p: json.dumps( - {"verdict": "pass", "missing": [], "explanation": "now reproducible"} - ), - reconsider=True, - ) - assert result["action"] == "reopened" - assert reopened["n"] == 7 - assert "reopened" in posted["body"].lower() - - def test_should_triage_issues_kind(self, triage_module, monkeypatch): - issue = { - "number": 7, - "title": "Bug: X is broken", - "body": "no detail", - "state": "open", - "author_association": "NONE", - "user": {"login": "mateo-berri"}, - } - monkeypatch.setattr(triage_module, "fetch_issue", lambda repo, n: issue) - # Grace already aged out -> close path. (Issues use the same - # GRACE_COMMENT_MARKER detection as PRs.) - self._stub_grace_aged_out(triage_module, monkeypatch) - closed = {} - posted = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update(body=body), - ) - monkeypatch.setattr( - triage_module, "close_issue", lambda repo, n: closed.update(n=n) - ) - - verdict = { - "verdict": "fail", - "kind": "bug", - "has_repro": False, - "missing": ["reproduction", "expected vs. actual"], - "explanation": "No repro provided.", - } - result = triage_module.triage( - repo="o/r", - kind="issue", - number=7, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "closed" - assert closed["n"] == 7 - assert "reproduction" in posted["body"] - - # ---- Grace-period flow ------------------------------------------------ - - def test_should_post_grace_warning_on_first_failing_run_in_close_mode( - self, triage_module, monkeypatch - ): - # First low-quality detection -> bot posts a warning comment with - # the GRACE_COMMENT_MARKER. The PR must NOT be closed yet. - pr = self._make_pr(body="just a sentence.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_no_warning(triage_module, monkeypatch) - posted = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"n": n, "body": body}), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail("must not close on first detection"), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Body too thin.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "warned-grace" - assert posted["n"] == 42 - # Pin the user-facing language pieces the user explicitly asked for. - assert "2 hours" in posted["body"] - assert "@agent-shin reconsider" in posted["body"] - assert "@greptileai" in posted["body"] - assert "even after the PR is closed" in posted["body"] - assert triage_module.GRACE_COMMENT_MARKER in posted["body"] - - def test_should_skip_close_inside_grace_window(self, triage_module, monkeypatch): - # A warning was posted recently; do nothing on this run regardless - # of close=True. The next run after `GRACE_PERIOD_SECONDS` elapses - # is the one that flips to actual close. - pr = self._make_pr(body="just a sentence.") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - monkeypatch.setattr( - triage_module, - "seconds_since_last_grace_warning", - lambda *a, **kw: 60.0, - ) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not comment during grace window"), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail("must not close during grace window"), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Body too thin.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=42, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "skip-in-grace-period" - assert result["grace_age_seconds"] == 60.0 - assert result["grace_period_seconds"] == triage_module.GRACE_PERIOD_SECONDS - - def test_should_dry_run_grace_warning_when_close_false( - self, triage_module, monkeypatch - ): - # In dry-run mode the FIRST failing detection returns - # `would-warn-grace` (with the previewed comment body) and never - # touches GitHub state. Lets a local operator preview the - # warning before flipping --close on. - pr = self._make_pr(body="thin") - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_no_warning(triage_module, monkeypatch) - monkeypatch.setattr( - triage_module, - "post_comment", - lambda *a, **kw: pytest.fail("must not post in dry-run grace warn"), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "thin", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=False, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "would-warn-grace" - assert "2 hours" in result["comment"] - - def test_should_warn_grace_for_swiftwinds_not_close_instantly( - self, triage_module, monkeypatch - ): - # Regression: SwiftWinds (the dogfood account) used to be in a - # now-removed `IMMEDIATE_CLOSE_LOGINS` bypass that skipped the grace - # window and closed on first detection. It must follow the SAME - # grace path as every other author: warn first, close only after the - # window elapses. A re-added instant-close bypass would call - # close_pr here and fail the test. - pr = self._make_pr(body="just a sentence.", user={"login": "SwiftWinds"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - self._stub_grace_no_warning(triage_module, monkeypatch) - posted = {} - monkeypatch.setattr( - triage_module, - "post_comment", - lambda repo, n, body: posted.update({"n": n, "body": body}), - ) - monkeypatch.setattr( - triage_module, - "close_pr", - lambda *a, **kw: pytest.fail( - "SwiftWinds must not close on first detection; it gets the grace window" - ), - ) - - verdict = { - "verdict": "fail", - "missing": ["QA proof"], - "explanation": "Body too thin.", - } - result = triage_module.triage( - repo="o/r", - kind="pr", - number=99, - close=True, - model="m", - judge=lambda p: json.dumps(verdict), - ) - assert result["action"] == "warned-grace" - assert "2 hours" in posted["body"] - - -class TestGraceWarningCommentText: - """Pin the user-facing promises in the grace warning so a future - refactor can't silently drop them.""" - - def test_pr_grace_warning_should_state_grace_window(self, triage_module): - body = triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": ["QA proof"], "explanation": "thin"} - ) - # The user explicitly asked: "specify in the comment" the grace window. - assert "2 hours" in body - - def test_pr_grace_warning_should_mention_reconsider_during_grace( - self, triage_module - ): - body = triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "@agent-shin reconsider" in body - - def test_pr_grace_warning_should_promise_greptileai_works_post_close( - self, triage_module - ): - body = triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - # Per user: comment should state @greptileai works even after close. - assert "@greptileai" in body - assert "even after the PR is closed" in body - - def test_pr_grace_warning_should_carry_grace_marker(self, triage_module): - # The marker is what `seconds_since_last_grace_warning` greps for - # on subsequent runs to detect that a warning has been posted. - # Dropping it would silently break the close-after-grace path. - body = triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert triage_module.GRACE_COMMENT_MARKER in body - - def test_issue_grace_warning_should_carry_grace_marker(self, triage_module): - body = triage_module.format_grace_warning_issue_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert triage_module.GRACE_COMMENT_MARKER in body - assert "2 hours" in body - # OSS authors can't reopen a bot-closed issue, so recovery is - # `@agent-shin reconsider` (the bot reopens), like the PR path. - assert "@agent-shin reconsider" in body - - def test_pr_close_comment_should_promise_greptileai_works_post_close( - self, triage_module - ): - # The standard close comment must ALSO point at @greptileai so - # contributors see the same options whether they read the warning - # or only catch the close comment. - body = triage_module.format_pr_close_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "@greptileai" in body - assert "even after the PR is closed" in body - - def test_pr_grace_warning_should_not_prompt_reconsider_during_grace_window( - self, triage_module - ): - # Per user feedback: during the 24h grace window, the contributor - # should just update the PR description. Asking them to also comment - # "@agent-shin reconsider" right away adds a step they don't need — - # the bot re-checks automatically on the next sweep. The reconsider - # trigger is reserved for the post-close recovery path. - # - # We pin this by checking that the grace section explicitly tells - # the contributor they don't need to ping the bot during the grace - # window. The presence of "@agent-shin reconsider" elsewhere in the - # comment (as the post-close path) is fine and required by other - # tests. - body = triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ) - assert "No need to ping" in body or "no need to ping" in body - - def test_grace_warnings_should_show_what_got_right(self, triage_module): - # The "What you got right" section must appear in the grace warning - # too, not only the close comment — the contributor sees the warning - # first and that's their best chance to know what to keep. - pr_body = triage_module.format_grace_warning_pr_comment( - { - "verdict": "fail", - "linked_issue": True, - "has_problem_description": True, - "has_expected_vs_actual": True, - "has_qa_proof": False, - "missing": ["QA proof"], - "explanation": "thin", - } - ) - assert "What you got right" in pr_body - assert "Linked a related GitHub issue" in pr_body - - issue_body = triage_module.format_grace_warning_issue_comment( - { - "verdict": "fail", - "kind": "feature", - "has_motivation_example": True, - "missing": ["concrete description"], - "explanation": "vague", - } - ) - assert "What you got right" in issue_body - assert "Motivation and concrete example" in issue_body - - def test_grace_warnings_should_use_softer_park_for_later_framing( - self, triage_module - ): - # Same softer-framing pin as the close comment, but for the warning - # — the contributor's first contact with the bot must not read as a - # hard deadline / ultimatum. - for body in ( - triage_module.format_grace_warning_pr_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - triage_module.format_grace_warning_issue_comment( - {"verdict": "fail", "missing": [], "explanation": ""} - ), - ): - assert "park this for later" in body - assert ( - "not a rejection" in body - or "isn't a rejection" in body - or ("isn't us saying" in body) - ) - - -class TestSecondsSinceLastGraceWarning: - """Mirror of TestSecondsSinceLastReconsiderVerdict for the new helper. - Both helpers share `_seconds_since_latest_marker_comment` underneath - so the parsing logic is exercised either way; these tests pin the - grace-marker-specific behavior.""" - - def _make_comment( - self, - *, - login: str, - body: str, - created_at: str | None = "2026-05-18T05:00:00Z", - ) -> dict: - comment: dict = {"user": {"login": login}, "body": body} - if created_at is not None: - comment["created_at"] = created_at - return comment - - def test_should_return_none_when_no_grace_marker(self, triage_module, monkeypatch): - comments = [ - self._make_comment( - login="github-actions[bot]", - body="Some other bot message", - ), - self._make_comment(login="random-user", body="ping?"), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_grace_warning("o/r", 1) is None - - def test_should_ignore_non_bot_comments_with_marker( - self, triage_module, monkeypatch - ): - # A user who quotes the marker in a question must NOT be treated - # as the bot warning; otherwise the close-after-grace path would - # never fire because the timer keeps resetting. - comments = [ - self._make_comment( - login="random-user", - body=f"What is {triage_module.GRACE_COMMENT_MARKER}?", - ) - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - assert triage_module.seconds_since_last_grace_warning("o/r", 1) is None - - def test_should_pick_latest_grace_marker(self, triage_module, monkeypatch): - comments = [ - self._make_comment( - login="github-actions[bot]", - body="old warning " + triage_module.GRACE_COMMENT_MARKER, - created_at="2026-05-18T03:00:00Z", - ), - self._make_comment( - login="github-actions[bot]", - body="newer warning " + triage_module.GRACE_COMMENT_MARKER, - created_at="2026-05-18T04:55:00Z", - ), - ] - monkeypatch.setattr( - triage_module, "_iter_paginated_json", lambda *a, **kw: iter(comments) - ) - - import datetime as real_dt - - class FrozenDateTime(real_dt.datetime): - @classmethod - def now(cls, tz=None): - return real_dt.datetime(2026, 5, 18, 5, 0, 0, tzinfo=tz) - - frozen_module = type(triage_module.dt)("datetime") - frozen_module.datetime = FrozenDateTime - frozen_module.timezone = real_dt.timezone - monkeypatch.setattr(triage_module, "dt", frozen_module) - - age = triage_module.seconds_since_last_grace_warning("o/r", 1) - # Newer warning is 5 minutes (300s) before "now". - assert age == 300.0 - - -class TestTriageAllowlist: - """The dogfood allowlist gates `triage`: while non-empty it is the sole - author filter (only the named accounts are acted on) and it bypasses the - internal-author exemption for them, so a maintainer can dogfood on their - own org account. Emptying it restores the internal-author skip.""" - - def _make_pr(self, **overrides): - base = { - "number": 1, - "title": "PR title", - "body": "Body with no linked issue and no QA proof.", - "state": "open", - "author_association": "NONE", - "user": {"login": "mateo-berri"}, - } - base.update(overrides) - return base - - def test_should_skip_author_not_on_allowlist(self, triage_module, monkeypatch): - pr = self._make_pr(user={"login": "random-oss-dev"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM must not run for non-allowlisted author"), - ) - assert result["action"] == "skip-not-allowlisted" - - def test_should_act_on_allowlisted_internal_author( - self, triage_module, monkeypatch - ): - pr = self._make_pr(author_association="MEMBER", user={"login": "mateo-berri"}) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: json.dumps( - {"verdict": "pass", "missing": [], "explanation": "ok"} - ), - ) - assert result["action"] == "pass-llm" - - def test_empty_allowlist_restores_internal_skip(self, triage_module, monkeypatch): - pr = self._make_pr( - author_association="MEMBER", user={"login": "krrishdholakia"} - ) - monkeypatch.setattr(triage_module, "fetch_pr", lambda repo, n: pr) - result = triage_module.triage( - repo="o/r", - kind="pr", - number=1, - close=True, - model="m", - judge=lambda p: pytest.fail("LLM must not run for internal author"), - allowlist=frozenset(), - ) - assert result["action"] == "skip-internal-author" - - def test_allowlist_constant_is_the_two_dogfood_accounts(self, triage_module): - assert triage_module.ALLOWLIST_LOGINS == frozenset( - {"mateo-berri", "swiftwinds"} - ) - for login in triage_module.ALLOWLIST_LOGINS: - assert login == login.lower(), login diff --git a/tests/test_litellm/test_github_triage_workflows.py b/tests/test_litellm/test_github_triage_workflows.py deleted file mode 100644 index f96c9b7e974..00000000000 --- a/tests/test_litellm/test_github_triage_workflows.py +++ /dev/null @@ -1,264 +0,0 @@ -"""Static guardrails for the Agent Shin + Greptile workflow YAML files. - -These workflows can post comments and close PRs/issues on -BerriAI/litellm, so the gating logic that decides "is this a real -close-on-fail run?" must fail-safe on any unexpected input. The risk -is mostly maintenance: someone edits the bash gate, drops a quote, -inverts a comparison, or uses `!= "false"` (which treats "True", -"yes", "1", and typos as enabling closure) and the regression isn't -caught until a real OSS contributor's PR gets auto-closed. - -The tests below pin a set of invariants. The first two apply to every -workflow that gates a destructive `--close`: - - 1. The gate uses the fail-safe `= "true"` comparison — not `!= "false"`, - not `!= ""`. Only the literal string "true" should ever enable - closure. - 2. The gate also requires `AGENT_SHIN_ENABLED = "true"` (or the - scheduled-job equivalent) — disabling the variable must always - force dry-run. - -A third invariant covers every workflow that installs the OpenAI client. -These run with a write-scoped `GITHUB_TOKEN`, so a compromised package -release would execute in that context; the install must therefore come -from the hash-pinned `.github/scripts/triage-requirements.txt` via -`pip --require-hashes`, never a floating `pip install openai>=...`. - -Static parsing of the YAML + bash text is the right level of test here: -the gating logic lives in a `run:` block, not in a Python module we can -import, and end-to-end testing a GitHub Actions workflow from CI is -infeasible. A YAML-level guardrail is exactly what would have caught -the original `!= "false"` regression at PR time. -""" - -from __future__ import annotations - -from pathlib import Path - -import pytest -import yaml - -REPO_ROOT = Path(__file__).resolve().parents[2] -WORKFLOWS_DIR = REPO_ROOT / ".github" / "workflows" - -# Map of workflow file -> the env var name that drives the destructive -# gate inside that workflow's `run:` block. Keeping this table explicit -# (rather than scraping every workflow file) means a new workflow file -# that bypasses the dry-run gating doesn't silently slip past this test. -DESTRUCTIVE_GATE_ENV: dict[str, str] = { - "close_low_quality_prs.yml": "CLOSE_FLAG", - # The reconsider workflow has no per-run "really do it?" knob — its - # only kill switch is `AGENT_SHIN_ENABLED`, which already serves as - # both the destructive gate and the global enablement gate. - "triage_reconsider.yml": "AGENT_SHIN_ENABLED", -} - - -# Privileged workflows that install the OpenAI client. They run with a -# write-scoped GITHUB_TOKEN, so the install must be hash-pinned: a poisoned -# release would otherwise execute in that context. A new workflow that -# installs the client must be added here and use the same pinned file. -LLM_CLIENT_INSTALLER_WORKFLOWS = ( - "triage_reconsider.yml", -) - -PINNED_INSTALL = "--require-hashes -r .github/scripts/triage-requirements.txt" -REQUIREMENTS_FILE = REPO_ROOT / ".github" / "scripts" / "triage-requirements.txt" - - -def _load_workflow(name: str) -> dict: - return yaml.safe_load((WORKFLOWS_DIR / name).read_text()) - - -def _all_run_blocks(workflow: dict) -> list[str]: - """Return every `run:` step's command text, joined.""" - commands: list[str] = [] - jobs = workflow.get("jobs") or {} - for job in jobs.values(): - for step in job.get("steps", []) or []: - if not isinstance(step, dict): - continue - run = step.get("run") - if isinstance(run, str): - commands.append(run) - return commands - - -@pytest.mark.parametrize("workflow_file,env_var", sorted(DESTRUCTIVE_GATE_ENV.items())) -def test_should_use_failsafe_equals_true_comparison(workflow_file: str, env_var: str) -> None: - """The destructive `--close` gate must use `= "true"` (fail-safe), not - `!= "false"` (which would treat "True", "yes", "1", or any typo as - enabling closure). - - Both bare `${ENV_VAR}` and `${ENV_VAR:-false}` (with a default) are - accepted forms — what matters is the comparison operator. The - Greptile closer relies on an outer `AGENT_SHIN_ENABLED` gate so it - can use the bare form; the Agent Shin workflows include `:-false` - for defense in depth. Either is fine. - """ - workflow = _load_workflow(workflow_file) - text = "\n".join(_all_run_blocks(workflow)) - assert env_var in text, ( - f"{workflow_file} no longer references {env_var}; was the gating env var renamed without updating this test?" - ) - accepted_patterns = ( - f'"${{{env_var}}}" = "true"', - f'"${{{env_var}:-false}}" = "true"', - ) - assert any(p in text for p in accepted_patterns), ( - f"{workflow_file} must gate the destructive --close flag on the " - f'EXACT string "true" (one of: {accepted_patterns!r}). Mirror ' - 'the Greptile closer pattern; do NOT use `!= "false"` which ' - 'fail-opens on unknown values like "True", "yes", "1", or typos.' - ) - forbidden_patterns = ( - f'"${{{env_var}}}" != "false"', - f'"${{{env_var}:-false}}" != "false"', - f'"${{{env_var}:-true}}" != "false"', - ) - for forbidden in forbidden_patterns: - assert forbidden not in text, ( - f"{workflow_file} uses the fail-open pattern {forbidden!r}. " - 'Switch to `= "true"` so unknown values stay dry-run.' - ) - - -@pytest.mark.parametrize("workflow_file", sorted(DESTRUCTIVE_GATE_ENV)) -def test_should_require_agent_shin_enabled_for_close(workflow_file: str) -> None: - """Every destructive gate must also gate on the global enablement - variable, so flipping `AGENT_SHIN_ENABLED` off is a kill switch - regardless of any per-run input. - - Two patterns are equally fine: - - Positive: `[ "${AGENT_SHIN_ENABLED:-false}" = "true" ]` to enter - the close branch (Agent Shin workflows). - - Negative: `[ "${AGENT_SHIN_ENABLED:-false}" != "true" ]` then - bail out / force dry-run (Greptile closer). - - What matters is that the comparison value is the literal "true"; - `!= "false"` or `= "1"` etc. would not be a true kill switch. - """ - workflow = _load_workflow(workflow_file) - text = "\n".join(_all_run_blocks(workflow)) - accepted_patterns = ( - '"${AGENT_SHIN_ENABLED:-false}" = "true"', - '"${AGENT_SHIN_ENABLED:-false}" != "true"', - ) - assert any(p in text for p in accepted_patterns), ( - f"{workflow_file} must gate destructive actions on " - '`AGENT_SHIN_ENABLED = "true"` (or the inverted `!= "true"` ' - "guard that forces dry-run). Without this, an unset repo " - "variable would not be treated as a kill switch." - ) - - -@pytest.mark.parametrize("workflow_file", LLM_CLIENT_INSTALLER_WORKFLOWS) -def test_llm_client_install_is_hash_pinned(workflow_file: str) -> None: - """Every privileged workflow installs the OpenAI client from the - hash-pinned requirements file, never by floating version. - - A bare `pip install "openai>=1.40.0"` resolves to whatever PyPI serves - at run time and executes during install/import while a write-scoped - `GITHUB_TOKEN` is in scope, so a compromised release runs in a - privileged context. This test fails if that floating form comes back or - if the `--require-hashes` install is loosened. - """ - blocks = _all_run_blocks(_load_workflow(workflow_file)) - assert PINNED_INSTALL in "\n".join(blocks), ( - f"{workflow_file} must install the client via `pip install " - f"{PINNED_INSTALL}`; a floating install runs unverified code with a " - "write-scoped token." - ) - offenders = [b for b in blocks if "pip install" in b and "openai" in b] - assert not offenders, ( - f"{workflow_file} installs openai by name ({offenders!r}); pin it " - "through the hash-locked requirements file so the version and " - "checksum are fixed." - ) - - -def test_triage_requirements_are_fully_hash_pinned() -> None: - """The shared requirements file pins every package to an exact version - with a sha256 hash, which is what `pip --require-hashes` enforces at - install time. A loosened pin or a missing hash here would silently widen - the supply-chain surface for all the installer workflows. - """ - assert REQUIREMENTS_FILE.exists(), ( - f"the hash-pinned requirements file the triage workflows install from is missing at {REQUIREMENTS_FILE}" - ) - joined = REQUIREMENTS_FILE.read_text().replace("\\\n", " ") - entries = [line.strip() for line in joined.splitlines() if line.strip() and not line.strip().startswith("#")] - assert any(e.split()[0].startswith("openai==") for e in entries), ( - "openai must be pinned to an exact version in the triage requirements" - ) - for entry in entries: - spec = entry.split()[0] - assert "==" in spec, ( - f"requirement {spec!r} is not pinned to an exact version; " - "--require-hashes needs every package pinned with ==" - ) - assert "--hash=sha256:" in entry, ( - f"requirement {spec!r} has no sha256 hash; every pin must carry " - "checksums so --require-hashes can verify the download" - ) - - -def _reconsider_steps() -> list[dict]: - workflow = _load_workflow("triage_reconsider.yml") - return workflow["jobs"]["reconsider"]["steps"] - - -def _index_of_run_step(steps: list[dict], needle: str) -> int: - for i, step in enumerate(steps): - run = step.get("run") - if isinstance(run, str) and needle in run: - return i - raise AssertionError(f"no run step contains {needle!r}") - - -def _reaction_steps(steps: list[dict], content: str) -> list[tuple[int, dict]]: - return [ - (i, s) - for i, s in enumerate(steps) - if isinstance(s.get("run"), str) and f"content={content}" in s["run"] and "/reactions" in s["run"] - ] - - -class TestReconsiderReactions: - """The reconsider workflow acknowledges the triggering comment with a 👀 - reaction the moment it accepts the trigger, and a 👍 once the run finishes, - so the contributor gets feedback immediately instead of waiting on a cron. - - Both reactions are gated on `AGENT_SHIN_ENABLED == 'true'` so a dry-run - leaves no visible trace, and both target the comment that fired the event - (`github.event.comment.id`). The ordering (👀 before the triage run, 👍 - after) is the whole point — these tests fail if a refactor reorders the - steps, drops a reaction, or stops gating them. - """ - - def test_eyes_reaction_is_posted_before_the_triage_run(self) -> None: - steps = _reconsider_steps() - run_idx = _index_of_run_step(steps, "triage_with_llm.py") - eyes = _reaction_steps(steps, "eyes") - assert len(eyes) == 1, "expected exactly one 👀 (eyes) reaction step" - idx, step = eyes[0] - assert idx < run_idx, "👀 must be posted BEFORE the slow triage run, not after" - assert "github.event.comment.id" in (step.get("env") or {}).get("COMMENT_ID", ""), ( - "👀 must react to the comment that triggered the workflow" - ) - assert "${COMMENT_ID}" in step["run"], "👀 must react to the triggering comment, not a hardcoded id" - assert "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"], ( - "👀 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert" - ) - - def test_thumbs_up_reaction_is_posted_after_a_successful_run(self) -> None: - steps = _reconsider_steps() - run_idx = _index_of_run_step(steps, "triage_with_llm.py") - thumbs = _reaction_steps(steps, "+1") - assert len(thumbs) == 1, "expected exactly one 👍 (+1) reaction step" - idx, step = thumbs[0] - assert idx > run_idx, "👍 must come AFTER the triage run" - assert "success()" in step["if"], "👍 must only fire when the reconsider run succeeded" - assert "vars.AGENT_SHIN_ENABLED == 'true'" in step["if"], ( - "👍 must be gated on AGENT_SHIN_ENABLED so dry-run stays inert" - ) From 8005856411ae43091c86ba2632d389916b8fa0ec Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:03:27 -0700 Subject: [PATCH 248/253] ci(build_and_test): seed the routing strategy through /config/update --- .circleci/config.yml | 6 ++++++ proxy_server_config.yaml | 1 - 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index df17a9e4402..87f1ee604cf 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1785,6 +1785,12 @@ jobs: - wait_for_service: url: http://localhost:4000 timeout: "300" + - run: + name: Seed the routing strategy through /config/update + command: | + curl --noproxy '*' -sSf -X POST http://localhost:4000/config/update \ + -H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \ + -d '{"router_settings": {"routing_strategy": "usage-based-routing-v2"}}' - run: name: Run tests command: | diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index 73990153227..703d56bc0cd 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -213,7 +213,6 @@ files_settings: api_key: os.environ/OPENAI_API_KEY router_settings: - routing_strategy: usage-based-routing-v2 redis_host: os.environ/REDIS_HOST redis_password: os.environ/REDIS_PASSWORD redis_port: os.environ/REDIS_PORT From a4624b6c6c4bb1ddf753bd57bbfd37b711598c90 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:03:45 -0700 Subject: [PATCH 249/253] fix(gemini): derive the finish reason key set from the Candidates type Candidates.finishReason listed eleven values while the mapping key set carried twenty-one, so typed fixtures could not spell the reasons this PR handles. GeminiFinishReason is now the one list, the key set derives from it, and a test checks every documented reason has an explicit mapping instead of falling through to "stop" --- .../vertex_and_google_ai_studio_gemini.py | 29 ++------------ litellm/types/llms/vertex_ai.py | 39 ++++++++++++------- ...test_vertex_and_google_ai_studio_gemini.py | 9 ++++- 3 files changed, 36 insertions(+), 41 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 a95b845718a..46f1b948026 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 @@ -6,7 +6,7 @@ import time from collections.abc import Callable, Mapping, Sequence from copy import deepcopy from functools import partial -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args import httpx @@ -57,6 +57,7 @@ from litellm.types.llms.vertex_ai import ( ContentType, FunctionCallingConfig, FunctionDeclaration, + GeminiFinishReason, GeminiThinkingConfig, GenerateContentResponseBody, HttpxPartType, @@ -1330,31 +1331,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.", } - _GEMINI_FINISH_REASON_KEYS = frozenset( - { - "STOP", - "MAX_TOKENS", - "SAFETY", - "RECITATION", - "FINISH_REASON_UNSPECIFIED", - "MALFORMED_FUNCTION_CALL", - "LANGUAGE", - "OTHER", - "BLOCKLIST", - "PROHIBITED_CONTENT", - "SPII", - "IMAGE_SAFETY", - "IMAGE_PROHIBITED_CONTENT", - "TOO_MANY_TOOL_CALLS", - "MALFORMED_RESPONSE", - "NO_IMAGE", - "IMAGE_RECITATION", - "IMAGE_OTHER", - "ESCALATION", - "UNEXPECTED_TOOL_CALL", - "MISSING_THOUGHT_SIGNATURE", - } - ) + _GEMINI_FINISH_REASON_KEYS: Final[frozenset[str]] = frozenset(get_args(GeminiFinishReason)) @staticmethod def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]: diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 3b95b786631..ce51e46ef15 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -425,22 +425,35 @@ class UrlContextMetadata(TypedDict, total=False): urlMetadata: list[UrlMetadata] +GeminiFinishReason = Literal[ + "FINISH_REASON_UNSPECIFIED", + "STOP", + "MAX_TOKENS", + "SAFETY", + "RECITATION", + "LANGUAGE", + "OTHER", + "BLOCKLIST", + "PROHIBITED_CONTENT", + "SPII", + "MALFORMED_FUNCTION_CALL", + "IMAGE_SAFETY", + "IMAGE_PROHIBITED_CONTENT", + "TOO_MANY_TOOL_CALLS", + "MALFORMED_RESPONSE", + "NO_IMAGE", + "IMAGE_RECITATION", + "IMAGE_OTHER", + "ESCALATION", + "UNEXPECTED_TOOL_CALL", + "MISSING_THOUGHT_SIGNATURE", +] + + class Candidates(TypedDict, total=False): index: int content: HttpxContentType - finishReason: Literal[ - "FINISH_REASON_UNSPECIFIED", - "STOP", - "MAX_TOKENS", - "SAFETY", - "RECITATION", - "OTHER", - "BLOCKLIST", - "PROHIBITED_CONTENT", - "SPII", - "MALFORMED_FUNCTION_CALL", - "IMAGE_SAFETY", - ] + finishReason: GeminiFinishReason safetyRatings: list[SafetyRatings] citationMetadata: CitationMetadata groundingMetadata: GroundingMetadata 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 4ee199bd9af..6c818016c87 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 Final, List, cast +from typing import Final, List, cast, get_args from unittest.mock import MagicMock, patch import httpx @@ -18,7 +18,7 @@ from litellm.llms.vertex_ai.common_utils import VertexAIError from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) -from litellm.types.llms.vertex_ai import UsageMetadata +from litellm.types.llms.vertex_ai import GeminiFinishReason, UsageMetadata from litellm.types.utils import ChoiceLogprobs, Usage from litellm.utils import CustomStreamWrapper @@ -940,6 +940,11 @@ def test_check_finish_reason(): ) +def test_every_documented_gemini_finish_reason_has_an_explicit_mapping(): + documented: Final = frozenset(get_args(GeminiFinishReason)) + assert set(VertexGeminiConfig.get_finish_reason_mapping()) == documented + + def test_finish_reason_unspecified_and_malformed_function_call(): """ Test that FINISH_REASON_UNSPECIFIED and MALFORMED_FUNCTION_CALL From 5db2a0c8850385b9e56b17bcd8053bab4fac0b59 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:05:42 -0700 Subject: [PATCH 250/253] test(proxy): type the sqlstate test parameters --- tests/test_litellm/proxy/db/test_exception_handler.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 26ac1ea65ad..f7cc5e3ed83 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -681,7 +681,7 @@ def test_is_deadlock_error_excludes_non_deadlocks(error): (httpx.ReadTimeout("no reply"), None), ], ) -def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error, sqlstate): +def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error: Exception, sqlstate: str | None): """Only a prisma data error carrying Postgres's own error code yields a SQLSTATE; a codeless or malformed payload, an engine-level error, and a transport error yield None.""" assert PrismaDBExceptionHandler.postgres_sqlstate(error) == sqlstate From 1704aeebb49eea499a863a2dfbd9c990b0e2b501 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 18 Sep 2026 23:06:47 +0000 Subject: [PATCH 251/253] fix(enterprise): resolve openai_moderations model at call time and default to omni-moderation-latest Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../enterprise_hooks/openai_moderation.py | 9 +++-- litellm/constants.py | 2 ++ .../guardrails/test_guardrail_coverage.py | 36 +++++++++++++++++++ 3 files changed, 42 insertions(+), 5 deletions(-) diff --git a/enterprise/enterprise_hooks/openai_moderation.py b/enterprise/enterprise_hooks/openai_moderation.py index 2162370804a..017f51bfabd 100644 --- a/enterprise/enterprise_hooks/openai_moderation.py +++ b/enterprise/enterprise_hooks/openai_moderation.py @@ -17,6 +17,7 @@ from fastapi import HTTPException import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_OPENAI_MODERATIONS_MODEL from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails._content_utils import iter_message_text @@ -24,11 +25,9 @@ from litellm.types.utils import CallTypesLiteral class _ENTERPRISE_OpenAI_Moderation(CustomLogger): - def __init__(self): - self.model_name = ( - litellm.openai_moderations_model_name or "text-moderation-latest" - ) # pass the model_name you initialized on litellm.Router() - pass + @property + def model_name(self) -> str: + return litellm.openai_moderations_model_name or DEFAULT_OPENAI_MODERATIONS_MODEL #### CALL HOOKS - proxy only #### diff --git a/litellm/constants.py b/litellm/constants.py index e7cb21a3a7d..a7d4eba0f15 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -158,6 +158,8 @@ DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL: Final = str( ) DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD", 0.75)) +DEFAULT_OPENAI_MODERATIONS_MODEL: Final = "omni-moderation-latest" + # MCP OAuth2 Client Credentials Defaults MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS: Final = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60")) MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200")) diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py index f25e83b1672..548677c70bc 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_coverage.py @@ -18,7 +18,9 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from httpx import Request, Response +import litellm from litellm import DualCache +from litellm.constants import DEFAULT_OPENAI_MODERATIONS_MODEL from litellm.proxy._types import UserAPIKeyAuth from litellm.types.utils import Choices, Message, ModelResponse @@ -764,6 +766,40 @@ async def test_openai_moderation_inspects_multimodal_content(monkeypatch, user_a assert seen_inputs == ["alpha beta"] +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("configured_after_init", "expected_model"), + [("omni-moderation-2024-09-26", "omni-moderation-2024-09-26"), (None, DEFAULT_OPENAI_MODERATIONS_MODEL)], +) +async def test_openai_moderation_reads_model_name_at_call_time( + monkeypatch, user_api_key, configured_after_init, expected_model +): + """``litellm_settings`` applies ``callbacks`` and ``openai_moderations_model_name`` in YAML + order, so the hook must resolve the model when it runs, not when it is constructed.""" + from enterprise.enterprise_hooks.openai_moderation import ( + _ENTERPRISE_OpenAI_Moderation, + ) + + monkeypatch.setattr(litellm, "openai_moderations_model_name", None) + guard = _ENTERPRISE_OpenAI_Moderation() + monkeypatch.setattr(litellm, "openai_moderations_model_name", configured_after_init) + + class FakeModeration: + results = [type("R", (), {"flagged": False})()] + + fake_router = MagicMock() + fake_router.amoderation = AsyncMock(return_value=FakeModeration()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router, raising=False) + + await guard.async_moderation_hook( + data={"messages": [{"role": "user", "content": "hello"}]}, + user_api_key_dict=user_api_key, + call_type="acompletion", + ) + + fake_router.amoderation.assert_awaited_once_with(model=expected_model, input="hello") + + # ── Google Text Moderation ──────────────────────────────────────────────────── From c6c8aed3f8594f89a86b91df553b84f6aab2fb20 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:22:04 -0700 Subject: [PATCH 252/253] fix(proxy): drop only the daily spend batch whose failure cannot be re-sent, requeue the unsent ones --- litellm/proxy/db/db_spend_update_writer.py | 30 ++++++----- .../proxy/db/test_db_spend_update_writer.py | 54 +++++++++++++++++++ 2 files changed, 71 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e9967fb0d67..37eac8604bd 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1332,15 +1332,7 @@ class DBSpendUpdateWriter: daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions), ) except Exception as e: # noqa: BLE001 # whatever failed here, the other tables must still flush - if not _daily_spend_commit_failure_is_requeue_safe(e): - spend_log_error( - "Spend tracking - dropped %d daily %s spend rows: the failed commit may have applied " - "or the database refused the data, so re-sending it is not safe. Error: %s", - len(transactions), - entity_type, - str(e), - exc=e, - ) + if not transactions: return spend_log_error( "Spend tracking - failed to commit daily %s spend updates. " @@ -2050,13 +2042,25 @@ class DBSpendUpdateWriter: sql, params = build_bulk_upsert(table=table, batch=merged_batch) await prisma_client.db.execute_raw(sql, *params) except Exception as batch_error: - # Log detailed error information for debugging batch upsert failures - # This helps diagnose issues like unique constraint violations + if _daily_spend_commit_failure_is_requeue_safe(batch_error): + spend_log_error( + "Daily %s spend batch upsert failed. Table: %s, Rows: %d, Error: %s", + entity_type, + table.name, + len(transactions_to_process), + str(batch_error), + exc=batch_error, + ) + raise + for key in transactions_to_process: + daily_spend_transactions.pop(key, None) spend_log_error( - "Daily %s spend batch upsert failed. Table: %s, Rows: %d, Error: %s", + "Spend tracking - dropped %d daily %s spend rows: the failed statement may have " + "applied or the database refused the data, so re-sending it is not safe. " + "Table: %s, Error: %s", + len(transactions_to_process), entity_type, table.name, - len(transactions_to_process), str(batch_error), exc=batch_error, ) diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 155bca656d5..20f1fa9d363 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1624,6 +1624,33 @@ async def test_update_daily_spend_keeps_failed_transactions_for_retry(): assert daily_spend_transactions == expected +@pytest.mark.asyncio +async def test_update_daily_spend_drops_the_batch_whose_failure_cannot_be_resent(): + """A reply lost after the statement was sent may already have applied, so the batch is + taken out of the caller's dict before the error propagates: whichever requeue the caller + runs afterwards, the Redis restore included, cannot send it a second time.""" + + def lose_the_reply() -> int: + raise httpx.ReadTimeout("no reply") + + prisma_client = _RecordingPrisma(execute_raw=lose_the_reply) + daily_spend_transactions = {"user-key": _daily_txn(user_id="user-1")} + proxy_logging_obj = MagicMock() + proxy_logging_obj.failure_handler = AsyncMock() + + with pytest.raises(httpx.ReadTimeout): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + ) + + assert daily_spend_transactions == {} + + @pytest.mark.asyncio async def test_commit_key_spend_updates_includes_last_active(): """ @@ -2875,6 +2902,33 @@ async def test_failed_daily_spend_commit_is_requeued_only_when_the_rows_are_prov assert db_writer.daily_spend_update_queue.update_queue.empty() +@pytest.mark.asyncio +async def test_failed_daily_spend_commit_drops_only_the_batch_that_was_sent(): + """A tick holding more than one batch of 100 rows sends them one statement at a time, and + a reply lost on one statement says nothing about the batches after it: only the batch that + was on the wire is dropped, the ones never sent go back on the queue and land next tick.""" + db_writer = DBSpendUpdateWriter() + await db_writer.daily_spend_update_queue.add_update( + {f"user-{i:03d}": _daily_txn(user_id=f"user-{i:03d}") for i in range(150)} + ) + db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend", failure=httpx.ReadTimeout("no reply")) + db_writer._flush_tool_discovery_queue = AsyncMock() + proxy_logging_obj = MagicMock() + proxy_logging_obj.failure_handler = AsyncMock() + + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + db.failing_table = None + await db_writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj + ) + + (upsert,) = _daily_upserts(db, "LiteLLM_DailyUserSpend") + assert _row_values(upsert, "user_id") == [f"user-{i:03d}" for i in range(100, 150)] + assert db_writer.daily_spend_update_queue.update_queue.empty() + + @pytest.mark.asyncio async def test_failed_daily_spend_commit_requeues_the_rows_and_flushes_the_other_tables(): """With the Redis buffer off, a daily batch that failed to commit was discarded along From 83d89aa134bbeac391ce4a49dd63223f52b97019 Mon Sep 17 00:00:00 2001 From: kerry Date: Sat, 19 Sep 2026 00:08:54 +0000 Subject: [PATCH 253/253] test(integration): cover off-peak pricing on a live proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/_support/client.py | 4 +- tests/integration/contracts.json | 6 ++ .../pricing/test_off_peak_pricing.py | 82 +++++++++++++++++++ 3 files changed, 90 insertions(+), 2 deletions(-) create mode 100644 tests/integration/pricing/test_off_peak_pricing.py diff --git a/tests/integration/_support/client.py b/tests/integration/_support/client.py index 97522e5728c..9f1118ab1e3 100644 --- a/tests/integration/_support/client.py +++ b/tests/integration/_support/client.py @@ -161,7 +161,7 @@ class Scenario: assert all(object_value(object_value(entry)["model_info"])["id"] != identity for entry in entries) assert read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,)) == [] - def model(self, **parameters: JsonValue) -> str: + def model(self, *, model_info: Mapping[str, JsonValue] | None = None, **parameters: JsonValue) -> str: name: Final = f"integration-{uuid.uuid4().hex}" created: Final = self.gateway.post( "/model/new", @@ -173,7 +173,7 @@ class Scenario: "api_base": f"{self.gateway.upstream_url}/v1", **parameters, }, - "model_info": {}, + "model_info": dict(model_info) if model_info is not None else {}, }, ) identity: Final = string_value(object_value(created["model_info"])["id"]) diff --git a/tests/integration/contracts.json b/tests/integration/contracts.json index 91b1bd86954..6958ade50f7 100644 --- a/tests/integration/contracts.json +++ b/tests/integration/contracts.json @@ -92,6 +92,12 @@ "tests/integration/pricing/test_price_precedence.py::test_same_upstream_aliases_keep_distinct_prices_after_reload": [ "quota_management.spend_tracking.alias_prices.remain_independent_on_reload" ], + "tests/integration/pricing/test_off_peak_pricing.py::test_open_off_peak_window_bills_off_peak_rates": [ + "quota_management.spend_tracking.off_peak_pricing.open_window_bills_off_peak_rates" + ], + "tests/integration/pricing/test_off_peak_pricing.py::test_closed_off_peak_window_bills_standard_rates": [ + "quota_management.spend_tracking.off_peak_pricing.closed_window_bills_standard_rates" + ], "tests/integration/spend/test_cache_and_quota.py::test_generated_cache_sequences_preserve_content_usage_and_zero_hit_cost": [ "quota_management.response_cache.generated_sequences_preserve_content_and_accounting" ], diff --git a/tests/integration/pricing/test_off_peak_pricing.py b/tests/integration/pricing/test_off_peak_pricing.py new file mode 100644 index 00000000000..5623356c078 --- /dev/null +++ b/tests/integration/pricing/test_off_peak_pricing.py @@ -0,0 +1,82 @@ +import json +from collections.abc import Mapping +from datetime import datetime, timedelta, timezone +from typing import Final + +import pytest +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows + +STANDARD_INPUT_RATE: Final = 0.001 +STANDARD_OUTPUT_RATE: Final = 0.002 +OFF_PEAK_INPUT_RATE: Final = 0.0001 +OFF_PEAK_OUTPUT_RATE: Final = 0.0002 + + +def off_peak_window(start_offset_hours: int, end_offset_hours: int) -> Mapping[str, JsonValue]: + now: Final = datetime.now(timezone.utc) + start: Final = now + timedelta(hours=start_offset_hours) + end: Final = now + timedelta(hours=end_offset_hours) + return { + "hours_utc": f"{start:%H:%M}-{end:%H:%M}", + "input_cost_per_token": OFF_PEAK_INPUT_RATE, + "output_cost_per_token": OFF_PEAK_OUTPUT_RATE, + } + + +def billed_model(scenario: Scenario, off_peak: Mapping[str, JsonValue]) -> str: + return scenario.model( + input_cost_per_token=STANDARD_INPUT_RATE, + output_cost_per_token=STANDARD_OUTPUT_RATE, + model_info={"off_peak_pricing": dict(off_peak)}, + ) + + +def assert_chat_bills_rates(gateway: Gateway, model: str, input_rate: float, output_rate: float) -> None: + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "off peak control"}]} + ) + assert response.status_code == 200, response.text + expected: Final = 20 * input_rate + 20 * output_rate + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6) + request_id: Final = string_value(object_value(response.json())["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["prompt_tokens"] == 20 + assert rows[0]["completion_tokens"] == 20 + assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + breakdown: Final = object_value(parsed["cost_breakdown"]) + assert float(breakdown["input_cost"]) == pytest.approx(20 * input_rate, rel=1e-6) + assert float(breakdown["output_cost"]) == pytest.approx(20 * output_rate, rel=1e-6) + + +@pytest.mark.covers("quota_management.spend_tracking.off_peak_pricing.open_window_bills_off_peak_rates") +def test_open_off_peak_window_bills_off_peak_rates(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = billed_model(scenario, off_peak_window(-1, 1)) + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + matching: Final = tuple(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model) + assert len(matching) == 1 + info: Final = object_value(matching[0]["model_info"]) + off_peak: Final = object_value(info["off_peak_pricing"]) + assert off_peak["input_cost_per_token"] == OFF_PEAK_INPUT_RATE + assert off_peak["output_cost_per_token"] == OFF_PEAK_OUTPUT_RATE + assert_chat_bills_rates(gateway, model, OFF_PEAK_INPUT_RATE, OFF_PEAK_OUTPUT_RATE) + + +@pytest.mark.covers("quota_management.spend_tracking.off_peak_pricing.closed_window_bills_standard_rates") +def test_closed_off_peak_window_bills_standard_rates(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = billed_model(scenario, off_peak_window(2, 3)) + assert_chat_bills_rates(gateway, model, STANDARD_INPUT_RATE, STANDARD_OUTPUT_RATE)