From bde00952b6aa68372e739e1f377cc9c32a9063f1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 18 Jul 2026 23:19:21 +0000 Subject: [PATCH 01/27] fix(proxy): requeue Redis spend buffer transactions when DB commit fails The Redis transaction buffer leader drains the spend buffers with a destructive lpop before committing to the database. When the DB commit failed after exhausting retries, the popped transactions were only logged and then lost, permanently undercounting key/user/team/org/end-user/ team-member/tag/agent and daily spend after a database outage. Track each popped category and re-push the ones that were not committed back to their Redis buffers so a later scheduler tick retries them. Categories that already committed are not re-queued, so their spend is not double-counted. The daily tag spend path gets the same treatment. --- litellm/proxy/db/db_spend_update_writer.py | 45 +++++- .../redis_update_buffer.py | 55 +++++++ .../test_redis_update_buffer.py | 46 ++++++ .../proxy/db/test_db_spend_update_writer.py | 151 +++++++++++++++++- 4 files changed, 289 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 54a4c2dad91..cc266019ff7 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -797,6 +797,12 @@ class DBSpendUpdateWriter: ): verbose_proxy_logger.debug("acquired lock for spend updates") + # Track everything popped from Redis. Each category is removed once it + # has been committed to the DB, so whatever is left after a failure can + # be re-queued for the next tick instead of being lost. Committed + # categories are never re-queued, so their spend is not double-counted. + uncommitted: dict[str, Any] = {} # mutable-ok: drives which popped categories still need re-queuing + try: ( db_spend_update_transactions, @@ -807,6 +813,15 @@ class DBSpendUpdateWriter: daily_agent_spend_update_transactions, ) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline() + uncommitted = { # mutable-ok: drives which popped categories still need re-queuing + "db_spend_update_transactions": db_spend_update_transactions, + "daily_spend_update_transactions": daily_spend_update_transactions, + "daily_team_spend_update_transactions": daily_team_spend_update_transactions, + "daily_org_spend_update_transactions": daily_org_spend_update_transactions, + "daily_end_user_spend_update_transactions": daily_end_user_spend_update_transactions, + "daily_agent_spend_update_transactions": daily_agent_spend_update_transactions, + } + if db_spend_update_transactions is not None: verbose_proxy_logger.info( "Spend tracking - committing spend updates from Redis to DB: " @@ -826,6 +841,7 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, db_spend_update_transactions=db_spend_update_transactions, ) + uncommitted.pop("db_spend_update_transactions", None) if daily_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_user_spend( @@ -834,6 +850,8 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_spend_update_transactions, ) + uncommitted.pop("daily_spend_update_transactions", None) + if daily_team_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_team_spend( n_retry_times=n_retry_times, @@ -841,6 +859,7 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_team_spend_update_transactions, ) + uncommitted.pop("daily_team_spend_update_transactions", None) if daily_org_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_org_spend( @@ -849,6 +868,7 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_org_spend_update_transactions, ) + uncommitted.pop("daily_org_spend_update_transactions", None) if daily_end_user_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_end_user_spend( @@ -857,6 +877,8 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_end_user_spend_update_transactions, ) + uncommitted.pop("daily_end_user_spend_update_transactions", None) + if daily_agent_spend_update_transactions is not None: await DBSpendUpdateWriter.update_daily_agent_spend( n_retry_times=n_retry_times, @@ -864,14 +886,20 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_agent_spend_update_transactions, ) + uncommitted.pop("daily_agent_spend_update_transactions", None) except Exception as e: spend_log_error( "Spend tracking - failed to commit spend updates from Redis to DB. " - "Data already popped from Redis may be lost. Error: %s", + "Re-queuing uncommitted transactions to Redis for retry on next tick. Error: %s", str(e), exc=e, ) finally: + to_restore = { # mutable-ok: transient kwargs payload consumed immediately below + name: txns for name, txns in uncommitted.items() if txns is not None + } + if to_restore: + await self.redis_update_buffer.restore_transactions_to_redis(**to_restore) await self.pod_lock_manager.release_lock( cronjob_id=DB_SPEND_UPDATE_JOB_NAME, ) @@ -1020,11 +1048,11 @@ class DBSpendUpdateWriter: cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ): verbose_proxy_logger.debug("acquired lock for daily tag spend updates") + daily_tag_spend_update_transactions = ( + await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() + ) + committed = False try: - daily_tag_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() - ) - if daily_tag_spend_update_transactions: await DBSpendUpdateWriter.update_daily_tag_spend( n_retry_times=n_retry_times, @@ -1032,14 +1060,19 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_tag_spend_update_transactions, ) + committed = True except Exception as e: spend_log_error( "Spend tracking - failed to commit daily tag spend updates from Redis to DB. " - "Data already popped from Redis may be lost. Error: %s", + "Re-queuing to Redis for retry on next tick. Error: %s", str(e), exc=e, ) finally: + if not committed and daily_tag_spend_update_transactions: + await self.redis_update_buffer.restore_transactions_to_redis( + daily_tag_spend_update_transactions=daily_tag_spend_update_transactions, + ) await self.pod_lock_manager.release_lock( cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ) 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 c924448669d..b30fadd86ab 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -8,6 +8,8 @@ import asyncio import json from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast +from redis.exceptions import RedisError + from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache from litellm.constants import ( @@ -374,6 +376,59 @@ class RedisUpdateBuffer: if daily_txns: await daily_queue.update_queue.put(daily_txns) + async def restore_transactions_to_redis( + self, + db_spend_update_transactions: DBSpendUpdateTransactions | None = None, + daily_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_team_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_org_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_end_user_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_agent_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_tag_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + ) -> None: + """ + Re-push transactions that were popped from Redis but not committed to the DB. + + The leader drains the buffers with a destructive ``lpop`` before committing to + the database. When a commit fails after its retries are exhausted, the popped + transactions must be pushed back so a later scheduler tick can retry them; + otherwise the aggregated spend is lost permanently. The re-pushed payloads use + the same JSON encoding as the store path, so the next drain parses them normally. + """ + if self.redis_cache is None: + return + + _configs = ( + (db_spend_update_transactions, REDIS_UPDATE_BUFFER_KEY), + (daily_spend_update_transactions, REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY), + (daily_team_spend_update_transactions, REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY), + (daily_org_spend_update_transactions, REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY), + (daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY), + (daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY), + (daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY), + ) + + rpush_list: list[RedisPipelineRpushOperation] = [ # mutable-ok: async_rpush_pipeline requires a list arg + RedisPipelineRpushOperation(key=redis_key, values=[safe_dumps(transactions)]) + for transactions, redis_key in _configs + if transactions + ] + if len(rpush_list) == 0: + return + + try: + await self.redis_cache.async_rpush_pipeline(rpush_list=rpush_list) + verbose_proxy_logger.info( + "Spend tracking - restored %d uncommitted transaction set(s) to Redis for retry on next tick.", + len(rpush_list), + ) + except RedisError as e: + verbose_proxy_logger.error( + "Spend tracking - failed to restore uncommitted transactions to Redis. " + "These spend updates are lost. Error: %s", + str(e), + ) + @staticmethod def _number_of_transactions_to_store_in_redis( db_spend_update_transactions: DBSpendUpdateTransactions, diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py index 33372e7794a..79909561683 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -270,6 +270,52 @@ async def test_get_all_transactions_from_redis_buffer_pipeline_no_redis(): assert result == (None, None, None, None, None, None) +@pytest.mark.asyncio +async def test_restore_transactions_to_redis_pushes_only_provided( + redis_update_buffer, mock_redis_cache +): + """ + restore_transactions_to_redis re-pushes only the transaction sets it was + given, to their matching buffer keys, so uncommitted spend can be retried. + """ + from litellm.constants import ( + REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, + REDIS_UPDATE_BUFFER_KEY, + ) + + mock_redis_cache.async_rpush_pipeline = AsyncMock(return_value=[1, 1]) + + db_spend = {"key_list_transactions": {"key1": 1.0}} + daily_user = {"user_key1": {"spend": 1.0}} + + await redis_update_buffer.restore_transactions_to_redis( + db_spend_update_transactions=db_spend, + daily_spend_update_transactions=daily_user, + ) + + mock_redis_cache.async_rpush_pipeline.assert_called_once() + rpush_list = mock_redis_cache.async_rpush_pipeline.call_args.kwargs["rpush_list"] + pushed_keys = {op["key"] for op in rpush_list} + assert pushed_keys == { + REDIS_UPDATE_BUFFER_KEY, + REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, + } + # Payloads round-trip through the same JSON encoding used on the store path + payloads = {op["key"]: json.loads(op["values"][0]) for op in rpush_list} + assert payloads[REDIS_UPDATE_BUFFER_KEY] == db_spend + assert payloads[REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY] == daily_user + + +@pytest.mark.asyncio +async def test_restore_transactions_to_redis_noop_when_empty( + redis_update_buffer, mock_redis_cache +): + """Nothing to restore -> no Redis call.""" + mock_redis_cache.async_rpush_pipeline = AsyncMock() + await redis_update_buffer.restore_transactions_to_redis() + mock_redis_cache.async_rpush_pipeline.assert_not_called() + + def test_validate_redis_transaction_buffer_raises_without_redis(): """ When use_redis_transaction_buffer=true but no Redis cache is configured, 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 4c17c5d3482..10544e82453 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 @@ -1532,9 +1532,9 @@ async def test_commit_spend_updates_uses_pipeline(): mock_redis_update_buffer = AsyncMock() mock_redis_update_buffer.store_in_memory_spend_updates_in_redis = AsyncMock() - # Return all-None tuple (no data to commit) + # Return all-None tuple (no data to commit); the pipeline yields 6 slots mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = ( - AsyncMock(return_value=(None, None, None, None, None, None, None)) + AsyncMock(return_value=(None, None, None, None, None, None)) ) db_writer.redis_update_buffer = mock_redis_update_buffer @@ -1565,6 +1565,153 @@ async def test_commit_spend_updates_uses_pipeline(): mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called() +@pytest.mark.asyncio +async def test_commit_with_redis_requeues_all_on_db_failure(): + """ + Regression for #33872: if the DB commit fails after the leader has already + popped transactions from Redis, the popped transactions must be re-queued to + Redis so a later tick can retry them, instead of being silently lost. + """ + db_writer = DBSpendUpdateWriter() + + db_spend = { + "user_list_transactions": {"user1": 1.5}, + "end_user_list_transactions": {}, + "key_list_transactions": {"key1": 1.5}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + daily_user = {"user_key1": {"spend": 1.5, "api_requests": 1}} + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock( + return_value=(db_spend, daily_user, None, None, None, None) + ) + mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock() + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + # Every DB write raises -> simulates a full database outage + db_writer._commit_spend_updates_to_db = AsyncMock(side_effect=Exception("db down")) + + with patch.object( + DBSpendUpdateWriter, + "update_daily_user_spend", + new=AsyncMock(side_effect=Exception("db down")), + ): + await db_writer._commit_spend_updates_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=0, + proxy_logging_obj=MagicMock(), + ) + + # Both failed categories must be re-queued to Redis, nothing lost + mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once() + _, kwargs = mock_redis_update_buffer.restore_transactions_to_redis.call_args + assert kwargs["db_spend_update_transactions"] == db_spend + assert kwargs["daily_spend_update_transactions"] == daily_user + # The lock must still be released + mock_pod_lock_manager.release_lock.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_commit_with_redis_only_requeues_failed_category(): + """ + A partial DB failure must not re-queue categories that already committed, + otherwise their spend would be double-counted on the next tick. + """ + db_writer = DBSpendUpdateWriter() + + db_spend = { + "user_list_transactions": {"user1": 1.5}, + "end_user_list_transactions": {}, + "key_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + daily_user = {"user_key1": {"spend": 1.5, "api_requests": 1}} + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock( + return_value=(db_spend, daily_user, None, None, None, None) + ) + mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock() + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + # db_spend commits fine; only the daily user commit fails + db_writer._commit_spend_updates_to_db = AsyncMock() + + with patch.object( + DBSpendUpdateWriter, + "update_daily_user_spend", + new=AsyncMock(side_effect=Exception("db down")), + ): + await db_writer._commit_spend_updates_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=0, + proxy_logging_obj=MagicMock(), + ) + + mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once() + _, kwargs = mock_redis_update_buffer.restore_transactions_to_redis.call_args + # Only the failed daily category is requeued; the committed db_spend is not + assert kwargs == {"daily_spend_update_transactions": daily_user} + + +@pytest.mark.asyncio +async def test_commit_with_redis_no_requeue_on_success(): + """When all commits succeed, nothing should be re-queued to Redis.""" + db_writer = DBSpendUpdateWriter() + + db_spend = { + "user_list_transactions": {"user1": 1.5}, + "end_user_list_transactions": {}, + "key_list_transactions": {}, + "team_list_transactions": {}, + "team_member_list_transactions": {}, + "org_list_transactions": {}, + "tag_list_transactions": {}, + "agent_list_transactions": {}, + } + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock( + return_value=(db_spend, None, None, None, None, None) + ) + mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock() + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + db_writer._commit_spend_updates_to_db = AsyncMock() + + await db_writer._commit_spend_updates_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=0, + proxy_logging_obj=MagicMock(), + ) + + mock_redis_update_buffer.restore_transactions_to_redis.assert_not_awaited() + + @pytest.mark.parametrize( "bucket_name,input_dict,table_attr,method_name,where_key,expected_order", [ From 118b47a8a39a79922f09c851364e43e3c73d0ce1 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 18 Jul 2026 23:45:23 +0000 Subject: [PATCH 02/27] fix(proxy): keep tag drain inside try and cover requeue paths with tests Move the destructive daily-tag Redis drain back inside the try so a Redis read failure still releases the pod lock via the finally block, and use a covariant Mapping for the restore signature. Add regression tests for the daily-tag requeue-on-failure/no-requeue-on-success paths and the RedisError swallow branch in restore_transactions_to_redis. --- litellm/proxy/db/db_spend_update_writer.py | 13 ++-- .../redis_update_buffer.py | 13 ++-- .../test_redis_update_buffer.py | 18 +++++ .../proxy/db/test_db_spend_update_writer.py | 72 +++++++++++++++++++ 4 files changed, 102 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index cc266019ff7..f13bf2e2105 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -797,11 +797,7 @@ class DBSpendUpdateWriter: ): verbose_proxy_logger.debug("acquired lock for spend updates") - # Track everything popped from Redis. Each category is removed once it - # has been committed to the DB, so whatever is left after a failure can - # be re-queued for the next tick instead of being lost. Committed - # categories are never re-queued, so their spend is not double-counted. - uncommitted: dict[str, Any] = {} # mutable-ok: drives which popped categories still need re-queuing + uncommitted: dict[str, Any] = {} # mutable-ok: tracks popped categories still needing commit try: ( @@ -1048,11 +1044,12 @@ class DBSpendUpdateWriter: cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ): verbose_proxy_logger.debug("acquired lock for daily tag spend updates") - daily_tag_spend_update_transactions = ( - await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() - ) + daily_tag_spend_update_transactions: dict[str, DailyTagSpendTransaction] | None = None committed = False try: + daily_tag_spend_update_transactions = ( + await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() + ) if daily_tag_spend_update_transactions: await DBSpendUpdateWriter.update_daily_tag_spend( n_retry_times=n_retry_times, 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 b30fadd86ab..660fd514d99 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -6,6 +6,7 @@ This is to prevent deadlocks and improve reliability import asyncio import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast from redis.exceptions import RedisError @@ -379,12 +380,12 @@ class RedisUpdateBuffer: async def restore_transactions_to_redis( self, db_spend_update_transactions: DBSpendUpdateTransactions | None = None, - daily_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, - daily_team_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, - daily_org_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, - daily_end_user_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, - daily_agent_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, - daily_tag_spend_update_transactions: dict[str, BaseDailySpendTransaction] | None = None, + daily_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, + daily_team_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, + daily_org_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, + daily_end_user_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, + daily_agent_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, + daily_tag_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None, ) -> None: """ Re-push transactions that were popped from Redis but not committed to the DB. diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py index 79909561683..3325893c5f6 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -316,6 +316,24 @@ async def test_restore_transactions_to_redis_noop_when_empty( mock_redis_cache.async_rpush_pipeline.assert_not_called() +@pytest.mark.asyncio +async def test_restore_transactions_to_redis_swallows_redis_error( + redis_update_buffer, mock_redis_cache +): + """A Redis failure during restore must not propagate to the caller's finally block.""" + from redis.exceptions import RedisError + + mock_redis_cache.async_rpush_pipeline = AsyncMock( + side_effect=RedisError("redis down") + ) + + await redis_update_buffer.restore_transactions_to_redis( + db_spend_update_transactions={"key_list_transactions": {"key1": 1.0}}, + ) + + mock_redis_cache.async_rpush_pipeline.assert_called_once() + + def test_validate_redis_transaction_buffer_raises_without_redis(): """ When use_redis_transaction_buffer=true but no Redis cache is configured, 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 10544e82453..06d06b50234 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 @@ -1712,6 +1712,78 @@ async def test_commit_with_redis_no_requeue_on_success(): mock_redis_update_buffer.restore_transactions_to_redis.assert_not_awaited() +@pytest.mark.asyncio +async def test_commit_daily_tag_spend_requeues_on_db_failure(): + """A failed daily tag commit must re-queue the popped tag transactions and release the lock.""" + db_writer = DBSpendUpdateWriter() + + daily_tag = {"tag_key1": {"spend": 1.5, "api_requests": 1}} + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.store_in_memory_daily_tag_spend_updates_in_redis = AsyncMock() + mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer = AsyncMock( + return_value=daily_tag + ) + mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock() + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + with patch.object( + DBSpendUpdateWriter, + "update_daily_tag_spend", + new=AsyncMock(side_effect=Exception("db down")), + ): + await db_writer._commit_daily_tag_spend_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=0, + proxy_logging_obj=MagicMock(), + ) + + mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once_with( + daily_tag_spend_update_transactions=daily_tag, + ) + mock_pod_lock_manager.release_lock.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_commit_daily_tag_spend_no_requeue_on_success(): + """A successful daily tag commit must not re-queue anything.""" + db_writer = DBSpendUpdateWriter() + + daily_tag = {"tag_key1": {"spend": 1.5, "api_requests": 1}} + + mock_redis_update_buffer = AsyncMock() + mock_redis_update_buffer.store_in_memory_daily_tag_spend_updates_in_redis = AsyncMock() + mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer = AsyncMock( + return_value=daily_tag + ) + mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock() + db_writer.redis_update_buffer = mock_redis_update_buffer + + mock_pod_lock_manager = AsyncMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + db_writer.pod_lock_manager = mock_pod_lock_manager + + with patch.object( + DBSpendUpdateWriter, + "update_daily_tag_spend", + new=AsyncMock(), + ): + await db_writer._commit_daily_tag_spend_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=0, + proxy_logging_obj=MagicMock(), + ) + + mock_redis_update_buffer.restore_transactions_to_redis.assert_not_awaited() + mock_pod_lock_manager.release_lock.assert_awaited_once() + + @pytest.mark.parametrize( "bucket_name,input_dict,table_attr,method_name,where_key,expected_order", [ From dc58c35bba52a259f99b2d867b4936eed043d43d Mon Sep 17 00:00:00 2001 From: shivam Date: Mon, 27 Jul 2026 23:43:21 +0000 Subject: [PATCH 03/27] fix(anthropic cost): apply regional geo uplift to cached tokens Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/anthropic/cost_calculation.py | 32 ++++--- tests/test_litellm/test_cost_calculator.py | 99 ++++++++++++++++++++++ 2 files changed, 117 insertions(+), 14 deletions(-) diff --git a/litellm/llms/anthropic/cost_calculation.py b/litellm/llms/anthropic/cost_calculation.py index 6a4de1c41b4..6d0a7f8000a 100644 --- a/litellm/llms/anthropic/cost_calculation.py +++ b/litellm/llms/anthropic/cost_calculation.py @@ -24,9 +24,10 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage", service_ti """ Return only the cache-related portion of the prompt cost (cache read + cache write). - These costs must NOT be scaled by geo/speed multipliers because the old + These costs must NOT be scaled by the ``fast`` speed multiplier because the old explicit ``fast/`` model entries carried unchanged cache rates while - multiplying only the regular input/output token costs. + multiplying only the regular input/output token costs. Regional pricing, by + contrast, uplifts every token type, so the geo multiplier does scale them. """ if usage.prompt_tokens_details is None: return 0.0 @@ -81,20 +82,23 @@ def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) model_info: Final = litellm.get_model_info(model=model, custom_llm_provider="anthropic") provider_specific_entry: Final[dict] = model_info.get("provider_specific_entry") or {} - multiplier = 1.0 - if ( - hasattr(usage, "inference_geo") - and usage.inference_geo - and usage.inference_geo.lower() not in ["global", "not_available"] - ): - multiplier *= provider_specific_entry.get(usage.inference_geo.lower(), 1.0) - if hasattr(usage, "speed") and usage.speed == "fast": - multiplier *= provider_specific_entry.get("fast", 1.0) + geo_multiplier: Final = ( + provider_specific_entry.get(usage.inference_geo.lower(), 1.0) + if getattr(usage, "inference_geo", None) and usage.inference_geo.lower() not in ("global", "not_available") + else 1.0 + ) + speed_multiplier: Final = ( + provider_specific_entry.get("fast", 1.0) if getattr(usage, "speed", None) == "fast" else 1.0 + ) - if multiplier != 1.0: + if speed_multiplier != 1.0: cache_cost: Final = _compute_cache_only_cost(model_info=model_info, usage=usage, service_tier=service_tier) - prompt_cost = (prompt_cost - cache_cost) * multiplier + cache_cost - completion_cost *= multiplier + prompt_cost = (prompt_cost - cache_cost) * speed_multiplier + cache_cost + completion_cost *= speed_multiplier + + if geo_multiplier != 1.0: + prompt_cost *= geo_multiplier + completion_cost *= geo_multiplier except Exception: pass diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 3f024e2fd03..16f69773151 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -2726,6 +2726,105 @@ def test_anthropic_cost_per_token_prices_cache_at_served_tier_with_multiplier(): assert completion_cost == pytest.approx(expected_completion) +def _register_anthropic_geo_cache_model(model: str) -> None: + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": 5e-6, + "output_cost_per_token": 25e-6, + "cache_creation_input_token_cost": 6.25e-6, + "cache_read_input_token_cost": 0.5e-6, + "litellm_provider": "anthropic", + "max_tokens": 8192, + "provider_specific_entry": {"us": 1.1, "fast": 2.0}, + } + } + ) + + +def test_anthropic_geo_multiplier_applies_to_cache_tokens(): + """ + Regression: the regional (geo) uplift must scale cache read and cache write + cost too, not just non-cache input and output. + + Anthropic's regional surcharge applies to every token type, so a cache-heavy + row (nearly all cache-creation tokens) must still come in 10% above the + global-priced row. Before the fix the uplift was applied only to the + non-cache portion, so cache-heavy spend was under-reported by ~10%. + """ + from litellm.llms.anthropic.cost_calculation import ( + cost_per_token as anthropic_cost_per_token, + ) + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "claude-test-geo-cache-model" + _register_anthropic_geo_cache_model(model) + + def make_usage() -> "Usage": + return Usage( + prompt_tokens=1_000_000, + completion_tokens=500, + total_tokens=1_000_500, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=200_000, + cache_creation_tokens=799_800, + ), + ) + + base_usage = make_usage() + base_prompt_cost, base_completion_cost = anthropic_cost_per_token(model=model, usage=base_usage) + + geo_usage = make_usage() + geo_usage.inference_geo = "us" + geo_prompt_cost, geo_completion_cost = anthropic_cost_per_token(model=model, usage=geo_usage) + + expected_base_prompt = 200 * 5e-6 + 200_000 * 0.5e-6 + 799_800 * 6.25e-6 + assert base_prompt_cost == pytest.approx(expected_base_prompt) + assert geo_prompt_cost == pytest.approx(expected_base_prompt * 1.1) + assert geo_completion_cost == pytest.approx(base_completion_cost * 1.1) + + +def test_anthropic_geo_and_fast_multipliers_compose(): + """ + The ``fast`` speed multiplier stays cache-exclusive (the old explicit + ``fast/`` entries kept base cache rates) while the geo multiplier scales the + whole cost, so a fast + regional row prices as + ``((non_cache * fast) + cache) * geo``. + """ + from litellm.llms.anthropic.cost_calculation import ( + cost_per_token as anthropic_cost_per_token, + ) + from litellm.types.utils import PromptTokensDetailsWrapper, Usage + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "claude-test-geo-fast-cache-model" + _register_anthropic_geo_cache_model(model) + + usage = Usage( + prompt_tokens=10_000, + completion_tokens=500, + total_tokens=10_500, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=2_000, + cache_creation_tokens=6_000, + ), + ) + usage.inference_geo = "us" + usage.speed = "fast" + + prompt_cost, completion_cost = anthropic_cost_per_token(model=model, usage=usage) + + cache_cost = 2_000 * 0.5e-6 + 6_000 * 6.25e-6 + non_cache_cost = 2_000 * 5e-6 + assert prompt_cost == pytest.approx((non_cache_cost * 2.0 + cache_cost) * 1.1) + assert completion_cost == pytest.approx(500 * 25e-6 * 2.0 * 1.1) + + def test_gemini_cache_tokens_details_no_negative_values(): """ Test for Issue #18750: Negative text_tokens with Gemini caching From 2351aaba74e5c25328fe7a709bf07f052e102a9d Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 28 Jul 2026 00:07:25 +0000 Subject: [PATCH 04/27] test(anthropic cost): scope local cost-map env flag with monkeypatch Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/test_cost_calculator.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 16f69773151..26ba485d796 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -2742,7 +2742,7 @@ def _register_anthropic_geo_cache_model(model: str) -> None: ) -def test_anthropic_geo_multiplier_applies_to_cache_tokens(): +def test_anthropic_geo_multiplier_applies_to_cache_tokens(monkeypatch): """ Regression: the regional (geo) uplift must scale cache read and cache write cost too, not just non-cache input and output. @@ -2757,7 +2757,7 @@ def test_anthropic_geo_multiplier_applies_to_cache_tokens(): ) from litellm.types.utils import PromptTokensDetailsWrapper, Usage - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") litellm.model_cost = litellm.get_model_cost_map(url="") model = "claude-test-geo-cache-model" @@ -2787,7 +2787,7 @@ def test_anthropic_geo_multiplier_applies_to_cache_tokens(): assert geo_completion_cost == pytest.approx(base_completion_cost * 1.1) -def test_anthropic_geo_and_fast_multipliers_compose(): +def test_anthropic_geo_and_fast_multipliers_compose(monkeypatch): """ The ``fast`` speed multiplier stays cache-exclusive (the old explicit ``fast/`` entries kept base cache rates) while the geo multiplier scales the @@ -2799,7 +2799,7 @@ def test_anthropic_geo_and_fast_multipliers_compose(): ) from litellm.types.utils import PromptTokensDetailsWrapper, Usage - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") litellm.model_cost = litellm.get_model_cost_map(url="") model = "claude-test-geo-fast-cache-model" From 66d9752db54da52868d0742a41d86541f44667c6 Mon Sep 17 00:00:00 2001 From: shivam Date: Tue, 28 Jul 2026 00:07:06 +0000 Subject: [PATCH 05/27] fix(anthropic): aggregate 5m/1h cache-write split across iterations path The iterations branch in AnthropicConfig.calculate_usage summed cache_creation_input_tokens but never aggregated the per-iteration cache_creation 5m/1h breakdown, leaving cache_creation_token_details as None. As a result all cache-creation tokens fell back to the flat 5m write rate, underbilling 1h cache writes by up to 2x. Aggregate the ephemeral_5m/ephemeral_1h split across iterations so 1h writes are priced at the 1h rate. Fixes LIT-4868 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/anthropic/chat/transformation.py | 18 ++++++- .../test_anthropic_chat_transformation.py | 52 +++++++++++++++++++ 2 files changed, 69 insertions(+), 1 deletion(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 1161c92232a..d8f2f426d8a 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1,6 +1,7 @@ import json import re import time +from collections.abc import Iterable, Mapping from typing import TYPE_CHECKING, Any, Final, NoReturn, cast import httpx @@ -2117,6 +2118,18 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return False return any(key in usage_object for key in ("cache_read_input_tokens", "cache_creation_input_tokens")) + @staticmethod + def _aggregate_cache_creation_token_details( + cache_creation_objects: Iterable[Mapping[str, Any] | None], + ) -> CacheCreationTokenDetails | None: + breakdowns: Final = tuple(c for c in cache_creation_objects if isinstance(c, Mapping)) + if not breakdowns: + return None + return CacheCreationTokenDetails( + ephemeral_5m_input_tokens=sum(int(c.get("ephemeral_5m_input_tokens") or 0) for c in breakdowns), + ephemeral_1h_input_tokens=sum(int(c.get("ephemeral_1h_input_tokens") or 0) for c in breakdowns), + ) + def calculate_usage( self, usage_object: dict, @@ -2150,6 +2163,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): cache_creation_input_tokens = sum(it.get("cache_creation_input_tokens", 0) or 0 for it in iterations) cache_read_input_tokens = sum(it.get("cache_read_input_tokens", 0) or 0 for it in iterations) prompt_tokens += cache_creation_input_tokens + cache_read_input_tokens + cache_creation_token_details = self._aggregate_cache_creation_token_details( + it.get("cache_creation") for it in iterations + ) if not iterations: if "cache_creation_input_tokens" in _usage and _usage["cache_creation_input_tokens"] is not None: @@ -2182,7 +2198,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if tool_search_count > 0: tool_search_requests = tool_search_count - if "cache_creation" in _usage and _usage["cache_creation"] is not None: + if cache_creation_token_details is None and "cache_creation" in _usage and _usage["cache_creation"] is not None: cache_creation_token_details = CacheCreationTokenDetails( ephemeral_5m_input_tokens=_usage["cache_creation"].get("ephemeral_5m_input_tokens"), ephemeral_1h_input_tokens=_usage["cache_creation"].get("ephemeral_1h_input_tokens"), 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 231d3b48754..79255d4f923 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 @@ -105,6 +105,58 @@ def test_calculate_usage(): assert usage._cache_read_input_tokens == 0 +def test_calculate_usage_aggregates_cache_creation_split_across_iterations(): + """ + In the iterations path each iteration can carry the 5m/1h cache_creation + breakdown. calculate_usage must aggregate it into cache_creation_token_details + so 1h writes are priced at the 1h rate instead of silently falling back to 5m. + + Regression for LIT-4868. + """ + from litellm.llms.anthropic.cost_calculation import cost_per_token + + config = AnthropicConfig() + usage_object = { + "input_tokens": 0, + "output_tokens": 5, + "iterations": [ + { + "type": "message", + "input_tokens": 0, + "output_tokens": 3, + "cache_creation_input_tokens": 10000, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 10000}, + }, + { + "type": "message", + "input_tokens": 0, + "output_tokens": 2, + "cache_creation_input_tokens": 10000, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 10000}, + }, + ], + } + + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None) + + details = usage.prompt_tokens_details.cache_creation_token_details + assert details is not None + assert details.ephemeral_5m_input_tokens == 0 + assert details.ephemeral_1h_input_tokens == 20000 + assert usage.prompt_tokens_details.cache_creation_tokens == 20000 + + info = litellm.get_model_info(model="claude-opus-4-8", custom_llm_provider="anthropic") + rate_5m = info["cache_creation_input_token_cost"] + rate_1h = info["cache_creation_input_token_cost_above_1hr"] + assert rate_1h > rate_5m + + prompt_cost, _ = cost_per_token(model="claude-opus-4-8", usage=usage) + assert prompt_cost == pytest.approx(20000 * rate_1h) + assert prompt_cost != pytest.approx(20000 * rate_5m) + + def test_calculate_usage_clamps_text_tokens_when_reasoning_estimate_exceeds_output(): config = AnthropicConfig() From efe5a3140082e117567bf6203380ea8a83525879 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 11 Aug 2026 01:32:41 +0000 Subject: [PATCH 06/27] refactor(anthropic): resolve cache-write split in one immutable step Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/anthropic/chat/transformation.py | 28 ++++++++++++------- 1 file changed, 18 insertions(+), 10 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index d8f2f426d8a..0dc877a700b 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -2130,6 +2130,23 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ephemeral_1h_input_tokens=sum(int(c.get("ephemeral_1h_input_tokens") or 0) for c in breakdowns), ) + @staticmethod + def _resolve_cache_creation_token_details(usage: Mapping[str, Any]) -> CacheCreationTokenDetails | None: + iterations: Final = usage.get("iterations") + if iterations: + aggregated: Final = AnthropicConfig._aggregate_cache_creation_token_details( + it.get("cache_creation") for it in iterations + ) + if aggregated is not None: + return aggregated + cache_creation: Final = usage.get("cache_creation") + if not isinstance(cache_creation, Mapping): + return None + return CacheCreationTokenDetails( + ephemeral_5m_input_tokens=cache_creation.get("ephemeral_5m_input_tokens"), + ephemeral_1h_input_tokens=cache_creation.get("ephemeral_1h_input_tokens"), + ) + def calculate_usage( self, usage_object: dict, @@ -2145,7 +2162,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): _usage: Final = usage_object cache_creation_input_tokens: int = 0 cache_read_input_tokens: int = 0 - cache_creation_token_details: CacheCreationTokenDetails | None = None + cache_creation_token_details: Final = self._resolve_cache_creation_token_details(_usage) web_search_requests: int | None = None tool_search_requests: int | None = None inference_geo: str | None = None @@ -2163,9 +2180,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): cache_creation_input_tokens = sum(it.get("cache_creation_input_tokens", 0) or 0 for it in iterations) cache_read_input_tokens = sum(it.get("cache_read_input_tokens", 0) or 0 for it in iterations) prompt_tokens += cache_creation_input_tokens + cache_read_input_tokens - cache_creation_token_details = self._aggregate_cache_creation_token_details( - it.get("cache_creation") for it in iterations - ) if not iterations: if "cache_creation_input_tokens" in _usage and _usage["cache_creation_input_tokens"] is not None: @@ -2198,12 +2212,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if tool_search_count > 0: tool_search_requests = tool_search_count - if cache_creation_token_details is None and "cache_creation" in _usage and _usage["cache_creation"] is not None: - cache_creation_token_details = CacheCreationTokenDetails( - ephemeral_5m_input_tokens=_usage["cache_creation"].get("ephemeral_5m_input_tokens"), - ephemeral_1h_input_tokens=_usage["cache_creation"].get("ephemeral_1h_input_tokens"), - ) - raw_input_tokens: Final = prompt_tokens - cache_read_input_tokens - cache_creation_input_tokens prompt_tokens_details: Final = PromptTokensDetailsWrapper( cached_tokens=cache_read_input_tokens, From d3fae8a260031a8aaa2cefeba0b3fa8aeb460d94 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 13 Aug 2026 03:02:54 +0000 Subject: [PATCH 07/27] refactor(proxy): extract redis tag spend drain and commit into a helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 51 +++++++++++++++------- 1 file changed, 35 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index cb3ed4c9520..14604dbc04e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1107,20 +1107,12 @@ class DBSpendUpdateWriter: cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ): verbose_proxy_logger.debug("acquired lock for daily tag spend updates") - daily_tag_spend_update_transactions: dict[str, DailyTagSpendTransaction] | None = None - committed = False try: - daily_tag_spend_update_transactions: Final = ( - await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() + await self._drain_and_commit_daily_tag_spend_from_redis( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + 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, - ) - committed = True except Exception as e: spend_log_error( "Spend tracking - failed to commit daily tag spend updates from Redis to DB. " @@ -1129,14 +1121,41 @@ class DBSpendUpdateWriter: exc=e, ) finally: - if not committed and daily_tag_spend_update_transactions: - await self.redis_update_buffer.restore_transactions_to_redis( - daily_tag_spend_update_transactions=daily_tag_spend_update_transactions, - ) await self.pod_lock_manager.release_lock( cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME, ) + async def _drain_and_commit_daily_tag_spend_from_redis( + self, + prisma_client: PrismaClient, + n_retry_times: int, + proxy_logging_obj: ProxyLogging, + ) -> None: + """ + Drain the Redis tag spend buffer and commit it, restoring the drained transactions if the commit fails. + + The drain is destructive, so a failed commit must push the transactions back for the next tick + or their spend is lost permanently. + """ + daily_tag_spend_update_transactions: Final = ( + await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() + ) + if not daily_tag_spend_update_transactions: + return + + try: + 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, + ) + except Exception: + await self.redis_update_buffer.restore_transactions_to_redis( + daily_tag_spend_update_transactions=daily_tag_spend_update_transactions, + ) + raise + async def _flush_tool_discovery_queue( self, prisma_client: PrismaClient, From 86b24befc11231ccfada401f31afbf676b06e819 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 13 Aug 2026 03:23:27 +0000 Subject: [PATCH 08/27] fix(proxy): stop discarding failed daily spend transactions before requeue Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/db_spend_update_writer.py | 3 -- .../proxy/db/test_db_spend_update_writer.py | 46 +++++++++++++++++++ 2 files changed, 46 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 14604dbc04e..356e45a9daa 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1655,9 +1655,6 @@ class DBSpendUpdateWriter: ) except Exception as e: - if "transactions_to_process" in locals(): - for key in transactions_to_process: - daily_spend_transactions.pop(key, None) _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) @staticmethod 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 62bfca73cb4..226cc62858a 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 @@ -1425,6 +1425,52 @@ async def test_update_daily_spend_re_raises_exception_after_logging(): ) +@pytest.mark.asyncio +async def test_update_daily_spend_keeps_failed_transactions_for_retry(): + """ + A failed batch must stay in the caller's transaction dict, otherwise the + Redis re-queue in _commit_spend_updates_to_db_with_redis has nothing left to + push back and the spend is lost permanently. + """ + + def raise_outage(): + raise ValueError("simulated database outage") + + prisma_client = _RecordingPrisma(execute_raw=raise_outage) + + daily_spend_transactions = { + "test_key": { + "user_id": "test-user", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + } + expected = dict(daily_spend_transactions) + + 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=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + ) + + assert daily_spend_transactions == expected + + @pytest.mark.asyncio async def test_commit_key_spend_updates_includes_last_active(): """ From 83efa9f630140134eaa0286415be4465378dbff5 Mon Sep 17 00:00:00 2001 From: Noah Nistler <60981020+noahnistler@users.noreply.github.com> Date: Fri, 17 Jul 2026 13:22:06 -0500 Subject: [PATCH 09/27] fix(azure_ai): recognize real Search doc endpoints so teams can read/write via passthrough The Azure AI Search vector store config declared its write endpoint as `PUT /docs` and its read endpoints as only `/docs/search`. The passthrough permission gate (`is_allowed_to_call_vector_store_endpoint`) derives a read/write permission type by matching the request route against those lists, and a route matching neither resolves to `None` and raises a 403 before the caller's `allowed_vector_store_indexes` grant is ever checked. Two real Azure routes fell through that gap for non-admins: document upload/merge/delete is `POST /docs/index` (not `PUT /docs`), and get index details is `GET /indexes/{name}` (no `/docs/search` suffix). So a team with a valid write or read grant still got 403 on upload and on reading index details, while admins slipped through because they skip the gate entirely. Correct the map: read is any GET under `/indexes/` (get details, stats, count, and the GET form of search) plus `POST /docs/search`; write is `POST /docs/index`. Index lifecycle (create/update/delete the index itself) stays proxy-admin only because it is handled first by the separate lifecycle check on POST/PUT/DELETE/PATCH, so this does not let a team create or delete indexes. Add regression tests that exercise the real AzureAIVectorStoreConfig map: a write-granted team may upload, a read-granted team may search and get index details, a team missing the matching grant is still denied, and a team cannot manage index lifecycle even with a write grant. --- .../azure_ai/vector_stores/transformation.py | 4 +- .../test_vector_store_endpoints.py | 92 +++++++++++++++++++ 2 files changed, 94 insertions(+), 2 deletions(-) diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index 5e16d759be1..0dc8bcb13a4 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -38,8 +38,8 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: return { - "read": [("GET", "/docs/search"), ("POST", "/docs/search")], - "write": [("PUT", "/docs")], + "read": [("GET", "/indexes/"), ("POST", "/docs/search")], + "write": [("POST", "/docs/index")], } def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials: diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 02ca64e5fb8..2a97a7df9d0 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -2928,3 +2928,95 @@ class TestUpdateVectorStoreAccessControlAndRedaction: params = response["vector_store"]["litellm_params"] assert params["api_key"] == REDACTED_BY_LITELM_STRING assert params["api_base"] == "https://api.openai.com/v1" + + +class TestAzureAIDocumentWritePassthroughPermission: + """Regression tests for the Azure AI Search passthrough write mapping. + + Azure's batch document write/merge/delete endpoint is + ``POST /indexes/{name}/docs/index``. A non-admin team holding a ``write`` + grant on the index must be allowed to call it, while index lifecycle + (create / update / delete the index itself) stays proxy-admin only. + + These exercise the real ``AzureAIVectorStoreConfig`` endpoint map on + purpose (no mocked provider config), so reverting the map to the old + ``("PUT", "/docs")`` entry makes ``test_team_with_write_grant_can_upload`` + fail. + """ + + INDEX = "my-index" + + def _request(self, method: str, path: str) -> MagicMock: + request = MagicMock(spec=Request) + request.method = method + request.url.path = path + return request + + def _team_member(self, permissions: list) -> MagicMock: + user = MagicMock(spec=UserAPIKeyAuth) + user.user_role = None + user.metadata = {"allowed_vector_store_indexes": [{"index_name": self.INDEX, "index_permissions": permissions}]} + user.team_metadata = None + return user + + def test_team_with_write_grant_can_upload(self): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"), + user_api_key_dict=self._team_member(["read", "write"]), + ) + assert result is True + + def test_team_without_write_grant_cannot_upload(self): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"), + user_api_key_dict=self._team_member(["read"]), + ) + assert exc_info.value.status_code == 403 + + def test_team_with_read_grant_can_search(self): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/search"), + user_api_key_dict=self._team_member(["read"]), + ) + assert result is True + + def test_team_with_read_grant_can_get_index_details(self): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"), + user_api_key_dict=self._team_member(["read"]), + ) + assert result is True + + def test_team_without_read_grant_cannot_get_index_details(self): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"), + user_api_key_dict=self._team_member(["write"]), + ) + assert exc_info.value.status_code == 403 + + @pytest.mark.parametrize( + "method, operation", + [("PUT", "update"), ("DELETE", "delete")], + ) + def test_team_cannot_manage_index_lifecycle_even_with_write_grant(self, method, operation): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request(method, f"/azure_ai/indexes/{self.INDEX}?api-version=2024-07-01"), + user_api_key_dict=self._team_member(["read", "write"]), + ) + assert exc_info.value.status_code == 403 + assert f"Only proxy admins can {operation}" in exc_info.value.detail From 23f50e1f343f576042c74e6a4e62d2959074b834 Mon Sep 17 00:00:00 2001 From: Noah Nistler <60981020+noahnistler@users.noreply.github.com> Date: Fri, 17 Jul 2026 14:08:12 -0500 Subject: [PATCH 10/27] fix(vector_stores): classify POST /indexes create as admin-only lifecycle with query string The service-level index-create guard checked normalized.endswith("/indexes") without stripping the query string, so Azure's real create request POST /indexes?api-version=... was never classified as a lifecycle request and fell through to the generic permission check instead of the explicit admin-only guard. Strip the query string before the suffix check, mirroring how the PUT/DELETE index paths already tolerate a trailing ?. Add the POST create path to the lifecycle regression parametrize so a non-admin team with a write grant is denied with the clear admin-only message. --- litellm/proxy/vector_store_endpoints/utils.py | 2 +- .../test_vector_store_endpoints.py | 12 ++++++++---- 2 files changed, 9 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 94ba7c06cad..afde5c787f1 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -86,7 +86,7 @@ def _is_vector_store_index_lifecycle_request( return True # POST /indexes (create index at service level; no index name in path). - normalized: Final = request_path.rstrip("/") + normalized: Final = request_path.split("?", 1)[0].rstrip("/") if request_method == "POST" and normalized.endswith("/indexes"): return True diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 2a97a7df9d0..ad86ce60eab 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -3007,15 +3007,19 @@ class TestAzureAIDocumentWritePassthroughPermission: assert exc_info.value.status_code == 403 @pytest.mark.parametrize( - "method, operation", - [("PUT", "update"), ("DELETE", "delete")], + "method, operation, path", + [ + ("PUT", "update", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"), + ("DELETE", "delete", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"), + ("POST", "create", "/azure_ai/indexes?api-version=2024-07-01"), + ], ) - def test_team_cannot_manage_index_lifecycle_even_with_write_grant(self, method, operation): + def test_team_cannot_manage_index_lifecycle_even_with_write_grant(self, method, operation, path): with pytest.raises(HTTPException) as exc_info: is_allowed_to_call_vector_store_endpoint( provider=LlmProviders.AZURE_AI, index_name=self.INDEX, - request=self._request(method, f"/azure_ai/indexes/{self.INDEX}?api-version=2024-07-01"), + request=self._request(method, path), user_api_key_dict=self._team_member(["read", "write"]), ) assert exc_info.value.status_code == 403 From bdc80b11accb5d1c455c7f5eea363fd096cc2489 Mon Sep 17 00:00:00 2001 From: Noah Nistler <60981020+noahnistler@users.noreply.github.com> Date: Fri, 17 Jul 2026 14:26:14 -0500 Subject: [PATCH 11/27] fix(azure_ai): authorize the targeted Search index, not any matching path segment The Azure passthrough scanned every URL segment for one matching a registered index, authorized against that, then forwarded the original path. A caller with a grant on a managed index named e.g. "index" or "docs" could send POST /azure_ai/indexes/{victim}/docs/index: the scan matched the trailing segment and authorized on the caller's own index while Azure applied the batch write to {victim} on the same Search service, enabling cross-index document uploads or deletions. Resolve the index positionally from the /indexes/{name} segment and require that exact name to be the one authorized and credentialed, so the authorized index and the physical target can never diverge. Add a pure helper plus regression tests covering positional extraction and the route-level cross-index attack. --- .../llm_passthrough_endpoints.py | 24 +++- .../test_llm_pass_through_endpoints.py | 132 ++++++++++++++++++ 2 files changed, 153 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index f84cdd0c222..423e9655d1a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -1234,6 +1234,22 @@ async def assemblyai_proxy_route( return received_value +def get_azure_ai_search_index_from_endpoint(endpoint: str) -> str | None: + """Return the index name in the ``/indexes/{name}`` position of an Azure AI + Search passthrough path, or ``None`` when the path targets no index. + + Only the segment immediately after ``indexes`` is the operable target. Any + other segment (for example the trailing ``index`` in ``.../docs/index``) must + never be treated as the index, otherwise a caller authorized on one index + could have Azure apply the operation to a different index on the same service. + """ + segments: Final = endpoint.split("?", 1)[0].strip("/").split("/") + for position, segment in enumerate(segments): + if segment == "indexes" and position + 1 < len(segments): + return segments[position + 1] or None + return None + + @router.api_route( "/azure_ai/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -1263,6 +1279,8 @@ async def azure_proxy_route( "/" ) # azure model is in the url - e.g. https://{endpoint}/openai/deployments/{deployment-id}/completions?api-version=2024-10-21 + search_index_name: Final = get_azure_ai_search_index_from_endpoint(endpoint) + if len(parts) > 1 and llm_router: for part in parts: # check if LLM MODEL @@ -1271,9 +1289,9 @@ async def azure_proxy_route( ) # check if vector store index is_vector_store_index = ( - (litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part)) - if litellm.vector_store_index_registry is not None - else False + part == search_index_name + and litellm.vector_store_index_registry is not None + and litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part) ) if is_router_model: 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 f631215c03d..7ecf2d510f6 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 @@ -19,9 +19,11 @@ import litellm from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( BaseOpenAIPassThroughHandler, RouteChecks, + azure_proxy_route, bedrock_llm_proxy_route, create_pass_through_route, cursor_proxy_route, + get_azure_ai_search_index_from_endpoint, get_vertex_base_url, llm_passthrough_factory_proxy_route, milvus_proxy_route, @@ -3249,3 +3251,133 @@ def test_is_passthrough_request_streaming_tolerates_non_object_bodies(request_bo ) assert is_passthrough_request_streaming(request_body) is expected + + +class TestGetAzureAISearchIndexFromEndpoint: + """The operable index is only the segment right after ``indexes``. + + A doc-write path ends in ``.../docs/index``; the trailing ``index`` must not + be mistaken for the target, otherwise a caller could be authorized on one + index while Azure applies the write to another. + """ + + @pytest.mark.parametrize( + "endpoint, expected", + [ + ("indexes/my-index/docs/index", "my-index"), + ("indexes/my-index/docs/search", "my-index"), + ("indexes/my-index", "my-index"), + ("indexes/my-index?api-version=2024-07-01", "my-index"), + ("/indexes/my-index/docs/index", "my-index"), + ("indexes/victim/docs/index", "victim"), + ("openai/deployments/gpt-4o/chat/completions", None), + ("indexes", None), + ("indexes/", None), + ], + ) + def test_extracts_positional_index_only(self, endpoint, expected): + assert get_azure_ai_search_index_from_endpoint(endpoint) == expected + + +class TestAzureProxyRouteCrossIndexAuthorization: + """Regression tests: the passthrough must authorize the index that the request + actually targets (the ``/indexes/{name}`` segment), never a different segment + that merely happens to match a managed index the caller can access. + """ + + def _request(self, method: str, path: str) -> MagicMock: + request = MagicMock(spec=Request) + request.method = method + request.headers = {"content-type": "application/json"} + request.url = MagicMock() + request.url.path = path + return request + + @pytest.mark.asyncio + async def test_authorizes_the_targeted_index(self): + index_object = MagicMock() + index_object.litellm_params.vector_store_name = "my-store" + vector_store = {"litellm_params": {"api_base": "https://svc.search.windows.net"}} + + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", + return_value=False, + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config" + ) as mock_get_config, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" + ) as mock_is_allowed, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.assert_user_can_access_vector_store", + new=AsyncMock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + new=AsyncMock(return_value=Response()), + ), + patch.object(litellm, "vector_store_index_registry") as mock_index_registry, + patch.object(litellm, "vector_store_registry") as mock_vector_registry, + ): + mock_get_config.return_value.get_auth_credentials.return_value = {"headers": {"api-key": "k"}} + mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: ( + vector_store_index_name == "my-index" + ) + mock_index_registry.get_vector_store_index_by_name.return_value = index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = vector_store + + await azure_proxy_route( + endpoint="indexes/my-index/docs/index", + request=self._request("POST", "/azure_ai/indexes/my-index/docs/index"), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + ) + + mock_is_allowed.assert_called_once() + assert mock_is_allowed.call_args.kwargs["index_name"] == "my-index" + mock_index_registry.get_vector_store_index_by_name.assert_called_once_with( + vector_store_index_name="my-index" + ) + + @pytest.mark.asyncio + async def test_trailing_index_segment_does_not_authorize_a_different_index(self): + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", + return_value=False, + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" + ) as mock_is_allowed, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + return_value="https://azure-openai.example.com", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="azure-key", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + new=AsyncMock(return_value=Response()), + ) as mock_handler, + patch.object(litellm, "vector_store_index_registry") as mock_index_registry, + ): + mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: ( + vector_store_index_name == "index" + ) + + await azure_proxy_route( + endpoint="indexes/victim/docs/index", + request=self._request("POST", "/azure_ai/indexes/victim/docs/index"), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + ) + + mock_is_allowed.assert_not_called() + mock_handler.assert_awaited_once() + assert mock_handler.await_args.kwargs["custom_llm_provider"] == litellm.LlmProviders.AZURE From c1125f0abb68c7bccef7210fb9550933a5e3a39f Mon Sep 17 00:00:00 2001 From: Noah Nistler <60981020+noahnistler@users.noreply.github.com> Date: Tue, 28 Jul 2026 10:27:46 -0500 Subject: [PATCH 12/27] fix(azure_ai): classify Search suggest, autocomplete, and analyze as reads The endpoint map covered document reads through the ("GET", "/indexes/") entry plus POST /docs/search, which left Azure's remaining POST query endpoints unclassified. POST /docs/suggest, POST /docs/autocomplete, and POST /analyze matched neither list, so the permission gate resolved permission_type to None and raised 403 before the caller's allowed_vector_store_indexes grant was consulted; a non-admin team with a read grant on the index still could not call them. Add the three as reads. They are query endpoints that never mutate the index, so a read grant is the right gate, and each needs its own literal entry because the write entry also matches on POST. Keep every pattern literal rather than a {placeholder} template: the matcher falls back to the substring before a {, which for these routes is always /indexes/, and reads are matched before writes, so a templated read would shadow the /docs/index write and let a read-only team upload. Extend the regression tests to the full non-lifecycle read surface (stats, GET-form search, $count, point lookup, and both forms of suggest and autocomplete, plus analyze), asserting a read grant reaches all of them and a write-only grant reaches none. --- .../azure_ai/vector_stores/transformation.py | 22 +++++++++++- .../test_vector_store_endpoints.py | 36 +++++++++++++++++++ 2 files changed, 57 insertions(+), 1 deletion(-) diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index 0dc8bcb13a4..f58d2f54d2e 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -37,8 +37,28 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): super().__init__() def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints: + """ + Every ``GET`` under ``/indexes/`` is a read: get details, stats, and the + document reads (GET-form search, ``$count``, point lookup, and the + GET forms of suggest and autocomplete). + + ``POST`` splits by endpoint. Search, suggest, autocomplete, and analyze + are query endpoints, so they read; ``/docs/index`` is the batch endpoint + carrying upload, merge, mergeOrUpload, and delete actions, so it writes. + + Patterns stay literal rather than ``{placeholder}`` templates because the + matcher falls back to the substring before a ``{``, which here is always + ``/indexes/`` -- broad enough that a templated read, matched first, would + shadow the ``/docs/index`` write. + """ return { - "read": [("GET", "/indexes/"), ("POST", "/docs/search")], + "read": [ + ("GET", "/indexes/"), + ("POST", "/docs/search"), + ("POST", "/docs/suggest"), + ("POST", "/docs/autocomplete"), + ("POST", "/analyze"), + ], "write": [("POST", "/docs/index")], } diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index ad86ce60eab..fac15c302f4 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -2946,6 +2946,21 @@ class TestAzureAIDocumentWritePassthroughPermission: INDEX = "my-index" + # Every non-lifecycle read Azure exposes for an index. The GET forms are all + # covered by the ("GET", "/indexes/") entry; the POST query endpoints each + # need their own, since the write entry also matches on POST. + READ_ROUTES = [ + ("GET", f"/azure_ai/indexes/{INDEX}/stats"), + ("GET", f"/azure_ai/indexes/{INDEX}/docs"), + ("GET", f"/azure_ai/indexes/{INDEX}/docs/$count"), + ("GET", f"/azure_ai/indexes/{INDEX}/docs/seed-doc-1"), + ("GET", f"/azure_ai/indexes/{INDEX}/docs/suggest"), + ("GET", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"), + ("POST", f"/azure_ai/indexes/{INDEX}/docs/suggest"), + ("POST", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"), + ("POST", f"/azure_ai/indexes/{INDEX}/analyze"), + ] + def _request(self, method: str, path: str) -> MagicMock: request = MagicMock(spec=Request) request.method = method @@ -3006,6 +3021,27 @@ class TestAzureAIDocumentWritePassthroughPermission: ) assert exc_info.value.status_code == 403 + @pytest.mark.parametrize("method, path", READ_ROUTES) + def test_team_with_read_grant_can_call_every_read_route(self, method, path): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request(method, path), + user_api_key_dict=self._team_member(["read"]), + ) + assert result is True + + @pytest.mark.parametrize("method, path", READ_ROUTES) + def test_team_without_read_grant_cannot_call_read_routes(self, method, path): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=self.INDEX, + request=self._request(method, path), + user_api_key_dict=self._team_member(["write"]), + ) + assert exc_info.value.status_code == 403 + @pytest.mark.parametrize( "method, operation, path", [ From f8fccec1080f378dacf80b2b8be36ab51661dca1 Mon Sep 17 00:00:00 2001 From: Noah Nistler <60981020+noahnistler@users.noreply.github.com> Date: Tue, 28 Jul 2026 10:55:51 -0500 Subject: [PATCH 13/27] fix(azure_ai): enforce admin-only index create on the passthrough route POST /azure_ai/indexes carries no index name, so get_azure_ai_search_index_from_endpoint returns None, is_vector_store_index never matches any segment, and the request falls through to the generic Azure passthrough on the proxy's own AZURE_API_BASE and AZURE_API_KEY without ever reaching is_allowed_to_call_vector_store_endpoint. A non-admin could therefore create a Search index whenever AZURE_API_BASE points at the Search service. The earlier lifecycle commit made this look covered. Its test asserts that POST /indexes?api-version=... is refused with "Only proxy admins can create", but it calls the permission gate directly, and that gate is exactly what the route skips for a path with no index name, so the guard was verified in isolation while the route stayed open. Gate the service-level create on the route itself, before the segment loop, with assert_proxy_admin_for_vector_store_index_management. Scope it to POST on a path whose last segment is indexes, mirroring the endswith("/indexes") branch the lifecycle helper already uses, so the managed-index paths and ordinary Azure OpenAI passthrough traffic are untouched. Add route-level tests: a non-admin is refused with the admin-only message and never reaches the passthrough handler, an admin still creates, and the new predicate is parametrized over the service-level, per-index, and non-Search paths. --- .../llm_passthrough_endpoints.py | 19 ++++ .../test_llm_pass_through_endpoints.py | 93 ++++++++++++++++++- 2 files changed, 110 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 423e9655d1a..8c76b9d4e1b 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -44,6 +44,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( ) from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( + assert_proxy_admin_for_vector_store_index_management, assert_user_can_access_vector_store, get_litellm_managed_vector_store, is_allowed_to_call_vector_store_endpoint, @@ -1250,6 +1251,21 @@ def get_azure_ai_search_index_from_endpoint(endpoint: str) -> str | None: return None +def is_azure_ai_search_service_level_index_create(method: str, endpoint: str) -> bool: + """Return True for ``POST /indexes``, Azure AI Search's service-level index create. + + No index name appears in that path, so ``get_azure_ai_search_index_from_endpoint`` + yields None and the managed-index branch can never claim the request. Without an + explicit guard it reaches the generic Azure passthrough on the proxy's own + credential, so a non-admin could create an index whenever ``AZURE_API_BASE`` + points at the Search service. + """ + if method != "POST": + return False + path: Final = endpoint.split("?", 1)[0].strip("/") + return path == "indexes" or path.endswith("/indexes") + + @router.api_route( "/azure_ai/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -1275,6 +1291,9 @@ async def azure_proxy_route( """ from litellm.proxy.proxy_server import llm_router + if is_azure_ai_search_service_level_index_create(method=request.method, endpoint=endpoint): + assert_proxy_admin_for_vector_store_index_management(user_api_key_dict, operation="create") + parts: Final = endpoint.split( "/" ) # azure model is in the url - e.g. https://{endpoint}/openai/deployments/{deployment-id}/completions?api-version=2024-10-21 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 7ecf2d510f6..8080ca71773 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 @@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest -from fastapi import Request, Response +from fastapi import HTTPException, Request, Response from fastapi.testclient import TestClient sys.path.insert( @@ -25,6 +25,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( cursor_proxy_route, get_azure_ai_search_index_from_endpoint, get_vertex_base_url, + is_azure_ai_search_service_level_index_create, llm_passthrough_factory_proxy_route, milvus_proxy_route, mistral_proxy_route, @@ -33,7 +34,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( vertex_proxy_route, vllm_proxy_route, ) -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -3381,3 +3382,91 @@ class TestAzureProxyRouteCrossIndexAuthorization: mock_is_allowed.assert_not_called() mock_handler.assert_awaited_once() assert mock_handler.await_args.kwargs["custom_llm_provider"] == litellm.LlmProviders.AZURE + + +class TestAzureProxyRouteServiceLevelIndexCreate: + """``POST /indexes`` carries no index name, so the managed-index branch cannot + claim it and it would otherwise reach the generic Azure passthrough on the + proxy's own credential. The admin-only index management guard has to be + enforced on the route itself, not just on the permission gate the route skips. + """ + + def _request(self, method: str, path: str) -> MagicMock: + request = MagicMock(spec=Request) + request.method = method + request.headers = {"content-type": "application/json"} + request.url = MagicMock() + request.url.path = path + return request + + @pytest.mark.parametrize( + "method, endpoint, expected", + [ + ("POST", "indexes", True), + ("POST", "indexes?api-version=2024-07-01", True), + ("POST", "/indexes/", True), + ("POST", "indexes/my-index", False), + ("POST", "indexes/my-index/docs/index", False), + ("GET", "indexes", False), + ("POST", "openai/deployments/gpt-4o/chat/completions", False), + ], + ) + def test_recognizes_service_level_create(self, method, endpoint, expected): + assert is_azure_ai_search_service_level_index_create(method=method, endpoint=endpoint) is expected + + @pytest.mark.asyncio + async def test_non_admin_cannot_create_an_index(self): + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + return_value="https://svc.search.windows.net", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + new=AsyncMock(return_value=Response()), + ) as mock_handler, + ): + with pytest.raises(HTTPException) as exc_info: + await azure_proxy_route( + endpoint="indexes?api-version=2024-07-01", + request=self._request("POST", "/azure_ai/indexes"), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth( + token="sk-team-token", + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + + assert exc_info.value.status_code == 403 + assert "Only proxy admins can create" in exc_info.value.detail + mock_handler.assert_not_awaited() + + @pytest.mark.asyncio + async def test_admin_can_still_create_an_index(self): + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str", + return_value="https://svc.search.windows.net", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", + return_value="azure-key", + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler", + new=AsyncMock(return_value=Response()), + ) as mock_handler, + ): + await azure_proxy_route( + endpoint="indexes?api-version=2024-07-01", + request=self._request("POST", "/azure_ai/indexes"), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth( + token="sk-admin-token", + user_role=LitellmUserRoles.PROXY_ADMIN, + ), + ) + + mock_handler.assert_awaited_once() From 08966c842b1b5a903a11a43a973ee760ce89c7f1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 14:12:09 -0700 Subject: [PATCH 14/27] test(vector_stores): drop redundant route-map comment --- .../vector_store_endpoints/test_vector_store_endpoints.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index fac15c302f4..8a227028b51 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -2946,9 +2946,6 @@ class TestAzureAIDocumentWritePassthroughPermission: INDEX = "my-index" - # Every non-lifecycle read Azure exposes for an index. The GET forms are all - # covered by the ("GET", "/indexes/") entry; the POST query endpoints each - # need their own, since the write entry also matches on POST. READ_ROUTES = [ ("GET", f"/azure_ai/indexes/{INDEX}/stats"), ("GET", f"/azure_ai/indexes/{INDEX}/docs"), From 5212e8c1f1a1306a27f7e2d3d75a11662725015a Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 14 Aug 2026 22:59:30 +0000 Subject: [PATCH 15/27] refactor(caching): accept read-only sequences for redis rpush pipeline payloads Keeps the spend buffer restore path free of mutable-collection construction so the type discipline gate stays within its LIT002 ceiling. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/redis_cache.py | 4 ++-- .../proxy/db/db_transaction_queue/redis_update_buffer.py | 6 +++--- litellm/types/caching.py | 3 ++- type-discipline-budget.json | 2 +- 4 files changed, 8 insertions(+), 7 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 5fedfc5bcce..a3936fd17e2 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -1572,7 +1572,7 @@ class RedisCache(BaseCache): async def _pipeline_rpush_helper( self, pipe: pipeline, - rpush_list: list[RedisPipelineRpushOperation], + rpush_list: Sequence[RedisPipelineRpushOperation], ) -> list[int]: """Helper function for pipeline rpush operations""" for rpush_op in rpush_list: @@ -1588,7 +1588,7 @@ class RedisCache(BaseCache): @_redis_circuit_breaker_guard async def async_rpush_pipeline( self, - rpush_list: list[RedisPipelineRpushOperation], + rpush_list: Sequence[RedisPipelineRpushOperation], ) -> list[int]: """ Use Redis Pipelines for bulk RPUSH operations 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 573da72b873..853c033c37e 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -407,11 +407,11 @@ class RedisUpdateBuffer: (daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY), ) - rpush_list: Final[list[RedisPipelineRpushOperation]] = [ # mutable-ok: async_rpush_pipeline requires a list arg - RedisPipelineRpushOperation(key=redis_key, values=[safe_dumps(transactions)]) + rpush_list: Final = tuple( + RedisPipelineRpushOperation(key=redis_key, values=(safe_dumps(transactions),)) for transactions, redis_key in restore_configs if transactions - ] + ) if len(rpush_list) == 0: return diff --git a/litellm/types/caching.py b/litellm/types/caching.py index 6616a2e9bac..427904d2fe2 100644 --- a/litellm/types/caching.py +++ b/litellm/types/caching.py @@ -1,3 +1,4 @@ +from collections.abc import Sequence from enum import Enum from typing import Any, Final, Literal, Optional, Union @@ -59,7 +60,7 @@ class RedisPipelineRpushOperation(TypedDict): """ key: str - values: list[Any] + values: Sequence[Any] class RedisPipelineLpopOperation(TypedDict): diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 894d99c92e0..94565199516 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22941 + "limit": 22938 }, "LIT002": { "limit": 27139 From 61334ec94acdd4bf3329ad7a38c23d3d6dc8bcc6 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 14 Aug 2026 16:32:27 -0700 Subject: [PATCH 16/27] fix(ui): match the MCP servers count badge to its sibling permission badges The Object Permissions section rendered the MCP Servers badge with shadcn's default variant (solid bg-primary), so a plain count showed up as a black pill next to the light Vector Stores and Agents counts. Counts now use secondary everywhere, and destructive stays reserved for the blocked state. --- .../permissions/MCPServerPermissions.test.tsx | 41 ++++++++++++++++++- .../permissions/MCPServerPermissions.tsx | 2 +- 2 files changed, 41 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx index 568c8b218bc..1d8a0a6dd84 100644 --- a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx @@ -3,7 +3,7 @@ import { render, screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import MCPServerPermissions from "./MCPServerPermissions"; import * as networking from "../networking"; -import { ALL_PROXY_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import { ALL_PROXY_MCP_SERVERS_SENTINEL, NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; vi.mock("../networking"); @@ -372,4 +372,43 @@ describe("MCPServerPermissions", () => { expect(screen.getByText("All")).toBeInTheDocument(); expect(screen.queryByText(ALL_PROXY_MCP_SERVERS_SENTINEL)).not.toBeInTheDocument(); }); + + it("should use the neutral badge variant unless MCP access is blocked", async () => { + /** + * The header badge sits next to the Vector Stores and Agents badges, which both render + * variant="secondary". "default" renders solid bg-primary (black), so it only belongs on + * the blocked state, which uses "destructive". + */ + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + + const { rerender } = render( + , + ); + expect(screen.getByText("0")).toHaveAttribute("data-variant", "secondary"); + + rerender( + , + ); + await waitFor(() => expect(screen.getByText("All")).toHaveAttribute("data-variant", "secondary")); + + rerender( + , + ); + await waitFor(() => expect(screen.getByText("Blocked")).toHaveAttribute("data-variant", "destructive")); + }); }); diff --git a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx index 02980cd4c3e..b00fd73c320 100644 --- a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx +++ b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.tsx @@ -112,7 +112,7 @@ export function MCPServerPermissions({

