From 5862be3e79154ce5ac8a736c77c7c3614fac40fb Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Tue, 7 Jul 2026 18:23:49 -0700 Subject: [PATCH 01/31] fix(proxy): resolve os.environ/ refs universally in DB-sourced models MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Root cause: PR #30867 removed request-time os.environ/ expansion in BaseAWSLLM.get_credentials. That is only safe if config-load pre-resolves os.environ/ refs so the value reaching get_credentials is already the real secret. The YAML config path has always done this. The DB-load path (ProxyConfig._resolve_db_litellm_param) only re-expanded keys in a hardcoded whitelist (_DB_LITELLM_PARAM_ENV_REF_KEYS) plus short-circuited env-ref resolution entirely for team-scoped rows. PR #32256 extended that whitelist to 18 keys to unblock a customer whose Bedrock model with aws_role_name: os.environ/BEDROCK_ASSUME_ROLE_ARN broke on v1.90+, but the whitelist is structurally fragile: every future auth field breaks the same way until someone remembers to add it Fix: remove the whitelist and the team-scope short-circuit. The DB-load resolver now expands os.environ/ on every string field, matching the YAML path. Trust boundary stays on the write side: only PROXY_ADMIN can create team_id=None rows, only team admins of a team can create rows scoped to that team, and the request-body vector is still blocked by _BANNED_REQUEST_BODY_PARAMS. Team-scoped rows now resolve env refs — this is a deliberate LIT-3831 threat-model expansion trusting team admins for env-var reads Regression tests in tests/test_litellm/proxy/proxy_server/test_proxy_config.py: - test_ProxyConfig__add_deployment_resolves_env_refs_after_db_decrypt pins admin-scoped rows resolve every field (previously api_base stayed literal) - test_ProxyConfig__add_deployment_resolves_team_env_refs pins team rows resolve env refs (previously stayed literal) - test_ProxyConfig__add_deployment_resolves_env_refs_on_arbitrary_field pins the no-whitelist invariant against a made-up field name - test_ProxyConfig__add_deployment_resolves_env_refs_for_aws_bedrock_auth_params (from #32256) still passes - Path B counterparts (decrypt_model_list_from_db) mirror the above Left as followups (not fixed here): - /model/info and /v2/model/info still echo resolved values for fields not in the current pop-list (aws_role_name, aws_sts_endpoint, api_base, etc.). Fix is to extend remove_sensitive_info_from_deployment; separate PR - Master-key rotation reads DB rows via decrypt_model_list_from_db which now resolves universally, so rotation collapses env-refs into hardcoded values. Pre-existing bug for the 6 previously-whitelisted fields; wider surface after this PR. Separate PR --- litellm/proxy/proxy_server.py | 58 ++------------- .../proxy/proxy_server/test_proxy_config.py | 71 ++++++++++--------- 2 files changed, 42 insertions(+), 87 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 38788d140e9..13e2d4f1252 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1117,45 +1117,6 @@ _OPENAPI_HTTP_METHODS = { # `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO # and cache endpoint files. _ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"} -_DB_LITELLM_PARAM_ENV_REF_KEYS = frozenset( - { - "api_key", - "client_secret", - "vertex_credentials", - "vertex_ai_credentials", - "aws_access_key_id", - "aws_secret_access_key", - "aws_session_token", - "aws_region_name", - "aws_session_name", - "aws_profile_name", - "aws_role_name", - "aws_web_identity_token", - "aws_sts_endpoint", - "aws_external_id", - "aws_bedrock_runtime_endpoint", - "aws_bedrock_project_id", - "aws_batch_role_arn", - "aws_workspace_id", - } -) - - -def _db_model_is_team_scoped(model: object) -> bool: - model_info = getattr(model, "model_info", None) - if isinstance(model_info, BaseModel): - return getattr(model_info, "team_id", None) is not None - if isinstance(model_info, str): - try: - model_info = json.loads(model_info) - except (TypeError, ValueError): - model_info = None - if isinstance(model_info, dict) and model_info.get("team_id") is not None: - return True - if getattr(model_info, "team_id", None) is not None: - return True - model_name = getattr(model, "model_name", None) - return isinstance(model_name, str) and model_name.startswith("model_name_") def _strip_operation_id_method_suffix(operation_id: str) -> str: @@ -5009,17 +4970,12 @@ class ProxyConfig: deleted_deployments += 1 return deleted_deployments - def _resolve_db_litellm_param(self, key: str, value: object, resolve_env_refs: bool = True) -> object: + def _resolve_db_litellm_param(self, key: str, value: object) -> object: if not isinstance(value, str): return value decrypted_value = decrypt_value_helper(value=value, key=key, return_original_value=True) - if ( - resolve_env_refs - and key in _DB_LITELLM_PARAM_ENV_REF_KEYS - and isinstance(decrypted_value, str) - and decrypted_value.startswith("os.environ/") - ): + if isinstance(decrypted_value, str) and decrypted_value.startswith("os.environ/"): return get_secret(decrypted_value) return decrypted_value @@ -5040,13 +4996,10 @@ class ProxyConfig: ## ADD MODEL LOGIC for m in db_models: _litellm_params = m.litellm_params - resolve_env_refs = not _db_model_is_team_scoped(m) if isinstance(_litellm_params, dict): # decrypt values for k, v in _litellm_params.items(): - _litellm_params[k] = self._resolve_db_litellm_param( - key=k, value=v, resolve_env_refs=resolve_env_refs - ) + _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) _litellm_params = LiteLLM_Params(**_litellm_params) else: @@ -5072,15 +5025,12 @@ class ProxyConfig: _model_list: list = [] for m in new_models: _litellm_params = m.litellm_params - resolve_env_refs = not _db_model_is_team_scoped(m) if isinstance(_litellm_params, BaseModel): _litellm_params = _litellm_params.model_dump() if isinstance(_litellm_params, dict): # decrypt values for k, v in _litellm_params.items(): - _litellm_params[k] = self._resolve_db_litellm_param( - key=k, value=v, resolve_env_refs=resolve_env_refs - ) + _litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v) _litellm_params = LiteLLM_Params(**_litellm_params) else: verbose_proxy_logger.error( diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 6cdaa17c0bf..98b270788dc 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -1102,6 +1102,11 @@ def test_ProxyConfig__add_deployment_invalid_litellm_params_skips(monkeypatch): def test_ProxyConfig__add_deployment_resolves_env_refs_after_db_decrypt(monkeypatch): + """Every ``os.environ/`` value on an admin-scoped DB row resolves at + load time, regardless of the field name. Replaces the earlier + behavior where only fields in ``_DB_LITELLM_PARAM_ENV_REF_KEYS`` + resolved: the whitelist has been removed so the resolver applies to + every string field.""" monkeypatch.setenv("LITELLM_DB_MODEL_API_KEY", "resolved-secret") monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret") monkeypatch.setattr( @@ -1129,19 +1134,21 @@ def test_ProxyConfig__add_deployment_resolves_env_refs_after_db_decrypt(monkeypa assert added == 1 assert deployment.litellm_params.api_key == "resolved-secret" - assert deployment.litellm_params.api_base == "os.environ/LITELLM_MASTER_KEY" + assert deployment.litellm_params.api_base == "master-secret" -def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch): - def fail_on_call(secret_name, *args, **kwargs): - raise AssertionError("team DB models should not resolve env refs") - +def test_ProxyConfig__add_deployment_resolves_team_env_refs(monkeypatch): + """Team-scoped DB rows now resolve ``os.environ/`` refs the same way + admin rows do. The prior team-scoped short-circuit and the + field-by-field whitelist have both been removed; the write-side team + auth check in ``ModelManagementAuthChecks.can_user_make_model_call`` + remains the single trust boundary. A literal (non-``os.environ/``) + value still passes through unchanged.""" monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret") monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", lambda value, key, return_original_value: value, ) - monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call) fake_router = MagicMock() fake_router.upsert_deployment = MagicMock(return_value=True) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) @@ -1153,7 +1160,7 @@ def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch): litellm_params={ "model": "openai/gpt-4o-mini", "api_key": "os.environ/LITELLM_MASTER_KEY", - "api_base": "https://attacker.example", + "api_base": "https://team.example", }, blocked=False, ) @@ -1162,8 +1169,8 @@ def test_ProxyConfig__add_deployment_keeps_team_env_refs_literal(monkeypatch): deployment = fake_router.upsert_deployment.call_args.kwargs["deployment"] assert added == 1 - assert deployment.litellm_params.api_key == "os.environ/LITELLM_MASTER_KEY" - assert deployment.litellm_params.api_base == "https://attacker.example" + assert deployment.litellm_params.api_key == "master-secret" + assert deployment.litellm_params.api_base == "https://team.example" def test_ProxyConfig__resolve_db_litellm_param_skips_non_string_values(monkeypatch): @@ -1242,31 +1249,26 @@ def test_ProxyConfig__add_deployment_resolves_env_refs_for_aws_bedrock_auth_para assert getattr(deployment.litellm_params, key) == expected, key -def test_ProxyConfig__add_deployment_keeps_team_aws_env_refs_literal(monkeypatch): - """Team-scoped DB models must NOT resolve env refs even for AWS auth - params: this is the LIT-3831 defense-in-depth path where a team admin - could otherwise craft a DB entry that reads the process environment.""" - - def fail_on_call(secret_name, *args, **kwargs): - raise AssertionError("team DB models should not resolve env refs") - - monkeypatch.setenv("BEDROCK_ASSUME_ROLE_ARN", "arn:aws:iam::123:role/should-not-leak") +def test_ProxyConfig__add_deployment_resolves_env_refs_on_arbitrary_field(monkeypatch): + """A made-up field name that was never on the removed whitelist still + resolves ``os.environ/`` refs. Pins the "no whitelist" invariant: + the resolver applies to every string field, not a curated list.""" + monkeypatch.setenv("SOME_CUSTOM_ENV", "resolved-custom-value") monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", lambda value, key, return_original_value: value, ) - monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call) fake_router = MagicMock() fake_router.upsert_deployment = MagicMock(return_value=True) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) pc = ProxyConfig() db_model = SimpleNamespace( model_id="model-1", - model_name="model_name_team-1_bedrock", - model_info={"id": "model-1", "team_id": "team-1"}, + model_name="custom-field-model", + model_info={"id": "model-1"}, litellm_params={ - "model": "bedrock/anthropic.claude-v2", - "aws_role_name": "os.environ/BEDROCK_ASSUME_ROLE_ARN", + "model": "openai/gpt-4o-mini", + "some_future_field": "os.environ/SOME_CUSTOM_ENV", }, blocked=False, ) @@ -1275,7 +1277,7 @@ def test_ProxyConfig__add_deployment_keeps_team_aws_env_refs_literal(monkeypatch deployment = fake_router.upsert_deployment.call_args.kwargs["deployment"] assert added == 1 - assert deployment.litellm_params.aws_role_name == "os.environ/BEDROCK_ASSUME_ROLE_ARN" + assert deployment.litellm_params.some_future_field == "resolved-custom-value" # --------------------------------------------------------------------------- @@ -1313,6 +1315,9 @@ def test_ProxyConfig_decrypt_model_list_from_db_returns_decrypted(monkeypatch): def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decrypt( monkeypatch, ): + """Path B (feeding /v2/model/info fallback and /model/info fallback) + resolves every ``os.environ/`` field on admin-scoped rows, mirroring + path A. Both paths now share the same universal-resolution shape.""" monkeypatch.setenv("LITELLM_DB_MODEL_API_KEY", "resolved-secret") monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret") monkeypatch.setattr( @@ -1341,21 +1346,21 @@ def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decry out = pc.decrypt_model_list_from_db(new_models=[m]) assert out[0]["litellm_params"]["api_key"] == "resolved-secret" - assert out[0]["litellm_params"]["api_base"] == "os.environ/LITELLM_MASTER_KEY" + assert out[0]["litellm_params"]["api_base"] == "master-secret" -def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_after_db_decrypt( +def test_ProxyConfig_decrypt_model_list_from_db_resolves_team_env_refs_after_db_decrypt( monkeypatch, ): - def fail_on_call(secret_name, *args, **kwargs): - raise AssertionError("team DB models should not resolve env refs") - + """Team-scoped rows on path B resolve ``os.environ/`` refs just like + admin rows do. Pairs with + ``test_ProxyConfig__add_deployment_resolves_team_env_refs`` on path + A — both paths now agree on the trust model.""" monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret") monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", lambda value, key, return_original_value: "os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value, ) - monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call) pc = ProxyConfig() m = SimpleNamespace( model_id="model-1", @@ -1363,7 +1368,7 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte model_info={"id": "model-1", "team_id": "team-1"}, litellm_params={ "api_key": "encrypted-env-ref", - "api_base": "https://attacker.example", + "api_base": "https://team.example", "model": "openai/gpt-4o-mini", }, blocked=False, @@ -1371,8 +1376,8 @@ def test_ProxyConfig_decrypt_model_list_from_db_keeps_team_env_refs_literal_afte out = pc.decrypt_model_list_from_db(new_models=[m]) - assert out[0]["litellm_params"]["api_key"] == "os.environ/LITELLM_MASTER_KEY" - assert out[0]["litellm_params"]["api_base"] == "https://attacker.example" + assert out[0]["litellm_params"]["api_key"] == "master-secret" + assert out[0]["litellm_params"]["api_base"] == "https://team.example" def test_ProxyConfig_decrypt_model_list_from_db_invalid_params_skips(): From 734fd29e00da887493856033152f01e268ad59dd Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 20:49:03 -0700 Subject: [PATCH 02/31] fix(utils): resolve bedrock regional inference profiles to regional pricing in get_model_info (LIT-4056) (#32389) * fix(utils): resolve bedrock regional inference profiles to regional pricing in get_model_info (LIT-4056) * test(register_model): use a triple provider prefix as the unresolvable-key fixture get_model_info now resolves bedrock/bedrock/... like a routing prefix, so the double-prefix fixture stopped exercising the register_model fallback path. Lock the new double-prefix resolution in as a model-info regression test --- litellm/utils.py | 38 +++++++++++-------- .../test_register_model_custom_pricing.py | 7 ++-- tests/test_litellm/test_utils.py | 37 ++++++++++++++++++ 3 files changed, 63 insertions(+), 19 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index f9bb84101e7..19c2fe16085 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2619,8 +2619,9 @@ _CACHE_PRICING_FIELDS = ( def _resolve_builtin_model_cost_entry(key: str, provider: str) -> Optional[Dict[str, Any]]: """Best-effort lookup of a built-in ``model_cost`` entry for a custom key - whose shape ``get_model_info`` cannot resolve (double provider prefixes - like ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region aliases). + whose shape ``get_model_info`` cannot resolve (repeated provider prefixes + like ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6`` or region + aliases). Returns a copy of the matching entry so the caller can inherit its defaults (most importantly cache pricing) without mutating the shared built-in. @@ -5052,9 +5053,9 @@ def _get_model_info_from_generalization( candidates = [ potential_model_names["combined_model_name"], model, + potential_model_names["split_model"], potential_model_names["combined_stripped_model_name"], potential_model_names["stripped_model_name"], - potential_model_names["split_model"], ] for candidate in candidates: generalized_info = match_fallback_generalization(candidate) @@ -5094,6 +5095,11 @@ def _get_potential_model_names( stripped_model_name, ) + if custom_llm_provider in ("bedrock", "bedrock_converse"): + from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix + + split_model = strip_bedrock_routing_prefix(split_model) + return PotentialModelNamesAndCustomLLMProvider( split_model=split_model, combined_model_name=combined_model_name, @@ -5261,9 +5267,9 @@ def _get_model_info_helper( Check if: (in order of specificity) 1. 'custom_llm_provider/model' in litellm.model_cost. Checks "groq/llama3-8b-8192" if model="llama3-8b-8192" and custom_llm_provider="groq" 2. 'model' in litellm.model_cost. Checks "gemini-1.5-pro-002" in litellm.model_cost if model="gemini-1.5-pro-002" and custom_llm_provider=None - 3. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given. - 4. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given. - 5. 'split_model' in litellm.model_cost. Checks "llama3-8b-8192" in litellm.model_cost if model="groq/llama3-8b-8192" + 3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8" + 4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given. + 5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given. """ _model_info: Optional[Dict[str, Any]] = None @@ -5289,6 +5295,16 @@ def _get_model_info_helper( custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None + if _model_info is None: + _matched_key = _get_model_cost_key(split_model) + if _matched_key is not None: + key = _matched_key + _model_info = _get_model_info_from_model_cost(key=cast(str, key)) + if not _check_provider_match( + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, + ): + _model_info = None if _model_info is None: _matched_key = _get_model_cost_key(combined_stripped_model_name) if _matched_key is not None: @@ -5309,16 +5325,6 @@ def _get_model_info_helper( custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None - if _model_info is None: - _matched_key = _get_model_cost_key(split_model) - if _matched_key is not None: - key = _matched_key - _model_info = _get_model_info_from_model_cost(key=cast(str, key)) - if not _check_provider_match( - model_info=_model_info, - custom_llm_provider=model_cost_custom_llm_provider, - ): - _model_info = None if _model_info is None: generalization = _get_model_info_from_generalization( diff --git a/tests/test_litellm/test_register_model_custom_pricing.py b/tests/test_litellm/test_register_model_custom_pricing.py index 8c3f690982b..ba82bfaadc6 100644 --- a/tests/test_litellm/test_register_model_custom_pricing.py +++ b/tests/test_litellm/test_register_model_custom_pricing.py @@ -320,8 +320,9 @@ def test_register_model_strips_none_litellm_provider_from_get_model_info(monkeyp def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key(): """Registering a custom override under a key shape that - ``get_model_info`` cannot resolve (e.g. a double provider prefix like - ``bedrock/bedrock/us.anthropic.claude-sonnet-4-6``) must still inherit + ``get_model_info`` cannot resolve (e.g. a triple provider prefix like + ``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6``; a double + prefix now resolves like a routing prefix) must still inherit the built-in cache pricing for the underlying model. Before the fix ``register_model`` fell back to an empty ``existing_model`` @@ -341,7 +342,7 @@ def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key(): litellm.model_cost = litellm.get_model_cost_map(url="") builtin_key = "us.anthropic.claude-sonnet-4-6" - registered_key = f"bedrock/bedrock/{builtin_key}" + registered_key = f"bedrock/bedrock/bedrock/{builtin_key}" builtin = litellm.model_cost[builtin_key] assert builtin["cache_creation_input_token_cost"] > 0 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index c35fb2fcbe2..053fe970d3e 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1063,6 +1063,43 @@ def test_get_model_info_gemini(): assert info.get("rpm") is not None, f"{model} does not have rpm" +def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_cost_map): + """Regression LIT-4056: with the bedrock/ routing prefix (plain, converse/, or + invoke/), the exact regional cost-map entry must win over the region-stripped + base entry, matching the unprefixed control form.""" + regional = litellm.model_cost["au.anthropic.claude-opus-4-8"] + base = litellm.model_cost["anthropic.claude-opus-4-8"] + assert regional["input_cost_per_token"] > base["input_cost_per_token"] + + for model in ( + "bedrock/au.anthropic.claude-opus-4-8", + "bedrock/converse/au.anthropic.claude-opus-4-8", + "bedrock/invoke/au.anthropic.claude-opus-4-8", + ): + info = litellm.get_model_info(model=model) + assert info["key"] == "au.anthropic.claude-opus-4-8", model + assert info["input_cost_per_token"] == regional["input_cost_per_token"], model + assert info["output_cost_per_token"] == regional["output_cost_per_token"], model + + control = litellm.get_model_info(model="au.anthropic.claude-opus-4-8", custom_llm_provider="bedrock") + assert control["key"] == "au.anthropic.claude-opus-4-8" + + +def test_get_model_info_bedrock_regional_profile_without_entry_falls_back_to_base(local_model_cost_map): + """A regional profile with no dedicated cost-map entry must still resolve to its + region-stripped base entry.""" + assert "jp.anthropic.claude-opus-4-8" not in litellm.model_cost + info = litellm.get_model_info(model="bedrock/jp.anthropic.claude-opus-4-8") + assert info["key"] == "anthropic.claude-opus-4-8" + + +def test_get_model_info_bedrock_double_provider_prefix_resolves(local_model_cost_map): + """A doubled bedrock/ prefix routes at runtime via strip_bedrock_routing_prefix, + so model info must resolve it to the same entry the request actually bills as.""" + info = litellm.get_model_info(model="bedrock/bedrock/us.anthropic.claude-sonnet-4-6") + assert info["key"] == "us.anthropic.claude-sonnet-4-6" + + def test_openai_models_in_model_info(): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") From b2e2a38bc0a71d7de65ede6a92ee7b1691800bdd Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 20:51:15 -0700 Subject: [PATCH 03/31] fix(passthrough): stream non-sse passthrough responses instead of buffering in memory (#32386) * fix(passthrough): stream non-sse passthrough responses instead of buffering in memory Non-SSE passthrough responses were fully read into proxy memory (content = await response.aread()) before the first byte reached the client. For large non-JSON bodies such as Anthropic batch results jsonl files this ballooned proxy RSS to a multiple of the file size and produced near-total TTFB dead air, letting intermediaries kill the silent connection and truncate the download. The upstream request is now sent with httpx stream semantics and the buffering decision is made from the response headers: application/json (and +json) bodies plus upstream errors keep the buffered behavior since spend logging, guardrails and managed-id rewriting inspect them, while every other 2xx body is relayed as a StreamingResponse that iterates upstream bytes without accumulating them, preserving status code and headers (including x-litellm-*) and firing the success-handler logging with response_body=None once the stream completes. * fix(passthrough): log client disconnects mid-stream and derive test client cache key from production code * test(passthrough): intercept AsyncClient.send in legacy passthrough tests and assert final wire params * test(passthrough): fail with a clear assert when the passthrough client cache scan misses --- .../pass_through_endpoints.py | 143 +++++- .../pass_through_endpoints/success_handler.py | 16 +- .../test_pass_through_endpoints.py | 35 +- .../passthrough/test_passthrough_main.py | 24 +- .../test_llm_pass_through_endpoints.py | 14 +- .../test_pass_through_endpoints.py | 418 +++++++++++++++++- 6 files changed, 582 insertions(+), 68 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 5a621163760..2aff663038b 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -7,7 +7,7 @@ import traceback from base64 import b64encode from datetime import datetime from itertools import groupby -from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast +from typing import Any, AsyncGenerator, Dict, List, Mapping, Optional, Tuple, Union, cast from urllib.parse import urlencode, urlparse import httpx @@ -389,18 +389,24 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): forward_multipart: bool = False, ) -> httpx.Response: """ - Handle non-streaming HTTP requests + Handle non-SSE HTTP requests - Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests + Handles special cases when GET requests, multipart/form-data requests, and generic httpx requests. + + GET and generic requests are sent with httpx stream semantics so the caller can + decide from the response headers whether to buffer the body (JSON, inspected for + logging/guardrails) or relay it to the client without materializing it in memory + (LIT-4009: large batch results files must not be buffered in proxy RSS). """ if request.method == "GET": - response = await async_client.request( - method=request.method, - url=url, + get_request = async_client.build_request( + request.method, + url, headers=headers, params=requested_query_params, ) - elif HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart: + return await async_client.send(get_request, stream=True) + if HttpPassThroughEndpointHelpers.is_multipart(request) is True and forward_multipart: # Forward multipart via make_multipart_http_request even when _parsed_body is # non-empty (pass_through_request always injects litellm_logging_obj, etc.). # forward_multipart is False when custom_body was supplied (JSON body despite @@ -412,16 +418,14 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): headers=headers, requested_query_params=requested_query_params, ) - else: - # Generic httpx method - response = await async_client.request( - method=request.method, - url=url, - headers=headers, - params=requested_query_params, - json=_parsed_body, - ) - return response + generic_request = async_client.build_request( + request.method, + url, + headers=headers, + params=requested_query_params, + json=_parsed_body, + ) + return await async_client.send(generic_request, stream=True) @staticmethod def is_multipart(request: Request) -> bool: @@ -1161,13 +1165,14 @@ async def pass_through_request( if state_raw_body is not None: # SigV4-signed callers (Bedrock) require the exact pre-signed bytes # to be forwarded so the signature/Content-Length stay valid. - response = await async_client.request( - method=request.method, - url=url, + raw_body_request = async_client.build_request( + request.method, + url, headers=headers, params=requested_query_params, content=state_raw_body, ) + response = await async_client.send(raw_body_request, stream=True) else: response = await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler( request=request, @@ -1223,6 +1228,40 @@ async def pass_through_request( status_code=response.status_code, ) + if not _should_buffer_passthrough_response(response): + relay_custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=user_api_key_dict, + call_id=litellm_call_id, + model_id=None, + cache_key=None, + api_base=str(url._uri_reference), + ) + relay_callback_headers = await proxy_logging_obj.post_call_response_headers_hook( + data=_parsed_body or {}, + user_api_key_dict=user_api_key_dict, + response=response, + request_headers=dict(request.headers), + ) + if relay_callback_headers: + relay_custom_headers.update(relay_callback_headers) + + return StreamingResponse( + _relay_passthrough_response_bytes( + response=response, + request_body=_parsed_body or {}, + url_route=str(url), + start_time=start_time, + logging_obj=logging_obj, + custom_llm_provider=custom_llm_provider, + success_handler_kwargs=kwargs, + ), + status_code=response.status_code, + headers=HttpPassThroughEndpointHelpers.get_response_headers( + headers=response.headers, + custom_headers=relay_custom_headers, + ), + ) + content = await response.aread() ## POST-CALL GUARDRAILS ## @@ -2211,6 +2250,70 @@ def _is_streaming_response(response: httpx.Response) -> bool: return False +def _should_buffer_passthrough_response(response: httpx.Response) -> bool: + """ + Decide from the response headers whether the body must be read into memory. + + JSON bodies (and upstream errors) stay buffered: spend logging, guardrails and + managed-id rewriting inspect them, and they are small in practice. Everything + else (jsonl batch results, octet-stream files, ...) is relayed to the client + chunk by chunk so a large body is never resident in full (LIT-4009). A missing + content-type is buffered because the body cannot be classified. + """ + if response.status_code >= 400: + return True + media_type = response.headers.get("content-type", "").split(";")[0].strip().lower() + return media_type in ("", "application/json") or media_type.endswith("+json") + + +async def _relay_passthrough_response_bytes( + response: httpx.Response, + request_body: dict, + url_route: str, + start_time: datetime, + logging_obj: LiteLLMLoggingObj, + custom_llm_provider: Optional[str], + success_handler_kwargs: dict, +) -> AsyncGenerator[bytes, None]: + """ + Yield upstream bytes to the client without accumulating them, then fire the + passthrough success handler with response_body=None (uninspected body). The + finally block also runs on client disconnect (GeneratorExit) so partial + downloads still produce a spend-log row, mirroring chunk_processor; a + disconnect additionally logs a warning with the number of bytes relayed so + partial deliveries are distinguishable from complete ones in proxy logs. + """ + bytes_relayed = 0 + upstream_fully_relayed = False + try: + async for chunk in response.aiter_bytes(): + bytes_relayed += len(chunk) + yield chunk + upstream_fully_relayed = True + finally: + if not upstream_fully_relayed: + verbose_proxy_logger.warning( + f"Passthrough stream for {url_route} ended before upstream body was fully relayed; " + f"{bytes_relayed} bytes were sent to the client" + ) + await response.aclose() + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + async_coroutine=pass_through_endpoint_logging.pass_through_async_success_handler( + httpx_response=response, + response_body=None, + url_route=url_route, + result="", + start_time=start_time, + end_time=datetime.now(), + logging_obj=logging_obj, + cache_hit=False, + request_body=request_body, + custom_llm_provider=custom_llm_provider, + **success_handler_kwargs, + ) + ) + + def _extract_model_from_vertex_ai_setup(setup_response: dict) -> Optional[str]: """ Extract the model name from Vertex AI Live setup response. diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index ee651a15afe..6a673f6bebb 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -34,6 +34,18 @@ from .llm_provider_handlers.vertex_passthrough_logging_handler import ( cohere_passthrough_logging_handler = CoherePassthroughLoggingHandler() +def _safe_response_text(httpx_response: httpx.Response) -> str: + """ + Streamed passthrough responses are relayed to the client without being read + into memory, so accessing .text on them raises ResponseNotRead. Their body is + intentionally uninspected; log an empty string instead of failing the row. + """ + try: + return httpx_response.text + except httpx.ResponseNotRead: + return "" + + class PassThroughEndpointLogging: def __init__(self): self.TRACKED_VERTEX_ROUTES = [ @@ -306,7 +318,9 @@ class PassThroughEndpointLogging: ] kwargs = normalized_llm_passthrough_logging_payload["kwargs"] if standard_logging_response_object is None: - standard_logging_response_object = StandardPassThroughResponseObject(response=httpx_response.text) + standard_logging_response_object = StandardPassThroughResponseObject( + response=_safe_response_text(httpx_response) + ) kwargs = self._set_cost_per_request( logging_obj=logging_obj, diff --git a/tests/local_testing/test_pass_through_endpoints.py b/tests/local_testing/test_pass_through_endpoints.py index eeb29dea531..793a60efc3f 100644 --- a/tests/local_testing/test_pass_through_endpoints.py +++ b/tests/local_testing/test_pass_through_endpoints.py @@ -22,10 +22,8 @@ from litellm.proxy.proxy_server import initialize_pass_through_endpoints # Mock the async_client used in the pass_through_request function -async def mock_request(*args, **kwargs): - mock_response = httpx.Response(200, json={"message": "Mocked response"}) - mock_response.request = Mock(spec=httpx.Request) - return mock_response +async def mock_request(self, request, **kwargs): + return httpx.Response(200, json={"message": "Mocked response"}, request=request) def remove_rerank_route(app): @@ -49,8 +47,8 @@ def client(): @pytest.mark.asyncio async def test_pass_through_endpoint_no_headers(client, monkeypatch): - # Mock the httpx.AsyncClient.request method - monkeypatch.setattr("httpx.AsyncClient.request", mock_request) + # Mock the httpx.AsyncClient.send method + monkeypatch.setattr("httpx.AsyncClient.send", mock_request) import litellm # Define a pass-through endpoint @@ -79,8 +77,8 @@ async def test_pass_through_endpoint_no_headers(client, monkeypatch): @pytest.mark.asyncio async def test_pass_through_endpoint(client, monkeypatch): - # Mock the httpx.AsyncClient.request method - monkeypatch.setattr("httpx.AsyncClient.request", mock_request) + # Mock the httpx.AsyncClient.send method + monkeypatch.setattr("httpx.AsyncClient.send", mock_request) import litellm # Define a pass-through endpoint @@ -181,7 +179,7 @@ async def test_pass_through_endpoint_rpm_limit( expected_status_codes, num_users, ): - monkeypatch.setattr("httpx.AsyncClient.request", mock_request) + monkeypatch.setattr("httpx.AsyncClient.send", mock_request) import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache @@ -285,7 +283,7 @@ async def test_pass_through_endpoint_rpm_limit( async def test_pass_through_endpoint_sequential_rpm_limit( client, monkeypatch, auth, rpm_limit, requests_to_make, expected_status_codes ): - monkeypatch.setattr("httpx.AsyncClient.request", mock_request) + monkeypatch.setattr("httpx.AsyncClient.send", mock_request) import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ProxyLogging, hash_token, user_api_key_cache @@ -504,10 +502,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch): captured_requests = [] - async def mock_bing_request(*args, **kwargs): + async def mock_bing_request(self, request, **kwargs): - captured_requests.append((args, kwargs)) - mock_response = httpx.Response( + captured_requests.append(request) + return httpx.Response( 200, json={ "_type": "SearchResponse", @@ -518,11 +516,10 @@ async def test_pass_through_endpoint_bing(client, monkeypatch): "value": [], }, }, + request=request, ) - mock_response.request = Mock(spec=httpx.Request) - return mock_response - monkeypatch.setattr("httpx.AsyncClient.request", mock_bing_request) + monkeypatch.setattr("httpx.AsyncClient.send", mock_bing_request) # Define a pass-through endpoint pass_through_endpoints = [ @@ -555,8 +552,8 @@ async def test_pass_through_endpoint_bing(client, monkeypatch): client.get("/bing/search?q=bob+barker") client.get("/bing/search-no-merge-params?q=bob+barker") - first_transformed_url = captured_requests[0][1]["url"] - second_transformed_url = captured_requests[1][1]["url"] + first_transformed_url = captured_requests[0].url + second_transformed_url = captured_requests[1].url # Parse URLs to compare query params order-independently # Parse first URL @@ -573,7 +570,7 @@ async def test_pass_through_endpoint_bing(client, monkeypatch): "setLang": ["en-US"], "mkt": ["en-US"], } - expected_second_params = {"setLang": ["en-US"], "mkt": ["en-US"]} + expected_second_params = {"q": ["bob barker"]} # Assert the response - compare base URL and params separately assert ( diff --git a/tests/test_litellm/passthrough/test_passthrough_main.py b/tests/test_litellm/passthrough/test_passthrough_main.py index 6e9c75e085a..0b5bfac87bb 100644 --- a/tests/test_litellm/passthrough/test_passthrough_main.py +++ b/tests/test_litellm/passthrough/test_passthrough_main.py @@ -387,8 +387,9 @@ async def test_pass_through_request_stream_param_no_override( # Create mocks for the async client mock_async_client = AsyncMock() - # Mock request to return the non-streaming response - mock_async_client.request.return_value = mock_response + # Mock build_request/send to return the non-streaming response + mock_async_client.build_request = Mock(return_value=Mock()) + mock_async_client.send.return_value = mock_response # Mock get_async_httpx_client to return our mock client mock_client_obj = Mock() @@ -420,20 +421,19 @@ async def test_pass_through_request_stream_param_no_override( stream=False, # Should be used since no stream in request body ) - # Verify that build_request was NOT called (no streaming path) - mock_async_client.build_request.assert_not_called() - - # Verify that send was NOT called (no streaming path) - mock_async_client.send.assert_not_called() - - # Verify that the non-streaming request method WAS called - mock_async_client.request.assert_called_once_with( - method="POST", - url=httpx.URL("https://api.anthropic.com/v1/messages"), + # Non-SSE requests are sent with stream semantics so large bodies can + # be relayed without buffering; the JSON response below is still + # buffered into a plain Response. + mock_async_client.request.assert_not_called() + mock_async_client.build_request.assert_called_once_with( + "POST", + httpx.URL("https://api.anthropic.com/v1/messages"), headers={"Authorization": "Bearer test-key"}, params={}, json=request_body, ) + mock_async_client.send.assert_called_once() + assert mock_async_client.send.call_args.kwargs.get("stream") is True # Verify response is a regular Response (not StreamingResponse) from fastapi.responses import Response, StreamingResponse diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 8bb7b52af14..cf3351c4ff8 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -1918,7 +1918,8 @@ class TestForwardHeaders: ): # Setup mock httpx client mock_client = MagicMock() - mock_client.request = AsyncMock(return_value=mock_httpx_response) + mock_client.build_request = MagicMock(return_value=MagicMock()) + mock_client.send = AsyncMock(return_value=mock_httpx_response) mock_client_obj = MagicMock() mock_client_obj.client = mock_client mock_get_client.return_value = mock_client_obj @@ -1942,10 +1943,10 @@ class TestForwardHeaders: ) # Verify the httpx client was called - assert mock_client.request.called + assert mock_client.send.called # Get the headers that were sent to the target - call_args = mock_client.request.call_args + call_args = mock_client.build_request.call_args sent_headers = call_args[1]["headers"] # Verify user headers were forwarded (except content-length and host) @@ -2019,7 +2020,8 @@ class TestForwardHeaders: ): # Setup mock httpx client mock_client = MagicMock() - mock_client.request = AsyncMock(return_value=mock_httpx_response) + mock_client.build_request = MagicMock(return_value=MagicMock()) + mock_client.send = AsyncMock(return_value=mock_httpx_response) mock_client_obj = MagicMock() mock_client_obj.client = mock_client mock_get_client.return_value = mock_client_obj @@ -2043,10 +2045,10 @@ class TestForwardHeaders: ) # Verify the httpx client was called - assert mock_client.request.called + assert mock_client.send.called # Get the headers that were sent to the target - call_args = mock_client.request.call_args + call_args = mock_client.build_request.call_args sent_headers = call_args[1]["headers"] # Verify only custom headers were sent diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 85211f392ee..89d100cc3a4 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -1,5 +1,6 @@ import asyncio import json +import logging import os import sys from contextlib import ExitStack @@ -1337,7 +1338,8 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): upstream_response.raise_for_status = MagicMock() async_client = MagicMock() - async_client.request = AsyncMock(return_value=upstream_response) + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) mock_get_client.return_value = MagicMock(client=async_client) async def _empty_chunks(*args, **kwargs): @@ -1361,7 +1363,7 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): stream=False, ) - async_client.request.assert_awaited_once() + async_client.send.assert_awaited_once() mock_chunk_processor.assert_called_once() logging_obj = mock_chunk_processor.call_args.kwargs[ @@ -3046,7 +3048,8 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod ) mock_async_client = AsyncMock() - mock_async_client.request = AsyncMock(return_value=upstream) + mock_async_client.build_request = MagicMock(return_value=MagicMock()) + mock_async_client.send = AsyncMock(return_value=upstream) mock_client_obj = MagicMock() mock_client_obj.client = mock_async_client @@ -3082,10 +3085,12 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod stream=False, ) - mock_async_client.request.assert_called_once() - req_kw = mock_async_client.request.call_args[1] - assert req_kw.get("content") == raw_signed - assert "json" not in req_kw + mock_async_client.build_request.assert_called_once() + build_kw = mock_async_client.build_request.call_args[1] + assert build_kw.get("content") == raw_signed + assert "json" not in build_kw + mock_async_client.send.assert_awaited_once() + assert mock_async_client.send.call_args.kwargs.get("stream") is True @pytest.mark.asyncio @@ -3826,7 +3831,8 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan mock_success_handler.return_value = None async_client = MagicMock() - async_client.request = AsyncMock(return_value=upstream_response) + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) @@ -3913,7 +3919,8 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa mock_success_handler.return_value = None async_client = MagicMock() - async_client.request = AsyncMock(return_value=upstream_response) + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) @@ -4043,7 +4050,8 @@ async def test_pass_through_request_non_streaming_success_unchanged(): mock_success_handler.return_value = None async_client = MagicMock() - async_client.request = AsyncMock(return_value=upstream_response) + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) mock_get_client.return_value = MagicMock(client=async_client) mock_request = MagicMock(spec=Request) @@ -4103,3 +4111,393 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio assert int(exc_info.value.code) == 500 assert "auth backend unavailable" in exc_info.value.message + + +class _RecordingUpstreamByteStream(httpx.AsyncByteStream): + def __init__(self, chunks): + self._chunks = chunks + self.chunks_served = 0 + self.closed = False + + async def __aiter__(self): + for chunk in self._chunks: + self.chunks_served += 1 + yield chunk + + async def aclose(self): + self.closed = True + + +class _FakeUpstreamTransport(httpx.AsyncBaseTransport): + def __init__(self, status_code, headers, stream): + self._status_code = status_code + self._headers = headers + self._stream = stream + + async def handle_async_request(self, request): + return httpx.Response( + status_code=self._status_code, + headers=self._headers, + stream=self._stream, + request=request, + ) + + +def _inject_fake_passthrough_client(transport, timeout): + """Dependency-inject a fake upstream via the client cache that + get_async_httpx_client resolves passthrough clients from (no monkeypatching + of the HTTP layer). The cache entry is located by calling the production + get_async_httpx_client and identity-scanning the cache for the handler it + returned, so the internal cache-key format is never duplicated here. Must + run inside the test's event loop because cache keys are loop-scoped. + Returns (client, cleanup).""" + import litellm + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.llms.custom_http import httpxSpecialProvider + + real_handler = get_async_httpx_client( + httpxSpecialProvider.PassThroughEndpoint, + params={"timeout": resolve_pass_through_request_timeout(timeout)}, + ) + cache = litellm.in_memory_llm_clients_cache + cache_key = next( + (key for key, cached in cache.cache_dict.items() if cached is real_handler), + None, + ) + assert cache_key is not None, ( + "PassThroughEndpoint client not found in in_memory_llm_clients_cache; " + "get_async_httpx_client may not be caching this provider." + ) + fake_client = httpx.AsyncClient(transport=transport) + cache.cache_dict[cache_key] = SimpleNamespace(client=fake_client) + + def _cleanup(): + cache.cache_dict.pop(cache_key, None) + + return fake_client, _cleanup + + +def _enter_relay_logging_mocks(stack, parsed_body): + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + mock_proxy_logging = stack.enter_context( + patch("litellm.proxy.proxy_server.proxy_logging_obj") + ) + mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_success_handler = stack.enter_context( + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler" + ) + ) + mock_success_handler.return_value = None + stack.enter_context( + patch.object( + GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock() + ) + ) + return mock_proxy_logging, mock_success_handler + + +def _relay_client_request(method="GET"): + mock_request = MagicMock(spec=Request) + mock_request.method = method + mock_request.url = "http://localhost:4000/passthrough-relay/results" + mock_request.body = AsyncMock(return_value=b"") + mock_request.headers = Headers({}) + mock_request.query_params = QueryParams({}) + return mock_request + + +@pytest.mark.asyncio +async def test_pass_through_request_relays_non_json_body_without_buffering(): + """ + Regression (LIT-4009): non-SSE passthrough responses used to be fully + buffered in proxy memory (content = await response.aread()) before a single + byte reached the client, ballooning proxy RSS to a multiple of the body size + for large non-JSON downloads (e.g. Anthropic batch results .jsonl files) and + producing near-total TTFB dead air that let intermediaries kill the silent + connection mid-download. + + A non-JSON 2xx body must be relayed as a StreamingResponse whose chunks are + pulled from the upstream one at a time, with zero chunks consumed before the + handler returns, upstream status/headers plus x-litellm-* headers preserved, + and the success-handler logging fired with response_body=None once the + stream completes. Pre-fix, the handler returned a plain Response after + reading the entire body, so these assertions fail on the old code. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + upstream_chunks = ( + b'{"custom_id": "a", "result": {}}\n', + b'{"custom_id": "b", "result": {}}\n', + b'{"custom_id": "c", "result": {}}\n', + ) + upstream_stream = _RecordingUpstreamByteStream(upstream_chunks) + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={ + "content-type": "application/x-jsonl", + "x-upstream-marker": "batch-results", + "content-length": str(sum(len(c) for c in upstream_chunks)), + }, + stream=upstream_stream, + ), + timeout=311.0, + ) + try: + with ExitStack() as stack: + _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/messages/batches/b1/results", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=311.0, + ) + + assert isinstance(response, StreamingResponse) + assert upstream_stream.chunks_served == 0 + mock_success_handler.assert_not_called() + + iterator = response.body_iterator + first_chunk = await iterator.__anext__() + assert first_chunk == upstream_chunks[0] + assert upstream_stream.chunks_served == 1 + + remaining = [chunk async for chunk in iterator] + assert b"".join([first_chunk, *remaining]) == b"".join(upstream_chunks) + assert upstream_stream.closed is True + + assert response.status_code == 200 + assert response.headers["x-upstream-marker"] == "batch-results" + assert "x-litellm-call-id" in response.headers + assert "content-length" not in response.headers + + mock_success_handler.assert_called_once() + success_kwargs = mock_success_handler.call_args.kwargs + assert success_kwargs["response_body"] is None + assert ( + success_kwargs["url_route"] + == "http://upstream.test/v1/messages/batches/b1/results" + ) + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_pass_through_request_json_response_stays_buffered_for_logging(): + """ + JSON responses (content-type application/json) must keep the buffered + behavior: spend logging and guardrails inspect the parsed body, so the + handler reads the full upstream body and passes the parsed dict to the + success handler. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + upstream_chunks = (b'{"id": "file-123"', b', "status": "processed"}') + upstream_stream = _RecordingUpstreamByteStream(upstream_chunks) + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "application/json"}, + stream=upstream_stream, + ), + timeout=312.0, + ) + try: + with ExitStack() as stack: + _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/files/file-123", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=312.0, + ) + + assert not isinstance(response, StreamingResponse) + assert response.status_code == 200 + assert response.body == b"".join(upstream_chunks) + assert upstream_stream.chunks_served == len(upstream_chunks) + + mock_success_handler.assert_called_once() + success_kwargs = mock_success_handler.call_args.kwargs + assert success_kwargs["response_body"] == { + "id": "file-123", + "status": "processed", + } + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_pass_through_request_upstream_error_body_stays_buffered(): + """ + Upstream errors are never relayed as a stream, whatever their content-type: + the body must stay available for the failure hook and reach the client + buffered with the upstream status code, exactly as before the fix. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + upstream_stream = _RecordingUpstreamByteStream((b"upstream ", b"exploded")) + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=502, + headers={"content-type": "application/x-jsonl"}, + stream=upstream_stream, + ), + timeout=313.0, + ) + try: + with ExitStack() as stack: + mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks( + stack, {} + ) + + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/messages/batches/b1/results", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=313.0, + ) + + assert not isinstance(response, StreamingResponse) + assert response.status_code == 502 + assert response.body == b"upstream exploded" + mock_proxy_logging.post_call_failure_hook.assert_called_once() + mock_success_handler.assert_not_called() + finally: + cleanup() + await fake_client.aclose() + + +_PARTIAL_RELAY_WARNING_MARKER = "ended before upstream body was fully relayed" + + +@pytest.mark.asyncio +async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(caplog): + """ + Regression: when the client disconnects mid-relay (GeneratorExit), the + proxy log must record that the upstream body was only partially delivered, + including the route and the byte count that reached the client, while the + success handler still fires so the partial delivery produces a spend-log + row. Pre-fix, the finally block fired the success handler silently and a + partial delivery was indistinguishable from a complete one. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n') + upstream_stream = _RecordingUpstreamByteStream(upstream_chunks) + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "application/x-jsonl"}, + stream=upstream_stream, + ), + timeout=314.0, + ) + try: + with ExitStack() as stack: + _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/messages/batches/b1/results", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=314.0, + ) + + assert isinstance(response, StreamingResponse) + iterator = response.body_iterator + first_chunk = await iterator.__anext__() + assert first_chunk == upstream_chunks[0] + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await iterator.aclose() + + partial_relay_warnings = [ + record.getMessage() + for record in caplog.records + if record.levelno == logging.WARNING + and _PARTIAL_RELAY_WARNING_MARKER in record.getMessage() + ] + assert len(partial_relay_warnings) == 1 + assert ( + "http://upstream.test/v1/messages/batches/b1/results" + in partial_relay_warnings[0] + ) + assert ( + f"{len(first_chunk)} bytes were sent to the client" + in partial_relay_warnings[0] + ) + + assert upstream_stream.closed is True + mock_success_handler.assert_called_once() + assert mock_success_handler.call_args.kwargs["response_body"] is None + finally: + cleanup() + await fake_client.aclose() + + +@pytest.mark.asyncio +async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning(caplog): + """ + A fully consumed relay must not be reported as a partial delivery: the + success handler fires and no partial-relay warning is logged. + """ + from fastapi.responses import StreamingResponse + + from litellm.proxy._types import UserAPIKeyAuth + + upstream_chunks = (b'{"custom_id": "a"}\n', b'{"custom_id": "b"}\n') + upstream_stream = _RecordingUpstreamByteStream(upstream_chunks) + fake_client, cleanup = _inject_fake_passthrough_client( + _FakeUpstreamTransport( + status_code=200, + headers={"content-type": "application/x-jsonl"}, + stream=upstream_stream, + ), + timeout=315.0, + ) + try: + with ExitStack() as stack: + _, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + + response = await pass_through_request( + request=_relay_client_request(), + target="http://upstream.test/v1/messages/batches/b1/results", + custom_headers={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-relay-test"), + timeout=315.0, + ) + + assert isinstance(response, StreamingResponse) + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + relayed = [chunk async for chunk in response.body_iterator] + + assert b"".join(relayed) == b"".join(upstream_chunks) + assert not any( + _PARTIAL_RELAY_WARNING_MARKER in record.getMessage() + for record in caplog.records + ) + mock_success_handler.assert_called_once() + finally: + cleanup() + await fake_client.aclose() From 7cc660866aea077508246d95e3f77cb8b940d212 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 7 Jul 2026 21:42:08 -0700 Subject: [PATCH 04/31] fix(ui/mcp): do not reset in-flight OAuth resume when create modal mounts closed (#32416) --- .../mcp_tools/create_mcp_server.test.tsx | 16 ++++++++++++++++ .../components/mcp_tools/create_mcp_server.tsx | 10 ++++++++-- 2 files changed, 24 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx index 9f381503d18..9b4d159ae0c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx @@ -1067,6 +1067,22 @@ describe("CreateMCPServer", () => { const reopenedUrlInput = screen.getByPlaceholderText("https://your-mcp-server.com") as HTMLInputElement; expect(reopenedUrlInput.value).toBe(""); }); + + it("does not reset an in-flight OAuth resume when mounted with the modal closed (post-redirect restore)", () => { + // After the "Authorize & Fetch Token" redirect the page reloads and this + // component mounts with isModalVisible=false while useMcpOAuthFlow is still + // exchanging the authorization code. Calling reset() during that mount bumps + // the hook's reset version and the fetched token is silently discarded, so + // the user sees no Connection Status / Tool Configuration and must authorize + // again after saving. + const { rerender } = render(); + expect(oauthHook.reset).not.toHaveBeenCalled(); + + // A real open -> closed transition must still reset (the #30000 leak fix). + rerender(); + rerender(); + expect(oauthHook.reset).toHaveBeenCalled(); + }); }); describe("when stdio transport is selected", () => { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index 815db5fb841..a4e09bd6f6b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -626,9 +626,15 @@ const CreateMCPServer: React.FC = ({ // Clear form, tools, and OAuth state when the modal closes so a previous server's // authorization, credentials, or tool list never bleed into the next "Add New MCP // Server" session, including when a parent dismisses the modal without routing - // through handleCancel or handleCreate. + // through handleCancel or handleCreate. Only a real open -> closed transition may + // trigger this: on the post-OAuth-redirect remount the modal starts closed while + // resumeOAuthFlow's token exchange is in flight, and resetting then discards the + // fetched token. + const wasModalVisibleRef = React.useRef(isModalVisible); React.useEffect(() => { - if (!isModalVisible) { + const wasVisible = wasModalVisibleRef.current; + wasModalVisibleRef.current = isModalVisible; + if (!isModalVisible && wasVisible) { form.resetFields(); setFormValues({}); setOauthAccessToken(null); From f922be32f0bb85cf014fd92f0b80cb2d8655f536 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 7 Jul 2026 21:46:36 -0700 Subject: [PATCH 05/31] fix(mcp): accept integer progressToken in host progress capture (#32402) --- .../proxy/_experimental/mcp_server/server.py | 4 +- .../mcp_server/test_mcp_tool_search.py | 41 +++++++++++++++++++ 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e3812522ded..fc847182a60 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -716,7 +716,7 @@ if MCP_AVAILABLE: if not (host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta): return None host_token = getattr(host_ctx.meta, "progressToken", None) - if not (host_token and hasattr(host_ctx, "session") and host_ctx.session): + if host_token is None or not (hasattr(host_ctx, "session") and host_ctx.session): return None host_session = host_ctx.session @@ -732,7 +732,7 @@ if MCP_AVAILABLE: except Exception as e: verbose_logger.error(f"Failed to forward progress to Host: {e}") - verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...") + verbose_logger.debug(f"Host progressToken captured: {str(host_token)[:8]}...") return forward_progress async def _build_virtual_call_logging_obj( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index f2b74d65059..5c2a04456b0 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -789,6 +789,47 @@ class TestCaptureHostProgressCallback: host.request_context.session = MagicMock() assert callable(_capture_host_progress_callback(host)) + def test_returns_callable_when_token_is_integer(self) -> None: + from litellm.proxy._experimental.mcp_server.server import ( + _capture_host_progress_callback, + ) + + host = MagicMock() + host.request_context.meta.progressToken = 12345 + host.request_context.session = MagicMock() + assert callable(_capture_host_progress_callback(host)) + + def test_returns_callable_when_token_is_zero(self) -> None: + from litellm.proxy._experimental.mcp_server.server import ( + _capture_host_progress_callback, + ) + + host = MagicMock() + host.request_context.meta.progressToken = 0 + host.request_context.session = MagicMock() + assert callable(_capture_host_progress_callback(host)) + + @pytest.mark.asyncio + async def test_forwarded_progress_token_preserves_integer_value(self) -> None: + from litellm.proxy._experimental.mcp_server.server import ( + _capture_host_progress_callback, + ) + + host = MagicMock() + host.request_context.meta.progressToken = 12345 + session = AsyncMock() + host.request_context.session = session + + callback = _capture_host_progress_callback(host) + assert callback is not None + await callback(0.5, 1.0) + + session.send_progress_notification.assert_awaited_once_with( + progress_token=12345, + progress=0.5, + total=1.0, + ) + class TestHandleListToolsVirtual: """Covers the protocol list_tools early-return when the flag is enabled.""" From c212c168529321d7fadc446b431001c2d88412d3 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 21:56:05 -0700 Subject: [PATCH 06/31] ci: ratchet LIT003 budget down to current count to remove suppression slack (#32423) --- type-discipline-budget.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 8d12528d8ab..2c44ab5a049 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -6,7 +6,7 @@ "limit": 27522 }, "LIT003": { - "limit": 422 + "limit": 292 }, "LIT004": { "limit": 44 From 404ec7fc2ee18edc885db3cb47c2bb682799bef2 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 21:57:46 -0700 Subject: [PATCH 07/31] ci(llm_responses_api_testing): bound live re-record calls and rerun timeout-only failures to stop 15m no-output kills (#32420) --- .circleci/config.yml | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index ce9aaa9be8a..b0a705966a2 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1029,6 +1029,8 @@ jobs: - *python312_image working_directory: ~/project resource_class: large + environment: + REQUEST_TIMEOUT: "180" steps: - checkout @@ -1058,7 +1060,8 @@ jobs: -v -x \ --junitxml=test-results/junit.xml \ --durations=5 \ - -n 8" + -n 8 \ + --reruns 1 --only-rerun Timeout" no_output_timeout: 15m # Store test results From 06a43d11c4810955f2319a243f3ab88dcfbebbae Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 21:58:14 -0700 Subject: [PATCH 08/31] test(responses): bound azure shell tool live call at 90s and skip on provider timeout (#32424) * ci(responses): bound azure shell tool e2e call and enforce per-test timeout The azure variant of test_responses_api_shell_tool always makes a live Azure call (its skip outcome means no VCR cassette is ever persisted). When Azure held the connection instead of answering, the call sat on litellm's 6000s responses deadline until CircleCI killed the whole job via no_output_timeout after 15m of silence (job 2013288). Bound the e2e call at 90s and skip on litellm.Timeout, matching the existing InternalServerError and BadRequestError skips, and give the llm_responses_api_testing job the same pytest-timeout guard the llm_translation_testing job already uses so no single hung test can consume the 15m no-output window again. * test(responses): drop job-level pytest timeout, keep shell tool 90s bound --- tests/llm_responses_api_testing/base_responses_api.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index 30f444b9acc..7d2e30f8372 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -765,7 +765,10 @@ class BaseResponsesAPITest(ABC): max_output_tokens=256, tools=tools, tool_choice="auto", + timeout=90, ) + except litellm.Timeout: + pytest.skip("Provider did not answer the shell tool request within 90s") except litellm.InternalServerError: pytest.skip("Skipping test due to litellm.InternalServerError") except litellm.BadRequestError as e: From 6df5e1b263a77a25a5bb483015fd13a79f3ef410 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 7 Jul 2026 22:10:58 -0700 Subject: [PATCH 09/31] ci: skip unit test workflows when only ui or markdown files change (#32422) * ci: skip unit test workflows when only docs or ui files change Mirror the CircleCI backend path filter (.circleci/scripts/classify_changes.sh) in the GitHub Actions unit test workflows by adding paths-ignore for ui/**, docs/**, *.md and *.mdx to every test-unit-*.yml pull_request trigger * ci: drop docs/** from unit test paths-ignore since the folder no longer exists --- .github/workflows/test-unit-core-utils.yml | 4 ++++ .github/workflows/test-unit-documentation.yml | 4 ++++ .github/workflows/test-unit-enterprise-routing.yml | 4 ++++ .github/workflows/test-unit-integrations.yml | 4 ++++ .github/workflows/test-unit-llm-providers.yml | 4 ++++ .github/workflows/test-unit-misc.yml | 4 ++++ .github/workflows/test-unit-proxy-auth.yml | 4 ++++ .github/workflows/test-unit-proxy-db.yml | 4 ++++ .github/workflows/test-unit-proxy-endpoints.yml | 4 ++++ .github/workflows/test-unit-proxy-infra.yml | 4 ++++ .github/workflows/test-unit-proxy-legacy.yml | 4 ++++ .github/workflows/test-unit-responses-caching-types.yml | 4 ++++ 12 files changed, 48 insertions(+) diff --git a/.github/workflows/test-unit-core-utils.yml b/.github/workflows/test-unit-core-utils.yml index d6d6353238f..e563679660b 100644 --- a/.github/workflows/test-unit-core-utils.yml +++ b/.github/workflows/test-unit-core-utils.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml index 4cef791a9b3..2c3d6e46618 100644 --- a/.github/workflows/test-unit-documentation.yml +++ b/.github/workflows/test-unit-documentation.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-enterprise-routing.yml b/.github/workflows/test-unit-enterprise-routing.yml index 13136c968d1..7a9b8b00f26 100644 --- a/.github/workflows/test-unit-enterprise-routing.yml +++ b/.github/workflows/test-unit-enterprise-routing.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-integrations.yml b/.github/workflows/test-unit-integrations.yml index c95ed4e7c24..b28ba3456ce 100644 --- a/.github/workflows/test-unit-integrations.yml +++ b/.github/workflows/test-unit-integrations.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-llm-providers.yml b/.github/workflows/test-unit-llm-providers.yml index df78564ab0c..fecdcbd3b95 100644 --- a/.github/workflows/test-unit-llm-providers.yml +++ b/.github/workflows/test-unit-llm-providers.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index 7c3b195f0ad..dbc3bfc8191 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-auth.yml b/.github/workflows/test-unit-proxy-auth.yml index 97dfaed6e81..ad534cc0098 100644 --- a/.github/workflows/test-unit-proxy-auth.yml +++ b/.github/workflows/test-unit-proxy-auth.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 2ac9a3b7c1c..35a1a9c78a0 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -5,6 +5,10 @@ on: branches: - main - litellm_internal_staging + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index cbb36eebdb9..7eb3d7719c0 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" workflow_dispatch: permissions: diff --git a/.github/workflows/test-unit-proxy-infra.yml b/.github/workflows/test-unit-proxy-infra.yml index 884d62289b9..cb944de5cf9 100644 --- a/.github/workflows/test-unit-proxy-infra.yml +++ b/.github/workflows/test-unit-proxy-infra.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-legacy.yml b/.github/workflows/test-unit-proxy-legacy.yml index 8db218cd1fc..9798a4e2277 100644 --- a/.github/workflows/test-unit-proxy-legacy.yml +++ b/.github/workflows/test-unit-proxy-legacy.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-responses-caching-types.yml b/.github/workflows/test-unit-responses-caching-types.yml index 2f177587997..7331544de24 100644 --- a/.github/workflows/test-unit-responses-caching-types.yml +++ b/.github/workflows/test-unit-responses-caching-types.yml @@ -7,6 +7,10 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" + paths-ignore: + - "ui/**" + - "**.md" + - "**.mdx" permissions: contents: read From d6cbf6e7e320f64138ccb0bcb847baae394fde2b Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 7 Jul 2026 22:47:03 -0700 Subject: [PATCH 10/31] feat(ui): expose MCP max_concurrent_requests in server create and edit forms (#32397) * feat(ui): expose MCP max_concurrent_requests in server create and edit forms The proxy has enforced a per-server outbound tool-call concurrency cap (max_concurrent_requests) across every MCP egress path since #31641, and the management API has accepted the field on create and update all along, but the dashboard offered no way to set it. Add an optional Max Concurrent Requests input to the MCP server create and edit forms; it applies to every auth type and transport, so it renders unconditionally rather than gated on auth mode. Clearing the field on edit sends null so the stored limit is unset. Also rebuild the per-server semaphore when the configured limit changes. Previously the semaphore was created once per server_id and never resized, so an edited limit only took effect after a proxy restart even though the new value was persisted and reloaded into the registry. * feat(ui): mark MCP max concurrent requests field label as optional * test(ui): stop OBO create-form tests from timing out on CI The token-exchange payload test and the Entra scope-required test filled five text fields with user.type, which dispatches a full keystroke sequence per character; every input event runs the antd form onValuesChange handler and re-renders the whole CreateMCPServer tree, roughly 120 renders per test. As the form grew the two tests reached 8s and 18s locally, which crosses the 30s vitest timeout on slower CI containers; ui_unit_tests failed twice this way. Switch the plain text fields to fireEvent.change (one input event per field), matching the existing stdio test pattern. Both tests assert form output, not keystroke behavior, and now run in about 3s each. --- .../mcp_server/mcp_server_manager.py | 15 ++-- .../test_mcp_max_concurrent_requests.py | 19 ++++ .../mcp_tools/create_mcp_server.test.tsx | 89 +++++++++++++++---- .../mcp_tools/create_mcp_server.tsx | 22 ++++- .../mcp_tools/mcp_server_edit.test.tsx | 80 +++++++++++++++++ .../components/mcp_tools/mcp_server_edit.tsx | 20 +++++ .../src/components/mcp_tools/types.tsx | 1 + 7 files changed, 220 insertions(+), 26 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 7c88c903324..39da1ba4a97 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -715,8 +715,10 @@ class MCPServerManager: # Per-server outbound tool-call concurrency limiters, lazily created from # each server's max_concurrent_requests. Keyed by server_id so the cap # survives the registry atomic-swap on config reload; a missing key means - # the server has no configured limit. - self._server_call_semaphores: dict[str, asyncio.Semaphore] = {} + # the server has no configured limit. The limit is cached alongside the + # semaphore so an edited limit rebuilds it instead of keeping the old cap + # until restart. + self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {} self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {} """ { @@ -3594,10 +3596,11 @@ class MCPServerManager: limit = mcp_server.max_concurrent_requests if limit is None or limit <= 0: return None - semaphore = self._server_call_semaphores.get(mcp_server.server_id) - if semaphore is None: - semaphore = asyncio.Semaphore(limit) - self._server_call_semaphores[mcp_server.server_id] = semaphore + cached = self._server_call_semaphores.get(mcp_server.server_id) + if cached is not None and cached[0] == limit: + return cached[1] + semaphore = asyncio.Semaphore(limit) + self._server_call_semaphores[mcp_server.server_id] = (limit, semaphore) return semaphore @asynccontextmanager diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py index 8c4d81223aa..e11897b65c2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_max_concurrent_requests.py @@ -164,6 +164,25 @@ async def test_openapi_backed_server_also_respects_the_cap(): assert tracker.peak_by_server["srv-openapi"] == 2 +@pytest.mark.asyncio +async def test_edited_limit_takes_effect_without_restart(): + """Editing max_concurrent_requests must rebuild the cached semaphore so the + new cap applies to subsequent calls immediately, not only after a restart.""" + manager = MCPServerManager() + server = _make_server("srv-edited", max_concurrent_requests=3) + + before_edit = _ConcurrencyTracker() + with _patch_client_with_tracker(manager, before_edit): + await _fire(manager, server, n=6) + assert before_edit.peak_by_server["srv-edited"] == 3 + + server.max_concurrent_requests = 1 + after_edit = _ConcurrencyTracker() + with _patch_client_with_tracker(manager, after_edit): + await _fire(manager, server, n=6) + assert after_edit.peak_by_server["srv-edited"] == 1 + + def test_semaphore_is_reused_per_server_and_distinct_across_servers(): manager = MCPServerManager() server_a = _make_server("srv-a", max_concurrent_requests=3) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx index 9b4d159ae0c..36af2f8d9fc 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.test.tsx @@ -388,16 +388,58 @@ describe("CreateMCPServer", () => { expect(screen.queryByText("Subject Token Type (optional)")).not.toBeInTheDocument(); }); - it("routes OAuth Token Exchange (OBO) config to the backend payload", async () => { + it("sends max_concurrent_requests in the create payload when set", async () => { await selectHttpTransport(); const user = userEvent.setup({ delay: null }); const nameInput = getServerNameInput(); - await user.type(nameInput, "TE_Server"); + await user.type(nameInput, "Limited_Server"); const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com"); - await user.type(urlInput, "https://upstream.example.com/mcp"); + await user.type(urlInput, "https://example.com/mcp"); + + await selectAntOption("Authentication", "None"); + + const limitInput = screen.getByPlaceholderText("e.g. 10"); + await user.type(limitInput, "5"); + + vi.mocked(networking.createMCPServer).mockResolvedValue({ + server_id: "new-server-1", + server_name: "Limited_Server", + alias: "Limited_Server", + url: "https://example.com/mcp", + transport: "http", + auth_type: "none", + created_at: "2024-01-01T00:00:00Z", + created_by: "user-1", + updated_at: "2024-01-01T00:00:00Z", + updated_by: "user-1", + }); + + const submitButton = screen.getByRole("button", { name: "Add MCP Server" }); + await act(async () => { + fireEvent.click(submitButton); + }); + + await waitFor(() => { + expect(networking.createMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0]; + expect(payload.max_concurrent_requests).toBe(5); + }); + + it("routes OAuth Token Exchange (OBO) config to the backend payload", async () => { + await selectHttpTransport(); + + // fireEvent.change over user.type: this test asserts payload shape, not + // keystroke behavior, and char-by-char typing re-renders the whole form + // per character, which pushed this test past the 30s CI timeout. + fireEvent.change(getServerNameInput(), { target: { value: "TE_Server" } }); + + const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com"); + fireEvent.change(urlInput, { target: { value: "https://upstream.example.com/mcp" } }); await selectAntOption("Authentication", "OAuth Token Exchange (OBO)"); @@ -405,12 +447,15 @@ describe("CreateMCPServer", () => { expect(screen.getByPlaceholderText("https://idp.example.com/oauth2/token")).toBeInTheDocument(); }); - await user.type( - screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), - "https://idp.example.com/oauth2/token", - ); - await user.type(screen.getByPlaceholderText("Enter OAuth client ID"), "te-client-id"); - await user.type(screen.getByPlaceholderText("Enter OAuth client secret"), "te-client-secret"); + fireEvent.change(screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), { + target: { value: "https://idp.example.com/oauth2/token" }, + }); + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client ID"), { + target: { value: "te-client-id" }, + }); + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client secret"), { + target: { value: "te-client-secret" }, + }); vi.mocked(networking.createMCPServer).mockResolvedValue({ server_id: "new-server-te", @@ -447,10 +492,13 @@ describe("CreateMCPServer", () => { it("makes scope required when the Entra OBO profile is selected", async () => { await selectHttpTransport(); - const user = userEvent.setup({ delay: null }); - - await user.type(getServerNameInput(), "Entra_Server"); - await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://upstream.example.com/mcp"); + // fireEvent.change over user.type for the same reason as the payload + // test above: char-by-char typing re-renders the whole form per + // character and pushes this test toward the 30s CI timeout. + fireEvent.change(getServerNameInput(), { target: { value: "Entra_Server" } }); + fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), { + target: { value: "https://upstream.example.com/mcp" }, + }); await selectAntOption("Authentication", "OAuth Token Exchange (OBO)"); await waitFor(() => { @@ -459,12 +507,15 @@ describe("CreateMCPServer", () => { await selectAntOption("Profile", "Microsoft Entra OBO"); - await user.type( - screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), - "https://login.microsoftonline.com/tenant/oauth2/v2.0/token", - ); - await user.type(screen.getByPlaceholderText("Enter OAuth client ID"), "entra-client"); - await user.type(screen.getByPlaceholderText("Enter OAuth client secret"), "entra-secret"); + fireEvent.change(screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), { + target: { value: "https://login.microsoftonline.com/tenant/oauth2/v2.0/token" }, + }); + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client ID"), { + target: { value: "entra-client" }, + }); + fireEvent.change(screen.getByPlaceholderText("Enter OAuth client secret"), { + target: { value: "entra-secret" }, + }); // Selecting Entra OBO makes the scope required; submitting without one is blocked by validation // (rfc8693 would not require it), which confirms the profile selection took effect. diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index a4e09bd6f6b..10668468c15 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -1,5 +1,5 @@ import React, { useState } from "react"; -import { Modal, Tooltip, Form, Select, Input, Switch, Collapse } from "antd"; +import { Modal, Tooltip, Form, Select, Input, InputNumber, Switch, Collapse } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Button, TextInput } from "@tremor/react"; import { createMCPServer, registerMCPServer, storeMCPOAuthUserCredential } from "../networking"; @@ -929,6 +929,26 @@ const CreateMCPServer: React.FC = ({ )} + + Max Concurrent Requests (optional) + + + + + } + name="max_concurrent_requests" + > + + + {/* Authentication - show for HTTP, SSE, and OpenAPI */} {transportType !== "stdio" && transportType !== "" && ( { }); }); }); + +describe("MCPServerEdit (max concurrent requests)", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + const limitedServer = { + ...interactiveOAuthServer, + auth_type: "none", + max_concurrent_requests: 5, + }; + + it("prefills the existing limit and sends an updated value in the payload", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...limitedServer, + max_concurrent_requests: 2, + }); + + render( + , + ); + + const limitInput = screen.getByPlaceholderText("e.g. 10") as HTMLInputElement; + expect(limitInput.value).toBe("5"); + + fireEvent.change(limitInput, { target: { value: "2" } }); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.max_concurrent_requests).toBe(2); + }); + + it("sends null when the limit is cleared so the backend unsets it", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...limitedServer, + max_concurrent_requests: null, + }); + + render( + , + ); + + const limitInput = screen.getByPlaceholderText("e.g. 10") as HTMLInputElement; + expect(limitInput.value).toBe("5"); + + fireEvent.change(limitInput, { target: { value: "" } }); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.max_concurrent_requests).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index faf7737e995..70632b459fc 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -852,6 +852,26 @@ const MCPServerEdit: React.FC = ({ )} + + Max Concurrent Requests (optional) + + + + + } + name="max_concurrent_requests" + > + + + {/* Authentication - for HTTP, SSE, and OpenAPI */} {!isStdioTransport && ( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 7e02b8779f5..9469d7bd89e 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -261,6 +261,7 @@ export interface MCPServer { available_on_public_internet?: boolean; delegate_auth_to_upstream?: boolean; oauth_passthrough?: boolean; + max_concurrent_requests?: number | null; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null; From 34db5f4813ab3449ef489a17b9d7b3da9d7c6635 Mon Sep 17 00:00:00 2001 From: Thibault Serot Date: Wed, 8 Jul 2026 16:26:47 +1000 Subject: [PATCH 11/31] feat(ui): add start time sort toggle to session logs sidebar --- .../LogDetailsDrawer.test.tsx | 105 ++++++++++++++++++ .../LogDetailsDrawer/LogDetailsDrawer.tsx | 39 ++++--- .../view_logs/LogDetailsDrawer/utils.test.ts | 30 +++++ .../view_logs/LogDetailsDrawer/utils.ts | 21 ++++ 4 files changed, 179 insertions(+), 16 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx new file mode 100644 index 00000000000..db0d168fcac --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx @@ -0,0 +1,105 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; +import { LogDetailsDrawer } from "./LogDetailsDrawer"; +import { sessionSpendLogsCall } from "../../networking"; +import { LogEntry } from "../columns"; + +vi.mock("../../networking", () => ({ + sessionSpendLogsCall: vi.fn(), +})); + +vi.mock("@/app/(dashboard)/hooks/logDetails/useLogDetails", () => ({ + useLogDetails: () => ({ data: null, isLoading: false }), +})); + +vi.mock("./LogDetailContent", () => ({ + LogDetailContent: () => null, + GuardrailJumpLink: () => null, +})); + +vi.mock("./DrawerHeader", () => ({ + DrawerHeader: () => null, +})); + +const makeLog = (overrides: Partial): LogEntry => ({ + request_id: "req", + api_key: "", + team_id: "", + model: "", + model_id: "", + call_type: "acompletion", + spend: 0, + total_tokens: 0, + prompt_tokens: 0, + completion_tokens: 0, + startTime: "2026-07-08T10:00:00.000Z", + endTime: "2026-07-08T10:00:01.000Z", + cache_hit: "false", + messages: [], + response: {}, + ...overrides, +}); + +const sessionLogs = [ + makeLog({ + request_id: "llm-early", + model: "llm-early", + startTime: "2026-07-08T10:00:00.000Z", + endTime: "2026-07-08T10:00:02.000Z", + }), + makeLog({ + request_id: "mcp-early", + model: "tool-early", + call_type: "call_mcp_tool", + startTime: "2026-07-08T10:00:01.000Z", + endTime: "2026-07-08T10:00:01.500Z", + }), + makeLog({ + request_id: "llm-late", + model: "llm-late", + startTime: "2026-07-08T10:00:02.000Z", + endTime: "2026-07-08T10:00:04.000Z", + }), + makeLog({ + request_id: "mcp-late", + model: "tool-late", + call_type: "call_mcp_tool", + startTime: "2026-07-08T10:00:03.000Z", + endTime: "2026-07-08T10:00:03.500Z", + }), +]; + +const renderSessionDrawer = () => { + vi.mocked(sessionSpendLogsCall).mockResolvedValue({ data: sessionLogs, total: 4, total_pages: 1 }); + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + render( + + {}} logEntry={null} sessionId="session-1" accessToken="token" /> + , + ); +}; + +const sidebarEventNames = () => + screen.queryAllByText(/^(llm-early|llm-late|tool-early|tool-late)$/).map((el) => el.textContent); + +describe("LogDetailsDrawer session sidebar sorting", () => { + it("defaults to grouped order: LLM calls newest first, MCP calls grouped last", async () => { + renderSessionDrawer(); + await waitFor(() => expect(sidebarEventNames()).toHaveLength(4)); + expect(sidebarEventNames()).toEqual(["llm-late", "llm-early", "tool-late", "tool-early"]); + }); + + it("switches to chronological order across LLM and MCP calls when Start time is selected", async () => { + renderSessionDrawer(); + await waitFor(() => expect(sidebarEventNames()).toHaveLength(4)); + + fireEvent.click(screen.getByText("Start time")); + + await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-early", "tool-early", "llm-late", "tool-late"])); + + fireEvent.click(screen.getByText("Grouped")); + + await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-late", "llm-early", "tool-late", "tool-early"])); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index bf3360a5371..36299139a13 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -1,5 +1,5 @@ import { useEffect, useMemo, useState } from "react"; -import { Button, Drawer } from "antd"; +import { Button, Drawer, Segmented } from "antd"; import { CheckOutlined, CopyOutlined, LeftOutlined, RightOutlined } from "@ant-design/icons"; import { Bot, Sparkles, Wrench } from "lucide-react"; import { LogEntry } from "../columns"; @@ -11,7 +11,7 @@ import { LogDetailContent, GuardrailJumpLink } from "./LogDetailContent"; import { sessionSpendLogsCall } from "../../networking"; import { useQuery } from "@tanstack/react-query"; import { getSpendString } from "@/utils/dataUtils"; -import { normalizeGuardrailEntries } from "./utils"; +import { normalizeGuardrailEntries, sortSessionLogs, SessionLogSortMode } from "./utils"; import { DRAWER_WIDTH } from "./constants"; import { useLogDetails } from "@/app/(dashboard)/hooks/logDetails/useLogDetails"; @@ -117,6 +117,7 @@ export function LogDetailsDrawer({ }: LogDetailsDrawerProps) { const isSessionMode = Boolean(sessionId); const [selectedSessionRequestId, setSelectedSessionRequestId] = useState(null); + const [sessionSortMode, setSessionSortMode] = useState("grouped"); const [isSidebarCollapsed, setIsSidebarCollapsed] = useState(false); const [copiedLeftPanelId, setCopiedLeftPanelId] = useState(false); @@ -152,26 +153,20 @@ export function LogDetailsDrawer({ // backend omits total, so the truncation note reflects what was fetched. const total: number = firstPage.total ?? rows.length; - const logs = rows - .map((row) => ({ - ...row, - request_duration_ms: row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime), - })) - .sort((a, b) => { - const aIsMcp = MCP_CALL_TYPES.includes(a.call_type) ? 1 : 0; - const bIsMcp = MCP_CALL_TYPES.includes(b.call_type) ? 1 : 0; - if (aIsMcp !== bIsMcp) return aIsMcp - bIsMcp; - // Newest first, matching the all-sessions logs overview. MCP calls - // stay grouped last (above), newest-first within that group too. - return new Date(b.startTime).getTime() - new Date(a.startTime).getTime(); - }); + const logs = rows.map((row) => ({ + ...row, + request_duration_ms: row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime), + })); return { logs, total }; }, enabled: Boolean(open && isSessionMode && sessionId && accessToken), }); - const sessionLogs: LogEntry[] = sessionData?.logs ?? []; + const sessionLogs: LogEntry[] = useMemo( + () => sortSessionLogs(sessionData?.logs ?? [], sessionSortMode), + [sessionData, sessionSortMode], + ); // total reported by the backend; when the page cap truncates the fetch this // exceeds sessionLogs.length, which drives the "showing most recent" note. const sessionTotalCount = sessionData?.total ?? sessionLogs.length; @@ -391,6 +386,18 @@ export function LogDetailsDrawer({ Showing most recent {logsForList.length} of {sessionTotalCount} )} + {isSessionMode && ( + setSessionSortMode(value as SessionLogSortMode)} + /> + )}
diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts new file mode 100644 index 00000000000..cbe12f5c101 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts @@ -0,0 +1,30 @@ +import { describe, expect, it } from "vitest"; +import { sortSessionLogs } from "./utils"; + +const llm = (id: string, startTime: string) => ({ request_id: id, call_type: "acompletion", startTime }); +const mcp = (id: string, startTime: string) => ({ request_id: id, call_type: "call_mcp_tool", startTime }); + +const ids = (rows: { request_id: string }[]) => rows.map((row) => row.request_id); + +describe("sortSessionLogs", () => { + const rows = [ + mcp("mcp-early", "2026-07-08T10:00:01.000Z"), + llm("llm-late", "2026-07-08T10:00:02.000Z"), + mcp("mcp-late", "2026-07-08T10:00:03.000Z"), + llm("llm-early", "2026-07-08T10:00:00.000Z"), + ]; + + it("grouped mode keeps MCP calls last, newest first within each group", () => { + expect(ids(sortSessionLogs(rows, "grouped"))).toEqual(["llm-late", "llm-early", "mcp-late", "mcp-early"]); + }); + + it("chronological mode interleaves all calls by start time, oldest first", () => { + expect(ids(sortSessionLogs(rows, "chronological"))).toEqual(["llm-early", "mcp-early", "llm-late", "mcp-late"]); + }); + + it("does not mutate the input array", () => { + const input = [...rows]; + sortSessionLogs(input, "chronological"); + expect(ids(input)).toEqual(ids(rows)); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts index 61301cf5b54..5a1a0e81f96 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts @@ -3,6 +3,27 @@ * These functions handle data formatting, validation, and guardrail calculations. */ +import { MCP_CALL_TYPES } from "../constants"; + +export type SessionLogSortMode = "grouped" | "chronological"; + +export function sortSessionLogs( + rows: T[], + mode: SessionLogSortMode, +): T[] { + if (mode === "chronological") { + return [...rows].sort((a, b) => new Date(a.startTime).getTime() - new Date(b.startTime).getTime()); + } + return [...rows].sort((a, b) => { + const aIsMcp = MCP_CALL_TYPES.includes(a.call_type) ? 1 : 0; + const bIsMcp = MCP_CALL_TYPES.includes(b.call_type) ? 1 : 0; + if (aIsMcp !== bIsMcp) return aIsMcp - bIsMcp; + // Newest first, matching the all-sessions logs overview. MCP calls + // stay grouped last (above), newest-first within that group too. + return new Date(b.startTime).getTime() - new Date(a.startTime).getTime(); + }); +} + /** * Formats data for display. If input is a string, attempts to parse as JSON. * @param input - Data to format (string or object) From bcd52754dead402827ce9e080d8a64ebe622c219 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 8 Jul 2026 09:43:47 +0300 Subject: [PATCH 12/31] feat(rate_limit): support per-tag rpm limiting on a single key (#31502) Add a tag_rpm_limit field to virtual keys so each request tag gets its own independent RPM counter on the v3 rate limiter. A key configured with per-tag limits tracks each tag/group separately, and requests whose tag has no configured limit fall back to the key-level limit. Includes the dashboard UI to manage per-tag limits on key create and edit. Resolves LIT-3147 --- litellm/proxy/_types.py | 2 + litellm/proxy/auth/auth_utils.py | 14 ++ .../hooks/parallel_request_limiter_v3.py | 51 ++++++- .../internal_user_endpoints.py | 2 + .../key_management_endpoints.py | 6 + .../proxy/auth/test_auth_utils.py | 15 ++ .../hooks/test_parallel_request_limiter_v3.py | 133 ++++++++++++++++++ .../test_key_management_endpoints.py | 23 +++ .../key_team_helpers/TagRateLimitEditor.tsx | 103 ++++++++++++++ .../organisms/create_key_button.tsx | 24 ++++ .../components/templates/key_edit_view.tsx | 27 ++++ .../components/templates/key_info_view.tsx | 7 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 36 +++++ 13 files changed, 442 insertions(+), 1 deletion(-) create mode 100644 ui/litellm-dashboard/src/components/key_team_helpers/TagRateLimitEditor.tsx diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9e390312fa1..b6bef568637 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1045,6 +1045,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): model_rpm_limit: Optional[dict] = None model_tpm_limit: Optional[dict] = None mcp_rpm_limit: Optional[Dict[str, int]] = None + tag_rpm_limit: Optional[dict[str, int]] = None guardrails: Optional[List[str]] = None policies: Optional[List[str]] = None prompts: Optional[List[str]] = None @@ -3869,6 +3870,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [ "model_rpm_limit", "model_tpm_limit", "mcp_rpm_limit", + "tag_rpm_limit", "rpm_limit_type", "tpm_limit_type", "enforced_params", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 3c508df0cc9..893e09ece6e 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -975,6 +975,20 @@ def get_team_mcp_rpm_limit( return None +def get_key_tag_rpm_limit( + user_api_key_dict: UserAPIKeyAuth, +) -> Optional[dict[str, int]]: + """ + Get the per-request-tag rpm limit configured on a given api key. + + The returned dict is keyed by request tag, so each tag/group tracked on + the key gets its own independent RPM counter. + """ + if user_api_key_dict.metadata: + return user_api_key_dict.metadata.get("tag_rpm_limit") + return None + + def get_project_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> Optional[Dict[str, int]]: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index ee0a0e1789d..7aedb74f2ea 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -31,8 +31,12 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_str_from_messages, ) from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata +from litellm.proxy.auth.auth_utils import ( + get_key_tag_rpm_limit, + get_model_rate_limit_from_metadata, +) from litellm.proxy.auth.budget_throttle import throttled_limit +from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, @@ -1300,6 +1304,43 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) ) + def _add_tag_per_key_rate_limit_descriptor( + self, + user_api_key_dict: UserAPIKeyAuth, + data: dict, + descriptors: list[RateLimitDescriptor], + ) -> None: + """ + Add per-request-tag rpm limit descriptors for the API key. + + Each tag carried on the request that has a configured limit gets its own + ``{api_key}:{tag}`` counter, so a burst on one tag/group never consumes + another's budget. Tags without a configured limit fall through to the + key-level descriptor. + """ + if not user_api_key_dict.api_key: + return + + tag_rpm_limit = get_key_tag_rpm_limit(user_api_key_dict) or {} + if not tag_rpm_limit: + return + + for tag in dict.fromkeys(get_tags_from_request_body(data)): + rpm_limit = tag_rpm_limit.get(tag) + if rpm_limit is None: + continue + descriptors.append( + RateLimitDescriptor( + key="tag_per_key", + value=f"{user_api_key_dict.api_key}:{tag}", + rate_limit={ + "requests_per_unit": rpm_limit, + "tokens_per_unit": None, + "window_size": self.window_size, + }, + ) + ) + def _add_mcp_per_key_rate_limit_descriptor( self, user_api_key_dict: UserAPIKeyAuth, @@ -1645,6 +1686,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptors=descriptors, ) + # Per-request-tag rate limits scoped to this key + self._add_tag_per_key_rate_limit_descriptor( + user_api_key_dict=user_api_key_dict, + data=data, + descriptors=descriptors, + ) + # REST MCP calls pass the raw body through this hook before server # resolution; only the later synthetic hook payload may carry this key. if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data: @@ -1961,6 +2009,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Org Level Rate Limits descriptors.extend(self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model)) + # Only check rate limits if we have descriptors with actual limits if descriptors: # First pass: RPM and max_parallel_requests sliding-window check. diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 989fad7cd0b..ccd15a68437 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -377,6 +377,7 @@ async def new_user( - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. + - tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user. - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). - agent_id: Optional[str] - The agent id associated with the user. @@ -1379,6 +1380,7 @@ async def user_update( - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. + - tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user. - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). - agent_id: Optional[str] - The agent id associated with the user. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 71cf2db3dfb..63f4b731871 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1496,6 +1496,7 @@ async def generate_key_fn( - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. + - tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit. - tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". - rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request @@ -2514,6 +2515,7 @@ async def update_key_fn( - rpm_limit: Optional[int] - Requests per minute limit - model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200} - mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200} + - tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit. - model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000} - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" @@ -3551,6 +3553,7 @@ async def generate_key_helper_fn( model_rpm_limit: Optional[dict] = None, model_tpm_limit: Optional[dict] = None, mcp_rpm_limit: Optional[dict] = None, + tag_rpm_limit: Optional[dict] = None, guardrails: Optional[list] = None, policies: Optional[list] = None, prompts: Optional[list] = None, @@ -3624,6 +3627,9 @@ async def generate_key_helper_fn( if mcp_rpm_limit is not None: metadata = metadata or {} metadata["mcp_rpm_limit"] = mcp_rpm_limit + if tag_rpm_limit is not None: + metadata = metadata or {} + metadata["tag_rpm_limit"] = tag_rpm_limit if guardrails is not None: metadata = metadata or {} metadata["guardrails"] = guardrails diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 21a001d77d7..042fc107f40 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -19,6 +19,7 @@ from litellm.proxy.auth.auth_utils import ( get_key_mcp_rpm_limit, get_key_model_rpm_limit, get_key_model_tpm_limit, + get_key_tag_rpm_limit, get_model_from_request, get_project_model_rpm_limit, get_project_model_tpm_limit, @@ -2393,3 +2394,17 @@ class TestIsRequestBodySafeBlocksModelList: ) is True ) + + +class TestGetKeyTagRateLimits: + """Tests for get_key_tag_rpm_limit.""" + + def test_reads_tag_rpm_limit_from_metadata(self): + key = UserAPIKeyAuth( + api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}} + ) + assert get_key_tag_rpm_limit(key) == {"cell-1": 5} + + def test_returns_none_when_unset(self): + key = UserAPIKeyAuth(api_key="sk-123") + assert get_key_tag_rpm_limit(key) is None diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 12f0a64a179..d150591c8de 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -3573,3 +3573,136 @@ async def test_pre_call_hook_skips_reservation_when_disabled(monkeypatch): ) assert TPM_RESERVED_TOKENS_KEY not in (data.get("metadata") or {}) + + +@pytest.mark.asyncio +async def test_per_tag_rate_limit_independent_counters_v3(monkeypatch): + """ + A single key with per-tag RPM limits tracks each tag independently: a tag + at its limit returns 429 while a different (unlimited) tag keeps flowing, + governed only by the generous key-level limit. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-per-tag-rpm") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + rpm_limit=100, + metadata={"tag_rpm_limit": {"cell-1": 2}}, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + async def call(tag: str) -> None: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-3.5-turbo", "metadata": {"tags": [tag]}}, + call_type="", + ) + + await call("cell-1") + await call("cell-1") + with pytest.raises(HTTPException) as exc_info: + await call("cell-1") + assert exc_info.value.status_code == 429 + assert "tag_per_key" in str(exc_info.value.detail) + + # cell-2 has no configured tag limit, so cell-1's exhausted counter must + # not block it; only the generous key-level limit applies. + for _ in range(5): + await call("cell-2") + + +@pytest.mark.asyncio +async def test_per_tag_descriptor_creation_v3(): + """ + _create_rate_limit_descriptors emits a tag_per_key descriptor carrying the + configured RPM limit only for request tags present in the configured map. + """ + _api_key = hash_token("sk-per-tag-desc") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + metadata={"tag_rpm_limit": {"cell-1": 5}}, + ) + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + + descriptors = handler._create_rate_limit_descriptors( + user_api_key_dict=user_api_key_dict, + data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1", "cell-2"]}}, + rpm_limit_type=None, + tpm_limit_type=None, + model_has_failures=False, + ) + + tag_descriptors = [d for d in descriptors if d["key"] == "tag_per_key"] + assert len(tag_descriptors) == 1, "only the configured tag yields a descriptor" + descriptor = tag_descriptors[0] + assert descriptor["value"] == f"{_api_key}:cell-1" + assert descriptor["rate_limit"]["requests_per_unit"] == 5 + + +@pytest.mark.asyncio +async def test_per_tag_descriptor_absent_without_config_v3(): + """No tag_per_key descriptor is created when the key has no tag limits.""" + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-no-tag"), + rpm_limit=10, + ) + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + + descriptors = handler._create_rate_limit_descriptors( + user_api_key_dict=user_api_key_dict, + data={"model": "gpt-3.5-turbo", "metadata": {"tags": ["cell-1"]}}, + rpm_limit_type=None, + tpm_limit_type=None, + model_has_failures=False, + ) + + assert not [d for d in descriptors if d["key"] == "tag_per_key"] + + +@pytest.mark.asyncio +async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch): + """ + Per-tag limits are opt-in sub-limits under the key-level ceiling, not a + standalone enforcement boundary: a request that carries no tag (or a tag + without a configured limit) is not rejected by any tag counter, but it is + still bounded by the key-level rpm_limit. This pins the documented + untagged-fallback behavior so a future "fail closed on missing tag" change + would fail here instead of silently breaking it. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _api_key = hash_token("sk-untagged-fallback") + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + rpm_limit=3, + metadata={"tag_rpm_limit": {"cell-1": 2}}, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + async def call(metadata: dict) -> None: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-3.5-turbo", "metadata": metadata}, + call_type="", + ) + + # Untagged and unconfigured-tag requests share the key-level budget of 3 + # and never hit a tag_per_key counter. + await call({}) + await call({"tags": ["cell-99"]}) + await call({}) + with pytest.raises(HTTPException) as exc_info: + await call({"tags": ["cell-99"]}) + assert exc_info.value.status_code == 429 + assert "tag_per_key" not in str(exc_info.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 4fb3df52cf6..d707421aeb6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -15,8 +15,11 @@ from unittest.mock import AsyncMock, MagicMock, patch from fastapi import HTTPException +import inspect + from litellm.proxy._types import ( GenerateKeyRequest, + NewUserRequest, LiteLLM_BudgetTable, LiteLLM_OrganizationTable, LiteLLM_TeamTableCachedObj, @@ -14480,3 +14483,23 @@ async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_g assert int(exc.value.code) == 403 assert "permissions" in str(exc.value.message) assert "Enterprise" not in str(exc.value.message) + + +def test_generate_key_helper_fn_accepts_per_tag_rate_limits(): + """ + Regression: new_user / SSO sign-in forward NewUserRequest fields to + generate_key_helper_fn via `**data_json`. The per-tag limit field must be + an accepted kwarg, otherwise user creation 500s with + "generate_key_helper_fn() got an unexpected keyword argument 'tag_rpm_limit'". + """ + params = inspect.signature(generate_key_helper_fn).parameters + assert "tag_rpm_limit" in params + + # The field exists on the request model that new_user forwards via **data_json. + assert "tag_rpm_limit" in NewUserRequest.model_fields + + # Binding the per-tag kwarg must not raise an unexpected-keyword TypeError. + inspect.signature(generate_key_helper_fn).bind_partial( + request_type="user", + tag_rpm_limit={"cell-1": 5}, + ) diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/TagRateLimitEditor.tsx b/ui/litellm-dashboard/src/components/key_team_helpers/TagRateLimitEditor.tsx new file mode 100644 index 00000000000..ee022ee9a75 --- /dev/null +++ b/ui/litellm-dashboard/src/components/key_team_helpers/TagRateLimitEditor.tsx @@ -0,0 +1,103 @@ +import { Button, Input, InputNumber } from "antd"; +import React from "react"; + +export interface TagRateLimitEntry { + // Stable identity for React list keys so deleting a middle row doesn't shift + // the controlled inputs of the rows below it. + id: string; + tag: string; + rpm_limit: number | null; +} + +let nextRowId = 0; +const newRowId = (): string => `tag-row-${nextRowId++}`; + +export interface TagRateLimits { + tag_rpm_limit: Record; +} + +// Build the rpm limit map from editor rows. A tag only enters the map when its +// name is non-empty and the RPM cell holds a number. +export const tagRowsToLimits = (rows: TagRateLimitEntry[]): TagRateLimits => { + const tag_rpm_limit: Record = {}; + rows.forEach(({ tag, rpm_limit }) => { + const name = tag.trim(); + if (!name) return; + if (typeof rpm_limit === "number") tag_rpm_limit[name] = rpm_limit; + }); + return { tag_rpm_limit }; +}; + +// Coerce an untyped metadata value into a {tag: number} map, dropping anything +// that isn't a numeric entry. Key metadata is loosely typed, so validate here. +const toNumberMap = (raw: unknown): Record => { + if (!raw || typeof raw !== "object") return {}; + const out: Record = {}; + Object.entries(raw as Record).forEach(([tag, limit]) => { + if (typeof limit === "number") out[tag] = limit; + }); + return out; +}; + +// Reconstruct editor rows from the stored rpm map. +export const tagLimitsToRows = (tagRpmLimit?: unknown): TagRateLimitEntry[] => { + const rpm = toNumberMap(tagRpmLimit); + return Object.keys(rpm).map((tag) => ({ + id: newRowId(), + tag, + rpm_limit: rpm[tag], + })); +}; + +interface TagRateLimitEditorProps { + value: TagRateLimitEntry[]; + onChange: (v: TagRateLimitEntry[]) => void; +} + +export function TagRateLimitEditor({ value, onChange }: TagRateLimitEditorProps) { + const addRow = () => { + onChange([...value, { id: newRowId(), tag: "", rpm_limit: null }]); + }; + + const removeRow = (idx: number) => { + onChange(value.filter((_, i) => i !== idx)); + }; + + const updateRow = (idx: number, field: keyof TagRateLimitEntry, fieldValue: string | number | null) => { + onChange(value.map((row, i) => (i === idx ? { ...row, [field]: fieldValue } : row))); + }; + + return ( +
+ {value.map((row, idx) => ( +
+ updateRow(idx, "tag", e.target.value)} + placeholder="Tag (e.g. cell-1)" + style={{ width: 180 }} + /> + updateRow(idx, "rpm_limit", v ?? null)} + placeholder="RPM" + style={{ width: 120 }} + /> + +
+ ))} + +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 0f371b72efe..ef2ddab70ed 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -30,6 +30,7 @@ import ProjectDropdown from "../common_components/ProjectDropdown"; import { CreateUserButton } from "../CreateUserButton"; import { BudgetFallbacksEditor } from "../key_team_helpers/BudgetFallbacksEditor"; import { BudgetWindowEntry, BudgetWindowsEditor } from "../key_team_helpers/BudgetWindowsEditor"; +import { TagRateLimitEditor, TagRateLimitEntry, tagRowsToLimits } from "../key_team_helpers/TagRateLimitEditor"; import { excludeProxyWideSentinel, getModelDisplayName, @@ -202,6 +203,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp const [rotationInterval, setRotationInterval] = useState("30d"); const [routerSettings, setRouterSettings] = useState(null); const [budgetLimits, setBudgetLimits] = useState([]); + const [tagRateLimits, setTagRateLimits] = useState([]); const [budgetFallbacks, setBudgetFallbacks] = useState>({}); const [budgetFallbacksKey, setBudgetFallbacksKey] = useState(0); const [routerSettingsKey, setRouterSettingsKey] = useState(0); @@ -223,6 +225,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp setSelectedOrganizationId(null); setSelectedProjectId(null); setBudgetLimits([]); + setTagRateLimits([]); setBudgetFallbacks({}); setBudgetFallbacksKey((k) => k + 1); }; @@ -244,6 +247,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp setSelectedOrganizationId(null); setSelectedProjectId(null); setBudgetLimits([]); + setTagRateLimits([]); setBudgetFallbacks({}); setBudgetFallbacksKey((k) => k + 1); }; @@ -543,6 +547,12 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp formValues.budget_limits = validWindows; } + // Add per-tag rate limits (only when at least one row is configured) + const { tag_rpm_limit } = tagRowsToLimits(tagRateLimits); + if (Object.keys(tag_rpm_limit).length > 0) { + formValues.tag_rpm_limit = tag_rpm_limit; + } + if (Object.keys(budgetFallbacks).length > 0) { formValues.budget_fallbacks = budgetFallbacks; } @@ -567,6 +577,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp NotificationsManager.success("Virtual Key Created"); form.resetFields(); setBudgetLimits([]); + setTagRateLimits([]); setBudgetFallbacks({}); setBudgetFallbacksKey((k) => k + 1); localStorage.removeItem("userData" + userID); @@ -1177,6 +1188,19 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp form={form} showDetailedDescriptions={true} /> + + Per-Tag Rate Limits{" "} + + + + + } + > + + ( Array.isArray(keyData.budget_limits) ? keyData.budget_limits : [], ); + const [tagRateLimits, setTagRateLimits] = useState( + tagLimitsToRows(keyData.metadata?.tag_rpm_limit), + ); const [budgetFallbacks, setBudgetFallbacks] = useState>( keyData.budget_fallbacks && typeof keyData.budget_fallbacks === "object" ? keyData.budget_fallbacks : {}, ); @@ -311,6 +320,11 @@ export function KeyEditView({ values.budget_limits = []; } + // Always send the current per-tag limit map so removing every row + // clears the stored limits ({} overwrites the metadata field). + const { tag_rpm_limit } = tagRowsToLimits(tagRateLimits); + values.tag_rpm_limit = tag_rpm_limit; + const hadExistingFallbacks = keyData.budget_fallbacks != null && Object.keys(keyData.budget_fallbacks).length > 0; if (Object.keys(budgetFallbacks).length > 0) { values.budget_fallbacks = budgetFallbacks; @@ -553,6 +567,19 @@ export function KeyEditView({ + + Per-Tag Rate Limits{" "} + + + + + } + > + + + {accessToken && ( + + Tag RPM Limits:{" "} + {currentKeyData.metadata?.tag_rpm_limit && + Object.keys(currentKeyData.metadata.tag_rpm_limit).length > 0 + ? JSON.stringify(currentKeyData.metadata.tag_rpm_limit) + : "Unlimited"} +
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d9bba85bf4b..23ada336710 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -6509,6 +6509,7 @@ export interface paths { * - model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit. * - model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit. * - mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit. + * - tag_rpm_limit: Optional[dict] - key-specific per-request-tag rpm limit, keyed by request tag. Example - {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; requests whose tag is absent fall back to the key-level rpm limit. * - tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". * - rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput". * - allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request @@ -6896,6 +6897,7 @@ export interface paths { * - rpm_limit: Optional[int] - Requests per minute limit * - model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200} * - mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200} + * - tag_rpm_limit: Optional[dict] - Per-request-tag RPM limits, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Each tag gets an independent counter; absent tags fall back to the key-level rpm limit. * - model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000} * - tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" * - rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic" @@ -14623,6 +14625,7 @@ export interface paths { * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) * - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. + * - tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user. * - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) * - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). * - agent_id: Optional[str] - The agent id associated with the user. @@ -14704,6 +14707,7 @@ export interface paths { * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) * - mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user. + * - tag_rpm_limit: Optional[dict] - Per-request-tag rpm limit, keyed by request tag {"cell-1": 1000, "cell-2": 500}. Enforced for keys only; values set on a user are stored but not enforced per user. * - model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) * - spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo"). * - agent_id: Optional[str] - The agent id associated with the user. @@ -23705,6 +23709,10 @@ export interface components { * @default 0 */ spend: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Tags */ tags?: string[] | null; /** Team Id */ @@ -23847,6 +23855,10 @@ export interface components { * @default 0 */ spend: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Tags */ tags?: string[] | null; /** Team Id */ @@ -28032,6 +28044,10 @@ export interface components { spend: number | null; /** Sso User Id */ sso_user_id?: string | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Team Id */ team_id?: string | null; /** Teams */ @@ -28186,6 +28202,10 @@ export interface components { * @default 0 */ spend: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Tags */ tags?: string[] | null; /** Team Id */ @@ -29817,6 +29837,10 @@ export interface components { soft_budget?: number | null; /** Spend */ spend?: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Tags */ tags?: string[] | null; /** Team Id */ @@ -31762,6 +31786,10 @@ export interface components { rpm_limit_type?: ("guaranteed_throughput" | "best_effort_throughput" | "dynamic") | null; /** Spend */ spend?: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Tags */ tags?: string[] | null; /** Team Id */ @@ -32218,6 +32246,10 @@ export interface components { rpm_limit?: number | null; /** Spend */ spend?: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Team Id */ team_id?: string | null; /** Tpm Limit */ @@ -32320,6 +32352,10 @@ export interface components { rpm_limit?: number | null; /** Spend */ spend?: number | null; + /** Tag Rpm Limit */ + tag_rpm_limit?: { + [key: string]: number; + } | null; /** Team Id */ team_id?: string | null; /** Tpm Limit */ From df2d44bab1e2c2ffd3acf68bbc0abf6ab1160f85 Mon Sep 17 00:00:00 2001 From: Thibault Serot Date: Wed, 8 Jul 2026 16:54:39 +1000 Subject: [PATCH 13/31] feat(ui): sort session sidebar by duration or start time --- .../LogDetailsDrawer.test.tsx | 12 +++---- .../LogDetailsDrawer/LogDetailsDrawer.tsx | 30 ++++++++-------- .../view_logs/LogDetailsDrawer/utils.test.ts | 36 +++++++++++++------ .../view_logs/LogDetailsDrawer/utils.ts | 23 +++++------- 4 files changed, 55 insertions(+), 46 deletions(-) diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx index db0d168fcac..1d23fecb5da 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx @@ -53,13 +53,13 @@ const sessionLogs = [ model: "tool-early", call_type: "call_mcp_tool", startTime: "2026-07-08T10:00:01.000Z", - endTime: "2026-07-08T10:00:01.500Z", + endTime: "2026-07-08T10:00:06.000Z", }), makeLog({ request_id: "llm-late", model: "llm-late", startTime: "2026-07-08T10:00:02.000Z", - endTime: "2026-07-08T10:00:04.000Z", + endTime: "2026-07-08T10:00:05.000Z", }), makeLog({ request_id: "mcp-late", @@ -84,10 +84,10 @@ const sidebarEventNames = () => screen.queryAllByText(/^(llm-early|llm-late|tool-early|tool-late)$/).map((el) => el.textContent); describe("LogDetailsDrawer session sidebar sorting", () => { - it("defaults to grouped order: LLM calls newest first, MCP calls grouped last", async () => { + it("defaults to duration order, longest call first across LLM and MCP calls", async () => { renderSessionDrawer(); await waitFor(() => expect(sidebarEventNames()).toHaveLength(4)); - expect(sidebarEventNames()).toEqual(["llm-late", "llm-early", "tool-late", "tool-early"]); + expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"]); }); it("switches to chronological order across LLM and MCP calls when Start time is selected", async () => { @@ -98,8 +98,8 @@ describe("LogDetailsDrawer session sidebar sorting", () => { await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-early", "tool-early", "llm-late", "tool-late"])); - fireEvent.click(screen.getByText("Grouped")); + fireEvent.click(screen.getByText("Duration")); - await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-late", "llm-early", "tool-late", "tool-early"])); + await waitFor(() => expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"])); }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index 36299139a13..79592216942 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -117,7 +117,7 @@ export function LogDetailsDrawer({ }: LogDetailsDrawerProps) { const isSessionMode = Boolean(sessionId); const [selectedSessionRequestId, setSelectedSessionRequestId] = useState(null); - const [sessionSortMode, setSessionSortMode] = useState("grouped"); + const [sessionSortMode, setSessionSortMode] = useState("duration"); const [isSidebarCollapsed, setIsSidebarCollapsed] = useState(false); const [copiedLeftPanelId, setCopiedLeftPanelId] = useState(false); @@ -173,9 +173,9 @@ export function LogDetailsDrawer({ const sessionTruncated = sessionTotalCount > sessionLogs.length; // Default selection for a freshly opened session: the most recent log (latest - // startTime). The list is sorted newest-first, but MCP calls are grouped last, - // so the latest log by time is not necessarily sessionLogs[0]; compute it - // explicitly. A clicked/remembered log still wins over this default. + // startTime). The list is ordered by the selected sort mode, so the latest + // log by time is not necessarily sessionLogs[0]; compute it explicitly. + // A clicked/remembered log still wins over this default. const mostRecentLog = useMemo( () => sessionLogs.reduce( @@ -387,16 +387,18 @@ export function LogDetailsDrawer({
)} {isSessionMode && ( - setSessionSortMode(value as SessionLogSortMode)} - /> +
+ Sort by + setSessionSortMode(value as SessionLogSortMode)} + /> +
)} diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts index cbe12f5c101..f59a20529d0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.test.ts @@ -1,30 +1,44 @@ import { describe, expect, it } from "vitest"; import { sortSessionLogs } from "./utils"; -const llm = (id: string, startTime: string) => ({ request_id: id, call_type: "acompletion", startTime }); -const mcp = (id: string, startTime: string) => ({ request_id: id, call_type: "call_mcp_tool", startTime }); +const log = (id: string, startTime: string, endTime: string, request_duration_ms?: number) => ({ + request_id: id, + startTime, + endTime, + request_duration_ms, +}); const ids = (rows: { request_id: string }[]) => rows.map((row) => row.request_id); describe("sortSessionLogs", () => { const rows = [ - mcp("mcp-early", "2026-07-08T10:00:01.000Z"), - llm("llm-late", "2026-07-08T10:00:02.000Z"), - mcp("mcp-late", "2026-07-08T10:00:03.000Z"), - llm("llm-early", "2026-07-08T10:00:00.000Z"), + log("mid-duration", "2026-07-08T10:00:01.000Z", "2026-07-08T10:00:01.500Z", 2000), + log("longest", "2026-07-08T10:00:02.000Z", "2026-07-08T10:00:02.500Z", 5000), + log("shortest", "2026-07-08T10:00:03.000Z", "2026-07-08T10:00:03.500Z", 300), + log("earliest-no-duration-field", "2026-07-08T10:00:00.000Z", "2026-07-08T10:00:04.000Z"), ]; - it("grouped mode keeps MCP calls last, newest first within each group", () => { - expect(ids(sortSessionLogs(rows, "grouped"))).toEqual(["llm-late", "llm-early", "mcp-late", "mcp-early"]); + it("duration mode sorts longest call first, deriving duration from timestamps when the field is missing", () => { + expect(ids(sortSessionLogs(rows, "duration"))).toEqual([ + "longest", + "earliest-no-duration-field", + "mid-duration", + "shortest", + ]); }); - it("chronological mode interleaves all calls by start time, oldest first", () => { - expect(ids(sortSessionLogs(rows, "chronological"))).toEqual(["llm-early", "mcp-early", "llm-late", "mcp-late"]); + it("start_time mode sorts calls in the order they started", () => { + expect(ids(sortSessionLogs(rows, "start_time"))).toEqual([ + "earliest-no-duration-field", + "mid-duration", + "longest", + "shortest", + ]); }); it("does not mutate the input array", () => { const input = [...rows]; - sortSessionLogs(input, "chronological"); + sortSessionLogs(input, "duration"); expect(ids(input)).toEqual(ids(rows)); }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts index 5a1a0e81f96..d313d07361c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/utils.ts @@ -3,25 +3,18 @@ * These functions handle data formatting, validation, and guardrail calculations. */ -import { MCP_CALL_TYPES } from "../constants"; +export type SessionLogSortMode = "duration" | "start_time"; -export type SessionLogSortMode = "grouped" | "chronological"; +type SortableSessionLog = { startTime: string; endTime: string; request_duration_ms?: number }; -export function sortSessionLogs( - rows: T[], - mode: SessionLogSortMode, -): T[] { - if (mode === "chronological") { +const durationMs = (row: SortableSessionLog): number => + row.request_duration_ms ?? Date.parse(row.endTime) - Date.parse(row.startTime); + +export function sortSessionLogs(rows: T[], mode: SessionLogSortMode): T[] { + if (mode === "start_time") { return [...rows].sort((a, b) => new Date(a.startTime).getTime() - new Date(b.startTime).getTime()); } - return [...rows].sort((a, b) => { - const aIsMcp = MCP_CALL_TYPES.includes(a.call_type) ? 1 : 0; - const bIsMcp = MCP_CALL_TYPES.includes(b.call_type) ? 1 : 0; - if (aIsMcp !== bIsMcp) return aIsMcp - bIsMcp; - // Newest first, matching the all-sessions logs overview. MCP calls - // stay grouped last (above), newest-first within that group too. - return new Date(b.startTime).getTime() - new Date(a.startTime).getTime(); - }); + return [...rows].sort((a, b) => durationMs(b) - durationMs(a)); } /** From 1fb2b4aef47493d836304bb74856ee2efea16718 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Tue, 7 Jul 2026 23:59:35 -0700 Subject: [PATCH 14/31] fix(mcp): drop the cached per-user OAuth token when the credential row changes (#32302) * fix(mcp): drop the cached per-user OAuth token when the credential row changes The v2 authorization_code chain Cached(Refreshing(V2PerUserTokenStore)) caches a positive token until its expires_at (or 300s without one), and CachedOAuthTokenStore.invalidate had no callers, so a re-authorization or revocation wrote the DB while egress kept serving the replaced token from the in-process cache until its TTL. LazyPerUserOAuthTokenStore now exposes invalidate, MCPServerManager threads it to the write side, and the three credential write sites (the OAuth callback, the Tools-tab persist endpoint, and the revoke endpoint) drop the cache entry after the row changes. The v2 refresher's own persist stays untouched; RefreshingTokenStore already feeds the rotated token back into the cache in the same fetch * test(mcp): pin cache invalidation on the revoke already-gone branch Greptile's review flagged that only the happy-path delete asserted the invalidate; a refactor moving the call inside the try block would silently skip the cache drop when the row was already deleted by a concurrent request while the cache still held the revoked token. The new test fails on exactly that mutation * test(mcp): cover invalidate on the redis-backed lazy store path Codecov flagged the redis fast path of LazyPerUserOAuthTokenStore.invalidate as unexercised; the existing invalidate tests only ran the no-redis chain. The new test builds the redis chain via a fetch and asserts a subsequent invalidate reaches the same store instance without a rebuild --- .../mcp_server/discoverable_endpoints.py | 6 + .../mcp_server/mcp_server_manager.py | 27 +++- .../outbound_credentials/oauth_token_store.py | 11 ++ .../per_user_oauth_store.py | 27 +++- .../mcp_management_endpoints.py | 10 ++ .../test_per_user_oauth_store.py | 96 +++++++++++- .../mcp_server/test_discoverable_endpoints.py | 101 +++++++++++++ .../mcp_server/test_mcp_server_manager.py | 33 ++++ .../test_mcp_management_endpoints.py | 142 ++++++++++++++++++ 9 files changed, 443 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 89d645b6f8a..fa1f73cea77 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -448,6 +448,12 @@ async def _store_per_user_token_server_side( ) return # Don't warm Redis if DB write failed + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + global_mcp_server_manager, + ) + + await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server.server_id) + # Warm the Redis cache so the first subsequent MCP call is a cache hit ttl = _compute_per_user_token_ttl(server, expires_in) await mcp_per_user_token_cache.set( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 39da1ba4a97..8e4b3c57bbb 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -70,6 +70,9 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import to_server_spec, to_subject, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + InvalidatableOAuthTokenStore, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import ( LazyPerUserOAuthTokenStore, ) @@ -689,9 +692,16 @@ class MCPServerManager: """ return auth_type == MCPAuth.oauth2_token_exchange and not (token_exchange_endpoint or token_url) - def __init__(self, cred_provider: Optional[UpstreamCredentialProvider] = None): + def __init__( + self, + cred_provider: Optional[UpstreamCredentialProvider] = None, + per_user_oauth_token_store: Optional[InvalidatableOAuthTokenStore] = None, + ): + self._per_user_oauth_token_store = per_user_oauth_token_store or LazyPerUserOAuthTokenStore( + self.get_mcp_server_by_id + ) self._cred_provider = cred_provider or UpstreamCredentialProvider( - oauth_token_store=LazyPerUserOAuthTokenStore(self.get_mcp_server_by_id), + oauth_token_store=self._per_user_oauth_token_store, token_exchanger=build_token_exchanger(), ) self.registry: dict[str, MCPServer] = {} @@ -3922,6 +3932,19 @@ class MCPServerManager: return False return await self._cred_provider.has_user_token(to_subject(user_api_key_auth, None), spec) + async def invalidate_user_oauth_token_cache(self, user_id: str, server_id: str) -> None: + """Drop the v2 chain's cached token for ``(user_id, server_id)`` after the credential row + changes (re-auth, revoke), so the next resolve reads the new row instead of serving the + replaced token until its cache TTL. Best-effort: a cache-drop failure is logged, never + raised, because the DB write already succeeded and the TTL remains the backstop. + """ + try: + await self._per_user_oauth_token_store.invalidate(user_id, server_id) + except Exception as exc: # noqa: BLE001 - cache drop is best-effort; TTL is the backstop + verbose_logger.warning( + "Failed to invalidate cached MCP OAuth token for user=%s server=%s: %s", user_id, server_id, exc + ) + async def _resolve_oauth2_headers_for_tool_call( self, mcp_server: MCPServer, diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py index fd2cb2f3e06..c1c70cf9050 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/oauth_token_store.py @@ -69,6 +69,17 @@ class OAuthTokenStore(Protocol): async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: ... +class InvalidatableOAuthTokenStore(OAuthTokenStore, Protocol): + """An ``OAuthTokenStore`` whose cached entry for a ``(user, server)`` pair can be dropped. + + The write side calls ``invalidate`` after a (re)authorization or revocation changes the + credential row, so reads stop serving the replaced token immediately instead of until its + cache TTL. ``CachedOAuthTokenStore`` (the top of the per-user chain) satisfies this. + """ + + async def invalidate(self, user_id: str, server_id: str) -> None: ... + + class TokenRefresher(Protocol): """Mints a fresh token from an expired one and persists it, returning the new token. diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py index 3bc10f1a0eb..21001c09f25 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/per_user_oauth_store.py @@ -24,8 +24,8 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.dual_cache_toke ) from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( CachedOAuthTokenStore, + InvalidatableOAuthTokenStore, OAuthToken, - OAuthTokenStore, RefreshCoordinator, RefreshingTokenStore, TokenCacheBackend, @@ -51,7 +51,7 @@ if TYPE_CHECKING: _DEFAULT_TTL_SECONDS = 300.0 ServerLookup = Callable[[str], "MCPServer | None"] -StoreBuilder = Callable[[ServerLookup], tuple[OAuthTokenStore, bool]] +StoreBuilder = Callable[[ServerLookup], tuple[InvalidatableOAuthTokenStore, bool]] async def _read_credential(user_id: str, server_id: str) -> dict[str, object] | None: @@ -185,7 +185,7 @@ class LazyPerUserOAuthTokenStore: self._server_lookup = server_lookup self._store_builder = store_builder self._redis_available = redis_available - self._store: OAuthTokenStore | None = None + self._store: InvalidatableOAuthTokenStore | None = None self._uses_redis = False self._fetch_lock = asyncio.Condition() self._local_fetches = 0 @@ -203,7 +203,26 @@ class LazyPerUserOAuthTokenStore: if not uses_redis: await self._finish_local_fetch() - async def _store_for_fetch(self) -> tuple[OAuthTokenStore, bool]: + async def invalidate(self, user_id: str, server_id: str) -> None: + """Drop the chain's cached entry for ``(user_id, server_id)`` after the credential row + changes (re-auth, revoke). Builds the chain if no fetch has run yet, so a shared (Redis) + cache entry written by another worker is dropped too; the in-process case is then a no-op + on an empty cache. + """ + if self._uses_redis: + store = self._store + if store is not None: + await store.invalidate(user_id, server_id) + return + + store, uses_redis = await self._store_for_fetch() + try: + await store.invalidate(user_id, server_id) + finally: + if not uses_redis: + await self._finish_local_fetch() + + async def _store_for_fetch(self) -> tuple[InvalidatableOAuthTokenStore, bool]: async with self._fetch_lock: while ( self._store is not None and not self._uses_redis and self._redis_available() and self._local_fetches > 0 diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index a66e2fcf618..c9952b245c7 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1913,6 +1913,11 @@ if MCP_AVAILABLE: expires_in=payload.expires_in, scopes=payload.scopes, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + global_mcp_server_manager, + ) + + await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id) # Read back the persisted record so the response reflects the stored # expires_at rather than recomputing it here (which could diverge by # milliseconds or if the storage logic ever adds a grace period). @@ -1953,6 +1958,11 @@ if MCP_AVAILABLE: await delete_user_credential(prisma_client, user_id, server_id) except RecordNotFoundError: pass # Already gone — treat as a successful delete + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 + global_mcp_server_manager, + ) + + await global_mcp_server_manager.invalidate_user_oauth_token_cache(user_id, server_id) return MCPOAuthUserCredentialStatus( server_id=server_id, has_credential=False, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py index f64ce594efa..ca32cf2bb8d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_per_user_oauth_store.py @@ -3,8 +3,8 @@ import asyncio import pytest from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + InvalidatableOAuthTokenStore, OAuthToken, - OAuthTokenStore, ) from litellm.proxy._experimental.mcp_server.outbound_credentials.per_user_oauth_store import ( LazyPerUserOAuthTokenStore, @@ -16,11 +16,15 @@ class _RecordingStore: def __init__(self, access_token: str) -> None: self._access_token = access_token self.calls: list[tuple[str, str]] = [] + self.invalidations: list[tuple[str, str]] = [] async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: self.calls.append((user_id, server_id)) return OAuthToken(access_token=self._access_token) + async def invalidate(self, user_id: str, server_id: str) -> None: + self.invalidations.append((user_id, server_id)) + class _BlockingStore: def __init__(self, access_token: str) -> None: @@ -28,6 +32,7 @@ class _BlockingStore: self.started = asyncio.Event() self.release = asyncio.Event() self.calls: list[tuple[str, str]] = [] + self.invalidations: list[tuple[str, str]] = [] async def fetch(self, user_id: str, server_id: str) -> OAuthToken | None: self.calls.append((user_id, server_id)) @@ -35,6 +40,9 @@ class _BlockingStore: await self.release.wait() return OAuthToken(access_token=self._access_token) + async def invalidate(self, user_id: str, server_id: str) -> None: + self.invalidations.append((user_id, server_id)) + class _RedisAvailability: def __init__(self) -> None: @@ -59,7 +67,7 @@ async def test_lazy_store_rebuilds_when_redis_becomes_available() -> None: redis_available = _RedisAvailability() build_calls = 0 - def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]: + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: nonlocal build_calls build_calls += 1 if redis_available.available: @@ -94,7 +102,7 @@ async def test_lazy_store_allows_concurrent_local_fetches_without_redis() -> Non redis_available = _RedisAvailability() build_calls = 0 - def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]: + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: nonlocal build_calls build_calls += 1 return local_store, False @@ -127,7 +135,7 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild() redis_store = _RecordingStore("redis") redis_available = _RedisAvailability() - def build_store(_server_lookup: ServerLookup) -> tuple[OAuthTokenStore, bool]: + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: if redis_available.available: return redis_store, True return local_store, False @@ -158,3 +166,83 @@ async def test_lazy_store_waits_for_in_flight_local_fetch_before_redis_rebuild() assert second is not None and second.access_token == "redis" assert local_store.calls == [("u", "s")] assert redis_store.calls == [("u", "s")] + + +@pytest.mark.asyncio +async def test_lazy_store_invalidate_builds_chain_and_delegates() -> None: + local_store = _RecordingStore("local") + build_calls = 0 + + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: + nonlocal build_calls + build_calls += 1 + return local_store, False + + def server_lookup(_server_id: str) -> None: + return None + + store = LazyPerUserOAuthTokenStore( + server_lookup, + store_builder=build_store, + redis_available=_RedisAvailability(), + ) + + await store.invalidate("u", "s") + + assert build_calls == 1 + assert local_store.invalidations == [("u", "s")] + + +@pytest.mark.asyncio +async def test_lazy_store_invalidate_reaches_the_store_fetch_reads() -> None: + local_store = _RecordingStore("local") + build_calls = 0 + + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: + nonlocal build_calls + build_calls += 1 + return local_store, False + + def server_lookup(_server_id: str) -> None: + return None + + store = LazyPerUserOAuthTokenStore( + server_lookup, + store_builder=build_store, + redis_available=_RedisAvailability(), + ) + + await store.fetch("u", "s") + await store.invalidate("u", "s") + + assert build_calls == 1 + assert local_store.calls == [("u", "s")] + assert local_store.invalidations == [("u", "s")] + + +@pytest.mark.asyncio +async def test_lazy_store_invalidate_works_after_redis_chain_is_built() -> None: + redis_store = _RecordingStore("redis") + redis_available = _RedisAvailability() + redis_available.available = True + build_calls = 0 + + def build_store(_server_lookup: ServerLookup) -> tuple[InvalidatableOAuthTokenStore, bool]: + nonlocal build_calls + build_calls += 1 + return redis_store, True + + def server_lookup(_server_id: str) -> None: + return None + + store = LazyPerUserOAuthTokenStore( + server_lookup, + store_builder=build_store, + redis_available=redis_available, + ) + + await store.fetch("u", "s") + await store.invalidate("u", "s") + + assert build_calls == 1 + assert redis_store.invalidations == [("u", "s")] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index fe90fd45856..c808b17678a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4055,3 +4055,104 @@ async def test_oauth_authorization_server_404_for_unknown_server_name(): mcp_server_name="does_not_exist", ) assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_store_per_user_token_server_side_invalidates_v2_token_cache(): + """A token stored by the OAuth callback (code exchange or refresh) drops the v2 per-user + token cache entry, so egress stops serving the replaced token immediately instead of + until its TTL.""" + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _store_per_user_token_server_side, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv-cb-1", + name="cb_server", + url="https://upstream.example/mcp", + transport="http", + auth_type=MCPAuth.oauth2, + ) + invalidate_mock = AsyncMock(return_value=None) + cache_set_mock = AsyncMock(return_value=None) + + with ( + patch( + "litellm.proxy.utils.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set", + new=cache_set_mock, + ), + patch.object( + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + await _store_per_user_token_server_side( + server=server, + user_id="user-cb-1", + token_response={"access_token": "fresh-tok", "expires_in": 3600}, + ) + + invalidate_mock.assert_awaited_once_with("user-cb-1", "srv-cb-1") + cache_set_mock.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_store_per_user_token_server_side_skips_invalidate_when_db_write_fails(): + """A failed DB write neither warms the v1 cache nor drops the v2 cache entry; the + previously stored token is still the truth.""" + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _store_per_user_token_server_side, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv-cb-2", + name="cb_server_2", + url="https://upstream.example/mcp", + transport="http", + auth_type=MCPAuth.oauth2, + ) + invalidate_mock = AsyncMock(return_value=None) + cache_set_mock = AsyncMock(return_value=None) + + with ( + patch( + "litellm.proxy.utils.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.db.store_user_oauth_credential", + new=AsyncMock(side_effect=RuntimeError("db down")), + ), + patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_per_user_token_cache.set", + new=cache_set_mock, + ), + patch.object( + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + await _store_per_user_token_server_side( + server=server, + user_id="user-cb-2", + token_response={"access_token": "fresh-tok", "expires_in": 3600}, + ) + + invalidate_mock.assert_not_awaited() + cache_set_mock.assert_not_awaited() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 6fa69b2f96c..a423f7c8b83 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -2940,6 +2940,39 @@ class TestMCPServerManager: assert await manager.has_user_oauth_token(server, user_auth) is False assert calls == [] # short-circuited on the None spec, never hit the resolver + @pytest.mark.asyncio + async def test_invalidate_user_oauth_token_cache_delegates_to_store(self): + """The write side's cache drop reaches the same per-user store the resolver reads.""" + + class _Store: + def __init__(self) -> None: + self.invalidations: list[tuple[str, str]] = [] + + async def fetch(self, user_id: str, server_id: str): + return None + + async def invalidate(self, user_id: str, server_id: str) -> None: + self.invalidations.append((user_id, server_id)) + + store = _Store() + manager = MCPServerManager(per_user_oauth_token_store=store) + await manager.invalidate_user_oauth_token_cache("alice", "srv-1") + assert store.invalidations == [("alice", "srv-1")] + + @pytest.mark.asyncio + async def test_invalidate_user_oauth_token_cache_swallows_store_errors(self): + """A cache-drop failure must not fail the credential write that triggered it.""" + + class _Store: + async def fetch(self, user_id: str, server_id: str): + return None + + async def invalidate(self, user_id: str, server_id: str) -> None: + raise RuntimeError("redis down") + + manager = MCPServerManager(per_user_oauth_token_store=_Store()) + await manager.invalidate_user_oauth_token_cache("alice", "srv-1") + @pytest.mark.asyncio async def test_resolve_oauth2_headers_no_user_id(self): """Skip lookup entirely when user_api_key_auth has no user_id.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 6976aa76a94..86bbce36de3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -3566,6 +3566,148 @@ async def test_delete_mcp_oauth_user_credential_only_deletes_oauth(): assert result.has_credential is False +@pytest.mark.asyncio +async def test_store_mcp_oauth_user_credential_invalidates_cached_token(): + """Re-authorizing via the Tools-tab persist drops the v2 per-user token cache entry, so + egress stops serving the replaced token immediately instead of until its TTL.""" + from litellm.proxy._types import MCPOAuthUserCredentialRequest + + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + store_mcp_oauth_user_credential, + ) + + server_id = "srv-inv-1" + user_id = "user-inv-1" + invalidate_mock = AsyncMock(return_value=None) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_server", + new=AsyncMock(return_value=generate_mock_mcp_server_db_record(server_id=server_id)), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._user_has_admin_view", + return_value=True, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.store_user_oauth_credential", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential", + new=AsyncMock(return_value={"type": "oauth2", "access_token": "new-tok"}), + ), + patch.object( + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + await store_mcp_oauth_user_credential( + server_id=server_id, + payload=MCPOAuthUserCredentialRequest(access_token="new-tok", expires_in=3600), + user_api_key_dict=_make_user_auth(user_id), + ) + + invalidate_mock.assert_awaited_once_with(user_id, server_id) + + +@pytest.mark.asyncio +async def test_delete_mcp_oauth_user_credential_invalidates_cached_token(): + """Revoking a stored OAuth credential drops the v2 per-user token cache entry, so the + revoked token stops flowing upstream immediately instead of until its TTL.""" + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_oauth_user_credential, + ) + + server_id = "srv-inv-2" + user_id = "user-inv-2" + invalidate_mock = AsyncMock(return_value=None) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential", + new=AsyncMock(return_value={"type": "oauth2", "access_token": "revoked-tok"}), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=AsyncMock(return_value=None), + ), + patch.object( + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + result = await delete_mcp_oauth_user_credential( + server_id=server_id, + user_api_key_dict=_make_user_auth(user_id), + ) + + invalidate_mock.assert_awaited_once_with(user_id, server_id) + assert result.has_credential is False + + +@pytest.mark.asyncio +async def test_delete_mcp_oauth_user_credential_invalidates_when_record_already_gone(): + """A concurrent delete can remove the row between the read and the delete; the cache may + still hold the revoked token, so the invalidate must fire even on RecordNotFoundError.""" + if not mgmt_endpoints.MCP_AVAILABLE: + pytest.skip("MCP module not installed") + + from litellm.proxy._experimental.mcp_server import mcp_server_manager as manager_module + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + delete_mcp_oauth_user_credential, + ) + + server_id = "srv-inv-3" + user_id = "user-inv-3" + invalidate_mock = AsyncMock(return_value=None) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=_make_prisma_client(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_user_oauth_credential", + new=AsyncMock(return_value={"type": "oauth2", "access_token": "revoked-tok"}), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.delete_user_credential", + new=AsyncMock(side_effect=mgmt_endpoints.RecordNotFoundError({}, message="already gone")), + ), + patch.object( + manager_module.global_mcp_server_manager, + "invalidate_user_oauth_token_cache", + new=invalidate_mock, + ), + ): + result = await delete_mcp_oauth_user_credential( + server_id=server_id, + user_api_key_dict=_make_user_auth(user_id), + ) + + invalidate_mock.assert_awaited_once_with(user_id, server_id) + assert result.has_credential is False + + @pytest.mark.asyncio async def test_list_mcp_user_credentials_batch_server_fetch(): """list_mcp_user_credentials uses a single batch DB call, not N+1 queries.""" From 6f6bd4568118ce7d15fa9b944554c18307f55a5f Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 8 Jul 2026 09:59:57 +0300 Subject: [PATCH 15/31] perf(auth): negative-cache missing user/key lookups on the request hot path (#32368) --- litellm/integrations/prometheus.py | 1 + litellm/proxy/auth/auth_checks.py | 8 +- litellm/proxy/management_endpoints/ui_sso.py | 16 +-- ...st_prometheus_budget_metrics_db_lookups.py | 93 ++++++++++++++++ .../test_auth_hot_path_network_requests.py | 101 ++++++++++++++++++ .../proxy/management_endpoints/test_ui_sso.py | 51 +++++++++ 6 files changed, 263 insertions(+), 7 deletions(-) create mode 100644 tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index e374068ca35..60fc021a6a8 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -3619,6 +3619,7 @@ class PrometheusLogger(CustomLogger): hashed_token=user_api_key_dict.token, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, + check_cache_only=True, ) if key_object: user_api_key_dict.budget_reset_at = key_object.budget_reset_at diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index ee548ba0a43..e7fee8d6eb2 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1474,7 +1474,7 @@ def _should_check_db(key: str, last_db_access_time: LimitedSizeOrderedDict, db_c elif last_db_access_time[key][0] is not None: # check db for non-null values (for refresh operations) return True elif last_db_access_time[key][0] is None: - if current_time - last_db_access_time[key] >= db_cache_expiry: + if current_time - last_db_access_time[key][1] >= db_cache_expiry: return True return False @@ -1649,6 +1649,12 @@ async def get_user_object( include={"organization_memberships": True}, ) else: + if should_check_db: + _update_last_db_access_time( + key=db_access_time_key, + value=None, + last_db_access_time=last_db_access_time, + ) raise Exception if response.organization_memberships is not None and len(response.organization_memberships) > 0: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 89c1a925eeb..dbf514d2298 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -2134,12 +2134,16 @@ async def cli_poll_key( models=session_data.get("models", []), ) - user_db_obj = await get_user_object( - user_id=user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - user_id_upsert=False, - ) + try: + user_db_obj = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + ) + except ValueError as e: + verbose_proxy_logger.debug(f"CLI poll: user lookup failed, proceeding without user budget: {e}") + user_db_obj = None user_budget = user_db_obj.max_budget if user_db_obj is not None else None team_budget: Optional[float] = None diff --git a/tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py b/tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py new file mode 100644 index 00000000000..ce446ae3a19 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_budget_metrics_db_lookups.py @@ -0,0 +1,93 @@ +""" +Unit tests for PrometheusLogger._assemble_key_object DB access. + +The post-request budget metrics run for every LLM API request. Auth has +already cached the key object for any real key in the same request, so the +metrics path must read the cache only. Falling through to the DB turns every +request whose token has no DB row (e.g. master-key requests, whose token is +an alias hash that never matches a stored key) into per-request +LiteLLM_VerificationToken and LiteLLM_DeprecatedVerificationToken queries. +""" + +import datetime +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from prometheus_client import REGISTRY + +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.integrations.prometheus import PrometheusLogger +from litellm.proxy._types import UserAPIKeyAuth + + +@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 + yield + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + try: + REGISTRY.unregister(collector) + except Exception: + pass + + +@pytest.fixture +def prometheus_logger(): + return PrometheusLogger() + + +@pytest.mark.asyncio +async def test_assemble_key_object_does_not_query_db_on_cache_miss(prometheus_logger): + mock_prisma = MagicMock() + mock_prisma.get_data = AsyncMock() + cache = DualCache(in_memory_cache=InMemoryCache()) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + ): + result = await prometheus_logger._assemble_key_object( + user_api_key="hashed-token-not-in-cache", + user_api_key_alias="", + key_max_budget=None, + key_spend=1.0, + response_cost=0.5, + ) + + mock_prisma.get_data.assert_not_called() + assert result.spend == 1.5 + assert result.budget_reset_at is None + + +@pytest.mark.asyncio +async def test_assemble_key_object_reads_budget_reset_at_from_cache(prometheus_logger): + hashed_token = "hashed-token-in-cache" + reset_at = datetime.datetime(2026, 8, 1, tzinfo=datetime.timezone.utc) + cached_key = UserAPIKeyAuth(token=hashed_token, budget_reset_at=reset_at) + + mock_prisma = MagicMock() + mock_prisma.get_data = AsyncMock() + cache = DualCache(in_memory_cache=InMemoryCache()) + await cache.async_set_cache(key=hashed_token, value=cached_key) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + ): + result = await prometheus_logger._assemble_key_object( + user_api_key=hashed_token, + user_api_key_alias="alias", + key_max_budget=10.0, + key_spend=1.0, + response_cost=0.5, + ) + + mock_prisma.get_data.assert_not_called() + assert result.budget_reset_at == reset_at diff --git a/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py b/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py index f1ff001777c..22752f767ce 100644 --- a/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py +++ b/tests/test_litellm/proxy/auth/test_auth_hot_path_network_requests.py @@ -516,3 +516,104 @@ async def test_full_hot_path_network_count(): assert ( summary["total_network_requests"] == 4 ), f"Expected 4 total network requests on warm path, got {summary['total_network_requests']}" + + +# ============================================================================ +# TEST: negative caching for entities that do not exist in the DB +# ============================================================================ + + +@pytest.mark.asyncio +async def test_get_user_object_missing_user_negative_cache(): + """ + A user_id with no DB row (e.g. the master key's default admin user_id) + must not trigger a DB query on every request. The first lookup hits the + DB; repeat lookups inside the db_cache_expiry window are throttled. + """ + user_id = "user-missing-negative-cache" + + cache = DualCache(in_memory_cache=InMemoryCache()) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_usertable = MagicMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + for _ in range(3): + with pytest.raises(ValueError): + await get_user_object( + user_id=user_id, + prisma_client=mock_prisma, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=None, + user_id_upsert=False, + ) + + assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1 + + +@pytest.mark.asyncio +async def test_get_user_object_missing_user_rechecks_after_expiry(): + """ + The negative cache must expire: a user created after a miss becomes + visible once the db_cache_expiry window has passed. + """ + from litellm.proxy.auth.auth_checks import db_cache_expiry, last_db_access_time + + user_id = "user-missing-expiry-recheck" + + cache = DualCache(in_memory_cache=InMemoryCache()) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_usertable = MagicMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + with pytest.raises(ValueError): + await get_user_object( + user_id=user_id, + prisma_client=mock_prisma, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=None, + user_id_upsert=False, + ) + assert mock_prisma.db.litellm_usertable.find_unique.call_count == 1 + + last_db_access_time[f"user_id:{user_id}"] = ( + None, + time.time() - (db_cache_expiry + 1), + ) + + with pytest.raises(ValueError): + await get_user_object( + user_id=user_id, + prisma_client=mock_prisma, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=None, + user_id_upsert=False, + ) + assert mock_prisma.db.litellm_usertable.find_unique.call_count == 2 + + +def test_should_check_db_negative_entry_throttles_then_expires(): + """ + A recorded miss (value=None) suppresses DB checks inside the expiry + window and allows them again after it. Exercises the timestamp element + of the stored (value, time) tuple directly. + """ + from litellm.caching.dual_cache import LimitedSizeOrderedDict + from litellm.proxy.auth.auth_checks import ( + _should_check_db, + _update_last_db_access_time, + ) + + tracker: LimitedSizeOrderedDict = LimitedSizeOrderedDict(max_size=10) + + _update_last_db_access_time(key="k", value=None, last_db_access_time=tracker) + assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is False + + tracker["k"] = (None, time.time() - 6) + assert _should_check_db(key="k", last_db_access_time=tracker, db_cache_expiry=5) is True diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 976048b9521..32229e3e64e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -7131,3 +7131,54 @@ async def test_legacy_login_page_hides_credentials_hint_via_general_settings(): assert response.status_code == 200 assert "Default Credentials" not in body assert "MASTER_KEY" not in body + + +@pytest.mark.asyncio +async def test_cli_poll_key_tolerates_missing_user_row(): + """The CLI poll must still mint the JWT when the user lookup raises, + e.g. the user row was created moments ago and a negative-cache window + from the pre-creation SSO existence check is still active on this pod.""" + from litellm.proxy.management_endpoints.ui_sso import ( + _hash_cli_sso_secret, + cli_poll_key, + ) + + session_key = "cli-session-missing-user" + session_data = { + "user_id": "just-created-user", + "user_role": "internal_user", + "teams": [], + "models": ["gpt-4"], + } + + mock_cache = MagicMock() + mock_cache.get_cache.return_value = { + "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), + "sso_complete": True, + "user_code_verified": True, + "session_data": session_data, + } + + mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.missing.user" + + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.prisma_client"), + patch( + "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token", + return_value=mock_jwt_token, + ), + patch( + "litellm.proxy.auth.auth_checks.get_user_object", + new=AsyncMock(side_effect=ValueError("User doesn't exist in db. 'user_id'=just-created-user")), + ), + ): + result = await cli_poll_key( + key_id=session_key, + team_id=None, + x_litellm_cli_poll_secret="poll-secret", + ) + + assert result["status"] == "ready" + assert result["key"] == mock_jwt_token + assert result["user_id"] == "just-created-user" From 5d89be551bbed49c493f2a1cea0bddf4ea8e468b Mon Sep 17 00:00:00 2001 From: Thibault Serot Date: Wed, 8 Jul 2026 17:06:38 +1000 Subject: [PATCH 16/31] fix(ui): fit session sort toggle inside sidebar column --- .../LogDetailsDrawer/LogDetailsDrawer.tsx | 23 +++++++++---------- 1 file changed, 11 insertions(+), 12 deletions(-) diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index 79592216942..92cf90fad63 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -387,18 +387,17 @@ export function LogDetailsDrawer({ )} {isSessionMode && ( -
- Sort by - setSessionSortMode(value as SessionLogSortMode)} - /> -
+ setSessionSortMode(value as SessionLogSortMode)} + /> )} From 6d2090a21b19d6277d44367983d9d70ee10e8a0c Mon Sep 17 00:00:00 2001 From: Thibault Serot Date: Wed, 8 Jul 2026 17:17:44 +1000 Subject: [PATCH 17/31] fix(ui): reset session sort mode when drawer closes --- .../LogDetailsDrawer.test.tsx | 21 ++++++++++++++++--- .../LogDetailsDrawer/LogDetailsDrawer.tsx | 1 + 2 files changed, 19 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx index 1d23fecb5da..5a49cccec70 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.test.tsx @@ -73,11 +73,13 @@ const sessionLogs = [ const renderSessionDrawer = () => { vi.mocked(sessionSpendLogsCall).mockResolvedValue({ data: sessionLogs, total: 4, total_pages: 1 }); const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); - render( + const drawer = (open: boolean) => ( - {}} logEntry={null} sessionId="session-1" accessToken="token" /> - , + {}} logEntry={null} sessionId="session-1" accessToken="token" /> + ); + const { rerender } = render(drawer(true)); + return { rerender, drawer }; }; const sidebarEventNames = () => @@ -102,4 +104,17 @@ describe("LogDetailsDrawer session sidebar sorting", () => { await waitFor(() => expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"])); }); + + it("resets the sort mode back to duration when the drawer is closed and reopened", async () => { + const { rerender, drawer } = renderSessionDrawer(); + await waitFor(() => expect(sidebarEventNames()).toHaveLength(4)); + + fireEvent.click(screen.getByText("Start time")); + await waitFor(() => expect(sidebarEventNames()).toEqual(["llm-early", "tool-early", "llm-late", "tool-late"])); + + rerender(drawer(false)); + rerender(drawer(true)); + + await waitFor(() => expect(sidebarEventNames()).toEqual(["tool-early", "llm-late", "llm-early", "tool-late"])); + }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx index 92cf90fad63..cc087a611b0 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx @@ -217,6 +217,7 @@ export function LogDetailsDrawer({ setIsSidebarCollapsed(false); } else { if (isSessionMode) setSelectedSessionRequestId(null); + setSessionSortMode("duration"); setCopiedLeftPanelId(false); } }, [open, isSessionMode]); From 684e3e1c2e98c8b43a6ecce0a0650842226db5e9 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 8 Jul 2026 00:17:55 -0700 Subject: [PATCH 18/31] test(vertex_ai): bump local_testing vertex tests from gemini-2.5-flash to gemini-3.5-flash (#32439) --- .../test_amazing_vertex_completion.py | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index 2382b8a5197..6e31166ad99 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -355,7 +355,11 @@ async def test_async_vertexai_response_basic(): user_message = "Hello, how are you?" messages = [{"content": user_message, "role": "user"}] response = await acompletion( - model="gemini-2.5-flash", messages=messages, temperature=0.7, timeout=5 + model="gemini-3.5-flash", + messages=messages, + temperature=0.7, + timeout=5, + vertex_location="global", ) print(f"response: {response}") except litellm.NotFoundError as e: @@ -388,7 +392,7 @@ async def test_async_vertexai_streaming_response(): ) test_models = random.sample(list(test_models), 1) test_models += list(litellm.vertex_language_models) # always test gemini-pro - test_models = ["gemini-2.5-flash"] + test_models = ["gemini-3.5-flash"] for model in test_models: if model in VERTEX_MODELS_TO_NOT_TEST or ( "gecko" in model @@ -412,6 +416,7 @@ async def test_async_vertexai_streaming_response(): temperature=0.7, timeout=5, stream=True, + vertex_location="global", ) print(f"response: {response}") complete_response: str = "" @@ -3840,10 +3845,11 @@ def test_vertex_schema_test(): } response = litellm.completion( - model="vertex_ai/gemini-2.5-flash", + model="vertex_ai/gemini-3.5-flash", messages=[{"role": "user", "content": "call the tool"}], tools=[tool], tool_choice="required", + vertex_location="global", ) print(response) @@ -3895,10 +3901,11 @@ def test_gemini_nullable_object_tool_schema_httpx(): ] response = litellm.completion( - model="vertex_ai/gemini-2.5-flash", + model="vertex_ai/gemini-3.5-flash", messages=[{"role": "user", "content": "call the tool"}], tools=tools, tool_choice="required", + vertex_location="global", ) print(response) From cd6e8cdf23186fad63b54744e8edd3bf6c2d53e2 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 8 Jul 2026 00:19:06 -0700 Subject: [PATCH 19/31] test(realtime): record and replay websocket traffic in redis vcr cassettes (#32390) * test(realtime): record and replay websocket traffic in redis vcr cassettes * style(realtime): ruff-format ws-vcr harness * fix(realtime): warn instead of silently disabling ws-vcr when the redis client cannot be built --- tests/_ws_vcr.py | 546 +++++++++++++++++++++ tests/llm_translation/conftest.py | 12 +- tests/llm_translation/realtime/conftest.py | 89 ++++ tests/llm_translation/test_ws_vcr.py | 275 +++++++++++ 4 files changed, 917 insertions(+), 5 deletions(-) create mode 100644 tests/_ws_vcr.py create mode 100644 tests/llm_translation/realtime/conftest.py create mode 100644 tests/llm_translation/test_ws_vcr.py diff --git a/tests/_ws_vcr.py b/tests/_ws_vcr.py new file mode 100644 index 00000000000..1f8843a23bd --- /dev/null +++ b/tests/_ws_vcr.py @@ -0,0 +1,546 @@ +"""Record and replay realtime WebSocket traffic in the shared VCR Redis store. + +The HTTP VCR layer (``tests/_vcr_redis_persister.py`` / +``tests/_vcr_conftest_common.py``) only intercepts httpx/aiohttp, so the +realtime suite always reached the live provider. This module intercepts at the +``websockets.connect`` boundary instead and caches whole WebSocket sessions +under a distinct ``litellm:vcr:wscassette:`` key, reusing the same Redis client, +24h TTL, save-on-pass, and best-effort degradation semantics. + +Record mode logs every frame in order with its direction, a text/binary flag, +and, for each server frame, the number of client frames seen before it. That +count is the causal gate for replay: a recorded server frame is only released +once the client has sent at least that many frames, so the deterministic replay +reproduces the same interleaving without a live connection. Client frames are +matched against the recording with volatile fields (ids, timestamps) normalized +away; a structurally different client frame is contract drift and raises loudly +rather than hanging, and every replay wait is bounded by a timeout. +""" + +from __future__ import annotations + +import asyncio +import base64 +import json +import logging +import os +import re +import warnings +from typing import AsyncIterator, Callable, Literal, Optional, Protocol, Union + +from pydantic import BaseModel, ConfigDict, ValidationError +from websockets.exceptions import ConnectionClosedOK + +from tests._vcr_redis_persister import ( + CASSETTE_TTL_SECONDS, + VCRCassetteCacheWarning, + _build_default_client, + _record_cache_failure, +) + +WS_REDIS_KEY_PREFIX = "litellm:vcr:wscassette:" +WS_MAX_SESSIONS_PER_CASSETTE = 20 +WS_MAX_FRAMES_PER_SESSION = 2000 +WS_REPLAY_TIMEOUT_ENV = "LITELLM_WS_VCR_REPLAY_TIMEOUT" +WS_DEFAULT_REPLAY_TIMEOUT_SECONDS = 15.0 +WS_CASSETTE_SCHEMA_VERSION = 1 + +_log = logging.getLogger(__name__) + +Message = Union[str, bytes] +Direction = Literal["client_to_server", "server_to_client"] +Opcode = Literal["text", "binary"] + + +class WsConnectionLike(Protocol): + async def recv(self, decode: Optional[bool] = None) -> Message: ... + + async def send(self, message: Message, *args: object, **kwargs: object) -> None: ... + + async def close(self, *args: object, **kwargs: object) -> None: ... + + def __aiter__(self) -> AsyncIterator[Message]: ... + + +class WsConnectContextLike(Protocol): + async def __aenter__(self) -> WsConnectionLike: ... + + async def __aexit__(self, *exc_info: object) -> Optional[bool]: ... + + +class RedisLike(Protocol): + def get(self, key: str) -> Optional[bytes]: ... + + def set(self, key: str, value: bytes, ex: int) -> object: ... + + +class WsFrame(BaseModel): + model_config = ConfigDict(frozen=True) + + direction: Direction + opcode: Opcode + text: Optional[str] = None + binary_b64: Optional[str] = None + client_frames_before: Optional[int] = None + + +class WsSession(BaseModel): + model_config = ConfigDict(frozen=True) + + frames: tuple[WsFrame, ...] + + +class WsCassette(BaseModel): + model_config = ConfigDict(frozen=True) + + schema_version: int = WS_CASSETTE_SCHEMA_VERSION + sessions: tuple[WsSession, ...] + + +class WsVcrReplayError(Exception): ... + + +class WsVcrContractDrift(WsVcrReplayError): ... + + +class WsVcrReplayTimeout(WsVcrReplayError): ... + + +_UUID_RE = re.compile(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}") +_OPENAI_ID_RE = re.compile( + r"\b(?:evt|event|item|msg|resp|response|sess|session|call|fc|rs|conv|ce|audio)_[A-Za-z0-9]{6,}" +) +_ISO_TS_RE = re.compile(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?") +_EPOCH_RE = re.compile(r"(? str: + scrubbed = _BEARER_RE.sub("Bearer ", text) + scrubbed = _OPENAI_KEY_RE.sub("", scrubbed) + scrubbed = _XAI_KEY_RE.sub("", scrubbed) + return scrubbed + + +def _normalize_json_for_match(obj: object) -> object: + if isinstance(obj, dict): + return { + str(key): ("" if key in _VOLATILE_KEYS else _normalize_json_for_match(value)) + for key, value in sorted(obj.items(), key=lambda kv: str(kv[0])) + } + if isinstance(obj, list): + return [_normalize_json_for_match(item) for item in obj] + if isinstance(obj, str): + return _normalize_scalar_string(obj) + return obj + + +def _normalize_scalar_string(text: str) -> str: + normalized = _UUID_RE.sub("", text) + normalized = _OPENAI_ID_RE.sub("", normalized) + normalized = _ISO_TS_RE.sub("", normalized) + normalized = _EPOCH_RE.sub("", normalized) + return normalized + + +def normalize_text_for_match(text: str) -> str: + try: + parsed = json.loads(text) + except (ValueError, TypeError): + return _normalize_scalar_string(text) + return json.dumps(_normalize_json_for_match(parsed), sort_keys=True, separators=(",", ":")) + + +def text_frames_match(recorded: str, incoming: str) -> bool: + return normalize_text_for_match(recorded) == normalize_text_for_match(incoming) + + +def _frame_payload(message: Message) -> tuple[Opcode, Optional[str], Optional[str]]: + if isinstance(message, str): + return "text", scrub_secrets(message), None + try: + decoded = message.decode("utf-8") + except UnicodeDecodeError: + return "binary", None, base64.b64encode(message).decode("ascii") + return "text", scrub_secrets(decoded), None + + +def ws_redis_key_for(nodeid: str) -> str: + rel = nodeid.replace("::", "/").replace("\\", "/").lstrip("./") + return f"{WS_REDIS_KEY_PREFIX}{rel}" + + +def replay_timeout_seconds() -> float: + raw = os.environ.get(WS_REPLAY_TIMEOUT_ENV) + if not raw: + return WS_DEFAULT_REPLAY_TIMEOUT_SECONDS + try: + return float(raw) + except ValueError: + return WS_DEFAULT_REPLAY_TIMEOUT_SECONDS + + +class WsSessionRecorder: + def __init__(self) -> None: + self._frames: list[WsFrame] = [] + self._client_count = 0 + + def record_client_frame(self, message: Message) -> None: + opcode, text, binary_b64 = _frame_payload(message) + self._frames.append(WsFrame(direction="client_to_server", opcode=opcode, text=text, binary_b64=binary_b64)) + self._client_count += 1 + + def record_server_frame(self, message: Message) -> None: + opcode, text, binary_b64 = _frame_payload(message) + self._frames.append( + WsFrame( + direction="server_to_client", + opcode=opcode, + text=text, + binary_b64=binary_b64, + client_frames_before=self._client_count, + ) + ) + + def to_session(self) -> WsSession: + return WsSession(frames=tuple(self._frames)) + + +class RecordingConnection: + def __init__(self, real: WsConnectionLike, recorder: WsSessionRecorder) -> None: + self._real = real + self._recorder = recorder + + async def recv(self, decode: Optional[bool] = None) -> Message: + result = await self._real.recv(decode=decode) + self._recorder.record_server_frame(result) + return result + + async def send(self, message: Message, *args: object, **kwargs: object) -> None: + self._recorder.record_client_frame(message) + await self._real.send(message, *args, **kwargs) + + async def close(self, *args: object, **kwargs: object) -> None: + await self._real.close(*args, **kwargs) + + def __aiter__(self) -> AsyncIterator[Message]: + return self._iterate() + + async def _iterate(self) -> AsyncIterator[Message]: + async for message in self._real: + self._recorder.record_server_frame(message) + yield message + + +class ReplayConnection: + def __init__( + self, + session: WsSession, + timeout: float, + on_error: Callable[[WsVcrReplayError], None], + ) -> None: + self._server_frames = tuple(f for f in session.frames if f.direction == "server_to_client") + self._client_frames = tuple(f for f in session.frames if f.direction == "client_to_server") + self._timeout = timeout + self._on_error = on_error + self._server_cursor = 0 + self._client_cursor = 0 + self._client_sent = 0 + self._closed = False + self._progress = asyncio.Event() + + async def recv(self, decode: Optional[bool] = None) -> Message: + want_bytes = decode is False + while True: + if self._closed or self._server_cursor >= len(self._server_frames): + raise ConnectionClosedOK(None, None) + frame = self._server_frames[self._server_cursor] + needed = frame.client_frames_before or 0 + if self._client_sent >= needed: + self._server_cursor += 1 + return _materialize_frame(frame, want_bytes) + await self._await_client_progress(needed) + + async def _await_client_progress(self, needed: int) -> None: + waiter = self._progress + try: + await asyncio.wait_for(waiter.wait(), timeout=self._timeout) + except asyncio.TimeoutError: + error = WsVcrReplayTimeout( + f"WS-VCR replay stalled: server frame #{self._server_cursor} needs " + f"{needed} client frame(s) but only {self._client_sent} were sent within " + f"{self._timeout}s. The client stopped driving the recorded session." + ) + self._on_error(error) + raise error + + async def send(self, message: Message, *args: object, **kwargs: object) -> None: + if self._client_cursor >= len(self._client_frames): + error = WsVcrContractDrift( + "WS-VCR contract drift: client sent frame " + f"#{self._client_cursor + 1} but the recording has only " + f"{len(self._client_frames)} client frame(s). Extra frame: {_preview(message)}" + ) + self._on_error(error) + raise error + recorded = self._client_frames[self._client_cursor] + if not _client_frame_matches(recorded, message): + error = WsVcrContractDrift( + "WS-VCR contract drift on client frame " + f"#{self._client_cursor + 1}:\n recorded: {_preview_frame(recorded)}\n" + f" got: {_preview(message)}" + ) + self._on_error(error) + raise error + self._client_cursor += 1 + self._client_sent += 1 + self._signal_progress() + + async def close(self, *args: object, **kwargs: object) -> None: + self._closed = True + self._signal_progress() + + def _signal_progress(self) -> None: + previous = self._progress + self._progress = asyncio.Event() + previous.set() + + def __aiter__(self) -> AsyncIterator[Message]: + return self._iterate() + + async def _iterate(self) -> AsyncIterator[Message]: + while True: + try: + yield await self.recv() + except ConnectionClosedOK: + return + + +def _materialize_frame(frame: WsFrame, want_bytes: bool) -> Message: + if frame.opcode == "text": + text = frame.text or "" + return text.encode("utf-8") if want_bytes else text + return base64.b64decode(frame.binary_b64 or "") + + +def _client_frame_matches(recorded: WsFrame, message: Message) -> bool: + opcode, text, binary_b64 = _frame_payload(message) + if recorded.opcode != opcode: + return False + if opcode == "text": + return text_frames_match(recorded.text or "", text or "") + return recorded.binary_b64 == binary_b64 + + +def _preview(message: Message) -> str: + text = message if isinstance(message, str) else message.decode("utf-8", errors="replace") + return scrub_secrets(text)[:200] + + +def _preview_frame(frame: WsFrame) -> str: + if frame.opcode == "text": + return (frame.text or "")[:200] + return f"" + + +class _RecordingConnect: + def __init__( + self, + real_context: WsConnectContextLike, + recorder: WsSessionRecorder, + on_done: Callable[[WsSessionRecorder], None], + ) -> None: + self._real_context = real_context + self._recorder = recorder + self._on_done = on_done + + async def __aenter__(self) -> RecordingConnection: + real = await self._real_context.__aenter__() + return RecordingConnection(real, self._recorder) + + async def __aexit__(self, *exc_info: object) -> Optional[bool]: + try: + return await self._real_context.__aexit__(*exc_info) + finally: + self._on_done(self._recorder) + + +class _ReplayConnect: + def __init__( + self, + session: WsSession, + timeout: float, + on_error: Callable[[WsVcrReplayError], None], + ) -> None: + self._session = session + self._timeout = timeout + self._on_error = on_error + + async def __aenter__(self) -> ReplayConnection: + return ReplayConnection(self._session, self._timeout, self._on_error) + + async def __aexit__(self, *exc_info: object) -> bool: + return False + + +class WsVcrController: + def __init__( + self, + original_connect: Callable[..., WsConnectContextLike], + cassette: Optional[WsCassette], + timeout: float, + ) -> None: + self._original_connect = original_connect + self._cassette = cassette + self._timeout = timeout + self._replay_cursor = 0 + self._recorded_sessions: list[WsSession] = [] + self._errors: list[WsVcrReplayError] = [] + self._replayed = False + self._recorded = False + + def connect(self, *args: object, **kwargs: object) -> object: + if self._cassette is not None and self._replay_cursor < len(self._cassette.sessions): + session = self._cassette.sessions[self._replay_cursor] + self._replay_cursor += 1 + self._replayed = True + return _ReplayConnect(session, self._timeout, self._errors.append) + self._recorded = True + recorder = WsSessionRecorder() + return _RecordingConnect(self._original_connect(*args, **kwargs), recorder, self._finish_recorder) + + def _finish_recorder(self, recorder: WsSessionRecorder) -> None: + self._recorded_sessions.append(recorder.to_session()) + + @property + def errors(self) -> tuple[WsVcrReplayError, ...]: + return tuple(self._errors) + + @property + def replayed(self) -> bool: + return self._replayed + + @property + def recorded(self) -> bool: + return self._recorded + + def built_cassette(self) -> Optional[WsCassette]: + if not self._recorded_sessions: + return None + return WsCassette(sessions=tuple(self._recorded_sessions)) + + def verdict(self) -> str: + if self._replayed and not self._recorded: + return f"[WS-VCR HIT] sessions={self._replay_cursor} frames={self._played_frame_count()}" + if self._recorded: + cassette = self.built_cassette() + frames = _cassette_frame_count(cassette) if cassette is not None else 0 + return f"[WS-VCR MISS] recorded sessions={len(self._recorded_sessions)} frames={frames}" + return "[WS-VCR NOOP] (no websocket traffic)" + + def _played_frame_count(self) -> int: + if self._cassette is None: + return 0 + return sum(len(s.frames) for s in self._cassette.sessions[: self._replay_cursor]) + + +def _cassette_frame_count(cassette: WsCassette) -> int: + return sum(len(s.frames) for s in cassette.sessions) + + +def load_ws_cassette(client: RedisLike, key: str) -> Optional[WsCassette]: + from redis.exceptions import RedisError + + try: + data = client.get(key) + except RedisError as exc: + _record_cache_failure("load", exc) + message = f"WS-VCR redis load failed for {key}; treating as cache miss: {type(exc).__name__}: {exc}" + _log.warning(message) + warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2) + return None + if data is None: + return None + try: + raw = data.decode("utf-8") if isinstance(data, (bytes, bytearray)) else data + return WsCassette.model_validate_json(raw) + except (ValidationError, ValueError, TypeError) as exc: + _record_cache_failure("load", exc) + message = ( + f"WS-VCR redis load failed for {key}; cached payload is corrupt, " + f"treating as cache miss: {type(exc).__name__}: {exc}" + ) + _log.warning(message) + warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2) + return None + + +def save_ws_cassette( + client: RedisLike, + key: str, + cassette: WsCassette, + passed: bool, + ttl_seconds: int = CASSETTE_TTL_SECONDS, +) -> bool: + from redis.exceptions import RedisError + + if not passed: + _log.info("WS-VCR redis save skipped for %s; test did not pass - leaving any prior cassette intact", key) + return False + if len(cassette.sessions) > WS_MAX_SESSIONS_PER_CASSETTE: + _log.warning( + "WS-VCR redis save refused for %s; %d sessions (> WS_MAX_SESSIONS_PER_CASSETTE=%d)", + key, + len(cassette.sessions), + WS_MAX_SESSIONS_PER_CASSETTE, + ) + return False + if any(len(session.frames) > WS_MAX_FRAMES_PER_SESSION for session in cassette.sessions): + _log.warning( + "WS-VCR redis save refused for %s; a session exceeds WS_MAX_FRAMES_PER_SESSION=%d", + key, + WS_MAX_FRAMES_PER_SESSION, + ) + return False + payload = cassette.model_dump_json().encode("utf-8") + try: + client.set(key, payload, ex=ttl_seconds) + except RedisError as exc: + _record_cache_failure("save", exc) + message = f"WS-VCR redis save failed for {key}; cassette not persisted: {type(exc).__name__}: {exc}" + _log.warning(message) + warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2) + return False + return True + + +def build_ws_cassette_client( + builder: Callable[[], RedisLike] = _build_default_client, +) -> Optional[RedisLike]: + try: + return builder() + except Exception as exc: + _record_cache_failure("load", exc) + message = ( + f"WS-VCR redis client unavailable; realtime tests fall back to live " + f"websocket traffic: {type(exc).__name__}: {exc}" + ) + _log.warning(message) + warnings.warn(message, VCRCassetteCacheWarning, stacklevel=2) + return None diff --git a/tests/llm_translation/conftest.py b/tests/llm_translation/conftest.py index 77fcb46a2b1..f5b71236e92 100644 --- a/tests/llm_translation/conftest.py +++ b/tests/llm_translation/conftest.py @@ -41,11 +41,13 @@ def fake_openai_endpoint(): # Per-item respx detection (``apply_vcr_auto_marker_to_items``) handles -# the vast majority of respx-vs-vcrpy conflicts automatically. The only -# entry below is the persister's own unit-test file, which exercises -# ``save_cassette`` / ``load_cassette`` against fakeredis and must not -# itself run under a live cassette context. -_VCR_AUTO_MARKER_SKIP_FILES = frozenset({"test_vcr_redis_persister.py"}) +# the vast majority of respx-vs-vcrpy conflicts automatically. The entries +# below are the persister's and the WebSocket VCR's own unit-test files, which +# exercise ``save_cassette`` / ``load_cassette`` against fakeredis and must not +# themselves run under a live cassette context. +_VCR_AUTO_MARKER_SKIP_FILES = frozenset( + {"test_vcr_redis_persister.py", "test_ws_vcr.py"} +) _VCR_INCOMPATIBLE_NODEID_SUFFIXES: tuple[str, ...] = () diff --git a/tests/llm_translation/realtime/conftest.py b/tests/llm_translation/realtime/conftest.py new file mode 100644 index 00000000000..e305131593e --- /dev/null +++ b/tests/llm_translation/realtime/conftest.py @@ -0,0 +1,89 @@ +"""WebSocket VCR wiring for the realtime suite. + +This directory inherits the HTTP VCR machinery from +``tests/llm_translation/conftest.py`` (which only intercepts httpx/aiohttp and +is therefore a no-op for realtime WebSocket traffic). The autouse fixture below +adds the WebSocket layer: it patches ``websockets.connect`` for the duration of +each test so realtime frames are recorded to, or replayed from, the same +cassette Redis under a ``litellm:vcr:wscassette:`` prefix. +""" + +from __future__ import annotations + +import os +import sys +from typing import Optional + +import pytest + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))) + +from tests._vcr_conftest_common import ( # noqa: E402 + vcr_disabled, + vcr_outcome_logging_enabled, +) +from tests._ws_vcr import ( # noqa: E402 + WsVcrController, + build_ws_cassette_client, + load_ws_cassette, + replay_timeout_seconds, + save_ws_cassette, + ws_redis_key_for, +) + +_ws_cassette_client: Optional[object] = None + + +def _get_ws_cassette_client() -> Optional[object]: + global _ws_cassette_client + if _ws_cassette_client is None: + _ws_cassette_client = build_ws_cassette_client() + return _ws_cassette_client + + +def _emit_verdict(request: pytest.FixtureRequest, verdict: str) -> None: + if os.environ.get("PYTEST_XDIST_WORKER"): + return + reporter = request.config.pluginmanager.getplugin("terminalreporter") + if reporter is None: + return + reporter.write_line(f"{verdict} :: {request.node.nodeid}") + + +@pytest.fixture(autouse=True) +def _ws_vcr(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch): + if vcr_disabled(): + yield + return + + import websockets + + client = _get_ws_cassette_client() + if client is None: + yield + return + + key = ws_redis_key_for(request.node.nodeid) + cassette = load_ws_cassette(client, key) + controller = WsVcrController( + original_connect=websockets.connect, + cassette=cassette, + timeout=replay_timeout_seconds(), + ) + monkeypatch.setattr(websockets, "connect", controller.connect) + + yield + + rep_call = getattr(request.node, "rep_call", None) + passed = bool(rep_call and rep_call.passed) + + if controller.recorded: + built = controller.built_cassette() + if built is not None: + save_ws_cassette(client, key, built, passed=passed) + + if vcr_outcome_logging_enabled(): + _emit_verdict(request, controller.verdict()) + + if controller.errors and passed: + raise controller.errors[0] diff --git a/tests/llm_translation/test_ws_vcr.py b/tests/llm_translation/test_ws_vcr.py new file mode 100644 index 00000000000..1a72d62d80f --- /dev/null +++ b/tests/llm_translation/test_ws_vcr.py @@ -0,0 +1,275 @@ +from __future__ import annotations + +import asyncio +import os +import sys +import warnings + +import fakeredis +import pytest +from websockets.exceptions import ConnectionClosedOK + +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))) + +from tests._vcr_redis_persister import ( # noqa: E402 + VCRCassetteCacheWarning, + cassette_cache_health, +) +from tests._ws_vcr import ( # noqa: E402 + CASSETTE_TTL_SECONDS, + RedisLike, + ReplayConnection, + WsCassette, + WsFrame, + WsSession, + WsSessionRecorder, + WsVcrContractDrift, + WsVcrReplayError, + WsVcrReplayTimeout, + build_ws_cassette_client, + load_ws_cassette, + save_ws_cassette, + scrub_secrets, + text_frames_match, + ws_redis_key_for, +) + + +def _server(text: str, client_frames_before: int) -> WsFrame: + return WsFrame( + direction="server_to_client", + opcode="text", + text=text, + client_frames_before=client_frames_before, + ) + + +def _client(text: str) -> WsFrame: + return WsFrame(direction="client_to_server", opcode="text", text=text) + + +def _collect_errors(): + errors: list[WsVcrReplayError] = [] + return errors, errors.append + + +def test_cassette_json_roundtrip_preserves_frames_and_gate(): + cassette = WsCassette( + sessions=( + WsSession( + frames=( + _server('{"type":"session.created"}', 0), + _client('{"type":"response.create"}'), + _server('{"type":"response.done"}', 1), + WsFrame( + direction="server_to_client", opcode="binary", binary_b64="dGVzdA==", client_frames_before=1 + ), + ) + ), + ) + ) + + restored = WsCassette.model_validate_json(cassette.model_dump_json()) + + assert restored == cassette + assert restored.sessions[0].frames[2].client_frames_before == 1 + assert restored.sessions[0].frames[3].opcode == "binary" + assert restored.sessions[0].frames[3].binary_b64 == "dGVzdA==" + + +def test_recorder_tracks_client_frame_count_as_causal_gate(): + recorder = WsSessionRecorder() + recorder.record_server_frame('{"type":"session.created"}') + recorder.record_client_frame('{"type":"conversation.item.create"}') + recorder.record_client_frame('{"type":"response.create"}') + recorder.record_server_frame('{"type":"response.done"}') + + session = recorder.to_session() + server_frames = [f for f in session.frames if f.direction == "server_to_client"] + + assert server_frames[0].client_frames_before == 0 + assert server_frames[1].client_frames_before == 2 + + +async def test_replay_recv_returns_bytes_when_decode_false_and_str_otherwise(): + session = WsSession(frames=(_server("hello", 0), _server("world", 0))) + _, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=1.0, on_error=on_error) + + as_bytes = await conn.recv(decode=False) + as_str = await conn.recv() + + assert as_bytes == b"hello" + assert as_str == "world" + + +async def test_replay_serves_server_frame_only_after_causal_client_count_met(): + session = WsSession( + frames=( + _server('{"type":"session.created"}', 0), + _client('{"type":"response.create"}'), + _server('{"type":"response.done"}', 1), + ) + ) + _, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=2.0, on_error=on_error) + + first = await conn.recv(decode=False) + assert first == b'{"type":"session.created"}' + + gated = asyncio.ensure_future(conn.recv(decode=False)) + await asyncio.sleep(0.1) + assert not gated.done(), "gated server frame was released before the recorded client frame was sent" + + await conn.send('{"type":"response.create"}') + released = await asyncio.wait_for(gated, timeout=1.0) + assert released == b'{"type":"response.done"}' + + +async def test_replay_exhausted_server_frames_raise_connection_closed(): + session = WsSession(frames=(_server("only", 0),)) + _, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=1.0, on_error=on_error) + + await conn.recv(decode=False) + with pytest.raises(ConnectionClosedOK): + await conn.recv(decode=False) + + +async def test_replay_timeout_raises_instead_of_hanging(): + session = WsSession( + frames=( + _server('{"type":"session.created"}', 0), + _server('{"type":"response.done"}', 5), + ) + ) + errors, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=0.15, on_error=on_error) + + await conn.recv(decode=False) + with pytest.raises(WsVcrReplayTimeout): + await asyncio.wait_for(conn.recv(decode=False), timeout=2.0) + assert errors and isinstance(errors[0], WsVcrReplayTimeout) + + +async def test_replay_accepts_client_frame_with_volatile_id_drift(): + recorded_client = _client('{"type":"conversation.item.create","item":{"id":"item_ABC12345","role":"user"}}') + session = WsSession(frames=(_server("s", 0), recorded_client, _server("done", 1))) + errors, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=1.0, on_error=on_error) + + await conn.recv(decode=False) + await conn.send('{"type":"conversation.item.create","item":{"role":"user","id":"item_ZZ99887766"}}') + + assert errors == [] + assert await conn.recv(decode=False) == b"done" + + +async def test_replay_rejects_structurally_different_client_frame(): + session = WsSession(frames=(_server("s", 0), _client('{"type":"response.create"}'), _server("done", 1))) + errors, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=1.0, on_error=on_error) + + await conn.recv(decode=False) + with pytest.raises(WsVcrContractDrift): + await conn.send('{"type":"session.update","session":{"voice":"alloy"}}') + assert errors and isinstance(errors[0], WsVcrContractDrift) + + +async def test_replay_rejects_extra_client_frame_beyond_recording(): + session = WsSession(frames=(_server("s", 0), _client('{"type":"response.create"}'))) + errors, on_error = _collect_errors() + conn = ReplayConnection(session, timeout=1.0, on_error=on_error) + + await conn.send('{"type":"response.create"}') + with pytest.raises(WsVcrContractDrift): + await conn.send('{"type":"response.create"}') + assert errors + + +def test_text_frames_match_normalizes_ids_and_timestamps_but_not_structure(): + assert text_frames_match( + '{"type":"x","event_id":"evt_111","ts":"2026-05-25T03:40:37.262045Z"}', + '{"type":"x","event_id":"evt_999","ts":"2026-06-01T10:00:00Z"}', + ) + assert not text_frames_match('{"type":"x","text":"hi"}', '{"type":"x","text":"bye"}') + assert not text_frames_match('{"type":"x"}', '{"type":"x","extra":1}') + + +def test_scrub_secrets_removes_auth_material(): + scrubbed = scrub_secrets("Authorization: Bearer sk-abcdef123456 and key xai-zzz99988877 raw sk-plainkey123") + assert "sk-abcdef123456" not in scrubbed + assert "xai-zzz99988877" not in scrubbed + assert "sk-plainkey123" not in scrubbed + assert "Bearer " in scrubbed + + +def test_recorder_scrubs_secrets_in_stored_frames(): + recorder = WsSessionRecorder() + recorder.record_client_frame('{"authorization":"Bearer sk-supersecretvalue"}') + stored = recorder.to_session().frames[0].text + assert stored is not None + assert "sk-supersecretvalue" not in stored + + +def _sample_cassette() -> WsCassette: + return WsCassette(sessions=(WsSession(frames=(_server('{"type":"session.created"}', 0),)),)) + + +def test_save_sets_24h_ttl_and_load_roundtrips(): + fake = fakeredis.FakeStrictRedis() + key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_y") + + assert save_ws_cassette(fake, key, _sample_cassette(), passed=True) is True + + ttl = fake.ttl(key) + assert CASSETTE_TTL_SECONDS - 5 <= ttl <= CASSETTE_TTL_SECONDS + loaded = load_ws_cassette(fake, key) + assert loaded == _sample_cassette() + + +def test_save_skipped_when_test_failed_leaves_no_key(): + fake = fakeredis.FakeStrictRedis() + key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_fail") + + assert save_ws_cassette(fake, key, _sample_cassette(), passed=False) is False + assert fake.get(key) is None + + +def test_save_skipped_when_test_failed_preserves_prior_cassette(): + fake = fakeredis.FakeStrictRedis() + key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::test_keep") + + save_ws_cassette(fake, key, _sample_cassette(), passed=True) + newer = WsCassette(sessions=(WsSession(frames=(_server('{"type":"other"}', 0),)),)) + + assert save_ws_cassette(fake, key, newer, passed=False) is False + assert load_ws_cassette(fake, key) == _sample_cassette() + + +def test_load_missing_key_returns_none(): + fake = fakeredis.FakeStrictRedis() + assert load_ws_cassette(fake, ws_redis_key_for("never/recorded")) is None + + +def test_ws_redis_key_uses_distinct_prefix(): + key = ws_redis_key_for("tests/llm_translation/realtime/test_x.py::TestY::test_z") + assert key.startswith("litellm:vcr:wscassette:") + assert "::" not in key + + +def test_build_ws_cassette_client_warns_and_counts_failure_instead_of_silently_disabling(): + def _broken_builder() -> RedisLike: + raise ValueError("invalid CASSETTE_REDIS_URL") + + failures_before = cassette_cache_health()["load_failures"] + with pytest.warns(VCRCassetteCacheWarning, match="fall back to live websocket traffic"): + assert build_ws_cassette_client(builder=_broken_builder) is None + assert cassette_cache_health()["load_failures"] == failures_before + 1 + + +def test_build_ws_cassette_client_returns_built_client_without_warning(): + fake = fakeredis.FakeStrictRedis() + with warnings.catch_warnings(): + warnings.simplefilter("error", VCRCassetteCacheWarning) + assert build_ws_cassette_client(builder=lambda: fake) is fake From bfff5e8d868312fcec9fe7fd9aaa3df14aa31ea3 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 8 Jul 2026 19:02:48 +0300 Subject: [PATCH 20/31] fix(mcp): log MCP tool calls returning isError=true as failures (#32238) An MCP tool call that completes with CallToolResult.isError=true correctly returns HTTP 200 per the MCP spec, but the shared post-call logging helper always fired async_success_handler, so the standard logging payload carried status=success and OTel (whose _parse_error only marks ERROR on status=failure) showed green spans for failed tools. The helper now checks the result after async_post_mcp_tool_call_hook runs (guardrails may flip isError there) and routes error results to the failure path: success gates are consumed so the @client wrapper cannot enqueue a success log, failure_handler and async_failure_handler fire with a new MCPToolResultError carrying the tool's first text content, and post_call_failure_hook records the failure the same way raised exceptions already do. Raised exceptions never reach the helper, so no double failure logging. HTTP wire behavior is unchanged Resolves LIT-4081 --- .../_experimental/mcp_server/exceptions.py | 15 + .../mcp_server/rest_endpoints.py | 36 ++- .../proxy/_experimental/mcp_server/server.py | 72 ++++- .../proxy/_experimental/mcp_server/utils.py | 19 ++ .../mcp_server/test_mcp_server.py | 304 ++++++++++++++++++ .../mcp_server/test_mcp_tool_search.py | 2 +- .../mcp_server/test_rest_endpoints.py | 6 +- 7 files changed, 440 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/exceptions.py b/litellm/proxy/_experimental/mcp_server/exceptions.py index b3f7ca9bbe2..3e3e549008d 100644 --- a/litellm/proxy/_experimental/mcp_server/exceptions.py +++ b/litellm/proxy/_experimental/mcp_server/exceptions.py @@ -73,3 +73,18 @@ class MCPUpstreamAuthError(Exception): detail=detail, headers={"www-authenticate": challenge} if challenge else None, ) + + +class MCPToolResultError(Exception): + """An MCP tool call completed with ``isError=True`` in its result. + + Never raised on the wire path: streamable HTTP MCP correctly returns tool + failures as HTTP 200 with ``result.isError: true`` per the MCP spec. This + exception only drives the standard failure logging (``status="failure"`` + payload, OTel ERROR span) for such results. + + Lives here rather than ``utils.py`` deliberately: tests reload ``utils`` + to re-read its env-derived constants, and a reload would fork this class + into two identities, breaking ``isinstance`` checks against instances + created before the reload. + """ diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index d482e537c5d..b917530dd52 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -8,6 +8,7 @@ from typing import ( Dict, List, Literal, + Mapping, Optional, Set, Tuple, @@ -78,7 +79,7 @@ if MCP_AVAILABLE: MCPInfo, MCPServer, _apply_toolset_scope, - _fire_mcp_success_logging, + _fire_mcp_tool_call_logging, _tool_name_matches, execute_mcp_tool, filter_tools_by_allowed_tools, @@ -86,23 +87,32 @@ if MCP_AVAILABLE: ######################################################## ############ MCP Server REST API Routes ################# - async def _safe_fire_mcp_success_logging( + async def _safe_fire_mcp_tool_call_logging( logging_obj: Optional[Any], result: Any, start_time: datetime, end_time: datetime, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + request_data: Optional[Mapping[str, object]] = None, ) -> None: if logging_obj is None: return logging_results = await asyncio.gather( - _fire_mcp_success_logging(logging_obj, result, start_time, end_time), + _fire_mcp_tool_call_logging( + logging_obj, + result, + start_time, + end_time, + user_api_key_auth=user_api_key_auth, + request_data=request_data, + ), return_exceptions=True, ) logging_error = logging_results[0] if isinstance(logging_error, asyncio.CancelledError): raise logging_error if isinstance(logging_error, BaseException): - verbose_logger.warning("MCP tool success logging failed (continuing): %s", logging_error) + verbose_logger.warning("MCP tool call logging failed (continuing): %s", logging_error) def _get_server_auth_header( server, @@ -872,7 +882,14 @@ if MCP_AVAILABLE: raw_headers=virtual_raw_headers, litellm_logging_obj=virtual_logging_obj, ) - await _safe_fire_mcp_success_logging(virtual_logging_obj, result, _tool_start_time, datetime.now()) + await _safe_fire_mcp_tool_call_logging( + virtual_logging_obj, + result, + _tool_start_time, + datetime.now(), + user_api_key_auth=user_api_key_dict, + request_data=data, + ) return result # Validate required parameters early @@ -955,7 +972,14 @@ if MCP_AVAILABLE: litellm_logging_obj=data.get("litellm_logging_obj"), requested_server_id=canonical_server_id, ) - await _safe_fire_mcp_success_logging(logging_obj, result, _tool_start_time, datetime.now()) + await _safe_fire_mcp_tool_call_logging( + logging_obj, + result, + _tool_start_time, + datetime.now(), + user_api_key_auth=user_api_key_dict, + request_data=data, + ) return result except MCPMissingUserEnvVarsError as e: verbose_logger.info( diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index fc847182a60..c03e49a1628 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -20,6 +20,7 @@ from typing import ( Callable, Dict, List, + Mapping, Optional, Set, Tuple, @@ -47,7 +48,10 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( get_request_base_url, ) -from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError +from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPToolResultError, + MCPUpstreamAuthError, +) from litellm.proxy._experimental.mcp_server.mcp_context import ( _mcp_active_toolset_id, _mcp_gateway_initialize_instructions, @@ -60,6 +64,7 @@ from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_VERSION, MCPMissingUserEnvVarsError, add_server_prefix_to_name, + extract_mcp_tool_result_error_message, get_server_prefix, iter_known_server_prefixes, ) @@ -2743,12 +2748,40 @@ if MCP_AVAILABLE: return response - async def _fire_mcp_success_logging( + _MCP_CREDENTIAL_REQUEST_FIELDS = frozenset( + { + "raw_headers", + "mcp_auth_header", + "mcp_server_auth_headers", + "oauth2_headers", + "user_api_key_auth", + } + ) + + async def _fire_mcp_tool_call_logging( logging_obj: LiteLLMLoggingObj, result: Any, start_time: datetime, end_time: datetime, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + request_data: Optional[Mapping[str, object]] = None, ) -> None: + """Fire post-call logging for an executed MCP tool call. + + A result with ``isError=True`` is logged as a failure (``status="failure"`` + payload, so OTel marks the span ERROR) while the HTTP wire behavior stays + 200 + ``isError: true`` per the MCP spec. The error check runs after + ``async_post_mcp_tool_call_hook`` because guardrails may flip the result + to ``isError=True`` in that hook. Raised exceptions never reach here (the + ``@client`` wrapper and ``call_mcp_tool``'s except path log those), so + this cannot double-log a failure. + + ``request_data`` may carry credential-bearing fields (the REST path puts + ``raw_headers``, ``mcp_auth_header``, ``mcp_server_auth_headers``, and + ``oauth2_headers`` at the top level of its data dict), so those are + stripped before the dict is handed to ``post_call_failure_hook`` + callbacks. + """ logging_obj.post_call(original_response=result) await logging_obj.async_post_mcp_tool_call_hook( kwargs=logging_obj.model_call_details, @@ -2757,7 +2790,31 @@ if MCP_AVAILABLE: end_time=end_time, ) logging_obj.call_type = CallTypes.call_mcp_tool.value - await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) + error_message = extract_mcp_tool_result_error_message(result) + if error_message is None: + await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) + return + + logging_obj.has_run_logging(event_type="sync_success") + logging_obj.has_run_logging(event_type="async_success") + tool_error = MCPToolResultError(error_message) + logging_obj.failure_handler(tool_error, "", start_time, end_time) + await logging_obj.async_failure_handler(tool_error, "", start_time, end_time) + + if user_api_key_auth is None: + return + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj: + sanitized_request_data = { + key: value for key, value in (request_data or {}).items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS + } + await proxy_logging_obj.post_call_failure_hook( + request_data=sanitized_request_data, + original_exception=tool_error, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + ) @client async def call_mcp_tool( @@ -2833,7 +2890,14 @@ if MCP_AVAILABLE: raise if litellm_logging_obj: - await _fire_mcp_success_logging(litellm_logging_obj, response, start_time, datetime.now()) + await _fire_mcp_tool_call_logging( + logging_obj=litellm_logging_obj, + result=response, + start_time=start_time, + end_time=datetime.now(), + user_api_key_auth=user_api_key_auth, + request_data=kwargs, + ) return response async def mcp_get_prompt( diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index c9c60030dbc..80a469b8c1a 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -415,6 +415,25 @@ def validate_mcp_server_name(server_name: str, raise_http_exception: bool = Fals raise Exception(error_message) +def extract_mcp_tool_result_error_message(result: object) -> Optional[str]: + """The first text content of an ``isError=True`` tool result, or ``None`` + when the result is not an error. + + Accepts both ``mcp.types.CallToolResult`` objects and their dict + equivalents, duck-typed so the ``mcp`` package is not required. + """ + is_error: object = result.get("isError") if isinstance(result, Mapping) else getattr(result, "isError", None) + if is_error is not True: + return None + content: object = result.get("content") if isinstance(result, Mapping) else getattr(result, "content", None) + if isinstance(content, (list, tuple)): + for item in content: + text: object = item.get("text") if isinstance(item, Mapping) else getattr(item, "text", None) + if isinstance(text, str) and text: + return text + return "MCP tool call returned isError=true" + + TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$") 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 b24457deabd..83c19dfd7ca 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 @@ -8,8 +8,10 @@ from fastapi import HTTPException from mcp import ReadResourceResult, Resource from mcp.types import ( BlobResourceContents, + CallToolResult, Prompt, ResourceTemplate, + TextContent, TextResourceContents, ) @@ -6598,6 +6600,308 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_ prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited() +# --------------------------------------------------------------------------- # +# MCP tool-call isError failure logging +# --------------------------------------------------------------------------- # + + +def _call_tool_result(is_error: bool, text: str) -> CallToolResult: + return CallToolResult(content=[TextContent(type="text", text=text)], isError=is_error) + + +def _mock_mcp_logging_obj() -> MagicMock: + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.async_post_mcp_tool_call_hook = AsyncMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.async_failure_handler = AsyncMock() + return logging_obj + + +def test_extract_mcp_tool_result_error_message(): + from litellm.proxy._experimental.mcp_server.utils import ( + extract_mcp_tool_result_error_message, + ) + + assert extract_mcp_tool_result_error_message(_call_tool_result(True, "boom")) == "boom" + assert extract_mcp_tool_result_error_message(_call_tool_result(False, "ok")) is None + assert ( + extract_mcp_tool_result_error_message(CallToolResult(content=[], isError=True)) + == "MCP tool call returned isError=true" + ) + assert ( + extract_mcp_tool_result_error_message({"isError": True, "content": [{"type": "text", "text": "denied"}]}) + == "denied" + ) + assert extract_mcp_tool_result_error_message({"isError": False, "content": []}) is None + assert extract_mcp_tool_result_error_message({}) is None + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_iserror_logs_failure(): + """Regression test: a CallToolResult with isError=True must go + down the failure logging path (async_failure_handler + post_call_failure_hook), + never async_success_handler.""" + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError + + logging_obj = _mock_mcp_logging_obj() + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=_call_tool_result(True, "upstream exploded"), + start_time=datetime.now(), + end_time=datetime.now(), + user_api_key_auth=user_auth, + request_data={"litellm_call_id": "cid"}, + ) + + logging_obj.async_success_handler.assert_not_awaited() + logging_obj.failure_handler.assert_called_once() + logging_obj.async_failure_handler.assert_awaited_once() + tool_error = logging_obj.async_failure_handler.await_args.args[0] + assert isinstance(tool_error, MCPToolResultError) + assert str(tool_error) == "upstream exploded" + logging_obj.has_run_logging.assert_any_call(event_type="sync_success") + logging_obj.has_run_logging.assert_any_call(event_type="async_success") + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + hook_kwargs = proxy_logging_mock.post_call_failure_hook.await_args.kwargs + assert hook_kwargs["route"] == "/mcp/call_tool" + assert hook_kwargs["original_exception"] is tool_error + assert hook_kwargs["user_api_key_dict"] is user_auth + logging_obj.async_post_mcp_tool_call_hook.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_success_path_unchanged(): + """isError=False must keep today's behavior: success handler fires, no + failure logging, no post_call_failure_hook.""" + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + logging_obj = _mock_mcp_logging_obj() + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + result = _call_tool_result(False, "all good") + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=result, + start_time=datetime.now(), + end_time=datetime.now(), + user_api_key_auth=UserAPIKeyAuth(api_key="test-key", user_id="test-user"), + request_data={}, + ) + + logging_obj.async_success_handler.assert_awaited_once() + assert logging_obj.async_success_handler.await_args.kwargs["result"] is result + logging_obj.async_failure_handler.assert_not_awaited() + logging_obj.failure_handler.assert_not_called() + proxy_logging_mock.post_call_failure_hook.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_iserror_without_auth_skips_failure_hook(): + """Without a UserAPIKeyAuth the failure handlers still fire but the proxy + post_call_failure_hook (which requires one) is skipped.""" + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + logging_obj = _mock_mcp_logging_obj() + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result={"isError": True, "content": [{"type": "text", "text": "denied"}]}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + logging_obj.async_success_handler.assert_not_awaited() + logging_obj.async_failure_handler.assert_awaited_once() + assert str(logging_obj.async_failure_handler.await_args.args[0]) == "denied" + proxy_logging_mock.post_call_failure_hook.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_strips_credentials_from_failure_hook(): + """Credential-bearing request_data fields (raw request headers, upstream MCP + auth headers, OAuth tokens) must never reach post_call_failure_hook + callbacks; non-credential fields must survive untouched.""" + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + logging_obj = _mock_mcp_logging_obj() + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + request_data = { + "name": "explode", + "litellm_call_id": "cid", + "raw_headers": {"authorization": "Bearer sk-caller-secret"}, + "mcp_auth_header": "upstream-secret", + "mcp_server_auth_headers": {"srv": {"authorization": "Bearer srv-secret"}}, + "oauth2_headers": {"authorization": "Bearer oauth-secret"}, + "user_api_key_auth": user_auth, + } + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_mock): + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=_call_tool_result(True, "boom"), + start_time=datetime.now(), + end_time=datetime.now(), + user_api_key_auth=user_auth, + request_data=request_data, + ) + + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + hook_request_data = proxy_logging_mock.post_call_failure_hook.await_args.kwargs["request_data"] + assert hook_request_data == {"name": "explode", "litellm_call_id": "cid"} + assert "secret" not in str(hook_request_data) + + +def _real_mcp_logging_obj(call_id: str): + from litellm.litellm_core_utils.litellm_logging import Logging + + start_time = datetime.now() + logging_obj = Logging( + model="MCP: weather/get_forecast", + messages=[{"role": "user", "content": "tool call"}], + stream=False, + call_type="call_mcp_tool", + start_time=start_time, + litellm_call_id=call_id, + function_id="test-fn", + ) + logging_obj.update_environment_variables( + model="MCP: weather/get_forecast", + user="", + optional_params={}, + litellm_params={"api_base": ""}, + ) + logging_obj.model_call_details["mcp_tool_call_metadata"] = { + "name": "get_forecast", + "arguments": {"city": "Paris"}, + "mcp_server_name": "weather", + } + return logging_obj, start_time + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_iserror_builds_failure_payload(monkeypatch): + """The standard logging payload for an isError=True result must carry + status='failure' with the tool's error text, so OTel (whose _parse_error + keys off status) marks the MCP span ERROR.""" + import litellm + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + + logging_obj, start_time = _real_mcp_logging_obj("test-mcp-iserror-payload") + + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=_call_tool_result(True, "upstream exploded"), + start_time=start_time, + end_time=datetime.now(), + ) + + payload = logging_obj.model_call_details["standard_logging_object"] + assert payload["status"] == "failure" + assert payload["error_str"] == "upstream exploded" + assert payload["error_information"]["error_class"] == "MCPToolResultError" + assert payload["metadata"]["mcp_tool_call_metadata"]["name"] == "get_forecast" + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_success_builds_success_payload(monkeypatch): + """isError=False still produces a status='success' payload.""" + import litellm + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + + logging_obj, start_time = _real_mcp_logging_obj("test-mcp-success-payload") + + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=_call_tool_result(False, "all good"), + start_time=start_time, + end_time=datetime.now(), + ) + + payload = logging_obj.model_call_details["standard_logging_object"] + assert payload["status"] == "success" + + +@pytest.mark.asyncio +async def test_fire_mcp_tool_call_logging_iserror_emits_otel_error_span(monkeypatch): + """End-to-end regression for the OTel symptom: an isError=True tool + result must reach OTel as an MCP span with StatusCode.ERROR and the tool's + error message, while isError=False stays non-error.""" + pytest.importorskip("opentelemetry") + from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, + ) + from opentelemetry.trace.status import StatusCode + + import litellm + from litellm.integrations.otel import OpenTelemetryV2Config + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.plumbing import providers + from litellm.proxy._experimental.mcp_server.server import ( + _fire_mcp_tool_call_logging, + ) + + cfg = OpenTelemetryV2Config(exporter="in_memory", legacy_compat=False) + exporter = InMemorySpanExporter() + tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter) + otel_logger = OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider) + + monkeypatch.setattr(litellm, "failure_callback", []) + monkeypatch.setattr(litellm, "_async_failure_callback", [otel_logger]) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", [otel_logger]) + + logging_obj, start_time = _real_mcp_logging_obj("test-mcp-iserror-otel") + + await _fire_mcp_tool_call_logging( + logging_obj=logging_obj, + result=_call_tool_result(True, "upstream exploded"), + start_time=start_time, + end_time=datetime.now(), + ) + + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call get_forecast" + assert span.status.status_code is StatusCode.ERROR + assert span.attributes["error.type"] == "MCPToolResultError" + assert "upstream exploded" in (span.status.description or "") + + @pytest.mark.asyncio async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): """ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index 5c2a04456b0..b8f0b205831 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -433,7 +433,7 @@ class TestCallToolRestApiVirtualTools: return_value=fake_result, ) as mock_execute, patch( - "litellm.proxy._experimental.mcp_server.rest_endpoints._fire_mcp_success_logging", + "litellm.proxy._experimental.mcp_server.rest_endpoints._fire_mcp_tool_call_logging", new_callable=AsyncMock, side_effect=RuntimeError("logging failed"), ) as mock_fire_logging, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index e114f46e866..3d9afd8f250 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1559,7 +1559,7 @@ class TestCallToolRestAPI: fire_logging = AsyncMock(side_effect=RuntimeError("logging failed")) monkeypatch.setattr( rest_endpoints, - "_fire_mcp_success_logging", + "_fire_mcp_tool_call_logging", fire_logging, raising=False, ) @@ -1590,13 +1590,13 @@ class TestCallToolRestAPI: fire_logging = AsyncMock(side_effect=asyncio.CancelledError()) monkeypatch.setattr( rest_endpoints, - "_fire_mcp_success_logging", + "_fire_mcp_tool_call_logging", fire_logging, raising=False, ) with pytest.raises(asyncio.CancelledError): - await rest_endpoints._safe_fire_mcp_success_logging( + await rest_endpoints._safe_fire_mcp_tool_call_logging( object(), {"result": "ok"}, datetime.now(), datetime.now() ) From c2d8a17692cb4dbaacccf9ecb3d678a8e4788db8 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 8 Jul 2026 10:01:41 -0700 Subject: [PATCH 21/31] test(responses): replace perma-skip azure shell e2e with offline coverage (#32444) --- .../base_responses_api.py | 9 +- .../test_azure_responses_api.py | 5 - .../azure_shell_tool.json | 14 ++ .../test_responses_api_request_body.py | 150 ++++++++++++++---- 4 files changed, 142 insertions(+), 36 deletions(-) create mode 100644 tests/test_litellm/expected_responses_api_request/azure_shell_tool.json diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index 7d2e30f8372..407091a65b3 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -746,7 +746,8 @@ class BaseResponsesAPITest(ABC): E2E test for Shell tool on OpenAI Responses API. Passes tools=[{"type": "shell", "environment": {"type": "container_auto"}}]; validates that the request is accepted and returns a valid response. - Only runs for OpenAI/Azure (Responses API with shell support). + Only runs for OpenAI; offline coverage for the Azure route lives in + tests/test_litellm/responses/test_responses_api_request_body.py. """ base_completion_call_args = self.get_base_completion_call_args() model = ( @@ -754,8 +755,10 @@ class BaseResponsesAPITest(ABC): or base_completion_call_args.get("model") or "" ) - if "openai/" not in str(model) and "azure/" not in str(model): - pytest.skip("Shell tool e2e is only run for OpenAI/Azure Responses API") + if "openai/" not in str(model): + pytest.skip( + "Shell tool e2e is OpenAI-only; no Azure deployment supports the shell tool yet, re-enable once one exists" + ) tools = [{"type": "shell", "environment": {"type": "container_auto"}}] input_msg = "List files in /mnt/data and show python --version." try: diff --git a/tests/llm_responses_api_testing/test_azure_responses_api.py b/tests/llm_responses_api_testing/test_azure_responses_api.py index fed9e9e11f0..ccef8cbf1e7 100644 --- a/tests/llm_responses_api_testing/test_azure_responses_api.py +++ b/tests/llm_responses_api_testing/test_azure_responses_api.py @@ -2,7 +2,6 @@ import os import sys import pytest import asyncio -from typing import Optional from unittest.mock import patch, AsyncMock sys.path.insert(0, os.path.abspath("../..")) @@ -30,10 +29,6 @@ class TestAzureResponsesAPITest(BaseResponsesAPITest): "api_version": "2025-03-01-preview", } - def get_advanced_model_for_shell_tool(self) -> Optional[str]: - """If specified, overrides the model used by test_responses_api_shell_tool_streaming_sees_shell_output (e.g. openai/gpt-5.2 for shell support).""" - return "azure/gpt-5-mini" - @pytest.mark.asyncio async def test_azure_responses_api_preview_api_version(): diff --git a/tests/test_litellm/expected_responses_api_request/azure_shell_tool.json b/tests/test_litellm/expected_responses_api_request/azure_shell_tool.json new file mode 100644 index 00000000000..b716c518106 --- /dev/null +++ b/tests/test_litellm/expected_responses_api_request/azure_shell_tool.json @@ -0,0 +1,14 @@ +{ + "model": "gpt-5-mini", + "input": "List files in /mnt/data and run python --version.", + "tools": [ + { + "type": "shell", + "environment": { + "type": "container_auto" + } + } + ], + "tool_choice": "auto", + "max_output_tokens": 256 +} diff --git a/tests/test_litellm/responses/test_responses_api_request_body.py b/tests/test_litellm/responses/test_responses_api_request_body.py index e312a11e893..c39ba75bd97 100644 --- a/tests/test_litellm/responses/test_responses_api_request_body.py +++ b/tests/test_litellm/responses/test_responses_api_request_body.py @@ -1,6 +1,7 @@ """ Test that litellm.responses() / litellm.aresponses() send the expected request body -over the wire. Expected JSON bodies are stored in expected_responses_api_request/. +over the wire and surface provider errors correctly. Expected JSON bodies are stored +in expected_responses_api_request/. """ import json @@ -18,24 +19,20 @@ def _expected_dir() -> Path: return Path(__file__).resolve().parent.parent / "expected_responses_api_request" -@pytest.mark.asyncio -async def test_aresponses_context_management_and_shell_request_body_matches_expected(): - """ - Call litellm.aresponses() with context_management and shell tool; - assert the httpx POST request body matches the expected JSON. - """ - expected_path = _expected_dir() / "context_management_and_shell.json" +def _load_expected_body(filename: str) -> dict: + expected_path = _expected_dir() / filename assert expected_path.exists(), f"Expected file not found: {expected_path}" with open(expected_path) as f: - expected_body = json.load(f) + return json.load(f) - # Minimal Responses API response so parsing succeeds - mock_response = { - "id": "resp_ctx_shell_test", + +def _minimal_responses_api_payload(response_id: str, model: str) -> dict: + return { + "id": response_id, "object": "response", "created_at": 1734366691, "status": "completed", - "model": "gpt-4o", + "model": model, "output": [ { "type": "message", @@ -69,21 +66,41 @@ async def test_aresponses_context_management_and_shell_request_body_matches_expe "user": None, } - class MockResponse: - def __init__(self, json_data, status_code=200): - self._json_data = json_data - self.status_code = status_code - self.text = json.dumps(json_data) - self.headers = httpx.Headers({}) - def json(self): - return self._json_data +class MockResponse: + def __init__(self, json_data, status_code=200): + self._json_data = json_data + self.status_code = status_code + self.text = json.dumps(json_data) + self.headers = httpx.Headers({}) + + def json(self): + return self._json_data + + +def _assert_request_body_matches(request_body: dict, expected_body: dict) -> None: + for key, expected_value in expected_body.items(): + assert key in request_body, f"Missing key in request body: {key}" + assert ( + request_body[key] == expected_value + ), f"Mismatch for key {key}: got {request_body[key]!r}, expected {expected_value!r}" + + +@pytest.mark.asyncio +async def test_aresponses_context_management_and_shell_request_body_matches_expected(): + """ + Call litellm.aresponses() with context_management and shell tool; + assert the httpx POST request body matches the expected JSON. + """ + expected_body = _load_expected_body("context_management_and_shell.json") with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", new_callable=AsyncMock, ) as mock_post: - mock_post.return_value = MockResponse(mock_response, 200) + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_ctx_shell_test", "gpt-4o"), 200 + ) await litellm.aresponses( model="openai/gpt-4o", @@ -95,10 +112,87 @@ async def test_aresponses_context_management_and_shell_request_body_matches_expe ) mock_post.assert_called_once() - request_body = mock_post.call_args.kwargs["json"] + _assert_request_body_matches(mock_post.call_args.kwargs["json"], expected_body) - for key, expected_value in expected_body.items(): - assert key in request_body, f"Missing key in request body: {key}" - assert ( - request_body[key] == expected_value - ), f"Mismatch for key {key}: got {request_body[key]!r}, expected {expected_value!r}" + +@pytest.mark.asyncio +async def test_aresponses_azure_shell_tool_request_body_matches_expected(): + """ + Call litellm.aresponses() on the Azure route with the shell tool; + assert the httpx POST request body carries the shell tool verbatim. + """ + expected_body = _load_expected_body("azure_shell_tool.json") + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = MockResponse( + _minimal_responses_api_payload("resp_azure_shell_test", "gpt-5-mini"), 200 + ) + + await litellm.aresponses( + model="azure/gpt-5-mini", + api_base="https://fake-resource.openai.azure.com", + api_key="fake-api-key", + api_version="2025-03-01-preview", + input=expected_body["input"], + tools=expected_body["tools"], + tool_choice=expected_body["tool_choice"], + max_output_tokens=expected_body["max_output_tokens"], + ) + + mock_post.assert_called_once() + _assert_request_body_matches(mock_post.call_args.kwargs["json"], expected_body) + + +@pytest.mark.asyncio +async def test_aresponses_azure_shell_tool_400_maps_to_bad_request_error(): + """ + Azure rejects the shell tool for unsupported deployments with a 400; + litellm must surface that as litellm.BadRequestError carrying the provider message. + """ + error_body = { + "error": { + "message": "Tool of type 'shell' is not supported with this model.", + "type": "invalid_request_error", + "param": "tools", + "code": None, + } + } + + def _raise_azure_400(*args, **kwargs): + response = httpx.Response( + status_code=400, + json=error_body, + request=httpx.Request( + "POST", + kwargs.get( + "url", + "https://fake-resource.openai.azure.com/openai/responses", + ), + ), + ) + response.raise_for_status() + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.side_effect = _raise_azure_400 + + with pytest.raises(litellm.BadRequestError) as excinfo: + await litellm.aresponses( + model="azure/gpt-5-mini", + api_base="https://fake-resource.openai.azure.com", + api_key="fake-api-key", + api_version="2025-03-01-preview", + input="List files in /mnt/data and run python --version.", + tools=[{"type": "shell", "environment": {"type": "container_auto"}}], + tool_choice="auto", + max_output_tokens=256, + ) + + assert excinfo.value.status_code == 400 + assert "shell" in str(excinfo.value).lower() + assert "not supported" in str(excinfo.value).lower() From f982b67d78c65d335144d54b8c7c831fcab903f3 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 8 Jul 2026 10:36:00 -0700 Subject: [PATCH 22/31] fix(proxy): harden secret name validation for external secret manager integrations (LIT-4201) (#32092) key_alias can become the secret name used by external secret manager integrations (HashiCorp Vault, CyberArk Conjur) when store_virtual_keys is enabled. Add raise_if_unsafe_secret_name, a shared validation check applied unconditionally before a secret name reaches either integration or the /key/generate, /key/update, and /key/regenerate API boundary, independent of the existing enable_key_alias_format_validation opt-in flag. Also hardens the Vault URL builder to percent-encode reserved characters in secret_name (preserving "/" and "@"), and switches the Conjur policy body to a real YAML serializer instead of raw string interpolation. --- .../key_management_endpoints.py | 24 +++++- .../secret_managers/base_secret_manager.py | 15 ++++ .../cyberark_secret_manager.py | 8 +- .../hashicorp_secret_manager.py | 3 +- tests/litellm_utils_tests/test_cyberark.py | 77 +++++++++++++++++++ tests/litellm_utils_tests/test_hashicorp.py | 27 +++++++ .../test_key_management_endpoints.py | 29 ++++++- .../test_base_secret_manager.py | 59 ++++++++++++++ 8 files changed, 234 insertions(+), 8 deletions(-) create mode 100644 tests/test_litellm/secret_managers/test_base_secret_manager.py diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 63f4b731871..bf64f537c7f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -111,6 +111,7 @@ from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) from litellm.router import Router +from litellm.secret_managers.base_secret_manager import raise_if_unsafe_secret_name from litellm.secret_managers.main import get_secret from litellm.types.proxy.management_endpoints.key_management_endpoints import ( BulkUpdateKeyRequest, @@ -6291,8 +6292,13 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None: """ Validate the format of the key_alias. - Gated behind ``litellm.enable_key_alias_format_validation`` (default **False**). - When disabled, no validation is performed so existing workflows are not broken. + A baseline validation always runs, regardless of + ``litellm.enable_key_alias_format_validation``. + + The remaining charset/length rules are gated behind + ``litellm.enable_key_alias_format_validation`` (default **False**). When disabled, + only the baseline validation above is performed, so existing workflows are not + broken. Rules (when enabled): - None is OK (no alias). @@ -6300,10 +6306,20 @@ def _validate_key_alias_format(key_alias: Optional[str]) -> None: - start/end with alphanumeric - only allow a-zA-Z0-9_-/.@ """ - if not litellm.enable_key_alias_format_validation: + if key_alias is None: return - if key_alias is None: + try: + raise_if_unsafe_secret_name(key_alias) + except ValueError: + raise ProxyException( + message="Invalid key_alias", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=400, + ) + + if not litellm.enable_key_alias_format_validation: return if not _KEY_ALIAS_PATTERN.match(key_alias): diff --git a/litellm/secret_managers/base_secret_manager.py b/litellm/secret_managers/base_secret_manager.py index d33d76093c9..2bb8dc73138 100644 --- a/litellm/secret_managers/base_secret_manager.py +++ b/litellm/secret_managers/base_secret_manager.py @@ -1,3 +1,4 @@ +import re from abc import ABC, abstractmethod from typing import Any, Dict, Optional, Union @@ -5,6 +6,20 @@ import httpx from litellm import verbose_logger +_UNSAFE_SECRET_NAME_PATTERN = re.compile(r"(^|/)\.\.(/|$)|[\x00-\x1f\x7f-\x9f…

]") + + +def raise_if_unsafe_secret_name(secret_name: str) -> None: + """ + Validate a secret name before it is used by a secret manager integration. + + Rejects ".." only as a path segment (bounded by "/" or the start/end of the + string, e.g. "../x", "x/..", or exactly ".."), not as a plain substring, so + names like "release-1.0..2" are not rejected. + """ + if _UNSAFE_SECRET_NAME_PATTERN.search(secret_name): + raise ValueError(f"Invalid secret_name {secret_name!r}") + class BaseSecretManager(ABC): """ diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index faf6224757f..2b888cb85f6 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -4,6 +4,7 @@ from typing import Any, Dict, Optional, Union from urllib.parse import quote import httpx +import yaml import litellm from litellm._logging import verbose_logger @@ -15,7 +16,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._types import KeyManagementSystem -from .base_secret_manager import BaseSecretManager +from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name from .main import str_to_bool @@ -125,8 +126,11 @@ class CyberArkSecretManager(BaseSecretManager): """ # In production, we'd check if the variable exists first # For now, we'll attempt to create it and ignore if it already exists + raise_if_unsafe_secret_name(secret_name) policy_url = f"{self.conjur_addr}/policies/{self.conjur_account}/policy/root" - policy_yaml = f"- !variable {secret_name}\n" + # Use a real YAML serializer to build the scalar safely. + quoted_name = yaml.safe_dump(secret_name, default_style='"').strip() + policy_yaml = f"- !variable {quoted_name}\n" try: client = _get_httpx_client(params={"ssl_verify": self.ssl_verify}) diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index bd1b1097347..039aecb9e58 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -14,7 +14,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._types import KeyManagementSystem -from .base_secret_manager import BaseSecretManager +from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name class HashicorpSecretManager(BaseSecretManager): @@ -220,6 +220,7 @@ class HashicorpSecretManager(BaseSecretManager): - With custom mount: http://127.0.0.1:8200/v1/kv/data/mykey - With path prefix: http://127.0.0.1:8200/v1/secret/data/myapp/mykey """ + raise_if_unsafe_secret_name(secret_name) resolved_namespace = self._sanitize_path_component(namespace if namespace is not None else self.vault_namespace) resolved_mount = self._sanitize_path_component(mount_name if mount_name is not None else self.vault_mount_name) if resolved_mount is None: diff --git a/tests/litellm_utils_tests/test_cyberark.py b/tests/litellm_utils_tests/test_cyberark.py index 67575d3e781..71daf35a265 100644 --- a/tests/litellm_utils_tests/test_cyberark.py +++ b/tests/litellm_utils_tests/test_cyberark.py @@ -5,6 +5,7 @@ Integration test for CyberArk Conjur Secret Manager. import os import sys import pytest +import yaml from dotenv import load_dotenv load_dotenv() @@ -42,6 +43,82 @@ def create_mock_response(status_code: int, text: str = ""): return mock_response +@pytest.mark.asyncio +async def test_cyberark_write_secret_rejects_yaml_injection(): + """ + Regression test: async_write_secret must reject a secret_name that is not + safe to embed in the Conjur policy body, before any HTTP call is made. + """ + with patch("litellm.proxy.proxy_server.premium_user", True): + malicious_secret_name = "foo\n- !grant\n role: !!admin\n member: attacker" + + mock_sync_client = MagicMock() + mock_async_client = AsyncMock() + + with ( + patch( + "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + return_value=mock_sync_client, + ), + patch( + "litellm.secret_managers.cyberark_secret_manager.get_async_httpx_client", + return_value=mock_async_client, + ), + ): + cyberark_manager = CyberArkSecretManager() + + response = await cyberark_manager.async_write_secret( + secret_name=malicious_secret_name, + secret_value="sk-1234", + ) + + assert response["status"] == "error" + assert "Invalid secret_name" in response["message"] + # The malicious policy YAML must never reach the wire. + mock_sync_client.client.post.assert_not_called() + mock_async_client.post.assert_not_called() + + +@pytest.mark.parametrize( + "secret_name", + [ + "foo: bar", + "foo # bar", + "plain-alias", + "team/user@example.com", + ], +) +def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secret_name): + """ + Regression test: _ensure_variable_exists must escape secret_name (not just + denylist-check it) so the policy body always parses back to exactly one + '!variable' scalar node holding the untouched secret_name. + """ + with patch("litellm.proxy.proxy_server.premium_user", True): + captured = {} + + def _capture_post(url, headers=None, content=None): + captured["content"] = content + return create_mock_response(status_code=201, text="") + + mock_sync_client = MagicMock() + mock_sync_client.client.post.side_effect = _capture_post + + with patch( + "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + return_value=mock_sync_client, + ): + cyberark_manager = CyberArkSecretManager() + cyberark_manager._ensure_variable_exists(secret_name) + + policy_yaml = captured["content"] + parsed = yaml.compose(policy_yaml) + assert len(parsed.value) == 1 + node = parsed.value[0] + assert node.tag == "!variable" + assert node.value == secret_name + + @pytest.mark.asyncio async def test_cyberark_write_and_read_secret(): """ diff --git a/tests/litellm_utils_tests/test_hashicorp.py b/tests/litellm_utils_tests/test_hashicorp.py index 3bdf11ea565..9aff7ddc10e 100644 --- a/tests/litellm_utils_tests/test_hashicorp.py +++ b/tests/litellm_utils_tests/test_hashicorp.py @@ -409,6 +409,33 @@ def test_hashicorp_custom_mount_and_prefix(hashicorp_secret_manager): hashicorp_secret_manager.vault_namespace = original_namespace +@pytest.mark.parametrize( + "malicious_secret_name", + [ + "../../../other-app/creds", + "litellm/../../secret", + "foo\nbar", + "foo
bar", + "foo
bar", + "foo\x85bar", + ], +) +def test_hashicorp_get_url_rejects_path_traversal(monkeypatch, malicious_secret_name): + """ + Regression test: get_url must reject an invalid secret_name instead of + building a URL from it. + + Uses monkeypatch + a directly-constructed manager (not the shared + hashicorp_secret_manager fixture) so this runs in CI without real Vault + credentials configured; get_url performs no I/O. + """ + monkeypatch.setenv("HCP_VAULT_TOKEN", "test-token-for-get-url-only") + manager = HashicorpSecretManager() + + with pytest.raises(ValueError): + manager.get_url(malicious_secret_name) + + mock_old_vault_response = { "request_id": "80fafb6a-e96a-4c5b-29fa-ff505ac72201", "lease_id": "", diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index d707421aeb6..2fe7725fd12 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -9023,7 +9023,7 @@ class TestValidateKeyAliasFormat: litellm.enable_key_alias_format_validation = False def test_validation_skipped_when_flag_disabled(self): - """When enable_key_alias_format_validation is False (default), no validation occurs.""" + """When enable_key_alias_format_validation is False (default), no charset/length validation occurs.""" from litellm.proxy.management_endpoints.key_management_endpoints import ( _validate_key_alias_format, ) @@ -9034,6 +9034,33 @@ class TestValidateKeyAliasFormat: _validate_key_alias_format("!invalid!") _validate_key_alias_format("a" * 256) + @pytest.mark.parametrize( + "unsafe_alias", + [ + "../../../other-app/creds", + "litellm/../../secret", + "foo\n- !grant\n role: !!admin\n member: attacker", + "foo\rbar", + "foo\x00bar", + ], + ) + def test_validate_key_alias_format_rejects_traversal_and_control_chars_even_when_flag_disabled( + self, unsafe_alias + ): + """ + Regression test: this check must reject an invalid key_alias unconditionally, + even when enable_key_alias_format_validation (the separate, opt-in charset + rule) is disabled. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_key_alias_format, + ) + + with pytest.raises(ProxyException) as exc: + _validate_key_alias_format(unsafe_alias) + assert str(exc.value.code) == "400" + assert "Invalid key_alias" in str(exc.value.message) + def test_validate_key_alias_format_valid(self): from litellm.proxy.management_endpoints.key_management_endpoints import ( _validate_key_alias_format, diff --git a/tests/test_litellm/secret_managers/test_base_secret_manager.py b/tests/test_litellm/secret_managers/test_base_secret_manager.py new file mode 100644 index 00000000000..cba6a99ab7f --- /dev/null +++ b/tests/test_litellm/secret_managers/test_base_secret_manager.py @@ -0,0 +1,59 @@ +""" +Test raise_if_unsafe_secret_name, the shared guard applied before secret_name +reaches a secret manager backend. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path + +from litellm.secret_managers.base_secret_manager import raise_if_unsafe_secret_name + + +@pytest.mark.parametrize( + "secret_name", + [ + "..", + "../../../other-app/creds", + "litellm/../../secret", + "foo/../bar", + "foo/..", + "../foo", + "foo\nbar", + "foo\rbar", + "foo\x00bar", + "foo\x7fbar", + "foo\x85bar", + "foo
bar", + "foo
bar", + ], +) +def test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks(secret_name): + with pytest.raises(ValueError): + raise_if_unsafe_secret_name(secret_name) + + +@pytest.mark.parametrize( + "secret_name", + [ + "plain-alias", + "my-key-123", + "prod/my-service-key", + "team/user@example.com", + "foo: bar", + "foo # bar", + "foo?evil=1", + "foo#bar", + "a" * 500, + "release-1.0..2", + "my..key", + "..foo", + "foo..", + "v2.0..1-beta", + ], +) +def test_raise_if_unsafe_secret_name_allows_legitimate_aliases(secret_name): + raise_if_unsafe_secret_name(secret_name) From c0327cded4f3631e09adf63ad8b488f1333a7ac1 Mon Sep 17 00:00:00 2001 From: David Katz Date: Wed, 8 Jul 2026 09:39:48 -0400 Subject: [PATCH 23/31] fix(mcp): pair token-endpoint client_secret with the same source as client_id On re-auth against a server with a persisted DCR client, register_client_with_server short-circuits and returns a placeholder client_secret ("dummy") that the browser echoes back to /token. exchange_token_with_server overrode the caller's client_id with the persisted one but still fell back to the caller's secret when the server had none stored, so a persisted public PKCE client (which has no secret) was paired with the literal string "dummy" and the IdP rejected the exchange with 401 on every re-authorization; the proxy surfaced that as a 500. First connects and brand-new servers worked because a real DCR registration ran and no placeholder existed. Resolve the secret from the server whenever the server's client_id wins, so a secretless public client sends no client_secret at all --- .../mcp_server/discoverable_endpoints.py | 6 +- .../mcp_server/test_discoverable_endpoints.py | 59 +++++++++++++++++++ 2 files changed, 64 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index fa1f73cea77..d4dbee37cdc 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -581,8 +581,12 @@ async def exchange_token_with_server( if mcp_server.token_url is None: raise HTTPException(status_code=400, detail="MCP server token url is not set") + # The id and secret must come from the same source. When the server-side client_id wins, + # falling back to the caller's secret pairs the persisted client with a foreign secret; the + # register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a + # persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s. resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id - resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret + resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret try: client_auth = build_token_endpoint_client_auth( auth_method=mcp_server.token_endpoint_auth_method, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c808b17678a..19d030f17c4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4156,3 +4156,62 @@ async def test_store_per_user_token_server_side_skips_invalidate_when_db_write_f invalidate_mock.assert_not_awaited() cache_set_mock.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_token_exchange_pairs_client_secret_with_server_client_id(): + """Re-auth regression: the register short-circuit hands the browser a placeholder + ``client_secret: "dummy"``, which the browser echoes back to /token. The server-side + persisted client_id wins the resolution, so the secret must come from the same (server) + source; pairing the persisted public PKCE client (no stored secret) with the caller's + placeholder makes the IdP reject the exchange with 401 on every re-auth.""" + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="srv-1", + name="srv-1", + server_name="srv-1", + alias="srv-1", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="persisted-client", + client_secret=None, + authorization_url="https://provider.example/oauth/authorize", + token_url="https://provider.example/oauth/token", + ) + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = {"access_token": "at", "token_type": "Bearer"} + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + await exchange_token_with_server( + request=mock_request, + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="https://litellm.example.com/ui/mcp/oauth/callback", + client_id="srv-1", + client_secret="dummy", + code_verifier="verifier", + ) + + sent = mock_async_client.post.call_args.kwargs["data"] + assert sent["client_id"] == "persisted-client" + assert "client_secret" not in sent From ad69d6f3f924a7619deab91d7d3d40f391a29854 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 8 Jul 2026 11:06:56 -0700 Subject: [PATCH 24/31] test: emit e2e coverage lines for loki (#32513) --- tests/e2e/CLAUDE.md | 4 +-- tests/e2e/coverage_registry/README.md | 25 +++++++++++----- tests/e2e/coverage_registry/collector.py | 29 +++++++++++++++---- tests/e2e/coverage_registry/schema.py | 15 ++++++++++ tests/e2e/coverage_registry/test_collector.py | 26 +++++++++++++++++ 5 files changed, 83 insertions(+), 16 deletions(-) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index a4c507ca5ea..502c881c764 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -63,7 +63,7 @@ The harness is fully typed and new code must not add `Any` or widen the basedpyr The set of tests we want is a registry checked into this repo, one row per behavior; that file is the definition of done and the denominator. Each e2e test declares what it covers with `@pytest.mark.covers("...")`, and a small collector diffs the registry against the tests and ships coverage to the existing Grafana. No Allure, no new dependencies -Coverage is organized as module > feature > test. Dashboard modules are Core LLMs, Non-Core LLMs, MCPs, Management/UI, Reliability & Performance, Logging & Guardrails, and Other. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works` +Coverage is organized as module > feature > test. Dashboard modules are `Core LLMs`, `Non-Core LLMs`, `MCPs`, `Management/UI`, `Reliability & Performance`, `Logging & Guardrails`, and `Other`. The Loki stdout formatter maps those display modules to log-safe labels (`core_llms`, `non_core_llms`, `mcp`, `management_ui`, `reliability_performance`, `logging_guardrails`, and `other`) without changing JSON or Prometheus labels. A feature is either an endpoint (`/chat/completions`) or a behavior (fallbacks, rate limits; config-driven, with no route of its own). A cell reads like `llm.chat_completions.bedrock_converse.tool_use.stream.works` The metric is coverage: the share of registry rows that have a passing covering test, reported to Grafana per module so a gap surfaces as an uncovered row rather than a silent absence @@ -71,7 +71,7 @@ Tests do not declare a dashboard module directly. They only declare the registry ### Naming grammar per module -LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix. `chat_completions`, `messages`, and `responses` are Core LLMs. Other LLM endpoints, including `batches` and `realtime`, roll up as Non-Core LLMs. +LLMs - endpoint features (subject = the route), seeded from the Claude Code compat matrix. `chat_completions`, `messages`, and `responses` roll up to `Core LLMs`. Other LLM endpoints, including `batches` and `realtime`, roll up to `Non-Core LLMs`. ``` llm..... diff --git a/tests/e2e/coverage_registry/README.md b/tests/e2e/coverage_registry/README.md index 4177cba7766..863ce34694d 100644 --- a/tests/e2e/coverage_registry/README.md +++ b/tests/e2e/coverage_registry/README.md @@ -9,18 +9,18 @@ note; the naming grammar lives in `tests/e2e/CLAUDE.md`. A **cell** is one customer-noticeable behavior a single e2e test can assert pass/fail on, for example `llm.chat_completions.bedrock_converse.tool_use.stream.works`. Cells are -grouped `module > feature > test`, with LLM cells split into Core LLMs and Non-Core -LLMs for dashboarding. Each cell carries a tier (P0/P1/P2), a source, and a +grouped `module > feature > test`, with LLM cells split into `Core LLMs` and +`Non-Core LLMs` for dashboarding. Each cell carries a tier (P0/P1/P2), a source, and a `fail_before_fix` flag. The rows live in per-prefix YAML files (`llm_*.yaml`, `mgmt.yaml`, `mcp.yaml`, `reliability.yaml`, `logging.yaml`, `guardrail.yaml`, `other.yaml`) and validate against the discriminated union in `schema.py`, so an LLM row cannot carry a guardrail field and vice versa. `llm` rows with `subject_endpoint` of `chat_completions`, `messages`, or -`responses` roll up to "Core LLMs"; all other LLM endpoints roll up to "Non-Core -LLMs". LLM endpoint, route, and capability values are typed in `schema.py`, so new -taxonomy values require an explicit schema change. `logging` and `guardrail` are two -id-prefixes that roll up into the single "Logging & Guardrails" dashboard module. +`responses` roll up to `Core LLMs`; all other LLM endpoints roll up to `Non-Core LLMs`. +LLM endpoint, route, and capability values are typed in `schema.py`, so new taxonomy +values require an explicit schema change. `logging` and `guardrail` are two id-prefixes +that roll up into the single `Logging & Guardrails` dashboard module. A test declares what it covers with a marker: @@ -40,8 +40,17 @@ proxy. Whether a covered cell currently passes or fails is a separate, live conc cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector ``` -Use `--format prometheus` or `--format json` for CI jobs that publish coverage to -Grafana. +Use `--format loki` after the e2e pytest run in the same Kubernetes job/pod to print +structured stdout lines for Loki: + +``` +cd tests/e2e && PYTHONPATH=. python -m coverage_registry.collector --format loki --strict +``` + +This emits exactly one `COVERAGE_TOTAL` line and one `COVERAGE_MODULE` line per module +in `MODULE_ORDER`, in that order. Loki uses log-safe `module=` labels from +`LOKI_MODULE_LABELS` (`core_llms`, `management_ui`, etc.) so existing JSON and +Prometheus consumers keep their human-readable module names unchanged. The headline is overall coverage. The collector also lists markers that point at ids not in the registry, so a typo or an unenumerated behavior surfaces instead of being diff --git a/tests/e2e/coverage_registry/collector.py b/tests/e2e/coverage_registry/collector.py index 3b577106605..f6e59ca4a88 100644 --- a/tests/e2e/coverage_registry/collector.py +++ b/tests/e2e/coverage_registry/collector.py @@ -21,7 +21,7 @@ from pathlib import Path import pytest from .registry import load_registry -from .schema import MODULE_ORDER, Cell, Tier, dashboard_module +from .schema import MODULE_ORDER, Cell, Tier, dashboard_module, loki_module_label E2E_DIR = Path(__file__).resolve().parent.parent @@ -239,13 +239,31 @@ def render_prometheus(report: CoverageReport) -> str: return "\n".join(lines) +def render_loki(report: CoverageReport) -> str: + lines = [ + ( + f"COVERAGE_TOTAL percent={report.coverage_percent:.1f} " + f"covered={report.covered} total={report.total}" + ) + ] + lines.extend( + ( + f"COVERAGE_MODULE module={loki_module_label(module.module)} " + f"percent={module.coverage_percent:.1f} " + f"covered={module.covered} total={module.total}" + ) + for module in report.modules + ) + return "\n".join(lines) + + def main() -> int: parser = ArgumentParser() parser.add_argument( "--format", - choices=("text", "json", "prometheus"), + choices=("text", "json", "prometheus", "loki"), default="text", - help="Output format. Use prometheus or json for Grafana ingestion jobs.", + help="Output format. Use loki for structured stdout lines in the e2e job.", ) parser.add_argument( "--strict", @@ -265,9 +283,8 @@ def main() -> int: "text": render, "json": render_json, "prometheus": render_prometheus, - }[ - args.format - ](report) + "loki": render_loki, + }[args.format](report) print(output) # noqa: T201 # CLI entrypoint output if args.strict and report.orphan_markers: return 1 diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index 2e2a00e78ba..bb27fbf0ea0 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -157,6 +157,16 @@ MODULE_ORDER: tuple[str, ...] = ( "Other", ) +LOKI_MODULE_LABELS: dict[str, str] = { + "Core LLMs": "core_llms", + "Non-Core LLMs": "non_core_llms", + "MCPs": "mcp", + "Management/UI": "management_ui", + "Reliability & Performance": "reliability_performance", + "Logging & Guardrails": "logging_guardrails", + "Other": "other", +} + def dashboard_module(cell: Cell) -> str: """Return the Grafana/reporting module for a registry cell.""" @@ -165,3 +175,8 @@ def dashboard_module(cell: Cell) -> str: return "Core LLMs" return "Non-Core LLMs" return PREFIX_ROLLUP[cell.module] + + +def loki_module_label(module: str) -> str: + """Return the log-safe Loki label for a dashboard module.""" + return LOKI_MODULE_LABELS[module] diff --git a/tests/e2e/coverage_registry/test_collector.py b/tests/e2e/coverage_registry/test_collector.py index 355bc52730d..079ee215866 100644 --- a/tests/e2e/coverage_registry/test_collector.py +++ b/tests/e2e/coverage_registry/test_collector.py @@ -15,6 +15,7 @@ from coverage_registry.collector import ( compute_coverage, render, render_json, + render_loki, render_prometheus, ) from coverage_registry.registry import load_registry @@ -24,6 +25,7 @@ from coverage_registry.schema import ( LlmEndpoint, LoggingCell, Tier, + loki_module_label, ) @@ -149,6 +151,30 @@ def test_prometheus_render_exposes_module_coverage_timeseries() -> None: assert "litellm_e2e_coverage_orphan_markers 0" in metrics +def test_loki_render_exposes_exact_stdout_lines_for_loki() -> None: + report = compute_coverage( + (_llm("llm.chat", Tier.P0), _llm("llm.batches", Tier.P0, "batches")), + frozenset({"llm.chat"}), + ) + + lines = render_loki(report).splitlines() + + assert len(lines) == 1 + len(report.modules) + assert lines[0] == "COVERAGE_TOTAL percent=50.0 covered=1 total=2" + assert ( + lines[1] == "COVERAGE_MODULE module=core_llms percent=100.0 covered=1 total=1" + ) + assert ( + lines[2] == "COVERAGE_MODULE module=non_core_llms percent=0.0 covered=0 total=1" + ) + assert [line.split("module=", 1)[1].split(" ", 1)[0] for line in lines[1:]] == [ + loki_module_label(module.module) for module in report.modules + ] + assert all( + " " not in line.split("module=", 1)[1].split(" ", 1)[0] for line in lines[1:] + ) + + def test_real_registry_loads_and_ids_are_unique() -> None: cells = load_registry() ids = [c.id for c in cells] From 93c047d52eaa665b30931aa4ca9bf7230c0ed74d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 8 Jul 2026 11:07:58 -0700 Subject: [PATCH 25/31] feat(proxy): make Microsoft Graph endpoint configurable for GCC High (LIT-4282) (#32517) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/proxy/management_endpoints/ui_sso.py | 22 ++++-- .../proxy/management_endpoints/test_ui_sso.py | 72 +++++++++++++++++++ 2 files changed, 89 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index dbf514d2298..43fdd3ed05a 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -113,7 +113,7 @@ from litellm.proxy.utils import ( from litellm.repositories.table_repositories import SSOConfigRepository from litellm.repositories.team_repository import TeamRepository from litellm.repositories.user_repository import UserRepository -from litellm.secret_managers.main import get_secret_bool, str_to_bool +from litellm.secret_managers.main import get_secret_bool, get_secret_str, str_to_bool from litellm.types.proxy.management_endpoints.ui_sso import * # noqa: F403 from litellm.types.proxy.management_endpoints.ui_sso import ( DefaultTeamSSOParams, @@ -3737,8 +3737,7 @@ class MicrosoftSSOHandler: Handles Microsoft SSO callback response and returns a CustomOpenID object """ - graph_api_base_url = "https://graph.microsoft.com/v1.0" - graph_api_user_groups_endpoint = f"{graph_api_base_url}/me/memberOf" + DEFAULT_GRAPH_API_BASE_URL = "https://graph.microsoft.com/v1.0" """ Constants @@ -3748,6 +3747,19 @@ class MicrosoftSSOHandler: # used for debugging to show the user groups litellm found from Graph API GRAPH_API_RESPONSE_KEY = "graph_api_user_groups" + @staticmethod + def get_graph_api_base_url() -> str: + """ + Returns the Microsoft Graph API base URL, configurable via the + `MICROSOFT_GRAPH_ENDPOINT` env var so non-default clouds such as Azure + Government (GCC High) can point at `https://graph.microsoft.us/v1.0` + """ + return get_secret_str("MICROSOFT_GRAPH_ENDPOINT") or MicrosoftSSOHandler.DEFAULT_GRAPH_API_BASE_URL + + @staticmethod + def get_graph_api_user_groups_endpoint() -> str: + return f"{MicrosoftSSOHandler.get_graph_api_base_url()}/me/memberOf" + @staticmethod async def get_microsoft_callback_response( request: Request, @@ -3924,7 +3936,7 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[str] = MicrosoftSSOHandler.graph_api_user_groups_endpoint + next_link: Optional[str] = MicrosoftSSOHandler.get_graph_api_user_groups_endpoint() auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 @@ -4007,7 +4019,7 @@ class MicrosoftSSOHandler: Users use Enterprise Applications to manage Groups and Users on Microsoft Entra ID """ - base_url = "https://graph.microsoft.com/v1.0" + base_url = MicrosoftSSOHandler.get_graph_api_base_url() # Endpoint to get app role assignments for the given service principal endpoint = f"/servicePrincipals/{service_principal_id}/appRoleAssignedTo" url = base_url + endpoint diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 32229e3e64e..642f20906a0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -389,6 +389,78 @@ async def test_get_user_groups_error_handling(): assert len(result) == 0 +@pytest.mark.asyncio +async def test_get_user_groups_uses_default_graph_endpoint(monkeypatch): + monkeypatch.delenv("MICROSOFT_GRAPH_ENDPOINT", raising=False) + + requested_urls: list[str] = [] + + async def mock_get(url, *args, **kwargs): + requested_urls.append(url) + mock = MagicMock() + mock.json.return_value = {"value": []} + return mock + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token") + + assert requested_urls == ["https://graph.microsoft.com/v1.0/me/memberOf"] + + +@pytest.mark.asyncio +async def test_get_user_groups_uses_configured_graph_endpoint(monkeypatch): + monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0") + + requested_urls: list[str] = [] + + async def mock_get(url, *args, **kwargs): + requested_urls.append(url) + mock = MagicMock() + mock.json.return_value = {"value": []} + return mock + + with patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client" + ) as mock_client: + mock_client.return_value = MagicMock() + mock_client.return_value.get = mock_get + + await MicrosoftSSOHandler.get_user_groups_from_graph_api(access_token="mock_token") + + assert requested_urls == ["https://graph.microsoft.us/v1.0/me/memberOf"] + + +@pytest.mark.asyncio +async def test_get_group_ids_from_service_principal_uses_configured_graph_endpoint(monkeypatch): + monkeypatch.setenv("MICROSOFT_GRAPH_ENDPOINT", "https://graph.microsoft.us/v1.0") + + requested_urls: list[str] = [] + + async def mock_get(url, *args, **kwargs): + requested_urls.append(url) + mock = MagicMock() + mock.json.return_value = {"value": []} + return mock + + async_client = MagicMock() + async_client.get = mock_get + + await MicrosoftSSOHandler.get_group_ids_from_service_principal( + service_principal_id="sp-123", + async_client=async_client, + access_token="mock_token", + ) + + assert requested_urls == [ + "https://graph.microsoft.us/v1.0/servicePrincipals/sp-123/appRoleAssignedTo" + ] + + def test_get_group_ids_from_graph_api_response(): # Arrange mock_response = MicrosoftGraphAPIUserGroupResponse( From 82fd456b94b36cbfee126e0d30c549b284741a04 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Wed, 8 Jul 2026 11:55:59 -0700 Subject: [PATCH 26/31] Revert "ci: skip unit test workflows when only ui or markdown files change (#32422)" This reverts commit 6df5e1b263a77a25a5bb483015fd13a79f3ef410. --- .github/workflows/test-unit-core-utils.yml | 4 ---- .github/workflows/test-unit-documentation.yml | 4 ---- .github/workflows/test-unit-enterprise-routing.yml | 4 ---- .github/workflows/test-unit-integrations.yml | 4 ---- .github/workflows/test-unit-llm-providers.yml | 4 ---- .github/workflows/test-unit-misc.yml | 4 ---- .github/workflows/test-unit-proxy-auth.yml | 4 ---- .github/workflows/test-unit-proxy-db.yml | 4 ---- .github/workflows/test-unit-proxy-endpoints.yml | 4 ---- .github/workflows/test-unit-proxy-infra.yml | 4 ---- .github/workflows/test-unit-proxy-legacy.yml | 4 ---- .github/workflows/test-unit-responses-caching-types.yml | 4 ---- 12 files changed, 48 deletions(-) diff --git a/.github/workflows/test-unit-core-utils.yml b/.github/workflows/test-unit-core-utils.yml index e563679660b..d6d6353238f 100644 --- a/.github/workflows/test-unit-core-utils.yml +++ b/.github/workflows/test-unit-core-utils.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml index 2c3d6e46618..4cef791a9b3 100644 --- a/.github/workflows/test-unit-documentation.yml +++ b/.github/workflows/test-unit-documentation.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-enterprise-routing.yml b/.github/workflows/test-unit-enterprise-routing.yml index 7a9b8b00f26..13136c968d1 100644 --- a/.github/workflows/test-unit-enterprise-routing.yml +++ b/.github/workflows/test-unit-enterprise-routing.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-integrations.yml b/.github/workflows/test-unit-integrations.yml index b28ba3456ce..c95ed4e7c24 100644 --- a/.github/workflows/test-unit-integrations.yml +++ b/.github/workflows/test-unit-integrations.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-llm-providers.yml b/.github/workflows/test-unit-llm-providers.yml index fecdcbd3b95..df78564ab0c 100644 --- a/.github/workflows/test-unit-llm-providers.yml +++ b/.github/workflows/test-unit-llm-providers.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index dbc3bfc8191..7c3b195f0ad 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-auth.yml b/.github/workflows/test-unit-proxy-auth.yml index ad534cc0098..97dfaed6e81 100644 --- a/.github/workflows/test-unit-proxy-auth.yml +++ b/.github/workflows/test-unit-proxy-auth.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 35a1a9c78a0..2ac9a3b7c1c 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -5,10 +5,6 @@ on: branches: - main - litellm_internal_staging - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index 7eb3d7719c0..cbb36eebdb9 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" workflow_dispatch: permissions: diff --git a/.github/workflows/test-unit-proxy-infra.yml b/.github/workflows/test-unit-proxy-infra.yml index cb944de5cf9..884d62289b9 100644 --- a/.github/workflows/test-unit-proxy-infra.yml +++ b/.github/workflows/test-unit-proxy-infra.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-legacy.yml b/.github/workflows/test-unit-proxy-legacy.yml index 9798a4e2277..8db218cd1fc 100644 --- a/.github/workflows/test-unit-proxy-legacy.yml +++ b/.github/workflows/test-unit-proxy-legacy.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read diff --git a/.github/workflows/test-unit-responses-caching-types.yml b/.github/workflows/test-unit-responses-caching-types.yml index 7331544de24..2f177587997 100644 --- a/.github/workflows/test-unit-responses-caching-types.yml +++ b/.github/workflows/test-unit-responses-caching-types.yml @@ -7,10 +7,6 @@ on: - litellm_internal_staging - litellm_oss_staging - "litellm_**" - paths-ignore: - - "ui/**" - - "**.md" - - "**.mdx" permissions: contents: read From c3dccb54cfd0666393e8874cfd007c72de8b33cc Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 8 Jul 2026 12:32:03 -0700 Subject: [PATCH 27/31] fix(health): bridge litellm_metadata into logging object in _batch_health_check (#32520) * fix(health): bridge litellm_metadata into logging object in _batch_health_check * Update litellm/litellm_core_utils/health_check_helpers.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * fix(health): address review - share single metadata copy, conditional api_base, add tests - Only set api_base in litellm_params when a value actually exists; providers like bedrock/vertex/gemini resolve it implicitly and an empty string overwrites their resolution. - Use a single .copy() for both metadata and litellm_metadata to prevent downstream drift between the two references. - Add 6 unit tests covering metadata bridging, api_base omission, guard conditions, and dispatch routing. Signed-off-by: pramod * refactor(health): use update_from_kwargs helper for metadata bridge Collapses the manual metadata/litellm_metadata plumbing in _batch_health_check into a single update_from_kwargs call, matching how the sibling batch/image/rerank/ocr surfaces bridge metadata onto the pre-injected logging object. Drops the bare Dict typing and the inline comment, and switches the tests to assert against the helper. --------- Signed-off-by: pramod Co-authored-by: pramod Co-authored-by: Pramod B <155433727+BPRMD18@users.noreply.github.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --- .../health_check_helpers.py | 11 ++ .../test_health_check_helpers.py | 137 ++++++++++++++++++ 2 files changed, 148 insertions(+) diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 405366382a1..42ac82abf8b 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -95,6 +95,17 @@ class HealthCheckHelpers: """ import litellm + logging_obj = filtered_model_params.get("litellm_logging_obj") + if logging_obj is not None: + api_base = filtered_model_params.get("api_base") + logging_obj.update_from_kwargs( + kwargs=filtered_model_params, + model=filtered_model_params.get("model"), + user=None, + optional_params={}, + litellm_params={"api_base": api_base} if api_base else None, + ) + if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS: return await litellm.alist_batches(**filtered_model_params) else: diff --git a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py b/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py index 02d72c89e80..e8ef8f15142 100644 --- a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py @@ -14,6 +14,7 @@ from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers from litellm.main import ahealth_check from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS def test_update_model_params_with_health_check_tracking_information(): @@ -140,3 +141,139 @@ async def test_ahealth_check_failure_masks_raw_request_headers(): assert headers["Content-Type"] == "application/json" print(f"Masked Authorization header: {headers.get('Authorization', 'NOT FOUND')}") + + +@pytest.mark.asyncio +async def test_batch_health_check_bridges_metadata_into_logging_obj(): + """_batch_health_check must call update_from_kwargs on the pre-injected + logging object so callbacks receive identity/tracking fields in + model_call_details["litellm_params"]["metadata"].""" + mock_logging_obj = MagicMock() + mock_logging_obj.update_from_kwargs = MagicMock() + + litellm_metadata = { + "tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME], + "user_api_key_alias": "health-check-key", + } + + filtered_model_params = { + "model": "openai/gpt-4", + "api_base": "https://api.openai.com", + "litellm_logging_obj": mock_logging_obj, + "litellm_metadata": litellm_metadata, + } + + with patch("litellm.alist_batches", new_callable=AsyncMock, return_value={}): + await HealthCheckHelpers._batch_health_check( + custom_llm_provider="openai", + model_params={"model": "openai/gpt-4"}, + filtered_model_params=filtered_model_params, + ) + + mock_logging_obj.update_from_kwargs.assert_called_once() + call_kwargs = mock_logging_obj.update_from_kwargs.call_args[1] + assert call_kwargs["model"] == "openai/gpt-4" + assert call_kwargs["kwargs"] is filtered_model_params + assert call_kwargs["litellm_params"] == {"api_base": "https://api.openai.com"} + + +@pytest.mark.asyncio +async def test_batch_health_check_omits_api_base_when_absent(): + """api_base must not appear in litellm_params when the provider resolves + it implicitly (bedrock, vertex, gemini).""" + mock_logging_obj = MagicMock() + mock_logging_obj.update_from_kwargs = MagicMock() + + litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]} + + filtered_model_params = { + "model": "bedrock/anthropic.claude-v2", + "litellm_logging_obj": mock_logging_obj, + "litellm_metadata": litellm_metadata, + } + + with patch("litellm.acompletion", new_callable=AsyncMock, return_value={}): + await HealthCheckHelpers._batch_health_check( + custom_llm_provider="bedrock", + model_params={"model": "bedrock/anthropic.claude-v2"}, + filtered_model_params=filtered_model_params, + ) + + call_kwargs = mock_logging_obj.update_from_kwargs.call_args[1] + assert call_kwargs["litellm_params"] is None + + +@pytest.mark.asyncio +async def test_batch_health_check_skips_bridge_when_no_logging_obj(): + """When litellm_logging_obj is absent, dispatch still proceeds.""" + litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]} + + filtered_model_params = { + "model": "openai/gpt-4", + "litellm_metadata": litellm_metadata, + } + + with patch( + "litellm.alist_batches", new_callable=AsyncMock, return_value={} + ) as mock_alist: + await HealthCheckHelpers._batch_health_check( + custom_llm_provider="openai", + model_params={"model": "openai/gpt-4"}, + filtered_model_params=filtered_model_params, + ) + mock_alist.assert_called_once() + + +@pytest.mark.asyncio +async def test_batch_health_check_uses_alist_batches_for_supported_providers(): + """Providers in LIST_BATCHES_SUPPORTED_PROVIDERS dispatch to alist_batches.""" + mock_logging_obj = MagicMock() + mock_logging_obj.update_from_kwargs = MagicMock() + + litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]} + + for provider in LIST_BATCHES_SUPPORTED_PROVIDERS: + filtered_model_params = { + "model": f"{provider}/some-model", + "litellm_logging_obj": mock_logging_obj, + "litellm_metadata": litellm_metadata, + } + + with patch( + "litellm.alist_batches", new_callable=AsyncMock, return_value={} + ) as mock_alist: + await HealthCheckHelpers._batch_health_check( + custom_llm_provider=provider, + model_params={"model": f"{provider}/some-model"}, + filtered_model_params=filtered_model_params, + ) + mock_alist.assert_called_once() + + +@pytest.mark.asyncio +async def test_batch_health_check_falls_back_to_acompletion_for_unsupported(): + """Providers not in LIST_BATCHES_SUPPORTED_PROVIDERS fall back to acompletion.""" + mock_logging_obj = MagicMock() + mock_logging_obj.update_from_kwargs = MagicMock() + + litellm_metadata = {"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME]} + + filtered_model_params = { + "model": "bedrock/anthropic.claude-v2", + "litellm_logging_obj": mock_logging_obj, + "litellm_metadata": litellm_metadata, + } + + model_params = {"model": "bedrock/anthropic.claude-v2", "messages": []} + + with ( + patch("litellm.alist_batches", new_callable=AsyncMock) as mock_alist, + patch("litellm.acompletion", new_callable=AsyncMock, return_value={}) as mock_acompletion, + ): + await HealthCheckHelpers._batch_health_check( + custom_llm_provider="bedrock", + model_params=model_params, + filtered_model_params=filtered_model_params, + ) + mock_alist.assert_not_called() + mock_acompletion.assert_called_once_with(**model_params) From 85d1fe6e2a535e9edfc1ae0b0854eb204573c7ba Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 8 Jul 2026 13:44:48 -0700 Subject: [PATCH 28/31] fix(otel): restore error.* span attributes on v2 error spans (LIT-4179) (#32524) The v2 emitter has never stamped error.message / error.code / error.stack_trace / error.llm_provider as span attributes; only error.type reached the wire. Backends that flatten span attributes into label indexes (Elastic APM labels.error_*, Datadog span tags) lost these four fields when v2 became the active integration on v1.90+ for otel_v2-flagged deployments. The pre-existing exception span event carrying the full message (LIT-3758) is unchanged; the message now rides both places at once, matching v1s shape. SpanError grows three optional detail fields; _parse_error threads them from StandardLoggingPayloadErrorInformation; the emitters error branch stamps them via a new module-level helper, guarded per field so guardrail-shape errors are not polluted with empty attributes. New semconv constants mirror open_inference.ErrorAttributes byte-for-byte, so v1 and v2 consumers read the same keys. Regression tests extend the mapped test files under tests/test_litellm/integrations/otel/. pytest reports 243 passed. --- litellm/integrations/otel/__init__.py | 2 + litellm/integrations/otel/emitter.py | 35 +++++- litellm/integrations/otel/model/payloads.py | 6 + litellm/integrations/otel/model/semconv.py | 20 ++++ .../otel/test_otel_v2_components.py | 110 ++++++++++++++++-- .../otel/test_otel_v2_sources_of_truth.py | 47 +++++++- 6 files changed, 203 insertions(+), 17 deletions(-) diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 7f78f7156b4..5e167e006ff 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -52,6 +52,7 @@ from litellm.integrations.otel.model.semconv import ( GenAIProvider, JsonRpc, LiteLLM, + LiteLLMError, MCPMethod, Metric, Network, @@ -87,6 +88,7 @@ __all__ = [ "HTTP", "JsonRpc", "LiteLLM", + "LiteLLMError", "MCP", "MCPMethod", "Metric", diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py index 8441cbae834..46aa166a8bb 100644 --- a/litellm/integrations/otel/emitter.py +++ b/litellm/integrations/otel/emitter.py @@ -16,9 +16,10 @@ from litellm.integrations.otel.model.payloads import ( MCPListToolsSpanData, MCPToolCallSpanData, ServiceSpanData, + SpanError, ) from litellm.integrations.otel.plumbing.providers import to_otel_span_kind -from litellm.integrations.otel.model.semconv import Error, ExceptionEvent +from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError from litellm.integrations.otel.model.spans import ( SPAN_REGISTRY, SpanRole, @@ -49,6 +50,27 @@ _NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = { _DEDUP_CACHE_MAX = 10_000 +def _stamp_otel_error_attributes(span: Span, error_type: str, resolved_message: str) -> None: + """Stamp the OTel-semconv error attributes (``error.type`` + ``error.message``). + ``error_type`` and ``resolved_message`` are ``finish_span``'s already-computed + fallback chains, so the pair on the status, event, and attributes stays in + lockstep.""" + span.set_attribute(Error.TYPE, error_type) + span.set_attribute(Error.MESSAGE, resolved_message) + + +def _stamp_litellm_error_attributes(span: Span, error: SpanError) -> None: + """Stamp litellm-specific error detail attributes. Emitted only when the + corresponding field is populated so guardrail-shape errors carrying only a + message aren't polluted with empty detail keys.""" + if error.code: + span.set_attribute(LiteLLMError.CODE, error.code) + if error.stack_trace: + span.set_attribute(LiteLLMError.STACK_TRACE, error.stack_trace) + if error.llm_provider: + span.set_attribute(LiteLLMError.LLM_PROVIDER, error.llm_provider) + + class SpanEmitter: def __init__( self, @@ -190,12 +212,13 @@ class SpanEmitter: if error and (error.error_type or error.message): error_type = error.error_type or "error" message = error.message or error.error_type or "error" - span.set_attribute(Error.TYPE, error_type) + _stamp_otel_error_attributes(span, error_type, message) + _stamp_litellm_error_attributes(span, error) span.set_status(Status(StatusCode.ERROR, message)) - # Carry the full message on the standard ``exception`` event so backends - # map it as full text under ``exception.message``. Setting it as a bare - # string attribute instead lets backends like Elasticsearch dynamic-map - # it to a ``keyword`` capped at 1024 chars, truncating the message. + # Also emit the semconv ``exception`` event so backends that + # dynamic-map unknown string span attrs to ``keyword`` (e.g. + # Elasticsearch with a 1024-char ``ignore_above``) still see the + # full untruncated message on the recognized event field. span.add_event( ExceptionEvent.NAME, {ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message}, diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index fcd710492f0..4a8f01858b5 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -141,6 +141,9 @@ class LLMCost: class SpanError: error_type: str | None = None message: str | None = None + code: str | None = None + stack_trace: str | None = None + llm_provider: str | None = None @dataclass(frozen=True) @@ -571,6 +574,9 @@ def _parse_error(payload: "StandardLoggingPayload") -> SpanError | None: return SpanError( error_type=as_str(info.get("error_class")) or as_str(info.get("error_code")), message=as_str(info.get("error_message")) or as_str(payload.get("error_str")), + code=as_str(info.get("error_code")), + stack_trace=as_str(info.get("traceback")), + llm_provider=as_str(info.get("llm_provider")), ) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 4e725ae0a29..69d1e454655 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -144,7 +144,27 @@ class Client: class Error: + """OTel-defined error attribute keys, from the semconv ``error.*`` registry. + ``MESSAGE`` is marked *Deprecated* upstream in favor of domain-specific + error message keys plus ``exception.message`` on the exception event, but + is still defined and stamped by litellm's v1 integration; keeping it here + for byte-for-byte parity.""" + TYPE: Final = "error.type" + MESSAGE: Final = "error.message" + + +class LiteLLMError: + """LiteLLM-specific error attribute keys. Emitted under the ``error.*`` + namespace (not ``litellm.*``) for byte-for-byte compat with the v1 + integration in ``opentelemetry.py``; consumers reading these keys on v1 + spans read the same keys on v2 spans. OTel semconv does not define any of + these three, and per its extension rules a namespace may carry additional + vendor keys as long as they don't collide with defined names.""" + + CODE: Final = "error.code" + STACK_TRACE: Final = "error.stack_trace" + LLM_PROVIDER: Final = "error.llm_provider" class ExceptionEvent: diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_components.py b/tests/test_litellm/integrations/otel/test_otel_v2_components.py index 19eef284b91..298047ec18b 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -579,13 +579,10 @@ def _exception_event(span): def test_error_message_recorded_as_full_exception_event_untruncated(): - """Regression for the Elasticsearch keyword/ignore_above:1024 truncation. - - A long error message must survive intact on the standard ``exception`` - event under ``exception.message`` — not get dropped onto a bare string - attribute that backends dynamic-map to a 1024-char ``keyword``. The SDK - must not truncate it either, so a 5000-char message stays 5000 chars. - """ + """The ``exception`` event carries the full untruncated message under + ``exception.message`` so backends that dynamic-map unknown string span + attrs to ``keyword`` (e.g. Elasticsearch with a 1024-char ``ignore_above``) + still see it in full via the semconv-recognized event field.""" from litellm.integrations.otel.model.semconv import Error, ExceptionEvent long_message = "boom: " + "x" * 5000 @@ -596,13 +593,108 @@ def test_error_message_recorded_as_full_exception_event_untruncated(): assert len(event.attributes[ExceptionEvent.MESSAGE]) == len(long_message) > 1024 assert event.attributes[ExceptionEvent.TYPE] == "litellm.APIError" - # error.type stays a low-cardinality attribute; the message does NOT become a - # bare string attribute (which is what got truncated). + # error.type stays a low-cardinality attribute; the exception EVENT field + # ``exception.message`` never becomes a bare string attribute. assert span.attributes[Error.TYPE] == "litellm.APIError" assert ExceptionEvent.MESSAGE not in span.attributes assert span.status.description == long_message +def test_error_details_stamped_as_span_attributes_for_labels_ingest(): + """OTel-defined keys and litellm-specific detail keys both ride span + attributes so backends that flatten attrs into label indexes (Elastic APM + ``labels.*``, Datadog span tags) render them. The exception event with the + full untruncated message stays alongside — both places, matching v1's + shape.""" + from litellm.integrations.otel.model.semconv import Error, ExceptionEvent, LiteLLMError + from litellm.integrations.otel.emitter import SpanEmitter + + cfg = OpenTelemetryV2Config(exporter="in_memory") + provider, exporter = providers.in_memory_provider(cfg) + engine = SpanEmitter(providers.get_tracer(provider, "t"), cfg) + data = LLMCallSpanData( + operation=GenAIOperation.CHAT, + provider="openai", + request_model="gpt-4o", + response_model=None, + response_id=None, + request_params=LLMRequestParams(), + usage=LLMUsage(), + finish_reasons=(), + error=SpanError( + error_type="litellm.BadRequestError", + message="400: violated moderation policy", + code="400", + stack_trace="File proxy_server.py line 8570 ...", + llm_provider="openai", + ), + response_cost=None, + server=None, + identity=RequestIdentity(call_id=None), + ) + engine.emit(SpanRole.LLM_CALL, data) + (span,) = exporter.get_finished_spans() + + # OTel-defined keys (from the ``error.*`` semconv registry). + assert span.attributes[Error.TYPE] == "litellm.BadRequestError" + assert span.attributes[Error.MESSAGE] == "400: violated moderation policy" + # LiteLLM-specific detail keys — vendor-namespaced under ``error.*`` + # for v1-parity, not defined by OTel semconv. + assert span.attributes[LiteLLMError.CODE] == "400" + assert span.attributes[LiteLLMError.STACK_TRACE] == "File proxy_server.py line 8570 ..." + assert span.attributes[LiteLLMError.LLM_PROVIDER] == "openai" + + # The exception event carries the same message on the span too. + event = _exception_event(span) + assert event.attributes[ExceptionEvent.MESSAGE] == "400: violated moderation policy" + + +def test_error_details_omitted_when_span_error_carries_only_message(): + """A guardrail-shape error (message only, no code/traceback/provider) must + not pollute the span with empty-string detail attributes. Only the keys + that carry real data land.""" + from litellm.integrations.otel.model.semconv import Error, LiteLLMError + + span = _emit_error_span("guardrail rejected", error_type="ContentFilter") + + assert span.attributes[Error.TYPE] == "ContentFilter" + assert span.attributes[Error.MESSAGE] == "guardrail rejected" + # LiteLLM-specific detail keys aren't stamped when the SpanError doesn't + # carry them. + assert LiteLLMError.CODE not in span.attributes + assert LiteLLMError.STACK_TRACE not in span.attributes + assert LiteLLMError.LLM_PROVIDER not in span.attributes + + +def test_v2_error_attribute_keys_match_v1_error_attributes_byte_for_byte(): + """v1 (``opentelemetry.py``) and v2 (``otel/`` package) stamp identical + span-attribute keys so consumers reading ``labels.error_message`` don't + care which integration produced the span. Renaming either side is a + breaking change for downstream dashboards; this test locks the vocabulary.""" + from litellm.integrations._types.open_inference import ErrorAttributes + from litellm.integrations.otel.model.semconv import Error, LiteLLMError + + assert Error.TYPE == ErrorAttributes.ERROR_TYPE + assert Error.MESSAGE == ErrorAttributes.ERROR_MESSAGE + assert LiteLLMError.CODE == ErrorAttributes.ERROR_CODE + assert LiteLLMError.STACK_TRACE == ErrorAttributes.ERROR_STACK_TRACE + assert LiteLLMError.LLM_PROVIDER == ErrorAttributes.ERROR_LLM_PROVIDER + + +def test_error_message_falls_back_to_error_type_when_message_absent(): + """A ``SpanError(error_type=..., message=None)`` still renders on the span: + the resolved message is the error_type, and it lands on ``error.message``, + the exception event, and the span-status description in lockstep so a + single-source-of-truth view isn't inconsistent.""" + from litellm.integrations.otel.model.semconv import Error, ExceptionEvent + + span = _emit_error_span(message=None, error_type="RateLimitError") + + assert span.attributes[Error.MESSAGE] == "RateLimitError" + assert _exception_event(span).attributes[ExceptionEvent.MESSAGE] == "RateLimitError" + assert span.status.description == "RateLimitError" + + def test_success_span_records_no_exception_event(): from litellm.integrations.otel.emitter import SpanEmitter from litellm.integrations.otel.model.semconv import ExceptionEvent diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 834a484090f..89aa73a6066 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -144,11 +144,13 @@ def _all_constants(cls): def test_attribute_keys_are_unique_across_namespaces(): - from litellm.integrations.otel import MCP, Client, JsonRpc, Network + from litellm.integrations.otel import MCP, Client, JsonRpc, LiteLLMError, Network # prefixes are allowed to be substrings; exact keys must not collide. + # ``LiteLLMError`` shares the ``error.*`` prefix with ``Error`` by design + # (v1-parity); the assert below is the guarantee they never overlap. exact = set() - for cls in (GenAI, Error, Server, HTTP, DB, MCP, JsonRpc, Network, Client): + for cls in (GenAI, Error, LiteLLMError, Server, HTTP, DB, MCP, JsonRpc, Network, Client): for key in _all_constants(cls): assert key not in exact, f"duplicate attribute key {key}" exact.add(key) @@ -342,6 +344,47 @@ def test_llm_call_adapter_failure_path(): assert data.error.message == "429 slow down" +def test_llm_call_adapter_carries_error_detail_fields(): + """``_parse_error`` threads the full detail set from ``error_information`` + (``error_code``, ``traceback``, ``llm_provider``) onto ``SpanError`` so the + emitter can stamp them as span attributes.""" + payload = _sample_payload( + status="failure", + error_information={ + "error_class": "BadRequestError", + "error_message": "400 violated moderation policy", + "error_code": "400", + "traceback": "File proxy_server.py line 8570 ...", + "llm_provider": "openai", + }, + ) + data = LLMCallSpanData.from_standard_logging_payload(payload) + assert data.error is not None + assert data.error.error_type == "BadRequestError" + assert data.error.message == "400 violated moderation policy" + assert data.error.code == "400" + assert data.error.stack_trace == "File proxy_server.py line 8570 ..." + assert data.error.llm_provider == "openai" + + +def test_llm_call_adapter_error_details_default_to_none_when_absent(): + """Guardrail-shape payloads carry only ``error_class`` + ``error_message``. + The detail fields must stay ``None`` so the emitter's ``if error.code:`` + guards skip stamping empty attributes.""" + payload = _sample_payload( + status="failure", + error_information={ + "error_class": "ContentFilter", + "error_message": "guardrail rejected", + }, + ) + data = LLMCallSpanData.from_standard_logging_payload(payload) + assert data.error is not None + assert data.error.code is None + assert data.error.stack_trace is None + assert data.error.llm_provider is None + + def test_adapter_is_resilient_to_minimal_payload(): data = LLMCallSpanData.from_standard_logging_payload({}) assert data.request_model == "" From e9e30dffb68264e497d846763853e4f1c96939e7 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 8 Jul 2026 13:47:53 -0700 Subject: [PATCH 29/31] refactor(ui): reskin shared DataTable from tremor onto shadcn table primitives (#32209) * test(ui): characterize DataTable behavior before shadcn reskin Pins the shared view_logs DataTable contract with library-agnostic queries ahead of the tremor-to-shadcn table migration: loading and empty states, TanStack column defs with custom cell renderers, onRowClick payload, both expansion render paths (colspan sub-component and sibling child rows), the getRowCanExpand gate, and client-side sorting on and off. These must pass unchanged after the reskin. * refactor(ui): reskin shared DataTable from tremor onto shadcn table primitives Swaps the view_logs DataTable's presentational layer from @tremor/react to the in-repo components/ui/table primitives and hardens the seam that every later table migration copies: - getRowId is injected instead of hardcoded to request_id through an any cast; identity defaults to the row index and the logs page now passes request_id explicitly, keeping expansion state attached to the right row across refetch reorders - one expansion render path: renderChildRows had zero consumers and is removed; renderSubComponent (colspan cell) is the single path - the four consumers passing dead no-op renderSubComponent and getRowCanExpand boilerplate drop it - loading and empty defaults become generic (Loading... / No results) instead of log-specific The characterization tests from the previous commit pass unchanged except the dead child-rows path test, replaced by a reorder-stability test for injected getRowId plus coverage of the new generic defaults. First tremor removal of the tables track; view_logs/table.tsx no longer imports @tremor/react. * test(ui): assert child rows hidden before expansion in DataTable test * fix(ui): suppress row hover on DataTable placeholder rows * feat(ui): polish DataTable with skeleton loading, header band, and numeric column alignment * feat(ui): shape DataTable skeletons per column and keep stale rows during refetch * revert(ui): drop DataTable skeleton loading, restore text loading row * fix(ui): clip DataTable to its rounded wrapper and right-align Duration/TTFT values --- ui/litellm-dashboard/eslint-metrics.json | 2 +- ui/litellm-dashboard/eslint-suppressions.json | 5 -- .../components/EntityUsage/TopKeyView.tsx | 9 +-- .../components/EntityUsage/TopModelView.tsx | 12 ++- .../components/mcp_tools/MCPToolsetsTab.tsx | 2 - .../src/components/pass_through_settings.tsx | 2 - .../src/components/view_logs/columns.tsx | 10 ++- .../src/components/view_logs/index.tsx | 1 + .../src/components/view_logs/table.test.tsx | 81 ++++++++++++++++--- .../src/components/view_logs/table.tsx | 78 +++++++++--------- 10 files changed, 124 insertions(+), 78 deletions(-) diff --git a/ui/litellm-dashboard/eslint-metrics.json b/ui/litellm-dashboard/eslint-metrics.json index 219cb0580e7..51cef1169f9 100644 --- a/ui/litellm-dashboard/eslint-metrics.json +++ b/ui/litellm-dashboard/eslint-metrics.json @@ -1,5 +1,5 @@ { - "@typescript-eslint/no-explicit-any": 1990, + "@typescript-eslint/no-explicit-any": 1988, "complexity": 128, "max-depth": 59, "no-console": 15 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 1c8f92b720f..67b19471aaf 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2077,11 +2077,6 @@ "count": 1 } }, - "src/components/view_logs/table.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/view_user_spend.tsx": { "react-hooks/set-state-in-effect": { "count": 2 diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopKeyView.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopKeyView.tsx index d0748583c30..40bc41b3e8c 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopKeyView.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopKeyView.tsx @@ -164,6 +164,7 @@ const TopKeyView: React.FC = ({ topKeys, teams, showTags = fals const spendColumn = { header: "Spend (USD)", accessorKey: "spend", + meta: { numeric: true }, cell: (info: any) => { const value = info.getValue(); return value > 0 && value < 0.01 ? "<$0.01" : `$${formatNumberWithCommas(value, 2)}`; @@ -247,13 +248,7 @@ const TopKeyView: React.FC = ({ topKeys, teams, showTags = fals ) : (
- <>} - getRowCanExpand={() => false} - isLoading={false} - /> +
)} diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopModelView.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopModelView.tsx index c69ba42f182..7562ef06a03 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopModelView.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/EntityUsage/TopModelView.tsx @@ -30,6 +30,7 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi { header: "Spend (USD)", accessorKey: "spend", + meta: { numeric: true }, cell: (info: any) => { const value = info.getValue(); return `$${formatNumberWithCommas(value, 2)}`; @@ -38,16 +39,19 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi { header: "Successful", accessorKey: "successful_requests", + meta: { numeric: true }, cell: (info: any) => {info.getValue()?.toLocaleString() || 0}, }, { header: "Failed", accessorKey: "failed_requests", + meta: { numeric: true }, cell: (info: any) => {info.getValue()?.toLocaleString() || 0}, }, { header: "Tokens", accessorKey: "tokens", + meta: { numeric: true }, cell: (info: any) => info.getValue()?.toLocaleString() || 0, }, ]; @@ -99,13 +103,7 @@ export default function TopModelView({ topModels, topModelsLimit, setTopModelsLi ) : (
- <>} - getRowCanExpand={() => false} - isLoading={false} - /> +
)} diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx index 0198b830229..546df9ebc4b 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPToolsetsTab.tsx @@ -509,8 +509,6 @@ export function MCPToolsetsTab({ accessToken, userRole }: MCPToolsetsTabProps) {
} - getRowCanExpand={() => false} isLoading={isLoading} noDataMessage="No toolsets yet. Click 'New Toolset' to create one." loadingMessage="Loading toolsets..." diff --git a/ui/litellm-dashboard/src/components/pass_through_settings.tsx b/ui/litellm-dashboard/src/components/pass_through_settings.tsx index 0fdd9c632bf..63fe0f92961 100644 --- a/ui/litellm-dashboard/src/components/pass_through_settings.tsx +++ b/ui/litellm-dashboard/src/components/pass_through_settings.tsx @@ -263,8 +263,6 @@ const PassThroughSettings: React.FC = ({
} - getRowCanExpand={() => false} isLoading={false} noDataMessage="No pass-through endpoints configured" /> diff --git a/ui/litellm-dashboard/src/components/view_logs/columns.tsx b/ui/litellm-dashboard/src/components/view_logs/columns.tsx index 0508c562df8..7452992ed59 100644 --- a/ui/litellm-dashboard/src/components/view_logs/columns.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/columns.tsx @@ -231,13 +231,14 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] : "Cost", accessorKey: "spend", size: 110, + meta: { numeric: true }, cell: (info: any) => { const row = info.row.original; const mcpCount = row.mcp_tool_call_count || 0; const mcpSpend = row.mcp_tool_call_spend || 0; return ( -
+
{getSpendString(info.getValue() || 0)} @@ -263,13 +264,14 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] ) : "Duration (s)", accessorKey: "request_duration_ms", + meta: { numeric: true }, cell: (info: any) => { const ms = info.getValue(); if (ms == null) return -; const seconds = (ms / 1000).toFixed(2); return ( - {seconds} + {seconds} ); }, @@ -287,6 +289,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] ) : "TTFT (s)", accessorKey: "completionStartTime", + meta: { numeric: true }, cell: (info: any) => { const row = info.row.original; const completionStartTime = info.getValue(); @@ -298,7 +301,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] const ttftSeconds = (ttftMs / 1000).toFixed(2); return ( - {ttftSeconds} + {ttftSeconds} ); }, @@ -395,6 +398,7 @@ export const createColumns = (sortProps?: LogsSortProps): ColumnDef[] : "Tokens", accessorKey: "total_tokens", size: 140, + meta: { numeric: true }, cell: (info: any) => { const row = info.row.original; return ( diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index 265e331e6a9..ee08712e56b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -287,6 +287,7 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p row.request_id} onRowClick={handleRowClick} isLoading={isLogsLoading} /> diff --git a/ui/litellm-dashboard/src/components/view_logs/table.test.tsx b/ui/litellm-dashboard/src/components/view_logs/table.test.tsx index 7299d280769..9a8469cfeba 100644 --- a/ui/litellm-dashboard/src/components/view_logs/table.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/table.test.tsx @@ -74,6 +74,35 @@ describe("DataTable states", () => { expect(screen.getByText("Nothing here")).toBeInTheDocument(); }); + it("falls back to generic loading and empty defaults", () => { + const { rerender } = render(); + expect(screen.getByText("Loading...")).toBeInTheDocument(); + + rerender(); + expect(screen.getByText("No results")).toBeInTheDocument(); + }); + + it("suppresses the primitive's row hover on loading, empty, and expansion placeholder rows", async () => { + const user = userEvent.setup(); + const { rerender } = render(); + expect(screen.getByText("Loading...").closest("tr")).toHaveClass("hover:bg-transparent"); + + rerender(); + expect(screen.getByText("No results").closest("tr")).toHaveClass("hover:bg-transparent"); + + rerender( + true} + renderSubComponent={({ row }) =>
details for {row.original.request_id}
} + />, + ); + await user.click(screen.getByRole("button", { name: "expand r1" })); + expect(screen.getByText("details for r1").closest("tr")).toHaveClass("hover:bg-transparent"); + expect(screen.getByText("alpha").closest("tr")).not.toHaveClass("hover:bg-transparent"); + }); + it("renders row data through plain TanStack column defs, including custom cell renderers", () => { const columns: ColumnDef[] = [ { header: "A", accessorKey: "a" }, @@ -84,6 +113,29 @@ describe("DataTable states", () => { expect(screen.getByText("alpha")).toBeInTheDocument(); expect(screen.getByText("custom:beta")).toBeInTheDocument(); }); + + it("clips the table to the rounded wrapper so the header band cannot bleed past the corners", () => { + const { container } = render(); + + const wrapper = container.firstElementChild; + expect(wrapper).toHaveClass("rounded-lg", "overflow-hidden"); + }); + + it("right-aligns headers and cells with tabular figures for numeric meta columns", () => { + const columns: ColumnDef[] = [ + { header: "A", accessorKey: "a" }, + { header: "B", accessorKey: "b", meta: { numeric: true } }, + ]; + render(); + + const headers = screen.getAllByRole("columnheader"); + expect(headers[1].querySelector("div")).toHaveClass("justify-end"); + expect(headers[0].querySelector("div")).not.toHaveClass("justify-end"); + + const cells = screen.getAllByRole("cell"); + expect(cells[1]).toHaveClass("text-right", "tabular-nums"); + expect(cells[0]).not.toHaveClass("text-right"); + }); }); describe("DataTable row interaction", () => { @@ -129,28 +181,33 @@ describe("DataTable expansion", () => { expect(screen.queryByText("details for r1")).not.toBeInTheDocument(); }); - it("renders child rows as sibling table rows (child-rows path)", async () => { + it("keeps expansion attached to the same row through data reorders when getRowId is injected", async () => { const user = userEvent.setup(); - render( + const { rerender } = render( row.request_id} getRowCanExpand={() => true} - renderChildRows={({ row }) => ( - - child of {row.original.request_id} - - )} + renderSubComponent={({ row }) =>
details for {row.original.request_id}
} />, ); - expect(screen.queryByText("child of r2")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "expand r1" })); + expect(screen.getByText("details for r1")).toBeInTheDocument(); - await user.click(screen.getByRole("button", { name: "expand r2" })); + rerender( + row.request_id} + getRowCanExpand={() => true} + renderSubComponent={({ row }) =>
details for {row.original.request_id}
} + />, + ); - const childCell = screen.getByText("child of r2"); - expect(childCell.closest("tr")).not.toBeNull(); - expect(within(screen.getByRole("table")).getByText("child of r2")).toBeInTheDocument(); + expect(screen.getByText("details for r1")).toBeInTheDocument(); + expect(screen.queryByText("details for r2")).not.toBeInTheDocument(); }); it("does not expand rows when getRowCanExpand is missing even if a renderer is provided", () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/table.tsx b/ui/litellm-dashboard/src/components/view_logs/table.tsx index 4510cc9a1f0..c96f34f9b93 100644 --- a/ui/litellm-dashboard/src/components/view_logs/table.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/table.tsx @@ -1,6 +1,7 @@ import { Fragment, useState } from "react"; import { ColumnDef, + RowData, flexRender, getCoreRowModel, getExpandedRowModel, @@ -10,16 +11,21 @@ import { SortingState, } from "@tanstack/react-table"; -import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; +import { Table, TableHeader, TableHead, TableBody, TableRow, TableCell } from "@/components/ui/table"; + +declare module "@tanstack/react-table" { + interface ColumnMeta { + numeric?: boolean; + } +} interface DataTableProps { data: TData[]; columns: ColumnDef[]; + getRowId?: (row: TData, index: number) => string; onRowClick?: (row: TData) => void; - /** Renders inside a single colspan cell (used by audit logs) */ + /** Renders inside a single colspan cell */ renderSubComponent?: (props: { row: Row }) => React.ReactElement; - /** Renders directly in tbody as sibling table rows (used by MCP children) */ - renderChildRows?: (props: { row: Row }) => React.ReactNode; getRowCanExpand?: (row: Row) => boolean; isLoading?: boolean; loadingMessage?: string; @@ -31,16 +37,16 @@ interface DataTableProps { export function DataTable({ data = [], columns, + getRowId, onRowClick, renderSubComponent, - renderChildRows, getRowCanExpand, isLoading = false, - loadingMessage = "🚅 Loading logs...", - noDataMessage = "No logs found", + loadingMessage = "Loading...", + noDataMessage = "No results", enableSorting = false, }: DataTableProps) { - const supportsExpansion = !!(renderSubComponent || renderChildRows) && !!getRowCanExpand; + const supportsExpansion = !!renderSubComponent && !!getRowCanExpand; const hasExplicitColumnSizes = columns.some((column) => column.size !== undefined); const [sorting, setSorting] = useState([]); @@ -55,58 +61,56 @@ export function DataTable({ enableSortingRemoval: false, }), ...(supportsExpansion && { getRowCanExpand }), - getRowId: (row: TData, index: number) => { - const _row: any = row as any; - return _row?.request_id ?? String(index); - }, + ...(getRowId && { getRowId }), getCoreRowModel: getCoreRowModel(), ...(enableSorting && { getSortedRowModel: getSortedRowModel() }), ...(supportsExpansion && { getExpandedRowModel: getExpandedRowModel() }), }); - const tableClassName = hasExplicitColumnSizes - ? "[&_td]:py-0.5 [&_th]:py-1 [&_table]:table-fixed" - : "[&_td]:py-0.5 [&_th]:py-1 table-fixed w-full box-border"; + const tableClassName = hasExplicitColumnSizes ? "table-fixed" : "table-fixed w-full box-border"; const tableStyle = hasExplicitColumnSizes ? { minWidth: table.getCenterTotalSize() } : { minWidth: "400px" }; return ( -
+
- + {table.getHeaderGroups().map((headerGroup) => ( - + {headerGroup.headers.map((header) => { const canSort = enableSorting && header.column.getCanSort(); const isSorted = header.column.getIsSorted(); + const numeric = header.column.columnDef.meta?.numeric; return ( - {header.isPlaceholder ? null : ( -
+
{flexRender(header.column.columnDef.header, header.getContext())} {canSort && ( - + {isSorted === "asc" ? "↑" : isSorted === "desc" ? "↓" : "⇅"} )}
)} - + ); })} ))} - + {isLoading ? ( - + -
+

{loadingMessage}

@@ -115,13 +119,15 @@ export function DataTable({ table.getRowModel().rows.map((row) => ( onRowClick?.(row.original)} > {row.getVisibleCells().map((cell) => ( {flexRender(cell.column.columnDef.cell, cell.getContext())} @@ -129,12 +135,8 @@ export function DataTable({ ))} - {/* Child rows rendered as real table rows (MCP children) */} - {supportsExpansion && row.getIsExpanded() && renderChildRows && renderChildRows({ row })} - - {/* Legacy sub-component in colspan cell (audit logs) */} - {supportsExpansion && row.getIsExpanded() && renderSubComponent && !renderChildRows && ( - + {supportsExpansion && row.getIsExpanded() && renderSubComponent && ( +
{renderSubComponent({ row })}
@@ -143,11 +145,9 @@ export function DataTable({
)) ) : ( - - -
-

{noDataMessage}

-
+ + +

{noDataMessage}

)} From 0f1e29b33486ba6e1600fb93de7214e57e54047d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 8 Jul 2026 14:00:24 -0700 Subject: [PATCH 30/31] fix(bedrock): preserve cache_control ttl on message-level cache points (#32538) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../prompt_templates/factory.py | 12 ++- ...llm_core_utils_prompt_templates_factory.py | 81 +++++++++++++++++++ 2 files changed, 89 insertions(+), 4 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index c1635158d3b..06abb591717 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -4377,6 +4377,7 @@ class BedrockConverseMessagesProcessor: _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", + model=model, ) if _cache_point_block is not None: _parts.append(_cache_point_block) @@ -4384,7 +4385,7 @@ class BedrockConverseMessagesProcessor: elif message_block["content"] and isinstance(message_block["content"], str): _part = BedrockContentBlock(text=messages[msg_i]["content"]) _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block, block_type="content_block" + message_block, block_type="content_block", model=model ) user_content.append(_part) if _cache_point_block is not None: @@ -4509,6 +4510,7 @@ class BedrockConverseMessagesProcessor: _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", + model=model, ) if _cache_point_block is not None: assistants_parts.append(_cache_point_block) @@ -4520,7 +4522,7 @@ class BedrockConverseMessagesProcessor: # If content is empty/whitespace, skip it (don't add a placeholder) # Add cache point block for assistant string content _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - assistant_message_block, block_type="content_block" + assistant_message_block, block_type="content_block", model=model ) if _cache_point_block is not None: assistant_content.append(_cache_point_block) @@ -4745,6 +4747,7 @@ def _bedrock_converse_messages_pt( _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", + model=model, ) if _cache_point_block is not None: _parts.append(_cache_point_block) @@ -4752,7 +4755,7 @@ def _bedrock_converse_messages_pt( elif message_block["content"] and isinstance(message_block["content"], str): _part = BedrockContentBlock(text=messages[msg_i]["content"]) _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - message_block, block_type="content_block" + message_block, block_type="content_block", model=model ) user_content.append(_part) if _cache_point_block is not None: @@ -4882,6 +4885,7 @@ def _bedrock_converse_messages_pt( _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( message_block=cast(OpenAIMessageContentListBlock, element), block_type="content_block", + model=model, ) if _cache_point_block is not None: assistants_parts.append(_cache_point_block) @@ -4892,7 +4896,7 @@ def _bedrock_converse_messages_pt( assistant_content.append(BedrockContentBlock(text=_assistant_content)) # Add cache point block for assistant string content _cache_point_block = litellm.AmazonConverseConfig()._get_cache_point_block( - assistant_message_block, block_type="content_block" + assistant_message_block, block_type="content_block", model=model ) if _cache_point_block is not None: assistant_content.append(_cache_point_block) diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 1d5289737f1..bcda88ea609 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -3085,3 +3085,84 @@ def test_bedrock_converse_messages_pt_document_rejects_url_source(): _bedrock_converse_messages_pt( messages, "anthropic.claude-sonnet-4-6", "bedrock" ) + + +def _collect_cache_points(blocks): + return [ + block["cachePoint"] + for message in blocks + for block in message["content"] + if "cachePoint" in block + ] + + +@pytest.mark.parametrize( + "messages", + [ + [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "conversation history", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + }, + ], + [ + {"role": "user", "content": "hello"}, + { + "role": "assistant", + "content": [ + { + "type": "text", + "text": "assistant reply", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + }, + ], + ], +) +def test_bedrock_converse_message_level_cache_point_preserves_ttl(messages): + """ + Regression for https://github.com/BerriAI/litellm/issues/32154: message-level + cache_control ttl was silently dropped because the message-level + _get_cache_point_block call sites never passed model=, so multi-turn prefixes + fell back to the 5m default while the system prompt kept 1h, churning the + cache every turn on models like Opus 4.8. + """ + result = _bedrock_converse_messages_pt( + messages=messages, + model="eu.anthropic.claude-opus-4-8", + llm_provider="bedrock", + ) + + cache_points = _collect_cache_points(result) + assert cache_points == [{"type": "default", "ttl": "1h"}] + + +@pytest.mark.asyncio +async def test_bedrock_converse_message_level_cache_point_preserves_ttl_async(): + messages = [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "conversation history", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ], + }, + ] + + result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=messages, + model="eu.anthropic.claude-opus-4-8", + llm_provider="bedrock", + ) + + assert _collect_cache_points(result) == [{"type": "default", "ttl": "1h"}] From 5973d9fd2b0d074f963e95fdae2b9c1aef3d88bc Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 8 Jul 2026 14:32:16 -0700 Subject: [PATCH 31/31] feat(ui): add eslint rules for nested ternaries, large inline object args, and long condition chains (#32415) * feat(ui): add eslint rules for nested ternaries, large inline object args, and long condition chains Adds three dashboard lint rules to keep new code readable. Nested ternaries are banned outright via the built-in no-nested-ternary, with the 265 existing occurrences grandfathered in eslint-suppressions.json so only new ones fail. Two custom rules ship as a small local plugin under scripts/eslint-rules: no-large-inline-object-arg flags object literals with 4+ properties passed straight into a call, nudging toward a named variable, and no-long-condition-chain flags boolean expressions that combine 4+ conditions, nudging toward a named boolean. Both are warnings tracked on the existing budget ratchet (eslint-budgets.json + eslint-metrics.json) with headroom above the current counts, so they ratchet down over time rather than freezing a baseline. Both thresholds are configurable rule options and covered by RuleTester unit tests. * fix(ui): scope no-long-condition-chain to boolean operators, not nullish Greptile flagged that the rule counted nullish-coalescing chains the same as &&/|| chains, so a 4-part `a ?? b ?? c ?? d` fallback surfaced "Boolean expression combines 4 conditions", which is inaccurate since a `??` fallback is value defaulting, not a condition. Restrict the visitor to && / || nodes so `??` chains are treated as leaves, while a boolean chain nested inside a `??` is still caught. Drops 6 miscounted occurrences (240 -> 234). * chore(ui): sync lint metrics and suppressions with staging Merge advanced the base branch, adding one no-large-inline-object-arg occurrence (508 -> 509) and making one grandfathered react-hooks suppression stale. Regenerate eslint-metrics.json and prune the suppression so the budget/drift gate passes. * chore(ui): sync lint metrics with staging Merge advanced the base, adding four no-large-inline-object-arg occurrences (509 -> 513). Regenerate eslint-metrics.json so the drift gate passes. --- ui/litellm-dashboard/eslint-budgets.json | 4 +- ui/litellm-dashboard/eslint-metrics.json | 2 + ui/litellm-dashboard/eslint-suppressions.json | 442 +++++++++++++++++- ui/litellm-dashboard/eslint.config.mjs | 6 +- .../scripts/eslint-rules/index.mjs | 11 + .../no-large-inline-object-arg.mjs | 41 ++ .../eslint-rules/no-long-condition-chain.mjs | 41 ++ .../no-large-inline-object-arg.test.ts | 46 ++ .../no-long-condition-chain.test.ts | 51 ++ 9 files changed, 641 insertions(+), 3 deletions(-) create mode 100644 ui/litellm-dashboard/scripts/eslint-rules/index.mjs create mode 100644 ui/litellm-dashboard/scripts/eslint-rules/no-large-inline-object-arg.mjs create mode 100644 ui/litellm-dashboard/scripts/eslint-rules/no-long-condition-chain.mjs create mode 100644 ui/litellm-dashboard/tests/eslint-rules/no-large-inline-object-arg.test.ts create mode 100644 ui/litellm-dashboard/tests/eslint-rules/no-long-condition-chain.test.ts diff --git a/ui/litellm-dashboard/eslint-budgets.json b/ui/litellm-dashboard/eslint-budgets.json index 8dedb9ac9ca..f08e1bb6160 100644 --- a/ui/litellm-dashboard/eslint-budgets.json +++ b/ui/litellm-dashboard/eslint-budgets.json @@ -2,5 +2,7 @@ "@typescript-eslint/no-explicit-any": { "max": 2040, "target": 1500 }, "no-console": { "max": 484, "target": 0 }, "complexity": { "max": 140, "target": 80 }, - "max-depth": { "max": 70, "target": 30 } + "max-depth": { "max": 70, "target": 30 }, + "local/no-large-inline-object-arg": { "max": 560, "target": 300 }, + "local/no-long-condition-chain": { "max": 265, "target": 120 } } diff --git a/ui/litellm-dashboard/eslint-metrics.json b/ui/litellm-dashboard/eslint-metrics.json index 51cef1169f9..f4dc89c5b80 100644 --- a/ui/litellm-dashboard/eslint-metrics.json +++ b/ui/litellm-dashboard/eslint-metrics.json @@ -1,6 +1,8 @@ { "@typescript-eslint/no-explicit-any": 1988, "complexity": 128, + "local/no-large-inline-object-arg": 513, + "local/no-long-condition-chain": 233, "max-depth": 59, "no-console": 15 } diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 67b19471aaf..b077338c75b 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1,4 +1,9 @@ { + "scripts/check-lint-budgets.mjs": { + "no-nested-ternary": { + "count": 1 + } + }, "src/app/(dashboard)/api-reference/APIReferenceView.tsx": { "no-restricted-imports": { "count": 1 @@ -64,6 +69,9 @@ } }, "src/app/(dashboard)/cost-tracking/components/cost_tracking_settings.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -133,11 +141,21 @@ "count": 1 } }, + "src/app/(dashboard)/guardrails-monitor/components/GuardrailDetail.tsx": { + "no-nested-ternary": { + "count": 3 + } + }, "src/app/(dashboard)/guardrails-monitor/components/GuardrailsMonitorView.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/app/(dashboard)/guardrails-monitor/components/GuardrailsOverview.tsx": { + "no-nested-ternary": { + "count": 8 + } + }, "src/app/(dashboard)/guardrails-monitor/components/ScoreChart.test.tsx": { "react/display-name": { "count": 1 @@ -343,6 +361,9 @@ } }, "src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -361,6 +382,9 @@ } }, "src/app/(dashboard)/playground/components/chat_ui/AgentBuilderView.tsx": { + "no-nested-ternary": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 5 } @@ -370,7 +394,15 @@ "count": 1 } }, + "src/app/(dashboard)/playground/components/chat_ui/ChatMessageBubble.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx": { + "no-nested-ternary": { + "count": 7 + }, "no-restricted-imports": { "count": 1 }, @@ -382,6 +414,9 @@ } }, "src/app/(dashboard)/playground/components/chat_ui/CodeInterpreterOutput.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-syntax": { "count": 2 } @@ -392,6 +427,9 @@ } }, "src/app/(dashboard)/playground/components/chat_ui/RealtimePlayground.tsx": { + "no-nested-ternary": { + "count": 2 + }, "react-hooks/immutability": { "count": 2 }, @@ -400,16 +438,27 @@ } }, "src/app/(dashboard)/playground/components/compareUI/CompareUI.tsx": { + "no-nested-ternary": { + "count": 4 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/playground/components/compareUI/components/MessageDisplay.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": { "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": { + "no-nested-ternary": { + "count": 8 + }, "react-hooks/preserve-manual-memoization": { "count": 3 } @@ -474,6 +523,9 @@ } }, "src/app/(dashboard)/projects/components/ProjectDetailsPage.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 } @@ -499,6 +551,9 @@ } }, "src/app/(dashboard)/prompts/components/index.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -517,6 +572,9 @@ } }, "src/app/(dashboard)/prompts/components/prompt_editor_view/PromptCodeSnippets.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -550,6 +608,9 @@ } }, "src/app/(dashboard)/prompts/components/prompt_editor_view/VersionHistorySidePanel.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/immutability": { "count": 1 } @@ -570,6 +631,9 @@ } }, "src/app/(dashboard)/prompts/components/prompt_info.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 }, @@ -578,6 +642,9 @@ } }, "src/app/(dashboard)/prompts/components/prompt_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -635,6 +702,9 @@ } }, "src/app/(dashboard)/users/_components/user_edit_view.test.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react/display-name": { "count": 1 } @@ -648,6 +718,9 @@ } }, "src/app/(dashboard)/users/_components/view_users.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -664,6 +737,9 @@ } }, "src/app/(dashboard)/users/_components/view_users/table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -677,6 +753,9 @@ } }, "src/app/(dashboard)/workflows/WorkflowRuns.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-syntax": { "count": 3 }, @@ -684,6 +763,11 @@ "count": 1 } }, + "src/app/chat/page.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/app/login/LoginPage.tsx": { "react-hooks/set-state-in-effect": { "count": 2 @@ -715,6 +799,9 @@ } }, "src/components/AIHub/ModelHubTable.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -746,6 +833,9 @@ } }, "src/components/AIHub/forms/MakeMCPPublicForm.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 }, @@ -778,11 +868,17 @@ } }, "src/components/DeletedKeysPage/DeletedKeysTable/DeletedKeysTable.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/components/DeletedTeamsPage/DeletedTeamsTable/DeletedTeamsTable.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -815,11 +911,26 @@ "count": 1 } }, + "src/components/GuardrailSettingsView.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/GuardrailsMonitor/LogViewer.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/HelpLink.test.tsx": { "unused-imports/no-unused-imports": { "count": 1 } }, + "src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/Navbar/BlogDropdown/BlogDropdown.test.tsx": { "max-nested-callbacks": { "count": 12 @@ -836,6 +947,9 @@ } }, "src/components/OldTeams.tsx": { + "no-nested-ternary": { + "count": 4 + }, "no-restricted-imports": { "count": 1 }, @@ -856,7 +970,15 @@ "count": 1 } }, + "src/components/Settings/AdminSettings/HashicorpVault/HashicorpVault.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -871,6 +993,11 @@ "count": 1 } }, + "src/components/Settings/AdminSettings/SSOSettings/RedactableField.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/Settings/AdminSettings/SSOSettings/SSOSettingsLoadingSkeleton.test.tsx": { "max-nested-callbacks": { "count": 4 @@ -881,7 +1008,15 @@ "count": 2 } }, + "src/components/Settings/AdminSettings/UISettings/UISettings.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/Settings/LoggingAndAlerts/LoggingCallbacks/LoggingCallbacksTable.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -907,12 +1042,20 @@ "count": 1 } }, + "src/components/TeamSSOSettings.test.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/ToolDetail.tsx": { "unused-imports/no-unused-imports": { "count": 2 } }, "src/components/ToolPolicies.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -932,6 +1075,9 @@ } }, "src/components/UsageIndicator.tsx": { + "no-nested-ternary": { + "count": 4 + }, "no-restricted-imports": { "count": 1 }, @@ -949,6 +1095,11 @@ "count": 1 } }, + "src/components/UsagePage/components/EndpointUsage/components/EndpointUsageTable.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/UsagePage/components/EntityUsage/EntityUsage.tsx": { "no-restricted-imports": { "count": 1 @@ -975,11 +1126,17 @@ } }, "src/components/UsagePage/components/UsageAIChatPanel.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/immutability": { "count": 1 } }, "src/components/UsagePage/components/UsagePageView.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -999,16 +1156,25 @@ } }, "src/components/VirtualKeysPage/VirtualKeysTable.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } }, "src/components/activity_metrics.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/components/add_model/AddModelForm.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1045,11 +1211,22 @@ } }, "src/components/add_model/litellm_model_name.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, + "src/components/add_model/model_connection_test.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, "src/components/add_model/provider_specific_fields.tsx": { + "no-nested-ternary": { + "count": 5 + }, "no-restricted-imports": { "count": 1 }, @@ -1079,6 +1256,9 @@ } }, "src/components/agents/add_agent_form.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 }, @@ -1102,7 +1282,15 @@ "count": 1 } }, + "src/components/agents/agent_form_fields.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/agents/agent_info.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1110,7 +1298,20 @@ "count": 1 } }, + "src/components/agents/agent_virtual_keys.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/agents/dynamic_agent_form_fields.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, "src/components/alerting/dynamic_form.tsx": { + "no-nested-ternary": { + "count": 4 + }, "no-restricted-imports": { "count": 1 } @@ -1123,6 +1324,31 @@ "count": 1 } }, + "src/components/chat/KeysPanel.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/chat/MCPAppsPanel.tsx": { + "no-nested-ternary": { + "count": 7 + } + }, + "src/components/chat/MCPConnectPicker.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/chat/MCPCredentialsTab.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/chat/UsagePanel.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, "src/components/claude_code_plugins.tsx": { "no-restricted-imports": { "count": 1 @@ -1145,6 +1371,9 @@ } }, "src/components/claude_code_plugins/plugin_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1235,6 +1464,9 @@ } }, "src/components/common_components/chartUtils.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1250,6 +1482,9 @@ } }, "src/components/common_components/simple_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1286,6 +1521,9 @@ } }, "src/components/general_settings.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 2 } @@ -1300,17 +1538,28 @@ "count": 1 } }, + "src/components/guardrails/GuardrailTestPlayground.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/guardrails/GuardrailTestResults.tsx": { "no-restricted-imports": { "count": 1 } }, "src/components/guardrails/TeamGuardrailsTab.tsx": { + "no-nested-ternary": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/guardrails/add_guardrail_form.tsx": { + "no-nested-ternary": { + "count": 4 + }, "react-hooks/set-state-in-effect": { "count": 1 }, @@ -1319,11 +1568,17 @@ } }, "src/components/guardrails/content_filter/CompetitorIntentConfiguration.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/guardrails/content_filter/ContentCategoryConfiguration.tsx": { + "no-nested-ternary": { + "count": 3 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1342,6 +1597,9 @@ } }, "src/components/guardrails/custom_code/CustomCodeModal.tsx": { + "no-nested-ternary": { + "count": 6 + }, "no-restricted-imports": { "count": 1 }, @@ -1372,16 +1630,25 @@ } }, "src/components/guardrails/guardrail_optional_params.tsx": { + "no-nested-ternary": { + "count": 5 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/guardrails/guardrail_provider_fields.tsx": { + "no-nested-ternary": { + "count": 5 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/guardrails/guardrail_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 2 } @@ -1407,6 +1674,9 @@ "src/components/llm_calls/chat_completion.tsx": { "max-params": { "count": 1 + }, + "no-nested-ternary": { + "count": 1 } }, "src/components/llm_calls/responses_api.tsx": { @@ -1444,7 +1714,15 @@ "count": 1 } }, + "src/components/mcp_tools/MCPToolArgumentsForm.tsx": { + "no-nested-ternary": { + "count": 5 + } + }, "src/components/mcp_tools/MCPToolsetsTab.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1456,11 +1734,17 @@ } }, "src/components/mcp_tools/McpCrudPermissionPanel.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 } }, "src/components/mcp_tools/OAuthFormFields.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1471,6 +1755,9 @@ } }, "src/components/mcp_tools/ToolTestPanel.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 }, @@ -1478,12 +1765,20 @@ "count": 1 } }, + "src/components/mcp_tools/UserEnvVarsModal.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, "src/components/mcp_tools/create_mcp_server.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, "react-hooks/set-state-in-effect": { - "count": 5 + "count": 4 } }, "src/components/mcp_tools/mcp_connect.tsx": { @@ -1495,6 +1790,9 @@ } }, "src/components/mcp_tools/mcp_connection_status.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 } @@ -1515,6 +1813,9 @@ } }, "src/components/mcp_tools/mcp_server_edit.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1531,6 +1832,9 @@ } }, "src/components/mcp_tools/mcp_servers.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1544,6 +1848,9 @@ } }, "src/components/mcp_tools/mcp_tools.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1575,6 +1882,9 @@ } }, "src/components/model_dashboard/HealthCheckComponent.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 }, @@ -1583,6 +1893,9 @@ } }, "src/components/model_dashboard/all_models_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1591,11 +1904,17 @@ "max-params": { "count": 1 }, + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } }, "src/components/model_dashboard/table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1619,6 +1938,9 @@ } }, "src/components/model_info_view.tsx": { + "no-nested-ternary": { + "count": 14 + }, "no-restricted-imports": { "count": 1 }, @@ -1627,6 +1949,9 @@ } }, "src/components/molecules/filter.tsx": { + "no-nested-ternary": { + "count": 2 + }, "react-hooks/use-memo": { "count": 1 } @@ -1643,6 +1968,9 @@ "max-params": { "count": 1 }, + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -1656,6 +1984,9 @@ "max-params": { "count": 23 }, + "no-nested-ternary": { + "count": 5 + }, "no-restricted-syntax": { "count": 154 } @@ -1731,6 +2062,9 @@ } }, "src/components/permissions/MCPServerPermissions.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 } @@ -1740,6 +2074,11 @@ "count": 1 } }, + "src/components/policies/PolicySelector.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/policies/add_attachment_form.tsx": { "no-restricted-imports": { "count": 1 @@ -1760,6 +2099,9 @@ } }, "src/components/policies/ai_suggestion_modal.tsx": { + "no-nested-ternary": { + "count": 10 + }, "no-restricted-imports": { "count": 1 }, @@ -1773,11 +2115,17 @@ } }, "src/components/policies/attachment_table.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/components/policies/guardrail_selection_modal.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1788,6 +2136,9 @@ } }, "src/components/policies/impact_popover.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1806,6 +2157,9 @@ } }, "src/components/policies/pipeline_flow_builder.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1827,6 +2181,9 @@ } }, "src/components/policies/policy_table.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -1856,6 +2213,9 @@ } }, "src/components/public_model_hub.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1865,12 +2225,25 @@ "count": 1 } }, + "src/components/router_settings/ReliabilityRetriesSection.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/routing_groups/index.tsx": { "react-hooks/preserve-manual-memoization": { "count": 1 } }, + "src/components/search_tools/SearchToolSelector.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/settings.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -1925,6 +2298,9 @@ } }, "src/components/team/EditMembership.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1935,6 +2311,9 @@ } }, "src/components/team/TeamInfo.tsx": { + "no-nested-ternary": { + "count": 3 + }, "no-restricted-imports": { "count": 1 }, @@ -1943,6 +2322,9 @@ } }, "src/components/team/TeamVirtualKeysTable.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -1966,6 +2348,9 @@ } }, "src/components/templates/key_edit_view.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } @@ -1976,6 +2361,9 @@ } }, "src/components/templates/key_info_view.tsx": { + "no-nested-ternary": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -2016,6 +2404,9 @@ } }, "src/components/vector_store_management/VectorStoreForm.tsx": { + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 }, @@ -2044,12 +2435,38 @@ "count": 1 } }, + "src/components/view_logs/EvalViewer/EvalViewer.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/view_logs/GuardrailViewer/CompliancePanel.tsx": { + "no-nested-ternary": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/view_logs/GuardrailViewer/ContentFilterDetails.tsx": { + "no-nested-ternary": { + "count": 1 + } + }, + "src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx": { + "no-nested-ternary": { + "count": 4 + } + }, + "src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx": { + "no-nested-ternary": { + "count": 4 + } + }, "src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": { + "no-nested-ternary": { + "count": 3 + }, "react-hooks/set-state-in-effect": { "count": 2 } @@ -2059,11 +2476,21 @@ "count": 2 } }, + "src/components/view_logs/LogDetailsDrawer/prettyMessagesUtils.ts": { + "no-nested-ternary": { + "count": 1 + } + }, "src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts": { "react-hooks/immutability": { "count": 2 } }, + "src/components/view_logs/LogsTableToolbar.tsx": { + "no-nested-ternary": { + "count": 4 + } + }, "src/components/view_logs/columns.tsx": { "no-restricted-imports": { "count": 1 @@ -2077,6 +2504,11 @@ "count": 1 } }, + "src/components/view_logs/table.tsx": { + "no-nested-ternary": { + "count": 2 + } + }, "src/components/view_user_spend.tsx": { "react-hooks/set-state-in-effect": { "count": 2 @@ -2113,6 +2545,9 @@ } }, "src/hooks/useTestMCPConnection.tsx": { + "no-nested-ternary": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -2130,6 +2565,11 @@ "count": 1 } }, + "src/lib/http/client.ts": { + "no-nested-ternary": { + "count": 1 + } + }, "src/utils/dataUtils.test.ts": { "max-nested-callbacks": { "count": 1 diff --git a/ui/litellm-dashboard/eslint.config.mjs b/ui/litellm-dashboard/eslint.config.mjs index acdd0c91309..0cf5b4ff655 100644 --- a/ui/litellm-dashboard/eslint.config.mjs +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -3,6 +3,7 @@ import tseslint from "typescript-eslint"; import nextCoreWebVitals from "eslint-config-next/core-web-vitals"; import prettier from "eslint-config-prettier/flat"; import unusedImports from "eslint-plugin-unused-imports"; +import local from "./scripts/eslint-rules/index.mjs"; const eslintConfig = [ { @@ -13,9 +14,11 @@ const eslintConfig = [ ...nextCoreWebVitals, prettier, { - plugins: { "unused-imports": unusedImports }, + plugins: { "unused-imports": unusedImports, local }, rules: { "unused-imports/no-unused-imports": "error", + "local/no-large-inline-object-arg": "warn", + "local/no-long-condition-chain": "warn", "@typescript-eslint/no-explicit-any": "warn", "no-console": ["warn", { allow: ["warn", "error"] }], "@typescript-eslint/no-unused-vars": "off", @@ -28,6 +31,7 @@ const eslintConfig = [ "no-useless-escape": "off", "no-self-assign": "error", "no-var": "error", + "no-nested-ternary": "error", "react/no-danger": "error", complexity: ["warn", 20], "max-depth": ["warn", 4], diff --git a/ui/litellm-dashboard/scripts/eslint-rules/index.mjs b/ui/litellm-dashboard/scripts/eslint-rules/index.mjs new file mode 100644 index 00000000000..150ba1d02e9 --- /dev/null +++ b/ui/litellm-dashboard/scripts/eslint-rules/index.mjs @@ -0,0 +1,11 @@ +import noLargeInlineObjectArg from "./no-large-inline-object-arg.mjs"; +import noLongConditionChain from "./no-long-condition-chain.mjs"; + +const plugin = { + rules: { + "no-large-inline-object-arg": noLargeInlineObjectArg, + "no-long-condition-chain": noLongConditionChain, + }, +}; + +export default plugin; diff --git a/ui/litellm-dashboard/scripts/eslint-rules/no-large-inline-object-arg.mjs b/ui/litellm-dashboard/scripts/eslint-rules/no-large-inline-object-arg.mjs new file mode 100644 index 00000000000..5c5ae170e23 --- /dev/null +++ b/ui/litellm-dashboard/scripts/eslint-rules/no-large-inline-object-arg.mjs @@ -0,0 +1,41 @@ +const DEFAULT_MIN_PROPERTIES = 4; + +const isArgumentOf = (node) => { + const parent = node.parent; + if (parent == null) return false; + if (parent.type !== "CallExpression" && parent.type !== "NewExpression") return false; + return parent.arguments.includes(node); +}; + +const rule = { + meta: { + type: "suggestion", + docs: { + description: + "Disallow passing a large object literal inline as a call argument; assign it to a named variable first.", + }, + schema: [ + { + type: "object", + properties: { minProperties: { type: "integer", minimum: 1 } }, + additionalProperties: false, + }, + ], + messages: { + tooLarge: + "Object literal with {{count}} properties passed inline as an argument; assign it to a named variable first.", + }, + }, + create(context) { + const minProperties = context.options[0]?.minProperties ?? DEFAULT_MIN_PROPERTIES; + return { + ObjectExpression(node) { + if (!isArgumentOf(node)) return; + if (node.properties.length < minProperties) return; + context.report({ node, messageId: "tooLarge", data: { count: node.properties.length } }); + }, + }; + }, +}; + +export default rule; diff --git a/ui/litellm-dashboard/scripts/eslint-rules/no-long-condition-chain.mjs b/ui/litellm-dashboard/scripts/eslint-rules/no-long-condition-chain.mjs new file mode 100644 index 00000000000..638e57442e2 --- /dev/null +++ b/ui/litellm-dashboard/scripts/eslint-rules/no-long-condition-chain.mjs @@ -0,0 +1,41 @@ +const DEFAULT_MIN_CONDITIONS = 4; + +const isBooleanLogical = (node) => + node?.type === "LogicalExpression" && (node.operator === "&&" || node.operator === "||"); + +const countConditions = (node) => + isBooleanLogical(node) ? countConditions(node.left) + countConditions(node.right) : 1; + +const rule = { + meta: { + type: "suggestion", + docs: { + description: + "Disallow logical expressions that combine many conditions; extract the condition into a named boolean.", + }, + schema: [ + { + type: "object", + properties: { minConditions: { type: "integer", minimum: 2 } }, + additionalProperties: false, + }, + ], + messages: { + tooMany: "Boolean expression combines {{count}} conditions; extract it into a named variable.", + }, + }, + create(context) { + const minConditions = context.options[0]?.minConditions ?? DEFAULT_MIN_CONDITIONS; + return { + LogicalExpression(node) { + if (!isBooleanLogical(node)) return; + if (isBooleanLogical(node.parent)) return; + const count = countConditions(node); + if (count < minConditions) return; + context.report({ node, messageId: "tooMany", data: { count } }); + }, + }; + }, +}; + +export default rule; diff --git a/ui/litellm-dashboard/tests/eslint-rules/no-large-inline-object-arg.test.ts b/ui/litellm-dashboard/tests/eslint-rules/no-large-inline-object-arg.test.ts new file mode 100644 index 00000000000..dfe22ea8266 --- /dev/null +++ b/ui/litellm-dashboard/tests/eslint-rules/no-large-inline-object-arg.test.ts @@ -0,0 +1,46 @@ +import { RuleTester } from "eslint"; +import rule from "../../scripts/eslint-rules/no-large-inline-object-arg.mjs"; + +const ruleTester = new RuleTester({ + languageOptions: { ecmaVersion: "latest", sourceType: "module" }, +}); + +ruleTester.run("no-large-inline-object-arg", rule as never, { + valid: [ + "foo({ a: 1, b: 2, c: 3 });", + "foo({});", + "const opts = { a: 1, b: 2, c: 3, d: 4 }; foo(opts);", + "const x = { a: 1, b: 2, c: 3, d: 4 };", + "function f() { return { a: 1, b: 2, c: 3, d: 4 }; }", + "const arr = [{ a: 1, b: 2, c: 3, d: 4 }];", + "foo(1, 2, { a: 1, b: 2 });", + { code: "foo({ a: 1, b: 2, c: 3, d: 4 });", options: [{ minProperties: 5 }] }, + ], + invalid: [ + { + code: "foo({ a: 1, b: 2, c: 3, d: 4 });", + errors: [{ messageId: "tooLarge", data: { count: 4 } }], + }, + { + code: "new Widget({ a: 1, b: 2, c: 3, d: 4, e: 5 });", + errors: [{ messageId: "tooLarge", data: { count: 5 } }], + }, + { + code: "foo(1, { a: 1, b: 2, c: 3, d: 4 });", + errors: [{ messageId: "tooLarge" }], + }, + { + code: "foo({ a: 1, ...rest, c: 3, d: 4 });", + errors: [{ messageId: "tooLarge", data: { count: 4 } }], + }, + { + code: "foo({ a: 1, b: 2, c: 3 });", + options: [{ minProperties: 3 }], + errors: [{ messageId: "tooLarge", data: { count: 3 } }], + }, + { + code: "outer({ a: 1, b: 2, c: 3, d: 4 }, inner({ e: 5, f: 6, g: 7, h: 8 }));", + errors: [{ messageId: "tooLarge" }, { messageId: "tooLarge" }], + }, + ], +}); diff --git a/ui/litellm-dashboard/tests/eslint-rules/no-long-condition-chain.test.ts b/ui/litellm-dashboard/tests/eslint-rules/no-long-condition-chain.test.ts new file mode 100644 index 00000000000..4a7fb677190 --- /dev/null +++ b/ui/litellm-dashboard/tests/eslint-rules/no-long-condition-chain.test.ts @@ -0,0 +1,51 @@ +import { RuleTester } from "eslint"; +import rule from "../../scripts/eslint-rules/no-long-condition-chain.mjs"; + +const ruleTester = new RuleTester({ + languageOptions: { ecmaVersion: "latest", sourceType: "module" }, +}); + +ruleTester.run("no-long-condition-chain", rule as never, { + valid: [ + "const x = a && b && c;", + "const x = a || b || c;", + "const x = a && (b || c);", + "const x = a && b;", + "if (a || b || c) {}", + "const x = a ?? b ?? c;", + "const url = a ?? b ?? c ?? d;", + "const x = (a && b) ?? c ?? d;", + { code: "const x = a && b && c && d;", options: [{ minConditions: 5 }] }, + ], + invalid: [ + { + code: "const x = a && b && c && d;", + errors: [{ messageId: "tooMany", data: { count: 4 } }], + }, + { + code: "const x = a || b || c || d || e;", + errors: [{ messageId: "tooMany", data: { count: 5 } }], + }, + { + code: "const x = a && b || c && d;", + errors: [{ messageId: "tooMany", data: { count: 4 } }], + }, + { + code: "if (!a && !b && !c && !d) {}", + errors: [{ messageId: "tooMany", data: { count: 4 } }], + }, + { + code: "const x = a && (b || c);", + options: [{ minConditions: 3 }], + errors: [{ messageId: "tooMany", data: { count: 3 } }], + }, + { + code: "const x = (a && b && c && d) || (e && f && g && h);", + errors: [{ messageId: "tooMany", data: { count: 8 } }], + }, + { + code: "const x = (a && b && c && d) ?? fallback;", + errors: [{ messageId: "tooMany", data: { count: 4 } }], + }, + ], +});