From b83d11351fc2db09444fa7457ce138017d22f6b8 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 6 May 2026 09:06:58 -0700 Subject: [PATCH 01/16] proxy: hot-reload config YAML when --reload is set (#27274) * proxy: hot-reload config YAML when --reload is set Uvicorn's --reload only watches *.py by default, so editing the --config YAML did not restart the proxy. _get_reload_options() now extends reload_dirs/reload_includes with the config file's directory and basename when --config is provided. * proxy: qualify reload_includes with absolute config path Address Greptile review on PR #27274. When the --config file lives outside cwd, reload_includes previously stored only the basename, which meant uvicorn/watchfiles would also reload on edits to any same-named file inside cwd. Use the absolute config path as the include pattern in that case so only the actual proxy config triggers a restart. Co-authored-by: Mateo Wang * fix(proxy): use basename for reload_includes config pattern Uvicorn's resolve_reload_patterns() calls pathlib.Path.glob(), which raises NotImplementedError on absolute patterns (uvicorn discussion 2156). Passing config_abs (an absolute path) when the config file lived outside cwd crashed startup under --reload. The config_dir is already added to reload_dirs, so using just the basename as the include pattern is sufficient to match the specific config file. * fix: make it reload app when yaml changes * style: remove unneeded comments --------- Co-authored-by: Claude Co-authored-by: Cursor Agent Co-authored-by: Mateo Wang --- litellm/proxy/proxy_cli.py | 68 +++++++++++++++++- tests/test_litellm/proxy/test_proxy_cli.py | 81 ++++++++++++++++++++++ 2 files changed, 147 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 71aeea67884..90bbfdd25a4 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -169,6 +169,66 @@ class ProxyInitializationHelpers: ) return uvicorn_args + @staticmethod + def _get_reload_options(config_path: Optional[str]) -> dict: + """Build uvicorn reload kwargs so --reload also reacts to YAML edits.""" + options: dict = {"reload": True} + if not config_path: + return options + config_abs = os.path.abspath(config_path) + config_dir = os.path.dirname(config_abs) + cwd = os.path.abspath(os.getcwd()) + reload_dirs = [cwd] + if config_dir and config_dir != cwd: + reload_dirs.append(config_dir) + options["reload_dirs"] = reload_dirs + # Must be a basename, not an absolute path: uvicorn's + # resolve_reload_patterns() calls pathlib.Path.glob(), which raises + # NotImplementedError on absolute patterns (uvicorn discussion #2156). + options["reload_includes"] = ["*.py", os.path.basename(config_abs)] + return options + + @staticmethod + def _patch_statreload_for_config(config_path: str) -> bool: + """Make uvicorn's StatReload reloader notice YAML config changes. + + Uvicorn uses WatchFilesReload when the optional `watchfiles` package + is installed, otherwise StatReload. StatReload hard-codes `*.py` in + `iter_py_files()` and silently ignores `reload_includes`, so the + kwargs from `_get_reload_options` alone don't trigger reloads on YAML + edits. We monkey-patch `iter_py_files` to also yield the config path. + + Idempotent across calls and a no-op for the WatchFilesReload path. + """ + try: + from uvicorn.supervisors.statreload import StatReload + except ImportError: # pragma: no cover - uvicorn is a hard dep + return False + + if not config_path: + return False + + from pathlib import Path + + config_abs = Path(config_path).resolve() + + patched_paths = getattr(StatReload, "_litellm_patched_config_paths", None) + if patched_paths is None: + original_iter = StatReload.iter_py_files + patched_paths = set() + + def _iter_with_config(self): # type: ignore[no-untyped-def] + yield from original_iter(self) + for path in StatReload._litellm_patched_config_paths: + if path.exists(): + yield path + + StatReload.iter_py_files = _iter_with_config # type: ignore[assignment] + StatReload._litellm_patched_config_paths = patched_paths # type: ignore[attr-defined] + + patched_paths.add(config_abs) + return True + @staticmethod def _init_hypercorn_server( app: FastAPI, @@ -619,7 +679,7 @@ class ProxyInitializationHelpers: "--reload", is_flag=True, default=False, - help="Enable uvicorn hot reload (dev only). Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", + help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) def run_server( # noqa: PLR0915 host, @@ -1028,7 +1088,11 @@ def run_server( # noqa: PLR0915 uvicorn_args["loop"] = loop_type if reload: - uvicorn_args["reload"] = True + uvicorn_args.update( + ProxyInitializationHelpers._get_reload_options(config) + ) + if config: + ProxyInitializationHelpers._patch_statreload_for_config(config) uvicorn.run( **uvicorn_args, diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 6fbce4a5458..f05e95f9e50 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -133,6 +133,87 @@ class TestProxyInitializationHelpers: ) assert args["timeout_worker_healthcheck"] == 15 + def test_get_reload_options_no_config(self): + opts = ProxyInitializationHelpers._get_reload_options(None) + assert opts == {"reload": True} + + def test_get_reload_options_with_config_in_cwd(self, tmp_path, monkeypatch): + config_file = tmp_path / "config.yaml" + config_file.write_text("model_list: []\n") + monkeypatch.chdir(tmp_path) + + opts = ProxyInitializationHelpers._get_reload_options("config.yaml") + + assert opts["reload"] is True + assert opts["reload_dirs"] == [str(tmp_path)] + assert opts["reload_includes"] == ["*.py", "config.yaml"] + + def test_get_reload_options_with_config_outside_cwd(self, tmp_path, monkeypatch): + cwd_dir = tmp_path / "work" + cwd_dir.mkdir() + elsewhere = tmp_path / "configs" + elsewhere.mkdir() + config_file = elsewhere / "proxy.yaml" + config_file.write_text("model_list: []\n") + monkeypatch.chdir(cwd_dir) + + opts = ProxyInitializationHelpers._get_reload_options(str(config_file)) + + assert opts["reload"] is True + assert opts["reload_dirs"] == [str(cwd_dir), str(elsewhere)] + assert opts["reload_includes"] == ["*.py", "proxy.yaml"] + + def test_patch_statreload_for_config_yields_yaml(self, tmp_path): + from pathlib import Path + + from uvicorn.supervisors.statreload import StatReload + + if hasattr(StatReload, "_litellm_patched_config_paths"): + StatReload._litellm_patched_config_paths.clear() + + config_file = tmp_path / "config.yaml" + config_file.write_text("model_list: []\n") + py_file = tmp_path / "module.py" + py_file.write_text("x = 1\n") + + applied = ProxyInitializationHelpers._patch_statreload_for_config( + str(config_file) + ) + assert applied is True + + fake_self = types.SimpleNamespace( + config=types.SimpleNamespace(reload_dirs=[tmp_path]) + ) + yielded_paths = {Path(p).resolve() for p in StatReload.iter_py_files(fake_self)} + + assert config_file.resolve() in yielded_paths + assert py_file.resolve() in yielded_paths + + def test_patch_statreload_for_config_is_idempotent(self, tmp_path): + from pathlib import Path + + from uvicorn.supervisors.statreload import StatReload + + if hasattr(StatReload, "_litellm_patched_config_paths"): + StatReload._litellm_patched_config_paths.clear() + + config_file = tmp_path / "config.yaml" + config_file.write_text("model_list: []\n") + py_file = tmp_path / "only.py" + py_file.write_text("x = 1\n") + + for _ in range(3): + ProxyInitializationHelpers._patch_statreload_for_config(str(config_file)) + + fake_self = types.SimpleNamespace( + config=types.SimpleNamespace(reload_dirs=[tmp_path]) + ) + yielded = list(StatReload.iter_py_files(fake_self)) + assert len(yielded) == len(set(map(str, yielded))) + yielded_paths = {Path(p).resolve() for p in yielded} + assert config_file.resolve() in yielded_paths + assert py_file.resolve() in yielded_paths + @patch("asyncio.run") @patch("builtins.print") def test_init_hypercorn_server(self, mock_print, mock_asyncio_run): From b1f577199a108d7ee9b3ee5616d9baab5059fa32 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 6 May 2026 11:39:15 -0700 Subject: [PATCH 02/16] fix(proxy): keep spend log cleanup running after batch failures and surface DB errors (#27303) Co-authored-by: Yassin Kortam --- litellm/constants.py | 6 + .../db_transaction_queue/spend_log_cleanup.py | 65 +++++-- .../proxy/test_spend_log_cleanup.py | 166 ++++++++++++++++++ 3 files changed, 225 insertions(+), 12 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 6918e40cad1..1e96f8eac9a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1457,6 +1457,12 @@ KEY_ROTATION_JOB_NAME = "litellm_key_rotation_job" EXPIRED_UI_SESSION_KEY_CLEANUP_JOB_NAME = "litellm_expired_ui_session_key_cleanup_job" SPEND_LOG_RUN_LOOPS = int(os.getenv("SPEND_LOG_RUN_LOOPS", 500)) SPEND_LOG_CLEANUP_BATCH_SIZE = int(os.getenv("SPEND_LOG_CLEANUP_BATCH_SIZE", 1000)) +SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int( + os.getenv("SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3) +) +SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float( + os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5) +) SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) SPEND_LOG_QUEUE_POLL_INTERVAL = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = int( diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index bc9efb52b0f..9475779cfdf 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -5,8 +5,10 @@ from typing import Optional from litellm._logging import verbose_proxy_logger from litellm.caching import RedisCache from litellm.constants import ( + SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS, SPEND_LOG_CLEANUP_BATCH_SIZE, SPEND_LOG_CLEANUP_JOB_NAME, + SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES, SPEND_LOG_RUN_LOOPS, ) from litellm.litellm_core_utils.duration_parser import duration_in_seconds @@ -74,6 +76,7 @@ class SpendLogCleanup: """ total_deleted = 0 run_count = 0 + consecutive_failures = 0 while True: if run_count > SPEND_LOG_RUN_LOOPS: verbose_proxy_logger.info( @@ -82,18 +85,50 @@ class SpendLogCleanup: break # Step 1: Find logs and delete them in one go without fetching to application # Delete in batches, limited by self.batch_size - deleted_result = await prisma_client.db.execute_raw( - """ - DELETE FROM "LiteLLM_SpendLogs" - WHERE "request_id" IN ( - SELECT "request_id" FROM "LiteLLM_SpendLogs" - WHERE "startTime" < $1::timestamptz - LIMIT $2 + try: + deleted_result = await prisma_client.db.execute_raw( + """ + DELETE FROM "LiteLLM_SpendLogs" + WHERE "request_id" IN ( + SELECT "request_id" FROM "LiteLLM_SpendLogs" + WHERE "startTime" < $1::timestamptz + LIMIT $2 + ) + """, + cutoff_date, + self.batch_size, ) - """, - cutoff_date, - self.batch_size, - ) + except Exception as batch_exc: + # A single batch failure (e.g. Prisma/DB timeout) must not abort + # the whole run — subsequent batches may still succeed. + consecutive_failures += 1 + verbose_proxy_logger.exception( + "Spend log cleanup batch failed " + "(run_count=%d, consecutive_failures=%d, batch_size=%d, " + "cutoff=%s, total_deleted_so_far=%d): %s: %s", + run_count, + consecutive_failures, + self.batch_size, + cutoff_date.isoformat(), + total_deleted, + type(batch_exc).__name__, + batch_exc, + ) + if ( + consecutive_failures + >= SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES + ): + verbose_proxy_logger.error( + "Aborting spend log cleanup after %d consecutive batch " + "failures; total deleted before abort: %d", + consecutive_failures, + total_deleted, + ) + break + await asyncio.sleep(SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS) + continue + + consecutive_failures = 0 deleted_count = 0 if isinstance(deleted_result, int): @@ -168,7 +203,13 @@ class SpendLogCleanup: verbose_proxy_logger.info(f"Deleted {total_deleted} logs") except Exception as e: - verbose_proxy_logger.error(f"Error during cleanup: {str(e)}") + # .exception() captures the traceback; str(e) alone on a Prisma/DB + # timeout is often empty and gives operators no signal to diagnose. + verbose_proxy_logger.exception( + "Error during spend log cleanup: %s: %s", + type(e).__name__, + e, + ) return # Return after error handling finally: # Only release the lock if it was actually acquired diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py index 4923d70a437..42bb919295f 100644 --- a/tests/test_litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -331,6 +331,172 @@ async def test_delete_old_logs_continues_on_valid_int_return(): assert total_deleted == 800 +@pytest.mark.asyncio +async def test_delete_old_logs_continues_after_single_batch_failure(monkeypatch): + """A single batch failure (e.g. DB timeout) must not abort the whole run — + subsequent batches should still execute and their counts accumulate.""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + # Zero out the failure backoff so the test doesn't take ~0.5s of real sleep. + monkeypatch.setattr( + cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0 + ) + + mock_prisma_client = MagicMock() + mock_db = MagicMock() + # batch 1 succeeds, batch 2 raises (one-off DB timeout), batches 3-4 succeed, + # batch 5 returns 0 → loop exits naturally. + mock_db.execute_raw = AsyncMock( + side_effect=[100, TimeoutError("simulated DB timeout"), 200, 50, 0] + ) + mock_prisma_client.db = mock_db + + cleaner = cleanup_module.SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "7d"} + ) + + cutoff_date = datetime.now(timezone.utc) - timedelta(days=7) + total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date) + + # All 5 batches should have been attempted; 100 + 200 + 50 = 350 deleted. + assert mock_db.execute_raw.call_count == 5 + assert total_deleted == 350 + + +@pytest.mark.asyncio +async def test_delete_old_logs_aborts_after_consecutive_failures(monkeypatch): + """If batch failures persist for SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES + in a row (e.g. DB is down), the loop must abort instead of hot-looping.""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + # Lower the threshold so the test is fast and deterministic. + monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3) + monkeypatch.setattr( + cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0 + ) + + mock_prisma_client = MagicMock() + mock_db = MagicMock() + # Every batch raises — must abort after exactly 3 attempts, not loop forever. + mock_db.execute_raw = AsyncMock( + side_effect=ConnectionError("simulated persistent DB outage") + ) + mock_prisma_client.db = mock_db + + cleaner = cleanup_module.SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "7d"} + ) + + cutoff_date = datetime.now(timezone.utc) - timedelta(days=7) + total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date) + + assert mock_db.execute_raw.call_count == 3 + assert total_deleted == 0 + + +@pytest.mark.asyncio +async def test_delete_old_logs_resets_consecutive_failures_on_success(monkeypatch): + """A success between failures must reset the consecutive-failure counter so + intermittent timeouts don't trip the abort threshold.""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 3) + monkeypatch.setattr( + cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0 + ) + + mock_prisma_client = MagicMock() + mock_db = MagicMock() + # Pattern: fail, fail, success (resets counter), fail, fail, success, done. + # Without reset, three of these would trip abort; with reset, they don't. + mock_db.execute_raw = AsyncMock( + side_effect=[ + TimeoutError("t1"), + TimeoutError("t2"), + 100, + TimeoutError("t3"), + TimeoutError("t4"), + 50, + 0, + ] + ) + mock_prisma_client.db = mock_db + + cleaner = cleanup_module.SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "7d"} + ) + + cutoff_date = datetime.now(timezone.utc) - timedelta(days=7) + total_deleted = await cleaner._delete_old_logs(mock_prisma_client, cutoff_date) + + assert mock_db.execute_raw.call_count == 7 + assert total_deleted == 150 + + +@pytest.mark.asyncio +async def test_cleanup_uses_logger_exception_for_full_traceback(monkeypatch): + """The outer error handler must call logger.exception() (not .error(str(e))) + so Prisma/DB timeouts surface a full traceback and exception type.""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + mock_logger = MagicMock() + monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger) + + mock_prisma_client = MagicMock() + # Force the outer try/except to fire by making _should_delete_spend_logs raise. + cleaner = cleanup_module.SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "7d"} + ) + cleaner.pod_lock_manager = None + + def boom(): + raise RuntimeError("simulated prisma timeout") + + cleaner._should_delete_spend_logs = boom # type: ignore[assignment] + + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + assert mock_logger.exception.called, "expected logger.exception() to be called" + # The exception type name must appear in the formatted args so operators can + # tell *what* failed, not just "Error during cleanup:". + call_args = mock_logger.exception.call_args + formatted = call_args[0][0] % call_args[0][1:] + assert "RuntimeError" in formatted + assert "simulated prisma timeout" in formatted + + +@pytest.mark.asyncio +async def test_cleanup_releases_lock_after_persistent_batch_failures(monkeypatch): + """Even when batch deletion aborts due to consecutive failures, the pod lock + must still be released so the next scheduled run isn't permanently blocked.""" + import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module + + monkeypatch.setattr(cleanup_module, "SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES", 2) + monkeypatch.setattr( + cleanup_module, "SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.0 + ) + + mock_prisma_client = MagicMock() + mock_db = MagicMock() + mock_db.execute_raw = AsyncMock(side_effect=TimeoutError("DB down")) + mock_prisma_client.db = mock_db + + mock_pod_lock_manager = MagicMock() + mock_pod_lock_manager.redis_cache = MagicMock() + mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True) + mock_pod_lock_manager.release_lock = AsyncMock() + + cleaner = cleanup_module.SpendLogCleanup( + general_settings={"maximum_spend_logs_retention_period": "7d"} + ) + cleaner.pod_lock_manager = mock_pod_lock_manager + + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + # Cleanup didn't crash; the abort-after-failures path returned cleanly. + mock_pod_lock_manager.release_lock.assert_awaited_once() + + def test_cleanup_batch_size_env_var(monkeypatch): """Ensure batch size is configurable via environment variable""" import importlib From c92a08a3077ce4601176d3c1ddf5dd38f825cb06 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 11:42:29 -0700 Subject: [PATCH 03/16] Fix team member budget enforcement without user row (#27273) * Fix team member budget enforcement without user row Co-authored-by: ishaan-berri * Clarify regenerated key budget repro Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- litellm/proxy/auth/auth_checks.py | 1 - .../proxy/auth/test_team_member_budget.py | 73 +++++++++++++++++++ 2 files changed, 73 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1b6bd3aff21..86b3dfd1fd1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3516,7 +3516,6 @@ async def _check_team_member_budget( if ( team_object is not None and team_object.team_id is not None - and user_object is not None and valid_token is not None and valid_token.user_id is not None ): diff --git a/tests/test_litellm/proxy/auth/test_team_member_budget.py b/tests/test_litellm/proxy/auth/test_team_member_budget.py index 11a28106e31..b38a953d189 100644 --- a/tests/test_litellm/proxy/auth/test_team_member_budget.py +++ b/tests/test_litellm/proxy/auth/test_team_member_budget.py @@ -306,6 +306,79 @@ async def test_team_member_budget_check_no_team_membership(): assert result is True +@pytest.mark.asyncio +async def test_team_member_budget_check_blocks_regenerated_key_after_old_key_exhausts_budget(): + """Deleting an exhausted key and creating a new key must not reset a user's team budget.""" + request_body = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + } + + team_object = LiteLLM_TeamTable( + team_id="test-team-1", + team_alias="Test Team", + spend=0.0, + max_budget=None, + ) + # The spend below represents usage accumulated by an earlier key that was + # later deleted. The new key must still be checked against the same + # user/team membership spend instead of receiving a fresh per-key budget. + regenerated_token = UserAPIKeyAuth( + token="new-regenerated-token", + user_id="test-user-1", + team_id="test-team-1", + models=["gpt-3.5-turbo"], + ) + team_membership = LiteLLM_TeamMembership( + user_id="test-user-1", + team_id="test-team-1", + spend=0.0000002, + litellm_budget_table=LiteLLM_BudgetTable( + max_budget=0.0000001, + ), + ) + + mock_request = MagicMock(spec=Request) + mock_prisma_client = MagicMock() + mock_user_api_key_cache = MagicMock() + mock_proxy_logging_obj = MagicMock() + + with ( + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + new_callable=AsyncMock, + return_value=team_membership, + ) as mock_get_team_membership, + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await common_checks( + request_body=request_body, + team_object=team_object, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/chat/completions", + llm_router=None, + proxy_logging_obj=mock_proxy_logging_obj, + valid_token=regenerated_token, + request=mock_request, + ) + + mock_get_team_membership.assert_any_await( + user_id="test-user-1", + team_id="test-team-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_user_api_key_cache, + proxy_logging_obj=mock_proxy_logging_obj, + ) + assert "Budget has been exceeded" in str(exc_info.value) + assert "test-user-1" in str(exc_info.value) + assert "test-team-1" in str(exc_info.value) + + @pytest.mark.asyncio async def test_team_member_budget_check_personal_key_not_team(): """Test that team member budget check is skipped for personal keys (no team).""" From d90cf562450e996d23c4cbee9703f871cb02858b Mon Sep 17 00:00:00 2001 From: oss-agent-shin Date: Wed, 6 May 2026 11:58:47 -0700 Subject: [PATCH 04/16] Fix SCIM user lookup filters (#27308) * Fix SCIM Okta userName lookup Co-authored-by: ishaan-berri * fix scim user filter typing Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- .../management_endpoints/scim/scim_v2.py | 35 +++-- .../scim/test_scim_v2_endpoints.py | 123 +++++++++++++++++- 2 files changed, 148 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 04699f19ffb..1f20764f837 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -4,6 +4,7 @@ This is an enterprise feature and requires a premium license. """ +import re from typing import Any, Dict, List, Optional, Set, Tuple from fastapi import ( @@ -843,6 +844,18 @@ async def get_service_provider_config(request: Request): return SCIMServiceProviderConfig(meta=meta) +def _parse_scim_eq_filter(scim_filter: str) -> Optional[Tuple[str, str]]: + """Parse the SCIM equality filters Okta uses before user lifecycle changes.""" + match = re.match( + r"""\s*([\w.]+)\s+eq\s+(['"]?)(.*?)\2\s*$""", + scim_filter, + flags=re.IGNORECASE, + ) + if not match: + return None + return match.group(1).lower(), match.group(3) + + # User Endpoints @scim_router.get( "/Users", @@ -867,15 +880,21 @@ async def get_users( try: prisma_client = await _get_prisma_client_or_raise_exception() # Parse filter if provided (basic support) - where_conditions = {} + where_conditions: Dict[str, Any] = {} if filter: - # Very basic filter support - only handling userName eq and emails.value eq - if "userName eq" in filter: - user_id = filter.split("userName eq ")[1].strip("\"'") - where_conditions["user_id"] = user_id - elif "emails.value eq" in filter: - email = filter.split("emails.value eq ")[1].strip("\"'") - where_conditions["user_email"] = email + # Okta locates users by userName before deprovisioning. LiteLLM + # exposes SCIM userName from user_email, while older SCIM-created + # users may still have user_id == userName, so support both. + parsed_filter = _parse_scim_eq_filter(filter) + if parsed_filter: + filter_attribute, filter_value = parsed_filter + if filter_attribute == "username": + where_conditions["OR"] = [ + {"user_email": filter_value}, + {"user_id": filter_value}, + ] + elif filter_attribute == "emails.value": + where_conditions["user_email"] = filter_value # Get users from database users: List[LiteLLM_UserTable] = ( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index ad53e87e555..ad893012807 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -4,6 +4,7 @@ import pytest from fastapi import HTTPException from litellm.proxy._types import ( + LiteLLM_UserTable, LitellmUserRoles, NewUserRequest, NewUserResponse, @@ -16,8 +17,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _process_group_patch_operations, create_group, create_user, + get_users, get_service_provider_config, - patch_group, patch_user, update_group, update_user, @@ -259,6 +260,124 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp ) +@pytest.mark.asyncio +async def test_get_users_filters_username_by_exposed_scim_username_for_okta(mocker): + """ + Okta deprovisioning first locates a user with `userName eq ""`. + LiteLLM exposes SCIM userName from user_email, so the lookup must match + user_email even when the internal user_id is a UUID. + """ + user = LiteLLM_UserTable( + user_id="internal-user-id", + user_email="okta.user@example.com", + user_alias="Okta User", + teams=[], + metadata={}, + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[user]) + mock_prisma_client.db.litellm_usertable.count = AsyncMock(return_value=1) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock( + return_value=SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="internal-user-id", + userName="okta.user@example.com", + emails=[SCIMUserEmail(value="okta.user@example.com")], + ) + ), + ) + + response = await get_users( + startIndex=1, + count=10, + filter='userName eq "okta.user@example.com"', + ) + + expected_where = { + "OR": [ + {"user_email": "okta.user@example.com"}, + {"user_id": "okta.user@example.com"}, + ] + } + mock_prisma_client.db.litellm_usertable.find_many.assert_awaited_once_with( + where=expected_where, + skip=0, + take=10, + order={"created_at": "desc"}, + ) + mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with( + where=expected_where + ) + assert response.totalResults == 1 + assert response.Resources[0].id == "internal-user-id" + + +@pytest.mark.asyncio +async def test_get_users_filters_email_value_by_user_email(mocker): + """ + SCIM clients can locate users with `emails.value eq ""`; keep that + filter as a direct user_email lookup alongside the userName fallback query. + """ + user = LiteLLM_UserTable( + user_id="internal-user-id", + user_email="scim.user@example.com", + user_alias="SCIM User", + teams=[], + metadata={}, + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=[user]) + mock_prisma_client.db.litellm_usertable.count = AsyncMock(return_value=1) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock( + return_value=SCIMUser( + schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], + id="internal-user-id", + userName="scim.user@example.com", + emails=[SCIMUserEmail(value="scim.user@example.com")], + ) + ), + ) + + response = await get_users( + startIndex=1, + count=10, + filter='emails.value eq "scim.user@example.com"', + ) + + expected_where = {"user_email": "scim.user@example.com"} + mock_prisma_client.db.litellm_usertable.find_many.assert_awaited_once_with( + where=expected_where, + skip=0, + take=10, + order={"created_at": "desc"}, + ) + mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with( + where=expected_where + ) + assert response.totalResults == 1 + assert response.Resources[0].id == "internal-user-id" + + @pytest.mark.asyncio async def test_handle_existing_user_by_email_no_email(mocker): """Should return None when new_user_request has no email""" @@ -1337,7 +1456,7 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true( ) # Execute the create_group function - should succeed - result = await create_group(group=scim_group) + await create_group(group=scim_group) # Verify users were created assert mock_create_user.call_count == 2 From 169c4366849ca4dc159758cd2f973d6053b4c82d Mon Sep 17 00:00:00 2001 From: Dibyo Mukherjee Date: Wed, 6 May 2026 15:05:22 -0400 Subject: [PATCH 05/16] Fix/member access group team (#27317) * fix(auth): pass team_id in member-level model access check _check_team_member_model_access calls _can_object_call_model without team_id, so access groups defined via model_info.access_groups cannot resolve for team-scoped DB models (their internal router name is model_name__, not the public name). The team-level check already passes team_id; this mirrors that. Co-Authored-By: Claude Opus 4.6 * test(auth): add tests for member-level access group resolution with team_id Eight tests covering _can_object_call_model and _check_team_member_model_access with team-scoped DB models: - access group resolves when team_id is passed - access group fails without team_id (pre-fix behavior) - literal model name still works with team_id (no regression) - denied model still denied with team_id - second model in group also reachable - end-to-end member access via access group (mocked membership) - end-to-end member denied for model not in allowed list - no-override member inherits team-level check Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Opus 4.6 --- litellm/proxy/auth/auth_checks.py | 1 + .../proxy/auth/test_auth_checks.py | 279 ++++++++++++++++++ 2 files changed, 280 insertions(+) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 86b3dfd1fd1..f6f99eb62c8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3618,6 +3618,7 @@ async def _check_team_member_model_access( llm_router=llm_router, models=member_allowed_models, object_type="team", + team_id=team_object.team_id, ) except ProxyException: raise ProxyException( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 3433e3e6d85..8a854bcd6a8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -906,6 +906,285 @@ def test_can_object_call_model_no_access_to_alias_or_underlying(): assert "my-fake-gpt" in str(exc_info.value.message) +# -- Team-member access-group resolution with team-scoped DB models ----------- + + +def _make_team_scoped_router(team_id: str = "team-a"): + """ + Build a Router whose model_list looks like what the proxy creates for + team-scoped BYOK DB models: the internal model_name is + ``__`` and the public name lives in + ``model_info.team_public_model_name``. Two models belong to the + access group ``fast-models``; one (``mock-power``) does not. + """ + from litellm import Router + + model_list = [ + { + "model_name": f"mock-fast-1_{team_id}_aaa", + "litellm_params": { + "model": "openai/mock-fast-1", + "api_key": "fake", + }, + "model_info": { + "id": f"demo-mock-fast-1-{team_id}", + "team_id": team_id, + "team_public_model_name": "mock-fast-1", + "access_groups": ["fast-models"], + }, + }, + { + "model_name": f"mock-fast-2_{team_id}_bbb", + "litellm_params": { + "model": "openai/mock-fast-2", + "api_key": "fake", + }, + "model_info": { + "id": f"demo-mock-fast-2-{team_id}", + "team_id": team_id, + "team_public_model_name": "mock-fast-2", + "access_groups": ["fast-models"], + }, + }, + { + "model_name": f"mock-power_{team_id}_ccc", + "litellm_params": { + "model": "openai/mock-power", + "api_key": "fake", + }, + "model_info": { + "id": f"demo-mock-power-{team_id}", + "team_id": team_id, + "team_public_model_name": "mock-power", + }, + }, + ] + return Router(model_list=model_list) + + +def test_can_object_call_model_access_group_with_team_id(): + """ + When team_id is passed, _can_object_call_model should resolve + model_info.access_groups for team-scoped DB models and allow + access via group name. + """ + from litellm.proxy.auth.auth_checks import _can_object_call_model + + router = _make_team_scoped_router() + + result = _can_object_call_model( + model="mock-fast-1", + llm_router=router, + models=["fast-models", "mock-power"], + object_type="team", + team_id="team-a", + ) + assert result is True + + +def test_can_object_call_model_access_group_without_team_id_fails(): + """ + Without team_id the router cannot find team-scoped DB models, so + access group resolution fails and the call is denied. + This is the pre-fix behavior. + """ + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + router = _make_team_scoped_router() + + with pytest.raises(ProxyException): + _can_object_call_model( + model="mock-fast-1", + llm_router=router, + models=["fast-models", "mock-power"], + object_type="team", + # team_id intentionally omitted + ) + + +def test_can_object_call_model_literal_name_with_team_id(): + """ + Literal model name matching should still work when team_id is + passed — no regression from adding team_id. + """ + from litellm.proxy.auth.auth_checks import _can_object_call_model + + router = _make_team_scoped_router() + + result = _can_object_call_model( + model="mock-power", + llm_router=router, + models=["fast-models", "mock-power"], + object_type="team", + team_id="team-a", + ) + assert result is True + + +def test_can_object_call_model_denied_model_with_team_id(): + """ + A model not in the allowed list (by name or access group) should + still be denied even when team_id is passed. + """ + from litellm.proxy._types import ProxyException + from litellm.proxy.auth.auth_checks import _can_object_call_model + + router = _make_team_scoped_router() + + with pytest.raises(ProxyException): + _can_object_call_model( + model="mock-vision", + llm_router=router, + models=["fast-models", "mock-power"], + object_type="team", + team_id="team-a", + ) + + +def test_can_object_call_model_second_group_member_with_team_id(): + """ + Both models in the access group should be reachable, not just + the first one. + """ + from litellm.proxy.auth.auth_checks import _can_object_call_model + + router = _make_team_scoped_router() + + result = _can_object_call_model( + model="mock-fast-2", + llm_router=router, + models=["fast-models"], + object_type="team", + team_id="team-a", + ) + assert result is True + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_with_access_group(): + """ + End-to-end test of _check_team_member_model_access: a member whose + allowed_models contains an access group name should be allowed to + call models in that group for team-scoped DB models. + """ + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + router = _make_team_scoped_router() + team = LiteLLM_TeamTable(team_id="team-a") + token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a") + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable( + allowed_models=["fast-models", "mock-power"], + ), + ) + + with patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + return_value=membership, + ): + # Should not raise — mock-fast-1 is in the fast-models group + await _check_team_member_model_access( + model="mock-fast-1", + team_object=team, + valid_token=token, + llm_router=router, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_denied_model(): + """ + A member with per-member allowed_models should be denied access to + a model that is neither listed by name nor covered by an access group. + """ + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + ProxyException, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + router = _make_team_scoped_router() + team = LiteLLM_TeamTable(team_id="team-a") + token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a") + membership = LiteLLM_TeamMembership( + user_id="alice", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable( + allowed_models=["fast-models", "mock-power"], + ), + ) + + with patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + return_value=membership, + ): + with pytest.raises(ProxyException) as exc_info: + await _check_team_member_model_access( + model="mock-vision", + team_object=team, + valid_token=token, + llm_router=router, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_no_override_inherits_team(): + """ + When a member has no allowed_models (empty budget table), the function + should return without raising — the team-level check applies instead. + """ + from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + UserAPIKeyAuth, + ) + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + + router = _make_team_scoped_router() + team = LiteLLM_TeamTable(team_id="team-a") + token = UserAPIKeyAuth(token="sk-test", user_id="bob", team_id="team-a") + membership = LiteLLM_TeamMembership( + user_id="bob", + team_id="team-a", + litellm_budget_table=LiteLLM_BudgetTable(), + ) + + with patch( + "litellm.proxy.auth.auth_checks.get_team_membership", + return_value=membership, + ): + # Should return without raising — no per-member restriction + await _check_team_member_model_access( + model="mock-vision", + team_object=team, + valid_token=token, + llm_router=router, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + # Tag Budget Enforcement Tests From c8e47dcb43cd293140de9f01dd84d9e50f35184b Mon Sep 17 00:00:00 2001 From: oss-agent-shin Date: Wed, 6 May 2026 12:29:11 -0700 Subject: [PATCH 06/16] Fix early proxy request size enforcement (#27311) * Add early proxy request size guard Co-authored-by: ishaan-berri * Address request size review feedback Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- .../request_size_limit_middleware.py | 121 ++++++++++++++++ litellm/proxy/proxy_server.py | 8 ++ .../test_request_size_limit_middleware.py | 135 ++++++++++++++++++ 3 files changed, 264 insertions(+) create mode 100644 litellm/proxy/middleware/request_size_limit_middleware.py create mode 100644 tests/proxy_unit_tests/test_request_size_limit_middleware.py diff --git a/litellm/proxy/middleware/request_size_limit_middleware.py b/litellm/proxy/middleware/request_size_limit_middleware.py new file mode 100644 index 00000000000..78a38e3572e --- /dev/null +++ b/litellm/proxy/middleware/request_size_limit_middleware.py @@ -0,0 +1,121 @@ +import json +from typing import Callable, Optional, Union + +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +MaxRequestSizeGetter = Callable[[], Optional[Union[int, float]]] +RequestSizeLimitEnabledGetter = Callable[[], bool] + + +class RequestEntityTooLarge(Exception): + pass + + +class RequestSizeLimitMiddleware: + """ + Reject oversized requests before downstream auth/routes parse the body. + + Content-Length can be rejected without reading any body bytes. Requests + without Content-Length are counted as the ASGI stream is consumed, limiting + memory exposure to the configured threshold plus the current chunk. + """ + + def __init__( + self, + app: ASGIApp, + get_max_request_size_mb: MaxRequestSizeGetter, + is_request_size_limit_enabled: RequestSizeLimitEnabledGetter, + ) -> None: + self.app = app + self.get_max_request_size_mb = get_max_request_size_mb + self.is_request_size_limit_enabled = is_request_size_limit_enabled + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await self.app(scope, receive, send) + return + + max_request_size_mb = self.get_max_request_size_mb() + max_request_size_bytes = _mb_to_bytes(max_request_size_mb) + if max_request_size_bytes is None or not self.is_request_size_limit_enabled(): + await self.app(scope, receive, send) + return + + content_length = _get_content_length(scope=scope) + if content_length is not None and content_length > max_request_size_bytes: + await _send_request_too_large( + send=send, max_request_size_mb=max_request_size_mb + ) + return + + received_body_bytes = 0 + response_started = False + + async def limited_receive() -> Message: + nonlocal received_body_bytes + + message = await receive() + if message["type"] != "http.request": + return message + + received_body_bytes += len(message.get("body", b"")) + if received_body_bytes > max_request_size_bytes: + raise RequestEntityTooLarge + return message + + async def tracking_send(message: Message) -> None: + nonlocal response_started + + if message["type"] == "http.response.start": + response_started = True + await send(message) + + try: + await self.app(scope, limited_receive, tracking_send) + except RequestEntityTooLarge: + if response_started: + raise + await _send_request_too_large( + send=send, max_request_size_mb=max_request_size_mb + ) + + +def _mb_to_bytes(max_request_size_mb: Optional[Union[int, float]]) -> Optional[int]: + if max_request_size_mb is None: + return None + if max_request_size_mb <= 0: + return None + return int(max_request_size_mb * 1024 * 1024) + + +def _get_content_length(scope: Scope) -> Optional[int]: + headers = dict(scope.get("headers") or []) + raw_content_length = headers.get(b"content-length") + if raw_content_length is None: + return None + + try: + return int(raw_content_length) + except ValueError: + return None + + +async def _send_request_too_large( + send: Send, + max_request_size_mb: Optional[Union[int, float]], +) -> None: + body = json.dumps( + {"error": f"Request size is too large. Max size is {max_request_size_mb} MB"}, + separators=(",", ":"), + ).encode("utf-8") + await send( + { + "type": "http.response.start", + "status": 413, + "headers": [ + (b"content-type", b"application/json"), + (b"content-length", str(len(body)).encode("latin-1")), + ], + } + ) + await send({"type": "http.response.body", "body": body, "more_body": False}) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a5905765c6e..5a379183e33 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -403,6 +403,9 @@ from litellm.proxy.middleware.in_flight_requests_middleware import ( InFlightRequestsMiddleware, ) from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware +from litellm.proxy.middleware.request_size_limit_middleware import ( + RequestSizeLimitMiddleware, +) from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, @@ -14881,6 +14884,11 @@ app.include_router(ui_discovery_endpoints_router) app.include_router(google_router) attach_lazy_features(app) +app.add_middleware( + RequestSizeLimitMiddleware, + get_max_request_size_mb=lambda: general_settings.get("max_request_size_mb"), + is_request_size_limit_enabled=lambda: premium_user is True, +) async def _stream_mcp_asgi_response( diff --git a/tests/proxy_unit_tests/test_request_size_limit_middleware.py b/tests/proxy_unit_tests/test_request_size_limit_middleware.py new file mode 100644 index 00000000000..3e8792e0179 --- /dev/null +++ b/tests/proxy_unit_tests/test_request_size_limit_middleware.py @@ -0,0 +1,135 @@ +import pytest +from starlette.responses import JSONResponse +from starlette.testclient import TestClient +from starlette.types import Message + +from litellm.proxy.middleware.request_size_limit_middleware import ( + RequestSizeLimitMiddleware, +) + + +def test_request_size_limit_middleware_rejects_content_length_before_body_read(): + downstream_called = False + + async def app(scope, receive, send): + nonlocal downstream_called + downstream_called = True + response = JSONResponse({"ok": True}) + await response(scope, receive, send) + + client = TestClient( + RequestSizeLimitMiddleware( + app, + get_max_request_size_mb=lambda: 1, + is_request_size_limit_enabled=lambda: True, + ) + ) + + response = client.post( + "/chat/completions", + content=b"x" * (1024 * 1024 + 1), + headers={"content-type": "application/json"}, + ) + + assert response.status_code == 413 + assert response.json() == {"error": "Request size is too large. Max size is 1 MB"} + assert response.headers["content-length"] == str(len(response.content)) + assert downstream_called is False + + +def test_request_size_limit_middleware_zero_limit_disables_guard(): + downstream_called = False + + async def app(scope, receive, send): + nonlocal downstream_called + downstream_called = True + response = JSONResponse({"ok": True}) + await response(scope, receive, send) + + client = TestClient( + RequestSizeLimitMiddleware( + app, + get_max_request_size_mb=lambda: 0, + is_request_size_limit_enabled=lambda: True, + ) + ) + + response = client.post( + "/chat/completions", + content=b"x", + headers={"content-type": "application/json"}, + ) + + assert response.status_code == 200 + assert response.json() == {"ok": True} + assert downstream_called is True + + +@pytest.mark.asyncio +async def test_request_size_limit_middleware_rejects_streamed_body_without_content_length(): + received_body_bytes = 0 + + async def app(scope, receive, send): + nonlocal received_body_bytes + while True: + message = await receive() + if message["type"] == "http.disconnect": + break + received_body_bytes += len(message.get("body", b"")) + if not message.get("more_body", False): + break + + response = JSONResponse({"ok": True}) + await response(scope, receive, send) + + middleware = RequestSizeLimitMiddleware( + app, + get_max_request_size_mb=lambda: 1, + is_request_size_limit_enabled=lambda: True, + ) + sent_messages: list[Message] = [] + receive_messages: list[Message] = [ + { + "type": "http.request", + "body": b"x" * (1024 * 1024), + "more_body": True, + }, + { + "type": "http.request", + "body": b"y", + "more_body": False, + }, + ] + + async def receive(): + return receive_messages.pop(0) + + async def send(message): + sent_messages.append(message) + + await middleware( + { + "type": "http", + "method": "POST", + "path": "/chat/completions", + "headers": [(b"content-type", b"application/json")], + }, + receive, + send, + ) + + expected_body = b'{"error":"Request size is too large. Max size is 1 MB"}' + assert sent_messages[0] == { + "type": "http.response.start", + "status": 413, + "headers": [ + (b"content-type", b"application/json"), + (b"content-length", str(len(expected_body)).encode("latin-1")), + ], + } + assert sent_messages[1] == { + "type": "http.response.body", + "body": expected_body, + "more_body": False, + } + assert received_body_bytes == 1024 * 1024 From 487479eff76da099b358d06cfea0b6508fac4ef0 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 13:35:13 -0700 Subject: [PATCH 07/16] perf: cap Prometheus end-user metric cardinality with TTL + LRU eviction (#27272) Co-authored-by: Yassin Kortam --- litellm/__init__.py | 3 + litellm/integrations/prometheus.py | 59 ++++++ .../bounded_prometheus_series_tracker.py | 107 +++++++++++ .../test_prometheus_end_user_cardinality.py | 181 ++++++++++++++++++ 4 files changed, 350 insertions(+) create mode 100644 litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py create mode 100644 tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py diff --git a/litellm/__init__.py b/litellm/__init__.py index e61ef25057f..617c102fb85 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -414,6 +414,9 @@ custom_prometheus_metadata_labels: List[str] = [] custom_prometheus_tags: List[str] = [] prometheus_metrics_config: Optional[List] = None prometheus_emit_stream_label: bool = False +prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000 +prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0 +prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0 disable_add_prefix_to_prompt: bool = ( False # used by anthropic, to disable adding prefix to prompt ) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 147420cd537..65b8a9fc4b1 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -25,6 +25,9 @@ from typing import ( import litellm from litellm._logging import print_verbose, verbose_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import ( + BoundedPrometheusSeriesTracker, +) from litellm.integrations.prometheus_helpers import ( PrometheusLabelFactoryContext, _get_cached_end_user_id_for_cost_tracking, @@ -81,6 +84,7 @@ class PrometheusLogger(CustomLogger): if _custom_buckets is not None else LATENCY_BUCKETS ) + self._bounded_prometheus_series_tracker = BoundedPrometheusSeriesTracker() # Create metric factory functions self._counter_factory = self._create_metric_factory(Counter) @@ -984,6 +988,40 @@ class PrometheusLogger(CustomLogger): return filtered_labels + def _track_end_user_metric_series( + self, + metric: Any, + metric_name: DEFINED_PROMETHEUS_METRICS, + labels: Dict[str, Optional[str]], + ) -> None: + """ + Cap the cardinality of metrics that include the ``end_user`` label. + + Called *after* ``metric.labels(...).inc()/observe()`` so the emission is + recorded in prometheus-client's child map before any eviction runs. + Series that get evicted before the next scrape lose updates accrued + since the last scrape — this is inherent to any cardinality cap. + """ + labelnames = self.get_labels_for_metric(metric_name) + if UserAPIKeyLabelNames.END_USER.value not in labelnames: + return + if labels.get(UserAPIKeyLabelNames.END_USER.value) is None: + return + + max_series = litellm.prometheus_end_user_metrics_max_series_per_metric + ttl_seconds = litellm.prometheus_end_user_metrics_ttl_seconds + if max_series is None and ttl_seconds is None: + return + + self._bounded_prometheus_series_tracker.track_series( + metric=metric, + metric_name=metric_name, + label_values=tuple(labels.get(label) for label in labelnames), + max_series=max_series, + ttl_seconds=ttl_seconds, + cleanup_interval_seconds=litellm.prometheus_end_user_metrics_cleanup_interval_seconds, + ) + def _inc_labeled_counter( self, counter: Any, @@ -998,6 +1036,7 @@ class PrometheusLogger(CustomLogger): label_context=label_context, ) counter.labels(**_labels).inc(amount) + self._track_end_user_metric_series(counter, metric_name, _labels) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): # Define prometheus client @@ -1479,6 +1518,11 @@ class PrometheusLogger(CustomLogger): self.litellm_llm_api_time_to_first_token_metric.labels( **_ttft_labels ).observe(time_to_first_token_seconds) + self._track_end_user_metric_series( + self.litellm_llm_api_time_to_first_token_metric, + "litellm_llm_api_time_to_first_token_metric", + _ttft_labels, + ) else: verbose_logger.debug( "Time to first token metric not emitted, stream option in model_parameters is not True" @@ -1499,6 +1543,11 @@ class PrometheusLogger(CustomLogger): self.litellm_llm_api_latency_metric.labels(**_labels).observe( api_call_total_time_seconds ) + self._track_end_user_metric_series( + self.litellm_llm_api_latency_metric, + "litellm_llm_api_latency_metric", + _labels, + ) # total request latency total_time_seconds = self._safe_duration_seconds( @@ -1516,6 +1565,11 @@ class PrometheusLogger(CustomLogger): self.litellm_request_total_latency_metric.labels(**_labels).observe( total_time_seconds ) + self._track_end_user_metric_series( + self.litellm_request_total_latency_metric, + "litellm_request_total_latency_metric", + _labels, + ) # request queue time (time from arrival to processing start) _litellm_params = kwargs.get("litellm_params", {}) or {} @@ -1533,6 +1587,11 @@ class PrometheusLogger(CustomLogger): self.litellm_request_queue_time_metric.labels(**_labels).observe( queue_time_seconds ) + self._track_end_user_metric_series( + self.litellm_request_queue_time_metric, + "litellm_request_queue_time_seconds", + _labels, + ) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug( diff --git a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py new file mode 100644 index 00000000000..d834ae20142 --- /dev/null +++ b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import time +from collections import OrderedDict +from threading import RLock +from typing import Any, Dict, Optional + + +class BoundedPrometheusSeriesTracker: + """ + Tracks Prometheus child series and removes stale/excess labelsets. + + The tracker is label-agnostic: callers decide which series should be tracked + and pass the full label tuple used by the Prometheus metric. + """ + + def __init__(self) -> None: + self._series: Dict[str, OrderedDict[tuple[Optional[str], ...], float]] = {} + self._last_ttl_cleanup: Dict[str, float] = {} + self.lock = RLock() + + def track_series( + self, + metric: Any, + metric_name: str, + label_values: tuple[Optional[str], ...], + max_series: Optional[int], + ttl_seconds: Optional[float], + cleanup_interval_seconds: Optional[float], + ) -> None: + if max_series is None and ttl_seconds is None: + return + + now = time.monotonic() + + with self.lock: + series = self._series.setdefault(metric_name, OrderedDict()) + series[label_values] = now + series.move_to_end(label_values) + + if ttl_seconds is not None and self._should_run_ttl_cleanup( + metric_name=metric_name, + now=now, + cleanup_interval_seconds=cleanup_interval_seconds, + ): + expired_label_values = [ + tracked_label_values + for tracked_label_values, last_seen in series.items() + if now - last_seen > ttl_seconds + ] + for tracked_label_values in expired_label_values: + self._remove_metric_series(metric, series, tracked_label_values) + + # max_series <= 0 is treated as "unlimited" so a misconfigured zero + # value cannot silently drop every emission for this metric. + if max_series is not None and max_series > 0: + while len(series) > max_series: + tracked_label_values = next(iter(series)) + if not self._remove_metric_child(metric, tracked_label_values): + break + del series[tracked_label_values] + + def _should_run_ttl_cleanup( + self, + metric_name: str, + now: float, + cleanup_interval_seconds: Optional[float], + ) -> bool: + if cleanup_interval_seconds is None or cleanup_interval_seconds <= 0: + self._last_ttl_cleanup[metric_name] = now + return True + + last_cleanup = self._last_ttl_cleanup.get(metric_name) + if last_cleanup is None or now - last_cleanup >= cleanup_interval_seconds: + self._last_ttl_cleanup[metric_name] = now + return True + return False + + def _remove_metric_series( + self, + metric: Any, + series: OrderedDict[tuple[Optional[str], ...], float], + label_values: tuple[Optional[str], ...], + ) -> None: + if self._remove_metric_child(metric, label_values): + series.pop(label_values, None) + + @staticmethod + def _remove_metric_child( + metric: Any, label_values: tuple[Optional[str], ...] + ) -> bool: + """ + Remove the Prometheus child for ``label_values`` and report whether the + tracker should commit the matching state change. + + Returns ``True`` when the child is no longer present in Prometheus + (either it was just removed or it was already gone), and ``False`` when + ``metric.remove()`` raised an unexpected error and the child likely + still exists. + """ + try: + metric.remove(*label_values) + return True + except KeyError: + return True + except (AttributeError, ValueError): + return False diff --git a/tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py b/tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py new file mode 100644 index 00000000000..868d86a6c24 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_end_user_cardinality.py @@ -0,0 +1,181 @@ +from time import monotonic + +import pytest +from prometheus_client import REGISTRY + +import litellm +from litellm.integrations.prometheus import PrometheusLogger +from litellm.integrations.prometheus_helpers import bounded_prometheus_series_tracker +from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import ( + BoundedPrometheusSeriesTracker, +) +from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + +@pytest.fixture(autouse=True) +def cleanup_prometheus_registry(): + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + + old_enable_end_user = litellm.enable_end_user_cost_tracking_prometheus_only + old_metrics_config = litellm.prometheus_metrics_config + old_max_series = litellm.prometheus_end_user_metrics_max_series_per_metric + old_ttl_seconds = litellm.prometheus_end_user_metrics_ttl_seconds + old_cleanup_interval_seconds = ( + litellm.prometheus_end_user_metrics_cleanup_interval_seconds + ) + + yield + + litellm.enable_end_user_cost_tracking_prometheus_only = old_enable_end_user + litellm.prometheus_metrics_config = old_metrics_config + litellm.prometheus_end_user_metrics_max_series_per_metric = old_max_series + litellm.prometheus_end_user_metrics_ttl_seconds = old_ttl_seconds + litellm.prometheus_end_user_metrics_cleanup_interval_seconds = ( + old_cleanup_interval_seconds + ) + + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +def test_prometheus_end_user_series_are_capped_per_metric(): + litellm.enable_end_user_cost_tracking_prometheus_only = True + litellm.prometheus_metrics_config = [ + { + "group": "end-user-spend", + "metrics": ["litellm_spend_metric"], + "include_labels": ["end_user"], + } + ] + litellm.prometheus_end_user_metrics_max_series_per_metric = 3 + litellm.prometheus_end_user_metrics_ttl_seconds = None + logger = PrometheusLogger() + + for index in range(6): + PrometheusLogger._inc_labeled_counter( + logger, + logger.litellm_spend_metric, + "litellm_spend_metric", + UserAPIKeyLabelValues(end_user=f"end-user-{index}"), + amount=0.01, + ) + + assert len(logger.litellm_spend_metric._metrics) == 3 + assert set(logger.litellm_spend_metric._metrics) == { + ("end-user-3",), + ("end-user-4",), + ("end-user-5",), + } + + +def test_bounded_prometheus_series_tracker_is_label_agnostic(): + class FakeMetric: + def __init__(self): + self.removed_label_values = [] + + def remove(self, *label_values): + self.removed_label_values.append(label_values) + + metric = FakeMetric() + tracker = BoundedPrometheusSeriesTracker() + + for index in range(4): + tracker.track_series( + metric=metric, + metric_name="generic_metric", + label_values=(f"route-{index}", "200"), + max_series=2, + ttl_seconds=None, + cleanup_interval_seconds=60.0, + ) + + assert metric.removed_label_values == [ + ("route-0", "200"), + ("route-1", "200"), + ] + + +def test_bounded_prometheus_series_tracker_treats_zero_max_as_unlimited(): + # A misconfigured ``max_series=0`` must not silently evict every emission. + class FakeMetric: + def __init__(self): + self.removed_label_values = [] + + def remove(self, *label_values): + self.removed_label_values.append(label_values) + + metric = FakeMetric() + tracker = BoundedPrometheusSeriesTracker() + + for index in range(3): + tracker.track_series( + metric=metric, + metric_name="generic_metric", + label_values=(f"end-user-{index}",), + max_series=0, + ttl_seconds=None, + cleanup_interval_seconds=60.0, + ) + + assert metric.removed_label_values == [] + + +def test_prometheus_end_user_series_expire_by_ttl(monkeypatch): + litellm.enable_end_user_cost_tracking_prometheus_only = True + litellm.prometheus_metrics_config = [ + { + "group": "end-user-spend", + "metrics": ["litellm_spend_metric"], + "include_labels": ["end_user"], + } + ] + litellm.prometheus_end_user_metrics_max_series_per_metric = None + litellm.prometheus_end_user_metrics_ttl_seconds = 10.0 + litellm.prometheus_end_user_metrics_cleanup_interval_seconds = 0.0 + logger = PrometheusLogger() + + current_time = [monotonic()] + monkeypatch.setattr( + bounded_prometheus_series_tracker.time, + "monotonic", + lambda: current_time[0], + ) + PrometheusLogger._inc_labeled_counter( + logger, + logger.litellm_spend_metric, + "litellm_spend_metric", + UserAPIKeyLabelValues(end_user="stale-end-user"), + amount=0.01, + ) + + current_time[0] += 11.0 + PrometheusLogger._inc_labeled_counter( + logger, + logger.litellm_spend_metric, + "litellm_spend_metric", + UserAPIKeyLabelValues(end_user="fresh-end-user"), + amount=0.01, + ) + + assert set(logger.litellm_spend_metric._metrics) == {("fresh-end-user",)} + + +def test_prometheus_end_user_not_tracked_by_default(): + litellm.enable_end_user_cost_tracking_prometheus_only = None + labels = PrometheusLogger().get_labels_for_metric("litellm_spend_metric") + assert "end_user" in labels + + label_values = UserAPIKeyLabelValues(end_user="not-exported") + from litellm.integrations.prometheus import prometheus_label_factory + + prometheus_labels = prometheus_label_factory(labels, label_values) + assert prometheus_labels["end_user"] is None From 924c1418437ed8c330811bef4c817bfdfd8b249b Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 15:15:21 -0700 Subject: [PATCH 08/16] Add new chat model metadata (#27313) * add new model metadata Co-authored-by: ishaan-berri * address review feedback Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- ...odel_prices_and_context_window_backup.json | 13 +++++++++ model_prices_and_context_window.json | 13 +++++++++ .../litellm/test_sambanova_model_metadata.py | 29 +++++++++++++++++++ 3 files changed, 55 insertions(+) create mode 100644 tests/litellm/test_sambanova_model_metadata.py diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 7946e2dceef..d609a3d1f4b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -28874,6 +28874,19 @@ "mode": "chat", "output_cost_per_token": 0.0 }, + "sambanova/MiniMax-M2.7": { + "input_cost_per_token": 3e-07, + "litellm_provider": "sambanova", + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://cloud.sambanova.ai/plans/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "sambanova/DeepSeek-R1": { "input_cost_per_token": 5e-06, "litellm_provider": "sambanova", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9c085c2c992..a8630eb0f32 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -28879,6 +28879,19 @@ "mode": "chat", "output_cost_per_token": 0.0 }, + "sambanova/MiniMax-M2.7": { + "input_cost_per_token": 3e-07, + "litellm_provider": "sambanova", + "max_input_tokens": 204800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://cloud.sambanova.ai/plans/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "sambanova/DeepSeek-R1": { "input_cost_per_token": 5e-06, "litellm_provider": "sambanova", diff --git a/tests/litellm/test_sambanova_model_metadata.py b/tests/litellm/test_sambanova_model_metadata.py new file mode 100644 index 00000000000..bc31bfb0af2 --- /dev/null +++ b/tests/litellm/test_sambanova_model_metadata.py @@ -0,0 +1,29 @@ +import json +from pathlib import Path + +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + +def test_sambanova_minimax_m27_model_info(): + model = "sambanova/MiniMax-M2.7" + json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json" + with open(json_path) as f: + model_cost = json.load(f) + + info = model_cost.get(model) + assert ( + info is not None + ), f"{model} not found in model_prices_and_context_window.json" + assert info["litellm_provider"] == "sambanova" + assert info["mode"] == "chat" + assert info["input_cost_per_token"] > 0 + assert info["output_cost_per_token"] > 0 + assert info["max_input_tokens"] == 204800 + assert info["max_output_tokens"] == 131072 + assert info["supports_function_calling"] is True + assert info["supports_reasoning"] is True + assert info["supports_tool_choice"] is True + + routed_model, provider, _, _ = get_llm_provider(model=model) + assert routed_model == "MiniMax-M2.7" + assert provider == "sambanova" From bd1a05aed97a2db216b71cd9fd83bbe51751d787 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 15:18:18 -0700 Subject: [PATCH 09/16] Fix MCP DB reload partial failures (#27314) * Fix MCP database reload partial failures Co-authored-by: ishaan-berri * Avoid staged MCP registry exposure Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- .../mcp_server/mcp_server_manager.py | 112 ++++++++----- .../mcp_server/test_mcp_server.py | 155 +++++++++++++++++- 2 files changed, 221 insertions(+), 46 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9923c3ce4bf..12fd1ef89a5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -768,7 +768,9 @@ class MCPServerManager: ) return new_server - async def _maybe_register_openapi_tools(self, server: MCPServer): + async def _maybe_register_openapi_tools( + self, server: MCPServer, *, initialize_mapping: bool = True + ): """Register OpenAPI tools if the server has a spec_path configured.""" if server.spec_path: verbose_logger.info( @@ -779,7 +781,8 @@ class MCPServerManager: server=server, base_url=server.url or "", ) - self.initialize_tool_name_to_mcp_server_name_mapping() + if initialize_mapping: + self.initialize_tool_name_to_mcp_server_name_mapping() async def add_server(self, mcp_server: LiteLLM_MCPServerTable): try: @@ -1978,7 +1981,11 @@ class MCPServerManager: _SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024 - def _assign_unique_short_prefix(self, server: MCPServer) -> None: + def _assign_unique_short_prefix( + self, + server: MCPServer, + registry: Optional[Dict[str, MCPServer]] = None, + ) -> None: """Resolve and cache a collision-free short tool prefix on ``server``. Called at registration time for every MCP server entering the @@ -2002,7 +2009,8 @@ class MCPServerManager: return used: Dict[str, str] = {} - for other in self.get_registry().values(): + registry_for_collision_check = registry or self.get_registry() + for other in registry_for_collision_check.values(): if other.server_id == server.server_id: continue if other.short_prefix: @@ -2916,46 +2924,72 @@ class MCPServerManager: # against the *full* set so dedup is deterministic regardless of # iteration order. for server in db_mcp_servers: - existing_server = previous_registry.get(server.server_id) + try: + existing_server = previous_registry.get(server.server_id) - if ( - existing_server is not None - and existing_server.updated_at is not None - and server.updated_at is not None - and existing_server.updated_at == server.updated_at - ): - # Re-use existing server instance to avoid re-running build_mcp_server_from_table() - # which can perform network discovery for OAuth2 servers. - new_registry[server.server_id] = existing_server - continue + if ( + existing_server is not None + and existing_server.updated_at is not None + and server.updated_at is not None + and existing_server.updated_at == server.updated_at + ): + # Re-use existing server instance to avoid re-running build_mcp_server_from_table() + # which can perform network discovery for OAuth2 servers. + new_registry[server.server_id] = existing_server + continue - _warn_on_server_name_fields( - server_id=server.server_id, - alias=getattr(server, "alias", None), - server_name=getattr(server, "server_name", None), - ) - verbose_logger.debug( - f"Building server from DB: {server.server_id} ({server.server_name})" - ) - new_server = await self.build_mcp_server_from_table(server) - # Carry the cached short_prefix from the previous registry entry - # (if any) so the prefix is stable across reloads. - if existing_server is not None and existing_server.short_prefix: - new_server.short_prefix = existing_server.short_prefix - new_registry[server.server_id] = new_server + _warn_on_server_name_fields( + server_id=server.server_id, + alias=getattr(server, "alias", None), + server_name=getattr(server, "server_name", None), + ) + verbose_logger.debug( + f"Building server from DB: {server.server_id} ({server.server_name})" + ) + new_server = await self.build_mcp_server_from_table(server) + # Carry the cached short_prefix from the previous registry entry + # (if any) so the prefix is stable across reloads. + if existing_server is not None and existing_server.short_prefix: + new_server.short_prefix = existing_server.short_prefix + new_registry[server.server_id] = new_server + except Exception as e: + verbose_logger.exception( + "Skipping MCP server %s (%s) during DB reload: %s", + server.server_id, + getattr(server, "alias", None), + e, + ) - # Swap in the new registry first so _assign_unique_short_prefix - # sees the complete set when checking for collisions. - self.registry = new_registry - for new_server in new_registry.values(): - self._assign_unique_short_prefix(new_server) - # Register OpenAPI tools *after* the final short prefix is assigned - # so the tools are stored in the global registry under the same - # prefix that lookups will use. - await self._maybe_register_openapi_tools(new_server) + # Assign short prefixes against the full candidate set without + # publishing the staged registry to concurrent callers. + registered_registry: Dict[str, MCPServer] = {} + registered_openapi_tools = False + for server_id, new_server in new_registry.items(): + try: + self._assign_unique_short_prefix(new_server, registry=new_registry) + # Register OpenAPI tools *after* the final short prefix is assigned + # so the tools are stored in the global registry under the same + # prefix that lookups will use. + await self._maybe_register_openapi_tools( + new_server, initialize_mapping=False + ) + registered_registry[server_id] = new_server + if new_server.spec_path: + registered_openapi_tools = True + except Exception as e: + verbose_logger.exception( + "Skipping MCP server %s (%s) during DB reload: %s", + new_server.server_id, + getattr(new_server, "alias", None), + e, + ) + + self.registry = registered_registry + if registered_openapi_tools: + self.initialize_tool_name_to_mcp_server_name_mapping() verbose_logger.debug( - "MCP registry refreshed (%s servers in registry)", len(new_registry) + "MCP registry refreshed (%s servers in registry)", len(registered_registry) ) def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 06f95159c08..d3e90246c37 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -2282,6 +2282,144 @@ class TestMCPServerManagerReload: mock_build.assert_awaited_once_with(db_row) assert manager.registry["server-1"] is rebuilt_server + @pytest.mark.asyncio + async def test_skips_server_when_build_from_database_fails(self, caplog): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + timestamp = datetime.utcnow() + healthy_row = _make_db_mcp_server("healthy-server", timestamp) + bad_row = _make_db_mcp_server("bad-server", timestamp) + another_healthy_row = _make_db_mcp_server("another-healthy-server", timestamp) + + healthy_server = MCPServer( + server_id="healthy-server", + name="healthy", + transport=MCPTransport.http, + updated_at=timestamp, + ) + another_healthy_server = MCPServer( + server_id="another-healthy-server", + name="another-healthy", + transport=MCPTransport.http, + updated_at=timestamp, + ) + + async def build_server(db_row): + if db_row.server_id == "bad-server": + raise RuntimeError("transient build failure") + if db_row.server_id == "healthy-server": + return healthy_server + return another_healthy_server + + mock_prisma = MagicMock() + mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock( + return_value=[healthy_row, bad_row, another_healthy_row] + ) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_prisma, + ), + patch.object( + manager, + "build_mcp_server_from_table", + AsyncMock(side_effect=build_server), + ), + patch.object(manager, "_maybe_register_openapi_tools", AsyncMock()), + caplog.at_level("ERROR", logger="LiteLLM"), + ): + await manager.reload_servers_from_database() + + assert set(manager.registry) == {"healthy-server", "another-healthy-server"} + assert manager.registry["healthy-server"] is healthy_server + assert manager.registry["another-healthy-server"] is another_healthy_server + assert "Skipping MCP server bad-server" in caplog.text + + @pytest.mark.asyncio + async def test_skips_server_when_openapi_registration_fails(self, caplog): + try: + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + except ImportError: + pytest.skip("MCP server not available") + + manager = MCPServerManager() + timestamp = datetime.utcnow() + healthy_row = _make_db_mcp_server("healthy-server", timestamp) + bad_openapi_row = _make_db_mcp_server("bad-openapi-server", timestamp) + existing_server = MCPServer( + server_id="existing-server", + name="existing", + transport=MCPTransport.http, + updated_at=timestamp, + ) + manager.registry = {existing_server.server_id: existing_server} + + healthy_server = MCPServer( + server_id="healthy-server", + name="healthy", + transport=MCPTransport.http, + updated_at=timestamp, + ) + bad_openapi_server = MCPServer( + server_id="bad-openapi-server", + name="bad-openapi", + transport=MCPTransport.http, + spec_path="https://example.invalid/openapi.json", + updated_at=timestamp, + ) + + async def build_server(db_row): + if db_row.server_id == "healthy-server": + return healthy_server + return bad_openapi_server + + observed_registries = [] + + async def register_openapi_tools(server, **kwargs): + observed_registries.append(set(manager.registry)) + assert kwargs == {"initialize_mapping": False} + if server.server_id == "bad-openapi-server": + raise RuntimeError("blocked address") + + mock_prisma = MagicMock() + mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock( + return_value=[healthy_row, bad_openapi_row] + ) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_prisma, + ), + patch.object( + manager, + "build_mcp_server_from_table", + AsyncMock(side_effect=build_server), + ), + patch.object( + manager, + "_maybe_register_openapi_tools", + AsyncMock(side_effect=register_openapi_tools), + ), + caplog.at_level("ERROR", logger="LiteLLM"), + ): + await manager.reload_servers_from_database() + + assert set(manager.registry) == {"healthy-server"} + assert manager.registry["healthy-server"] is healthy_server + assert observed_registries == [ + {"existing-server"}, + {"existing-server"}, + ] + assert "Skipping MCP server bad-openapi-server" in caplog.text + @pytest.mark.asyncio async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook(): @@ -2946,7 +3084,7 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): """ P1 Regression: list_tools path must apply _resolve_oauth2_flow to legacy DB rows where oauth2_flow is NULL but M2M credentials are present. - + Without this fix, has_client_credentials returns False and the caller's Authorization header is forwarded upstream instead of being blocked. """ @@ -3044,7 +3182,7 @@ async def test_call_tool_empty_extra_headers_returns_none(): """ P2 Regression: When all configured extra_headers are filtered out (e.g. Authorization for M2M), the resulting extra_headers should be None, not {}. - + Downstream code that checks `if extra_headers is None` will behave differently if an empty dict is passed instead. """ @@ -3071,7 +3209,10 @@ async def test_call_tool_empty_extra_headers_returns_none(): extra_headers=["Authorization"], # Will be filtered out for M2M ) - raw_headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + raw_headers = { + "Authorization": "Bearer sk-1234", + "Content-Type": "application/json", + } captured_extra_headers = None @@ -3108,8 +3249,8 @@ async def test_call_tool_empty_extra_headers_returns_none(): pass # We only care about the captured headers # With P2 fix: extra_headers should be None (not {}) when all headers filtered - assert captured_extra_headers is None, ( - "P2 API consistency issue: expected None for empty extra_headers, got: " - + str(captured_extra_headers) + assert ( + captured_extra_headers is None + ), "P2 API consistency issue: expected None for empty extra_headers, got: " + str( + captured_extra_headers ) - From c15718f9d1bcedd9ccd128d4112878a1278f2480 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 15:28:22 -0700 Subject: [PATCH 10/16] Fix Anthropic streaming reasoning token usage (#27319) * fix anthropic streaming reasoning token usage Co-authored-by: ishaan-berri * test anthropic streaming reasoning usage end to end Co-authored-by: ishaan-berri * address anthropic reasoning token text split Co-authored-by: ishaan-berri * harden anthropic reasoning usage for mocked tokens Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- litellm/llms/anthropic/chat/handler.py | 16 +- litellm/llms/anthropic/chat/transformation.py | 15 +- .../chat/test_anthropic_chat_handler.py | 287 ++++++++++++++++++ .../test_anthropic_chat_transformation.py | 28 ++ 4 files changed, 341 insertions(+), 5 deletions(-) diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 672de6a7054..2fb29b32a61 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -579,6 +579,10 @@ class ModelResponseIterator: # Accumulate compaction blocks for multi-turn reconstruction self.compaction_blocks: List[Dict[str, Any]] = [] + # Accumulate streamed thinking text so final usage can split reasoning + # tokens from regular output tokens. + self.reasoning_content_chunks: List[str] = [] + # Track server tool use inputs and results for code_interpreter_results self._server_tool_inputs: Dict[str, Any] = {} self.tool_results: List[Dict[str, Any]] = [] @@ -609,9 +613,14 @@ class ModelResponseIterator: return False def _handle_usage(self, anthropic_usage_chunk: Union[dict, UsageDelta]) -> Usage: + reasoning_content = ( + "".join(self.reasoning_content_chunks) + if self.reasoning_content_chunks + else None + ) return AnthropicConfig().calculate_usage( usage_object=cast(dict, anthropic_usage_chunk), - reasoning_content=None, + reasoning_content=reasoning_content, speed=self.speed, ) @@ -658,10 +667,13 @@ class ModelResponseIterator: "thinking" in content_block["delta"] or "signature" in content_block["delta"] ): + thinking_content = content_block["delta"].get("thinking") + if isinstance(thinking_content, str) and thinking_content: + self.reasoning_content_chunks.append(thinking_content) thinking_blocks = [ ChatCompletionThinkingBlock( type="thinking", - thinking=content_block["delta"].get("thinking") or "", + thinking=thinking_content or "", signature=str(content_block["delta"].get("signature") or ""), ) ] diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index dc3100f4670..2f11a3fccb5 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -2156,8 +2156,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): speed: Optional[str] = None, ) -> Usage: # NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this - prompt_tokens = usage_object.get("input_tokens", 0) or 0 - completion_tokens = usage_object.get("output_tokens", 0) or 0 + raw_prompt_tokens = usage_object.get("input_tokens", 0) or 0 + prompt_tokens: int = ( + int(raw_prompt_tokens) if isinstance(raw_prompt_tokens, (int, float)) else 0 + ) + raw_completion_tokens = usage_object.get("output_tokens", 0) or 0 + completion_tokens: int = ( + int(raw_completion_tokens) + if isinstance(raw_completion_tokens, (int, float)) + else 0 + ) _usage = usage_object cache_creation_input_tokens: int = 0 cache_read_input_tokens: int = 0 @@ -2226,11 +2234,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): text_tokens=raw_input_tokens, ) # Always populate completion_token_details, not just when there's reasoning_content - reasoning_tokens = ( + estimated_reasoning_tokens = ( token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0 ) + reasoning_tokens = min(estimated_reasoning_tokens, completion_tokens) completion_token_details = CompletionTokensDetailsWrapper( reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else 0, text_tokens=( diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index bf0461d89f1..2fdd639e74d 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -1,7 +1,11 @@ +import json +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from unittest.mock import AsyncMock, MagicMock import pytest +import litellm from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call from litellm.types.llms.openai import ( @@ -343,6 +347,289 @@ def test_text_only_streaming_has_index_zero(): ), f"Expected index=0, got {parsed.choices[0].index}" +def test_streaming_thinking_deltas_count_reasoning_tokens_in_usage(): + """Anthropic streaming usage should account for emitted thinking deltas.""" + chunks = [ + { + "type": "message_start", + "message": { + "id": "msg_123", + "type": "message", + "role": "assistant", + "content": [], + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "thinking", "thinking": ""}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "thinking_delta", + "thinking": "First I need to count the favorable outcomes. ", + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "thinking_delta", + "thinking": "Then I compare that count with all possible outcomes.", + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "signature_delta", "signature": "sig_123"}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "text", "text": ""}, + }, + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "text_delta", "text": "The probability is 3/8."}, + }, + {"type": "content_block_stop", "index": 1}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"output_tokens": 50}, + }, + ] + + iterator = ModelResponseIterator(None, sync_stream=True) + final_usage = None + reasoning_deltas = [] + + for chunk in chunks: + parsed = iterator.chunk_parser(chunk) + reasoning_content = getattr(parsed.choices[0].delta, "reasoning_content", None) + if reasoning_content: + reasoning_deltas.append(reasoning_content) + if parsed.usage is not None: + final_usage = parsed.usage + + assert reasoning_deltas == [ + "First I need to count the favorable outcomes. ", + "Then I compare that count with all possible outcomes.", + ] + assert final_usage is not None + completion_tokens_details = final_usage.completion_tokens_details + assert completion_tokens_details is not None + assert completion_tokens_details.reasoning_tokens > 0 + assert completion_tokens_details.text_tokens == ( + final_usage.completion_tokens - completion_tokens_details.reasoning_tokens + ) + + +def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinking(): + """The completion API should preserve Anthropic thinking usage in streaming mode.""" + thinking_parts = [ + "First I need to count the favorable outcomes. ", + "Then I compare that count with all possible outcomes.", + ] + thinking_text = "".join(thinking_parts) + answer_text = "The probability is 3/8." + requests_seen = [] + + class MockAnthropicHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, format, *args): # type: ignore[no-untyped-def] + return + + def do_POST(self): # type: ignore[no-untyped-def] + content_length = int(self.headers.get("content-length", "0")) + payload = json.loads(self.rfile.read(content_length).decode("utf-8")) + requests_seen.append( + { + "path": self.path, + "model": payload.get("model"), + "stream": payload.get("stream", False), + "thinking": payload.get("thinking"), + } + ) + + if payload.get("stream"): + events = [ + { + "type": "message_start", + "message": { + "id": "msg_mock", + "type": "message", + "role": "assistant", + "model": payload.get("model"), + "content": [], + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "thinking", "thinking": ""}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "thinking_delta", + "thinking": thinking_parts[0], + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "thinking_delta", + "thinking": thinking_parts[1], + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "signature_delta", + "signature": "sig_mock", + }, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "text", "text": ""}, + }, + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "text_delta", "text": answer_text}, + }, + {"type": "content_block_stop", "index": 1}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 50}, + }, + {"type": "message_stop"}, + ] + self._write_response( + content_type="text/event-stream", + body="".join( + f"data: {json.dumps(event)}\n\n" for event in events + ).encode("utf-8"), + ) + return + + self._write_response( + content_type="application/json", + body=json.dumps( + { + "id": "msg_mock", + "type": "message", + "role": "assistant", + "model": payload.get("model"), + "content": [ + { + "type": "thinking", + "thinking": thinking_text, + "signature": "sig_mock", + }, + {"type": "text", "text": answer_text}, + ], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 50}, + } + ).encode("utf-8"), + ) + + def _write_response(self, content_type: str, body: bytes) -> None: + self.send_response(200) + self.send_header("content-type", content_type) + self.send_header("content-length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + server = ThreadingHTTPServer(("127.0.0.1", 0), MockAnthropicHandler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + + try: + request_kwargs = { + "model": "anthropic/claude-sonnet-4-6", + "api_base": f"http://127.0.0.1:{server.server_port}", + "api_key": "test", + "messages": [ + { + "role": "user", + "content": "Solve a probability problem and show thinking.", + } + ], + "thinking": {"type": "adaptive"}, + "max_tokens": 128, + } + + non_stream_response = litellm.completion(**request_kwargs, stream=False) + non_stream_details = non_stream_response.usage.completion_tokens_details + assert non_stream_details is not None + assert non_stream_details.reasoning_tokens > 0 + + reasoning_chunks = [] + content_chunks = [] + stream_usage = None + for chunk in litellm.completion( + **request_kwargs, + stream=True, + stream_options={"include_usage": True}, + ): + chunk_dict = chunk.model_dump(exclude_none=True) + choices = chunk_dict.get("choices") or [] + if choices: + delta = choices[0].get("delta") or {} + if delta.get("reasoning_content"): + reasoning_chunks.append(delta["reasoning_content"]) + if delta.get("content"): + content_chunks.append(delta["content"]) + if chunk_dict.get("usage"): + stream_usage = chunk_dict["usage"] + + assert reasoning_chunks == thinking_parts + assert content_chunks == [answer_text] + assert stream_usage is not None + stream_completion_details = stream_usage["completion_tokens_details"] + assert ( + stream_completion_details["reasoning_tokens"] + == non_stream_details.reasoning_tokens + ) + assert stream_completion_details["text_tokens"] == ( + stream_usage["completion_tokens"] + - stream_completion_details["reasoning_tokens"] + ) + assert requests_seen == [ + { + "path": "/v1/messages", + "model": "claude-sonnet-4-6", + "stream": False, + "thinking": {"type": "adaptive"}, + }, + { + "path": "/v1/messages", + "model": "claude-sonnet-4-6", + "stream": True, + "thinking": {"type": "adaptive"}, + }, + ] + finally: + server.shutdown() + + def test_text_and_tool_streaming_has_index_zero(): """Test that mixed text and tool streaming responses have choice index=0""" chunks = [ 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 6f67d5417d2..e38698c9100 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 @@ -97,6 +97,34 @@ def test_calculate_usage(): assert usage._cache_read_input_tokens == 0 +def test_calculate_usage_clamps_text_tokens_when_reasoning_estimate_exceeds_output(): + config = AnthropicConfig() + + usage = config.calculate_usage( + usage_object={"input_tokens": 10, "output_tokens": 1}, + reasoning_content="This reasoning text intentionally tokenizes above one output token.", + ) + + assert usage.completion_tokens == 1 + assert usage.completion_tokens_details is not None + assert usage.completion_tokens_details.reasoning_tokens == usage.completion_tokens + assert usage.completion_tokens_details.text_tokens == 0 + + +def test_calculate_usage_handles_mocked_output_tokens_with_reasoning_content(): + config = AnthropicConfig() + + usage = config.calculate_usage( + usage_object={"input_tokens": 10, "output_tokens": MagicMock()}, + reasoning_content="mocked response reasoning", + ) + + assert usage.completion_tokens == 0 + assert usage.completion_tokens_details is not None + assert usage.completion_tokens_details.reasoning_tokens == 0 + assert usage.completion_tokens_details.text_tokens == 0 + + @pytest.mark.parametrize( "usage_object,expected_usage", [ From aba131d3cf43a2aff3ae466be18cb7ffc1efa33f Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 15:32:55 -0700 Subject: [PATCH 11/16] fix: Vertex Anthropic streaming status error hangs (#27310) * Fix streaming HTTP status error hangs Co-authored-by: ishaan-berri * Fix sync streaming HTTP status error hangs Co-authored-by: ishaan-berri * Cap sync streaming error read workers Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- litellm/llms/custom_httpx/http_handler.py | 53 +++++++-- .../llms/custom_httpx/test_http_handler.py | 102 ++++++++++++++++++ 2 files changed, 149 insertions(+), 6 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index dd955c23d23..af18c666679 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1,4 +1,5 @@ import asyncio +import concurrent.futures import inspect import os import socket @@ -133,6 +134,11 @@ _DEFAULT_TIMEOUT = httpx.Timeout( timeout=COMPLETION_HTTP_FALLBACK_SECONDS, connect=HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS, ) +_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS = 5.0 +_STREAMING_ERROR_BODY_READ_EXECUTOR = concurrent.futures.ThreadPoolExecutor( + max_workers=50, + thread_name_prefix="litellm-streaming-error-body-read", +) def _prepare_request_data_and_content( @@ -386,17 +392,30 @@ def _safe_get_response_text(response: httpx.Response) -> str: return "" -async def _safe_aread_response(response: httpx.Response) -> bytes: +async def _safe_aread_response( + response: httpx.Response, timeout: Optional[float] = None +) -> bytes: """Safely read async response body, falling back to empty bytes on errors.""" try: + if timeout is not None: + return await asyncio.wait_for(response.aread(), timeout=timeout) return await response.aread() except Exception: return b"" -def _safe_read_response(response: httpx.Response) -> bytes: +def _safe_read_response( + response: httpx.Response, timeout: Optional[float] = None +) -> bytes: """Safely read sync response body, falling back to empty bytes on errors.""" try: + if timeout is not None: + future = _STREAMING_ERROR_BODY_READ_EXECUTOR.submit(response.read) + try: + return future.result(timeout=timeout) + except Exception: + response.close() + return b"" return response.read() except Exception: return b"" @@ -405,8 +424,19 @@ def _safe_read_response(response: httpx.Response) -> bytes: def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None: """Raise a MaskedHTTPStatusError for sync HTTP handlers.""" if stream: - _body = mask_sensitive_info(_safe_read_response(e.response)) - raise MaskedHTTPStatusError(e, message=_body, text=_body) from None + try: + _body = mask_sensitive_info( + _safe_read_response( + e.response, + timeout=_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS, + ) + ) + raise MaskedHTTPStatusError(e, message=_body, text=_body) from None + finally: + try: + e.response.close() + except Exception: + pass _text = mask_sensitive_info(_safe_get_response_text(e.response)) raise MaskedHTTPStatusError(e, message=_text, text=_text) from None @@ -414,8 +444,19 @@ def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None: async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> None: """Raise a MaskedHTTPStatusError for async HTTP handlers.""" if stream: - _body = mask_sensitive_info(await _safe_aread_response(e.response)) - raise MaskedHTTPStatusError(e, message=_body, text=_body) from None + try: + _body = mask_sensitive_info( + await _safe_aread_response( + e.response, + timeout=_STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS, + ) + ) + raise MaskedHTTPStatusError(e, message=_body, text=_body) from None + finally: + try: + await e.response.aclose() + except Exception: + pass _text = mask_sensitive_info(_safe_get_response_text(e.response)) raise MaskedHTTPStatusError(e, message=_text, text=_text) from None diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index dd52304a703..26f50e8e492 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -1,8 +1,10 @@ +import asyncio import io import os import pathlib import ssl import sys +import threading from unittest.mock import MagicMock, patch import certifi @@ -18,11 +20,111 @@ from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, + MaskedHTTPStatusError, _get_httpx_client, get_ssl_configuration, ) +@pytest.mark.asyncio +async def test_async_post_streaming_status_error_should_not_wait_forever_for_body( + monkeypatch, +): + """ + Vertex Anthropic streamRawPredict can return a pre-stream 4xx where the + streamed error body never terminates. The handler must still surface the + status promptly instead of blocking the downstream client. + """ + + class HangingErrorStream(httpx.AsyncByteStream): + async def __aiter__(self): + await asyncio.Event().wait() + if False: + yield b"" + + async def mock_handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + request=request, + headers={"content-type": "application/json"}, + stream=HangingErrorStream(), + ) + + monkeypatch.setattr( + "litellm.llms.custom_httpx.http_handler._STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS", + 0.01, + ) + + litellm_handler = AsyncHTTPHandler() + await litellm_handler.client.aclose() + litellm_handler.client = httpx.AsyncClient( + transport=httpx.MockTransport(mock_handler) + ) + try: + with pytest.raises(MaskedHTTPStatusError) as exc_info: + await asyncio.wait_for( + litellm_handler.post( + "https://vertex.example/streamRawPredict", + stream=True, + ), + timeout=0.2, + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.response.status_code == 400 + finally: + await litellm_handler.close() + + +def test_sync_post_streaming_status_error_should_not_wait_forever_for_body( + monkeypatch, +): + """ + Keep the sync streaming error path aligned with the async path so a + non-terminating streamed error body cannot block a worker thread forever. + """ + + class HangingSyncErrorStream(httpx.SyncByteStream): + def __init__(self): + self.closed_event = threading.Event() + + def __iter__(self): + self.closed_event.wait() + if False: + yield b"" + + def close(self): + self.closed_event.set() + + def mock_handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 400, + request=request, + headers={"content-type": "application/json"}, + stream=HangingSyncErrorStream(), + ) + + monkeypatch.setattr( + "litellm.llms.custom_httpx.http_handler._STREAMING_ERROR_BODY_READ_TIMEOUT_SECONDS", + 0.01, + ) + + litellm_handler = HTTPHandler() + litellm_handler.client.close() + litellm_handler.client = httpx.Client(transport=httpx.MockTransport(mock_handler)) + try: + with pytest.raises(MaskedHTTPStatusError) as exc_info: + litellm_handler.post( + "https://vertex.example/streamRawPredict", + stream=True, + ) + + assert exc_info.value.status_code == 400 + assert exc_info.value.response.status_code == 400 + finally: + litellm_handler.close() + + @pytest.mark.asyncio async def test_ssl_security_level(monkeypatch): # Ensure aiohttp transport is enabled for this test From b318231fe9d8b2b98c61fe8ee7b2babc98939211 Mon Sep 17 00:00:00 2001 From: oss-agent-shin Date: Wed, 6 May 2026 15:50:06 -0700 Subject: [PATCH 12/16] Add Azure Sentinel audit log support (#27280) * Add Azure Sentinel audit log callback support Co-authored-by: ishaan-berri * Fix Azure Sentinel audit log batching Co-authored-by: ishaan-berri * Fix Azure Sentinel CI checks Co-authored-by: ishaan-berri --------- Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- .../azure_sentinel/azure_sentinel.py | 139 ++++++++++--- .../custom_logger_registry.py | 2 + litellm/proxy/_types.py | 13 ++ .../integrations/test_azure_sentinel.py | 186 +++++++++++++++++- 4 files changed, 305 insertions(+), 35 deletions(-) diff --git a/litellm/integrations/azure_sentinel/azure_sentinel.py b/litellm/integrations/azure_sentinel/azure_sentinel.py index dd508e6c6c2..0cfd49cda37 100644 --- a/litellm/integrations/azure_sentinel/azure_sentinel.py +++ b/litellm/integrations/azure_sentinel/azure_sentinel.py @@ -14,16 +14,18 @@ For batching specific details see CustomBatchLogger class import asyncio import os +import time import traceback -from typing import List, Optional +from typing import List, Optional, Union from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.types.utils import StandardLoggingPayload +from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload class AzureSentinelLogger(CustomBatchLogger): @@ -39,6 +41,7 @@ class AzureSentinelLogger(CustomBatchLogger): tenant_id: Optional[str] = None, client_id: Optional[str] = None, client_secret: Optional[str] = None, + audit_stream_name: Optional[str] = None, **kwargs, ): """ @@ -57,57 +60,77 @@ class AzureSentinelLogger(CustomBatchLogger): If not provided, will use AZURE_SENTINEL_CLIENT_ID or AZURE_CLIENT_ID env var. client_secret (str, optional): Azure Client Secret for OAuth2 authentication. If not provided, will use AZURE_SENTINEL_CLIENT_SECRET or AZURE_CLIENT_SECRET env var. + audit_stream_name (str, optional): Stream name from DCR for audit logs. + If not provided, audit logs use the standard stream name. """ self.async_httpx_client = get_async_httpx_client( llm_provider=httpxSpecialProvider.LoggingCallback ) - self.dcr_immutable_id = dcr_immutable_id or os.getenv( + resolved_dcr_immutable_id = dcr_immutable_id or os.getenv( "AZURE_SENTINEL_DCR_IMMUTABLE_ID" ) - self.stream_name = stream_name or os.getenv( - "AZURE_SENTINEL_STREAM_NAME", "Custom-LiteLLM" + resolved_stream_name = ( + stream_name or os.getenv("AZURE_SENTINEL_STREAM_NAME") or "Custom-LiteLLM" ) - self.endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT") - self.tenant_id = ( + resolved_audit_stream_name = audit_stream_name or resolved_stream_name + resolved_endpoint = endpoint or os.getenv("AZURE_SENTINEL_ENDPOINT") + resolved_tenant_id = ( tenant_id or os.getenv("AZURE_SENTINEL_TENANT_ID") or os.getenv("AZURE_TENANT_ID") ) - self.client_id = ( + resolved_client_id = ( client_id or os.getenv("AZURE_SENTINEL_CLIENT_ID") or os.getenv("AZURE_CLIENT_ID") ) - self.client_secret = ( + resolved_client_secret = ( client_secret or os.getenv("AZURE_SENTINEL_CLIENT_SECRET") or os.getenv("AZURE_CLIENT_SECRET") ) - if not self.dcr_immutable_id: + if not resolved_dcr_immutable_id: raise ValueError( "AZURE_SENTINEL_DCR_IMMUTABLE_ID is required. Set it as an environment variable or pass dcr_immutable_id parameter." ) - if not self.endpoint: + if not resolved_endpoint: raise ValueError( "AZURE_SENTINEL_ENDPOINT is required. Set it as an environment variable or pass endpoint parameter." ) - if not self.tenant_id: + if not resolved_tenant_id: raise ValueError( "AZURE_SENTINEL_TENANT_ID or AZURE_TENANT_ID is required. Set it as an environment variable or pass tenant_id parameter." ) - if not self.client_id: + if not resolved_client_id: raise ValueError( "AZURE_SENTINEL_CLIENT_ID or AZURE_CLIENT_ID is required. Set it as an environment variable or pass client_id parameter." ) - if not self.client_secret: + if not resolved_client_secret: raise ValueError( "AZURE_SENTINEL_CLIENT_SECRET or AZURE_CLIENT_SECRET is required. Set it as an environment variable or pass client_secret parameter." ) + self.dcr_immutable_id = resolved_dcr_immutable_id + self.stream_name = resolved_stream_name + self.audit_stream_name = resolved_audit_stream_name + self.endpoint = resolved_endpoint + self.tenant_id = resolved_tenant_id + self.client_id = resolved_client_id + self.client_secret = resolved_client_secret + # Build API endpoint: {Endpoint}/dataCollectionRules/{DCR Immutable ID}/streams/{Stream Name}?api-version=2023-01-01 - self.api_endpoint = f"{self.endpoint.rstrip('/')}/dataCollectionRules/{self.dcr_immutable_id}/streams/{self.stream_name}?api-version=2023-01-01" + self.api_endpoint = self._build_api_endpoint( + endpoint=resolved_endpoint, + dcr_immutable_id=resolved_dcr_immutable_id, + stream_name=resolved_stream_name, + ) + self.audit_api_endpoint = self._build_api_endpoint( + endpoint=resolved_endpoint, + dcr_immutable_id=resolved_dcr_immutable_id, + stream_name=resolved_audit_stream_name, + ) # OAuth2 scope for Azure Monitor self.oauth_scope = "https://monitor.azure.com/.default" @@ -118,6 +141,13 @@ class AzureSentinelLogger(CustomBatchLogger): super().__init__(**kwargs, flush_lock=self.flush_lock) asyncio.create_task(self.periodic_flush()) self.log_queue: List[StandardLoggingPayload] = [] + self.audit_log_queue: List[StandardAuditLogPayload] = [] + + @staticmethod + def _build_api_endpoint( + endpoint: str, dcr_immutable_id: str, stream_name: str + ) -> str: + return f"{endpoint.rstrip('/')}/dataCollectionRules/{dcr_immutable_id}/streams/{stream_name}?api-version=2023-01-01" async def _get_oauth_token(self) -> str: """ @@ -126,9 +156,6 @@ class AzureSentinelLogger(CustomBatchLogger): Returns: Bearer token string """ - # Check if we have a valid cached token - import time - if ( self.oauth_token and self.oauth_token_expires_at @@ -170,9 +197,6 @@ class AzureSentinelLogger(CustomBatchLogger): if not self.oauth_token: raise Exception("OAuth2 token response did not contain access_token") - # Cache token expiry time - import time - self.oauth_token_expires_at = time.time() + expires_in return self.oauth_token @@ -246,6 +270,34 @@ class AzureSentinelLogger(CustomBatchLogger): ) pass + async def async_log_audit_log_event( + self, audit_log: StandardAuditLogPayload + ) -> None: + """ + Async log LiteLLM audit log events to Azure Sentinel. + + Audit logs are queued separately from standard LLM logs so mixed callback + usage never sends schema-mismatched records in the same ingestion batch. + """ + try: + verbose_logger.debug( + "Azure Sentinel: Logging audit event id=%s action=%s table=%s", + audit_log.get("id"), + audit_log.get("action"), + audit_log.get("table_name"), + ) + + self.audit_log_queue.append(audit_log) + + if len(self.audit_log_queue) >= self.batch_size: + await self.async_send_audit_batch() + + except Exception as e: + verbose_logger.exception( + f"Azure Sentinel Audit Log Layer Error - {str(e)}\n{traceback.format_exc()}" + ) + pass + async def async_send_batch(self): """ Sends the batch of logs to Azure Monitor Logs Ingestion API @@ -253,22 +305,42 @@ class AzureSentinelLogger(CustomBatchLogger): Raises: Raises a NON Blocking verbose_logger.exception if an error occurs """ + await self._async_send_batch_to_api( + log_queue=self.log_queue, + api_endpoint=self.api_endpoint, + log_type="logs", + ) + + async def async_send_audit_batch(self): + """ + Sends the batch of audit logs to Azure Monitor Logs Ingestion API + """ + await self._async_send_batch_to_api( + log_queue=self.audit_log_queue, + api_endpoint=self.audit_api_endpoint, + log_type="audit logs", + ) + + async def _async_send_batch_to_api( + self, + log_queue: List[Union[StandardLoggingPayload, StandardAuditLogPayload]], + api_endpoint: str, + log_type: str, + ) -> None: try: - if not self.log_queue: + if not log_queue: return verbose_logger.debug( - "Azure Sentinel - about to flush %s events", len(self.log_queue) + "Azure Sentinel - about to flush %s %s", len(log_queue), log_type ) - from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - # Get OAuth2 token bearer_token = await self._get_oauth_token() # Convert log queue to JSON array format expected by Logs Ingestion API # Each log entry should be a JSON object in the array - body = safe_dumps(self.log_queue) + body = safe_dumps(log_queue) # Set headers for Logs Ingestion API headers = { @@ -278,7 +350,7 @@ class AzureSentinelLogger(CustomBatchLogger): # Send the request response = await self.async_httpx_client.post( - url=self.api_endpoint, data=body.encode("utf-8"), headers=headers + url=api_endpoint, data=body.encode("utf-8"), headers=headers ) if response.status_code not in [200, 204]: @@ -301,4 +373,15 @@ class AzureSentinelLogger(CustomBatchLogger): f"Azure Sentinel Error sending batch API - {str(e)}\n{traceback.format_exc()}" ) finally: - self.log_queue.clear() + log_queue.clear() + + async def flush_queue(self): + if self.flush_lock is None: + return + + async with self.flush_lock: + if self.log_queue: + await self.async_send_batch() + if self.audit_log_queue: + await self.async_send_audit_batch() + self.last_flush_time = time.time() diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index f873bfeece5..fd402b90d88 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -14,6 +14,7 @@ from litellm import _custom_logger_compatible_callbacks_literal from litellm.integrations.agentops import AgentOps from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook from litellm.integrations.argilla import ArgillaLogger +from litellm.integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLogger from litellm.integrations.bitbucket import BitBucketPromptManager from litellm.integrations.braintrust_logging import BraintrustLogger @@ -73,6 +74,7 @@ class CustomLoggerRegistry: "opik": OpikLogger, "argilla": ArgillaLogger, "opentelemetry": OpenTelemetry, + "azure_sentinel": AzureSentinelLogger, "azure_storage": AzureBlobStorageLogger, "humanloop": HumanloopLogger, # OTEL compatible loggers diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c6653a722d6..2c976479798 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3348,6 +3348,19 @@ class AllCallbacks(LiteLLMPydanticObjectBase): ], ) + azure_sentinel: CallbackOnUI = CallbackOnUI( + litellm_callback_name="azure_sentinel", + ui_callback_name="Azure Sentinel", + litellm_callback_params=[ + "AZURE_SENTINEL_DCR_IMMUTABLE_ID", + "AZURE_SENTINEL_ENDPOINT", + "AZURE_SENTINEL_TENANT_ID", + "AZURE_SENTINEL_CLIENT_ID", + "AZURE_SENTINEL_CLIENT_SECRET", + "AZURE_SENTINEL_STREAM_NAME", + ], + ) + openmeter: CallbackOnUI = CallbackOnUI( litellm_callback_name="openmeter", ui_callback_name="OpenMeter", diff --git a/tests/test_litellm/integrations/test_azure_sentinel.py b/tests/test_litellm/integrations/test_azure_sentinel.py index 031b85211f7..30b246202fc 100644 --- a/tests/test_litellm/integrations/test_azure_sentinel.py +++ b/tests/test_litellm/integrations/test_azure_sentinel.py @@ -2,13 +2,18 @@ Test Azure Sentinel logging integration """ -import datetime -from unittest.mock import AsyncMock, patch +import json +from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger -from litellm.types.utils import StandardLoggingPayload +from litellm.types.utils import StandardAuditLogPayload, StandardLoggingPayload + + +def _close_periodic_flush_task(coro): + coro.close() + return None @pytest.mark.asyncio @@ -20,7 +25,7 @@ async def test_azure_sentinel_oauth_and_send_batch(): test_client_id = "test-client-id" test_client_secret = "test-client-secret" - with patch("asyncio.create_task"): + with patch("asyncio.create_task", side_effect=_close_periodic_flush_task): logger = AzureSentinelLogger( dcr_immutable_id=test_dcr_id, endpoint=test_endpoint, @@ -42,9 +47,6 @@ async def test_azure_sentinel_oauth_and_send_batch(): # Add to queue logger.log_queue.append(standard_payload) - # Mock OAuth token response - from unittest.mock import MagicMock - mock_token_response = MagicMock() mock_token_response.status_code = 200 mock_token_response.json = MagicMock( @@ -91,3 +93,173 @@ async def test_azure_sentinel_oauth_and_send_batch(): # Verify queue is cleared assert len(logger.log_queue) == 0 + + +@pytest.mark.asyncio +async def test_azure_sentinel_queues_audit_log_event(): + """Test that Azure Sentinel supports direct audit log callbacks""" + with patch("asyncio.create_task", side_effect=_close_periodic_flush_task): + logger = AzureSentinelLogger( + dcr_immutable_id="dcr-test123456789", + endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com", + tenant_id="test-tenant-id", + client_id="test-client-id", + client_secret="test-client-secret", + ) + + logger.batch_size = 2 + logger.async_send_audit_batch = AsyncMock() + + audit_log = StandardAuditLogPayload( + id="audit-123", + updated_at="2026-05-06T04:39:00+00:00", + changed_by="user-1", + changed_by_api_key="sk-test", + action="created", + table_name="LiteLLM_TeamTable", + object_id="team-1", + before_value=None, + updated_values='{"team_alias": "sentinel-demo"}', + ) + + await logger.async_log_audit_log_event(audit_log) + + assert logger.audit_log_queue == [audit_log] + logger.async_send_audit_batch.assert_not_called() + + await logger.async_log_audit_log_event(audit_log) + + assert logger.audit_log_queue == [audit_log, audit_log] + logger.async_send_audit_batch.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_azure_sentinel_sends_audit_log_payload_to_ingestion_api(): + """Test that queued audit logs are sent to Azure Monitor Logs Ingestion""" + with patch("asyncio.create_task", side_effect=_close_periodic_flush_task): + logger = AzureSentinelLogger( + dcr_immutable_id="dcr-test123456789", + endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com", + tenant_id="test-tenant-id", + client_id="test-client-id", + client_secret="test-client-secret", + ) + + audit_log = StandardAuditLogPayload( + id="audit-123", + updated_at="2026-05-06T04:39:00+00:00", + changed_by="user-1", + changed_by_api_key="sk-test", + action="created", + table_name="LiteLLM_TeamTable", + object_id="team-1", + before_value=None, + updated_values='{"team_alias": "sentinel-demo"}', + ) + await logger.async_log_audit_log_event(audit_log) + + mock_token_response = MagicMock() + mock_token_response.status_code = 200 + mock_token_response.json = MagicMock( + return_value={ + "access_token": "test-bearer-token", + "expires_in": 3600, + } + ) + mock_token_response.text = "Success" + + mock_api_response = MagicMock() + mock_api_response.status_code = 204 + mock_api_response.text = "Success" + + async def mock_post(*args, **kwargs): + if "oauth2/v2.0/token" in kwargs.get("url", ""): + return mock_token_response + return mock_api_response + + logger.async_httpx_client.post = AsyncMock(side_effect=mock_post) + + await logger.flush_queue() + + api_call_args = logger.async_httpx_client.post.call_args_list[-1] + body = json.loads(api_call_args.kwargs["data"].decode("utf-8")) + assert body == [audit_log] + assert "dcr-test123456789" in api_call_args.kwargs["url"] + assert "Custom-LiteLLM" in api_call_args.kwargs["url"] + assert len(logger.audit_log_queue) == 0 + + +@pytest.mark.asyncio +async def test_azure_sentinel_flushes_standard_and_audit_logs_separately(): + """Test mixed callback roles do not send schema-mismatched batches.""" + with patch("asyncio.create_task", side_effect=_close_periodic_flush_task): + logger = AzureSentinelLogger( + dcr_immutable_id="dcr-test123456789", + stream_name="Custom-LiteLLM-Standard", + audit_stream_name="Custom-LiteLLM-Audit", + endpoint="https://test-dce.eastus-1.ingest.monitor.azure.com", + tenant_id="test-tenant-id", + client_id="test-client-id", + client_secret="test-client-secret", + ) + + standard_payload = StandardLoggingPayload( + id="standard-123", + call_type="completion", + model="gpt-3.5-turbo", + status="success", + messages=[{"role": "user", "content": "Hello"}], + response={"choices": [{"message": {"content": "Hi"}}]}, + ) + audit_log = StandardAuditLogPayload( + id="audit-123", + updated_at="2026-05-06T04:39:00+00:00", + changed_by="user-1", + changed_by_api_key="sk-test", + action="created", + table_name="LiteLLM_TeamTable", + object_id="team-1", + before_value=None, + updated_values='{"team_alias": "sentinel-demo"}', + ) + + logger.log_queue.append(standard_payload) + await logger.async_log_audit_log_event(audit_log) + + mock_token_response = MagicMock() + mock_token_response.status_code = 200 + mock_token_response.json = MagicMock( + return_value={ + "access_token": "test-bearer-token", + "expires_in": 3600, + } + ) + mock_token_response.text = "Success" + + mock_api_response = MagicMock() + mock_api_response.status_code = 204 + mock_api_response.text = "Success" + + async def mock_post(*args, **kwargs): + if "oauth2/v2.0/token" in kwargs.get("url", ""): + return mock_token_response + return mock_api_response + + logger.async_httpx_client.post = AsyncMock(side_effect=mock_post) + + await logger.flush_queue() + + ingestion_calls = [ + call + for call in logger.async_httpx_client.post.call_args_list + if "dataCollectionRules" in call.kwargs["url"] + ] + assert len(ingestion_calls) == 2 + + standard_call, audit_call = ingestion_calls + assert "Custom-LiteLLM-Standard" in standard_call.kwargs["url"] + assert json.loads(standard_call.kwargs["data"].decode("utf-8")) == [ + standard_payload + ] + assert "Custom-LiteLLM-Audit" in audit_call.kwargs["url"] + assert json.loads(audit_call.kwargs["data"].decode("utf-8")) == [audit_log] From a3a42c6c479d46733cef96a4de1bfd5fd4ca6acd Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 6 May 2026 16:34:45 -0700 Subject: [PATCH 13/16] [Chore] CI: Assign test_request_size_limit_middleware To Proxy-Runtime Shard (#27341) The assert-shard-coverage guard in test-unit-proxy-db.yml failed because test_request_size_limit_middleware.py was added under tests/proxy_unit_tests/ but not referenced by any matrix entry. Assigning it to the proxy-runtime shard, which already covers other server-runtime tests (proxy_routes, proxy_gunicorn, server_root_path). --- .github/workflows/test-unit-proxy-db.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index d5781f767f6..5a9688db9c4 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -141,6 +141,7 @@ jobs: tests/proxy_unit_tests/test_server_root_path.py tests/proxy_unit_tests/test_proxy_pass_user_config.py tests/proxy_unit_tests/test_proxy_token_counter.py + tests/proxy_unit_tests/test_request_size_limit_middleware.py workers: 4 dist: loadscope timeout: 15 From f1c91d754dc3658c8a35622951451ce5118b43cf Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 6 May 2026 16:41:50 -0700 Subject: [PATCH 14/16] [Chore] CI: Block PRs that drop overall code coverage (#27340) * [Chore] CI: Block PRs that drop overall code coverage Tighten Codecov project status threshold from 1% to 0% so any drop in overall project coverage relative to the base commit fails the codecov/project check. target: auto keeps the bar floating with the codebase, no manual maintenance needed as coverage moves up over time. * [Chore] CI: Always post Codecov status regardless of CI outcome Set codecov.require_ci_to_pass: false and codecov.notify.wait_for_ci: false so Codecov posts the codecov/project and codecov/patch checks as soon as the expected uploads arrive, instead of withholding them when unrelated CI jobs fail. The coverage-regression check is independent of test pass/fail, and CI failures are already enforced by their own required-status checks. --- codecov.yaml | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/codecov.yaml b/codecov.yaml index 09fccc6b995..8609d3143d6 100644 --- a/codecov.yaml +++ b/codecov.yaml @@ -1,3 +1,8 @@ +codecov: + require_ci_to_pass: false # post coverage status even if CI has unrelated failures + notify: + wait_for_ci: false # post as soon as expected uploads arrive, don't wait on CI + component_management: individual_components: - component_id: "Router" @@ -28,7 +33,7 @@ coverage: project: default: target: auto - threshold: 1% # at maximum allow project coverage to drop by 1% + threshold: 0% # do not allow project coverage to drop patch: default: target: auto From 854456f58ef68d7f772e13c6a9266bfe45a5277f Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 6 May 2026 17:22:20 -0700 Subject: [PATCH 15/16] Fix Prometheus remaining metric zero values (#27348) Co-authored-by: oss-agent-shin <279349115+oss-agent-shin@users.noreply.github.com> Co-authored-by: ishaan-berri --- litellm/integrations/prometheus.py | 12 +++---- ...prometheus_custom_metadata_label_counts.py | 34 +++++++++++++++++++ 2 files changed, 40 insertions(+), 6 deletions(-) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 65b8a9fc4b1..f9b1c666439 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -1443,12 +1443,12 @@ class PrometheusLogger(CustomLogger): ) remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}" - remaining_requests = ( - metadata.get(remaining_requests_variable_name, sys.maxsize) or sys.maxsize - ) - remaining_tokens = ( - metadata.get(remaining_tokens_variable_name, sys.maxsize) or sys.maxsize - ) + remaining_requests = metadata.get(remaining_requests_variable_name) + if remaining_requests is None: + remaining_requests = sys.maxsize + remaining_tokens = metadata.get(remaining_tokens_variable_name) + if remaining_tokens is None: + remaining_tokens = sys.maxsize enum_values = UserAPIKeyLabelValues( hashed_api_key=user_api_key, diff --git a/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py b/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py index 05d1eab8cc8..99eb5abb7b5 100644 --- a/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py +++ b/tests/test_litellm/integrations/test_prometheus_custom_metadata_label_counts.py @@ -1,4 +1,5 @@ import logging +import sys import pytest from prometheus_client import REGISTRY @@ -123,3 +124,36 @@ def test_virtual_key_rate_limit_metrics_accept_custom_metadata_labels( and sample.value == 3 for sample in samples ) + + +def test_virtual_key_rate_limit_metrics_preserve_zero_remaining_values( + monkeypatch: pytest.MonkeyPatch, +): + prometheus_logger = _create_prometheus_logger_with_custom_labels(monkeypatch) + metadata = { + "model_group": "gpt-4o-mini", + "litellm-key-remaining-requests-gpt-4o-mini": 0, + "litellm-key-remaining-tokens-gpt-4o-mini": 0, + } + kwargs = { + "litellm_params": { + "metadata": metadata, + }, + "standard_logging_object": _standard_logging_payload_with_requester_metadata(), + } + + prometheus_logger._set_virtual_key_rate_limit_metrics( + user_api_key="test-hash", + user_api_key_alias="test-alias", + kwargs=kwargs, + metadata=metadata, + model_id="model-123", + ) + + request_samples = _metric_samples("litellm_remaining_api_key_requests_for_model") + token_samples = _metric_samples("litellm_remaining_api_key_tokens_for_model") + + assert any(sample.value == 0 for sample in request_samples) + assert any(sample.value == 0 for sample in token_samples) + assert not any(sample.value == sys.maxsize for sample in request_samples) + assert not any(sample.value == sys.maxsize for sample in token_samples) From a67b7a7e87f11bed01f9e073125a7f8f180105a2 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 6 May 2026 17:39:38 -0700 Subject: [PATCH 16/16] Refactor Bedrock response stream shape handling (#27257) * Refactor Bedrock response stream shape handling - Introduced a module-level constant `BEDROCK_RESPONSE_STREAM_SHAPE` to cache the response stream shape, eliminating the need for per-instance caching in `BedrockEventStreamDecoderBase`. - Updated relevant methods to utilize the new constant, improving performance by avoiding redundant loading of the shape. - Added tests to ensure the shape is loaded correctly at import time and is consistent across different modules. - Added a new mock server script for testing Bedrock pass-through functionality. * Refactor response parsing for Bedrock and SageMaker - Improved code readability by formatting the parsing method calls in `AWSEventStreamDecoder` for both Bedrock and SageMaker response stream shapes. - Added blank lines for better separation of code blocks in `invoke_handler.py` and `common_utils.py` to enhance maintainability. * Enhance error handling for Bedrock and SageMaker response stream shape loading - Wrapped the loading logic in `_load_bedrock_response_stream_shape` and `_load_sagemaker_response_stream_shape` with try-except blocks to gracefully handle exceptions. - Added logging to warn when the response stream shape cannot be pre-loaded, ensuring the module imports cleanly. - Updated tests to verify that loading failures return `None` instead of propagating exceptions. * Implement error handling for missing response stream shapes in Bedrock and SageMaker - Added checks in `_parse_message_from_event` methods to raise appropriate errors when `BEDROCK_RESPONSE_STREAM_SHAPE` or `SAGEMAKER_RESPONSE_STREAM_SHAPE` is None, ensuring clearer error reporting. - Updated logging messages to reflect the unavailability of event-stream decoding for both Bedrock and SageMaker. - Enhanced unit tests to verify that the correct exceptions are raised when the response stream shapes are not loaded. --- .gitignore | 3 +- .../chat/invoke_agent/transformation.py | 24 +- litellm/llms/bedrock/chat/invoke_handler.py | 34 +-- litellm/llms/bedrock/common_utils.py | 58 ++-- litellm/llms/sagemaker/common_utils.py | 51 ++-- scripts/mock_bedrock_passthrough_target.py | 276 ++++++++++++++++++ .../llms/bedrock/test_bedrock_common_utils.py | 104 +++++++ .../sagemaker/test_sagemaker_common_utils.py | 96 ++++++ 8 files changed, 568 insertions(+), 78 deletions(-) create mode 100644 scripts/mock_bedrock_passthrough_target.py diff --git a/.gitignore b/.gitignore index 59812ed6ed4..20355a8e4ef 100644 --- a/.gitignore +++ b/.gitignore @@ -100,4 +100,5 @@ STABILIZATION_TODO.md **/playwright-report **/*.storageState.json **/coverage -test-config \ No newline at end of file +test-config +.vscode \ No newline at end of file diff --git a/litellm/llms/bedrock/chat/invoke_agent/transformation.py b/litellm/llms/bedrock/chat/invoke_agent/transformation.py index 4c667b0ce39..e4072c24557 100644 --- a/litellm/llms/bedrock/chat/invoke_agent/transformation.py +++ b/litellm/llms/bedrock/chat/invoke_agent/transformation.py @@ -299,29 +299,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM): ) def _get_response_stream_shape(self): - """Get the response stream shape for parsing, reusing existing logic.""" - try: - # Try to reuse the cached shape from the existing decoder - from litellm.llms.bedrock.chat.invoke_handler import ( - get_response_stream_shape, - ) + from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE - return get_response_stream_shape() - except ImportError: - # Fallback: create our own shape - try: - from botocore.loaders import Loader - from botocore.model import ServiceModel - - loader = Loader() - bedrock_service_dict = loader.load_service_model( - "bedrock-runtime", "service-2" - ) - bedrock_service_model = ServiceModel(bedrock_service_dict) - return bedrock_service_model.shape_for("ResponseStream") - except Exception as e: - verbose_logger.warning(f"Could not load response stream shape: {e}") - return None + return BEDROCK_RESPONSE_STREAM_SHAPE def _extract_response_content(self, events: InvokeAgentEventList) -> str: """Extract the final response content from parsed events.""" diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 9dfada7c418..92ca75db95b 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -67,9 +67,13 @@ from litellm.types.utils import ( from litellm.utils import CustomStreamWrapper, get_secret from ..base_aws_llm import BaseAWSLLM -from ..common_utils import BedrockError, ModelResponseIterator, get_bedrock_tool_name +from ..common_utils import ( + BEDROCK_RESPONSE_STREAM_SHAPE, + BedrockError, + ModelResponseIterator, + get_bedrock_tool_name, +) -_response_stream_shape_cache = None bedrock_tool_name_mappings: InMemoryCache = InMemoryCache( max_size_in_memory=50, default_ttl=600 ) @@ -1391,20 +1395,6 @@ class BedrockLLM(BaseAWSLLM): return None -def get_response_stream_shape(): - global _response_stream_shape_cache - if _response_stream_shape_cache is None: - from botocore.loaders import Loader - from botocore.model import ServiceModel - - loader = Loader() - bedrock_service_dict = loader.load_service_model("bedrock-runtime", "service-2") - bedrock_service_model = ServiceModel(bedrock_service_dict) - _response_stream_shape_cache = bedrock_service_model.shape_for("ResponseStream") - - return _response_stream_shape_cache - - class AWSEventStreamDecoder: def __init__(self, model: str, json_mode: Optional[bool] = False) -> None: from botocore.parsers import EventStreamJSONParser @@ -1838,8 +1828,18 @@ class AWSEventStreamDecoder: yield self._chunk_parser(chunk_data=_data) def _parse_message_from_event(self, event) -> Optional[str]: + if BEDROCK_RESPONSE_STREAM_SHAPE is None: + raise BedrockError( + status_code=500, + message=( + "Bedrock event-stream shape could not be loaded from botocore. " + "Ensure botocore is correctly installed." + ), + ) response_dict = event.to_response_dict() - parsed_response = self.parser.parse(response_dict, get_response_stream_shape()) + parsed_response = self.parser.parse( + response_dict, BEDROCK_RESPONSE_STREAM_SHAPE + ) if response_dict["status_code"] != 200: decoded_body = response_dict["body"].decode() diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 9a97a134cc4..856a525f773 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -14,6 +14,7 @@ if TYPE_CHECKING: import httpx import litellm +from litellm import verbose_logger from litellm.llms.base_llm.anthropic_messages.transformation import ( BaseAnthropicMessagesConfig, ) @@ -917,38 +918,57 @@ def get_bedrock_chat_config(model: str): return litellm.AmazonInvokeConfig() +def _load_bedrock_response_stream_shape(): + """ + Load the ResponseStream shape from botocore's bundled bedrock-runtime schema. + + Called once at module import time; the result is stored in + ``BEDROCK_RESPONSE_STREAM_SHAPE`` and reused for the process lifetime. + Returns ``None`` if botocore is unavailable or the service model cannot be + loaded, so the module still imports cleanly. + """ + try: + from botocore.loaders import Loader + from botocore.model import ServiceModel + + loader = Loader() + service_dict = loader.load_service_model("bedrock-runtime", "service-2") + return ServiceModel(service_dict).shape_for("ResponseStream") + except Exception as e: + verbose_logger.warning( + "litellm: could not pre-load bedrock-runtime response stream shape " + "— Bedrock event-stream decoding will be unavailable. Error: %s", + e, + ) + return None + + +# Eagerly resolved once per process — avoids per-instance or per-request disk I/O. +BEDROCK_RESPONSE_STREAM_SHAPE = _load_bedrock_response_stream_shape() + + class BedrockEventStreamDecoderBase: """ Base class for event stream decoding for Bedrock """ - _response_stream_shape_cache = None - def __init__(self): from botocore.parsers import EventStreamJSONParser self.parser = EventStreamJSONParser() - def get_response_stream_shape(self): - if self._response_stream_shape_cache is None: - from botocore.loaders import Loader - from botocore.model import ServiceModel - - loader = Loader() - bedrock_service_dict = loader.load_service_model( - "bedrock-runtime", "service-2" - ) - bedrock_service_model = ServiceModel(bedrock_service_dict) - self._response_stream_shape_cache = bedrock_service_model.shape_for( - "ResponseStream" - ) - - return self._response_stream_shape_cache - def _parse_message_from_event(self, event) -> Optional[str]: + if BEDROCK_RESPONSE_STREAM_SHAPE is None: + raise BedrockError( + status_code=500, + message=( + "Bedrock event-stream shape could not be loaded from botocore. " + "Ensure botocore is correctly installed." + ), + ) response_dict = event.to_response_dict() parsed_response = self.parser.parse( - response_dict, self.get_response_stream_shape() + response_dict, BEDROCK_RESPONSE_STREAM_SHAPE ) if response_dict["status_code"] != 200: diff --git a/litellm/llms/sagemaker/common_utils.py b/litellm/llms/sagemaker/common_utils.py index ad6b24d85a3..50c8ee4220e 100644 --- a/litellm/llms/sagemaker/common_utils.py +++ b/litellm/llms/sagemaker/common_utils.py @@ -9,7 +9,27 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.utils import GenericStreamingChunk as GChunk from litellm.types.utils import StreamingChatCompletionChunk -_response_stream_shape_cache = None + +def _load_sagemaker_response_stream_shape(): + try: + from botocore.loaders import Loader + from botocore.model import ServiceModel + + loader = Loader() + service_dict = loader.load_service_model("sagemaker-runtime", "service-2") + return ServiceModel(service_dict).shape_for( + "InvokeEndpointWithResponseStreamOutput" + ) + except Exception as e: + verbose_logger.warning( + "litellm: could not pre-load sagemaker-runtime response stream shape " + "— SageMaker event-stream decoding will be unavailable. Error: %s", + e, + ) + return None + + +SAGEMAKER_RESPONSE_STREAM_SHAPE = _load_sagemaker_response_stream_shape() class SagemakerError(BaseLLMException): @@ -187,8 +207,18 @@ class AWSEventStreamDecoder: verbose_logger.error(f"Final error parsing accumulated JSON: {e}") def _parse_message_from_event(self, event) -> Optional[str]: + if SAGEMAKER_RESPONSE_STREAM_SHAPE is None: + raise SagemakerError( + status_code=500, + message=( + "SageMaker event-stream shape could not be loaded from botocore. " + "Ensure botocore is correctly installed." + ), + ) response_dict = event.to_response_dict() - parsed_response = self.parser.parse(response_dict, get_response_stream_shape()) + parsed_response = self.parser.parse( + response_dict, SAGEMAKER_RESPONSE_STREAM_SHAPE + ) if response_dict["status_code"] != 200: raise ValueError(f"Bad response code, expected 200: {response_dict}") @@ -204,20 +234,3 @@ class AWSEventStreamDecoder: return None return chunk.decode() # type: ignore[no-any-return] - - -def get_response_stream_shape(): - global _response_stream_shape_cache - if _response_stream_shape_cache is None: - from botocore.loaders import Loader - from botocore.model import ServiceModel - - loader = Loader() - sagemaker_service_dict = loader.load_service_model( - "sagemaker-runtime", "service-2" - ) - sagemaker_service_model = ServiceModel(sagemaker_service_dict) - _response_stream_shape_cache = sagemaker_service_model.shape_for( - "InvokeEndpointWithResponseStreamOutput" - ) - return _response_stream_shape_cache diff --git a/scripts/mock_bedrock_passthrough_target.py b/scripts/mock_bedrock_passthrough_target.py new file mode 100644 index 00000000000..e993cd99bde --- /dev/null +++ b/scripts/mock_bedrock_passthrough_target.py @@ -0,0 +1,276 @@ +#!/usr/bin/env python3 +""" +Minimal HTTP target for testing LiteLLM **Bedrock pass-through** (`/bedrock/...` on the proxy). + +What it does + - Serves a tiny Converse-shaped JSON (and optional invoke-shaped) response so the proxy can + complete a round trip without calling AWS. + - Does **not** verify SigV4 (Bedrock does); any Authorization header is accepted. + +How to run + uv run python scripts/mock_bedrock_passthrough_target.py --host 127.0.0.1 --port 9999 + +Wire LiteLLM to this host (use **one** of these patterns): + + 1) model_list (recommended) — set the Bedrock runtime base to the mock: + + model_list: + - model_name: mock-bedrock-claude + litellm_params: + model: bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0 + custom_llm_provider: bedrock + aws_region_name: us-west-2 + api_base: "http://127.0.0.1:9999" + + 2) Environment (see litellm BaseAWSLLM.get_runtime_endpoint):: + + export AWS_BEDROCK_RUNTIME_ENDPOINT="http://127.0.0.1:9999" + +Then call the proxy, e.g. (model_name must match config):: + + curl -sS -X POST "http://127.0.0.1:4000/bedrock/model/mock-bedrock-claude/converse" \ + -H "Authorization: Bearer $LITELLM_KEY" -H "Content-Type: application/json" \ + -d '{"messages":[{"role":"user","content":[{"text":"hi"}]}]}' + +The proxy will forward to: {api_base}/model//converse (SigV4-signed). +This mock implements POST .../converse and returns a minimal valid Converse response. + +Notes + - `invoke-with-response-stream` returns a real **binary** AWS event stream + (`application/vnd.amazon.eventstream`) with Anthropic-style JSON payloads inside each + `PayloadPart`, matching Bedrock's InvokeModelWithResponseStream wire format. See + https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InvokeModelWithResponseStream.html + and https://docs.aws.amazon.com/awstreams/latest/devguide/message-formats.html + - `converse-stream` is still JSON-only placeholder (different inner event shapes). + - Use real (or any non-empty) AWS creds in the environment of the **proxy**; signing still runs. +""" +from __future__ import annotations + +import argparse +import base64 +import json +from binascii import crc32 +from struct import pack +from typing import Any, Dict, Iterator, List + +from fastapi import FastAPI, Request +from fastapi.responses import JSONResponse +from starlette.responses import StreamingResponse + +app = FastAPI(title="Mock Bedrock runtime (pass-through test target)") + + +# Minimal structure compatible with Converse: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_Converse.html +def _converse_response_body() -> Dict[str, Any]: + return { + "output": { + "message": { + "role": "assistant", + "content": [ + {"text": "mock: ok from mock_bedrock_passthrough_target.py"} + ], + } + }, + "stopReason": "end_turn", + "usage": { + "inputTokens": 1, + "outputTokens": 2, + "totalTokens": 3, + }, + } + + +# Minimal invoke (Anthropic messages on bedrock) style — adjust if you test /invoke +def _invoke_response_body() -> Dict[str, Any]: + return { + "id": "msg_mock", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "mock invoke response"}], + "model": "mock", + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 2}, + } + + +def _encode_event_stream_message(headers: Dict[str, str], payload: bytes) -> bytes: + """Single AWS binary event-stream frame (same layout botocore's ``EventStreamBuffer`` parses).""" + header_blob = b"" + for name, value in headers.items(): + nb = name.encode("utf-8") + vb = value.encode("utf-8") + header_blob += bytes([len(nb)]) + nb + bytes([7]) + pack("!H", len(vb)) + vb + headers_length = len(header_blob) + payload_length = len(payload) + total_length = 12 + headers_length + payload_length + 4 + prelude_wo_crc = pack("!II", total_length, headers_length) + prelude_crc_val = crc32(prelude_wo_crc) & 0xFFFFFFFF + prelude = prelude_wo_crc + pack("!I", prelude_crc_val) + wo_msg_crc = prelude + header_blob + payload + msg_crc_val = crc32(wo_msg_crc[8:], prelude_crc_val) & 0xFFFFFFFF + return wo_msg_crc + pack("!I", msg_crc_val) + + +def _bedrock_payload_part(inner_event: Dict[str, Any]) -> bytes: + """Outer JSON expected by bedrock-runtime ``ResponseStream`` / ``PayloadPart``.""" + inner_bytes = json.dumps(inner_event, separators=(",", ":")).encode("utf-8") + outer = { + "chunk": { + "bytes": base64.b64encode(inner_bytes).decode("ascii"), + } + } + return json.dumps(outer, separators=(",", ":")).encode("utf-8") + + +def _anthropic_invoke_stream_events( + model_id: str, assistant_text: str +) -> List[Dict[str, Any]]: + """ + Minimal Anthropic Messages stream events as returned inside Bedrock stream chunks. + Mirrors the sequence Amazon emits for Claude on ``invoke-with-response-stream``. + """ + msg_id = "msg_mock_bedrock_stream" + input_tokens = 3 + output_tokens = max(1, len(assistant_text) // 4) + events: List[Dict[str, Any]] = [ + { + "type": "message_start", + "message": { + "model": model_id, + "id": msg_id, + "type": "message", + "role": "assistant", + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": { + "input_tokens": input_tokens, + "output_tokens": 1, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": { + "ephemeral_5m_input_tokens": 0, + "ephemeral_1h_input_tokens": 0, + }, + }, + }, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + ] + # Split text into small deltas so downstream streaming behavior is visible. + step = 24 + for i in range(0, len(assistant_text), step): + events.append( + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "text_delta", + "text": assistant_text[i : i + step], + }, + } + ) + events.append({"type": "content_block_stop", "index": 0}) + events.append( + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + }, + } + ) + events.append( + { + "type": "message_stop", + "amazon-bedrock-invocationMetrics": { + "inputTokenCount": input_tokens, + "outputTokenCount": output_tokens, + "invocationLatency": 42, + "firstByteLatency": 10, + }, + } + ) + return events + + +def _iter_invoke_with_response_stream(model_id: str) -> Iterator[bytes]: + text = ( + "mock streaming: ok from scripts/mock_bedrock_passthrough_target.py " + "(invoke-with-response-stream)." + ) + headers = { + ":event-type": "chunk", + ":content-type": "application/json", + ":message-type": "event", + } + for ev in _anthropic_invoke_stream_events(model_id, text): + yield _encode_event_stream_message(headers, _bedrock_payload_part(ev)) + + +@app.get("/health") +def health() -> Dict[str, str]: + return {"status": "ok"} + + +@app.post("/model/{model_path:path}/converse") +async def converse(model_path: str, request: Request) -> JSONResponse: + # Optional: log body for debugging + _ = await request.body() + return JSONResponse(content=_converse_response_body()) + + +@app.post("/model/{model_path:path}/converse-stream") +async def converse_stream(model_path: str, request: Request) -> JSONResponse: + """ + Not a real AWS event stream — returns JSON for quick smoke tests only. + """ + _ = await request.body() + return JSONResponse( + content={ + "note": "This mock does not implement application/vnd.amazon.eventstream; use /converse for basic tests." + } + ) + + +@app.post("/model/{model_path:path}/invoke") +async def invoke(model_path: str, request: Request) -> JSONResponse: + _ = await request.body() + return JSONResponse(content=_invoke_response_body()) + + +@app.post("/model/{model_path:path}/invoke-with-response-stream") +async def invoke_with_response_stream( + model_path: str, request: Request +) -> StreamingResponse: + """ + Binary ``application/vnd.amazon.eventstream`` body compatible with boto3/botocore + ``InvokeModelWithResponseStream`` / LiteLLM's Bedrock invoke streaming path. + """ + _ = await request.body() + return StreamingResponse( + _iter_invoke_with_response_stream(model_id=model_path), + media_type="application/vnd.amazon.eventstream", + ) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="127.0.0.1") + parser.add_argument("--port", type=int, default=9999) + args = parser.parse_args() + + import uvicorn + + uvicorn.run(app, host=args.host, port=args.port, log_level="info") + + +if __name__ == "__main__": + main() diff --git a/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py b/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py index c356f866b07..8fa9290d3de 100644 --- a/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/test_litellm/llms/bedrock/test_bedrock_common_utils.py @@ -13,6 +13,110 @@ sys.path.insert( from litellm.llms.bedrock.common_utils import BedrockModelInfo +# --------------------------------------------------------------------------- # +# BEDROCK_RESPONSE_STREAM_SHAPE eager-load tests # +# --------------------------------------------------------------------------- # + + +def test_bedrock_response_stream_shape_loaded_at_import(): + """ + BEDROCK_RESPONSE_STREAM_SHAPE is resolved at module import time. + In a standard environment with botocore installed it must be non-None. + """ + from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE + + assert BEDROCK_RESPONSE_STREAM_SHAPE is not None + + +def test_bedrock_response_stream_shape_load_failure_returns_none(): + """ + If botocore's Loader raises (e.g. missing data files), _load_bedrock_response_stream_shape + should return None rather than propagating the exception, so the module + still imports cleanly. + """ + from unittest.mock import patch + + import litellm.llms.bedrock.common_utils as mod + + with patch( + "botocore.loaders.Loader.load_service_model", + side_effect=Exception("no data"), + ): + shape = mod._load_bedrock_response_stream_shape() + assert shape is None + + +def test_bedrock_response_stream_shape_is_structure_shape(): + """ + The loaded shape should be the botocore StructureShape for ResponseStream, + not a plain dict or any other type. + """ + from botocore.model import StructureShape + + from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE + + assert BEDROCK_RESPONSE_STREAM_SHAPE is not None, ( + "BEDROCK_RESPONSE_STREAM_SHAPE is None — botocore may not be installed" + ) + shape: StructureShape = BEDROCK_RESPONSE_STREAM_SHAPE # remove Optional + assert isinstance(shape, StructureShape) + assert shape.name == "ResponseStream" + + +def test_bedrock_response_stream_shape_same_object_across_imports(): + """ + Both bedrock modules that use the shape must reference the identical object — + confirming the constant is not re-loaded per import. + """ + from litellm.llms.bedrock.chat.invoke_handler import ( + BEDROCK_RESPONSE_STREAM_SHAPE as invoke_shape, + ) + from litellm.llms.bedrock.common_utils import ( + BEDROCK_RESPONSE_STREAM_SHAPE as common_shape, + ) + + assert common_shape is invoke_shape + + +def test_bedrock_event_stream_decoder_base_uses_module_shape(): + """ + BedrockEventStreamDecoderBase instances no longer carry their own + per-instance cache — _parse_message_from_event uses the module constant + directly, so there is no instance-level _response_stream_shape_cache attr. + """ + from litellm.llms.bedrock.common_utils import BedrockEventStreamDecoderBase + + decoder_a = BedrockEventStreamDecoderBase() + decoder_b = BedrockEventStreamDecoderBase() + + assert "_response_stream_shape_cache" not in decoder_a.__dict__ + assert "_response_stream_shape_cache" not in decoder_b.__dict__ + + +def test_bedrock_parse_message_from_event_raises_on_none_shape(): + """ + When BEDROCK_RESPONSE_STREAM_SHAPE is None (botocore unavailable), + _parse_message_from_event must raise BedrockError before touching the + botocore parser — not an opaque AttributeError from inside botocore. + """ + from unittest.mock import MagicMock, patch + + import litellm.llms.bedrock.common_utils as mod + from litellm.llms.bedrock.common_utils import BedrockError, BedrockEventStreamDecoderBase + + decoder = BedrockEventStreamDecoderBase() + mock_event = MagicMock() + + with patch.object(mod, "BEDROCK_RESPONSE_STREAM_SHAPE", None): + with pytest.raises(BedrockError) as exc_info: + decoder._parse_message_from_event(mock_event) + + assert exc_info.value.status_code == 500 + assert "botocore" in str(exc_info.value.message).lower() + # The botocore parser must never have been called + mock_event.to_response_dict.assert_not_called() + + def test_deepseek_cris(): """ Test that DeepSeek models with cross-region inference prefix use converse route diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py index 70a8d86cb1b..9d7706557b5 100644 --- a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py @@ -11,6 +11,102 @@ from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder from litellm.llms.sagemaker.completion.transformation import SagemakerConfig +# --------------------------------------------------------------------------- # +# SAGEMAKER_RESPONSE_STREAM_SHAPE eager-load tests # +# --------------------------------------------------------------------------- # + + +def test_sagemaker_response_stream_shape_loaded_at_import(): + """ + SAGEMAKER_RESPONSE_STREAM_SHAPE is resolved at module import time. + In a standard environment with botocore installed it must be non-None. + """ + from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + + assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None + + +def test_sagemaker_response_stream_shape_load_failure_returns_none(): + """ + If botocore's Loader raises (e.g. missing data files), _load_sagemaker_response_stream_shape + should return None rather than propagating the exception, so the module + still imports cleanly. + """ + from unittest.mock import patch + + import litellm.llms.sagemaker.common_utils as mod + + with patch( + "botocore.loaders.Loader.load_service_model", + side_effect=Exception("no data"), + ): + shape = mod._load_sagemaker_response_stream_shape() + assert shape is None + + +def test_sagemaker_response_stream_shape_is_structure_shape(): + """ + The loaded shape should be the botocore StructureShape for + InvokeEndpointWithResponseStreamOutput, not a plain dict or any other type. + """ + from botocore.model import StructureShape + + from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + + assert SAGEMAKER_RESPONSE_STREAM_SHAPE is not None, ( + "SAGEMAKER_RESPONSE_STREAM_SHAPE is None — botocore may not be installed" + ) + shape: StructureShape = SAGEMAKER_RESPONSE_STREAM_SHAPE # remove Optional + assert isinstance(shape, StructureShape) + assert shape.name == "InvokeEndpointWithResponseStreamOutput" + + +def test_sagemaker_response_stream_shape_not_reloaded_on_new_decoder(): + """ + Creating multiple AWSEventStreamDecoder instances must not trigger + additional botocore Loader calls — the shape is resolved once at import + time and reused. + """ + from litellm.llms.sagemaker.common_utils import SAGEMAKER_RESPONSE_STREAM_SHAPE + + decoder_a = AWSEventStreamDecoder(model="test-model-a") + decoder_b = AWSEventStreamDecoder(model="test-model-b") + + # Both decoders should use the same pre-loaded shape object (identity check) + assert "_response_stream_shape_cache" not in decoder_a.__dict__ + assert "_response_stream_shape_cache" not in decoder_b.__dict__ + # The module constant is still the same object + from litellm.llms.sagemaker.common_utils import ( + SAGEMAKER_RESPONSE_STREAM_SHAPE as shape_after, + ) + + assert SAGEMAKER_RESPONSE_STREAM_SHAPE is shape_after + + +def test_sagemaker_parse_message_from_event_raises_on_none_shape(): + """ + When SAGEMAKER_RESPONSE_STREAM_SHAPE is None (botocore unavailable), + _parse_message_from_event must raise ValueError before touching the + botocore parser — not an opaque AttributeError from inside botocore. + """ + from unittest.mock import MagicMock, patch + + import litellm.llms.sagemaker.common_utils as mod + from litellm.llms.sagemaker.common_utils import SagemakerError + + decoder = AWSEventStreamDecoder(model="test-model") + mock_event = MagicMock() + + with patch.object(mod, "SAGEMAKER_RESPONSE_STREAM_SHAPE", None): + with pytest.raises(SagemakerError) as exc_info: + decoder._parse_message_from_event(mock_event) + + assert exc_info.value.status_code == 500 + assert "botocore" in str(exc_info.value.message).lower() + # The botocore parser must never have been called + mock_event.to_response_dict.assert_not_called() + + @pytest.mark.asyncio async def test_aiter_bytes_unicode_decode_error(): """