MCP Servers

- + {blocksAllMcpServers ? "Blocked" : grantsAllProxyMcpServers ? "All" : totalCount}
From 94e943144ea5f237fc158d85f00c67ef9fe72c08 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 14 Aug 2026 16:40:16 -0700 Subject: [PATCH 17/27] refactor(ui): drop the explanatory comment from the badge variant test --- .../src/components/permissions/MCPServerPermissions.test.tsx | 5 ----- 1 file changed, 5 deletions(-) diff --git a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx index 1d8a0a6dd84..c2df945f367 100644 --- a/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/permissions/MCPServerPermissions.test.tsx @@ -374,11 +374,6 @@ describe("MCPServerPermissions", () => { }); it("should use the neutral badge variant unless MCP access is blocked", async () => { - /** - * The header badge sits next to the Vector Stores and Agents badges, which both render - * variant="secondary". "default" renders solid bg-primary (black), so it only belongs on - * the blocked state, which uses "destructive". - */ vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); const { rerender } = render( From 0ab23f5ce9eb5f5db4240ac7519a54e8cf73ad9e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 16:41:34 -0700 Subject: [PATCH 18/27] fix(anthropic): bill undetailed iteration cache writes at the 5m rate --- litellm/llms/anthropic/chat/transformation.py | 18 ++++--- .../test_anthropic_chat_transformation.py | 50 +++++++++++++++++++ 2 files changed, 60 insertions(+), 8 deletions(-) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 0dc877a700b..31713f8f085 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1,7 +1,7 @@ import json import re import time -from collections.abc import Iterable, Mapping +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, NoReturn, cast import httpx @@ -2120,23 +2120,25 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _aggregate_cache_creation_token_details( - cache_creation_objects: Iterable[Mapping[str, Any] | None], + iterations: Sequence[Mapping[str, Any]], ) -> CacheCreationTokenDetails | None: - breakdowns: Final = tuple(c for c in cache_creation_objects if isinstance(c, Mapping)) + breakdowns: Final = tuple(c for c in (it.get("cache_creation") for it in iterations) if isinstance(c, Mapping)) if not breakdowns: return None + detailed_5m: Final = sum(int(c.get("ephemeral_5m_input_tokens") or 0) for c in breakdowns) + detailed_1h: Final = sum(int(c.get("ephemeral_1h_input_tokens") or 0) for c in breakdowns) + total: Final = sum(int(it.get("cache_creation_input_tokens") or 0) for it in iterations) + undetailed: Final = max(total - detailed_5m - detailed_1h, 0) return CacheCreationTokenDetails( - ephemeral_5m_input_tokens=sum(int(c.get("ephemeral_5m_input_tokens") or 0) for c in breakdowns), - ephemeral_1h_input_tokens=sum(int(c.get("ephemeral_1h_input_tokens") or 0) for c in breakdowns), + ephemeral_5m_input_tokens=detailed_5m + undetailed, + ephemeral_1h_input_tokens=detailed_1h, ) @staticmethod def _resolve_cache_creation_token_details(usage: Mapping[str, Any]) -> CacheCreationTokenDetails | None: iterations: Final = usage.get("iterations") if iterations: - aggregated: Final = AnthropicConfig._aggregate_cache_creation_token_details( - it.get("cache_creation") for it in iterations - ) + aggregated: Final = AnthropicConfig._aggregate_cache_creation_token_details(iterations) if aggregated is not None: return aggregated cache_creation: Final = usage.get("cache_creation") 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 79255d4f923..867b148bfc3 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 @@ -157,6 +157,56 @@ def test_calculate_usage_aggregates_cache_creation_split_across_iterations(): assert prompt_cost != pytest.approx(20000 * rate_5m) +def test_calculate_usage_bills_undetailed_iteration_cache_writes_at_5m_rate(): + """ + When only some iterations carry the cache_creation breakdown, the writes + without a breakdown must still be billed (at the default 5m rate) instead + of silently priced at zero once details exist. + + Regression for the Cursor Bugbot finding on the LIT-4868 fix. + """ + from litellm.llms.anthropic.cost_calculation import cost_per_token + + config = AnthropicConfig() + usage_object = { + "input_tokens": 0, + "output_tokens": 5, + "iterations": [ + { + "type": "message", + "input_tokens": 0, + "output_tokens": 3, + "cache_creation_input_tokens": 10000, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 10000}, + }, + { + "type": "message", + "input_tokens": 0, + "output_tokens": 2, + "cache_creation_input_tokens": 7000, + "cache_read_input_tokens": 0, + }, + ], + } + + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None) + + details = usage.prompt_tokens_details.cache_creation_token_details + assert details is not None + assert details.ephemeral_5m_input_tokens == 7000 + assert details.ephemeral_1h_input_tokens == 10000 + assert usage.prompt_tokens_details.cache_creation_tokens == 17000 + + info = litellm.get_model_info(model="claude-opus-4-8", custom_llm_provider="anthropic") + rate_5m = info["cache_creation_input_token_cost"] + rate_1h = info["cache_creation_input_token_cost_above_1hr"] + + prompt_cost, _ = cost_per_token(model="claude-opus-4-8", usage=usage) + assert prompt_cost == pytest.approx(7000 * rate_5m + 10000 * rate_1h) + assert prompt_cost != pytest.approx(10000 * rate_1h) + + def test_calculate_usage_clamps_text_tokens_when_reasoning_estimate_exceeds_output(): config = AnthropicConfig() From e94a97fcfc5e793fc1de5f4291a782d4ba36dc98 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 16:44:22 -0700 Subject: [PATCH 19/27] fix(cost_calculator): mirror the anthropic geo uplift in the token-type cost breakdown --- .../litellm_core_utils/llm_cost_calc/utils.py | 25 +++++++ litellm/llms/anthropic/cost_calculation.py | 7 +- .../llm_cost_calc/test_llm_cost_calc_utils.py | 69 +++++++++++++++++++ 3 files changed, 96 insertions(+), 5 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index b94851794f0..38bdf89981e 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -694,6 +694,23 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str | return 1.0 +def get_provider_specific_geo_multiplier(model_info: ModelInfo, usage: Usage) -> float: + """ + Resolve the provider-specific regional pricing multiplier for the geo the + request was served from (``usage.inference_geo``), e.g. Anthropic's ``us: 1.1`` + stored under ``provider_specific_entry``. The regional surcharge applies to + every token type, so per-type cost breakdowns must scale by it too. + + Returns 1.0 when the request was served globally or the model carries no + multiplier for the geo. + """ + inference_geo: Final = getattr(usage, "inference_geo", None) + if not isinstance(inference_geo, str) or inference_geo.lower() in ("global", "not_available"): + return 1.0 + provider_specific_entry: Final[dict[str, float]] = model_info.get("provider_specific_entry") or {} + return float(provider_specific_entry.get(inference_geo.lower(), 1.0)) + + def _resolve_reasoning_token_cost( model_info: ModelInfo, service_tier: str | None, @@ -981,6 +998,14 @@ def get_token_type_cost_breakdown( cache_read_cost *= uplift cache_creation_cost *= uplift + # Mirror the provider-specific geo uplift (e.g. Anthropic us: 1.1) the totals + # apply, so cache and reasoning line items stay reconciled with them. + geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage) + if geo_multiplier != 1.0: + reasoning_cost *= geo_multiplier + cache_read_cost *= geo_multiplier + cache_creation_cost *= geo_multiplier + return TokenTypeCostBreakdown( reasoning_cost=reasoning_cost, cache_read_cost=cache_read_cost, diff --git a/litellm/llms/anthropic/cost_calculation.py b/litellm/llms/anthropic/cost_calculation.py index 6d0a7f8000a..e792f69622c 100644 --- a/litellm/llms/anthropic/cost_calculation.py +++ b/litellm/llms/anthropic/cost_calculation.py @@ -13,6 +13,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import ( _parse_prompt_tokens_details, calculate_cache_writing_cost, generic_cost_per_token, + get_provider_specific_geo_multiplier, ) if TYPE_CHECKING: @@ -82,11 +83,7 @@ def cost_per_token(model: str, usage: "Usage", service_tier: str | None = None) model_info: Final = litellm.get_model_info(model=model, custom_llm_provider="anthropic") provider_specific_entry: Final[dict] = model_info.get("provider_specific_entry") or {} - geo_multiplier: Final = ( - provider_specific_entry.get(usage.inference_geo.lower(), 1.0) - if getattr(usage, "inference_geo", None) and usage.inference_geo.lower() not in ("global", "not_available") - else 1.0 - ) + geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage) speed_multiplier: Final = ( provider_specific_entry.get("fast", 1.0) if getattr(usage, "speed", None) == "fast" else 1.0 ) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 3aa41e18f1e..d22a139ba79 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -2558,6 +2558,75 @@ def test_token_type_cost_breakdown_applies_regional_uplift(): assert text_input_cost + eu.cache_read_cost == pytest.approx(prompt_cost) +def test_token_type_cost_breakdown_applies_anthropic_geo_multiplier(monkeypatch): + """ + Anthropic's regional (geo) uplift lives in provider_specific_entry and is + applied to every token type in the totals, so the per-type breakdown must + scale its cache and reasoning line items by it too. Otherwise the logged + cache costs stay at the base rate and the cache uplift is misattributed to + plain input for exactly the cache-heavy regional traffic the uplift targets. + """ + from litellm.llms.anthropic.cost_calculation import ( + cost_per_token as anthropic_cost_per_token, + ) + + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "claude-test-geo-breakdown-model" + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": 5e-6, + "output_cost_per_token": 25e-6, + "cache_creation_input_token_cost": 6.25e-6, + "cache_read_input_token_cost": 0.5e-6, + "litellm_provider": "anthropic", + "max_tokens": 8192, + "provider_specific_entry": {"us": 1.1}, + } + } + ) + + def make_usage() -> Usage: + return Usage( + prompt_tokens=10_000, + completion_tokens=500, + total_tokens=10_500, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=2_000, + cache_creation_tokens=6_000, + ), + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=200, text_tokens=300 + ), + ) + + base_usage = make_usage() + geo_usage = make_usage() + geo_usage.inference_geo = "us" + + base = get_token_type_cost_breakdown( + model=model, custom_llm_provider="anthropic", usage=base_usage + ) + geo = get_token_type_cost_breakdown( + model=model, custom_llm_provider="anthropic", usage=geo_usage + ) + + assert base.cache_read_cost == pytest.approx(2_000 * 0.5e-6) + assert base.cache_creation_cost == pytest.approx(6_000 * 6.25e-6) + assert geo.cache_read_cost == pytest.approx(base.cache_read_cost * 1.1) + assert geo.cache_creation_cost == pytest.approx(base.cache_creation_cost * 1.1) + assert geo.reasoning_cost == pytest.approx(base.reasoning_cost * 1.1) + + # The uplifted breakdown must still reconcile with the uplifted totals. + prompt_cost, completion_cost = anthropic_cost_per_token(model=model, usage=geo_usage) + text_input_cost = 2_000 * 5e-6 * 1.1 + text_output_cost = 300 * 25e-6 * 1.1 + assert text_input_cost + geo.cache_read_cost + geo.cache_creation_cost == pytest.approx(prompt_cost) + assert text_output_cost + geo.reasoning_cost == pytest.approx(completion_cost) + + @pytest.mark.parametrize("details_as_dict", [True, False]) def test_image_response_input_image_tokens_priced_at_image_rate(details_as_dict): """ From 2959465ea087b1aae2e118b7dc37f83d4fdf629c Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 14 Aug 2026 16:51:41 -0700 Subject: [PATCH 20/27] fix(openai,azure): return a length-truncated 200 when the output budget fits no token (#36859) OpenAI and Azure GPT-5.x answer a chat request whose output budget cannot fit a single visible token with a 400, while the same models return a length-truncated 200 one or two tokens higher. Agents that probe a model with a hardcoded max_tokens of 1 read that 400 as "model unavailable". The four chat request helpers now recognise the provider's own sentence and hand back the length-truncated response the provider gives at a slightly larger budget: finish_reason "length", empty content, zero completion tokens. Any other 400 still raises. Streaming is covered by the same seam, and the caller's budget is never raised on their behalf. The provider bills the prompt it processed but sends no usage object with the 400, so the prompt tokens are estimated with the same token_counter every other usage-less path uses. Reporting zero would let a caller send an arbitrarily large prompt with max_tokens 1 and be charged nothing. --- litellm/llms/azure/azure.py | 13 ++ litellm/llms/openai/common_utils.py | 82 ++++++++++ litellm/llms/openai/openai.py | 10 ++ .../llms/openai/test_openai_common_utils.py | 145 ++++++++++++++++++ 4 files changed, 250 insertions(+) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 3438e835faf..c8f94b575ad 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -10,6 +10,7 @@ from openai import ( AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, + BadRequestError, OpenAI, ) @@ -37,6 +38,10 @@ from litellm.utils import ( from ...types.llms.openai import HttpxBinaryResponseContent from ..base import BaseLLM +from ..openai.common_utils import ( + build_output_token_limit_response, + is_output_token_limit_error, +) from .common_utils import ( AzureOpenAIError, BaseAzureLLM, @@ -147,6 +152,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): headers: Final = dict(raw_response.headers) response: Final = raw_response.parse() return headers, response + except BadRequestError as e: + if not is_output_token_limit_error(e): + raise + return build_output_token_limit_response(e=e, data=data, is_async=False) except Exception as e: raise e @@ -175,6 +184,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): time_delta: Final = round(end_time - start_time, 2) e.message += f" - timeout value={timeout}, time taken={time_delta} seconds" raise e + except BadRequestError as e: + if not is_output_token_limit_error(e): + raise + return build_output_token_limit_response(e=e, data=data, is_async=True) except Exception as e: raise e diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 82ebee3962e..1b1ab80e85d 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -7,16 +7,25 @@ import inspect import json import os import ssl +import time +import uuid +from collections.abc import AsyncIterator, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Optional import httpx import openai from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI +from openai.types.chat import ChatCompletion, ChatCompletionChunk, ChatCompletionMessage +from openai.types.chat.chat_completion import Choice +from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice +from openai.types.chat.chat_completion_chunk import ChoiceDelta +from openai.types.completion_usage import CompletionUsage if TYPE_CHECKING: from aiohttp import ClientSession import litellm +from litellm.litellm_core_utils.token_counter import token_counter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( _DEFAULT_TTL_FOR_HTTPX_CLIENTS, @@ -111,6 +120,79 @@ def drop_params_from_unprocessable_entity_error( return new_data +_OUTPUT_TOKEN_LIMIT_ERROR_MARKER: Final[str] = ( + "could not finish the message because max_tokens or model output limit was reached" +) + + +def is_output_token_limit_error(e: openai.BadRequestError) -> bool: + """ + True when OpenAI/Azure rejected a chat request because the output budget could not fit a single visible token. + + GPT-5.x turns that case into a 400 while returning a length-truncated 200 for marginally larger budgets, so the + match has to stay pinned to the full provider sentence to avoid swallowing genuine bad requests. + """ + return _OUTPUT_TOKEN_LIMIT_ERROR_MARKER in e.message.lower() + + +def _output_token_limit_completion(model: str, prompt_tokens: int) -> ChatCompletion: + return ChatCompletion( + id=f"chatcmpl-{uuid.uuid4()}", + choices=( + Choice( + index=0, + finish_reason="length", + message=ChatCompletionMessage(role="assistant", content=""), + ), + ), + created=int(time.time()), + model=model, + object="chat.completion", + usage=CompletionUsage(completion_tokens=0, prompt_tokens=prompt_tokens, total_tokens=prompt_tokens), + ) + + +def _output_token_limit_chunk(model: str) -> ChatCompletionChunk: + return ChatCompletionChunk( + id=f"chatcmpl-{uuid.uuid4()}", + choices=( + ChunkChoice( + index=0, + finish_reason="length", + delta=ChoiceDelta(role="assistant", content=""), + ), + ), + created=int(time.time()), + model=model, + object="chat.completion.chunk", + ) + + +def _iter_once(chunk: ChatCompletionChunk) -> Iterator[ChatCompletionChunk]: + yield chunk + + +async def _aiter_once(chunk: ChatCompletionChunk) -> AsyncIterator[ChatCompletionChunk]: + yield chunk + + +def build_output_token_limit_response( + e: openai.BadRequestError, data: Mapping[str, object], is_async: bool +) -> tuple[httpx.Headers, ChatCompletion | Iterator[ChatCompletionChunk] | AsyncIterator[ChatCompletionChunk]]: + """Synthesize the length-truncated response the provider itself returns for slightly larger output budgets. + + The provider billed the prompt it processed but sends no usage object with the 400, so the prompt is estimated + the way every other usage-less path estimates it: reporting zero would spend input tokens against no budget. + """ + model: Final[str] = str(data.get("model", "")) + messages: Final = data.get("messages") + prompt_tokens: Final = token_counter(model=model, messages=messages) if isinstance(messages, list) else 0 + if not data.get("stream"): + return e.response.headers, _output_token_limit_completion(model, prompt_tokens) + chunk: Final = _output_token_limit_chunk(model) + return e.response.headers, (_aiter_once(chunk) if is_async else _iter_once(chunk)) + + class BaseOpenAILLM: """ Base class for OpenAI LLMs for getting their httpx clients and SSL verification settings diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index e96b61d8204..4fc6655ca54 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -46,7 +46,9 @@ from .chat.o_series_transformation import OpenAIOSeriesConfig from .common_utils import ( BaseOpenAILLM, OpenAIError, + build_output_token_limit_response, drop_params_from_unprocessable_entity_error, + is_output_token_limit_error, ) openaiOSeriesConfig: Final = OpenAIOSeriesConfig() @@ -436,6 +438,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): time_delta: Final = round(end_time - start_time, 2) e.message += f" - timeout value={timeout}, time taken={time_delta} seconds" raise e + except openai.BadRequestError as e: + if not is_output_token_limit_error(e): + raise + return build_output_token_limit_response(e=e, data=data, is_async=True) except Exception as e: raise e @@ -469,6 +475,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): return headers, response except OpenAIError: raise + except openai.BadRequestError as e: + if not is_output_token_limit_error(e): + raise + return build_output_token_limit_response(e=e, data=data, is_async=False) except Exception as e: if raw_response is not None: raise Exception( diff --git a/tests/test_litellm/llms/openai/test_openai_common_utils.py b/tests/test_litellm/llms/openai/test_openai_common_utils.py index a099b5c659f..a28e133700e 100644 --- a/tests/test_litellm/llms/openai/test_openai_common_utils.py +++ b/tests/test_litellm/llms/openai/test_openai_common_utils.py @@ -2,6 +2,8 @@ import os import sys from unittest.mock import MagicMock, call, patch +import httpx +import openai import pytest sys.path.insert( @@ -9,6 +11,7 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm +from litellm.litellm_core_utils.token_counter import token_counter from litellm.llms.openai.common_utils import BaseOpenAILLM # Test parameters for different API functions @@ -247,3 +250,145 @@ def test_a_client_litellm_built_its_own_http_client_for_is_still_closed(monkeypa closer.reap() assert wrapper.is_closed() is True + + +OUTPUT_LIMIT_400_MESSAGE = ( + "Could not finish the message because max_tokens or model output limit was reached. " + "Please try again with higher max_tokens." +) +GENUINE_400_MESSAGE = "Invalid value for 'max_tokens': integer above maximum value. Expected <= 128000, got 999999999." +LONG_PROMPT = "please summarise the following notes for me: " + ("token " * 200) + +CALL_KWARGS_BY_PROVIDER = { + "openai": {"model": "gpt-5.6-sol", "api_key": "sk-not-a-real-key"}, + "azure": { + "model": "azure/gpt-5.6-sol", + "api_key": "not-a-real-key", + "api_base": "https://not-a-real-resource.openai.azure.com", + "api_version": "2024-10-21", + }, +} + + +def _transport(message: str) -> httpx.MockTransport: + def _handler(_request: httpx.Request) -> httpx.Response: + return httpx.Response(400, json={"error": {"message": message, "type": "invalid_request_error"}}) + + return httpx.MockTransport(_handler) + + +def _sync_client_raising(provider: str, message: str): + http_client = httpx.Client(transport=_transport(message)) + if provider == "azure": + return openai.AzureOpenAI( + api_key="not-a-real-key", + azure_endpoint="https://not-a-real-resource.openai.azure.com", + api_version="2024-10-21", + http_client=http_client, + ) + return openai.OpenAI(api_key="sk-not-a-real-key", http_client=http_client) + + +def _async_client_raising(provider: str, message: str): + http_client = httpx.AsyncClient(transport=_transport(message)) + if provider == "azure": + return openai.AsyncAzureOpenAI( + api_key="not-a-real-key", + azure_endpoint="https://not-a-real-resource.openai.azure.com", + api_version="2024-10-21", + http_client=http_client, + ) + return openai.AsyncOpenAI(api_key="sk-not-a-real-key", http_client=http_client) + + +def _completion_kwargs(provider: str, client, **overrides) -> dict: + return { + **CALL_KWARGS_BY_PROVIDER[provider], + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 1, + "client": client, + **overrides, + } + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +def test_sync_output_limit_400_maps_to_length_truncated_response(provider): + response = litellm.completion( + **_completion_kwargs(provider, _sync_client_raising(provider, OUTPUT_LIMIT_400_MESSAGE)) + ) + + assert response.choices[0].finish_reason == "length" + assert response.choices[0].message.content == "" + assert response.usage.completion_tokens == 0 + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +def test_mapped_response_still_bills_the_prompt_the_provider_processed(provider): + messages = [{"role": "user", "content": LONG_PROMPT}] + expected_prompt_tokens = token_counter(model="gpt-5.6-sol", messages=messages) + assert expected_prompt_tokens > 100, "the fixture prompt must be big enough for a zeroed count to stand out" + + response = litellm.completion( + **_completion_kwargs(provider, _sync_client_raising(provider, OUTPUT_LIMIT_400_MESSAGE), messages=messages) + ) + + assert response.usage.prompt_tokens == expected_prompt_tokens + assert response.usage.completion_tokens == 0 + assert litellm.completion_cost(completion_response=response) > 0 + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +@pytest.mark.asyncio +async def test_async_output_limit_400_maps_to_length_truncated_response(provider): + response = await litellm.acompletion( + **_completion_kwargs(provider, _async_client_raising(provider, OUTPUT_LIMIT_400_MESSAGE)) + ) + + assert response.choices[0].finish_reason == "length" + assert response.choices[0].message.content == "" + assert response.usage.completion_tokens == 0 + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +def test_sync_streaming_output_limit_400_maps_to_length_truncated_stream(provider): + stream = litellm.completion( + **_completion_kwargs(provider, _sync_client_raising(provider, OUTPUT_LIMIT_400_MESSAGE), stream=True) + ) + chunks = list(stream) + + assert [c.choices[0].finish_reason for c in chunks].count("length") == 1 + assert all(not c.choices[0].delta.content for c in chunks) + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +@pytest.mark.asyncio +async def test_async_streaming_output_limit_400_maps_to_length_truncated_stream(provider): + stream = await litellm.acompletion( + **_completion_kwargs(provider, _async_client_raising(provider, OUTPUT_LIMIT_400_MESSAGE), stream=True) + ) + chunks = [chunk async for chunk in stream] + + assert [c.choices[0].finish_reason for c in chunks].count("length") == 1 + assert all(not c.choices[0].delta.content for c in chunks) + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +@pytest.mark.parametrize("stream", [False, True]) +def test_sync_genuine_bad_request_still_raises(provider, stream): + with pytest.raises(litellm.BadRequestError): + result = litellm.completion( + **_completion_kwargs(provider, _sync_client_raising(provider, GENUINE_400_MESSAGE), stream=stream) + ) + list(result) + + +@pytest.mark.parametrize("provider", ["openai", "azure"]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.asyncio +async def test_async_genuine_bad_request_still_raises(provider, stream): + with pytest.raises(litellm.BadRequestError): + result = await litellm.acompletion( + **_completion_kwargs(provider, _async_client_raising(provider, GENUINE_400_MESSAGE), stream=stream) + ) + async for _ in result: + pass From eb4b847268fb5cf6876d59bb3426104ee004a45d Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Fri, 14 Aug 2026 16:52:39 -0700 Subject: [PATCH 21/27] fix(proxy): always emit the Anthropic /v1/models token limits, null when unknown (#36961) Anthropic's Models API declares max_input_tokens and max_tokens as nullable, not optional, and the live vendor endpoint returns both keys on every entry. The merged Anthropic-native listing dropped either key whenever LiteLLM could not resolve a limit, so a client validating against a nullable-but-required schema saw a malformed entry for any model the cost map does not know. --- litellm/llms/anthropic/common_utils.py | 10 +++--- .../anthropic/test_anthropic_common_utils.py | 16 ++++++--- .../proxy/proxy_server/test_routes_models.py | 36 +++++++++++++++++-- 3 files changed, 49 insertions(+), 13 deletions(-) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index b444c77d718..1cdbd60f943 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -1227,16 +1227,13 @@ def process_anthropic_headers(headers: httpx.Headers | dict) -> dict: def _anthropic_model_entry(model: ModelInfoResponse, created_at: str) -> Mapping[str, object]: - token_limits: Final = ( - ("max_input_tokens", model.get("max_input_tokens")), - ("max_tokens", model.get("max_output_tokens")), - ) return { # mutable-ok: JSON response body, serialized by the route and never mutated "type": "model", "id": model["id"], "display_name": model["id"], "created_at": created_at, - **{name: limit for name, limit in token_limits if limit is not None}, # mutable-ok: merged into the body above + "max_input_tokens": model.get("max_input_tokens"), + "max_tokens": model.get("max_output_tokens"), } @@ -1246,7 +1243,8 @@ def create_anthropic_model_list_response(models: Sequence[ModelInfoResponse]) -> Clients that send an anthropic-version header parse the Anthropic Models API shape (type/display_name/created_at plus has_more/first_id/last_id) and filter the list themselves, so every model is returned here. The token limits carry - over from the OpenAI-shaped listing, named as the Messages API names them + over from the OpenAI-shaped listing, named as the Messages API names them, and + are always present because the vendor shape declares them nullable, not optional """ created_at: Final = ( datetime.fromtimestamp(DEFAULT_MODEL_CREATED_AT_TIME, tz=timezone.utc).isoformat().replace("+00:00", "Z") diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index 431030bcf2e..d205a903063 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -2056,11 +2056,14 @@ def test_create_anthropic_model_list_response_shape(): # ISO 8601 with a Z suffix, as the Anthropic Models API returns. assert entry["created_at"].endswith("Z") assert "+00:00" not in entry["created_at"] - assert "max_input_tokens" not in entry - assert "max_tokens" not in entry + assert entry["max_input_tokens"] is None + assert entry["max_tokens"] is None def test_create_anthropic_model_list_response_carries_token_limits(): + """max_input_tokens and max_tokens are nullable in the Anthropic Models shape, + not optional, so both keys are emitted for every entry and carry null when the + limit is unknown.""" from litellm.llms.anthropic.common_utils import ( create_anthropic_model_list_response, ) @@ -2091,9 +2094,12 @@ def test_create_anthropic_model_list_response_carries_token_limits(): assert opus["max_tokens"] == 64000 assert "max_output_tokens" not in opus assert input_only["max_input_tokens"] == 8192 - assert "max_tokens" not in input_only - assert "max_input_tokens" not in unknown - assert "max_tokens" not in unknown + assert input_only["max_tokens"] is None + assert unknown["max_input_tokens"] is None + assert unknown["max_tokens"] is None + for entry in response["data"]: + assert "max_input_tokens" in entry + assert "max_tokens" in entry def test_create_anthropic_model_list_response_empty(): diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_models.py b/tests/test_litellm/proxy/proxy_server/test_routes_models.py index f18c5998b8c..2b126b1ea95 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_models.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_models.py @@ -15,6 +15,8 @@ import pytest import litellm from litellm.proxy import proxy_server +from litellm.proxy import utils as proxy_utils +from litellm.proxy.utils import create_model_info_response from .conftest import normalize # type: ignore[import-not-found] @@ -151,8 +153,38 @@ def test_anthropic_format_exposes_token_limits( assert claude["max_input_tokens"] == 200000 assert claude["max_tokens"] == 64000 assert "max_output_tokens" not in claude - assert "max_input_tokens" not in gpt_4 - assert "max_tokens" not in gpt_4 + assert gpt_4["max_input_tokens"] is None + assert gpt_4["max_tokens"] is None + + +@pytest.mark.parametrize("path", ["/v1/models", "/models"]) +def test_anthropic_format_carries_router_configured_token_limits(client, auth_as, patched_models, monkeypatch, path): + """Pins the whole resolution chain, not just the formatter: a deployment's + configured limits beat the cost map, and the configured output budget is what + lands on the Anthropic ``max_tokens``. All eight limits differ, so an entry + built from another entry's lookup shows up as the wrong numbers.""" + + def _configured(model_name): + return (300000, 32000) if model_name == "gpt-4" else (500000, 4096) + + def _cost_map_lookup(model_id): + max_input, max_output = (200000, 64000) if model_id == "gpt-4" else (100000, 8000) + return {"max_input_tokens": max_input, "max_output_tokens": max_output, "mode": "chat"} + + patched_models.get_configured_token_limits = MagicMock(side_effect=_configured) + + def _resolved(**kwargs): + return create_model_info_response(**kwargs, get_model_info=_cost_map_lookup) + + monkeypatch.setattr(proxy_utils, "create_model_info_response", _resolved) + + with auth_as(): + response = client.get(path, headers={"anthropic-version": "2023-06-01"}) + + assert response.status_code == 200 + gpt_4, claude = response.json()["data"] + assert (gpt_4["max_input_tokens"], gpt_4["max_tokens"]) == (300000, 32000) + assert (claude["max_input_tokens"], claude["max_tokens"]) == (500000, 4096) @pytest.mark.parametrize("path", ["/v1/models", "/models"]) From f9704497fbe9c9d823dd0d7ebc47ad05f72b55bb Mon Sep 17 00:00:00 2001 From: Louis Vauterin <34511287+Louis-Vauterin@users.noreply.github.com> Date: Sat, 15 Aug 2026 01:53:10 +0200 Subject: [PATCH 22/27] feat(helm): add startupProbe and hpa.behavior to the componentized chart (#36382) Two small pod-spec passthroughs the componentized chart was missing, both additive and empty by default so existing renders are unchanged: - gateway/backend/ui deployments gain a `startupProbe` knob (same `{{- with }}` toYaml pattern as liveness/readiness), to gate liveness during a slow cold start without a kill loop. - gateway/backend/ui HPAs gain an `hpa.behavior` passthrough rendered verbatim under spec.behavior (scaleUp/scaleDown policies + stabilization windows). Tests: extend probe_tests.yaml (startupProbe absent by default / renders verbatim) and add hpa_behavior_tests.yaml. Full chart suite: 76 tests pass. Signed-off-by: Louis Vauterin Co-authored-by: Claude Opus 4.8 --- .../litellm/templates/backend/deployment.yaml | 4 ++ helm/litellm/templates/backend/hpa.yaml | 4 ++ .../litellm/templates/gateway/deployment.yaml | 4 ++ helm/litellm/templates/gateway/hpa.yaml | 4 ++ helm/litellm/templates/ui/deployment.yaml | 4 ++ helm/litellm/templates/ui/hpa.yaml | 4 ++ helm/litellm/tests/hpa_behavior_tests.yaml | 58 +++++++++++++++++++ helm/litellm/tests/probe_tests.yaml | 27 +++++++++ helm/litellm/values.yaml | 24 ++++++++ 9 files changed, 133 insertions(+) create mode 100644 helm/litellm/tests/hpa_behavior_tests.yaml diff --git a/helm/litellm/templates/backend/deployment.yaml b/helm/litellm/templates/backend/deployment.yaml index c5d799a0faf..5c0431fc0bd 100644 --- a/helm/litellm/templates/backend/deployment.yaml +++ b/helm/litellm/templates/backend/deployment.yaml @@ -81,6 +81,10 @@ spec: readinessProbe: {{- toYaml . | nindent 12 }} {{- end }} + {{- with .Values.backend.startupProbe }} + startupProbe: + {{- toYaml . | nindent 12 }} + {{- end }} {{- with .Values.backend.lifecycle }} lifecycle: {{- toYaml . | nindent 12 }} diff --git a/helm/litellm/templates/backend/hpa.yaml b/helm/litellm/templates/backend/hpa.yaml index d02f011d0bb..a414092fb39 100644 --- a/helm/litellm/templates/backend/hpa.yaml +++ b/helm/litellm/templates/backend/hpa.yaml @@ -30,4 +30,8 @@ spec: type: Utilization averageUtilization: {{ .Values.backend.hpa.targetMemoryUtilizationPercentage }} {{- end }} + {{- with .Values.backend.hpa.behavior }} + behavior: + {{- toYaml . | nindent 4 }} + {{- end }} {{- end }} diff --git a/helm/litellm/templates/gateway/deployment.yaml b/helm/litellm/templates/gateway/deployment.yaml index 7d16134a53d..d5363d0096e 100644 --- a/helm/litellm/templates/gateway/deployment.yaml +++ b/helm/litellm/templates/gateway/deployment.yaml @@ -83,6 +83,10 @@ spec: readinessProbe: {{- toYaml . | nindent 12 }} {{- end }} + {{- with .Values.gateway.startupProbe }} + startupProbe: + {{- toYaml . | nindent 12 }} + {{- end }} {{- with .Values.gateway.lifecycle }} lifecycle: {{- toYaml . | nindent 12 }} diff --git a/helm/litellm/templates/gateway/hpa.yaml b/helm/litellm/templates/gateway/hpa.yaml index 27c4f05ba59..e97cef95ffb 100644 --- a/helm/litellm/templates/gateway/hpa.yaml +++ b/helm/litellm/templates/gateway/hpa.yaml @@ -30,4 +30,8 @@ spec: type: Utilization averageUtilization: {{ .Values.gateway.hpa.targetMemoryUtilizationPercentage }} {{- end }} + {{- with .Values.gateway.hpa.behavior }} + behavior: + {{- toYaml . | nindent 4 }} + {{- end }} {{- end }} diff --git a/helm/litellm/templates/ui/deployment.yaml b/helm/litellm/templates/ui/deployment.yaml index b4129dbc8ac..91d6de39ea6 100644 --- a/helm/litellm/templates/ui/deployment.yaml +++ b/helm/litellm/templates/ui/deployment.yaml @@ -69,6 +69,10 @@ spec: readinessProbe: {{- toYaml . | nindent 12 }} {{- end }} + {{- with .Values.ui.startupProbe }} + startupProbe: + {{- toYaml . | nindent 12 }} + {{- end }} {{- with .Values.ui.lifecycle }} lifecycle: {{- toYaml . | nindent 12 }} diff --git a/helm/litellm/templates/ui/hpa.yaml b/helm/litellm/templates/ui/hpa.yaml index b43eda5ac4a..a9b0b51129e 100644 --- a/helm/litellm/templates/ui/hpa.yaml +++ b/helm/litellm/templates/ui/hpa.yaml @@ -30,4 +30,8 @@ spec: type: Utilization averageUtilization: {{ .Values.ui.hpa.targetMemoryUtilizationPercentage }} {{- end }} + {{- with .Values.ui.hpa.behavior }} + behavior: + {{- toYaml . | nindent 4 }} + {{- end }} {{- end }} diff --git a/helm/litellm/tests/hpa_behavior_tests.yaml b/helm/litellm/tests/hpa_behavior_tests.yaml new file mode 100644 index 00000000000..84d0ff8a2ae --- /dev/null +++ b/helm/litellm/tests/hpa_behavior_tests.yaml @@ -0,0 +1,58 @@ +suite: test HPA scaling behavior passthrough +templates: + - gateway/hpa.yaml + - backend/hpa.yaml + - ui/hpa.yaml +values: + - ./values/required.yaml +tests: + - it: HPA omits spec.behavior by default, so Kubernetes' default scaling applies + templates: + - gateway/hpa.yaml + - backend/hpa.yaml + asserts: + - isKind: + of: HorizontalPodAutoscaler + - notExists: + path: spec.behavior + + - it: gateway HPA renders spec.behavior verbatim when configured + template: gateway/hpa.yaml + set: + gateway.hpa.behavior: + scaleDown: + stabilizationWindowSeconds: 300 + policies: + - { type: Percent, value: 50, periodSeconds: 60 } + scaleUp: + stabilizationWindowSeconds: 0 + selectPolicy: Max + policies: + - { type: Percent, value: 100, periodSeconds: 30 } + - { type: Pods, value: 2, periodSeconds: 30 } + asserts: + - equal: + path: spec.behavior + value: + scaleDown: + stabilizationWindowSeconds: 300 + policies: + - { type: Percent, value: 50, periodSeconds: 60 } + scaleUp: + stabilizationWindowSeconds: 0 + selectPolicy: Max + policies: + - { type: Percent, value: 100, periodSeconds: 30 } + - { type: Pods, value: 2, periodSeconds: 30 } + + - it: behavior passthrough works on every autoscaled component (ui parity) + template: ui/hpa.yaml + set: + ui.hpa.enabled: true + ui.hpa.behavior: + scaleUp: + stabilizationWindowSeconds: 0 + asserts: + - equal: + path: spec.behavior.scaleUp.stabilizationWindowSeconds + value: 0 diff --git a/helm/litellm/tests/probe_tests.yaml b/helm/litellm/tests/probe_tests.yaml index a04709db2f5..a2866bb7648 100644 --- a/helm/litellm/tests/probe_tests.yaml +++ b/helm/litellm/tests/probe_tests.yaml @@ -104,3 +104,30 @@ tests: periodSeconds: 15 timeoutSeconds: 4 failureThreshold: 3 + + - it: no startupProbe by default, so existing installs are unchanged + templates: + - gateway/deployment.yaml + - backend/deployment.yaml + asserts: + - notExists: + path: spec.template.spec.containers[0].startupProbe + + - it: startupProbe renders verbatim when configured, gating a slow cold start + template: gateway/deployment.yaml + set: + gateway.startupProbe: + httpGet: { path: /health/readiness, port: http } + failureThreshold: 30 + periodSeconds: 10 + timeoutSeconds: 5 + asserts: + - equal: + path: spec.template.spec.containers[0].startupProbe + value: + httpGet: + path: /health/readiness + port: http + failureThreshold: 30 + periodSeconds: 10 + timeoutSeconds: 5 diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index cd377667602..7820a898ef1 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -223,12 +223,28 @@ gateway: initialDelaySeconds: 5 periodSeconds: 10 timeoutSeconds: 10 + # Optional startupProbe. Empty by default, so existing installs are unchanged + # and liveness/readiness apply from container start. Set it to gate + # liveness/readiness until a slow cold start finishes — a high failureThreshold + # tolerates long first-boot times without a liveness-kill loop, e.g.: + # httpGet: { path: /health/readiness, port: http } + # failureThreshold: 30 + # periodSeconds: 10 + startupProbe: {} hpa: enabled: true minReplicas: 1 maxReplicas: 10 targetCPUUtilizationPercentage: 70 targetMemoryUtilizationPercentage: 80 + # Optional autoscaling/v2 scaling behavior (scaleUp / scaleDown policies and + # stabilization windows). Empty by default -> Kubernetes' default behavior. + # Rendered verbatim under spec.behavior, e.g.: + # scaleUp: + # stabilizationWindowSeconds: 0 + # policies: + # - { type: Percent, value: 100, periodSeconds: 30 } + behavior: {} # PodDisruptionBudget for the gateway pods. Set exactly one of # `minAvailable` / `maxUnavailable` (minAvailable wins if both are set; # enabling without either falls back to `maxUnavailable: 1`). Disabled by @@ -319,11 +335,15 @@ backend: initialDelaySeconds: 5 periodSeconds: 10 timeoutSeconds: 10 + # Optional startupProbe; same shape as gateway.startupProbe. Empty by default. + startupProbe: {} hpa: enabled: true minReplicas: 1 maxReplicas: 4 targetCPUUtilizationPercentage: 70 + # Optional autoscaling/v2 scaling behavior; same shape as gateway.hpa.behavior. + behavior: {} # Same shape as gateway.pdb. pdb: enabled: false @@ -379,11 +399,15 @@ ui: httpGet: { path: /, port: http } initialDelaySeconds: 2 periodSeconds: 10 + # Optional startupProbe; same shape as gateway.startupProbe. Empty by default. + startupProbe: {} hpa: enabled: false minReplicas: 1 maxReplicas: 3 targetCPUUtilizationPercentage: 80 + # Optional autoscaling/v2 scaling behavior; same shape as gateway.hpa.behavior. + behavior: {} # Same shape as gateway.pdb. pdb: enabled: false From b14c4a8d458a4021bbeddcfebcec69d36eef9535 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 16:53:48 -0700 Subject: [PATCH 23/27] fix(vector_stores): classify write endpoints before reads on substring collisions --- .../azure_ai/vector_stores/transformation.py | 7 ++- litellm/proxy/vector_store_endpoints/utils.py | 20 ++++--- .../test_vector_store_endpoints.py | 55 +++++++++++++++++++ 3 files changed, 71 insertions(+), 11 deletions(-) diff --git a/litellm/llms/azure_ai/vector_stores/transformation.py b/litellm/llms/azure_ai/vector_stores/transformation.py index f58d2f54d2e..5e61d0a1dd9 100644 --- a/litellm/llms/azure_ai/vector_stores/transformation.py +++ b/litellm/llms/azure_ai/vector_stores/transformation.py @@ -48,8 +48,11 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM): Patterns stay literal rather than ``{placeholder}`` templates because the matcher falls back to the substring before a ``{``, which here is always - ``/indexes/`` -- broad enough that a templated read, matched first, would - shadow the ``/docs/index`` write. + ``/indexes/``. The matcher is substring-based, so an index name may + itself contain a read fragment (an index named ``analyze*`` puts + ``/analyze`` inside the batch-write path); writes are classified before + reads, so such a path demands the write grant rather than being + shadowed into a read. """ return { "read": [ diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index afde5c787f1..93f1510bf22 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -387,17 +387,19 @@ def is_allowed_to_call_vector_store_endpoint( ) return True - # Determine the permission type based on the request + # Writes are classified before reads so a path matching both patterns + # requires the stronger grant (e.g. the azure batch write on an index + # named "analyze*" also contains the "/analyze" read fragment) permission_type = None - for endpoint in provider_vector_store_endpoints["read"]: + for endpoint in provider_vector_store_endpoints["write"]: if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route): - permission_type = "read" + permission_type = "write" break if permission_type is None: - for endpoint in provider_vector_store_endpoints["write"]: + for endpoint in provider_vector_store_endpoints["read"]: if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route): - permission_type = "write" + permission_type = "read" break if permission_type is None: @@ -454,15 +456,15 @@ def is_allowed_to_call_vector_store_files_endpoint( request_route: Final = get_request_route(request) permission_type: str | None = None - for endpoint in provider_vector_store_endpoints.get("read", ()): + for endpoint in provider_vector_store_endpoints.get("write", ()): if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route): - permission_type = "read" + permission_type = "write" break if permission_type is None: - for endpoint in provider_vector_store_endpoints.get("write", ()): + for endpoint in provider_vector_store_endpoints.get("read", ()): if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route): - permission_type = "write" + permission_type = "read" break if permission_type is None: diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 8a227028b51..20b2f68bb0c 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -3057,3 +3057,58 @@ class TestAzureAIDocumentWritePassthroughPermission: ) assert exc_info.value.status_code == 403 assert f"Only proxy admins can {operation}" in exc_info.value.detail + + +class TestAzureAIAnalyzeNamedIndexClassification: + """Regression tests for write-before-read endpoint classification. + + The endpoint matcher is substring-based, so the batch-write path of an + index named ``analyze*`` contains the ``("POST", "/analyze")`` read + fragment. Reads-first classification labeled that write a read, letting a + read-only grant upload, merge, and delete documents (and refusing + legitimate write-only grants). Writes are classified first now, so an + ambiguous path demands the stronger grant. + """ + + def _request(self, method: str, path: str) -> MagicMock: + request = MagicMock(spec=Request) + request.method = method + request.url.path = path + return request + + def _team_member(self, index: str, permissions: list) -> MagicMock: + user = MagicMock(spec=UserAPIKeyAuth) + user.user_role = None + user.metadata = {"allowed_vector_store_indexes": [{"index_name": index, "index_permissions": permissions}]} + user.team_metadata = None + return user + + @pytest.mark.parametrize("index", ["analyze", "analyzer-reports"]) + def test_read_only_grant_cannot_upload_to_analyze_named_index(self, index): + with pytest.raises(HTTPException) as exc_info: + is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=index, + request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"), + user_api_key_dict=self._team_member(index, ["read"]), + ) + assert exc_info.value.status_code == 403 + + @pytest.mark.parametrize("index", ["analyze", "analyzer-reports"]) + def test_write_grant_can_upload_to_analyze_named_index(self, index): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name=index, + request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"), + user_api_key_dict=self._team_member(index, ["write"]), + ) + assert result is True + + def test_read_only_grant_can_still_analyze_on_analyze_named_index(self): + result = is_allowed_to_call_vector_store_endpoint( + provider=LlmProviders.AZURE_AI, + index_name="analyze", + request=self._request("POST", "/azure_ai/indexes/analyze/analyze"), + user_api_key_dict=self._team_member("analyze", ["read"]), + ) + assert result is True From d4d6bc25771484e278c367740d841f42e5c5c114 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 14 Aug 2026 17:04:32 -0700 Subject: [PATCH 24/27] fix(proxy): serve aggregate MCP endpoint on bare /mcp instead of 307-redirecting (#34845) The MCP sub-app is attached with app.mount("/mcp", ...) and a Starlette mount never matches its bare prefix, so POST /mcp fell through to the router's redirect_slashes 307. Behind a TLS-terminating ingress whose peer address is not in uvicorn's forwarded-allow-ips (default: loopback only) the redirect Location is built from the socket scheme as http://, and MCP clients strip the Authorization header on the cross-origin follow, so reconnects fail with ECONNRESET right after a successful OAuth flow. The redirect also fires before auth, so the bare spelling never returns the RFC 9728 WWW-Authenticate challenge that OAuth clients need to start the flow. Add an explicit /mcp route beside the existing /toolset/{name}/mcp and /{name}/mcp spellings, forwarding to handle_streamable_http_mcp with the same scope rewrite those routes already use (path=/mcp, _original_path preserved for OAuth challenge URL selection). When the mcp package is unavailable the route 404s, matching what the bare sub-app serves on /mcp/ in that state. /mcp/, /mcp/{server}, /{server}/mcp and /toolset/{name}/mcp spellings are unchanged; the exact-match route and the mount have disjoint match sets so registration order cannot matter. --- backend/routes/allowlist.py | 2 + litellm/proxy/proxy_server.py | 23 ++ .../proxy/test_dynamic_mcp_route.py | 71 +++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 198 ++++++++++++++++++ 4 files changed, 294 insertions(+) diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 8ccd439979b..3f7bf788a1b 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -146,11 +146,13 @@ BACKEND_EXACT_PATHS: frozenset[str] = frozenset( "/docs/oauth2-redirect", "/redoc", "/fallback/login", + "/mcp", # bare spelling of the aggregate MCP endpoint; /mcp/ prefix covers the rest } ) BACKEND_MOUNT_PATHS: frozenset[str] = frozenset( { "/swagger", # API documentation static assets belong to the backend + "/mcp", # lazily-mounted MCP sub-app serves on the backend component } ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 359187f81cb..377963f1b91 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -17274,6 +17274,29 @@ async def _stream_mcp_asgi_response(handle_fn, scope: dict, receive) -> "Streami ######################################################## +@app.api_route( + BASE_MCP_ROUTE, + methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"], +) +async def aggregate_mcp_route(request: Request): + """Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + MCP clients behind TLS-terminating proxies.""" + from litellm.proxy._experimental.mcp_server.utils import is_mcp_available + + if not is_mcp_available(): + raise HTTPException(status_code=404, detail="Not Found") + + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + ) + + scope = dict(request.scope) + scope["_original_path"] = scope.get("path", "") + scope["path"] = BASE_MCP_ROUTE + return await _stream_mcp_asgi_response(handle_streamable_http_mcp, scope, request.receive) + + # Toolset-namespaced MCP routes - handle /toolset/{toolset_name}/mcp # Must be declared BEFORE /{mcp_server_name}/mcp to avoid being swallowed by the catchall. @app.api_route( diff --git a/tests/test_litellm/proxy/test_dynamic_mcp_route.py b/tests/test_litellm/proxy/test_dynamic_mcp_route.py index 592cebd957c..da7b8e01f46 100644 --- a/tests/test_litellm/proxy/test_dynamic_mcp_route.py +++ b/tests/test_litellm/proxy/test_dynamic_mcp_route.py @@ -540,3 +540,74 @@ async def test_toolset_mcp_route_unexpected_exception_returns_500_without_traceb assert exc_info.value.detail == "Internal server error" assert "db-host" not in str(exc_info.value.detail) assert "traceback" not in str(exc_info.value.detail).lower() + + +# --------------------------------------------------------------------------- +# 7. Aggregate /mcp without a trailing slash (bare mount prefix) +# --------------------------------------------------------------------------- + +_IS_MCP_AVAILABLE = "litellm.proxy._experimental.mcp_server.utils.is_mcp_available" + + +def _test_client(): + from fastapi.testclient import TestClient + + from litellm.proxy.proxy_server import app + + return TestClient(app, follow_redirects=False) + + +@pytest.mark.parametrize("method", ["GET", "POST", "DELETE"]) +def test_aggregate_mcp_route_bare_path_is_served_not_redirected(method): + """Bare /mcp must dispatch to the MCP handler with aggregate semantics, + never 307-redirect. Driven through the real app router so a lost route + registration (not just a broken handler body) fails this test.""" + captured_scope: dict = {} + + async def capturing_handle(scope, receive, send): + captured_scope.update(scope) + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b"{}"}) + + with patch(_HANDLE_HTTP, new=capturing_handle): + response = _test_client().request(method, "/mcp") + + assert response.status_code == 200 + assert captured_scope.get("path") == "/mcp" + assert captured_scope.get("_original_path") == "/mcp" + + +def test_aggregate_mcp_route_requires_exact_path(): + """The bare-path route must match exactly /mcp; a sibling path like /mcpx + must not reach the MCP handler through it.""" + calls = [] + + async def marking_handle(scope, receive, send): + calls.append(scope.get("path")) + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b"{}"}) + + with patch(_HANDLE_HTTP, new=marking_handle): + response = _test_client().post("/mcpx") + + assert calls == [] + assert response.status_code != 200 + + +def test_aggregate_mcp_route_returns_404_when_mcp_unavailable(): + """When the mcp package is unavailable the canonical /mcp/ sub-app is a + bare FastAPI that 404s, so the bare spelling must 404 identically instead + of erroring on the handler import.""" + handler_calls = [] + + async def marking_handle(scope, receive, send): + handler_calls.append(scope.get("path")) + + with ( + patch(_IS_MCP_AVAILABLE, new=MagicMock(return_value=False)), + patch(_HANDLE_HTTP, new=marking_handle), + ): + response = _test_client().post("/mcp") + + assert response.status_code == 404 + assert handler_calls == [] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1c5fe408203..2cbc7fd6220 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -7616,6 +7616,64 @@ export interface paths { patch?: never; trace?: never; }; + "/mcp": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + get: operations["aggregate_mcp_route_mcp_get"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + put: operations["aggregate_mcp_route_mcp_put"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + post: operations["aggregate_mcp_route_mcp_post"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + delete: operations["aggregate_mcp_route_mcp_delete"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + options: operations["aggregate_mcp_route_mcp_options"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + head: operations["aggregate_mcp_route_mcp_head"]; + /** + * Aggregate Mcp Route + * @description Serve the aggregate MCP endpoint on the bare ``/mcp`` spelling: the + * ``/mcp`` mount cannot match its bare prefix, and the resulting 307 breaks + * MCP clients behind TLS-terminating proxies. + */ + patch: operations["aggregate_mcp_route_mcp_patch"]; + trace?: never; + }; "/mcp-rest/test/connection": { parameters: { query?: never; @@ -45790,6 +45848,146 @@ export interface operations { }; }; }; + aggregate_mcp_route_mcp_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_put: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_post: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_delete: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_options: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_head: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; + aggregate_mcp_route_mcp_patch: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + }; + }; test_connection_mcp_rest_test_connection_post: { parameters: { query?: never; From 2d3c3e30986a2d5050ae781fb8e633776f890b6b Mon Sep 17 00:00:00 2001 From: tin-berri Date: Fri, 14 Aug 2026 17:05:55 -0700 Subject: [PATCH 25/27] feat(shadow_eval): add reverse-direction shadow eval jobs (#36865) Shadow eval only answered "should this key adopt this auto-router". Once a key is on the router it is invisible to the feature, because the sampling gate skips any request the shadowed router already served, so post-adoption quality regressions go unmeasured. Reverse mode inverts the arms: sample the traffic the router did serve and duplicate it against a fixed baseline_model, judged by the same blind pairwise judge. Same job table, same attempt rows, same aggregates. real_* stays the arm the caller was served and shadow_* the duplicated one, so in reverse real_model is the router's pick and shadow_model is the baseline. The active-job slot becomes one per (key, direction) so both directions can run at once, and tier attribution in reverse reads the control request's routing decision rather than the shadow call's write-back. --- .../migration.sql | 8 + .../litellm_proxy_extras/schema.prisma | 13 +- litellm/integrations/shadow_eval_logger.py | 205 +++++++++++------ .../auto_router_endpoints.py | 67 ++++-- litellm/proxy/schema.prisma | 13 +- .../auto_router_endpoints.py | 57 ++++- schema.prisma | 13 +- .../integrations/test_shadow_eval_logger.py | 214 ++++++++++++++++-- .../test_auto_router_endpoints.py | 79 ++++++- .../_components/ShadowEvalSection.test.tsx | 1 + .../_components/ShadowEvalSection.tsx | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 47 +++- 12 files changed, 575 insertions(+), 143 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260813180408_add_shadow_eval_direction/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260813180408_add_shadow_eval_direction/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260813180408_add_shadow_eval_direction/migration.sql new file mode 100644 index 00000000000..57c9abab07d --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260813180408_add_shadow_eval_direction/migration.sql @@ -0,0 +1,8 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN "baseline_model" TEXT, +ADD COLUMN "direction" TEXT NOT NULL DEFAULT 'forward'; + +DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key"; + +CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key_direction" + ON "LiteLLM_ShadowEvalJob"("api_key_id", "direction") WHERE "stopped_at" IS NULL; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 79d778fb464..71345d2ccde 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } -// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic. -// A sampled slice of requests is duplicated through the router in a detached task and an -// LLM judge compares real vs shadow responses blind. The job row is immutable config plus +// Shadow eval: evaluation of an auto-router against a key's live traffic, in either +// direction. forward duplicates the requests the key did not route through the router +// through it, answering whether the key should adopt it; reverse duplicates the requests +// the router did serve against a fixed baseline model, answering whether a key already on +// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge +// compares real vs shadow responses blind. The job row is immutable config plus // stopped_at; every count, status, and spend figure is derived from the append-only // attempt rows, so nothing can disagree across pods or stop races. model LiteLLM_ShadowEvalJob { id String @id @default(cuid()) api_key_id String // hashed virtual key whose traffic is shadowed - router_name String + router_name String // the auto-router under evaluation, in either direction + direction String @default("forward") // forward | reverse + baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float max_turns Int // sample budget: judge at most this many turns diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index c7b89e0e9b0..ca9b6982414 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -1,5 +1,6 @@ """Shadow Eval Logger: samples a shadowed key's successful chat requests, duplicates each -through the auto-router in a detached task, blind-judges real vs shadow, and appends one +against the job's other arm in a detached task (the auto-router for a forward job, the +fixed baseline model for a reverse one), blind-judges real vs shadow, and appends one ``LiteLLM_ShadowEvalAttempt`` row (verdict or error) as the feature's only hot-path write. Counts, status, and spend derive from those rows at read time, so nothing can disagree across pods or stop races; the hook reads active jobs through a short-TTL cache.""" @@ -10,10 +11,12 @@ import random from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timezone +from itertools import groupby +from operator import itemgetter from types import MappingProxyType from typing import TYPE_CHECKING, Final -from pydantic import BaseModel +from pydantic import BaseModel, ConfigDict, ValidationError, field_validator, model_validator from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache @@ -28,6 +31,7 @@ from litellm.litellm_core_utils.llm_judge import ( parse_json_verdict, ) from litellm.litellm_core_utils.redact_messages import should_redact_message_logging +from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN if TYPE_CHECKING: @@ -161,13 +165,26 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: return False +def _routing_decision(metadata: Mapping[str, object]) -> Mapping[str, object]: + """The routing decision a pre-routing strategy wrote to a call's metadata, empty when + a plain model served it. Read off the sampled request for the control arm, and off the + shadow call's own write-back for the shadow arm.""" + decision: Final = metadata.get("routing_decision") + return decision if isinstance(decision, Mapping) else _EMPTY_METADATA + + +def _routed_tier(metadata: Mapping[str, object]) -> str | None: + decision: Final = _routing_decision(metadata) + raw: Final = decision.get("tier_label") or decision.get("tier") + return str(raw) if raw is not None else None + + def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: str) -> bool: - """Duplicating a request the shadowed router already served compares the router to - itself: guaranteed ties, judge spend for zero information.""" - decision: Final = request_metadata.get("routing_decision") - if not isinstance(decision, Mapping): - return False - return decision.get("router_model_name") == router_name + """Whether the router under evaluation served this request, which is what decides + the direction it belongs to. A forward job skips its own router's traffic, since + duplicating it would compare the router to itself: guaranteed ties, judge spend for + zero information. A reverse job samples exactly that traffic and nothing else.""" + return _routing_decision(request_metadata).get("router_model_name") == router_name @dataclass(frozen=True, slots=True) @@ -197,22 +214,53 @@ class _JudgeVerdict: cost: float -@dataclass(frozen=True, slots=True) -class ActiveShadowEvalJob: - """One active job as the sampling path needs it: immutable config plus the attempt - count as of the cache fill (the turn budget's staleness is bounded by the cache TTL).""" +class ActiveShadowEvalJob(BaseModel): + """One active job as the sampling path needs it, validated straight off the untyped + job row: immutable config plus the attempt count as of the cache fill (the turn + budget's staleness is bounded by the cache TTL). Every way a row can be unsamplable + is a validation error here, so a bad row is skipped rather than sampled wrongly.""" + + model_config = ConfigDict(frozen=True, from_attributes=True) id: str router_name: str + direction: ShadowEvalDirection = "forward" + baseline_model: str | None = None shadow_percentage: float judge_model: str max_turns: int ends_at: datetime - attempts: int + attempts: int = 0 + + @field_validator("ends_at") + @classmethod + def _as_utc(cls, value: datetime) -> datetime: + return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value + + @model_validator(mode="after") + def _baseline_model_matches_direction(self) -> "ActiveShadowEvalJob": + if (self.baseline_model is not None) != (self.direction == "reverse"): + raise ValueError("baseline_model is set for exactly the reverse jobs") + return self + + @property + def shadow_target(self) -> str: + """The model the duplicated arm calls: the router itself for a forward job, the + fixed baseline for a reverse one. Total because the validator above pins + baseline_model to reverse jobs and only those.""" + return self.baseline_model or self.router_name -def _as_utc(value: datetime) -> datetime: - return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value +def _as_active_job(record: object, attempts: int) -> ActiveShadowEvalJob | None: + """The sampling path's view of one job row, or None for a row it cannot sample: an + unknown direction, or a reverse job with no baseline model to duplicate against. + Failing closed here is what keeps the dispatch path total.""" + try: + job: Final = ActiveShadowEvalJob.model_validate(record) + except ValidationError as e: + verbose_logger.debug("shadow_eval: skipping unsamplable job row: %s", e) + return None + return job.model_copy(update={"attempts": attempts}) _jobs_cache: Final = InMemoryCache(max_size_in_memory=4, default_ttl=_JOBS_CACHE_TTL_SECONDS) @@ -238,8 +286,9 @@ class ShadowEvalLogger(CustomLogger): # generation; the refill absorbs written rows and resets. self._job_starts: dict[str, int] = {} # mutable-ok: per-generation counter - async def _active_jobs(self) -> Mapping[str, ActiveShadowEvalJob]: - """Active jobs by api_key_id, cache-first. A DB fault returns empty without + async def _active_jobs(self) -> Mapping[str, tuple[ActiveShadowEvalJob, ...]]: + """Active jobs by api_key_id, cache-first. A key holds at most one job per + direction, so the value is a collection. A DB fault returns empty without caching, so sampling pauses for that request and the next one retries.""" cached: Final = await self._jobs_cache.async_get_cache(_JOBS_CACHE_KEY) if cached is not None: @@ -264,18 +313,19 @@ class ShadowEvalLogger(CustomLogger): else () ) attempt_counts: Final = {str(row["job_id"]): int(row["_count"]["_all"]) for row in grouped or []} - jobs: Final = { - str(record.api_key_id): ActiveShadowEvalJob( - id=str(record.id), - router_name=str(record.router_name), - shadow_percentage=float(record.shadow_percentage), - judge_model=str(record.judge_model), - max_turns=int(record.max_turns), - ends_at=_as_utc(record.ends_at), - attempts=attempt_counts.get(str(record.id), 0), + by_key: Final = tuple( + sorted( + ( + (str(record.api_key_id), job) + for record in records or [] + if (job := _as_active_job(record, attempt_counts.get(str(record.id), 0))) is not None + ), + key=itemgetter(0), ) - for record in records or [] - } + ) + jobs: Final = MappingProxyType( + {key: tuple(job for _, job in group) for key, group in groupby(by_key, key=itemgetter(0))} + ) await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs) self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill return jobs @@ -308,43 +358,46 @@ class ShadowEvalLogger(CustomLogger): api_key_hash: Final = metadata.get("user_api_key_hash") if not api_key_hash: return - job: Final = (await self._active_jobs()).get(str(api_key_hash)) - if job is None: - return - if datetime.now(timezone.utc) >= job.ends_at: - return - if job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns: - return request_id: Final = payload.get("id") or "" if not request_id: return - if not _sample_hits(request_id, job.id, job.shadow_percentage): - return if payload.get("call_type") not in _SAMPLED_CALL_TYPES: return # only known chat-shaped traffic is comparable; unknown or missing types fail closed - if _request_was_routed_by(request_metadata, job.router_name): - return - if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS: - return raw_messages: Final = kwargs.get("messages") - self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1 - self._inflight_shadow_tasks += 1 - task: Final = asyncio.create_task( - self._run_shadow_eval( - job=job, - request_id=request_id, - messages=tuple(m for m in raw_messages if isinstance(m, Mapping)) - if isinstance(raw_messages, Sequence) - else (), - response_obj=response_obj, - real_model=payload.get("model") or "", - model_parameters=MappingProxyType( - dict(payload.get("model_parameters") or {}) # mutable-ok: frozen snapshot - ), - parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot - ) + messages: Final = ( + tuple(m for m in raw_messages if isinstance(m, Mapping)) if isinstance(raw_messages, Sequence) else () ) - task.add_done_callback(self._release_shadow_slot) + control_tier: Final = _routed_tier(request_metadata) + # A key can hold one job per direction, and a request routed by one job's + # router while bypassing the other's qualifies for both. Each is separately + # budgeted, so both fire. + for job in (await self._active_jobs()).get(str(api_key_hash), ()): + if datetime.now(timezone.utc) >= job.ends_at: + continue + if job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns: + continue + if not _sample_hits(request_id, job.id, job.shadow_percentage): + continue + if _request_was_routed_by(request_metadata, job.router_name) != (job.direction == "reverse"): + continue + if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS: + return + self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1 + self._inflight_shadow_tasks += 1 + asyncio.create_task( + self._run_shadow_eval( + job=job, + request_id=request_id, + messages=messages, + response_obj=response_obj, + real_model=payload.get("model") or "", + control_tier=control_tier, + model_parameters=MappingProxyType( + dict(payload.get("model_parameters") or {}) # mutable-ok: frozen snapshot + ), + parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot + ) + ).add_done_callback(self._release_shadow_slot) except Exception as e: # noqa: BLE001 # logging hooks must never fail the request verbose_logger.debug("shadow_eval: failed to schedule task: %s", e) @@ -360,6 +413,7 @@ class ShadowEvalLogger(CustomLogger): messages: Sequence[Mapping[str, object]], response_obj: object, real_model: str, + control_tier: str | None, model_parameters: Mapping[str, object], parent_metadata: Mapping[str, object], ) -> None: @@ -376,9 +430,11 @@ class ShadowEvalLogger(CustomLogger): if await _key_or_team_is_over_budget(parent_metadata): return - shadow: Final = await self._call_router_shadow(job.router_name, messages, model_parameters, parent_metadata) + shadow: Final = await self._call_router_shadow( + job.shadow_target, messages, model_parameters, parent_metadata + ) if isinstance(shadow, _CallFailure): - await self._record_attempt(prisma, job, request_id, outcome="error", error=shadow.error) + await self._record_attempt(prisma, job, request_id, control_tier, outcome="error", error=shadow.error) return verdict: Final = await self._call_judge( @@ -393,6 +449,7 @@ class ShadowEvalLogger(CustomLogger): prisma, job, request_id, + control_tier, outcome="error", error=verdict.error, shadow=shadow, @@ -403,6 +460,7 @@ class ShadowEvalLogger(CustomLogger): prisma, job, request_id, + control_tier, outcome=verdict.preference, shadow=shadow, real_model=real_model, @@ -411,13 +469,16 @@ class ShadowEvalLogger(CustomLogger): ) except Exception as e: # noqa: BLE001 # detached task: record what happened, never raise verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e) - await self._record_attempt(prisma, job, request_id, outcome="error", error=f"pipeline error: {e}") + await self._record_attempt( + prisma, job, request_id, control_tier, outcome="error", error=f"pipeline error: {e}" + ) @staticmethod async def _record_attempt( prisma: "PrismaClient | None", job: ActiveShadowEvalJob, request_id: str, + control_tier: str | None, *, outcome: str, shadow: _ShadowResponse | None = None, @@ -434,7 +495,7 @@ class ShadowEvalLogger(CustomLogger): "job_id": job.id, "request_id": request_id, "outcome": outcome, - "tier": shadow.tier if shadow else None, + "tier": control_tier if job.direction == "reverse" else (shadow.tier if shadow else None), "real_model": real_model or None, "shadow_model": shadow.model if shadow else None, "confidence": confidence, @@ -447,14 +508,15 @@ class ShadowEvalLogger(CustomLogger): async def _call_router_shadow( self, - router_name: str, + target_model: str, messages: Sequence[Mapping[str, object]], model_parameters: Mapping[str, object], parent_metadata: Mapping[str, object], ) -> "_ShadowResponse | _CallFailure": - """Send the prompt through the auto-router being evaluated. The metadata carries - the shadowed key's identity (spend attribution) and receives the router's routing - decision write-back, read back for tier attribution.""" + """Send the prompt through the arm nobody was served: the auto-router under + evaluation, or a reverse job's fixed baseline model. The metadata carries the + shadowed key's identity (spend attribution) and receives a routing decision + write-back, which a plain baseline model simply never makes.""" router: Final = self._router_provider() if router is None: return _CallFailure("no router configured on this pod") @@ -466,7 +528,7 @@ class ShadowEvalLogger(CustomLogger): } try: response: Final = await router.acompletion( - model=router_name, + model=target_model, messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts metadata=shadow_metadata, num_retries=0, @@ -479,13 +541,10 @@ class ShadowEvalLogger(CustomLogger): text: Final = self._extract_response_text(response) if not text: return _CallFailure("shadow router returned an empty response") - raw_decision: Final = shadow_metadata.get("routing_decision") - routing_decision: Final = raw_decision if isinstance(raw_decision, Mapping) else _EMPTY_METADATA - raw_tier: Final = routing_decision.get("tier_label") or routing_decision.get("tier") return _ShadowResponse( text=text, - model=str(getattr(response, "model", None) or routing_decision.get("routed_model") or ""), - tier=str(raw_tier) if raw_tier is not None else None, + model=str(getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""), + tier=_routed_tier(shadow_metadata), ) async def _call_judge( @@ -552,7 +611,7 @@ class ShadowEvalLogger(CustomLogger): return extract_text_from_content(content) -_EMPTY_JOBS: Final[Mapping[str, ActiveShadowEvalJob]] = MappingProxyType({}) +_EMPTY_JOBS: Final[Mapping[str, tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({}) def _default_prisma_provider() -> "PrismaClient | None": diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index cb0e8dba62a..4b2569fa9fa 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -464,35 +464,38 @@ def _is_configured_pre_routing_strategy(llm_router: "Router", router_name: str) ) -def _validate_judge_model(llm_router: "Router | None", judge_model: str) -> None: - """Reject a judge model the dispatch path cannot resolve, at start rather than as a - silently growing error count once the job is already sampling and billing.""" - if llm_router is not None and _is_configured_pre_routing_strategy(llm_router, judge_model): +def _validate_plain_model(llm_router: "Router | None", model: str, field_name: str) -> None: + """Reject a model the dispatch path cannot resolve, at start rather than as a silently + growing error count once the job is already sampling and billing. Both the judge and a + reverse job's baseline must be plain models: an auto-router in either slot would + re-route per turn, so the comparison would have no fixed arm to attribute results to.""" + if llm_router is not None and _is_configured_pre_routing_strategy(llm_router, model): raise HTTPException( status_code=400, - detail=f"judge_model '{judge_model}' is an auto-router; the judge must be a plain model", + detail=f"{field_name} '{model}' is an auto-router; it must be a plain model", ) - if router_resolves_model(llm_router, judge_model): + if router_resolves_model(llm_router, model): return import litellm try: - litellm.get_llm_provider(model=judge_model) + litellm.get_llm_provider(model=model) except Exception as e: raise HTTPException( status_code=400, detail=( - f"judge_model '{judge_model}' is neither a model configured on this proxy nor a " + f"{field_name} '{model}' is neither a model configured on this proxy nor a " "provider-qualified public model name (e.g. 'anthropic/claude-sonnet-5')" ), ) from e def _is_unique_violation(error: Exception) -> bool: - """Whether a Prisma create failed on a unique index. One active job per key lives in - a partial unique index (raw SQL in the migration; schema.prisma cannot express partial - indexes), so the read-then-create check above it is advisory: two concurrent starts - pass the read, and the loser must surface as the same 409 rather than a 500.""" + """Whether a Prisma create failed on a unique index. One active job per key and + direction lives in a partial unique index (raw SQL in the migration; schema.prisma + cannot express partial indexes), so the read-then-create check above it is advisory: + two concurrent starts pass the read, and the loser must surface as the same 409 + rather than a 500.""" try: from prisma.errors import UniqueViolationError except ImportError: @@ -573,8 +576,10 @@ def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]: async def _shadow_eval_results(prisma_client: "PrismaClient", job_id: str) -> ShadowEvalResult | None: """Both stratifications of one job's verdicts. Tier answers "where does the router do - well"; current-model answers "which of the models this key uses today would the router - beat". Reads are bounded by the job's own attempts (<= max_turns) via the job_id index.""" + well"; the model stratification groups by whichever model served the real arm, so it + answers "which of the models this key uses today would the router beat" forward, and + "for the turns the router sent to X, did X beat the baseline" in reverse. Reads are + bounded by the job's own attempts (<= max_turns) via the job_id index.""" by_tier: Final = _ATTEMPT_AGG_ROWS.validate_python( await prisma_client.db.query_raw(_ATTEMPT_AGG_BY_TIER_SQL, job_id) or () ) @@ -604,9 +609,15 @@ async def start_shadow_eval( user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], ) -> ShadowEvalJobResponse: """ - Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic - through an auto-router, judge real vs. shadow responses blind, and stratify win rates - by the router's tier classification and by the incumbent model. + Start a shadow eval: duplicate a sampled slice of a key's live traffic against a second + arm, judge the two responses blind, and stratify win rates by tier and by the model that + served the real arm. + + A forward job answers whether the key should adopt router_name: it samples the requests + the router did not serve and duplicates them through it. A reverse job answers whether a + key already on the router still gains from it: it samples the requests the router did + serve and duplicates them against baseline_model. A key can hold one active job per + direction, so both questions can run at once. Shadow responses are never served to users. The job samples until it has judged max_turns turns, reaches the end of its window, or is stopped; sampling changes @@ -620,7 +631,9 @@ async def start_shadow_eval( raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, data.router_name): raise HTTPException(status_code=400, detail=f"'{data.router_name}' is not a configured auto-router") - _validate_judge_model(llm_router, data.judge_model) + _validate_plain_model(llm_router, data.judge_model, "judge_model") + if data.baseline_model is not None: + _validate_plain_model(llm_router, data.baseline_model, "baseline_model") key_row: Final = await prisma_client.db.litellm_verificationtoken.find_unique( where={"token": data.api_key_id} # mutable-ok: Prisma filter ) @@ -634,16 +647,20 @@ async def start_shadow_eval( ) # A job that expired or exhausted its turn budget stopped sampling on its own, but - # still holds the one-active-per-key partial unique index until stamped; free it so - # a new eval can start. + # still holds its slot in the per-key, per-direction partial unique index until + # stamped; free it so a new eval can start. Sweeping both directions is deliberate. await prisma_client.db.execute_raw(_SWEEP_FINISHED_JOBS_SQL, data.api_key_id) active: Final = await prisma_client.db.litellm_shadowevaljob.find_first( - where={"api_key_id": data.api_key_id, "stopped_at": None}, # mutable-ok: Prisma filter + where={ # mutable-ok: Prisma filter + "api_key_id": data.api_key_id, + "direction": data.direction, + "stopped_at": None, + }, ) if active is not None: raise HTTPException( status_code=409, - detail=f"Key already has an active shadow eval job ({active.id}). Stop it first.", + detail=f"Key already has an active {data.direction} shadow eval job ({active.id}). Stop it first.", ) now: Final = datetime.now(timezone.utc) try: @@ -651,6 +668,8 @@ async def start_shadow_eval( data={ # mutable-ok: Prisma payload "api_key_id": data.api_key_id, "router_name": data.router_name, + "direction": data.direction, + "baseline_model": data.baseline_model, "judge_model": data.judge_model, "shadow_percentage": data.shadow_percentage, "max_turns": data.max_turns, @@ -663,7 +682,9 @@ async def start_shadow_eval( raise raise HTTPException( status_code=409, - detail="Key already has an active shadow eval job (started concurrently). Stop it first.", + detail=( + f"Key already has an active {data.direction} shadow eval job (started concurrently). Stop it first." + ), ) from e return ShadowEvalJobResponse.model_validate(job, from_attributes=True) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 79d778fb464..71345d2ccde 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } -// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic. -// A sampled slice of requests is duplicated through the router in a detached task and an -// LLM judge compares real vs shadow responses blind. The job row is immutable config plus +// Shadow eval: evaluation of an auto-router against a key's live traffic, in either +// direction. forward duplicates the requests the key did not route through the router +// through it, answering whether the key should adopt it; reverse duplicates the requests +// the router did serve against a fixed baseline model, answering whether a key already on +// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge +// compares real vs shadow responses blind. The job row is immutable config plus // stopped_at; every count, status, and spend figure is derived from the append-only // attempt rows, so nothing can disagree across pods or stop races. model LiteLLM_ShadowEvalJob { id String @id @default(cuid()) api_key_id String // hashed virtual key whose traffic is shadowed - router_name String + router_name String // the auto-router under evaluation, in either direction + direction String @default("forward") // forward | reverse + baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float max_turns Int // sample budget: judge at most this many turns diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index bf8a3d34098..1b0c7476fc3 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -6,7 +6,7 @@ from collections.abc import Mapping from datetime import datetime, timezone from typing import Final, Literal, TypeAlias -from pydantic import AliasChoices, BaseModel, ConfigDict, Field, computed_field, field_validator +from pydantic import AliasChoices, BaseModel, ConfigDict, Field, computed_field, field_validator, model_validator from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig from litellm.types.utils import StandardLoggingRoutingDecision @@ -146,11 +146,13 @@ class AutoRouterBenchmarksResponse(BaseModel): ShadowEvalStatus: TypeAlias = Literal["running", "completed", "stopped"] +ShadowEvalDirection: TypeAlias = Literal["forward", "reverse"] + DEFAULT_SHADOW_EVAL_JUDGE_MODEL: Final[str] = "anthropic/claude-sonnet-5" class StartShadowEvalRequest(BaseModel): - """Start shadowing a key's traffic through an auto-router for blind comparison.""" + """Start duplicating a key's traffic for blind comparison against an auto-router.""" api_key_id: str = Field( description=( @@ -158,7 +160,23 @@ class StartShadowEvalRequest(BaseModel): "key's traffic; requests made with any other key are not sampled." ) ) - router_name: str = Field(description="The auto-router config to shadow requests through") + router_name: str = Field(description="The auto-router under evaluation, in either direction") + direction: ShadowEvalDirection = Field( + default="forward", + description=( + "forward answers 'should this key adopt router_name': it samples the requests the key did NOT " + "route through the router and duplicates them through it. reverse answers 'is the router still " + "worth it for a key already on it': it samples the requests the router did serve and duplicates " + "them against baseline_model. The response the caller received is always the real arm" + ), + ) + baseline_model: str | None = Field( + default=None, + description=( + "Required when direction is reverse and rejected otherwise: the fixed model the router's own " + "responses are judged against. Must be a plain model rather than another auto-router" + ), + ) shadow_percentage: float = Field( ge=0.1, le=100.0, @@ -193,15 +211,33 @@ class StartShadowEvalRequest(BaseModel): def _round_percentage(cls, value: float) -> float: return round(value, 2) + @model_validator(mode="after") + def _baseline_model_matches_direction(self) -> "StartShadowEvalRequest": + if self.direction == "reverse" and self.baseline_model is None: + raise ValueError("baseline_model is required when direction is 'reverse'") + if self.direction == "forward" and self.baseline_model is not None: + raise ValueError("baseline_model is only meaningful when direction is 'reverse'") + return self + class ShadowEvalSlice(BaseModel): """Judge outcomes for one slice of a job's verdicts (a router tier, or one of the - models the shadowed key currently uses).""" + models that served the real arm).""" group: str turn_count: int - real_win_rate_pct: float = Field(description="Share of judged turns where the real (control) model won") - shadow_win_rate_pct: float = Field(description="Share of judged turns where the shadowed router's pick won") + real_win_rate_pct: float = Field( + description=( + "Share of judged turns the real arm won, meaning the response the caller actually received: " + "the key's own model in forward mode, the router's pick in reverse" + ) + ) + shadow_win_rate_pct: float = Field( + description=( + "Share of judged turns the shadow arm won, meaning the duplicated response nobody was served: " + "the router's pick in forward mode, baseline_model in reverse" + ) + ) tie_rate_pct: float avg_judge_confidence: float @@ -210,7 +246,12 @@ class ShadowEvalResult(BaseModel): """Stratified results of a shadow-eval job's verdicts so far.""" by_tier: tuple[ShadowEvalSlice, ...] - by_current_model: tuple[ShadowEvalSlice, ...] + by_current_model: tuple[ShadowEvalSlice, ...] = Field( + description=( + "Sliced by the model that served the real arm: the key's incumbent models in forward mode, " + "and in reverse the models the router itself picked" + ) + ) overall_shadow_win_rate_pct: float overall_tie_rate_pct: float @@ -226,6 +267,8 @@ class ShadowEvalJobResponse(BaseModel): job_id: str = Field(validation_alias=AliasChoices("id", "job_id")) api_key_id: str = Field(description="The hashed virtual key whose traffic this job evaluates, and only that key's") router_name: str + direction: ShadowEvalDirection = "forward" + baseline_model: str | None = None judge_model: str shadow_percentage: float max_turns: int diff --git a/schema.prisma b/schema.prisma index 79d778fb464..71345d2ccde 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession { @@index([last_turn_at], map: "idx_autorouter_session_last_turn") } -// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic. -// A sampled slice of requests is duplicated through the router in a detached task and an -// LLM judge compares real vs shadow responses blind. The job row is immutable config plus +// Shadow eval: evaluation of an auto-router against a key's live traffic, in either +// direction. forward duplicates the requests the key did not route through the router +// through it, answering whether the key should adopt it; reverse duplicates the requests +// the router did serve against a fixed baseline model, answering whether a key already on +// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge +// compares real vs shadow responses blind. The job row is immutable config plus // stopped_at; every count, status, and spend figure is derived from the append-only // attempt rows, so nothing can disagree across pods or stop races. model LiteLLM_ShadowEvalJob { id String @id @default(cuid()) api_key_id String // hashed virtual key whose traffic is shadowed - router_name String + router_name String // the auto-router under evaluation, in either direction + direction String @default("forward") // forward | reverse + baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float max_turns Int // sample budget: judge at most this many turns diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index e1c56db21af..3a69340109d 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -6,6 +6,7 @@ from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock import pytest +from pydantic import ValidationError from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY @@ -19,7 +20,7 @@ from litellm.integrations.shadow_eval_logger import ( _sample_hits, _unmask_preference, ) -from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN +from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN, ModelResponse def _job(**overrides) -> ActiveShadowEvalJob: @@ -51,6 +52,8 @@ def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock: id=job.id, api_key_id=api_key_id, router_name=job.router_name, + direction=job.direction, + baseline_model=job.baseline_model, shadow_percentage=job.shadow_percentage, judge_model=job.judge_model, max_turns=job.max_turns, @@ -61,40 +64,55 @@ def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock: def _router(shadow_text="shadow answer", judge_json='{"preference": "A", "confidence": 0.9, "reasoning": "x"}'): - """One mock router serving the shadow call first, the judge call second. The shadow - call's metadata receives the routing decision write-back, like the real router.""" + """One mock router serving the shadow call first, the judge call second, told apart by + the internal-origin stamp rather than the model, since a reverse job's shadow arm names + a plain model. Only the auto-router writes a routing decision back, and only a plain + model reports the model it served on the response, which is how each direction learns + which model answered.""" router = MagicMock() router.model_group_alias = {} router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}]) async def acompletion(**kwargs): + if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN: + return {"choices": [{"message": {"content": judge_json}}]} if kwargs["model"] == "my-router": kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"} return {"choices": [{"message": {"content": shadow_text}}], "usage": {"completion_tokens": 5}} - return {"choices": [{"message": {"content": judge_json}}]} + return ModelResponse( + model=kwargs["model"], + choices=[{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": shadow_text}}], + ) router.acompletion = MagicMock(side_effect=acompletion) return router -def _logger(router=None, prisma=None, job=None) -> ShadowEvalLogger: +def _logger(router=None, prisma=None, jobs=()) -> ShadowEvalLogger: cache = InMemoryCache(max_size_in_memory=4, default_ttl=60) logger = ShadowEvalLogger( router_provider=lambda: router, prisma_provider=lambda: prisma, jobs_cache=cache, ) - if job is not None: - cache.set_cache("shadow_eval:active_jobs", {"key-hash": job}) + if jobs: + cache.set_cache("shadow_eval:active_jobs", {"key-hash": tuple(jobs)}) return logger -def _success_kwargs(request_id="req-1", api_key_hash="key-hash", request_metadata=None, call_type="acompletion"): +def _routed_by(router_name="my-router", tier="COMPLEX"): + """Metadata as a pre-routing strategy leaves it on the request it served.""" + return {"routing_decision": {"router_model_name": router_name, "tier_label": tier, "routed_model": "router-pick"}} + + +def _success_kwargs( + request_id="req-1", api_key_hash="key-hash", request_metadata=None, call_type="acompletion", model="claude-opus" +): return { "standard_logging_object": { "id": request_id, "call_type": call_type, - "model": "claude-opus", + "model": model, "metadata": {"user_api_key_hash": api_key_hash}, "model_parameters": {"temperature": 0.5, "stream": True}, }, @@ -164,7 +182,7 @@ class TestSuccessHookSkipChain: monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005) prisma = _prisma() router = _router() - logger = _logger(router=router, prisma=prisma, job=_job()) + logger = _logger(router=router, prisma=prisma, jobs=(_job(),)) await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None) await _drain(logger) @@ -209,7 +227,7 @@ class TestSuccessHookSkipChain: async def test_skip_paths_store_nothing(self, kwargs_mutation, job_mutation): starts = job_mutation.pop("_starts", 0) prisma = _prisma() - logger = _logger(router=_router(), prisma=prisma, job=_job(**job_mutation)) + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(**job_mutation),)) logger._job_starts = {"job-1": starts} await logger.async_log_success_event(_success_kwargs(**kwargs_mutation), RESPONSE, None, None) @@ -222,7 +240,7 @@ class TestSuccessHookSkipChain: """A finished pipeline frees its concurrency slot but not its slice of the turn budget; the budget only reopens when a cache refill absorbs the written rows.""" prisma = _prisma() - logger = _logger(router=_router(), prisma=prisma, job=_job(attempts=199, max_turns=200)) + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(attempts=199, max_turns=200),)) await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None) await _drain(logger) @@ -237,7 +255,7 @@ class TestSuccessHookSkipChain: identity to the shadow and judge calls.""" prisma = _prisma() router = _router() - logger = _logger(router=router, prisma=prisma, job=_job()) + logger = _logger(router=router, prisma=prisma, jobs=(_job(),)) hook_kwargs = _success_kwargs() hook_kwargs["litellm_params"] = { @@ -256,7 +274,7 @@ class TestSuccessHookSkipChain: predicate, so every redaction source counts.""" prisma = _prisma() router = _router() - logger = _logger(router=router, prisma=prisma, job=_job()) + logger = _logger(router=router, prisma=prisma, jobs=(_job(),)) hook_kwargs = _success_kwargs() hook_kwargs["standard_callback_dynamic_params"] = {"turn_off_message_logging": True} @@ -268,7 +286,7 @@ class TestSuccessHookSkipChain: async def test_inflight_cap_sheds_instead_of_queueing(self): prisma = _prisma() - logger = _logger(router=_router(), prisma=prisma, job=_job()) + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),)) logger._inflight_shadow_tasks = _MAX_CONCURRENT_SHADOW_TASKS await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None) @@ -291,8 +309,8 @@ class TestActiveJobsCache: first = await logger._active_jobs() second = await logger._active_jobs() - assert first["key-hash"].id == "job-1" - assert second["key-hash"].attempts == 7 + assert [job.id for job in first["key-hash"]] == ["job-1"] + assert second["key-hash"][0].attempts == 7 assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1 where = prisma.db.litellm_shadowevaljob.find_many.call_args.kwargs["where"] assert where["stopped_at"] is None @@ -353,6 +371,7 @@ class TestShadowPipeline: messages=({"role": "user", "content": "hi"},), response_obj=RESPONSE, real_model="claude-opus", + control_tier=None, model_parameters={}, parent_metadata={}, ) @@ -381,6 +400,7 @@ class TestShadowPipeline: messages=({"role": "user", "content": "hi"},), response_obj=RESPONSE, real_model="claude-opus", + control_tier=None, model_parameters={}, parent_metadata={"user_api_key_auth": UserAPIKeyAuth(api_key="sk-abc", max_budget=10.0)}, ) @@ -411,6 +431,7 @@ class TestShadowPipeline: messages=({"role": "user", "content": "hi"},), response_obj=RESPONSE, real_model="claude-opus", + control_tier=None, model_parameters={}, parent_metadata={}, ) @@ -438,6 +459,7 @@ class TestShadowPipeline: messages=({"role": "user", "content": "hi"},), response_obj=RESPONSE, real_model="claude-opus", + control_tier=None, model_parameters={"stream": True, "temperature": 0.2, "metadata": {"x": 1}}, parent_metadata=parent_metadata, ) @@ -458,6 +480,164 @@ class TestShadowPipeline: assert judge_call["max_tokens"] == JUDGE_MAX_OUTPUT_TOKENS +def _reverse_job(**overrides) -> ActiveShadowEvalJob: + return _job(**{"direction": "reverse", "baseline_model": "baseline-model", **overrides}) + + +class TestJobValidation: + @pytest.mark.parametrize( + "overrides", + [ + {"direction": "reverse"}, + {"baseline_model": "baseline-model"}, + {"direction": "sideways", "baseline_model": "baseline-model"}, + ], + ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction"], + ) + def test_unsamplable_shapes_are_rejected(self, overrides): + with pytest.raises(ValidationError): + _job(**overrides) + + def test_shadow_target_follows_direction(self): + assert _job().shadow_target == "my-router" + assert _reverse_job().shadow_target == "baseline-model" + + +@pytest.mark.asyncio +class TestDirection: + @pytest.mark.parametrize( + "job,routed_by,sampled", + [ + (_job(), None, True), + (_job(), "my-router", False), + (_job(), "other-router", True), + (_reverse_job(), "my-router", True), + (_reverse_job(), None, False), + (_reverse_job(), "other-router", False), + ], + ids=[ + "forward-samples-unrouted", + "forward-skips-its-own-router", + "forward-samples-another-router", + "reverse-samples-its-own-router", + "reverse-skips-unrouted", + "reverse-skips-another-router", + ], + ) + async def test_direction_decides_which_traffic_is_sampled(self, job, routed_by, sampled): + """The two directions partition the key's traffic: whatever one samples, the other + skips, so a key running both never judges the same turn twice for the same reason.""" + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma, jobs=(job,)) + + await logger.async_log_success_event( + _success_kwargs(request_metadata=_routed_by(routed_by) if routed_by else {}), RESPONSE, None, None + ) + await _drain(logger) + + assert prisma.db.litellm_shadowevalattempt.create.await_count == int(sampled) + + async def test_reverse_duplicates_against_the_baseline_model(self): + prisma = _prisma() + router = _router() + logger = _logger(router=router, prisma=prisma, jobs=(_reverse_job(),)) + + await logger.async_log_success_event( + _success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None + ) + await _drain(logger) + + assert router.acompletion.call_args_list[0].kwargs["model"] == "baseline-model" + + async def test_reverse_row_orients_arms_and_reads_tier_off_the_served_request(self): + """real is what the caller received, so in reverse it is the router's own pick and + the tier that produced it; only the shadow arm moves to the baseline.""" + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma, jobs=(_reverse_job(),)) + + await logger.async_log_success_event( + _success_kwargs(request_metadata=_routed_by(tier="COMPLEX"), model="router-pick"), RESPONSE, None, None + ) + await _drain(logger) + + row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + assert row["real_model"] == "router-pick" + assert row["shadow_model"] == "baseline-model" + assert row["tier"] == "COMPLEX" + + async def test_forward_row_still_reads_tier_off_the_shadow_call(self): + """A forward job's tier describes the arm being evaluated, which is the shadow one, + so a routing decision on the incumbent request must not leak into it.""" + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),)) + + await logger.async_log_success_event( + _success_kwargs(request_metadata=_routed_by("other-router", tier="CONTROL_TIER")), RESPONSE, None, None + ) + await _drain(logger) + + row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + assert row["tier"] == "SIMPLE" + assert row["shadow_model"] == "cheap-model" + + async def test_a_key_running_both_directions_dispatches_both(self): + """One request can qualify for a forward job on a router that did not serve it and a + reverse job on the router that did. The two are separately budgeted experiments, so + both fire rather than one silently losing the turn.""" + prisma = _prisma() + logger = _logger( + router=_router(), + prisma=prisma, + jobs=(_job(id="forward-job", router_name="other-router"), _reverse_job(id="reverse-job")), + ) + + await logger.async_log_success_event( + _success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None + ) + await _drain(logger) + + rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.call_args_list] + assert sorted(row["job_id"] for row in rows) == ["forward-job", "reverse-job"] + assert logger._job_starts == {"forward-job": 1, "reverse-job": 1} + + +@pytest.mark.asyncio +class TestActiveJobsFailClosed: + async def test_a_row_the_sampler_cannot_read_is_dropped_not_guessed(self): + """A reverse row with no baseline model has no second arm to call, so it is skipped + rather than silently dispatched at the router it is supposed to be judging.""" + broken = _job_record(_job(id="job-broken")) + broken.direction = "reverse" + broken.baseline_model = None + prisma = _prisma(jobs=[broken, _job_record(_job(id="job-ok"))], attempt_counts=[("job-ok", 1)]) + logger = ShadowEvalLogger( + router_provider=lambda: None, + prisma_provider=lambda: prisma, + jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60), + ) + + assert [job.id for job in (await logger._active_jobs())["key-hash"]] == ["job-ok"] + + async def test_both_of_a_key_s_jobs_survive_the_lookup(self): + records = [ + _job_record(_job(id="job-forward")), + _job_record(_reverse_job(id="job-reverse")), + _job_record(_job(id="job-other"), api_key_id="other-key"), + ] + prisma = _prisma(jobs=records, attempt_counts=[("job-reverse", 3)]) + logger = ShadowEvalLogger( + router_provider=lambda: None, + prisma_provider=lambda: prisma, + jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60), + ) + + jobs = await logger._active_jobs() + + assert sorted(job.id for job in jobs["key-hash"]) == ["job-forward", "job-reverse"] + assert [job.id for job in jobs["other-key"]] == ["job-other"] + assert {job.id: job.attempts for job in jobs["key-hash"]}["job-reverse"] == 3 + + def _failing_router(): router = MagicMock() router.model_group_alias = {} diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index 16e82bc3bda..dbde7c461b8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -559,7 +559,7 @@ def _start_request(**overrides: object) -> StartShadowEvalRequest: @pytest.mark.asyncio async def test_start_shadow_eval_creates_job_and_frees_expired_or_exhausted_ones(monkeypatch: pytest.MonkeyPatch): """Expiry and turn-budget exhaustion both end sampling on their own; either must - release the one-active-per-key index so a new eval can start.""" + release the key's slot in the active-job index so a new eval can start.""" import litellm.proxy.proxy_server as proxy_server prisma = _shadow_prisma() @@ -592,8 +592,21 @@ async def test_start_shadow_eval_creates_job_and_frees_expired_or_exhausted_ones (ADMIN, {"judge_model": "not/a real model!"}, None, 400), (ADMIN, {"judge_model": "my-router"}, None, 400), (ADMIN, {}, "active", 409), + (ADMIN, {"direction": "reverse", "baseline_model": "my-router"}, None, 400), + (ADMIN, {"direction": "reverse", "baseline_model": "not/a real model!"}, None, 400), + (ADMIN, {"direction": "reverse", "baseline_model": "openai/gpt-4o", "router_name": "not-a-router"}, None, 400), + ], + ids=[ + "non-admin", + "view-only", + "unknown-router", + "unresolvable-judge", + "router-as-judge", + "already-active", + "router-as-baseline", + "unresolvable-baseline", + "reverse-still-needs-an-auto-router", ], - ids=["non-admin", "view-only", "unknown-router", "unresolvable-judge", "router-as-judge", "already-active"], ) async def test_start_shadow_eval_rejections( monkeypatch: pytest.MonkeyPatch, caller, request_overrides, active, expected_status @@ -609,6 +622,68 @@ async def test_start_shadow_eval_rejections( assert exc.value.status_code == expected_status +@pytest.mark.parametrize( + "overrides", + [ + {"direction": "reverse"}, + {"baseline_model": "openai/gpt-4o"}, + {"direction": "sideways", "baseline_model": "openai/gpt-4o"}, + ], + ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction"], +) +def test_start_request_pins_baseline_model_to_reverse(overrides): + """A forward job has no second arm to name and a reverse job cannot run without one, + so neither shape reaches the endpoint to be half-validated there.""" + with pytest.raises(ValidationError): + _start_request(**overrides) + + +@pytest.mark.asyncio +async def test_start_shadow_eval_reverse_records_its_arms_and_holds_its_own_slot(monkeypatch: pytest.MonkeyPatch): + """The two directions ask opposite questions of the same key, so a forward job holding + the slot must not block a reverse one. The second reverse start still 409s.""" + import litellm.proxy.proxy_server as proxy_server + + prisma = _shadow_prisma() + active = {"forward": _job_record()} + prisma.db.litellm_shadowevaljob.find_first = AsyncMock( + side_effect=lambda where, **_: active.get(str(where.get("direction"))) + ) + prisma.db.litellm_shadowevaljob.create = AsyncMock( + return_value=_job_record(direction="reverse", baseline_model="openai/gpt-4o") + ) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + monkeypatch.setattr(proxy_server, "llm_router", _shadow_router()) + + reverse = _start_request(direction="reverse", baseline_model="openai/gpt-4o") + response = await start_shadow_eval(reverse, ADMIN) + + assert (response.direction, response.baseline_model) == ("reverse", "openai/gpt-4o") + create_data = prisma.db.litellm_shadowevaljob.create.call_args.kwargs["data"] + assert create_data["direction"] == "reverse" + assert create_data["baseline_model"] == "openai/gpt-4o" + + active["reverse"] = _job_record(id="job-2", direction="reverse") + with pytest.raises(HTTPException) as exc: + await start_shadow_eval(reverse, ADMIN) + assert exc.value.status_code == 409 + + +@pytest.mark.asyncio +async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkeypatch: pytest.MonkeyPatch): + import litellm.proxy.proxy_server as proxy_server + + prisma = _shadow_prisma() + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + monkeypatch.setattr(proxy_server, "llm_router", _shadow_router()) + + await start_shadow_eval(_start_request(), ADMIN) + + create_data = prisma.db.litellm_shadowevaljob.create.call_args.kwargs["data"] + assert create_data["direction"] == "forward" + assert create_data["baseline_model"] is None + + @pytest.mark.asyncio async def test_start_shadow_eval_rejects_a_key_this_proxy_does_not_know(monkeypatch: pytest.MonkeyPatch): """A typo'd api_key_id would otherwise create a job no traffic can ever match.""" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx index 467439122dd..d4d26650086 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx @@ -358,6 +358,7 @@ describe("ShadowEvalSection", () => { const expectedBody = { api_key_id: "hash-alpha", router_name: "gpt-auto", + direction: "forward", shadow_percentage: 10, duration_days: 7, max_turns: 200, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx index 6bb00933218..711fc1af539 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx @@ -308,6 +308,7 @@ const StartForm: React.FC = () => { const startBody = { api_key_id: apiKeyId, router_name: routerName, + direction: "forward" as const, shadow_percentage: parsedPct, duration_days: Number.parseInt(durationDays, 10), max_turns: parsedMaxTurns, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 2cbc7fd6220..eeea16f3ccd 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -838,9 +838,15 @@ export interface paths { put?: never; /** * Start Shadow Eval - * @description Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic - * through an auto-router, judge real vs. shadow responses blind, and stratify win rates - * by the router's tier classification and by the incumbent model. + * @description Start a shadow eval: duplicate a sampled slice of a key's live traffic against a second + * arm, judge the two responses blind, and stratify win rates by tier and by the model that + * served the real arm. + * + * A forward job answers whether the key should adopt router_name: it samples the requests + * the router did not serve and duplicates them through it. A reverse job answers whether a + * key already on the router still gains from it: it samples the requests the router did + * serve and duplicates them against baseline_model. A key can hold one active job per + * direction, so both questions can run at once. * * Shadow responses are never served to users. The job samples until it has judged * max_turns turns, reaches the end of its window, or is stopped; sampling changes @@ -32737,11 +32743,19 @@ export interface components { * @description The hashed virtual key whose traffic this job evaluates, and only that key's */ api_key_id: string; + /** Baseline Model */ + baseline_model?: string | null; /** * Created At * Format: date-time */ created_at: string; + /** + * Direction + * @default forward + * @enum {string} + */ + direction: "forward" | "reverse"; /** * Ends At * Format: date-time @@ -32794,7 +32808,10 @@ export interface components { * @description Stratified results of a shadow-eval job's verdicts so far. */ ShadowEvalResult: { - /** By Current Model */ + /** + * By Current Model + * @description Sliced by the model that served the real arm: the key's incumbent models in forward mode, and in reverse the models the router itself picked + */ by_current_model: components["schemas"]["ShadowEvalSlice"][]; /** By Tier */ by_tier: components["schemas"]["ShadowEvalSlice"][]; @@ -32806,7 +32823,7 @@ export interface components { /** * ShadowEvalSlice * @description Judge outcomes for one slice of a job's verdicts (a router tier, or one of the - * models the shadowed key currently uses). + * models that served the real arm). */ ShadowEvalSlice: { /** Avg Judge Confidence */ @@ -32815,12 +32832,12 @@ export interface components { group: string; /** * Real Win Rate Pct - * @description Share of judged turns where the real (control) model won + * @description Share of judged turns the real arm won, meaning the response the caller actually received: the key's own model in forward mode, the router's pick in reverse */ real_win_rate_pct: number; /** * Shadow Win Rate Pct - * @description Share of judged turns where the shadowed router's pick won + * @description Share of judged turns the shadow arm won, meaning the duplicated response nobody was served: the router's pick in forward mode, baseline_model in reverse */ shadow_win_rate_pct: number; /** Tie Rate Pct */ @@ -33003,7 +33020,7 @@ export interface components { }; /** * StartShadowEvalRequest - * @description Start shadowing a key's traffic through an auto-router for blind comparison. + * @description Start duplicating a key's traffic for blind comparison against an auto-router. */ StartShadowEvalRequest: { /** @@ -33011,6 +33028,18 @@ export interface components { * @description The hashed virtual key whose traffic will be shadowed. Shadow evaluation runs ONLY on this key's traffic; requests made with any other key are not sampled. */ api_key_id: string; + /** + * Baseline Model + * @description Required when direction is reverse and rejected otherwise: the fixed model the router's own responses are judged against. Must be a plain model rather than another auto-router + */ + baseline_model?: string | null; + /** + * Direction + * @description forward answers 'should this key adopt router_name': it samples the requests the key did NOT route through the router and duplicates them through it. reverse answers 'is the router still worth it for a key already on it': it samples the requests the router did serve and duplicates them against baseline_model. The response the caller received is always the real arm + * @default forward + * @enum {string} + */ + direction: "forward" | "reverse"; /** * Duration Days * @description How many days the job samples traffic before completing on its own @@ -33031,7 +33060,7 @@ export interface components { max_turns: number; /** * Router Name - * @description The auto-router config to shadow requests through + * @description The auto-router under evaluation, in either direction */ router_name: string; /** From f99d0a4b389c6142977c21f4d7e5d9bf9a051c8f Mon Sep 17 00:00:00 2001 From: Ilan Chemla Date: Sat, 15 Aug 2026 03:09:58 +0300 Subject: [PATCH 26/27] feat(search): add Nimble as a search provider (#36347) * feat(search): add Nimble as a search provider Adds `NimbleSearchConfig` so `search_provider: nimble` works across the SDK, the proxy /v1/search endpoint, the Search Tools dashboard, and spend tracking. Nimble's /v2/search already uses the Perplexity unified spec's parameter names, so the request transform is close to a pass-through. `search_domain_filter` splits into include_domains/exclude_domains on the spec's `-` prefix, `country` is upper-cased to the ISO form Nimble documents, and everything else is forwarded so focus, search_depth, time_range and the rest stay reachable. On the response side, snippet prefers `content` and falls back to `description`, and a malformed body raises an attributed error rather than reporting an empty search. Also tightens `BaseSearchConfig.get_supported_perplexity_optional_params` to return `frozenset[str]` instead of a bare mutable `set`, which every caller already treats as read-only. * fix(search): surface Nimble error bodies instead of empty results Greptile flagged that a null or absent `results` degraded to a successful empty search. A search with no hits comes back as `"results": []`, verified against the live API, so the field is now required and anything else raises the attributed schema error the other malformed bodies already take. Also unwraps Nimble's second error envelope. Collection failures return `{"success", "task_id", "message"}` rather than the `{"detail"}` shape validation errors use, and only the latter was being read. Drops comments that restated the adjacent code. * docs(search): drop the Nimble param list from the transform docstring It restated the vendor's API reference, which the module docstring already links, and would go stale the moment Nimble adds a focus mode. --- .../llms/base_llm/search/transformation.py | 19 +- litellm/llms/nimble/__init__.py | 3 + litellm/llms/nimble/search/__init__.py | 3 + litellm/llms/nimble/search/transformation.py | 264 ++++++++++++++++++ ...odel_prices_and_context_window_backup.json | 8 + litellm/types/utils.py | 1 + litellm/utils.py | 2 + model_prices_and_context_window.json | 8 + provider_endpoints_support.json | 7 + .../enforce_llms_folder_style.py | 1 + tests/search_tests/test_nimble_search.py | 155 ++++++++++ .../search/test_base_search_transformation.py | 3 + .../test_nimble_search_transformation.py | 251 +++++++++++++++++ .../public/assets/logos/nimble.png | Bin 0 -> 6579 bytes .../_components/CreateSearchTools.tsx | 2 + 15 files changed, 720 insertions(+), 7 deletions(-) create mode 100644 litellm/llms/nimble/__init__.py create mode 100644 litellm/llms/nimble/search/__init__.py create mode 100644 litellm/llms/nimble/search/transformation.py create mode 100644 tests/search_tests/test_nimble_search.py create mode 100644 tests/test_litellm/llms/nimble/search/test_nimble_search_transformation.py create mode 100644 ui/litellm-dashboard/public/assets/logos/nimble.png diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index 6987e261d4e..dee67e0b100 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -18,6 +18,16 @@ else: LiteLLMLoggingObj = Any +_PERPLEXITY_UNIFIED_PARAMS: Final[frozenset[str]] = frozenset( + ( + "max_results", + "search_domain_filter", + "country", + "max_tokens_per_page", + ) +) + + def _search_host(url: str) -> str: return urlsplit(url).netloc.lower() @@ -96,7 +106,7 @@ class BaseSearchConfig: return "POST" @staticmethod - def get_supported_perplexity_optional_params() -> set: + def get_supported_perplexity_optional_params() -> frozenset[str]: """ Get the set of Perplexity unified search parameters. These are the standard parameters that providers should transform from. @@ -104,12 +114,7 @@ class BaseSearchConfig: Returns: Set of parameter names that are part of the unified spec """ - return { - "max_results", - "search_domain_filter", - "country", - "max_tokens_per_page", - } + return _PERPLEXITY_UNIFIED_PARAMS def _assert_trusted_api_base_for_server_credential( self, diff --git a/litellm/llms/nimble/__init__.py b/litellm/llms/nimble/__init__.py new file mode 100644 index 00000000000..05272cb1230 --- /dev/null +++ b/litellm/llms/nimble/__init__.py @@ -0,0 +1,3 @@ +from litellm.llms.nimble.search.transformation import NimbleSearchConfig + +__all__ = ("NimbleSearchConfig",) diff --git a/litellm/llms/nimble/search/__init__.py b/litellm/llms/nimble/search/__init__.py new file mode 100644 index 00000000000..05272cb1230 --- /dev/null +++ b/litellm/llms/nimble/search/__init__.py @@ -0,0 +1,3 @@ +from litellm.llms.nimble.search.transformation import NimbleSearchConfig + +__all__ = ("NimbleSearchConfig",) diff --git a/litellm/llms/nimble/search/transformation.py b/litellm/llms/nimble/search/transformation.py new file mode 100644 index 00000000000..7485686d230 --- /dev/null +++ b/litellm/llms/nimble/search/transformation.py @@ -0,0 +1,264 @@ +""" +Calls Nimble's /v2/search endpoint to search the web. + +Nimble API Reference: https://docs.nimbleway.com/api-reference/search/search +""" + +from __future__ import annotations + +from collections.abc import Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +import httpx +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, + SearchResult, +) +from litellm.secret_managers.main import get_secret_str + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +_NIMBLE_DOCS_URL: Final = "https://docs.nimbleway.com/api-reference/search/search" + + +class _NimbleResult(BaseModel): + """One entry of Nimble's `results` array. Every field is optional so a single degraded + result degrades to empty strings instead of failing the whole call.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + title: str | None = None + url: str | None = None + content: str | None = None + description: str | None = None + # Free-form per Nimble's schema, so an unexpected shape must not fail the search. + additional_data: object = None + + +class _NimbleSearchResponse(BaseModel): + """Nimble's /v2/search response envelope.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + # Required: a search with no hits returns `[]`, so a null or absent `results` means the + # body is not a search response and must not be reported as a successful empty search. + results: tuple[_NimbleResult, ...] + + +class _AdditionalData(BaseModel): + """The slice of a result's free-form `additional_data` that maps onto SearchResult.""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + publish_date: str | None = None + + +class _ErrorEnvelope(BaseModel): + """Nimble reports errors as either `{"detail": ...}` (validation) or + `{"success": "false", "task_id": ..., "message": ...}` (collection).""" + + model_config = ConfigDict(extra="ignore", frozen=True) + + detail: str | None = None + message: str | None = None + + +_DomainListAdapter: Final = TypeAdapter(tuple[str, ...]) + +_NOTHING: Final[Mapping[str, object]] = MappingProxyType({}) + + +def _optional(key: str, value: object) -> Mapping[str, object]: + """A one-entry mapping to spread into a payload, or nothing when the value is absent.""" + return MappingProxyType({key: value}) if value is not None else _NOTHING + + +class NimbleSearchConfig(BaseSearchConfig): + NIMBLE_API_BASE = "https://sdk.nimbleway.com/v2" + + @staticmethod + def ui_friendly_name() -> str: + return "Nimble" + + def validate_environment( + self, + headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature + api_key: str | None = None, + api_base: str | None = None, + **kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment signature + ) -> dict[str, str]: # mutable-ok: the http handler passes this straight to httpx as headers + """ + Validate environment and return headers. + + Returns a new dict rather than mutating ``headers``: the http handler calls this + a second time after ``litellm/search/main.py`` already did, so it has to be idempotent. + """ + resolved_api_key: Final = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("NIMBLE_API_KEY",), + base_env_var="NIMBLE_API_BASE", + default_api_base=self.NIMBLE_API_BASE, + ) + if not resolved_api_key: + raise ValueError("NIMBLE_API_KEY is not set. Set `NIMBLE_API_KEY` environment variable.") + return { # mutable-ok: httpx requires a plain dict of headers + **headers, + "Authorization": f"Bearer {resolved_api_key}", + "Content-Type": "application/json", + # Nimble's client-attribution header: names the calling software, nothing else. + "X-Client-Source": "litellm", + } + + def get_complete_url( + self, + api_base: str | None, + optional_params: dict[str, object], # mutable-ok: BaseSearchConfig.get_complete_url signature + data: dict[str, object] | list[dict[str, object]] | None = None, # mutable-ok: base signature + **kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url signature + ) -> str: + resolved_base: Final = (api_base or get_secret_str("NIMBLE_API_BASE") or self.NIMBLE_API_BASE).rstrip("/") + if resolved_base.endswith("/search"): + return resolved_base + return f"{resolved_base}/search" + + def transform_search_request( + self, + query: str | list[str], # mutable-ok: BaseSearchConfig.transform_search_request signature + optional_params: dict[str, object], # mutable-ok: base signature + **kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request signature + ) -> dict[str, object]: # mutable-ok: the http handler passes this straight to httpx as the JSON body + """ + Transform Search request to Nimble API format. + + Nimble already uses the Perplexity unified spec's names, so this is close to a pass-through: + - query -> query (a list is joined with spaces; Nimble takes a single string) + - max_results -> max_results (sent unclamped so Nimble's own 1-100 validation reports the error) + - country -> country, upper-cased to the ISO form Nimble documents + - search_domain_filter -> include_domains, with `-`-prefixed entries going to exclude_domains + - max_tokens_per_page -> dropped (no Nimble equivalent) + + Everything else is forwarded as-is, so the rest of Nimble's surface stays reachable + without LiteLLM tracking it. + """ + unified_params: Final = self.get_supported_perplexity_optional_params() + country: Final = optional_params.get("country") + + # Spread after the derived domain filters so an explicitly supplied `include_domains` + # or `exclude_domains` wins over anything read out of `search_domain_filter`. + passthrough: Final = MappingProxyType( + {param: value for param, value in optional_params.items() if param not in unified_params} + ) + + return { # mutable-ok: httpx requires a plain dict for the JSON body + **_domain_filters(optional_params.get("search_domain_filter")), + **passthrough, + "query": " ".join(query) if isinstance(query, list) else query, + **_optional("max_results", optional_params.get("max_results")), + **_optional("country", country.upper() if isinstance(country, str) else None), + } + + def transform_search_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response signature + ) -> SearchResponse: + """ + Transform Nimble API response to LiteLLM unified SearchResponse format. + + `date` carries only the absolute `publish_date`. News results often carry a relative + `publish_date_raw` ("1 day ago") instead, which is not a date, so the whole + `additional_data` object rides through as an extra on `SearchResult` and nothing is lost. + + Nimble ranks results itself via metadata.position, so the order is preserved as received. + A body that does not match the documented schema raises an attributed error rather than + being reported as a successful empty search. Parsing the response bytes rather than + `.json()` covers the non-JSON case through that same path. + """ + try: + parsed: Final = _NimbleSearchResponse.model_validate_json(raw_response.content) + except ValidationError as e: + raise self.get_error_class( + error_message=f"response does not match the documented /v2/search schema: {e}", + status_code=raw_response.status_code, + headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature + ) + + return SearchResponse( + results=[ # mutable-ok: SearchResponse.results is declared list[SearchResult] + SearchResult( + title=result.title or "", + url=result.url or "", + snippet=result.content or result.description or "", + date=_publish_date(result.additional_data), + last_updated=None, + **_optional("additional_data", result.additional_data), + ) + for result in parsed.results + ], + object="search", + ) + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: dict[str, str], # mutable-ok: BaseSearchConfig.get_error_class signature + ) -> Exception: + detail: Final = _unwrap_error_detail(error_message).rstrip(". ") + return BaseLLMException( + status_code=status_code, + message=f"Nimble Search: {detail}. See {_NIMBLE_DOCS_URL} for details.", + headers=headers, + ) + + +def _unwrap_error_detail(error_message: str) -> str: + """ + Surface the human-readable message inside Nimble's error envelopes. + + Falls back to the raw body for anything else (CDN HTML pages, plain text, other shapes). + """ + try: + body: Final = _ErrorEnvelope.model_validate_json(error_message) + except ValidationError: + return error_message + return body.detail or body.message or error_message + + +def _domain_filters(search_domain_filter: object) -> Mapping[str, object]: + """ + Split the unified `search_domain_filter` into Nimble's include/exclude lists. + + Follows the Perplexity unified spec, where a `-` prefix means "exclude this domain". + Anything that is not a list of strings is ignored rather than raising, since it only + ever narrows a search that is otherwise valid. + """ + try: + domains: Final = _DomainListAdapter.validate_python(search_domain_filter) + except ValidationError: + return _NOTHING + return MappingProxyType( + { + key: value + for key, value in ( + ("include_domains", tuple(d for d in domains if d and not d.startswith("-"))), + ("exclude_domains", tuple(d[1:] for d in domains if d.startswith("-") and len(d) > 1)), + ) + if value + } + ) + + +def _publish_date(additional_data: object) -> str | None: + try: + return _AdditionalData.model_validate(additional_data).publish_date + except ValidationError: + return None diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b288269b0a2..0eb9f6119ff 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -16295,6 +16295,14 @@ "notes": "TinyFish Search API" } }, + "nimble/search": { + "input_cost_per_query": 0.005, + "litellm_provider": "nimble", + "mode": "search", + "metadata": { + "notes": "Nimble Search API pay-as-you-go list price: $5 per 1,000 searches, up to 100 results per search. Volume plans price differently." + } + }, "elevenlabs/scribe_v1": { "input_cost_per_second": 6.11e-05, "litellm_provider": "elevenlabs", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d9ef538d530..220826ccbca 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3758,6 +3758,7 @@ class SearchProviders(str, Enum): YOU_COM = "you_com" APISERPENT = "apiserpent" TINYFISH = "tinyfish" + NIMBLE = "nimble" # Create a set of all search provider values for quick lookup diff --git a/litellm/utils.py b/litellm/utils.py index 79372f00284..f011c1ff62f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9064,6 +9064,7 @@ class ProviderConfigManager: from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig from litellm.llms.linkup.search.transformation import LinkupSearchConfig + from litellm.llms.nimble.search.transformation import NimbleSearchConfig from litellm.llms.parallel_ai.search.transformation import ( ParallelAISearchConfig, ) @@ -9093,6 +9094,7 @@ class ProviderConfigManager: SearchProviders.YOU_COM: YouComSearchConfig, SearchProviders.APISERPENT: APISerpentSearchConfig, SearchProviders.TINYFISH: TinyfishSearchConfig, + SearchProviders.NIMBLE: NimbleSearchConfig, } config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None) if config_class is None: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b288269b0a2..0eb9f6119ff 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16295,6 +16295,14 @@ "notes": "TinyFish Search API" } }, + "nimble/search": { + "input_cost_per_query": 0.005, + "litellm_provider": "nimble", + "mode": "search", + "metadata": { + "notes": "Nimble Search API pay-as-you-go list price: $5 per 1,000 searches, up to 100 results per search. Volume plans price differently." + } + }, "elevenlabs/scribe_v1": { "input_cost_per_second": 6.11e-05, "litellm_provider": "elevenlabs", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 65db63dc045..0712e8e383d 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2423,6 +2423,13 @@ "search": true } }, + "nimble": { + "display_name": "Nimble (`nimble`)", + "url": "https://docs.nimbleway.com/api-reference/search/search", + "endpoints": { + "search": true + } + }, "triton": { "display_name": "Triton (`triton`)", "url": "https://docs.litellm.ai/docs/providers/triton-inference-server", diff --git a/tests/code_coverage_tests/enforce_llms_folder_style.py b/tests/code_coverage_tests/enforce_llms_folder_style.py index 2cbd445365e..04a95b45196 100644 --- a/tests/code_coverage_tests/enforce_llms_folder_style.py +++ b/tests/code_coverage_tests/enforce_llms_folder_style.py @@ -22,6 +22,7 @@ SEARCH_PROVIDERS = [ "serper", "apiserpent", "tinyfish", + "nimble", ] ALLOWED_FILES_IN_LLMS_FOLDER = [ diff --git a/tests/search_tests/test_nimble_search.py b/tests/search_tests/test_nimble_search.py new file mode 100644 index 00000000000..c83b7236a09 --- /dev/null +++ b/tests/search_tests/test_nimble_search.py @@ -0,0 +1,155 @@ +""" +Tests for Nimble Search API integration. +""" + +import json +import os +import sys +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm +from tests.search_tests.base_search_unit_tests import BaseSearchTest + +MOCK_NIMBLE_RESPONSE = { + "request_id": "0f8b3a1c-1d2e-4f5a-9b0c-6d7e8f9a0b1c", + "total_results": 2, + "results": [ + { + "title": "Nimble Web API", + "description": "Short SERP description", + "url": "https://nimbleway.com/", + "content": "Full markdown content for the first result", + "metadata": {"position": 1, "entity_type": "organic", "country": "US", "locale": "en"}, + "additional_data": {"publish_date": "2026-07-15"}, + }, + { + "title": "Nimble Docs", + "description": "Only a description here", + "url": "https://docs.nimbleway.com/", + "content": "", + "metadata": {"position": 2, "entity_type": "organic"}, + "additional_data": None, + }, + ], + "serp_data": None, +} + + +def _mock_response(): + response = Mock() + response.status_code = 200 + response.headers = {} + response.content = json.dumps(MOCK_NIMBLE_RESPONSE).encode() + return response + + +@pytest.mark.skip(reason="Local only tested search providers") +class TestNimbleSearch(BaseSearchTest): + """ + E2E tests for Nimble Search functionality that make real API calls. + Inherits from BaseSearchTest to run standard search tests. + """ + + def get_search_provider(self) -> str: + return "nimble" + + +class TestNimbleSearchTransformation: + """ + Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked. + Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/. + """ + + @pytest.fixture(autouse=True) + def _server_key(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NIMBLE_API_KEY", "test-api-key") + monkeypatch.delenv("NIMBLE_API_BASE", raising=False) + + def test_nimble_search_request_and_response(self): + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ) as mock_post: + response = litellm.search( + query="nimble web scraping", + search_provider="nimble", + max_results=2, + country="us", + search_domain_filter=["nimbleway.com", "-spam.example"], + ) + + assert mock_post.called + call_kwargs = mock_post.call_args.kwargs + assert call_kwargs["url"] == "https://sdk.nimbleway.com/v2/search" + assert call_kwargs["headers"]["Authorization"] == "Bearer test-api-key" + assert call_kwargs["headers"]["X-Client-Source"] == "litellm" + + request_body = call_kwargs["json"] + assert request_body["query"] == "nimble web scraping" + assert request_body["max_results"] == 2 + assert request_body["country"] == "US" + assert request_body["include_domains"] == ("nimbleway.com",) + assert request_body["exclude_domains"] == ("spam.example",) + + assert response.object == "search" + assert len(response.results) == 2 + assert response.results[0].title == "Nimble Web API" + assert response.results[0].url == "https://nimbleway.com/" + assert response.results[0].snippet == "Full markdown content for the first result" + assert response.results[0].date == "2026-07-15" + # Second result has no `content`, so the SERP description is the snippet. + assert response.results[1].snippet == "Only a description here" + assert response.results[1].date is None + + def test_provider_specific_params_survive_to_the_wire(self): + """Nimble-native params must not be eaten by `filter_out_litellm_params`.""" + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ) as mock_post: + litellm.search( + query="test query", + search_provider="nimble", + focus="news", + search_depth="deep", + time_range="week", + locale="fr", + output_format="plain_text", + max_subagents=5, + ) + + request_body = mock_post.call_args.kwargs["json"] + assert request_body["focus"] == "news" + assert request_body["search_depth"] == "deep" + assert request_body["time_range"] == "week" + assert request_body["locale"] == "fr" + assert request_body["output_format"] == "plain_text" + assert request_body["max_subagents"] == 5 + + @pytest.mark.asyncio + async def test_nimble_asearch(self): + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new=AsyncMock(return_value=_mock_response()), + ) as mock_post: + response = await litellm.asearch( + query="latest ai developments", + search_provider="nimble", + focus="news", + ) + + assert mock_post.call_args.kwargs["json"]["focus"] == "news" + assert len(response.results) == 2 + + def test_nimble_search_tracks_cost(self): + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", + return_value=_mock_response(), + ): + response = litellm.search(query="pricing check", search_provider="nimble") + + assert response._hidden_params["response_cost"] == pytest.approx(0.005) diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py index a1353d57038..b93ffdb0b44 100644 --- a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py +++ b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py @@ -27,6 +27,7 @@ from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig from litellm.llms.linkup.search.transformation import LinkupSearchConfig +from litellm.llms.nimble.search.transformation import NimbleSearchConfig from litellm.llms.parallel_ai.search.transformation import ParallelAISearchConfig from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig from litellm.llms.searchapi.search.transformation import SearchAPIConfig @@ -57,6 +58,7 @@ _BASE_ENV_VARS = ( "DATAFORSEO_API_BASE", "TINYFISH_API_BASE", "CRW_API_BASE", + "NIMBLE_API_BASE", ) @@ -96,6 +98,7 @@ PROVIDERS: Tuple[ProviderSpec, ...] = ( ), (TinyfishSearchConfig, {"TINYFISH_API_KEY": "srv"}, "caller-key", {}), (FastCRWSearchConfig, {"CRW_API_KEY": "srv"}, "caller-key", {}), + (NimbleSearchConfig, {"NIMBLE_API_KEY": "srv"}, "caller-key", {}), ) _IDS = tuple(spec[0].__name__ for spec in PROVIDERS) diff --git a/tests/test_litellm/llms/nimble/search/test_nimble_search_transformation.py b/tests/test_litellm/llms/nimble/search/test_nimble_search_transformation.py new file mode 100644 index 00000000000..d6292c9cf3e --- /dev/null +++ b/tests/test_litellm/llms/nimble/search/test_nimble_search_transformation.py @@ -0,0 +1,251 @@ +import json +from unittest.mock import Mock + +import pytest + +from litellm.llms.nimble.search.transformation import NimbleSearchConfig + + +def _config() -> NimbleSearchConfig: + return NimbleSearchConfig() + + +def _resp(payload, status_code: int = 200): + r = Mock() + r.status_code = status_code + r.headers = {} + r.content = (payload if isinstance(payload, str) else json.dumps(payload)).encode() + return r + + +def _result(**overrides): + base = { + "title": "Test Title", + "description": "Test description", + "url": "https://example.com", + "content": "Test content", + "metadata": {"position": 1, "entity_type": "organic"}, + "additional_data": None, + } + return {**base, **overrides} + + +def test_ui_friendly_name(): + assert _config().ui_friendly_name() == "Nimble" + + +def test_validate_environment_with_explicit_key(): + headers = _config().validate_environment({}, api_key="explicit-key") + assert headers["Authorization"] == "Bearer explicit-key" + assert headers["Content-Type"] == "application/json" + assert headers["X-Client-Source"] == "litellm" + + +def test_validate_environment_reads_env_key(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NIMBLE_API_KEY", "env-key") + assert _config().validate_environment({})["Authorization"] == "Bearer env-key" + + +def test_validate_environment_missing_key_raises(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NIMBLE_API_KEY", raising=False) + with pytest.raises(ValueError, match="NIMBLE_API_KEY"): + _config().validate_environment({}) + + +def test_validate_environment_does_not_mutate_and_is_idempotent(): + """The http handler re-runs validate_environment after search/main.py already did.""" + config = _config() + caller_headers = {"X-Custom": "keep-me"} + + once = config.validate_environment(caller_headers, api_key="k") + twice = config.validate_environment(once, api_key="k") + + assert caller_headers == {"X-Custom": "keep-me"} + assert once == twice + assert once["X-Custom"] == "keep-me" + + +def test_get_complete_url_default_base(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("NIMBLE_API_BASE", raising=False) + assert _config().get_complete_url(None, {}) == "https://sdk.nimbleway.com/v2/search" + + +def test_get_complete_url_reads_env_base(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("NIMBLE_API_BASE", "https://env-base.local/v2") + assert _config().get_complete_url(None, {}) == "https://env-base.local/v2/search" + + +@pytest.mark.parametrize( + "api_base", + [ + "https://self-hosted.local/v2", + "https://self-hosted.local/v2/", + "https://self-hosted.local/v2/search", + "https://self-hosted.local/v2/search/", + ], +) +def test_get_complete_url_appends_search_exactly_once(api_base: str): + assert _config().get_complete_url(api_base, {}) == "https://self-hosted.local/v2/search" + + +def test_transform_search_request_joins_list_query(): + assert _config().transform_search_request(["foo", "bar"], {})["query"] == "foo bar" + + +def test_transform_search_request_max_results_is_not_clamped(): + """Nimble validates 1-100 itself; a clearer error beats silently rewriting the request.""" + assert _config().transform_search_request("q", {"max_results": 500})["max_results"] == 500 + + +def test_transform_search_request_uppercases_country(): + assert _config().transform_search_request("q", {"country": "us"})["country"] == "US" + + +def test_transform_search_request_drops_max_tokens_per_page(): + assert "max_tokens_per_page" not in _config().transform_search_request("q", {"max_tokens_per_page": 1024}) + + +def test_transform_search_request_splits_domain_filter(): + data = _config().transform_search_request("q", {"search_domain_filter": ["arxiv.org", "-spam.com", "nature.com"]}) + assert data["include_domains"] == ("arxiv.org", "nature.com") + assert data["exclude_domains"] == ("spam.com",) + + +def test_transform_search_request_omits_empty_domain_lists(): + data = _config().transform_search_request("q", {"search_domain_filter": ["arxiv.org"]}) + assert data["include_domains"] == ("arxiv.org",) + assert "exclude_domains" not in data + + +def test_transform_search_request_ignores_non_list_domain_filter(): + assert "include_domains" not in _config().transform_search_request("q", {"search_domain_filter": "arxiv.org"}) + + +@pytest.mark.parametrize("native_key", ["include_domains", "exclude_domains"]) +def test_transform_search_request_native_domains_win(native_key: str): + """An explicit provider-native value must not be silently clobbered by the unified param.""" + data = _config().transform_search_request( + "q", + {"search_domain_filter": ["derived.com", "-derived-ex.com"], native_key: ["native.com"]}, + ) + assert data[native_key] == ["native.com"] + + +def test_transform_search_response_prefers_content(): + resp = _config().transform_search_response(_resp({"results": [_result()]}), logging_obj=Mock()) + assert resp.results[0].snippet == "Test content" + + +def test_transform_search_response_falls_back_to_description(): + resp = _config().transform_search_response(_resp({"results": [_result(content="")]}), logging_obj=Mock()) + assert resp.results[0].snippet == "Test description" + + +def test_transform_search_response_reads_publish_date(): + resp = _config().transform_search_response( + _resp({"results": [_result(additional_data={"publish_date": "2026-08-01"})]}), + logging_obj=Mock(), + ) + assert resp.results[0].date == "2026-08-01" + + +@pytest.mark.parametrize("additional_data", [{}, "not-a-dict"]) +def test_transform_search_response_date_is_none_without_usable_publish_date(additional_data): + resp = _config().transform_search_response( + _resp({"results": [_result(additional_data=additional_data)]}), logging_obj=Mock() + ) + assert resp.results[0].date is None + + +def test_transform_search_response_keeps_additional_data(): + """News results often carry only a relative `publish_date_raw`, which is not a date; + it must still reach the caller rather than being dropped on the floor.""" + resp = _config().transform_search_response( + _resp({"results": [_result(additional_data={"publish_date_raw": "1 day ago"})]}), + logging_obj=Mock(), + ) + assert resp.results[0].date is None + assert resp.results[0].additional_data == {"publish_date_raw": "1 day ago"} + + +def test_transform_search_response_omits_additional_data_when_absent(): + resp = _config().transform_search_response(_resp({"results": [_result()]}), logging_obj=Mock()) + assert not hasattr(resp.results[0], "additional_data") + + +def test_transform_search_response_preserves_order(): + resp = _config().transform_search_response( + _resp({"results": [_result(title=t) for t in ("first", "second", "third")]}), + logging_obj=Mock(), + ) + assert [r.title for r in resp.results] == ["first", "second", "third"] + + +def test_transform_search_response_degraded_result_does_not_fail_the_call(): + resp = _config().transform_search_response( + _resp({"results": [{"url": "https://example.com"}, _result()]}), logging_obj=Mock() + ) + assert len(resp.results) == 2 + assert resp.results[0].title == "" + assert resp.results[0].snippet == "" + assert resp.results[1].title == "Test Title" + + +def test_transform_search_response_zero_hits(): + """A search with no hits really does come back as `"results": []`.""" + payload = {"request_id": "abc", "total_results": 0, "results": []} + assert _config().transform_search_response(_resp(payload), logging_obj=Mock()).results == [] + + +@pytest.mark.parametrize( + "body", + [ + "502 Bad Gateway", # non-JSON body + '{"results": ["garbage"]}', # right key, wrong element shape + '{"results": {"unexpected": "shape"}}', + '{"results": null}', # must not degrade to a successful empty search + "{}", # ditto for an absent key + ], +) +def test_transform_search_response_malformed_body_raises_instead_of_reporting_empty(body: str): + """A body LiteLLM cannot parse must not be reported as a successful zero-result search.""" + with pytest.raises(Exception, match="Nimble Search"): + _config().transform_search_response(_resp(body, status_code=502), logging_obj=Mock()) + + +def test_get_error_class_attributes_the_provider(): + error = _config().get_error_class(error_message="quota exceeded", status_code=429, headers={}) + assert error.status_code == 429 + assert "Nimble Search: quota exceeded" in str(error) + assert "docs.nimbleway.com" in str(error) + + +def test_get_error_class_unwraps_nimble_detail_envelope(): + """Verbatim body from a live 422; the raw JSON envelope should not reach the user.""" + error = _config().get_error_class( + error_message='{"detail":"search_depth=\'fast\' is only supported with focus=\'general\'."}', + status_code=422, + headers={}, + ) + assert ( + str(error) == "Nimble Search: search_depth='fast' is only supported with focus='general'. " + "See https://docs.nimbleway.com/api-reference/search/search for details." + ) + + +def test_get_error_class_unwraps_nimble_message_envelope(): + """Verbatim body from a live collection failure, which uses a different envelope.""" + error = _config().get_error_class( + error_message='{"success":"false","task_id":"4f74af04","message":"can\'t download the query response"}', + status_code=500, + headers={}, + ) + assert ( + str(error) == "Nimble Search: can't download the query response. " + "See https://docs.nimbleway.com/api-reference/search/search for details." + ) + + +@pytest.mark.parametrize("body", ["502 Bad Gateway", '{"detail": null}']) +def test_get_error_class_falls_back_to_the_raw_body(body: str): + assert f"Nimble Search: {body}." in str(_config().get_error_class(body, status_code=500, headers={})) diff --git a/ui/litellm-dashboard/public/assets/logos/nimble.png b/ui/litellm-dashboard/public/assets/logos/nimble.png new file mode 100644 index 0000000000000000000000000000000000000000..6ad2ff611e787fdb6d04f4fd074280be2e87976a GIT binary patch literal 6579 zcmXw8cTf}E*WQFcLJ>mm5I}ktktPs&2)(?5G(kjAlrA7OAs|hvG*No*pfssbKtzz< zrHQD3^xpHu-^};N&OL3;bN1eybDp^yZEUD>je?B=0Dx<{C{0rU01-tH03#zVHeRI< zi3_<0>aI5cP}2W|BW11V9b_R+HRF>`o%JujtdvH0kNljp6uP;z zQp_BXB<`Pkn6FJ?B^lGFRmRNq(o(E`xLci4TwGu{44S9i2B%?oW#>=9_}1otV@H5b z@DnJ`?Jrhc_FH$&_4mhC+joaDFAa|dHntBd8x|4>G(O3F(+{`X7YOG;7mI)XgS^%5 zT(mm72zV^Z_gFY4y?#`@=utGo<{fmyi8F(V>;4~Ti1Nz~XC)HVHv}u!`MwZuT~RAD zt7mjs4+z7vhdE|C^9y(Rmc(K(hZ$NO@Sgse?KC+iIj7D$RT~NErFJ(!Um#)Tf)i)r zqic^dqr0lDl|&(xP_R8~+jPwwvR@=7jW55Iw`lg2@6*7M3NfnaNoSwTYP)0BD&G6G zlrZGCov`f%7VecSTP+{NIV!Es^`P8v?PakK|K{j=b?BB4Y@7^|dw+nMyJFRzp$s}V zKVS3(qGonv6cF&k%t|fXSR9U_SkekTEM7IP94DD^e;!$*>XFyWXhhF38rq}_UDO|3i|3x5d( zOM%6LR5Aag+G|7!3#WpDhoK4~EJ#=*T^mqIxUA?ePg2pdJjwGc_mh2om zp%nixmLTS^i=!1Xt@?cbLPb@S0k)t6u(yWng%Mu)Vy@?DlGy10Q@+7sAyVxB6m8(l z{hsR~)6~75Pf)rOj%)tRAdiU7rs@1hAUS%O_)Onms z^+kFB55Tb}7&|Ywv@K6?}@i9tYU< zyQZ}5O9AQlJ_f|lv0v|9cW%cI7{<4cOIK=RpMm~BfK!!PT8qD6!y;f~y=xhI>(eu= zzjVOQy+tc=FPJj~1NAM{j_IMP>|ubu|5PtiRz?rTr(S7ubgBs*fzI!Kr5LMexzF~_ z4j>b`1mra;$KOl%z$zuczuj|9dtrdZ$eHE-@8fJ9y}COG@&X?yMh|j@0I}Q;7ECn#(j$0AF%!P$ z4TB__4Vpkuw(1P0V+;*Y@0rqEHhTGNfWDL;yyU=Tc>z4Zm%suQxjoDjQYn!H3`F1* zs7=K^Z$%G=lsi%9ot}KbCwWtH8A^RcdLs0zwg(l9Vst_YmSqJ_Z^F^Gwp#mF zox~6C`{IXeyO1JXyh_OF-@mK)&IKSzBBAK>ee-pN*P|*Qf-*{w3q!m+CrOeHO?4^W zf0k_uhuS&mg-9BL4nJ}BmR?8G0@Rm9Lr0EHb<9j_B**iZ7NoA`T@ifxvN3teyUeq^ zPWkDl8!!HXg%;9qkFOW8tag(kMGm+lW}YU}O_RqoIhh!WAkKtSY0nbE%))vTHa_7V z@)Y~`Z>l47Sj#Z0>%yW3+&%p@VkTq905)sW$| z_lfkRLQ2v}vqeU4Tca25%qElZ(1MnRRbxsg1BZ<<@0jDqXOHu6e5CB*g@Y7PLA>)1 zk29{>g)-!B@JlW{d}@=?qG_3|C~E9qw6j&i8;~Lf+ENR2pp@3PRTp}unUAov1Oj4h z|NcJhQg)HEVMgvY%F~g4RL=pJuf8nr+Ii5-_TRgz0MQB=ww{k)0a)S?E3axoEgUDO31$2?|adF zsM4zL_t>>1k6=KbCIS9Njv~XAvAOx2=6=BF7v;^aC}6aK^|cH|#$CW>;SwtCSd|z! zOFr5#lR=Cm+7^j9=f><3tG|8WFv06@S4y@Ld_96D?DSh+3_ZNTK{~KbnV}0jd3S-v z{j50pw8p0LM)>szvs43Mb5%Mq*iql-ZEx1VeD%ZW#ix~!<*J9ExYVsjH@`xd}vAS>%40*V?x;L)1DSdFw;9msLwzFA-Jz&`2s&~NBs zI{l-!Cme=_Vk>6EFDS8xx zr6)x+oU{?-{hNNwncd*W1RBTdaM>{%kQ={e{lp;tu?}n>xa-a)&o6)eyCh#2G>5xU zi>i5;6X%1*gX1y{2KNf@FHH+Jf8S{ABhlXy&Z5IB3*mIh3pF!I*<|0mEvk6r<1$C{ zGB`5G>*`u#K_L-=+K{~2jqf!bGIyVLndasBeiAtG-u|z={e3N=H2Ewy~p>m;b_7iu)Y6ZRo5gLDKUsH+`kO-89@&GVg_ z3D)degTeQVP+^JU-ASfU?W>tu_cOLDngXgwHe&rNNmfTQsA32LUnkV3 zjKqiiud+gl@X!6URw=j2*`*HF4o>F?@`ap_J!rmq4$ZgQKu#(kd3U39iVl?Qf>q63 z%yLmmkTRYLark80eb|uXcnQ5L5Y^f>eV~u(F zcNEZn1$NVJ+Q$Ng=qB=)5uA_ThN3~oxg28P1afZiqru^#gEWJXSU^Y%3L6K{;e40h z1!TzwDkve$Xlj8*?xbtg{t zJ;NRpz%wY3@4j(DHpu>A@Gf=#8NIj9E*9T=IG((^3P>DNvewtt%*@}}~ z34tP3d|z*FshNq>h+EZNdYfdD@dBpeUL5oxrw&tZuQAWzc(R{K<-}%*+e|En>RkF_ zn6;20UBPv(gE5p_AK^oaJAdv!qV$WN3K`#Yt0xKKt>3O=nX_1WaOc`Z_~gEDpM= z@IofGlGd1q%dSz>ir?pPxDkqY)u~9smMWn;p0T(2A3(DxYdEL;r(*N{g-6j%U%b_Q zhi9&4FBM(4nc9x{e(8cd=y;S?pM-9zWHtDa+nvue=7!aQA~$lqe>8a4hy5I4q#M~n z2~%vbczikT`M!|QyZeV?Wc|p#=+?^lvaJWr@`<;srq$RcgHYbAIjhASF7HgCTy61% zU)m?mzRj?};S+_o3J?FXz%eR32`j!Il~)I(t$ShbVDd4EI$o{DL6jf-t;ue)QEji> zvS}=1YUdb>SbihCaGy7s3XY*|b}m7YTe`$~QPniky1% z>u`udheSeN&|sHI_~XS=)*DRht^Ydsj|Xy83%D8Dgr_31(qfx3Qf{mIPZf8|zFFUz z&kWgpFM6=Qk(1xN#N~689G{5=1^`tT)xj9tTW98Lbrs8XVif;6Q3RnwqF(odf8dx~ zeo2l9(`h(_$bsegm*9c2*5+Co8|6!#+uVd5g?Yj^2}Xza0(;y?lbR>{rz12qD>YM$ zN!9d?;83K!gnYLzKbt+R|5}73#ddrf-p1qCytki`=0TA5t1Sdr2xc*FJNs09ma|B6 z$Qc2EGg9!W&~F>@=N;L`u-aAKakZfED*QG*9Lrsw9zuwa5&rmkN|o}V-_5F#ZIyaM ze69)-f*815T>srFy8TXKDcOZF!n%Ms3>4hsn4amKfnzZI_T^jE8`awN>f}iuWJP|Y z`EV624)Fr*+PMt=e!ptRQd96nlMhYQ1KTxp*GyAL)!`VrnI_3TKMM;lqn}(<{F=UH zsiN0vE>Vvrc#K8bzI$w~q0w-ws0}Mo3kth)lW4V%Y05qD($s!~OoG2-t|sqiHsPfyA72GXtjtO<%J6$sr7B4pto!mdlynJWo;?C8 zYPTbV(!HLV4N&!1Z(E7_QBV4B9>*bDG}1nbaH#WqEt-?jbZ_q*8Ty<#M2aL5_PZ!l zgiy$^%_E(N6hWEl2Cotx{#16r0Al}C5K*X}@<``{2da0aFONgWi-<;BWpUd70&r1Da4UZ&j&?w+x$pA`K)^#%-AGKFzyM@_geOD(^Y zw(+Ia85v~A9a&Wn1<0zk9p_%z`cJ7IW5^DDDe#|P19D>+HL0yWLv$U^d2v6Bi98u5V3XRXCYq)A2!zPR+(WrJs^gZ5)rkp0K=KPag%Fbm$sW=NEt)J zBRI591naWB=>Y%`ULa7sH?JTEQym(F57P)T?s2I_-}iFl;TSO?#C>*B?vAgTKb8vmb_^U8je1S!Z&6Xd9Zo$0EK z%m5Zj&ykB)W?|c$S9DxAd9IfrdH7dfjrKV#a-DTcXs-=(A2CLcxz=2kz1AW5vk*`I z1tA*@6;^93>+SnmyG_H02`QC2P~G_#bn524cnw7kN%DF2B9qSJq_h3fndQ{5(2B@n zvq|@B85_!Z`4hE9Cs^;J6RUnDWwj?6t*WPRl>kZa{&{pB8|Leht4ELfqzp;`JjEa! zimF8ujYwdn5|svC@X0Qkg?au^l{rd{93{Twc&=U{g>xgrMcD25c26JFFAISLi$>U!L>%7^vP%z80-5n(mFI6k-w zWAB`;H1FiRzChC-QeipD_$n`QmE$!;@1x@Rj)nPIeqSyQb_>=%cMZWJhace)raG4Q zLYzRtSVZ#AFTp7Cn`?v1;#twq^WuQBcPh)|;5VpEc@+6d5{J%U5?dgd4a1uG7?BW` zTQHRi+H;ud#IpXYKAun2PhyZyR${rOyLo|tFjMxlp&B^2xtg{lR9HO!Z@;`a6m%ZX zeSXGo`lk5AIpge6EM0_M6--NhXK?-ZWZnAUqjONU*bbE);t7&FBV?$o8zF+okx$SYXr8 z-r~Rob)hg1MFx1z@PXt)JfD$s!mu+FCc=6!XW-KdYHak9FK=KHmyVs>=~hmaL1xm2 z`v^_Ydk3-Z=7wpN>3$y1+*04s({*oeav&cw^P`r&iRpXv3N~Lc!|&MncXb$oB7?Ip zX$r|yqWZ??YHMx|JZ$ln5iG-_u*df0z6yIZiPnE6500H<4%L<@GuX>ceUV&!Bw9tU zn5=OT5MT>`KjYLF4m$V9i|ha69kX$z?;seyrDa5y!M-?aqT&$9F85y_tq95iWO=fu zRB`jyjH47=nFAA&i+!x&uHi+TMmOb5RI+1unOX#qqOY&8{X3fHCe1tP#v-Pn^t3=W zbFM5!7#)gJoNN!3eIllL=j7G&C2QRnoKmi{e;<4Y#_%~4dX1{*C0`Ansp#y{_xM#$ zAqI$3lA;yWio>xz=QFHk(Qz^&AB-pcXzp4F4)d@DnC>f;wc!eWsfXfDu&X_97M0R2 z4RD?MnsglEjqGC|g=I{C4tGA^Z|y~8mnKh?$L!ilq_#Au`LVs#mtA}gM{>focRs^jzEX)-3`JTIk@QK{yO2$ru?=M$@x+% z7uWNX^>H(7Loq$2T#jVWv@=~Q4AU`fh4u+z`q>T2(9b@u@KLqiD0?8n`f^b6`W=Ho zCu17G{`T{Qdp*m}BCV`EHPAvxg%sepBx~xspo%K8C|HSPF=l|+%g@Xv-d>ft)-b0s zbJ@FK@m>;24%A}v&AuuB;Vz6R`}Mo;1)BP&9SFF!lz$JGY_+l90K*7Vs6&veD`UYT zVrNfFCMj2^2v-wM6R5>6As;k>@Lnq;{_~3ex%&U?WAmfSZ*PO}Waj9|zcv^8Rbc$i zJoxTklHN!)W0$JzCv#p`3={gBV0^AJEG+DMFuTl6*OEy#yl)Hh)@^MxH$M=>Z++p! zwVlOS0l{ElfP3AOU71O8fP@xUaQLz8T*&|sxo+FX(VihoQGiMlkWc7|&8f2jSfT(A zrhaBs$96#KFELquN=O>u0Xb?k4}Ry)^%f2kAMY;U6!V#3jSlv5u7%^cZ_3-%i6ruT zz{`%4hCBJK?)LR8wX`bJ8X_Sc4nsvH@VFO?Cdj{K0v=xeY>zb8oZ~@y!p-~GXm)(^Mpdc* zkYdzq0rTjJ(<`aB6(B0d_Y%2l5Rfr%kYmC>TU!)Bw)W;lvjipuK)ov#SL~IH@B`x! z{J^`_iEK3YfzuRsw;o6;R}Y~0>7N=(XGaO(T!4W{y9p7MjpPb`_;_`X5D#{paZ@+KMFk%7&t^Rb12xSXn0A3Ve-DZ`G8I;io%e{_} z6po6&kE>hA?*LR{kiRW~w|N6FNU3m^EgC+5kQ723Mq+n@E_1A=+7M!g;cp=zA;~m< zQSWar?sh4SA^~lHXNN@WViySr06ZS_Vi#+~3tVzakd}{7U>7yKmTBeWJqVPMwAG{Z zU2!6ESfxHTW!0x}+*9m=c-1@5_aY(JPs;R+yb#(@r)`zxqS^j4?`Ss`41(wSd2WBa zS%izC5OENvpKF5u?kCv_EFtmO*MJjDUGZT9dzO zjti);JeUB3N?!+QgPx_2nK!*RD%fmAxQx0x5v=l1Rkwr{B1rE^!h!xJSe{PYuj>Zs z@6mn&u$EutLOWa4&sy?zt3I35$8avFAb@y0S)o2vTC%z1E7Sg9hvv9-?$GRCn~K(} z0kY@kmAiw63{I_qO}HL6RXEW5BPvht8&j-CBivfi$pLY7m~3iNL2&#n8oJaJce(iU z?FC!bMWEv5&hp^Z69cVsdZ&A4F^ZNeTl<8f-_bJzRvpvCe=2~kmZ4^~x_#LH0h9G2 Az5oCK literal 0 HcmV?d00001 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx index 1eeff00cb1b..6c8cef0b1a1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx @@ -12,6 +12,7 @@ import { AvailableSearchProvider, SearchTool } from "./types"; import dataforseoLogo from "../../../../../public/assets/logos/dataforseo.png"; import exaAiLogo from "../../../../../public/assets/logos/exa_ai.png"; import googlePseLogo from "../../../../../public/assets/logos/google_pse.png"; +import nimbleLogo from "../../../../../public/assets/logos/nimble.png"; import parallelAiLogo from "../../../../../public/assets/logos/parallel_ai.png"; import perplexityLogo from "../../../../../public/assets/logos/perplexity.png"; import tavilyLogo from "../../../../../public/assets/logos/tavily.png"; @@ -25,6 +26,7 @@ const searchProviderLogoMap: Record = { exa_ai: exaAiLogo.src, google_pse: googlePseLogo.src, dataforseo: dataforseoLogo.src, + nimble: nimbleLogo.src, }; interface SearchProviderLabelProps { From f9f5c03884fc2a98a77ae66b1107d233b25dce16 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Fri, 14 Aug 2026 17:21:07 -0700 Subject: [PATCH 27/27] fix(mcp): drop caller host and configured upstream headers from logged metadata (#36901) * fix(mcp): drop caller host and configured upstream headers from logged metadata The synthetic request that carries MCP client headers into add_litellm_data_to_request forwarded the caller's Host header, and Request.url is built from it, so a caller chose the proxy_server_request url and the metadata endpoint that every logging callback records. _upstream_credential_headers also only knew the configured client side auth header and the x-mcp- prefix family, so a header name declared in mcp_servers..extra_headers reached logging metadata in cleartext. Those names are admin chosen, so no prefix rule can recognize them; read them off the server registry instead. The header is still forwarded upstream, which is what extra_headers is for. authorization is left out because clean_headers already strips it and claiming it here would move authenticated_with_header on the oauth passthrough config. The Responses bridge tests stub the server manager, so their fakes gain the registry accessor the sanitizer now reads. * fix(mcp): drop caller host from the sanitized header mapping too The synthetic request stopped forwarding host, but the parallel sanitizer did not, so a forged hostname still reached the guardrail payload and the list_tools spend row. Drop it there as well. Exempt the configured identity headers from the upstream credential set. get_user_from_headers resolves end user attribution off the same request this module reconstructs, and it only fills end_user_id when auth left it unset, so claiming user_header_name or a user_header_mappings name would lose attribution on the MCP paths that authenticate upstream. Drop the isinstance guard on extra_headers entries: the field is typed list[str], so the check is dead and basedpyright scores it. * fix(mcp): accept a bare user_header_mappings entry when exempting identity headers get_internal_user_header_from_mapping and get_customer_user_header_from_mapping both normalize a single mapping to a one element list, and config_settings.md documents the key as a dict. Iterating the bare form yields its keys instead, so the exemption silently matched nothing and an identity header also named in an MCP server's extra_headers was dropped after all. --- .../proxy/_experimental/mcp_server/utils.py | 68 +++++++++- .../_experimental/mcp_server/test_utils.py | 122 ++++++++++++++++++ .../mcp/test_litellm_proxy_mcp_handler.py | 4 + .../mcp/test_mcp_streaming_iterator.py | 1 + 4 files changed, 189 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 4cf84dd0725..83883664df5 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -880,7 +880,9 @@ _HOP_BY_HOP_HEADERS: Final = frozenset( } ) -_SYNTHETIC_REQUEST_EXCLUDED_HEADERS: Final = _HOP_BY_HOP_HEADERS | frozenset({"content-type", "x-forwarded-for"}) +_SYNTHETIC_REQUEST_EXCLUDED_HEADERS: Final = _HOP_BY_HOP_HEADERS | frozenset( + {"content-type", "host", "x-forwarded-for"} +) _SYNTHETIC_REQUEST_SERVER: Final = ("127.0.0.1", 4000) @@ -908,10 +910,57 @@ def _mcp_client_side_auth_header_name() -> str: return MCPRequestHandler.LITELLM_MCP_AUTH_HEADER_NAME +def _identity_header_names() -> frozenset[str]: + """Lowercased header names the deployment reads the caller's identity out of. A name here + is a claim about who the caller is rather than a secret, and ``get_user_from_headers`` + resolves it off the request this module reconstructs, so dropping one would lose end user + attribution on the MCP paths that leave ``end_user_id`` unset at connect time. + + ``user_header_mappings`` is accepted as a bare mapping as well as a list of them, matching + ``get_internal_user_header_from_mapping`` and ``get_customer_user_header_from_mapping``. + Iterating the bare form without normalizing yields its keys, which would silently exempt + nothing.""" + try: + from litellm.proxy.proxy_server import general_settings + except ImportError: + return frozenset() + if not general_settings: + return frozenset() + user_header: Final = general_settings.get("user_header_name") + configured: Final = general_settings.get("user_header_mappings") + mappings: Final = configured if isinstance(configured, list) else (configured,) if configured else () + mapped: Final = (mapping.get("header_name") for mapping in mappings if isinstance(mapping, Mapping)) + return frozenset(name.lower() for name in (user_header, *mapped) if isinstance(name, str) and name) + + +def _forwarded_upstream_header_names() -> frozenset[str]: + """Lowercased header names that a configured MCP server forwards upstream through its + ``extra_headers`` allowlist. The names are chosen by the admin, so no prefix rule can + recognize them, and a caller supplied value under one of them is an upstream credential. + + ``authorization`` is left out because ``clean_headers`` already strips it, and claiming it + here would change which header ``authenticated_with_header`` resolves to on the oauth + passthrough config, which lists it in ``extra_headers`` by design. Identity headers are + left out for the same reason: naming one in ``extra_headers`` forwards the caller's + identity upstream, it does not turn that identity into a secret.""" + try: + from .mcp_server_manager import global_mcp_server_manager + except ImportError: + return frozenset() + exempt: Final = _identity_header_names() | frozenset({"authorization"}) + return frozenset( + name.lower() + for server in global_mcp_server_manager.get_registry().values() + for name in (server.extra_headers or ()) + if name.lower() not in exempt + ) + + def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]: """Lowercased names of the headers in ``header_names`` that carry an upstream MCP - credential rather than request context: the configured client side auth header and - the per-server ``x-mcp-{alias}-{header}`` family. ``clean_headers`` only knows the + credential rather than request context: the configured client side auth header, any + header name a configured server forwards upstream via ``extra_headers``, and the + per-server ``x-mcp-{alias}-{header}`` family. ``clean_headers`` only knows the credential headers of the chat completions path, so these are dropped on top of it. """ from .auth.user_api_key_auth_mcp import MCPRequestHandler @@ -923,10 +972,13 @@ def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]: } ) client_side_auth: Final = _mcp_client_side_auth_header_name().lower() + forwarded_upstream: Final = _forwarded_upstream_header_names() return frozenset( name for name in (raw_name.lower() for raw_name in header_names) - if name == client_side_auth or (name.startswith(_MCP_SERVER_AUTH_HEADER_PREFIX) and name not in non_credential) + if name == client_side_auth + or name in forwarded_upstream + or (name.startswith(_MCP_SERVER_AUTH_HEADER_PREFIX) and name not in non_credential) ) @@ -944,7 +996,9 @@ def build_synthetic_mcp_request( ``proxy_server_request``, header-based tags, guardrails and trace correlation exactly as on the chat completions path. Hop-by-hop headers describe the original HTTP framing rather than the logical request, so they are dropped, and - ``x-forwarded-for`` comes from the resolved ``client_ip`` to avoid spoofing. Upstream + ``x-forwarded-for`` comes from the resolved ``client_ip`` to avoid spoofing. ``host`` is + dropped for the same reason: it is what ``Request.url`` is built from, so forwarding it + would let a caller choose the URL every logging callback records. Upstream MCP credentials and the deployment's proxy key header, including a custom ``litellm_key_header_name``, are dropped so they cannot reach a callback or a guardrail through the derived metadata even when a caller omits ``general_settings``. @@ -991,7 +1045,8 @@ def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[s too: these headers are read back out of the metadata to change proxy behaviour, so leaving one in place would let any MCP client turn off the redaction an admin configured. This path carries no key or team object to authorize an opt-out with, so - it always strips them.""" + it always strips them. ``host`` goes too, so that a caller cannot name the deployment in + the guardrail payload and the spend row the way it could once name the request URL.""" from starlette.datastructures import Headers from litellm.proxy.litellm_pre_call_utils import ( @@ -1003,6 +1058,7 @@ def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[s excluded: Final = ( _upstream_credential_headers(raw_headers.keys() if raw_headers else ()) | UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS + | frozenset({"host"}) ) cleaned: Final = clean_headers( Headers(raw_headers), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py index 00ed4e91efa..0252fb9843d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_utils.py @@ -4,12 +4,32 @@ import pytest from fastapi import HTTPException from litellm.proxy._experimental.mcp_server.utils import ( + _upstream_credential_headers, build_synthetic_mcp_request, logging_safe_mcp_headers, validate_and_normalize_mcp_server_payload, validate_tool_display_names, ) from litellm.proxy._types import NewMCPServerRequest +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _server_forwarding(*header_names: str) -> MCPServer: + return MCPServer( + server_id="srv-1", + name="deepwiki", + transport="http", + url="https://mcp.example.com/mcp", + extra_headers=list(header_names), + ) + + +def _configured_servers(*servers: MCPServer): + return patch.dict( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers", + {server.server_id: server for server in servers}, + clear=False, + ) class TestValidateToolDisplayNames: @@ -114,6 +134,70 @@ class TestLoggingSafeMcpHeaders: assert safe == {"x-nuid": "nuid-1"} + def test_strips_headers_a_server_forwards_upstream(self): + """mcp_servers..extra_headers names the headers the proxy relays upstream, so a + caller supplied value under one of them is an upstream credential no prefix rule can spot. + Config is written in canonical casing while the wire header arrives lowercased.""" + with _configured_servers(_server_forwarding("X-GitHub-Token", "X-Tenant")): + safe = logging_safe_mcp_headers({"x-github-token": "ghp_secret", "x-tenant": "acct-1", "x-nuid": "nuid-1"}) + + assert safe == {"x-nuid": "nuid-1"} + + def test_strips_caller_asserted_host(self): + """This mapping reaches the guardrail payload and the list_tools spend row, so a caller + must not be able to name the deployment there either.""" + safe = logging_safe_mcp_headers({"host": "evil.attacker.example", "x-nuid": "nuid-1"}) + + assert safe == {"x-nuid": "nuid-1"} + + def test_keeps_identity_header_a_server_also_forwards(self): + """get_user_from_headers resolves end user attribution off this same request, so a header + the deployment reads identity from stays even when a server forwards it upstream.""" + with patch.dict( + "litellm.proxy.proxy_server.general_settings", + {"user_header_name": "x-user-email"}, + clear=False, + ): + with _configured_servers(_server_forwarding("x-user-email", "x-github-token")): + safe = logging_safe_mcp_headers({"x-user-email": "alice@corp.example", "x-github-token": "ghp_secret"}) + + assert safe == {"x-user-email": "alice@corp.example"} + + @pytest.mark.parametrize( + "configured", + [ + [{"header_name": "X-User", "litellm_user_role": "customer"}], + {"header_name": "X-User", "litellm_user_role": "customer"}, + ], + ids=["list-of-mappings", "bare-mapping"], + ) + def test_keeps_identity_header_from_user_header_mappings(self, configured): + """get_internal_user_header_from_mapping and get_customer_user_header_from_mapping both + accept a bare mapping as well as a list, and config_settings.md documents the key as a + dict, so the exemption has to read both shapes.""" + with patch.dict( + "litellm.proxy.proxy_server.general_settings", + {"user_header_mappings": configured}, + clear=False, + ): + with _configured_servers(_server_forwarding("X-User", "X-GitHub-Token")): + safe = logging_safe_mcp_headers({"x-user": "alice", "x-github-token": "ghp_secret"}) + + assert safe == {"x-user": "alice"} + + def test_keeps_authorization_classification_for_oauth_passthrough(self): + """clean_headers already strips authorization, and claiming it here would change which + header authenticated_with_header resolves to on a config that lists it by design.""" + with _configured_servers(_server_forwarding("Authorization", "X-GitHub-Token")): + assert "authorization" not in _upstream_credential_headers(["authorization", "x-github-token"]) + assert "x-github-token" in _upstream_credential_headers(["authorization", "x-github-token"]) + + def test_keeps_headers_when_no_server_forwards_them(self): + with _configured_servers(_server_forwarding("x-github-token")): + safe = logging_safe_mcp_headers({"x-other-token": "not-forwarded", "x-nuid": "nuid-1"}) + + assert safe == {"x-other-token": "not-forwarded", "x-nuid": "nuid-1"} + class TestBuildSyntheticMcpRequest: def test_forwards_client_headers_without_upstream_credentials(self): @@ -147,3 +231,41 @@ class TestBuildSyntheticMcpRequest: assert request.headers.get("x-nuid") == "nuid-1" assert "x-company-key" not in request.headers + + def test_drops_caller_host_so_the_logged_url_is_not_client_steerable(self): + """add_litellm_data_to_request records str(request.url) as proxy_server_request.url, and + Request.url is built from the host header, so forwarding it hands the caller that value.""" + request = build_synthetic_mcp_request( + path="/mcp/tools/call", + raw_headers={"host": "evil.attacker.example", "x-nuid": "nuid-1"}, + ) + + assert "evil.attacker.example" not in str(request.url) + assert "host" not in request.headers + assert request.headers.get("x-nuid") == "nuid-1" + + def test_drops_headers_a_server_forwards_upstream(self): + with _configured_servers(_server_forwarding("x-github-token")): + request = build_synthetic_mcp_request( + path="/mcp/tools/call", + raw_headers={"x-github-token": "ghp_secret", "x-nuid": "nuid-1"}, + ) + + assert "x-github-token" not in request.headers + assert request.headers.get("x-nuid") == "nuid-1" + + def test_keeps_identity_header_so_end_user_attribution_survives(self): + """add_litellm_data_to_request reads user_header_name off this request to fill + end_user_id, so forwarding that header upstream must not remove it here.""" + with patch.dict( + "litellm.proxy.proxy_server.general_settings", + {"user_header_name": "x-user-email"}, + clear=False, + ): + with _configured_servers(_server_forwarding("x-user-email")): + request = build_synthetic_mcp_request( + path="/mcp/tools/call", + raw_headers={"x-user-email": "alice@corp.example"}, + ) + + assert request.headers.get("x-user-email") == "alice@corp.example" diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index 4981caa10c3..87525273911 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -28,6 +28,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module) fake_manager = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), call_tool=AsyncMock(return_value=_DummyMCPResult()), # Newer logging path calls this to enrich spend logs metadata _get_mcp_server_from_tool_name=MagicMock(return_value=None), @@ -373,6 +374,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey post_call_failure_hook = _setup_proxy_logging(monkeypatch) fake_manager = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) ) monkeypatch.setattr( @@ -464,6 +466,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch # Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields. fake_manager = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), get_allowed_mcp_servers=AsyncMock(return_value=[]), get_mcp_servers_from_ids=MagicMock(return_value=[]), get_mcp_server_by_name=MagicMock(return_value=None), @@ -516,6 +519,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch): mock_get_tools, ) fake_manager = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), get_allowed_mcp_servers=AsyncMock(return_value=[]), get_mcp_servers_from_ids=MagicMock(return_value=[]), get_mcp_server_by_name=MagicMock(return_value=None), diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index 24edf12fffe..aacd614abb9 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -69,6 +69,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock: """Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests.""" call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)) fake_manager = types.SimpleNamespace( + get_registry=MagicMock(return_value={}), call_tool=call_tool, _get_mcp_server_from_tool_name=MagicMock(return_value=None), get_mcp_server_by_name=MagicMock(return_value=None),