From 3aaa9bbd4519b136651900a12e48cb69772de18f Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 9 Jul 2026 08:10:08 +0000 Subject: [PATCH 01/54] test(managed-files): lock in store_unified_file_id idempotency on retrieve --- .../proxy/test_managed_files_hook.py | 51 +++++++++++++++++++ 1 file changed, 51 insertions(+) diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 3169b9b08e0..1526aad7a24 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -385,3 +385,54 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri(): message = str(exc_info.value) assert unified_file_id in message assert s3_uri not in message + + +def _make_real_managed_files_instance(): + """Create a _PROXY_LiteLLMManagedFiles with a real store_unified_file_id but + an AsyncMock prisma client, so the DB write path itself can be asserted.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + mock_cache = MagicMock() + mock_cache.async_set_cache = AsyncMock() + + mock_prisma = MagicMock() + mock_prisma.db.litellm_managedfiletable.upsert = AsyncMock() + mock_prisma.db.litellm_managedfiletable.create = AsyncMock( + side_effect=AssertionError( + "store_unified_file_id must upsert, not create, on the retrieve path" + ) + ) + + return ( + _PROXY_LiteLLMManagedFiles( + internal_usage_cache=mock_cache, + prisma_client=mock_prisma, + ), + mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_store_unified_file_id_is_idempotent_via_upsert(): + """Regression test for the managed-batch retrieve 500 (UniqueViolationError on + unified_file_id): re-registering an already-stored output file id must upsert on + unified_file_id, never do an unconditional create that raises on conflict.""" + managed_files, mock_prisma = _make_real_managed_files_instance() + file_id = "litellm_proxy_unified_output_id_abc" + + await managed_files.store_unified_file_id( + file_id=file_id, + file_object=_make_file_object(), + litellm_parent_otel_span=None, + model_mappings={"model-deploy-xyz": "file-output-abc"}, + user_api_key_dict=_make_user_api_key_dict(), + ) + + mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited() + mock_prisma.db.litellm_managedfiletable.upsert.assert_awaited_once() + assert ( + mock_prisma.db.litellm_managedfiletable.upsert.await_args.kwargs["where"] + == {"unified_file_id": file_id} + ) From aca7f5732438057b76314c2ecdc39208815b401e Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 22 Jul 2026 04:23:35 +0000 Subject: [PATCH 02/54] fix(jwt_auth): allow /v1/messages for JWT teams by default Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 2 +- .../proxy/auth/test_handle_jwt.py | 59 +++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7df725bf965..53b6e51aaef 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4254,7 +4254,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): team_id_upsert: bool = False team_ids_jwt_field: Optional[str] = None upsert_sso_user_to_team: bool = False - team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"] + team_allowed_routes: List[str] = ["openai_routes", "anthropic_routes", "info_routes", "mcp_routes"] team_id_default: Optional[str] = Field( default=None, description="If no team_id given, default permissions/spend-tracking to this team.s", diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index ffc5241d027..35aaa0f1254 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1240,6 +1240,65 @@ async def test_find_team_with_model_access_model_group(monkeypatch): assert team_obj.team_id == "team-1" +@pytest.mark.asyncio +async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatch): + """Regression for #31189: a single-team JWT that grants the requested model + through an access group must resolve on /v1/messages without an explicit + x-litellm-team-id header. /v1/messages lives in `anthropic_routes`, so when a + team has no `team_allowed_routes` configured the default allowlist must cover + it just like /chat/completions and /v1/responses; otherwise the internal route + check fails and surfaces a misleading "No team has access to the requested + model" 403.""" + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + router = Router( + model_list=[ + { + "model_name": "claude-sonnet-4-6", + "litellm_params": {"model": "claude-sonnet-4-6"}, + "model_info": {"access_groups": ["coding_only_models"]}, + } + ] + ) + import sys + import types + + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + team = LiteLLM_TeamTable(team_id="coding-team", models=["coding_only_models"]) + + async def mock_get_team_object(*args, **kwargs): # type: ignore + return team + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + team_id, team_obj = await JWTAuthManager.find_team_with_model_access( + team_ids={"coding-team"}, + requested_model="claude-sonnet-4-6", + route="/v1/messages", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + assert team_id == "coding-team" + assert team_obj.team_id == "coding-team" + + @pytest.mark.asyncio async def test_auth_builder_returns_team_membership_object(): """ From 221b1859db12f76e730e3c3a8eb1a410814744b0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 22 Jul 2026 17:53:38 +0000 Subject: [PATCH 03/54] fix(proxy): stop serving stale team model allowlist after /team/update get_team_object consults proxy_logging_obj.internal_usage_cache before user_api_key_cache, but _cache_team_object (the refresh every team mutation goes through) only wrote user_api_key_cache. With enable_redis_auth_cache both caches share one Redis, so any request backfills the internal cache's in-memory tier with the team object and that copy keeps shadowing the freshly written team until its TTL expires. The auth builder then wrote the team object it had just read back into the cache after check 6, clobbering the fresh Redis value with the stale one, which made the staleness self-sustaining under traffic: keys with models=["all-team-models"] kept getting 403 team_model_access_denied for models added via /team/update, and kept access to removed ones. _cache_team_object now deletes the internal usage cache entry before writing the refreshed team, and the auth-time write-back is removed so only authoritative writers (DB reads and team mutations) populate the team cache, mirroring how key objects already handle this (see test_auth_does_not_rewrite_cached_key_object_back_into_cache). The LIT-4000 test pinning the removed write-back is deleted; its concern (team object cached under the canonical key) is handled by _cache_team_object inside get_team_object's DB path and pinned by test_cache_team_object_writes_team_id_and_invalidates_team_alias Resolves LIT-4391 --- litellm/proxy/auth/auth_checks.py | 7 +- litellm/proxy/auth/user_api_key_auth.py | 11 - .../proxy/auth/test_auth_checks.py | 121 +++++++++- .../proxy/auth/test_user_api_key_auth.py | 212 ++++++++++-------- .../test_team_endpoints.py | 1 + 5 files changed, 241 insertions(+), 111 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 99a867a5d07..4c998f997d9 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1764,9 +1764,14 @@ async def _cache_team_object( ## CACHE REFRESH TIME! team_table.last_refreshed_at = time.time() + key = "team_id:{}".format(team_id) + + if proxy_logging_obj is not None: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) + # team_id is the table primary key — guaranteed unique, safe to write. await _cache_management_object( - key="team_id:{}".format(team_id), + key=key, value=team_table, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 83a8a69511b..620669729f7 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1945,17 +1945,6 @@ async def _user_api_key_auth_builder( else: valid_token.team_object_permission = None - # Cache under the canonical "team_id:{id}" key so get_team_object and - # _update_team_cache serve this write from the L2 cache. The guard keeps a - # non-team (personal) key, whose team_id is None, from reaching the cache - # layer, which Redis rejects with a NoneType key error. - if valid_token.team_id is not None and _team_obj is not None: - await user_api_key_cache.async_set_cache( - key=f"team_id:{valid_token.team_id}", - value=_team_obj, - model_type=LiteLLM_TeamTableCachedObj, - ) - # Fetch project object if key belongs to a project _project_obj = None if valid_token.project_id is not None: diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 5e07d1bcbc5..b7bf5728cbd 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -50,6 +50,7 @@ from litellm.proxy.auth.auth_checks import ( vector_store_access_check, ) from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache @@ -4156,6 +4157,11 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): len(teams)==1 before populating the cache. 3. When team_alias is None, NO alias-key operation happens (no delete of an empty-keyed entry, no spurious write). + 4. DELETES the team_id-keyed entry from the internal usage cache + BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache` + consults the internal usage cache first, so a leftover copy there + (backfilled from a Redis shared with `user_api_key_cache`) would + keep serving the pre-update team allowlist. """ from unittest.mock import AsyncMock, MagicMock @@ -4202,9 +4208,14 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): # (2) team_alias-keyed entry is deleted in BOTH the in-memory cache # and the Redis dual cache (mirrors _delete_cache_key_object pattern). cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity") - logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with( - key="team_alias:H-Capacity" - ) + + # (4) internal usage cache: team_id entry deleted BEFORE the fresh + # write, alias entry deleted as before. + internal_deleted_keys = [ + (c.kwargs.get("key") or c.args[0]) + for c in logging_obj.internal_usage_cache.dual_cache.async_delete_cache.await_args_list + ] + assert internal_deleted_keys == ["team_id:team-1234", "team_alias:H-Capacity"] # ===== team_alias is None: no alias-key operation ===== aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None}) @@ -4222,7 +4233,9 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): ) cache2.delete_cache.assert_not_called() - logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_not_awaited() + logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with( + key="team_id:team-no-alias" + ) written_keys_aliasless = [ (c.kwargs.get("key") or c.args[0]) for c in cache2.async_set_cache.await_args_list @@ -4230,6 +4243,106 @@ async def test_cache_team_object_writes_team_id_and_invalidates_team_alias(): assert written_keys_aliasless == ["team_id:team-no-alias"] +class _SharedFakeRedis(RedisCache): + """Dict-backed stand-in for the single Redis that both + ``user_api_key_cache`` (enable_redis_auth_cache) and + ``proxy_logging_obj.internal_usage_cache.dual_cache`` share in the + LIT-4391 deployment topology. Only the methods DualCache calls are + implemented; ``super().__init__`` is skipped intentionally.""" + + def __init__(self): + self._store: dict = {} + + async def async_set_cache(self, key, value, **kwargs): + self._store[key] = json.dumps(value) + + async def async_get_cache(self, key, **kwargs): + raw = self._store.get(key) + return json.loads(raw) if raw is not None else None + + async def async_delete_cache(self, key): + self._store.pop(key, None) + + +@pytest.mark.asyncio +async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): + """ + Regression test for LIT-4391: keys with models=["all-team-models"] kept + getting 403 team_model_access_denied for models added via /team/update. + + `_get_team_object_from_cache` consults the internal usage cache BEFORE + `user_api_key_cache`. When both share one Redis (enable_redis_auth_cache), + any team read backfills the internal cache's in-memory tier with the team + object. `_cache_team_object` (the /team/update refresh) only wrote + `user_api_key_cache`, so that backfilled copy kept shadowing the update + until its TTL expired — and the auth-time write-back then pushed the stale + copy back into the shared Redis, making the staleness self-sustaining. + + Pins: + 1. After `_cache_team_object` writes an updated team, `get_team_object` + returns the UPDATED model list even though the internal usage cache's + in-memory tier was backfilled with the pre-update team. + 2. The shared Redis still holds the updated team afterwards — the + internal-cache invalidation must happen BEFORE the fresh write, or it + would wipe the value it just wrote. + """ + from litellm.caching.dual_cache import DualCache + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object + + team_id = "team-lit-4391" + shared_redis = _SharedFakeRedis() + user_api_key_cache = UserApiKeyCache(redis_cache=shared_redis) + proxy_logging_obj = MagicMock() + proxy_logging_obj.internal_usage_cache.dual_cache = DualCache( + redis_cache=shared_redis, + default_in_memory_ttl=300, + ) + prisma_client = MagicMock() + + await _cache_team_object( + team_id=team_id, + team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + primed = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert primed is not None and primed.models == ["model-a"] + + await _cache_team_object( + team_id=team_id, + team_table=LiteLLM_TeamTableCachedObj( + team_id=team_id, models=["model-a", "model-b"] + ), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + refreshed = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert refreshed is not None and refreshed.models == ["model-a", "model-b"], ( + "get_team_object served a stale team allowlist after _cache_team_object " + f"refreshed it. Got models={refreshed.models if refreshed else None}" + ) + + redis_copy = await shared_redis.async_get_cache(f"team_id:{team_id}") + assert redis_copy is not None and redis_copy["models"] == ["model-a", "model-b"], ( + "The shared Redis lost the refreshed team object — the internal-cache " + "invalidation must run BEFORE the fresh write, not after. " + f"Got: {redis_copy}" + ) + + MODEL_DISCOVERY_ROUTES = [ "/v1/models", "/models", diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 2c1948adca1..d70e53f2c8a 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2635,6 +2635,123 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_auth_flow_never_persists_fallback_team_object_lit_4391(): + """ + Regression test for LIT-4391 (stale team allowlist poisoning). + + When `get_team_object` fails at the "Check 6" team-auth step (cache miss + inside the DB-throttle window, DB blip, ...), the builder falls back to a + team object reconstructed from the CACHED token's team_* snapshot — which + can be arbitrarily stale (e.g. pre-/team/update models). + + The builder used to write that team object back into `user_api_key_cache` + under "team_id:" after Check 6. Writing a cache-read (or worse, a + token-snapshot) value back into the shared cache re-poisons it — with + enable_redis_auth_cache it clobbered the fresh team `/team/update` had + just written to Redis, making the stale allowlist self-sustaining across + requests. Only authoritative writers (`_cache_team_object` on DB reads and + team mutations) may populate the team cache. + + Pins: the auth flow completes on the fallback path WITHOUT writing any + "team_id:*" cache entry. + """ + from starlette.datastructures import URL + from starlette.requests import Request + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder + + api_key = "sk-test-lit-4391-no-team-writeback" + valid_token = UserAPIKeyAuth( + api_key=api_key, + token=api_key, + user_role=LitellmUserRoles.INTERNAL_USER, + team_id="team-lit-4391", + team_models=["model-a"], + models=["all-team-models"], + ) + + mock_cache = AsyncMock() + mock_cache.async_get_cache = AsyncMock(return_value=valid_token) + mock_cache.async_set_cache = AsyncMock(return_value=None) + + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.internal_usage_cache = MagicMock() + mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock() + mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs = { + "prisma_client": MagicMock(), + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _originals = {k: getattr(_proxy_server_mod, k, None) for k in _attrs} + + try: + for k, v in _attrs.items(): + setattr(_proxy_server_mod, k, v) + + request = Request(scope={"type": "http"}) + request._url = URL(url="/chat/completions") + + with ( + patch( + "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", + new_callable=AsyncMock, + return_value=valid_token, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.get_team_object", + new_callable=AsyncMock, + side_effect=HTTPException( + status_code=404, + detail={"error": "Team doesn't exist in db."}, + ), + ), + ): + result = await _user_api_key_auth_builder( + request=request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + + assert result.team_id == "team-lit-4391" + + team_cache_writes = [ + key + for c in mock_cache.async_set_cache.await_args_list + if isinstance(key := (c.kwargs.get("key") if "key" in c.kwargs else c.args[0]), str) + and key.startswith("team_id:") + ] + assert team_cache_writes == [], ( + "The auth flow wrote a team object into the cache. Fallback/" + "cache-read team objects must never be persisted — only " + "_cache_team_object (DB reads and team mutations) may write " + f"'team_id:*' entries. Got writes: {team_cache_writes}" + ) + + finally: + for k, v in _originals.items(): + setattr(_proxy_server_mod, k, v) + + # --------------------------------------------------------------------------- # _run_centralized_common_checks — centralized authz gate @@ -4073,101 +4190,6 @@ async def test_real_jwt_still_requires_license_when_jwt_auth_enabled(monkeypatch assert "enterprise only feature" in message -@pytest.mark.asyncio -async def test_auth_path_caches_team_object_under_canonical_team_id_key(): - """Regression for LIT-4000: the auth builder must cache the team object under - the canonical ``team_id:{id}`` key that ``get_team_object`` and - ``_update_team_cache`` read, never under the raw ``team_id`` (and never under - a ``None`` key, which Redis rejects with a NoneType key error). A raw or None - key is silently dropped by Redis / never served back, so every request - re-hits Postgres for the team object instead of the L2 cache. - - Drives the real builder for a team-scoped key against a real in-memory - ``UserApiKeyCache`` and reads the team object back. Mutating the cache key at - the write site to the raw ``valid_token.team_id`` (or ``None``) makes the - canonical-key read miss and fails this test. - """ - from fastapi import Request - from starlette.datastructures import URL - - import litellm.proxy.proxy_server as _proxy_server_mod - from litellm.proxy._types import LiteLLM_TeamTableCachedObj - from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder - from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache - from litellm.proxy.proxy_server import hash_token - - team_id = "team-lit-4000" - api_key = "sk-lit-4000-team-key" - cache = UserApiKeyCache() - - team_token = UserAPIKeyAuth(token=hash_token(api_key), team_id=team_id) - team_obj = LiteLLM_TeamTableCachedObj(team_id=team_id) - - proxy_logging_obj = MagicMock() - proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) - attrs = { - "prisma_client": MagicMock(), - "user_api_key_cache": cache, - "proxy_logging_obj": proxy_logging_obj, - "master_key": "sk-test-master", - "general_settings": {"allow_requests_on_db_unavailable": False}, - "llm_model_list": [], - "llm_router": None, - "open_telemetry_logger": None, - "model_max_budget_limiter": MagicMock(), - "user_custom_auth": None, - "jwt_handler": None, - "litellm_proxy_admin_name": "admin", - } - originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} - try: - for k, v in attrs.items(): - setattr(_proxy_server_mod, k, v) - request = Request(scope={"type": "http"}) - request._url = URL(url="/chat/completions") - with ( - patch( - "litellm.proxy.auth.resolvers.store.IdentityStore._resolve_key", - AsyncMock(return_value=team_token), - ), - patch( - "litellm.proxy.auth.user_api_key_auth.get_team_object", - AsyncMock(return_value=team_obj), - ), - patch( - "litellm.proxy.auth.user_api_key_auth._enforce_key_and_fallback_model_access", - new_callable=AsyncMock, - ), - patch( - "litellm.proxy.auth.user_api_key_auth._return_user_api_key_auth_obj", - new_callable=AsyncMock, - return_value=team_token, - ), - patch( - "litellm.proxy.auth.auth_exception_handler.seed_request_identity", - ), - ): - await _user_api_key_auth_builder( - request=request, - api_key=f"Bearer {api_key}", - azure_api_key_header="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - request_data={}, - ) - finally: - for k, v in originals.items(): - setattr(_proxy_server_mod, k, v) - - served = cache.get_cache( - key=f"team_id:{team_id}", model_type=LiteLLM_TeamTableCachedObj - ) - assert served is not None and served.team_id == team_id - assert cache.get_cache(key=team_id) is None - assert cache.get_cache(key=None) is None - - @pytest.mark.asyncio async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): """A cache-hit auth must not write the token back into the cache. diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 4936191c344..2b2e0681ee1 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -6500,6 +6500,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(): return_value=mock_existing_team ) mock_cache.async_set_cache = AsyncMock() + mock_logging.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock() # Mock team update mock_updated_team = MagicMock(spec=LiteLLM_TeamTable) From 769f434bd252f62b26146d42b46407771e0cb2c0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 22 Jul 2026 18:06:39 +0000 Subject: [PATCH 04/54] fix(proxy): make team cache invalidation best-effort Greptile review: _cache_team_object runs after a successful DB fetch in get_team_object and after every team mutation's DB write, but DualCache.async_delete_cache propagates backend errors, so a Redis blip during invalidation would turn a healthy team lookup into a 404 and a committed /team/update into a 500. Both the internal usage cache delete and the alias-key invalidation now log a warning and continue on failure, matching how DualCache.async_set_cache already swallows write errors. Worst case on failure is worker-local staleness bounded by the internal cache's in-memory TTL, the same bound other workers already have --- litellm/proxy/auth/auth_checks.py | 24 ++++++++++-- .../proxy/auth/test_auth_checks.py | 39 +++++++++++++++++++ 2 files changed, 59 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 4c998f997d9..21b455c5d44 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1767,7 +1767,15 @@ async def _cache_team_object( key = "team_id:{}".format(team_id) if proxy_logging_obj is not None: - await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) + try: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to invalidate internal usage cache entry %s; " + "a stale team object may be served until its TTL expires: %s", + key, + e, + ) # team_id is the table primary key — guaranteed unique, safe to write. await _cache_management_object( @@ -1793,9 +1801,17 @@ async def _cache_team_object( # the cache from a verified single row. if team_table.team_alias: alias_key = "team_alias:{}".format(team_table.team_alias) - user_api_key_cache.delete_cache(key=alias_key) - if proxy_logging_obj is not None: - await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key) + try: + user_api_key_cache.delete_cache(key=alias_key) + if proxy_logging_obj is not None: + await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to invalidate cached team alias entry %s; " + "a stale team object may be served until its TTL expires: %s", + alias_key, + e, + ) async def _cache_key_object( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index b7bf5728cbd..fa7ee5b531a 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -4343,6 +4343,45 @@ async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391(): ) +@pytest.mark.asyncio +async def test_cache_team_object_tolerates_cache_invalidation_failures(): + """ + Greptile review on the LIT-4391 fix: `_cache_team_object` runs after a + successful DB fetch (inside `get_team_object`) and after every team + mutation's DB write. A cache-backend error during the best-effort + invalidations must NOT fail those operations — otherwise a Redis blip + turns a healthy team lookup into a 404 and a committed /team/update into + a 500. The authoritative team_id-keyed write must still happen. + """ + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import _cache_team_object + + cache = MagicMock() + cache.async_set_cache = AsyncMock() + cache.delete_cache = MagicMock(side_effect=Exception("redis down")) + logging_obj = MagicMock() + logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock( + side_effect=Exception("redis down") + ) + + await _cache_team_object( + team_id="team-cache-outage", + team_table=LiteLLM_TeamTableCachedObj( + team_id="team-cache-outage", + team_alias="cache-outage-alias", + models=["model-a"], + ), + user_api_key_cache=cache, + proxy_logging_obj=logging_obj, + ) + + written_keys = [ + (c.kwargs.get("key") or c.args[0]) + for c in cache.async_set_cache.await_args_list + ] + assert written_keys == ["team_id:team-cache-outage"] + + MODEL_DISCOVERY_ROUTES = [ "/v1/models", "/models", From a1504452d9937c9e31a470016f93853382d2a795 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 22 Jul 2026 18:12:32 +0000 Subject: [PATCH 05/54] fix(proxy): suppress BLE001 for best-effort cache invalidation guards The ruff strict gate flagged the two new blind excepts. They are deliberate: the guards exist so that any cache backend failure, not just an enumerable set of Redis errors, leaves the authoritative team write and the mutation response intact --- litellm/proxy/auth/auth_checks.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 21b455c5d44..72f59238ec3 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1769,7 +1769,7 @@ async def _cache_team_object( if proxy_logging_obj is not None: try: await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key) - except Exception as e: + except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write verbose_proxy_logger.warning( "Failed to invalidate internal usage cache entry %s; " "a stale team object may be served until its TTL expires: %s", @@ -1805,7 +1805,7 @@ async def _cache_team_object( user_api_key_cache.delete_cache(key=alias_key) if proxy_logging_obj is not None: await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key) - except Exception as e: + except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation verbose_proxy_logger.warning( "Failed to invalidate cached team alias entry %s; " "a stale team object may be served until its TTL expires: %s", From b5c363e016cf4bf6c6af89109f084eaa621b62f0 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Wed, 15 Jul 2026 08:24:09 +1000 Subject: [PATCH 06/54] fix(router_strategy): serialize latency for non-chat responses in lowest-latency routing log_success_event/async_log_success_event only converted the end_time - start_time timedelta to float seconds inside the isinstance(response_obj, ModelResponse) branch, so every embedding / speech / image response appended a raw timedelta to the latency list and broke the Redis cache sync with 'Object of type timedelta is not JSON serializable' (no cross-replica latency sharing for those model groups + error-log spam). Normalize response_ms to float seconds up-front in both handlers. Completes the partial fix from #14040. Fixes #33169 Co-Authored-By: Claude Fable 5 --- litellm/router_strategy/lowest_latency.py | 14 +++ .../router_strategy/test_lowest_latency.py | 96 +++++++++++++++++++ 2 files changed, 110 insertions(+) create mode 100644 tests/test_litellm/router_strategy/test_lowest_latency.py diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 23476fe7dcc..14fca48c0d6 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -73,6 +73,13 @@ class LowestLatencyLoggingHandler(CustomLogger): precise_minute = f"{current_date}-{current_hour}-{current_minute}" response_ms = end_time - start_time + if isinstance(response_ms, timedelta): + # normalize to float seconds up-front: non-chat responses + # (embeddings, speech, image) skip the ModelResponse branch + # below, and a raw timedelta appended to the latency list + # breaks JSON serialization when the router cache syncs to + # Redis (issue #33169) + response_ms = response_ms.total_seconds() time_to_first_token_response_time = None if kwargs.get("stream", None) is not None and kwargs["stream"] is True: @@ -262,6 +269,13 @@ class LowestLatencyLoggingHandler(CustomLogger): precise_minute = f"{current_date}-{current_hour}-{current_minute}" response_ms = end_time - start_time + if isinstance(response_ms, timedelta): + # normalize to float seconds up-front: non-chat responses + # (embeddings, speech, image) skip the ModelResponse branch + # below, and a raw timedelta appended to the latency list + # breaks JSON serialization when the router cache syncs to + # Redis (issue #33169) + response_ms = response_ms.total_seconds() time_to_first_token_response_time = None if kwargs.get("stream", None) is not None and kwargs["stream"] is True: # only log ttft for streaming request diff --git a/tests/test_litellm/router_strategy/test_lowest_latency.py b/tests/test_litellm/router_strategy/test_lowest_latency.py new file mode 100644 index 00000000000..4ec89f281e8 --- /dev/null +++ b/tests/test_litellm/router_strategy/test_lowest_latency.py @@ -0,0 +1,96 @@ +#### What this tests #### +# Latency values recorded by lowest-latency routing must be JSON +# serializable for non-chat responses too (embeddings/speech/image skip +# the ModelResponse branch, so the raw timedelta used to leak into the +# latency list and break the Redis cache sync). Issue #33169. + +import json +import os +import sys +from datetime import datetime, timedelta + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.caching.caching import DualCache +from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler + +DEPLOYMENT_ID = "9876" +KWARGS = { + "litellm_params": { + "metadata": { + "model_group": "gemini-embedding-001", + "deployment": "vertex_ai/gemini-embedding-001", + }, + "model_info": {"id": DEPLOYMENT_ID}, + } +} + + +def _embedding_response(): + return litellm.EmbeddingResponse( + model="gemini-embedding-001", + data=[{"embedding": [0.1, 0.2], "index": 0, "object": "embedding"}], + object="list", + usage=litellm.Usage(prompt_tokens=5, completion_tokens=0, total_tokens=5), + ) + + +def _recorded_latencies(cache: DualCache): + cached = cache.get_cache(key="gemini-embedding-001_map") or {} + return cached.get(DEPLOYMENT_ID, {}).get("latency", []) + + +def test_sync_embedding_latency_is_json_serializable(): + """log_success_event with datetime start/end (as the proxy passes) must not + record a raw timedelta for non-ModelResponse results.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + start_time = datetime(2026, 1, 1, 12, 0, 0) + end_time = datetime(2026, 1, 1, 12, 0, 2) + + handler.log_success_event( + response_obj=_embedding_response(), + kwargs=KWARGS, + start_time=start_time, + end_time=end_time, + ) + + latencies = _recorded_latencies(cache) + assert latencies, "expected a latency entry to be recorded" + assert all( + not isinstance(value, timedelta) for value in latencies + ), f"raw timedelta leaked into latency list: {latencies}" + assert latencies[-1] == pytest.approx(2.0) + # the exact failure mode from production: redis cache sync json.dumps + json.dumps({"latency": latencies}) + + +@pytest.mark.asyncio +async def test_async_embedding_latency_is_json_serializable(): + """async_log_success_event is the path the proxy actually hits.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + start_time = datetime(2026, 1, 1, 12, 0, 0) + end_time = datetime(2026, 1, 1, 12, 0, 3) + + await handler.async_log_success_event( + response_obj=_embedding_response(), + kwargs=KWARGS, + start_time=start_time, + end_time=end_time, + ) + + latencies = _recorded_latencies(cache) + assert latencies, "expected a latency entry to be recorded" + assert all( + not isinstance(value, timedelta) for value in latencies + ), f"raw timedelta leaked into latency list: {latencies}" + assert latencies[-1] == pytest.approx(3.0) + json.dumps({"latency": latencies}) From b9a7b807b08e72f5390af89fdca7b55186fa4b81 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Wed, 15 Jul 2026 09:30:53 +1000 Subject: [PATCH 07/54] refactor(router_strategy): drop now-dead timedelta branch after up-front normalization response_ms is normalized to float seconds at the top of both handlers, so the isinstance(response_ms, timedelta) guard inside the ModelResponse branch was unreachable and the Union[float, timedelta] annotation on final_value was wider than reality. Review follow-up, no behavior change. Co-Authored-By: Claude Fable 5 --- litellm/router_strategy/lowest_latency.py | 30 +++++++++-------------- 1 file changed, 12 insertions(+), 18 deletions(-) diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 14fca48c0d6..da81534f389 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -86,7 +86,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # only log ttft for streaming request time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time - final_value: Union[float, timedelta] = response_ms + final_value: float = response_ms time_to_first_token: Optional[float] = None total_tokens = 0 @@ -96,15 +96,12 @@ class LowestLatencyLoggingHandler(CustomLogger): completion_tokens = _usage.completion_tokens total_tokens = _usage.total_tokens - # Handle both timedelta and float response times - if isinstance(response_ms, timedelta): - response_seconds = response_ms.total_seconds() - else: - response_seconds = response_ms + # response_ms is already normalized to float seconds above + response_seconds = response_ms - final_value = safe_divide_seconds(response_seconds, completion_tokens) - if final_value is not None: - final_value = float(final_value) + normalized_value = safe_divide_seconds(response_seconds, completion_tokens) + if normalized_value is not None: + final_value = float(normalized_value) else: final_value = response_seconds @@ -281,7 +278,7 @@ class LowestLatencyLoggingHandler(CustomLogger): # only log ttft for streaming request time_to_first_token_response_time = kwargs.get("completion_start_time", end_time) - start_time - final_value: Union[float, timedelta] = response_ms + final_value: float = response_ms total_tokens = 0 time_to_first_token: Optional[float] = None @@ -291,15 +288,12 @@ class LowestLatencyLoggingHandler(CustomLogger): completion_tokens = _usage.completion_tokens total_tokens = _usage.total_tokens - # Handle both timedelta and float response times - if isinstance(response_ms, timedelta): - response_seconds = response_ms.total_seconds() - else: - response_seconds = response_ms + # response_ms is already normalized to float seconds above + response_seconds = response_ms - final_value = safe_divide_seconds(response_seconds, completion_tokens) - if final_value is not None: - final_value = float(final_value) + normalized_value = safe_divide_seconds(response_seconds, completion_tokens) + if normalized_value is not None: + final_value = float(normalized_value) else: final_value = response_ms From 3cde96d84822316b8ea69330e7585b155ded8592 Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Wed, 15 Jul 2026 12:14:32 +1000 Subject: [PATCH 08/54] refactor(router_strategy): align async else-branch with sync (response_seconds) Review nit: no behavioral difference (response_seconds = response_ms at that point), symmetry only. Co-Authored-By: Claude Fable 5 --- litellm/router_strategy/lowest_latency.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index da81534f389..ffe8245b012 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -295,7 +295,7 @@ class LowestLatencyLoggingHandler(CustomLogger): if normalized_value is not None: final_value = float(normalized_value) else: - final_value = response_ms + final_value = response_seconds if time_to_first_token_response_time is not None: if isinstance(time_to_first_token_response_time, timedelta): From ac8ee71512658a9911a638d8bbfb8746c5d848fe Mon Sep 17 00:00:00 2001 From: Mihidum Hettiyahandi <55163074+mihidumh@users.noreply.github.com> Date: Thu, 16 Jul 2026 12:46:54 +1000 Subject: [PATCH 09/54] test(router_strategy): cover chat-path normalization branches in both handlers Codecov flagged the ModelResponse-branch lines as uncovered: add async per-token normalization, plus zero-completion-token fallback tests for both handlers (the else branch storing plain float seconds). Co-Authored-By: Claude Fable 5 --- .../router_strategy/test_lowest_latency.py | 74 +++++++++++++++++++ 1 file changed, 74 insertions(+) diff --git a/tests/test_litellm/router_strategy/test_lowest_latency.py b/tests/test_litellm/router_strategy/test_lowest_latency.py index 4ec89f281e8..4edc1e21d6b 100644 --- a/tests/test_litellm/router_strategy/test_lowest_latency.py +++ b/tests/test_litellm/router_strategy/test_lowest_latency.py @@ -94,3 +94,77 @@ async def test_async_embedding_latency_is_json_serializable(): ), f"raw timedelta leaked into latency list: {latencies}" assert latencies[-1] == pytest.approx(3.0) json.dumps({"latency": latencies}) + + +def _chat_response(completion_tokens: int): + return litellm.ModelResponse( + model="gpt-4o-mini", + choices=[ + litellm.Choices( + finish_reason="stop", + index=0, + message=litellm.Message(content="hi", role="assistant"), + ) + ], + usage=litellm.Usage( + prompt_tokens=10, + completion_tokens=completion_tokens, + total_tokens=10 + completion_tokens, + ), + ) + + +@pytest.mark.asyncio +async def test_async_chat_latency_normalized_per_token(): + """Chat responses go through the per-token normalization branch — with the + up-front timedelta conversion the stored value must be seconds/token.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + await handler.async_log_success_event( + response_obj=_chat_response(completion_tokens=4), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 2), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(0.5) # 2s / 4 tokens + json.dumps({"latency": latencies}) + + +@pytest.mark.asyncio +async def test_async_chat_zero_completion_tokens_falls_back_to_seconds(): + """safe_divide_seconds returns None for zero tokens — the fallback branch + must store plain float seconds, not a timedelta.""" + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + await handler.async_log_success_event( + response_obj=_chat_response(completion_tokens=0), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 3), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(3.0) + assert not isinstance(latencies[-1], timedelta) + json.dumps({"latency": latencies}) + + +def test_sync_chat_zero_completion_tokens_falls_back_to_seconds(): + cache = DualCache() + handler = LowestLatencyLoggingHandler(router_cache=cache) + + handler.log_success_event( + response_obj=_chat_response(completion_tokens=0), + kwargs=KWARGS, + start_time=datetime(2026, 1, 1, 12, 0, 0), + end_time=datetime(2026, 1, 1, 12, 0, 2), + ) + + latencies = _recorded_latencies(cache) + assert latencies and latencies[-1] == pytest.approx(2.0) + assert not isinstance(latencies[-1], timedelta) + json.dumps({"latency": latencies}) From acd414f18604e99188fc76ce31312bfd1795f7a7 Mon Sep 17 00:00:00 2001 From: Lukas Geiger Date: Sat, 25 Jul 2026 05:18:35 +0000 Subject: [PATCH 10/54] fix(vertex_ai): forward function_call id on Vertex Gemini 3+ tool turns Vertex AI now accepts and returns `id` on functionCall and functionResponse parts for Gemini 3+ on the v1 endpoint, so the provider check added in #28324 is stale. It silently drops the id for every Vertex caller, which breaks strict tool-call matching Gate the id on model version alone, which is what the code did before #28324 and what Google AI Studio already does. `_forward_gemini_function_call_id` no longer takes `custom_llm_provider`, and the decision is resolved once in `_gemini_convert_messages_with_history` and passed to both converters as a bool rather than re-derived independently in each. The context caching path is covered by the same change, since it already passes `model` and the gate needs nothing else The `id` comments on `FunctionCall`, `FunctionResponse` and `HttpxFunctionCall` were also written by #28324 and asserted the opposite of current behaviour, so they are corrected here --- .../prompt_templates/factory.py | 19 +- .../llms/vertex_ai/gemini/transformation.py | 9 +- .../vertex_and_google_ai_studio_gemini.py | 8 +- litellm/types/llms/vertex_ai.py | 10 +- ...test_vertex_and_google_ai_studio_gemini.py | 199 ++++++++++-------- 5 files changed, 134 insertions(+), 111 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index f7ff4d6b16f..dc05a20c37c 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1266,7 +1266,7 @@ def _get_dummy_thought_signature() -> str: def convert_to_gemini_tool_call_invoke( message: ChatCompletionAssistantMessage, model: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + forward_function_call_id: bool = False, ) -> List[VertexPartType]: """ OpenAI tool invokes: @@ -1316,16 +1316,12 @@ def convert_to_gemini_tool_call_invoke( VertexGeminiConfig, ) - forward_tool_call_id = bool( - model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider) - ) - if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper( function_call_params=tool["function"], - tool_call_id=(tool.get("id") if forward_tool_call_id else None), + tool_call_id=(tool.get("id") if forward_function_call_id else None), ) if gemini_function_call is not None: part_dict: VertexPartType = {"function_call": gemini_function_call} @@ -1377,8 +1373,7 @@ def convert_to_gemini_tool_call_invoke( def convert_to_gemini_tool_call_result( message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], last_message_with_tool_calls: Optional[dict], - model: Optional[str] = None, - custom_llm_provider: Optional[str] = None, + forward_function_call_id: bool = False, ) -> Union[VertexPartType, List[VertexPartType]]: """ OpenAI message with a tool result looks like: @@ -1500,14 +1495,8 @@ def convert_to_gemini_tool_call_result( name = tool.get("function", {}).get("name", "") # Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix). - # Only Google AI Studio Gemini 3+ accepts `id` on function_response parts. - # Vertex AI and older Gemini models reject the field with HTTP 400. - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - gemini_call_id: Optional[str] = None - if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider): + if forward_function_call_id: raw_tool_call_id = message.get("tool_call_id") if raw_tool_call_id and isinstance(raw_tool_call_id, str): stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 0db1118a7b4..cbca57c5e62 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history( vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project") vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials") + from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig + + forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "") + try: while msg_i < len(messages): user_content: List[PartType] = [] @@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history( gemini_tool_call_parts = convert_to_gemini_tool_call_invoke( assistant_msg, model=model, - custom_llm_provider=custom_llm_provider, + forward_function_call_id=forward_function_call_id, ) ## check if gemini_tool_call already exists in assistant_content for gemini_tool_call_part in gemini_tool_call_parts: @@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history( _part = convert_to_gemini_tool_call_result( messages[msg_i], # type: ignore last_message_with_tool_calls, # type: ignore - model=model, - custom_llm_provider=custom_llm_provider, + forward_function_call_id=forward_function_call_id, ) msg_i += 1 # Handle both single part and list of parts (for Computer Use with images) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 624190a0b61..2661c0a546f 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool: + def _forward_gemini_function_call_id(model: str) -> bool: """ Whether to include `id` on function_call / function_response parts. - Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict - tool-call matching. Vertex AI rejects the field with HTTP 400. + Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and + Google AI Studio alike. Older Gemini models reject the field with HTTP 400. """ - if custom_llm_provider != "gemini": - return False return VertexGeminiConfig._is_gemini_3_or_newer(model) def _supports_penalty_parameters(self, model: str) -> bool: diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index fb3ddeebf52..da1ff7eda67 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -16,7 +16,7 @@ GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]] class FunctionResponse(TypedDict, total=False): # `id` correlates this response with the originating `functionCall` part. - # Supported on Google AI Studio Gemini 3.5+; Vertex AI rejects this field. + # Supported on Gemini 3+; older Gemini models reject this field. id: str name: Required[str] response: Optional[dict] @@ -24,8 +24,8 @@ class FunctionResponse(TypedDict, total=False): class FunctionCall(TypedDict, total=False): - # `id` correlates the corresponding `functionResponse` on Google AI Studio - # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field. + # `id` correlates the corresponding `functionResponse` on Gemini 3+. + # Older Gemini models omit/reject this field. id: str name: Required[str] args: Optional[dict] @@ -58,8 +58,8 @@ class PartType(TypedDict, total=False): class HttpxFunctionCall(TypedDict, total=False): - # `id` correlates the corresponding `functionResponse` on Google AI Studio - # Gemini 3.5+. Vertex AI and older Gemini models omit/reject this field. + # `id` correlates the corresponding `functionResponse` on Gemini 3+. + # Older Gemini models omit/reject this field. id: str name: Required[str] args: dict diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 95e8e6561f1..8b9c2fbc0a2 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -2273,82 +2273,8 @@ def test_is_gemini_3_or_newer(): assert VertexGeminiConfig._is_gemini_3_or_newer("") == False -def test_forward_gemini_function_call_id_vertex_vs_google_ai_studio(): - """Vertex AI rejects `id` on function_call/function_response; Google AI Studio accepts it on Gemini 3.5+.""" - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - - model = "gemini-3.5-flash" - assert ( - VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai") is False - ) - assert ( - VertexGeminiConfig._forward_gemini_function_call_id(model, "vertex_ai_beta") - is False - ) - assert VertexGeminiConfig._forward_gemini_function_call_id(model, "gemini") is True - assert VertexGeminiConfig._forward_gemini_function_call_id(model, None) is False - assert ( - VertexGeminiConfig._forward_gemini_function_call_id( - "gemini-2.5-flash", "gemini" - ) - is False - ) - - -def test_vertex_ai_gemini_35_tool_calls_omit_function_call_id(): - """Regression: Vertex must not send OpenAI tool_call id inside Gemini function_call parts.""" - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - messages = [ - {"role": "user", "content": "Explore this directory"}, - { - "role": "assistant", - "content": "", - "tool_calls": [ - { - "id": "call_50e7e0fe0989464a89f188eda443", - "type": "function", - "function": { - "name": "read", - "arguments": '{"filePath": "/tmp"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "call_50e7e0fe0989464a89f188eda443", - "content": "ok", - }, - ] - - contents = _gemini_convert_messages_with_history( - messages=messages, - model="gemini-3.5-flash", - custom_llm_provider="vertex_ai", - ) - - for content in contents: - for part in content.get("parts", []): - fc = part.get("function_call") - if fc is not None: - assert "id" not in fc, f"Vertex payload must not include id: {fc}" - fr = part.get("function_response") - if fr is not None: - assert "id" not in fr, f"Vertex payload must not include id: {fr}" - - -def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) - - tool_call_id = "call_50e7e0fe0989464a89f188eda443" - messages = [ +def _tool_call_messages(tool_call_id: str): + return [ {"role": "user", "content": "hi"}, { "role": "assistant", @@ -2371,12 +2297,8 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): }, ] - contents = _gemini_convert_messages_with_history( - messages=messages, - model="gemini-3.5-flash", - custom_llm_provider="gemini", - ) +def _collect_function_call_ids(contents): function_call_ids = [] function_response_ids = [] for content in contents: @@ -2387,9 +2309,120 @@ def test_google_ai_studio_gemini_35_tool_calls_include_function_call_id(): fr = part.get("function_response") if fr is not None: function_response_ids.append(fr.get("id")) + return function_call_ids, function_response_ids - assert function_call_ids == [tool_call_id] - assert function_response_ids == [tool_call_id] + +def test_forward_gemini_function_call_id_is_gated_on_model_version_only(): + """Gemini 3+ takes `id` on Vertex AI and Google AI Studio alike; older models reject it.""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False + assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False + + +@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"]) +def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider): + """Vertex AI accepts `id` on Gemini 3+, so it must be sent there and not just on AI Studio. + + Both parts are asserted together: Vertex pairs a result to its call by id, so emitting one + side without the other would break strict tool-call matching. + """ + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + tool_call_id = "call_50e7e0fe0989464a89f188eda443" + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages(tool_call_id), + model="gemini-3.5-flash", + custom_llm_provider=custom_llm_provider, + ) + + assert _collect_function_call_ids(contents) == ([tool_call_id], [tool_call_id]) + + +@pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "gemini"]) +def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider): + """Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"), + model="gemini-2.5-flash", + custom_llm_provider=custom_llm_provider, + ) + + for content in contents: + for part in content.get("parts", []): + fc = part.get("function_call") + if fc is not None: + assert "id" not in fc, f"gemini-2.5 payload must not include id: {fc}" + fr = part.get("function_response") + if fr is not None: + assert "id" not in fr, f"gemini-2.5 payload must not include id: {fr}" + + +def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix(): + """The thought signature rides along on the OpenAI id but must not reach Vertex. + + Vertex now sees this code path for the first time, so the suffix has to be stripped here too. + """ + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + from litellm.litellm_core_utils.prompt_templates.factory import ( + THOUGHT_SIGNATURE_SEPARATOR, + ) + + bare_id = "call_50e7e0fe0989464a89f188eda443" + contents = _gemini_convert_messages_with_history( + messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"), + model="gemini-3.5-flash", + custom_llm_provider="vertex_ai", + ) + + _, function_response_ids = _collect_function_call_ids(contents) + assert function_response_ids == [bare_id] + + +@pytest.mark.parametrize("model", ["gemini-3.5-flash", "gemini-2.5-flash"]) +def test_tool_response_without_matching_tool_call_is_rejected(model): + """An unpairable tool result must raise, not ship a functionResponse with no matching call.""" + from litellm.llms.vertex_ai.gemini.transformation import ( + _gemini_convert_messages_with_history, + ) + + messages = [ + {"role": "user", "content": "hi"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_50e7e0fe0989464a89f188eda443", + "type": "function", + "function": { + "name": "read", + "arguments": '{"filePath": "/tmp"}', + }, + } + ], + }, + {"role": "tool", "content": "ok"}, + ] + + with pytest.raises(Exception, match="Missing corresponding tool call"): + _gemini_convert_messages_with_history( + messages=messages, + model=model, + custom_llm_provider="vertex_ai", + ) def test_reasoning_effort_maps_to_thinking_level_gemini_3(): From 5f50791e99187ce58428c375cc17c55e07fa0a83 Mon Sep 17 00:00:00 2001 From: shivam Date: Mon, 27 Jul 2026 23:33:01 +0000 Subject: [PATCH 11/54] fix(vertex_ai): source managed-file read bucket + credentials from per-model litellm_params Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/vertex_ai/files/handler.py | 46 ++++- .../files/test_vertex_ai_files_handler.py | 183 +++++++++++++----- 2 files changed, 177 insertions(+), 52 deletions(-) diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index 3bc09139f8f..4d2a1e18eb5 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -1,7 +1,9 @@ import asyncio +import json +import os import time from urllib.parse import unquote -from typing import Any, Coroutine, Optional, Tuple, Union +from typing import Any, Coroutine, Mapping, Optional, Tuple, Union import httpx @@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import ( GCSBucketBase, GCSLoggingConfig, ) +from litellm.types.utils import StandardCallbackDynamicParams from litellm.litellm_core_utils.cloud_storage_security import ( VERTEX_AI_MANAGED_GCS_PREFIX, should_allow_legacy_cloud_file_ids, @@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase): llm_provider=LlmProviders.VERTEX_AI, ) + def _resolve_read_gcs_config( + self, + litellm_params: Mapping[str, object] | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + ) -> tuple[str | None, str | None]: + """ + Resolve the GCS bucket and service-account credentials for the read/content path. + + Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` / + ``bucket_name`` and ``vertex_credentials``), mirroring the write path in + ``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global + ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch + run entirely at the model-group level, so output written to a per-model bucket is + readable without setting the global env vars. + """ + params: Mapping[str, object] = litellm_params or {} + bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name") + configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME") + + credentials = params.get("vertex_credentials") or vertex_credentials + if isinstance(credentials, dict): + path_service_account: str | None = json.dumps(credentials) + elif isinstance(credentials, str): + path_service_account = credentials + else: + path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT") + + return configured_bucket_name, path_service_account + def _extract_bucket_and_object_from_file_id( self, file_id: str, @@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase): if not file_id: raise ValueError("file_id is required in file_content_request") - gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={}) + configured_bucket_name, path_service_account = self._resolve_read_gcs_config( + litellm_params=litellm_params, + vertex_credentials=vertex_credentials, + ) + dynamic_params = StandardCallbackDynamicParams( + gcs_bucket_name=configured_bucket_name, + gcs_path_service_account=path_service_account, + ) + gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config( + kwargs={"standard_callback_dynamic_params": dynamic_params} + ) bucket_name, object_path = self._extract_bucket_and_object_from_file_id( file_id=file_id, configured_bucket_name=gcs_logging_config["bucket_name"], diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py index 453a0c14bf9..5e854bbad70 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_handler.py @@ -31,10 +31,7 @@ class TestVertexAIFilesHandler: def test_extract_bucket_and_object_from_file_id_standard_path(self): """Test extraction of bucket and object from URL-encoded file_id with standard path""" # Sample file_id with nested folder structure - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Ftest-folder%2Fsub-folder%2Ftest-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Ftest-folder%2Fsub-folder%2Ftest-file.txt" bucket_name, object_path = self.handler._extract_bucket_and_object_from_file_id( file_id=file_id, @@ -105,21 +102,14 @@ class TestVertexAIFilesHandler: async def test_afile_content_success(self): """Test successful async file content retrieval""" # Setup test data - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Fuploads%2Fabc-test-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock the download_gcs_object method with ( - patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download, + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, patch.object( self.handler, "get_gcs_logging_config", @@ -148,15 +138,9 @@ class TestVertexAIFilesHandler: # Verify the download was called with correct parameters mock_download.assert_called_once() call_args = mock_download.call_args - assert ( - call_args.kwargs["object_name"] - == "litellm-vertex-files/uploads/abc-test-file.txt" - ) + assert call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-test-file.txt" assert "standard_callback_dynamic_params" in call_args.kwargs - assert ( - call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] - == "test-bucket" - ) + assert call_args.kwargs["standard_callback_dynamic_params"]["gcs_bucket_name"] == "test-bucket" @pytest.mark.asyncio async def test_afile_content_missing_file_id(self): @@ -164,9 +148,7 @@ class TestVertexAIFilesHandler: file_content_request = FileContentRequest(extra_headers=None, extra_body=None) # Should raise ValueError for missing file_id - with pytest.raises( - ValueError, match="file_id is required in file_content_request" - ): + with pytest.raises(ValueError, match="file_id is required in file_content_request"): await self.handler.afile_content( file_content_request=file_content_request, vertex_credentials=None, @@ -179,20 +161,13 @@ class TestVertexAIFilesHandler: @pytest.mark.asyncio async def test_afile_content_download_failure(self): """Test async file content retrieval when download fails""" - file_id = ( - "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files" - "%2Fuploads%2Fabc-test-file.txt" - ) + file_id = "gs%3A%2F%2Ftest-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-test-file.txt" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock download to return None (failure) with ( - patch.object( - self.handler, "download_gcs_object", new_callable=AsyncMock - ) as mock_download, + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, patch.object( self.handler, "get_gcs_logging_config", @@ -216,14 +191,130 @@ class TestVertexAIFilesHandler: max_retries=3, ) + def test_resolve_read_gcs_config_prefers_per_model_bucket(self, monkeypatch): + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json") + + bucket, service_account = self.handler._resolve_read_gcs_config( + litellm_params={ + "gcs_bucket_name": "my-model-bucket", + "vertex_credentials": "/model/sa.json", + }, + vertex_credentials=None, + ) + + assert bucket == "my-model-bucket" + assert service_account == "/model/sa.json" + + def test_resolve_read_gcs_config_falls_back_to_env(self, monkeypatch): + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/env/sa.json") + + bucket, service_account = self.handler._resolve_read_gcs_config(litellm_params={}, vertex_credentials=None) + + assert bucket == "env-default-bucket" + assert service_account == "/env/sa.json" + + def test_resolve_read_gcs_config_serializes_dict_credentials(self, monkeypatch): + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + _, service_account = self.handler._resolve_read_gcs_config( + litellm_params={"gcs_bucket_name": "my-model-bucket"}, + vertex_credentials={"type": "service_account", "project_id": "p"}, + ) + + assert service_account == '{"type": "service_account", "project_id": "p"}' + + @pytest.mark.asyncio + async def test_afile_content_honors_per_model_bucket_over_env(self, monkeypatch): + """ + Regression for #32640: a batch output written to a per-model gcs_bucket_name must be + readable even when the global GCS_BUCKET_NAME points at a different bucket. Before the + fix the read path resolved the bucket from env only and raised + "file_id bucket does not match the configured storage bucket". + """ + monkeypatch.setenv("GCS_BUCKET_NAME", "env-default-bucket") + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl" + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) + + with ( + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, + patch.object( + self.handler, + "get_or_create_vertex_instance", + new_callable=AsyncMock, + return_value=object(), + ), + ): + mock_download.return_value = b"batch output" + + result = await self.handler.afile_content( + file_content_request=file_content_request, + vertex_credentials="/model/sa.json", + vertex_project="test-project", + vertex_location="us-central1", + timeout=60.0, + max_retries=0, + litellm_params={ + "gcs_bucket_name": "my-model-bucket", + "vertex_credentials": "/model/sa.json", + }, + ) + + assert isinstance(result, HttpxBinaryResponseContent) + assert result.response.content == b"batch output" + + dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"] + assert dynamic_params["gcs_bucket_name"] == "my-model-bucket" + assert dynamic_params["gcs_path_service_account"] == "/model/sa.json" + assert mock_download.call_args.kwargs["object_name"] == "litellm-vertex-files/uploads/abc-batch-output.jsonl" + + @pytest.mark.asyncio + async def test_afile_content_reads_without_global_env_bucket(self, monkeypatch): + """ + Regression for #32640: with no global GCS_BUCKET_NAME set, a model-group-level + deployment (per-model gcs_bucket_name) must still be readable. Before the fix the read + path raised "GCS_BUCKET_NAME is not set in the environment". + """ + monkeypatch.delenv("GCS_BUCKET_NAME", raising=False) + monkeypatch.delenv("GCS_PATH_SERVICE_ACCOUNT", raising=False) + + file_id = "gs%3A%2F%2Fmy-model-bucket%2Flitellm-vertex-files%2Fuploads%2Fabc-batch-output.jsonl" + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) + + with ( + patch.object(self.handler, "download_gcs_object", new_callable=AsyncMock) as mock_download, + patch.object( + self.handler, + "get_or_create_vertex_instance", + new_callable=AsyncMock, + return_value=object(), + ), + ): + mock_download.return_value = b"batch output" + + result = await self.handler.afile_content( + file_content_request=file_content_request, + vertex_credentials="/model/sa.json", + vertex_project="test-project", + vertex_location="us-central1", + timeout=60.0, + max_retries=0, + litellm_params={"gcs_bucket_name": "my-model-bucket"}, + ) + + assert isinstance(result, HttpxBinaryResponseContent) + dynamic_params = mock_download.call_args.kwargs["standard_callback_dynamic_params"] + assert dynamic_params["gcs_bucket_name"] == "my-model-bucket" + def test_file_content_sync_success(self): """Test successful sync file content retrieval""" file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Create expected response mock_response = httpx.Response( @@ -261,25 +352,17 @@ class TestVertexAIFilesHandler: file_id = "gs%3A%2F%2Ftest-bucket%2Ftest-file.txt" expected_content = b"test file content" - file_content_request = FileContentRequest( - file_id=file_id, extra_headers=None, extra_body=None - ) + file_content_request = FileContentRequest(file_id=file_id, extra_headers=None, extra_body=None) # Mock the afile_content method - with patch.object( - self.handler, "afile_content", new_callable=AsyncMock - ) as mock_afile_content: + with patch.object(self.handler, "afile_content", new_callable=AsyncMock) as mock_afile_content: mock_response = httpx.Response( status_code=200, content=expected_content, headers={"content-type": "application/octet-stream"}, - request=httpx.Request( - method="GET", url="gs://test-bucket/test-file.txt" - ), - ) - mock_afile_content.return_value = HttpxBinaryResponseContent( - response=mock_response + request=httpx.Request(method="GET", url="gs://test-bucket/test-file.txt"), ) + mock_afile_content.return_value = HttpxBinaryResponseContent(response=mock_response) # Call the method with _is_async=True result = self.handler.file_content( From c2f0014a632f465f5c242085718873f1de3e6d4f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 27 Jul 2026 17:15:09 -0700 Subject: [PATCH 12/54] fix(jwt_auth): grant only /v1/messages routes to JWT teams by default, not all anthropic_routes --- litellm/proxy/_types.py | 8 ++++++- .../proxy/auth/test_handle_jwt.py | 22 +++++++++++++++++++ 2 files changed, 29 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 53b6e51aaef..1608811b5a1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4254,7 +4254,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): team_id_upsert: bool = False team_ids_jwt_field: Optional[str] = None upsert_sso_user_to_team: bool = False - team_allowed_routes: List[str] = ["openai_routes", "anthropic_routes", "info_routes", "mcp_routes"] + team_allowed_routes: List[str] = [ + "openai_routes", + "info_routes", + "mcp_routes", + "/v1/messages", + "/v1/messages/count_tokens", + ] team_id_default: Optional[str] = Field( default=None, description="If no team_id given, default permissions/spend-tracking to this team.s", diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 35aaa0f1254..3840c90d691 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1299,6 +1299,28 @@ async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatc assert team_obj.team_id == "coding-team" +@pytest.mark.parametrize( + "route,expected", + [ + ("/v1/messages", True), + ("/v1/messages/count_tokens", True), + ("/v1/skills", False), + ("/v1/skills/skill_abc123", False), + ], +) +def test_default_team_allowed_routes_cover_messages_but_not_skills(route, expected): + from litellm.proxy.auth.auth_checks import allowed_routes_check + + assert ( + allowed_routes_check( + user_role=LitellmUserRoles.TEAM, + user_route=route, + litellm_proxy_roles=LiteLLM_JWTAuth(), + ) + is expected + ) + + @pytest.mark.asyncio async def test_auth_builder_returns_team_membership_object(): """ From ddaee8df164781833f2358fd914aff5b776c63ec Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 9 Jul 2026 06:15:18 +0000 Subject: [PATCH 13/54] fix(auth): resolve managed batch/file deployment model_id to model name for team access checks --- litellm/proxy/auth/auth_utils.py | 7 +++- litellm/router.py | 7 +++- .../test_router_helper_utils.py | 17 +++++++++ .../proxy/auth/test_auth_utils.py | 37 +++++++++++++++++++ 4 files changed, 65 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index ecb37e67c14..644253ceac7 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id( ) _append_model_candidates( candidates=candidates, - value=get_model_id_from_unified_batch_id(unified_file_id), + value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router), ) except Exception as e: verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e)) @@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id( parsed_id = parse_unified_id(resource_id) if parsed_id: - _append_model_candidates(candidates=candidates, value=parsed_id.get("model_id")) + _append_model_candidates( + candidates=candidates, + value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router), + ) _append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names")) except Exception as e: verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e)) diff --git a/litellm/router.py b/litellm/router.py index 487d6a31226..e2ed320f089 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9519,7 +9519,12 @@ class Router: return None # Strategy 1: Check if model_id directly matches a model_name or deployment ID - if model_id in self.model_names or self.has_model_id(model_id): + if model_id in self.model_names: + return model_id + if self.has_model_id(model_id): + deployment = self.get_deployment(model_id=model_id) + if deployment is not None and deployment.model_name: + return deployment.model_name return model_id # Strategy 2: Search through router's model_list to find by litellm_params.model diff --git a/tests/router_unit_tests/test_router_helper_utils.py b/tests/router_unit_tests/test_router_helper_utils.py index a969d21a681..bcc70fae67c 100644 --- a/tests/router_unit_tests/test_router_helper_utils.py +++ b/tests/router_unit_tests/test_router_helper_utils.py @@ -2659,6 +2659,23 @@ def test_resolve_model_name_from_model_id(): result = router.resolve_model_name_from_model_id("gpt-5-mini") assert result == "gpt-5-mini" + # Test case 10: model_id is a deployment ID (hash) that differs from the + # public model_name. Regression for #32580: managed batch/file IDs embed the + # deployment model_id, and it must resolve back to the public model_name so + # team model-access checks compare against the model group, not the hash. + model_list = [ + { + "model_name": "bedrock-batch-model", + "litellm_params": { + "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0", + }, + "model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"}, + }, + ] + router = Router(model_list=model_list) + result = router.resolve_model_name_from_model_id("8d0eaa7e6c6f54a425dfd0062cb6b0dc") + assert result == "bedrock-batch-model" + def test_get_valid_args(): """Test get_valid_args static method returns valid Router.__init__ arguments""" diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 9f24c662581..8b523c34e84 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -569,6 +569,43 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): ) +def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(): + """Regression for #32580: managed batch retrieve/cancel encode the deployment + model_id (a sha256 hash) into the batch id. The auth layer must resolve that + hash back to the public model group name so team model-access checks compare + against the model group, not the raw deployment hash.""" + import base64 + + from litellm.router import Router + + router = Router( + model_list=[ + { + "model_name": "bedrock-batch-model", + "litellm_params": { + "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0", + }, + "model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"}, + } + ] + ) + + decoded_batch_id = ( + "litellm_proxy;model_id:8d0eaa7e6c6f54a425dfd0062cb6b0dc;" + "llm_batch_id:provider-batch-123" + ) + batch_id = base64.urlsafe_b64encode(decoded_batch_id.encode()).decode().rstrip("=") + + assert ( + get_model_from_request( + request_data={"batch_id": batch_id}, + route="/v1/batches/{batch_id}", + llm_router=router, + ) + == "bedrock-batch-model" + ) + + def test_get_model_from_request_resolves_character_id_model_with_router(): from litellm.types.videos.utils import encode_character_id_with_provider From e3559cf1b701980eb40948253315f69c00b30905 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 28 Jul 2026 15:39:50 +0000 Subject: [PATCH 14/54] test(auth): cover managed batch/file team access denial end to end Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/auth/test_auth_utils.py | 94 +++++++++++++++---- 1 file changed, 77 insertions(+), 17 deletions(-) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 8b523c34e84..1610d76efb7 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -569,43 +569,103 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): ) -def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(): - """Regression for #32580: managed batch retrieve/cancel encode the deployment - model_id (a sha256 hash) into the batch id. The auth layer must resolve that - hash back to the public model group name so team model-access checks compare - against the model group, not the raw deployment hash.""" - import base64 +_BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc" + +def _managed_batch_router(): from litellm.router import Router - router = Router( + return Router( model_list=[ { "model_name": "bedrock-batch-model", "litellm_params": { "model": "bedrock/global.anthropic.claude-haiku-4-5-20251001-v1:0", }, - "model_info": {"id": "8d0eaa7e6c6f54a425dfd0062cb6b0dc"}, - } + "model_info": {"id": _BATCH_DEPLOYMENT_ID}, + }, + { + "model_name": "some-other-model", + "litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"}, + "model_info": {"id": "a-different-deployment-id"}, + }, ] ) - decoded_batch_id = ( - "litellm_proxy;model_id:8d0eaa7e6c6f54a425dfd0062cb6b0dc;" - "llm_batch_id:provider-batch-123" - ) - batch_id = base64.urlsafe_b64encode(decoded_batch_id.encode()).decode().rstrip("=") +def _encode_managed_id(decoded: str) -> str: + return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=") + + +_MANAGED_BATCH_ID = _encode_managed_id( + f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123" +) +_MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id( + f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;" + "llm_output_file_id:provider-file-456" +) + + +@pytest.mark.parametrize( + "route, request_data", + [ + ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/files/{file_id}", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ], +) +def test_get_model_from_request_resolves_batch_id_deployment_to_model_name(route, request_data): + """Regression for #32580: managed batch retrieve/cancel and managed batch output + file reads encode the deployment model_id into the resource id. The auth layer must + resolve that id back to the public model group name so model-access checks compare + against the model group, not the raw deployment id.""" assert ( get_model_from_request( - request_data={"batch_id": batch_id}, - route="/v1/batches/{batch_id}", - llm_router=router, + request_data=request_data, + route=route, + llm_router=_managed_batch_router(), ) == "bedrock-batch-model" ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "route, request_data", + [ + ("/v1/batches/{batch_id}", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/batches/{batch_id}/cancel", {"batch_id": _MANAGED_BATCH_ID}), + ("/v1/files/{file_id}/content", {"file_id": _MANAGED_BATCH_OUTPUT_FILE_ID}), + ], +) +async def test_managed_batch_routes_pass_team_model_access_check(route, request_data): + """End-to-end regression for #32580: a team scoped to the batch model group got + ``team_model_access_denied`` on retrieve/cancel because the deployment id, not the + model group, was authorized. Fails pre-fix with the deployment id in the message.""" + from litellm.proxy._types import LiteLLM_TeamTable + from litellm.proxy.auth.auth_checks import can_team_access_model + + llm_router = _managed_batch_router() + model = get_model_from_request(request_data=request_data, route=route, llm_router=llm_router) + + assert ( + await can_team_access_model( + model=model, + team_object=LiteLLM_TeamTable(team_id="team-batch", models=["bedrock-batch-model"]), + llm_router=llm_router, + ) + is True + ) + + with pytest.raises(Exception, match="team not allowed to access model"): + await can_team_access_model( + model=model, + team_object=LiteLLM_TeamTable(team_id="team-other", models=["some-other-model"]), + llm_router=llm_router, + ) + + def test_get_model_from_request_resolves_character_id_model_with_router(): from litellm.types.videos.utils import encode_character_id_with_provider From 5e1d9705dbe09e059f418d222b790278fcb16291 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 19:45:00 +0000 Subject: [PATCH 15/54] fix(proxy): skip team model aliases that point at deleted deployments A team's model_aliases can map a public name like gpt-4 to the internal routing key (model_name_{team_id}_{uuid}) of a team deployment that has since been deleted, e.g. after replacing per-team duplicates with one gateway-level model. The pre-call rewrite then sent every request to a name the router cannot serve, failing with "no healthy deployments for model_name_..." even though the requested name still resolves at the gateway level. The rewrite is now skipped when the alias target has no live deployment in the router delete_model also skipped the team alias scan for internal-shaped names on the assumption they can never be alias values, which is exactly the shape legacy team model aliases have, so deleting a legacy team model left the stale alias behind. The scan now always runs, and a public name that still resolves to a live router deployment (e.g. a shared gateway-level model group) stays in team.models so the delete does not revoke the team's access to it --- litellm/proxy/litellm_pre_call_utils.py | 99 ++++++++++++------- .../model_management_endpoints.py | 45 ++++++--- tests/proxy_unit_tests/test_proxy_utils.py | 6 +- .../test_model_management_endpoints.py | 95 +++++++++++++++++- .../proxy/test_litellm_pre_call_utils.py | 61 ++++++++++++ 5 files changed, 250 insertions(+), 56 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d94fed0ee5b..c40718e41fd 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1829,6 +1829,15 @@ async def add_litellm_data_to_request( return data +def _warn_stale_team_alias_once(warning_key: str, message: str, *args: str) -> None: + if warning_key in _STALE_TEAM_ALIAS_WARNING_KEYS: + return + _STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None + while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS: + _STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False) + verbose_proxy_logger.warning(message, *args) + + def _update_model_if_team_alias_exists( data: dict, user_api_key_dict: UserAPIKeyAuth, @@ -1848,49 +1857,63 @@ def _update_model_if_team_alias_exists( Note: model_aliases for team models are deprecated. This function only applies to legacy non-team-scoped aliases. Team-scoped deployments use team_public_model_name and are resolved via map_team_model in route_llm_request. + + An alias that targets a team-scoped internal name (``model_name_{team_id}_{uuid}``) + with no live deployment behind it is never applied: the deployment was deleted, so + the rewrite could only fail with an error naming a model the caller never sent. + Keeping the requested model name lets it resolve against the deployments that still + exist (e.g. a gateway-level model group shared with the team). """ _model = data.get("model") - if _model and user_api_key_dict.team_model_aliases and _model in user_api_key_dict.team_model_aliases: - from litellm.proxy.proxy_server import llm_router + if not _model or not user_api_key_dict.team_model_aliases or _model not in user_api_key_dict.team_model_aliases: + return - # Skip alias rewrite if this model resolves to team-specific deployments - # (team models use team_public_model_name, not model_aliases) - aliased_target = user_api_key_dict.team_model_aliases[_model] + from litellm.proxy.proxy_server import llm_router - # Optional bypass for stale aliases from pre-PR deployments: - # only enabled via feature flag to preserve backwards compatibility. - # Cached at module level to avoid hot-path secret lookups on every request. - global _ENABLE_TEAM_STALE_ALIAS_BYPASS - if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None: - _ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False) - enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS - # Check if the alias points to a team-scoped UUID name - # (format: "model_name_{team_id}_{uuid}") - is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_") - if is_stale_team_alias and llm_router: - # This is a stale alias from pre-PR deployments. - # Check if current team deployments exist for the public name. - key = (user_api_key_dict.team_id, _model) - if key in llm_router.team_model_to_deployment_indices: - if enable_stale_alias_bypass: - # Team deployments exist; skip stale alias - return - warning_key = f"{user_api_key_dict.team_id}:{_model}:{aliased_target}" - if warning_key not in _STALE_TEAM_ALIAS_WARNING_KEYS: - _STALE_TEAM_ALIAS_WARNING_KEYS[warning_key] = None - while len(_STALE_TEAM_ALIAS_WARNING_KEYS) > _MAX_STALE_ALIAS_WARNING_KEYS: - _STALE_TEAM_ALIAS_WARNING_KEYS.popitem(last=False) - verbose_proxy_logger.warning( - "Stale team model alias detected for model='%s', team_id='%s'. " - "New sibling deployments may be unreachable. " - "Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable " - "team-scoped sibling routing.", - _sanitize_for_log(_model), - user_api_key_dict.team_id, - ) + # Skip alias rewrite if this model resolves to team-specific deployments + # (team models use team_public_model_name, not model_aliases) + aliased_target = user_api_key_dict.team_model_aliases[_model] - data["model"] = aliased_target - return + # Optional bypass for stale aliases from pre-PR deployments: + # only enabled via feature flag to preserve backwards compatibility. + # Cached at module level to avoid hot-path secret lookups on every request. + global _ENABLE_TEAM_STALE_ALIAS_BYPASS + if _ENABLE_TEAM_STALE_ALIAS_BYPASS is None: + _ENABLE_TEAM_STALE_ALIAS_BYPASS = get_secret_bool("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", False) + enable_stale_alias_bypass = _ENABLE_TEAM_STALE_ALIAS_BYPASS + # Check if the alias points to a team-scoped UUID name + # (format: "model_name_{team_id}_{uuid}") + is_stale_team_alias = aliased_target.startswith(f"model_name_{user_api_key_dict.team_id}_") + if is_stale_team_alias and llm_router: + if aliased_target not in llm_router.model_name_to_deployment_indices: + _warn_stale_team_alias_once( + f"deleted:{user_api_key_dict.team_id}:{_model}:{aliased_target}", + "Team model alias for model='%s', team_id='%s' targets '%s', which has no live " + "deployment. Routing with the requested model name instead; remove the stale " + "entry from the team's model_aliases to silence this warning.", + _sanitize_for_log(_model), + _sanitize_for_log(user_api_key_dict.team_id), + _sanitize_for_log(aliased_target), + ) + return + # This is a stale alias from pre-PR deployments. + # Check if current team deployments exist for the public name. + key = (user_api_key_dict.team_id, _model) + if key in llm_router.team_model_to_deployment_indices: + if enable_stale_alias_bypass: + # Team deployments exist; skip stale alias + return + _warn_stale_team_alias_once( + f"{user_api_key_dict.team_id}:{_model}:{aliased_target}", + "Stale team model alias detected for model='%s', team_id='%s'. " + "New sibling deployments may be unreachable. " + "Set LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS=true to enable " + "team-scoped sibling routing.", + _sanitize_for_log(_model), + _sanitize_for_log(user_api_key_dict.team_id), + ) + + data["model"] = aliased_target def _update_model_if_key_alias_exists( diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 91f0b9b2790..ddb7f08ad15 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -52,6 +52,7 @@ from litellm.proxy.utils import PrismaClient from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ModelTableRepository from litellm.repositories.team_repository import TeamRepository +from litellm.router import Router from litellm.types.proxy.management_endpoints.model_management_endpoints import ( UpdateUsefulLinksRequest, ) @@ -788,6 +789,7 @@ async def _remove_unbacked_team_models( prisma_client: PrismaClient, user_api_key_cache: Any, proxy_logging_obj: Any, + llm_router: Router | None = None, ) -> None: """ Strip a deleted team model's public name(s) from team.models and refresh the cache. @@ -795,26 +797,40 @@ async def _remove_unbacked_team_models( Must be called after the deployment row is deleted: a public name is removed only when no remaining team deployment still backs it, so a load-balanced replica isn't revoked while siblings serve it, and concurrent deletes can't leave a ghost. + + Legacy team models (created before team_public_model_name existed) store a + ``{public_name: "model_name_{team_id}_{uuid}"}`` entry in the team's model_aliases, + so the alias scan runs for every team model; skipping it for internal-shaped names + left stale aliases that rewrote requests to deployments that no longer exist. + + A public name that still resolves to a live router deployment (e.g. a gateway-level + model group shared with the team) is kept in team.models, so deleting a per-team + duplicate does not revoke the team's access to the shared deployment. """ team_id = model_params.model_info.team_id if team_id is None: return - # BYOK models carry an internal `model_name_{team_id}_{uuid}` name that can never - # be a team alias value, so skip the full litellm_modeltable scan for them. - removed_model_aliases: List[Tuple[str, str]] = [] - if not model_params.model_name.startswith(f"model_name_{team_id}_"): - removed_model_aliases = await delete_team_model_alias( - public_model_name=model_params.model_name, - prisma_client=prisma_client, - ) - names_to_remove = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id} - if model_params.model_info.team_public_model_name is not None: - names_to_remove.add(model_params.model_info.team_public_model_name) - - if names_to_remove: - names_to_remove -= await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client) + removed_model_aliases = await delete_team_model_alias( + public_model_name=model_params.model_name, + prisma_client=prisma_client, + ) + removed_alias_names = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id} + candidate_names = ( + removed_alias_names | {model_params.model_info.team_public_model_name} + if model_params.model_info.team_public_model_name is not None + else removed_alias_names + ) + if not candidate_names: + return + team_backed_names = await _get_team_public_model_names(team_id=team_id, prisma_client=prisma_client) + router_served_names = ( + frozenset(name for name in candidate_names if name in llm_router.model_name_to_deployment_indices) + if llm_router is not None + else frozenset() + ) + names_to_remove = candidate_names - team_backed_names - router_served_names if not names_to_remove: return @@ -1120,6 +1136,7 @@ async def delete_model( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + llm_router=llm_router, ) ## CREATE AUDIT LOG ## diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index d2218b08386..ad852c16905 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -2184,7 +2184,8 @@ def test_team_alias_stale_bypass_disabled_by_default(monkeypatch): pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None class _MockRouter: - team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]} + model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]} + team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]} test_data = {"model": "gpt-4o"} user_api_key_dict = UserAPIKeyAuth( @@ -2209,7 +2210,8 @@ def test_team_alias_stale_bypass_enabled_by_flag(monkeypatch): pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None class _MockRouter: - team_model_to_deployment_indices = {("team-1", "gpt-4o"): [0]} + model_name_to_deployment_indices = {"model_name_team-1_legacy-uuid": [0]} + team_model_to_deployment_indices = {("team-1", "gpt-4o"): [1]} test_data = {"model": "gpt-4o"} user_api_key_dict = UserAPIKeyAuth( diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index bd5eda0197b..54c17a845d2 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -2031,8 +2031,7 @@ class TestDeleteTeamBYOKModelGhost: mock_refresh.assert_awaited_once() assert mock_refresh.await_args.kwargs["team_row"] is updated_team_row - # BYOK internal name can't be an alias value -> the alias-table scan is skipped. - mock_prisma.db.litellm_modeltable.find_many.assert_not_awaited() + mock_prisma.db.litellm_modeltable.find_many.assert_awaited() @pytest.mark.asyncio async def test_delete_non_internal_team_model_still_scans_aliases(self): @@ -2186,6 +2185,98 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db.litellm_teamtable.update.assert_not_awaited() mock_refresh.assert_not_awaited() + @pytest.mark.asyncio + async def test_delete_legacy_team_model_scrubs_stale_alias_and_keeps_gateway_access( + self, + ): + """Regression: legacy team models store {public_name: internal model_name} in + the team's model_aliases. delete_model skipped the alias scan for + internal-shaped names, so the stale alias kept rewriting requests for the + public name to a deployment that no longer existed ("no healthy deployments + for model_name_{team_id}_..."). Deleting the deployment must scrub the + alias, and the public name must stay in team.models while a gateway-level + deployment still serves it.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + team_id = "team-legacy-alias" + model_id = "legacy-alias-model-1" + public_name = "gpt-4" + internal_name = f"model_name_{team_id}_abc-uuid" + + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=internal_name, + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={"id": model_id, "team_id": team_id}, + created_by="admin", + updated_by="admin", + ) + team_row = LiteLLM_TeamTable( + team_id=team_id, + team_alias="legacy-alias-team", + members_with_roles=[Member(user_id="admin", role="admin")], + models=[public_name], + ) + alias_row = MagicMock( + id="alias-row-1", model_aliases={public_name: internal_name} + ) + alias_row.team = MagicMock() + alias_row.team.team_id = team_id + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock( + return_value=[alias_row] + ) + mock_prisma.db.litellm_modeltable.update = AsyncMock() + + mock_router = MagicMock() + mock_router.model_name_to_deployment_indices = {public_name: [0]} + + admin_user = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()) as mock_refresh, + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=admin_user, + ) + + assert "deleted successfully" in result["message"] + + mock_prisma.db.litellm_modeltable.update.assert_awaited_once() + alias_update_kwargs = mock_prisma.db.litellm_modeltable.update.await_args.kwargs + assert alias_update_kwargs["where"] == {"id": "alias-row-1"} + assert json.loads(alias_update_kwargs["data"]["model_aliases"]) == {} + + # A gateway-level deployment still serves the public name -> team access stays. + mock_prisma.db.litellm_teamtable.update.assert_not_awaited() + mock_refresh.assert_not_awaited() + class TestDeleteModelTeamAuth: """Team auth on the /model/delete path. diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 143884c8a0d..36fde1b10d8 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5457,3 +5457,64 @@ def test_get_sanitized_user_information_from_key_drops_callback_config(): # UserAPIKeyAuth is the live auth object; the per-key callbacks are resolved # from it during pre-call, so it must not be mutated by building the log view assert "logging" in (user_api_key_dict.metadata or {}) + + +def test_team_alias_targeting_deleted_team_deployment_keeps_requested_model(monkeypatch): + """ + Regression: a team's model_aliases can point at the internal routing key + (model_name_{team_id}_{uuid}) of a team deployment that was since deleted, + e.g. after an admin replaces per-team duplicates with one gateway-level + model. Rewriting to the dead internal name made every request fail with + "no healthy deployments for model_name_..." even though the requested + public name resolves at the gateway level. The rewrite must be skipped + when the alias target has no live deployment. + """ + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists + + monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False) + pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None + + class _MockRouter: + model_name_to_deployment_indices = {"gpt-4": [0]} + team_model_to_deployment_indices = {} + + test_data = {"model": "gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_key", + team_id="team-1", + team_model_aliases={"gpt-4": "model_name_team-1_dead-uuid"}, + ) + + with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()): + _update_model_if_team_alias_exists( + data=test_data, user_api_key_dict=user_api_key_dict + ) + + assert test_data.get("model") == "gpt-4" + + +def test_team_alias_targeting_live_team_deployment_still_rewrites(monkeypatch): + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + from litellm.proxy.litellm_pre_call_utils import _update_model_if_team_alias_exists + + monkeypatch.delenv("LITELLM_ENABLE_TEAM_STALE_ALIAS_BYPASS", raising=False) + pre_call_utils._ENABLE_TEAM_STALE_ALIAS_BYPASS = None + + class _MockRouter: + model_name_to_deployment_indices = {"model_name_team-1_live-uuid": [0]} + team_model_to_deployment_indices = {} + + test_data = {"model": "gpt-4"} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_key", + team_id="team-1", + team_model_aliases={"gpt-4": "model_name_team-1_live-uuid"}, + ) + + with patch("litellm.proxy.proxy_server.llm_router", _MockRouter()): + _update_model_if_team_alias_exists( + data=test_data, user_api_key_dict=user_api_key_dict + ) + + assert test_data.get("model") == "model_name_team-1_live-uuid" From d409fec6de277521f76c170bae42bdb9584f9e19 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 20:07:30 +0000 Subject: [PATCH 16/54] fix(proxy): keep team model aliases while a surviving replica serves the deleted name Scrub aliases on delete only when the deleted deployment's model_name no longer resolves in the router. A legacy load-balanced team model can have several deployment rows sharing one internal name; deleting one replica must not remove aliases that still route to the survivors, in any team --- .../model_management_endpoints.py | 16 +++- .../test_model_management_endpoints.py | 73 +++++++++++++++++++ 2 files changed, 86 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ddb7f08ad15..bdaa69cd8fe 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -802,6 +802,9 @@ async def _remove_unbacked_team_models( ``{public_name: "model_name_{team_id}_{uuid}"}`` entry in the team's model_aliases, so the alias scan runs for every team model; skipping it for internal-shaped names left stale aliases that rewrote requests to deployments that no longer exist. + Aliases are scrubbed only when the deleted deployment's name no longer resolves in + the router, so deleting one replica of a load-balanced group never breaks aliases + that still route to the surviving replicas (in any team). A public name that still resolves to a live router deployment (e.g. a gateway-level model group shared with the team) is kept in team.models, so deleting a per-team @@ -811,9 +814,16 @@ async def _remove_unbacked_team_models( if team_id is None: return - removed_model_aliases = await delete_team_model_alias( - public_model_name=model_params.model_name, - prisma_client=prisma_client, + deleted_name_still_served = ( + llm_router is not None and model_params.model_name in llm_router.model_name_to_deployment_indices + ) + removed_model_aliases: List[Tuple[str, str]] = ( + [] + if deleted_name_still_served + else await delete_team_model_alias( + public_model_name=model_params.model_name, + prisma_client=prisma_client, + ) ) removed_alias_names = {alias for alias_team_id, alias in removed_model_aliases if alias_team_id == team_id} candidate_names = ( diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 54c17a845d2..b14c9d1e490 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -2277,6 +2277,79 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db.litellm_teamtable.update.assert_not_awaited() mock_refresh.assert_not_awaited() + @pytest.mark.asyncio + async def test_delete_replica_keeps_alias_while_surviving_replica_serves_it(self): + """Deleting one replica of a load-balanced legacy team model (several + deployment rows sharing one internal model_name) must not scrub the team + alias: the surviving replicas still serve the aliased name, so removing + the alias would break routing that works.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + team_id = "team-lb-legacy" + model_id = "lb-replica-1" + internal_name = f"model_name_{team_id}_shared-uuid" + + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=internal_name, + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={"id": model_id, "team_id": team_id}, + created_by="admin", + updated_by="admin", + ) + team_row = LiteLLM_TeamTable( + team_id=team_id, + team_alias="lb-legacy-team", + members_with_roles=[Member(user_id="admin", role="admin")], + models=["gpt-4"], + ) + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_modeltable.update = AsyncMock() + + mock_router = MagicMock() + mock_router.model_name_to_deployment_indices = {internal_name: [0]} + + admin_user = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=admin_user, + ) + + assert "deleted successfully" in result["message"] + mock_prisma.db.litellm_modeltable.find_many.assert_not_awaited() + mock_prisma.db.litellm_modeltable.update.assert_not_awaited() + mock_prisma.db.litellm_teamtable.update.assert_not_awaited() + class TestDeleteModelTeamAuth: """Team auth on the /model/delete path. From 9f4e3c6009086ba8222d26df95af0c186bc394ac Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Mon, 27 Jul 2026 17:13:53 -0700 Subject: [PATCH 17/54] fix(proxy): report when a model write does not survive the post-write reload Every model-write endpoint returned 200 off the DB write alone; a model the reload dropped (ignore_invalid_deployments, or a wholesale reload failure) stayed invisible on every channel at once, which is how the registry-leak defect went undiagnosed for three weeks. ProxyConfig.add_deployment and clear_cache now return whether the reload pass completed, and each write endpoint verifies the rows it wrote are live in this pod's router afterwards, distinguishing a deliberately environment-inactive model via the same predicate the Router's own gate uses. The access-group writers return the mutated id set instead of discarding it --- ...model_access_group_management_endpoints.py | 251 ++++++++++++------ .../model_management_endpoints.py | 173 ++++++++++-- litellm/router.py | 63 +++-- .../test_access_group_management.py | 112 ++++++++ .../test_model_management_endpoints.py | 107 +++++++- .../test_litellm/test_model_block_unblock.py | 52 +++- tests/test_litellm/test_router.py | 19 ++ 7 files changed, 635 insertions(+), 142 deletions(-) diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index 1ca29eee89d..5e3ff8eb7f8 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -6,6 +6,7 @@ Endpoints here: """ import json +from collections.abc import Mapping, Sequence from typing import Any, Dict, List, Tuple from fastapi import APIRouter, Depends, HTTPException @@ -16,6 +17,9 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth # Clear cache and reload models to pick up the access group changes from litellm.proxy.management_endpoints.model_management_endpoints import ( + live_model_ids_snapshot, + model_info_as_mapping, + reload_serving_verdict, clear_cache, ) from litellm.proxy.utils import PrismaClient @@ -72,11 +76,92 @@ def add_access_group_to_deployment(model_info: Dict[str, Any], access_group: str return model_info, True +def _raise_http_if_reload_degraded_serving( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + access_group: str, +) -> None: + """Same verdict as the model-write endpoints, expressed through this file's + HTTPException error convention, with the metadata-only obligation: these writes + change group membership, not the models themselves, so a row that was already not + serving before the reload is never blamed here; only a model this reload stopped + serving is reported.""" + missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=False) + gone = tuple(dict.fromkeys((*missing, *collateral))) + if not gone: + return + raise HTTPException( + status_code=500, + detail={ + "error": ( + f"Access group '{access_group}' was saved to the database, but model id(s) {list(gone)} that " + "this pod was serving are no longer live after the reload it triggered. Other pods reload on " + "their own interval. Check server logs for 'Error upserting deployment' for the cause." + ) + }, + ) + + +async def _tag_deployment_with_access_group( + model_id: str, + model_info: object, + access_group: str, + prisma_client: PrismaClient, +) -> tuple[str, Mapping[str, object]] | None: + """Write `access_group` into one deployment's model_info; returns the + (model_id, updated model_info) pair when a write happened, None when the + deployment already carried the group.""" + updated_model_info, was_modified = add_access_group_to_deployment( + model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)), + access_group=access_group, + ) + if not was_modified: + return None + await ModelRepository(prisma_client).table.update( + where={"model_id": model_id}, + data={"model_info": json.dumps(updated_model_info)}, + ) + verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}") + return (model_id, updated_model_info) + + +def _readable_model_info_or_raise(model_id: str, model_info: object) -> Mapping[str, object]: + """These helpers rewrite the model_info column wholesale, so a present-but-unreadable + value must refuse loudly rather than be silently replaced with a fresh object; an + absent value stays a legitimate empty start.""" + parsed = model_info_as_mapping(model_info) + if parsed is None and model_info is not None: + raise ValueError(f"model_info for deployment {model_id} is not a readable JSON object; refusing to rewrite it") + return parsed or {} + + +async def _strip_access_group_from_deployment( + model_id: str, + model_info: object, + access_group: str, + prisma_client: PrismaClient, +) -> tuple[str, Mapping[str, object]] | None: + """Remove `access_group` from one deployment's model_info; returns the + (model_id, updated model_info) pair when a write happened, None when the + deployment did not carry the group.""" + updated_model_info, was_modified = remove_access_group_from_deployment( + model_info=dict(_readable_model_info_or_raise(model_id=model_id, model_info=model_info)), + access_group=access_group, + ) + if not was_modified: + return None + await ModelRepository(prisma_client).table.update( + where={"model_id": model_id}, + data={"model_info": json.dumps(updated_model_info)}, + ) + return (model_id, updated_model_info) + + async def update_deployments_with_access_group( model_names: List[str], access_group: str, prisma_client: PrismaClient, -) -> int: +) -> tuple[tuple[str, Mapping[str, object]], ...]: """ Update all deployments for the given model names to include the access group. @@ -86,20 +171,15 @@ async def update_deployments_with_access_group( prisma_client: Database client Returns: - int: Number of deployments updated + The (model_id, updated model_info) pair of every deployment actually written, + so callers can verify each one survived the post-write reload """ - models_updated = 0 + deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": {"in": model_names}}) + verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_names: {model_names}") + found_names = {deployment.model_name for deployment in deployments} for model_name in model_names: - verbose_proxy_logger.debug(f"Updating deployments for model_name: {model_name}") - - # Get all deployments with this model_name - deployments = await ModelRepository(prisma_client).table.find_many(where={"model_name": model_name}) - - verbose_proxy_logger.debug(f"Found {len(deployments)} deployments for model_name: {model_name}") - - # If no deployments found, this is a config model (not in DB) - if len(deployments) == 0: + if model_name not in found_names: raise HTTPException( status_code=400, detail={ @@ -107,65 +187,52 @@ async def update_deployments_with_access_group( }, ) - # Update each deployment - for deployment in deployments: - model_info = deployment.model_info or {} - - # Add access group using helper - updated_model_info, was_modified = add_access_group_to_deployment( - model_info=model_info, - access_group=access_group, - ) - - # Only update in DB if modified - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - - models_updated += 1 - verbose_proxy_logger.debug( - f"Updated deployment {deployment.model_id} with access group: {access_group}" - ) - - return models_updated + tagged = [ + await _tag_deployment_with_access_group( + model_id=deployment.model_id, + model_info=deployment.model_info, + access_group=access_group, + prisma_client=prisma_client, + ) + for deployment in deployments + ] + return tuple(pair for pair in tagged if pair is not None) async def update_specific_deployments_with_access_group( model_ids: List[str], access_group: str, prisma_client: PrismaClient, -) -> int: +) -> tuple[tuple[str, Mapping[str, object]], ...]: """ Update specific deployments (by model_id) to include the access group. Unlike update_deployments_with_access_group which tags ALL deployments sharing a model_name, this function only tags the specific deployments identified by - their unique model_id. + their unique model_id. Returns the (model_id, updated model_info) pair of every + deployment actually written. """ - models_updated = 0 - for model_id in model_ids: - verbose_proxy_logger.debug(f"Updating specific deployment model_id: {model_id}") - deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id}) - if deployment is None: - raise HTTPException( - status_code=400, - detail={"error": f"Deployment with model_id '{model_id}' not found in Database."}, - ) - model_info = deployment.model_info or {} - updated_model_info, was_modified = add_access_group_to_deployment( - model_info=model_info, + verbose_proxy_logger.debug(f"Updating specific deployment model_ids: {model_ids}") + tagged = [ + await _tag_deployment_with_access_group( + model_id=model_id, + model_info=(await _find_deployment_or_400(model_id=model_id, prisma_client=prisma_client)), access_group=access_group, + prisma_client=prisma_client, ) - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - models_updated += 1 - verbose_proxy_logger.debug(f"Updated deployment {model_id} with access group: {access_group}") - return models_updated + for model_id in model_ids + ] + return tuple(pair for pair in tagged if pair is not None) + + +async def _find_deployment_or_400(model_id: str, prisma_client: PrismaClient) -> Mapping[str, object] | None: + deployment = await ModelRepository(prisma_client).table.find_unique(where={"model_id": model_id}) + if deployment is None: + raise HTTPException( + status_code=400, + detail={"error": f"Deployment with model_id '{model_id}' not found in Database."}, + ) + return deployment.model_info def remove_access_group_from_deployment(model_info: Dict[str, Any], access_group: str) -> Tuple[Dict[str, Any], bool]: @@ -335,20 +402,28 @@ async def create_model_group( # Update deployments using the appropriate method if use_model_ids: assert data.model_ids is not None - models_updated = await update_specific_deployments_with_access_group( + updated_pairs = await update_specific_deployments_with_access_group( model_ids=data.model_ids, access_group=data.access_group, prisma_client=prisma_client, ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( + updated_pairs = await update_deployments_with_access_group( model_names=data.model_names, access_group=data.access_group, prisma_client=prisma_client, ) + models_updated = len(updated_pairs) + + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=updated_pairs, + access_group=data.access_group, + ) verbose_proxy_logger.info( f"Successfully created access group '{data.access_group}' with {models_updated} models updated" @@ -573,38 +648,42 @@ async def update_access_group( # Step 1: Remove access group from ALL DB deployments (skip config models) all_deployments = await ModelRepository(prisma_client).table.find_many() - for deployment in all_deployments: - model_info = deployment.model_info or {} - - updated_model_info, was_modified = remove_access_group_from_deployment( - model_info=model_info, + stripped = [ + await _strip_access_group_from_deployment( + model_id=deployment.model_id, + model_info=deployment.model_info, access_group=access_group, + prisma_client=prisma_client, ) - - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) + for deployment in all_deployments + ] + stripped_pairs = tuple(pair for pair in stripped if pair is not None) # Step 2: Add access group using the appropriate method if use_model_ids: assert data.model_ids is not None - models_updated = await update_specific_deployments_with_access_group( + updated_pairs = await update_specific_deployments_with_access_group( model_ids=data.model_ids, access_group=access_group, prisma_client=prisma_client, ) else: assert data.model_names is not None - models_updated = await update_deployments_with_access_group( + updated_pairs = await update_deployments_with_access_group( model_names=data.model_names, access_group=access_group, prisma_client=prisma_client, ) + models_updated = len(updated_pairs) # Clear cache and reload models to pick up the access group changes + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=list({**dict(stripped_pairs), **dict(updated_pairs)}.items()), + access_group=access_group, + ) verbose_proxy_logger.info( f"Successfully updated access group '{access_group}' with {models_updated} models updated" @@ -686,25 +765,27 @@ async def delete_access_group( try: # Remove access group from all DB deployments (skip config models) all_deployments = await ModelRepository(prisma_client).table.find_many() - models_updated = 0 - for deployment in all_deployments: - model_info = deployment.model_info or {} - - updated_model_info, was_modified = remove_access_group_from_deployment( - model_info=model_info, + removed = [ + await _strip_access_group_from_deployment( + model_id=deployment.model_id, + model_info=deployment.model_info, access_group=access_group, + prisma_client=prisma_client, ) - - if was_modified: - await ModelRepository(prisma_client).table.update( - where={"model_id": deployment.model_id}, - data={"model_info": json.dumps(updated_model_info)}, - ) - models_updated += 1 + for deployment in all_deployments + ] + removed_pairs = tuple(pair for pair in removed if pair is not None) + models_updated = len(removed_pairs) # Clear cache and reload models to pick up the access group changes + live_before_reload = live_model_ids_snapshot() await clear_cache() + _raise_http_if_reload_degraded_serving( + before=live_before_reload, + written_models=removed_pairs, + access_group=access_group, + ) verbose_proxy_logger.info( f"Successfully deleted access group '{access_group}' from {models_updated} deployments" diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 91f0b9b2790..e3f79af8f43 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -13,6 +13,7 @@ model/{model_id}/update - PATCH endpoint for model update. import asyncio import datetime import json +from collections.abc import Mapping, Sequence from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast from fastapi import APIRouter, Depends, HTTPException, Header, Request, status @@ -272,6 +273,7 @@ async def patch_model( ) # Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates) + live_before_reload = live_model_ids_snapshot() await clear_cache() ## CREATE AUDIT LOG ## @@ -288,6 +290,12 @@ async def patch_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(model_id, getattr(updated_model, "model_info", None))], + action="update", + ) + return updated_model except Exception as e: @@ -370,6 +378,7 @@ async def _set_model_blocked_status( }, ) + live_before_reload = live_model_ids_snapshot() await clear_cache() asyncio.create_task( @@ -387,6 +396,12 @@ async def _set_model_blocked_status( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(data.model_id, getattr(updated_model, "model_info", None))], + action=action, + ) + return updated_model except Exception as e: @@ -713,13 +728,8 @@ async def _get_team_deployments( # Confirm team_id in model_info (defensive check) result = [] for row in response: - model_info = row.model_info - if isinstance(model_info, str): - try: - model_info = json.loads(model_info) - except (TypeError, ValueError): - continue - if isinstance(model_info, dict) and model_info.get("team_id") == team_id: + model_info = model_info_as_mapping(row.model_info) + if model_info is not None and model_info.get("team_id") == team_id: result.append(row) return result @@ -770,13 +780,8 @@ async def _get_team_public_model_names( deployments = await _get_team_deployments(team_id, prisma_client) public_names: Set[str] = set() for row in deployments: - model_info = row.model_info - if isinstance(model_info, str): - try: - model_info = json.loads(model_info) - except (TypeError, ValueError): - continue - if isinstance(model_info, dict): + model_info = model_info_as_mapping(row.model_info) + if model_info is not None: public_name = model_info.get("team_public_model_name") if public_name: public_names.add(public_name) @@ -853,18 +858,11 @@ async def _update_existing_team_model_assignment( def _get_team_public_model_name( model_info: Optional[Union[dict, str]], ) -> Optional[str]: - if isinstance(model_info, dict): - value = model_info.get("team_public_model_name") - return value if isinstance(value, str) else None - if isinstance(model_info, str): - try: - parsed = json.loads(model_info) - except (TypeError, ValueError): - return None - if isinstance(parsed, dict): - value = parsed.get("team_public_model_name") - return value if isinstance(value, str) else None - return None + parsed = model_info_as_mapping(model_info) + if parsed is None: + return None + value = parsed.get("team_public_model_name") + return value if isinstance(value, str) else None old_public_name = db_model.model_info.team_public_model_name if db_model.model_info else None @@ -1275,6 +1273,7 @@ async def add_new_model( - store keys separately """ + live_before_reload = live_model_ids_snapshot() try: _original_litellm_model_name = model_params.model_name if model_params.model_info.team_id is None: @@ -1330,6 +1329,12 @@ async def add_new_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(model_response.model_id, getattr(model_response, "model_info", None))], + action="create", + ) + return model_response except Exception as e: @@ -1450,8 +1455,8 @@ async def update_model( ) # Clear cache and reload models (uses config setting or defaults to preserving config models for DB updates) + live_before_reload = live_model_ids_snapshot() await clear_cache() - ## CREATE AUDIT LOG ## asyncio.create_task( create_object_audit_log( @@ -1474,6 +1479,12 @@ async def update_model( ) ) + raise_if_reload_degraded_serving( + before=live_before_reload, + written_models=[(_model_id, getattr(model_response, "model_info", None))], + action="update", + ) + return model_response except Exception as e: verbose_proxy_logger.exception( @@ -1677,6 +1688,114 @@ def _deduplicate_litellm_router_models(models: List[Dict]) -> List[Dict]: return unique_models +def model_info_as_mapping(model_info: object) -> Mapping[str, object] | None: + """A DB row's model_info column arrives as a dict or as its JSON string depending on + the query path, and every consumer needs the mapping. Single owner of that parse: + returns None when no usable mapping exists (None, an unparseable string, or JSON + that is not an object), and callers choose what None means for them.""" + if isinstance(model_info, Mapping): + return model_info + if not isinstance(model_info, str): + return None + try: + parsed = json.loads(model_info) + except (TypeError, ValueError): + return None + return parsed if isinstance(parsed, Mapping) else None + + +def _expects_liveness_on_this_pod(model_info: object) -> bool: + from litellm.router import model_info_is_active_for_environment + + try: + return model_info_is_active_for_environment(model_info=model_info_as_mapping(model_info)) + except ValueError: + return True + + +def live_model_ids_snapshot() -> frozenset[str]: + """The ids this pod's router is currently serving, read fresh from the module global + because a reload can rebind it. The empirical ground truth every verdict below is + computed from; an absent router serves nothing.""" + from litellm.proxy.proxy_server import llm_router + + if llm_router is None: + return frozenset() + return frozenset(llm_router.get_model_ids()) + + +def reload_serving_verdict( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + written_must_serve: bool, +) -> tuple[tuple[str, ...], tuple[str, ...]]: + """Judge a write-triggered reload by diffing the router's serving state instead of + trusting any layer of the reload stack to report its own failure. + + The full cell matrix, per id: + - written, must-serve (the write's purpose is this model's serving state): live now + is fine; not live is reported unless the row is deliberately inactive for this + pod's LITELLM_ENVIRONMENT; a row whose model_info cannot be read counts as + expecting to serve, so its drop is still reported + - written, metadata-only (must_not_degrade): live before and gone now is reported; + a row that was already not serving stays silent, because its deadness predates + this write and blaming it would block unrelated metadata fixes + - not written but live before and gone now: collateral degradation of this pod + caused by the reload this request triggered (a wholesale re-add failure, or a + newly introduced conflict), always reported + + Returns (written ids violating their obligation, collateral ids no longer served). + Best effort under concurrent admin writes: the snapshot spans only this request. + """ + now = live_model_ids_snapshot() + written_ids = frozenset(model_id for model_id, _ in written_models) + if written_must_serve: + missing = tuple( + model_id + for model_id, model_info in written_models + if model_id not in now and _expects_liveness_on_this_pod(model_info) + ) + else: + missing = tuple(model_id for model_id, _ in written_models if model_id in before and model_id not in now) + collateral = tuple(sorted(before - now - written_ids)) + return (missing, collateral) + + +def raise_if_reload_degraded_serving( + before: frozenset[str], + written_models: Sequence[tuple[str, object]], + action: str, +) -> None: + """The caller-visible error this pod's model-write endpoints owe their caller when + the model they wrote is not being served after the reload they triggered. The DB + write is durable either way and every other pod reloads on its own interval; this + speaks only for the handling pod.""" + missing, collateral = reload_serving_verdict(before=before, written_models=written_models, written_must_serve=True) + if not missing and not collateral: + return + missing_clause = ( + f"the model id(s) {list(missing)} are not live in this pod's router after the reload and are not " + "being served by this pod." + if missing + else "the reload it triggered degraded this pod's serving state." + ) + collateral_clause = ( + f" Previously served model id(s) {list(collateral)} are also no longer being served by this pod." + if collateral + else "" + ) + raise ProxyException( + message=( + f"Model {action} was saved to the database, but {missing_clause}{collateral_clause} " + "Other pods reload on their own interval. Check server logs for 'Error upserting deployment' or " + "'Error creating deployment' for the cause." + ), + type=ProxyErrorTypes.internal_server_error, + code=status.HTTP_500_INTERNAL_SERVER_ERROR, + param=None, + ) + + async def clear_cache(): """ Clear router caches and reload models. diff --git a/litellm/router.py b/litellm/router.py index 487d6a31226..d3caa069b28 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -19,6 +19,7 @@ import threading import time import traceback from collections import defaultdict +from collections.abc import Mapping from functools import lru_cache from typing import ( TYPE_CHECKING, @@ -270,6 +271,43 @@ def _cost_value_as_float(value: Union[str, int, float, None]) -> Optional[float] return None +def model_info_is_active_for_environment(model_info: Mapping[str, object] | None) -> bool: + """Single owner of the environment-gating rule: a deployment whose model_info names + `supported_environments` loads only on pods whose LITELLM_ENVIRONMENT is in that list. + `Router.deployment_is_active_for_environment` delegates here, and the model-write + endpoints consult the same rule to tell a deliberately inactive model from one that + was dropped by a failed reload.""" + if model_info is None: + return True + supported_environments = model_info.get("supported_environments") + if supported_environments is None: + return True + if not isinstance(supported_environments, (list, tuple)): + raise ValueError( + f"supported_environments must be a list of {VALID_LITELLM_ENVIRONMENTS}. " + f"but set as: {supported_environments} for model_info: {model_info}" + ) + litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT") + if litellm_environment is None: + raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env") + + if litellm_environment not in VALID_LITELLM_ENVIRONMENTS: + raise ValueError( + f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}" + ) + + for _env in supported_environments: + if _env not in VALID_LITELLM_ENVIRONMENTS: + raise ValueError( + f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} " + f"for model_info: {model_info}" + ) + + if litellm_environment in supported_environments: + return True + return False + + _PreRoutingStrategyT = TypeVar("_PreRoutingStrategyT") @@ -7982,30 +8020,7 @@ class Router: - ValueError: If LITELLM_ENVIRONMENT is not set in .env or not one of the valid values - ValueError: If supported_environments is not set in model_info or not one of the valid values """ - if ( - deployment.model_info is None - or "supported_environments" not in deployment.model_info - or deployment.model_info["supported_environments"] is None - ): - return True - litellm_environment = get_secret_str(secret_name="LITELLM_ENVIRONMENT") - if litellm_environment is None: - raise ValueError("Set 'supported_environments' for model but not 'LITELLM_ENVIRONMENT' set in .env") - - if litellm_environment not in VALID_LITELLM_ENVIRONMENTS: - raise ValueError( - f"LITELLM_ENVIRONMENT must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {litellm_environment}" - ) - - for _env in deployment.model_info["supported_environments"]: - if _env not in VALID_LITELLM_ENVIRONMENTS: - raise ValueError( - f"supported_environments must be one of {VALID_LITELLM_ENVIRONMENTS}. but set as: {_env} for deployment: {deployment}" - ) - - if litellm_environment in deployment.model_info["supported_environments"]: - return True - return False + return model_info_is_active_for_environment(model_info=deployment.model_info) def set_model_list(self, model_list: list): original_model_list = copy.deepcopy(model_list) diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py index 8eed696e77d..3240ad20edb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_management.py @@ -160,7 +160,9 @@ async def test_create_access_group_with_model_names_tags_all_deployments(): { "model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o", "api_key": "fake-key"}, + "model_info": {"id": deployment_id, "db_model": True}, } + for deployment_id in ("deploy-A", "deploy-B", "deploy-C") ] ) @@ -318,3 +320,113 @@ async def test_create_access_group_invalid_model_id_returns_400(): await create_model_group(data=request_data, user_api_key_dict=mock_user) assert exc_info.value.status_code == 400 assert "non-existent-id" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_create_access_group_surfaces_dropped_models(): + """An access-group write whose reload does not leave the tagged models live on this + pod must report the drop through this file's HTTPException contract, not a 200.""" + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + create_model_group, + ) + from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + NewModelGroupRequest, + ) + + deploy_a = MagicMock(model_id="deploy-A", model_name="gpt-4o", model_info={}) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deploy_a) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + mock_user = UserAPIKeyAuth(user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + wiped_router = MagicMock() + wiped_router.get_model_ids.side_effect = [["deploy-A"], []] + with ( + patch("litellm.proxy.proxy_server.llm_router", wiped_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch( + "litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(HTTPException) as exc_info: + await create_model_group( + data=NewModelGroupRequest(access_group="production-models", model_ids=["deploy-A"]), + user_api_key_dict=mock_user, + ) + + assert exc_info.value.status_code == 500 + assert "deploy-A" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_tag_deployment_parses_string_model_info_and_refuses_corrupt(): + """The model_info column can arrive as its JSON string; tagging must parse it rather + than crash, and must refuse to rewrite a present-but-unreadable value.""" + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + _tag_deployment_with_access_group, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + pair = await _tag_deployment_with_access_group( + model_id="deploy-str", + model_info='{"access_groups": ["existing"]}', + access_group="new-group", + prisma_client=mock_prisma, + ) + assert pair is not None + assert pair[0] == "deploy-str" + assert pair[1]["access_groups"] == ["existing", "new-group"] + + with pytest.raises(ValueError, match="deploy-corrupt"): + await _tag_deployment_with_access_group( + model_id="deploy-corrupt", + model_info="{not json", + access_group="new-group", + prisma_client=mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_delete_access_group_ignores_models_that_were_already_dead(): + """A metadata-only strip over a model this pod never served must not fail the write; + the model's deadness predates the request, and blaming it here would make a broken + model block every access-group fix that touches it.""" + deploy_broken = MagicMock( + model_id="deploy-broken", model_name="broken-model", model_info={"access_groups": ["doomed-group"]} + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[deploy_broken]) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( + delete_access_group, + ) + + never_served_router = MagicMock() + never_served_router.get_model_ids.return_value = [] + with ( + patch("litellm.proxy.proxy_server.llm_router", never_served_router), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch( + "litellm.proxy.management_endpoints.model_access_group_management_endpoints.clear_cache", + new=AsyncMock(return_value=None), + ), + ): + response = await delete_access_group( + access_group="doomed-group", + user_api_key_dict=UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert response.models_updated == 1 + mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index bd5eda0197b..62f6c0c8fa1 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -819,6 +819,7 @@ class TestUpdateModel: ) mock_router = MagicMock() + mock_router.get_model_ids.return_value = [model_id] admin_user = UserAPIKeyAuth( user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN ) @@ -838,7 +839,7 @@ class TestUpdateModel: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock(return_value=None), + new=AsyncMock(return_value=True), ) as mock_clear_cache, ): await update_model( @@ -1888,6 +1889,7 @@ class TestAddAndDeleteModelLifecycle: mock_router = MagicMock() mock_router.delete_deployment = MagicMock() + mock_router.get_model_ids.return_value = [model_id] _PS = "litellm.proxy.proxy_server" _ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper" @@ -3002,7 +3004,7 @@ class TestPatchModelBlockedAuthGate: with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.premium_user", True), patch( @@ -3048,7 +3050,7 @@ class TestPatchModelBlockedAuthGate: with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), - patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]})), patch("litellm.proxy.proxy_server.store_model_in_db", True), patch("litellm.proxy.proxy_server.premium_user", True), patch( @@ -3057,7 +3059,7 @@ class TestPatchModelBlockedAuthGate: ), patch( "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock(return_value=None), + new=AsyncMock(return_value=True), ), ): result = await patch_model( @@ -3067,3 +3069,100 @@ class TestPatchModelBlockedAuthGate: ) assert result is updated_row mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + + +class TestWriteSurfacesReloadDrop: + """A model-write endpoint may report success only if every row it wrote is, after the + reload it triggered, live in this pod's router or deliberately environment-inactive.""" + + def test_reload_serving_verdict_matrix(self, monkeypatch): + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + reload_serving_verdict, + ) + + live_router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "gpt-4o"}, + "model_info": {"id": "m-live", "db_model": True}, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", live_router) + monkeypatch.setenv("LITELLM_ENVIRONMENT", "development") + + written = [ + ("m-live", {"id": "m-live"}), + ("m-gone", {"id": "m-gone"}), + ("m-env", {"id": "m-env", "supported_environments": ["production"]}), + ("m-env-str", '{"id": "m-env-str", "supported_environments": ["production"]}'), + ("m-env-misconfigured", {"id": "m-env-misconfigured", "supported_environments": ["bogus"]}), + ("m-corrupt", "{not json"), + ] + missing, collateral = reload_serving_verdict( + before=frozenset({"m-live", "m-collateral"}), written_models=written, written_must_serve=True + ) + assert missing == ("m-gone", "m-env-misconfigured", "m-corrupt") + assert collateral == ("m-collateral",) + + missing, collateral = reload_serving_verdict( + before=frozenset({"m-live", "m-was-live"}), + written_models=[("m-live", None), ("m-was-live", None), ("m-never-lived", None)], + written_must_serve=False, + ) + assert missing == ("m-was-live",) + assert collateral == () + + def test_raise_if_reload_degraded_serving_contract(self, monkeypatch): + import litellm + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + raise_if_reload_degraded_serving, + ) + + live_router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "gpt-4o"}, + "model_info": {"id": "m-live", "db_model": True}, + } + ] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", live_router) + + assert ( + raise_if_reload_degraded_serving( + before=frozenset({"m-live"}), written_models=[("m-live", None)], action="update" + ) + is None + ) + + with pytest.raises(ProxyException, match="m-gone"): + raise_if_reload_degraded_serving( + before=frozenset(), written_models=[("m-gone", None)], action="update" + ) + + with pytest.raises(ProxyException, match="m-collateral"): + raise_if_reload_degraded_serving( + before=frozenset({"m-live", "m-collateral"}), written_models=[("m-live", None)], action="update" + ) + + +class TestModelInfoAsMapping: + """The model_info column reaches consumers as a dict or as its JSON string; this is + the single owner of that parse, and None means no usable mapping.""" + + def test_contract(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + model_info_as_mapping, + ) + + assert model_info_as_mapping({"id": "m1"}) == {"id": "m1"} + assert model_info_as_mapping('{"id": "m1"}') == {"id": "m1"} + assert model_info_as_mapping(None) is None + assert model_info_as_mapping("{not json") is None + assert model_info_as_mapping('["a", "b"]') is None + assert model_info_as_mapping(42) is None diff --git a/tests/test_litellm/test_model_block_unblock.py b/tests/test_litellm/test_model_block_unblock.py index 318b2f519c4..dc0098e405e 100644 --- a/tests/test_litellm/test_model_block_unblock.py +++ b/tests/test_litellm/test_model_block_unblock.py @@ -34,9 +34,9 @@ def _setup_model_block_mocks(monkeypatch, *, updated_blocked: bool): mock_prisma_client.db.litellm_proxymodeltable = model_table mock_router = MagicMock() - mock_router.get_deployment.return_value = None + mock_router.get_model_ids.return_value = [model_id] - mock_clear_cache = AsyncMock(return_value=None) + mock_clear_cache = AsyncMock(return_value=True) mock_audit_log = AsyncMock(return_value=None) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -197,3 +197,51 @@ async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch assert exc_info.value.status_code == 403 assert "Model is blocked" in exc_info.value.message + + +@pytest.mark.asyncio +async def test_model_block_surfaces_wholesale_reload_failure(monkeypatch): + """The write endpoints owe the caller an error when the pod failed to reload at all; + the DB row is saved but this pod is not serving the change.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import block_model + + model_id, model_table, updated_row, mock_clear_cache, mock_audit_log = _setup_model_block_mocks( + monkeypatch, updated_blocked=True + ) + wiped_router = MagicMock() + wiped_router.get_model_ids.side_effect = [[model_id], []] + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", wiped_router) + + with pytest.raises(ProxyException, match=model_id): + await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by="operator@example.com", + ) + + assert mock_audit_log.call_args.kwargs["object_id"] == model_id + + +@pytest.mark.asyncio +async def test_model_block_surfaces_model_dropped_by_reload(monkeypatch): + """A reload that completes but drops the written model (ignore_invalid_deployments + swallowed its re-add) must not produce an unqualified success.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import block_model + + model_id, model_table, updated_row, mock_clear_cache, _ = _setup_model_block_mocks( + monkeypatch, updated_blocked=True + ) + dropped_router = MagicMock() + dropped_router.get_model_ids.return_value = [] + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", dropped_router) + + with pytest.raises(ProxyException, match=model_id): + await block_model( + data=BlockModelRequest(model_id=model_id), + http_request=MagicMock(), + user_api_key_dict=_proxy_admin(), + litellm_changed_by="operator@example.com", + ) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ad4e430c603..6db19d63a01 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -6379,3 +6379,22 @@ class TestPreRoutingStrategyRegistryLifecycle: litellm_params=LiteLLM_Params(**params) ) assert actual is expected, params["model"] + + +def test_model_info_is_active_for_environment_matrix(monkeypatch): + """The model-write endpoints consult this predicate to tell a deliberately + environment-inactive model from one dropped by a failed reload; the Router's own + deployment gate delegates to it, so the two can never diverge.""" + from litellm.router import model_info_is_active_for_environment + + assert model_info_is_active_for_environment(model_info=None) is True + assert model_info_is_active_for_environment(model_info={"id": "m1"}) is True + assert model_info_is_active_for_environment(model_info={"supported_environments": None}) is True + + monkeypatch.setenv("LITELLM_ENVIRONMENT", "development") + assert model_info_is_active_for_environment(model_info={"supported_environments": ["development"]}) is True + assert model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) is False + + monkeypatch.delenv("LITELLM_ENVIRONMENT") + with pytest.raises(ValueError, match="LITELLM_ENVIRONMENT"): + model_info_is_active_for_environment(model_info={"supported_environments": ["production"]}) From 17ce2c4e9211a4caf27bfaf58781c9b5f32b633e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 18:57:15 -0700 Subject: [PATCH 18/54] fix(proxy): resolve named credentials on provider-only batch and files calls --- litellm/proxy/batches_endpoints/endpoints.py | 25 ++ .../openai_files_endpoints/common_utils.py | 24 ++ .../openai_files_endpoints/files_endpoints.py | 36 ++- .../proxy/batches_endpoints/test_endpoints.py | 153 ++++++++++- .../test_files_endpoint.py | 257 ++++++++++++++++++ 5 files changed, 488 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index a2dc5e1caf5..a91b29002e3 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -23,6 +23,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + apply_team_provider_credentials, decode_model_from_file_id, encode_batch_response_ids, encode_file_id_with_model, @@ -295,6 +296,12 @@ async def create_batch( verbose_proxy_logger.debug(f"Created batch using model: {model_param}") else: # SCENARIO 3: Fallback to custom_llm_provider (uses env variables) + apply_team_provider_credentials( + data=cast(dict, _create_batch_data), # cast-ok: TypedDict is a dict at runtime + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.acreate_batch( custom_llm_provider=custom_llm_provider, **_create_batch_data, # type: ignore @@ -525,6 +532,12 @@ async def retrieve_batch( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.aretrieve_batch( custom_llm_provider=custom_llm_provider, **data, # type: ignore @@ -718,6 +731,12 @@ async def list_batches( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.alist_batches( custom_llm_provider=custom_llm_provider, # type: ignore after=after, @@ -908,6 +927,12 @@ async def cancel_batch( # Extract batch_id from data to avoid "multiple values for keyword argument" error # data was cast from CancelBatchRequest which already contains batch_id data.pop("batch_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) _cancel_batch_data = CancelBatchRequest(batch_id=batch_id, **data) response = await litellm.acancel_batch( custom_llm_provider=custom_llm_provider, # type: ignore diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index efd1d6b6cee..d4ac45559ee 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -14,6 +14,7 @@ from litellm.types.utils import SpecialEnums if TYPE_CHECKING: from fastapi import Request + from litellm.proxy._types import UserAPIKeyAuth from litellm.router import Router @@ -373,6 +374,29 @@ def get_team_provider_credentials( return None +def apply_team_provider_credentials( + data: dict, # mutable-ok: credentials are merged into the request payload in place, same contract as prepare_data_with_credentials + llm_router: Optional["Router"], + user_api_key_dict: "UserAPIKeyAuth", + custom_llm_provider: str, +) -> None: + """ + Resolve credentials for a provider-only request (no model pinned) via + ``get_team_provider_credentials`` and merge them into ``data`` in-place. + Leaves ``data`` untouched when no authorized deployment matches, so the + caller falls back to environment-variable credentials exactly as before. + """ + credentials = get_team_provider_credentials( + llm_router=llm_router, + team_models=user_api_key_dict.team_models or [], + custom_llm_provider=custom_llm_provider, + team_id=user_api_key_dict.team_id, + ) + if credentials is None: + return + prepare_data_with_credentials(data=data, credentials=credentials) + + def prepare_data_with_credentials( data: dict, credentials: dict, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 0c9aa667751..f1bcfbafe58 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -43,10 +43,10 @@ from litellm.litellm_core_utils.cloud_storage_security import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + apply_team_provider_credentials, encode_file_id_with_model, extract_file_creation_params, get_credentials_for_model, - get_team_provider_credentials, handle_model_based_routing, prepare_data_with_credentials, validate_managed_files_requirement, @@ -253,6 +253,12 @@ async def route_create_file( _create_file_request=_create_file_request, ) else: + apply_team_provider_credentials( + data=cast(dict, _create_file_request), # cast-ok: TypedDict is a plain dict at runtime; merged in place + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) # get configs for custom_llm_provider llm_provider_config = get_files_provider_config(custom_llm_provider=custom_llm_provider) if llm_provider_config is not None: @@ -735,6 +741,14 @@ async def get_file_content( check_file_id_encoding=True, ) + if not should_route: + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) + from litellm.proxy.openai_files_endpoints.file_content_streaming_handler import ( FileContentStreamingHandler, ) @@ -983,6 +997,12 @@ async def get_file( # Remove file_id from data to avoid "multiple values for keyword argument" error # data was initialized with {"file_id": file_id} data.pop("file_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.afile_retrieve( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1183,6 +1203,12 @@ async def delete_file( ) else: data.pop("file_id", None) + apply_team_provider_credentials( + data=data, + llm_router=llm_router, + user_api_key_dict=user_api_key_dict, + custom_llm_provider=custom_llm_provider, + ) response = await litellm.afile_delete( custom_llm_provider=custom_llm_provider, file_id=file_id, @@ -1354,14 +1380,12 @@ async def list_files( # No model/target_model_names pinned: resolve upstream credentials from # the team's deployment for this provider so the call is authenticated # against the team's own account (e.g. the team's openai deployment). - team_credentials = get_team_provider_credentials( + apply_team_provider_credentials( + data=data, llm_router=llm_router, - team_models=user_api_key_dict.team_models or [], + user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, - team_id=user_api_key_dict.team_id, ) - if team_credentials is not None: - prepare_data_with_credentials(data=data, credentials=team_credentials) response = await litellm.afile_list( custom_llm_provider=custom_llm_provider, diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index 9573fddd435..6a185988c9b 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -51,7 +51,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( from litellm.proxy.utils import ProxyLogging from litellm.router import Router from litellm.types.llms.openai import BatchJobStatus -from litellm.types.utils import LiteLLMBatch +from litellm.types.utils import CredentialItem, LiteLLMBatch from fastapi import Response @@ -2091,3 +2091,154 @@ async def test_retrieve__unified_no_router_500(retrieve_harness): assert exc.value.code == "500" retrieve_harness.router_aretrieve.assert_not_called() retrieve_harness.litellm_aretrieve.assert_not_called() + + +# =========================================================================== # +# SCENARIO 3 + configured deployments: a provider-only call (custom-llm-provider +# header, no model anywhere) must resolve the gateway/team deployment's named +# credential for that provider and attach it to the provider call kwargs, +# instead of silently falling through to the host environment's default +# credentials (regression: vertex batch jobs landing in the hosting env's GCP +# project because litellm_credential_name never reached the call). +# =========================================================================== # + +VERTEX_NAMED_CREDENTIAL = CredentialItem( + credential_name="vertex-named-cred", + credential_info={}, + credential_values={ + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + }, +) + + +def vertex_named_credential_router() -> Router: + return Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "litellm_credential_name": "vertex-named-cred", + }, + } + ] + ) + + +@pytest.mark.asyncio +async def test_create__provider_only_resolves_named_vertex_credentials(harness): + """Provider-only create must attach the configured named credential, and must + NOT turn the call into a model-routed one (no model kwarg injected).""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_create(harness) + + assert harness.acreate_kwargs() == { + "custom_llm_provider": "vertex_ai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_create__provider_only_ignores_other_provider_deployments(harness): + """A provider-only vertex call must not pick up credentials from deployments + of a different provider; with no vertex deployment the payload is exactly the + pre-fix env-var fallback.""" + set_body( + harness, + { + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + }, + ) + harness.provider_from_headers.return_value = "vertex_ai" + openai_only_router = Router( + model_list=[ + { + "model_name": "gpt-4o", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-global-openai"}, + } + ] + ) + + with patch.object(proxy_server, "llm_router", openai_only_router): + await call_create(harness) + + assert harness.acreate_kwargs() == { + "custom_llm_provider": "vertex_ai", + "input_file_id": "file-plain", + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "metadata": None, + } + + +@pytest.mark.asyncio +async def test_retrieve__provider_only_resolves_named_vertex_credentials(retrieve_harness): + retrieve_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_retrieve(retrieve_harness, "batch-raw-xyz") + + assert retrieve_harness.aretrieve_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_list__provider_only_resolves_named_vertex_credentials(list_harness): + list_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_list(list_harness) + + assert list_harness.alist_kwargs() == { + "custom_llm_provider": "vertex_ai", + "after": None, + "limit": None, + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } + + +@pytest.mark.asyncio +async def test_cancel__provider_only_resolves_named_vertex_credentials(cancel_harness): + cancel_harness.provider_from_headers.return_value = "vertex_ai" + + with patch.object(litellm, "credential_list", [VERTEX_NAMED_CREDENTIAL]): + with patch.object(proxy_server, "llm_router", vertex_named_credential_router()): + await call_cancel(cancel_harness, "batch-raw-xyz") + + assert cancel_harness.acancel_kwargs() == { + "custom_llm_provider": "vertex_ai", + "batch_id": "batch-raw-xyz", + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + } diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 46ecb31e1c8..23cd1c71cbe 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2610,3 +2610,260 @@ def test_list_files_with_all_proxy_models_team_uses_openai_deployment( assert captured_kwargs.get("api_key") == "team-openai-key" assert captured_kwargs.get("custom_llm_provider") == "openai" proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _setup_vertex_named_credential_router(monkeypatch) -> Router: + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="vertex-named-cred", + credential_info={}, + credential_values={ + "vertex_project": "customer-project", + "vertex_location": "us-central1", + "vertex_credentials": "/creds/customer-sa.json", + }, + ) + ], + ) + return Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "litellm_credential_name": "vertex-named-cred", + }, + } + ] + ) + + +def _assert_vertex_named_credentials_attached(captured_kwargs: dict) -> None: + assert captured_kwargs.get("custom_llm_provider") == "vertex_ai" + assert captured_kwargs.get("vertex_project") == "customer-project" + assert captured_kwargs.get("vertex_location") == "us-central1" + assert captured_kwargs.get("vertex_credentials") == "/creds/customer-sa.json" + assert captured_kwargs.get("model") is None + + +def test_create_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + """ + POST /v1/files with only a custom-llm-provider header (no model, no + target_model_names) must attach the configured named vertex credential to + the upstream call instead of falling through to google.auth.default(), + which uploads into the hosting environment's GCP project. + """ + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_acreate_file(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-vertex-123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", b"{}", "application/jsonl")}, + data={"purpose": "batch"}, + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_get_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_retrieve(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-abc123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "afile_retrieve", _mock_afile_retrieve) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files/file-abc123", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_get_file_content_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_content(**kwargs): + captured_kwargs.update(kwargs) + return HttpxBinaryResponseContent( + response=httpx.Response( + status_code=200, + content=b"vertex-bytes", + headers={"content-type": "application/octet-stream"}, + ) + ) + + monkeypatch.setattr(litellm, "afile_content", _mock_afile_content) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.get( + "/v1/files/file-abc123/content", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert response.content == b"vertex-bytes" + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_delete_file_provider_only_resolves_named_vertex_credentials( + mocker: MockerFixture, monkeypatch +): + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = _setup_vertex_named_credential_router(monkeypatch) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_delete(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-abc123", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "afile_delete", _mock_afile_delete) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + + try: + response = client.delete( + "/v1/files/file-abc123", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("file_id") == "file-abc123" + _assert_vertex_named_credentials_attached(captured_kwargs) + proxy_logging_obj.post_call_failure_hook.assert_not_called() From b592a37b8d2ecb5ab0f29796917e59a02c0817a1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 19:39:30 -0700 Subject: [PATCH 19/54] test(proxy): cover stale-alias warning dedup and key-cache eviction --- .../proxy/test_litellm_pre_call_utils.py | 30 +++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 36fde1b10d8..d6c5e9b1b81 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -5518,3 +5518,33 @@ def test_team_alias_targeting_live_team_deployment_still_rewrites(monkeypatch): ) assert test_data.get("model") == "model_name_team-1_live-uuid" + + +def test_warn_stale_team_alias_once_logs_once_per_key(monkeypatch): + from collections import OrderedDict + + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + + monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict()) + + with patch.object(pre_call_utils.verbose_proxy_logger, "warning") as mock_warning: + pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4") + pre_call_utils._warn_stale_team_alias_once("team-1:gpt-4", "stale alias %s", "gpt-4") + + assert mock_warning.call_count == 1 + + +def test_warn_stale_team_alias_once_evicts_oldest_key_beyond_cap(monkeypatch): + from collections import OrderedDict + + import litellm.proxy.litellm_pre_call_utils as pre_call_utils + + monkeypatch.setattr(pre_call_utils, "_STALE_TEAM_ALIAS_WARNING_KEYS", OrderedDict()) + monkeypatch.setattr(pre_call_utils, "_MAX_STALE_ALIAS_WARNING_KEYS", 2) + + with patch.object(pre_call_utils.verbose_proxy_logger, "warning"): + pre_call_utils._warn_stale_team_alias_once("key-1", "stale alias") + pre_call_utils._warn_stale_team_alias_once("key-2", "stale alias") + pre_call_utils._warn_stale_team_alias_once("key-3", "stale alias") + + assert list(pre_call_utils._STALE_TEAM_ALIAS_WARNING_KEYS) == ["key-2", "key-3"] From 1e04aee089e2495a2917ab0f7972fffd6959e167 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 12:57:10 -0700 Subject: [PATCH 20/54] fix(proxy): reject model writes that corrupt an auto-router pseudo-model An auto-router deployment's litellm_params.model (auto_router/...) is the discriminator the router loads it by, but the model management endpoints accepted any client-supplied value verbatim; a doubled or stripped prefix made router init fail on the next load and ignore_invalid_deployments silently dropped the deployment. Validate writes that supply litellm_params.model at all three endpoints against the merged params and reject incoherent values with an actionable 400. Classification is extracted to router_utils/auto_router_model_naming.py so the Router predicates and the validation share one source --- .../model_management_endpoints.py | 59 +++++ litellm/router.py | 23 +- .../router_utils/auto_router_model_naming.py | 101 +++++++++ .../test_model_management_endpoints.py | 213 ++++++++++++++++++ .../test_auto_router_model_naming.py | 77 +++++++ 5 files changed, 457 insertions(+), 16 deletions(-) create mode 100644 litellm/router_utils/auto_router_model_naming.py create mode 100644 tests/test_litellm/router_utils/test_auto_router_model_naming.py diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 28c406edf76..1c0e7211493 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -57,10 +57,15 @@ from litellm.router import Router from litellm.types.proxy.management_endpoints.model_management_endpoints import ( UpdateUsefulLinksRequest, ) +from litellm.router_utils.auto_router_model_naming import ( + STRATEGY_ROUTER_PARAM_FIELDS, + validate_strategy_router_model_write, +) from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, Deployment, DeploymentTypedDict, + GenericLiteLLMParams, LiteLLMParamsTypedDict, updateDeployment, ) @@ -98,6 +103,45 @@ async def get_db_model(model_id: str, prisma_client: PrismaClient) -> Optional[D return deployment_pydantic_obj +def _strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> str | None: + """Reject writes that would corrupt a strategy router's pseudo-model. + + An auto-router deployment's ``litellm_params.model`` (``auto_router/...``) is + the discriminator the router loads it by; a write that mangles it makes the + router drop the deployment silently under ``ignore_invalid_deployments``. + Only writes that supply ``litellm_params.model`` are judged, against the + merged (stored + incoming) params, so partial patches and restores of an + already-corrupted row stay legal. Returns the violation, or None. + """ + if incoming_params is None or incoming_params.model is None: + return None + present_fields = frozenset( + field + for field in STRATEGY_ROUTER_PARAM_FIELDS + for source in (incoming_params, existing_params) + if source is not None and getattr(source, field, None) is not None + ) + return validate_strategy_router_model_write(model=incoming_params.model, present_fields=present_fields) + + +def _raise_on_strategy_router_write_violation( + incoming_params: GenericLiteLLMParams | None, + existing_params: GenericLiteLLMParams | None, +) -> None: + violation = _strategy_router_write_violation(incoming_params=incoming_params, existing_params=existing_params) + if violation is None: + return + raise ProxyException( + message=violation, + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_params.model", + ) + + def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel: merged_deployment_dict = DeploymentTypedDict( model_name=db_model.model_name, @@ -255,6 +299,11 @@ async def patch_model( param="blocked", ) + _raise_on_strategy_router_write_violation( + incoming_params=patch_data.litellm_params, + existing_params=db_model.litellm_params, + ) + # Handle team model updates with proper alias management update_data = await _update_team_model_in_db( db_model=db_model, @@ -1292,6 +1341,11 @@ async def add_new_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=None, + ) + model_response: Optional[LiteLLM_ProxyModelTable] = None # update DB if store_model_in_db is True: @@ -1446,6 +1500,11 @@ async def update_model( premium_user=premium_user, ) + _raise_on_strategy_router_write_violation( + incoming_params=model_params.litellm_params, + existing_params=deployment.litellm_params, + ) + # update DB if store_model_in_db is True: _existing_litellm_params_dict = dict(_existing_litellm_params.litellm_params) diff --git a/litellm/router.py b/litellm/router.py index d3caa069b28..759d6a0024c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -110,6 +110,9 @@ from litellm.router_utils.batch_utils import ( replace_model_in_jsonl, should_replace_model_in_jsonl, ) +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, +) from litellm.router_utils.client_initalization_utils import InitalizeCachedClient from litellm.router_utils.clientside_credential_handler import ( get_dynamic_litellm_params, @@ -7623,15 +7626,7 @@ class Router: but NOT "auto_router/complexity_router" or "auto_router/adaptive_router" (which use the complexity-router and adaptive-router strategies). """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return False # This is handled by complexity_router - if litellm_params.model.startswith("auto_router/adaptive_router"): - return False # This is handled by adaptive_router - if litellm_params.model.startswith("auto_router/quality_router"): - return False # This is handled by quality_router - if litellm_params.model.startswith("auto_router/"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "semantic" @staticmethod def _deployment_tags(deployment: Deployment) -> tuple[str, ...]: @@ -7686,9 +7681,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/complexity_router" """ - if litellm_params.model.startswith("auto_router/complexity_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "complexity" def init_complexity_router_deployment(self, deployment: Deployment): """ @@ -7738,7 +7731,7 @@ class Router: def _is_adaptive_router_deployment(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment opts in via the `auto_router/adaptive_router` model prefix.""" - return litellm_params.model.startswith("auto_router/adaptive_router") + return classify_strategy_router_model(litellm_params.model) == "adaptive" def _deployment_participates_in_adaptive_routing(self, litellm_params: LiteLLM_Params) -> bool: """True when this deployment owns an `adaptive_routers` entry once finalized: @@ -7964,9 +7957,7 @@ class Router: Returns True if the litellm_params model starts with "auto_router/quality_router". """ - if litellm_params.model.startswith("auto_router/quality_router"): - return True - return False + return classify_strategy_router_model(litellm_params.model) == "quality" def init_quality_router_deployment(self, deployment: Deployment): """ diff --git a/litellm/router_utils/auto_router_model_naming.py b/litellm/router_utils/auto_router_model_naming.py new file mode 100644 index 00000000000..72865501030 --- /dev/null +++ b/litellm/router_utils/auto_router_model_naming.py @@ -0,0 +1,101 @@ +"""Naming contract for strategy-router (auto-router) pseudo-models. + +A deployment whose ``litellm_params.model`` starts with ``auto_router/`` does not +name a provider model; the string is the discriminator that selects which +pre-routing strategy owns the deployment. This module is the single source of +truth for classifying that string (``Router._is_*_router_deployment`` delegates +here) and for checking that a client-supplied write leaves the deployment +coherent, so management endpoints can reject corruption with a 400 instead of +the router silently dropping the deployment at load time under +``ignore_invalid_deployments``. +""" + +from typing import Literal, Mapping + +AUTO_ROUTER_MODEL_PREFIX = "auto_router/" + +StrategyRouterKind = Literal["semantic", "complexity", "adaptive", "quality"] + +STRATEGY_ROUTER_PARAM_FIELDS: frozenset[str] = frozenset( + { + "auto_router_config", + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + "complexity_router_config", + "complexity_router_default_model", + "adaptive_router_config", + "quality_router_config", + "quality_router_default_model", + } +) + +_REQUIRED_FIELD_GROUPS: Mapping[StrategyRouterKind, tuple[tuple[str, ...], ...]] = { + "semantic": ( + ("auto_router_config", "auto_router_config_path"), + ("auto_router_default_model",), + ("auto_router_embedding_model",), + ), + "complexity": (("complexity_router_config", "complexity_router_default_model"),), + "adaptive": (("adaptive_router_config",),), + "quality": (("quality_router_config", "quality_router_default_model"),), +} + + +def classify_strategy_router_model(model: str) -> StrategyRouterKind | None: + """Classify a ``litellm_params.model`` string the way the Router does. + + Returns None for regular provider models. Mirrors Router registration + exactly: reserved names are matched by prefix, everything else under + ``auto_router/`` is a semantic router. + """ + if not model.startswith(AUTO_ROUTER_MODEL_PREFIX): + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith("complexity_router"): + return "complexity" + if remainder.startswith("adaptive_router"): + return "adaptive" + if remainder.startswith("quality_router"): + return "quality" + return "semantic" + + +def validate_strategy_router_model_write(model: str, present_fields: frozenset[str]) -> str | None: + """Check that writing ``model`` leaves a deployment the router can load. + + ``present_fields`` is the set of strategy-router param fields that are + non-None on the deployment after the write (stored fields merged with the + incoming ones). Returns a human-readable violation, or None when coherent. + """ + kind = classify_strategy_router_model(model) + if kind is None: + offending = sorted(present_fields & STRATEGY_ROUTER_PARAM_FIELDS) + if offending: + return ( + f"litellm_params.model='{model}' does not start with '{AUTO_ROUTER_MODEL_PREFIX}' but the " + f"deployment carries auto-router settings ({', '.join(offending)}), so the router could not " + f"load it. Keep the '{AUTO_ROUTER_MODEL_PREFIX}' prefix; to change the name clients call, " + "edit the public model_name instead." + ) + return None + remainder = model[len(AUTO_ROUTER_MODEL_PREFIX) :] + if remainder.startswith(AUTO_ROUTER_MODEL_PREFIX): + return ( + f"litellm_params.model='{model}' repeats the '{AUTO_ROUTER_MODEL_PREFIX}' prefix, so the router " + f"could not load it. Use '{remainder}'; to change the name clients call, edit the public " + "model_name instead." + ) + if not remainder: + return ( + f"litellm_params.model='{model}' is missing the router name after the '{AUTO_ROUTER_MODEL_PREFIX}' prefix." + ) + missing = tuple( + " or ".join(group) for group in _REQUIRED_FIELD_GROUPS[kind] if not any(f in present_fields for f in group) + ) + if missing: + return ( + f"litellm_params.model='{model}' selects the {kind} router, which requires " + f"{'; '.join(missing)} in litellm_params." + ) + return None diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 3bd83ade20a..1e8add52f74 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -3330,3 +3330,216 @@ class TestModelInfoAsMapping: assert model_info_as_mapping("{not json") is None assert model_info_as_mapping('["a", "b"]') is None assert model_info_as_mapping(42) is None + + +class TestStrategyRouterWriteValidation: + """Management write paths must reject litellm_params.model values that would + corrupt a strategy router's pseudo-model (LIT-4663). The router loads these + deployments by the auto_router/ discriminator, so a mangled string makes it + drop the deployment silently under ignore_invalid_deployments; the mistake + has to fail loudly at the API boundary instead.""" + + def _stored_complexity_params(self) -> LiteLLM_Params: + return LiteLLM_Params( + model="auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + + def _db_complexity_router(self, model_id: str) -> Deployment: + return Deployment( + model_name="my-auto-router", + litellm_params=self._stored_complexity_params(), + model_info={"id": model_id}, + ) + + def test_double_prefix_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "repeats" in violation + + def test_prefix_strip_rejected_against_stored_params(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + violation = _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="complexity_router"), + existing_params=self._stored_complexity_params(), + ) + assert violation is not None + assert "does not start with" in violation + + def test_patch_without_model_is_not_judged(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(rpm=10), + existing_params=self._stored_complexity_params(), + ) + is None + ) + assert _strategy_router_write_violation(incoming_params=None, existing_params=None) is None + + def test_restore_of_corrupted_row_is_allowed(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + from litellm.types.router import updateLiteLLMParams + + corrupted = LiteLLM_Params( + model="auto_router/auto_router/complexity_router", + complexity_router_config={"tiers": {"SIMPLE": "gpt-4o-mini"}}, + ) + assert ( + _strategy_router_write_violation( + incoming_params=updateLiteLLMParams(model="auto_router/complexity_router"), + existing_params=corrupted, + ) + is None + ) + + def test_create_semantic_router_missing_embedding_rejected(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _strategy_router_write_violation, + ) + + violation = _strategy_router_write_violation( + incoming_params=LiteLLM_Params( + model="auto_router/my-router", + auto_router_config="{}", + auto_router_default_model="gpt-4o-mini", + ), + existing_params=None, + ) + assert violation is not None + assert "auto_router_embedding_model" in violation + + @pytest.mark.asyncio + async def test_patch_model_rejects_double_prefix(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + from litellm.types.router import updateLiteLLMParams + + model_id = "strategy-router-patch-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.get_db_model", + new=AsyncMock(return_value=self._db_complexity_router(model_id)), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints._update_team_model_in_db", + new=AsyncMock(), + ) as mock_update, + ): + with pytest.raises(ProxyException) as exc_info: + await patch_model( + model_id=model_id, + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(model="auto_router/auto_router/complexity_router") + ), + user_api_key_dict=admin, + ) + assert "repeats" in str(exc_info.value.message) + mock_update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_add_new_model_rejects_prefixed_model_without_config(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-auto-router", + litellm_params=LiteLLM_Params(model="auto_router/complexity_router"), + model_info={"id": "strategy-router-create-test"}, + ), + user_api_key_dict=admin, + ) + assert "requires" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_update_model_rejects_prefix_strip(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + from litellm.types.router import ModelInfo, updateLiteLLMParams + + model_id = "strategy-router-update-test" + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "model_name": "my-auto-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_config": {"tiers": {"SIMPLE": "gpt-4o-mini"}}, + }, + "model_info": {"id": model_id}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(model="complexity_router"), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=admin, + ) + assert "does not start with" in str(exc_info.value.message) + mock_prisma.db.litellm_proxymodeltable.update.assert_not_awaited() diff --git a/tests/test_litellm/router_utils/test_auto_router_model_naming.py b/tests/test_litellm/router_utils/test_auto_router_model_naming.py new file mode 100644 index 00000000000..ca290caac0b --- /dev/null +++ b/tests/test_litellm/router_utils/test_auto_router_model_naming.py @@ -0,0 +1,77 @@ +import pytest + +from litellm.router_utils.auto_router_model_naming import ( + classify_strategy_router_model, + validate_strategy_router_model_write, +) + +COMPLEXITY_FIELDS = frozenset({"complexity_router_config"}) +SEMANTIC_FIELDS = frozenset( + {"auto_router_config", "auto_router_default_model", "auto_router_embedding_model"} +) + + +@pytest.mark.parametrize( + "model,expected", + [ + ("anthropic/claude-sonnet-5", None), + ("complexity_router", None), + ("autorouter/complexity_router", None), + ("auto_router/my-router", "semantic"), + ("auto_router/complexity_router", "complexity"), + ("auto_router/complexity_router-eu", "complexity"), + ("auto_router/adaptive_router", "adaptive"), + ("auto_router/quality_router", "quality"), + ("auto_router/auto_router/complexity_router", "semantic"), + ("auto_router/", "semantic"), + ], +) +def test_classify_strategy_router_model(model, expected): + assert classify_strategy_router_model(model) == expected + + +@pytest.mark.parametrize( + "model,present_fields,expected_fragment", + [ + ("auto_router/auto_router/complexity_router", COMPLEXITY_FIELDS, "repeats"), + ("complexity_router", COMPLEXITY_FIELDS, "does not start with"), + ("anthropic/claude-sonnet-5", COMPLEXITY_FIELDS, "does not start with"), + ("auto_router/", frozenset(), "missing the router name"), + ("auto_router/complexity_router", frozenset(), "requires"), + ("auto_router/my-router", frozenset({"auto_router_config"}), "requires"), + ("auto_router/adaptive_router", frozenset(), "requires"), + ("auto_router/quality_router", frozenset(), "requires"), + ], +) +def test_validate_rejects_incoherent_writes(model, present_fields, expected_fragment): + violation = validate_strategy_router_model_write(model=model, present_fields=present_fields) + assert violation is not None + assert expected_fragment in violation + + +@pytest.mark.parametrize( + "model,present_fields", + [ + ("anthropic/claude-sonnet-5", frozenset()), + ("openai/gpt-4o-mini", frozenset({"api_key"})), + ("auto_router/complexity_router", COMPLEXITY_FIELDS), + ("auto_router/complexity_router", frozenset({"complexity_router_default_model"})), + ("auto_router/complexity_router-eu", COMPLEXITY_FIELDS), + ("auto_router/my-router", SEMANTIC_FIELDS), + ( + "auto_router/my-router", + frozenset( + { + "auto_router_config_path", + "auto_router_default_model", + "auto_router_embedding_model", + } + ), + ), + ("auto_router/adaptive_router", frozenset({"adaptive_router_config"})), + ("auto_router/quality_router", frozenset({"quality_router_default_model"})), + ("auto_router/quality_router", frozenset({"quality_router_config"})), + ], +) +def test_validate_accepts_coherent_writes(model, present_fields): + assert validate_strategy_router_model_write(model=model, present_fields=present_fields) is None From 6d607ca3c228060aa9b706a46e5b00a487798047 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 20:47:40 -0700 Subject: [PATCH 21/54] fix(router): never resolve another team's deployment credentials for shared model names --- .../openai_files_endpoints/common_utils.py | 2 +- litellm/router.py | 42 ++++++- .../test_files_endpoint.py | 84 +++++++++++++ tests/test_litellm/test_router.py | 119 ++++++++++++++++++ 4 files changed, 243 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index d4ac45559ee..dca4b2a3773 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -319,7 +319,7 @@ def get_team_provider_credentials( return None def _provider_credentials(model_id: str) -> Optional[dict]: - credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id) + credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id) if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider: return credentials return None diff --git a/litellm/router.py b/litellm/router.py index 487d6a31226..38450fe4ef1 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -30,6 +30,7 @@ from typing import ( Generator, List, Literal, + Mapping, Optional, Set, Tuple, @@ -8630,6 +8631,33 @@ class Router: raise Exception("Model Name invalid - {}".format(type(model))) return None + @staticmethod + def _deployment_usable_by_team(model: Union[Mapping, Deployment], team_id: str | None) -> bool: + """ + A team-scoped deployment (``model_info.team_id`` set) is only usable by + callers from that same team; deployments without a team owner are shared. + """ + model_info = model.get("model_info") if isinstance(model, dict) else model.model_info + owner_team_id = model_info.get("team_id") if model_info is not None else None + return owner_team_id is None or owner_team_id == team_id + + def _get_model_group_deployment_usable_by_team( + self, model_group_name: str, team_id: str | None + ) -> Deployment | None: + """ + Like ``get_deployment_by_model_group_name``, but skips deployments owned + by other teams so a shared model name never resolves another team's + credentials. + """ + indices = self.model_name_to_deployment_indices.get(model_group_name) or () + usable = ( + self.model_list[idx] for idx in indices if self._deployment_usable_by_team(self.model_list[idx], team_id) + ) + first_usable = next(usable, None) + if first_usable is None: + return None + return Deployment(**first_usable) if isinstance(first_usable, dict) else first_usable + def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]": """ Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete @@ -8664,7 +8692,10 @@ class Router: model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm") team_id: Optional team id of the caller. When set, team-scoped deployments (indexed by team public model name, including team - wildcard models like "openai/*") are also considered. + wildcard models like "openai/*") are also considered. Name and + wildcard lookups never resolve a deployment owned by a + different team, so shared model names can't leak another + team's credentials. Returns: Dictionary containing api_key, api_base, custom_llm_provider, etc. @@ -8681,7 +8712,7 @@ class Router: # If not found, try by model_group_name if deployment is None: - deployment = self.get_deployment_by_model_group_name(model_group_name=model_id) + deployment = self._get_model_group_deployment_usable_by_team(model_group_name=model_id, team_id=team_id) # If not found, check team-scoped deployments whose team public model # name exactly matches model_id (wildcard team names are matched via @@ -8698,7 +8729,12 @@ class Router: if deployment is None: team_pattern_router = self.team_pattern_routers.get(team_id) if team_id is not None else None team_wildcard_models = (team_pattern_router.route(model_id) or []) if team_pattern_router else [] - potential_wildcard_models = team_wildcard_models or self.pattern_router.route(model_id) or [] + global_wildcard_models = [ + wildcard_model + for wildcard_model in (self.pattern_router.route(model_id) or []) + if self._deployment_usable_by_team(wildcard_model, team_id) + ] + potential_wildcard_models = team_wildcard_models or global_wildcard_models if potential_wildcard_models: # Use the first matching wildcard deployment deployment_dict = potential_wildcard_models[0] diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 23cd1c71cbe..17cd3f6172a 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2867,3 +2867,87 @@ def test_delete_file_provider_only_resolves_named_vertex_credentials( assert captured_kwargs.get("file_id") == "file-abc123" _assert_vertex_named_credentials_attached(captured_kwargs) proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def test_create_file_provider_only_skips_other_team_vertex_deployment( + mocker: MockerFixture, monkeypatch +): + """ + Regression: with a team-scoped vertex deployment indexed before a global + one under the same model name, a provider-only upload from a different + team must use the global deployment's credentials, never the other + team's. + """ + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + router = Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ] + ) + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_acreate_file(**kwargs): + captured_kwargs.update(kwargs) + return OpenAIFileObject( + id="file-vertex-456", + object="file", + bytes=2, + created_at=1234567890, + filename="batch.jsonl", + purpose="batch", + status="uploaded", + ) + + monkeypatch.setattr(litellm, "acreate_file", _mock_acreate_file) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + team_id="team-a", + team_models=["gemini-2.5-pro"], + ) + + try: + response = client.post( + "/v1/files", + files={"file": ("batch.jsonl", b"{}", "application/jsonl")}, + data={"purpose": "batch"}, + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "vertex_ai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + assert captured_kwargs.get("vertex_project") == "shared-project" + proxy_logging_obj.post_call_failure_hook.assert_not_called() diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index ad4e430c603..8dcefc0405a 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3755,6 +3755,125 @@ def test_get_deployment_credentials_with_provider_team_wildcard_priority(): assert global_credentials["api_key"] == "global-key" +def test_get_deployment_credentials_with_provider_skips_other_team_deployment(): + """ + Regression: a team-scoped deployment sharing a model_name with a global + deployment must never resolve for another team's (or an unscoped) caller, + even when it is indexed first; the shared global deployment wins instead. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ], + ) + + other_team_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-a" + ) + assert other_team_credentials is not None + assert other_team_credentials["vertex_project"] == "shared-project" + + unscoped_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro" + ) + assert unscoped_credentials is not None + assert unscoped_credentials["vertex_project"] == "shared-project" + + owner_credentials = router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-b" + ) + assert owner_credentials is not None + assert owner_credentials["vertex_project"] == "team-b-project" + + +def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only_name(): + """ + When the only deployments under a model name belong to another team, other + callers must get None (env fallback) instead of that team's credentials. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + ], + ) + + assert ( + router.get_deployment_credentials_with_provider( + model_id="gemini-2.5-pro", team_id="team-a" + ) + is None + ) + assert ( + router.get_deployment_credentials_with_provider(model_id="gemini-2.5-pro") + is None + ) + + +def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): + """ + Global wildcard resolution must skip a team-scoped wildcard deployment for + callers outside that team, falling through to the shared wildcard entry. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": {"model": "openai/*", "api_key": "team-b-key"}, + "model_info": { + "id": "team-b-wildcard", + "team_id": "team-b", + "team_public_model_name": "openai/*", + }, + }, + { + "model_name": "openai/*", + "litellm_params": {"model": "openai/*", "api_key": "global-key"}, + }, + ], + ) + + other_team_credentials = router.get_deployment_credentials_with_provider( + model_id="openai/gpt-5.2", team_id="team-a" + ) + assert other_team_credentials is not None + assert other_team_credentials["api_key"] == "global-key" + + owner_credentials = router.get_deployment_credentials_with_provider( + model_id="openai/gpt-5.2", team_id="team-b" + ) + assert owner_credentials is not None + assert owner_credentials["api_key"] == "team-b-key" + + def test_team_wildcard_credentials_not_usable_after_delete_deployment(): """ Regression: team_pattern_routers retained deleted deployments, so a team From 47a9fabb5a3dc3029ed504a959a635f12dea3b1c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 21:11:50 -0700 Subject: [PATCH 22/54] fix(proxy): honor key-level model allowlist in provider-only credential resolution --- .../openai_files_endpoints/common_utils.py | 80 ++++++++++---- .../test_files_endpoint.py | 100 ++++++++++++++++++ 2 files changed, 159 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index dca4b2a3773..b2e36188681 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -295,9 +295,8 @@ def get_credentials_for_model( def get_team_provider_credentials( llm_router: Optional["Router"], - team_models: List[str], + user_api_key_dict: "UserAPIKeyAuth", custom_llm_provider: str, - team_id: Optional[str] = None, ) -> Optional[dict]: """ Resolve upstream credentials for a provider-scoped file operation @@ -305,19 +304,59 @@ def get_team_provider_credentials( Priority: 1. The team's own (BYOK) deployment for this provider — a deployment whose - ``model_info.team_id`` matches ``team_id``. This keeps team-scoped listings - on the team's own provider account/key instead of a shared global one. - 2. Fallback: any deployment the team is granted access to for this provider, - expanding wildcard routes and the all-proxy-models sentinel. + ``model_info.team_id`` matches the caller's team. This keeps team-scoped + listings on the team's own provider account/key instead of a shared + global one. + 2. Fallback: any deployment the caller is granted access to for this + provider, expanding wildcard routes and the all-proxy-models sentinel. - Credential lookup is always scoped to the team's allowlist, so a team can - never resolve a provider key for a deployment it isn't authorized to use. + Credential lookup is scoped to both the team's allowlist and the key's own + model allowlist (``user_api_key_dict.models``), so neither a team nor a + restricted key within a team can resolve a provider key for a deployment + it isn't authorized to use. A key restricted to an explicit model list + only narrows the team scope; sentinel-bearing keys (all-proxy-models / + all-team-models) defer to the team scope instead of widening past it. Returns None when the router is unavailable or no authorized deployment matches, so the caller can fall back to default credential resolution. """ if llm_router is None: return None + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.model_checks import get_complete_model_list, get_key_models + + team_id = user_api_key_dict.team_id + team_models = user_api_key_dict.team_models or [] + + proxy_model_list = llm_router.get_model_names(team_id=team_id) + model_access_groups = llm_router.get_model_access_groups() + + raw_key_models = user_api_key_dict.models or [] + sentinel_values = { + SpecialModelNames.all_proxy_models.value, + SpecialModelNames.all_team_models.value, + } + key_is_restricted = bool(raw_key_models) and not (set(raw_key_models) & sentinel_values) + key_model_allowlist = ( + tuple( + dict.fromkeys( + get_key_models( + user_api_key_dict=user_api_key_dict, + proxy_model_list=proxy_model_list, + model_access_groups=model_access_groups, + ) + ) + ) + if key_is_restricted + else () + ) + key_model_allowlist_set = frozenset(key_model_allowlist) + + def _key_may_use(public_model_name: Optional[str]) -> bool: + if not key_model_allowlist_set: + return True + return public_model_name is not None and public_model_name in key_model_allowlist_set + def _provider_credentials(model_id: str) -> Optional[dict]: credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id) if credentials is not None and credentials.get("custom_llm_provider") == custom_llm_provider: @@ -333,27 +372,27 @@ def get_team_provider_credentials( deployment_id = model_info.get("id") if deployment_id is None: continue + if not _key_may_use(model_info.get("team_public_model_name") or deployment.get("model_name")): + continue credentials = _provider_credentials(deployment_id) if credentials is not None: return credentials - # 2. Fall back to deployments the team is allowed to access. The - # all-proxy-models sentinel isn't expanded by get_complete_model_list, so - # normalize it to an empty allowlist, which defers to the team-scoped - # proxy model list. A team with a restricted allowlist (e.g. anthropic - # only) therefore never resolves another provider's key. - from litellm.proxy._types import SpecialModelNames - from litellm.proxy.auth.model_checks import get_complete_model_list - + # 2. Fall back to deployments the caller is allowed to access. The key's + # effective allowlist (sentinels and access groups already expanded by + # get_key_models) wins when set; otherwise the team's allowlist applies. + # The all-proxy-models sentinel isn't expanded by + # get_complete_model_list, so normalize it to an empty allowlist, which + # defers to the team-scoped proxy model list. A team or key with a + # restricted allowlist (e.g. anthropic only) therefore never resolves + # another provider's key. grants_all_models = SpecialModelNames.all_proxy_models.value in team_models effective_team_models = [] if grants_all_models else team_models - proxy_model_list = llm_router.get_model_names(team_id=team_id) - model_access_groups = llm_router.get_model_access_groups() models_to_try = list( dict.fromkeys( get_complete_model_list( - key_models=[], + key_models=list(key_model_allowlist), team_models=effective_team_models, proxy_model_list=proxy_model_list, user_model=None, @@ -388,9 +427,8 @@ def apply_team_provider_credentials( """ credentials = get_team_provider_credentials( llm_router=llm_router, - team_models=user_api_key_dict.team_models or [], + user_api_key_dict=user_api_key_dict, custom_llm_provider=custom_llm_provider, - team_id=user_api_key_dict.team_id, ) if credentials is None: return diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 17cd3f6172a..ac01c6ae1d1 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -2951,3 +2951,103 @@ def test_create_file_provider_only_skips_other_team_vertex_deployment( assert response.status_code == 200, response.text assert captured_kwargs.get("vertex_project") == "shared-project" proxy_logging_obj.post_call_failure_hook.assert_not_called() + + +def _team_openai_plus_global_anthropic_router() -> Router: + return Router( + model_list=[ + { + "model_name": "team-gpt", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "team-openai-key", + }, + "model_info": { + "id": "team-a-openai", + "team_id": "team-a", + "team_public_model_name": "team-gpt", + }, + }, + { + "model_name": "claude-opus-4-6", + "litellm_params": { + "model": "anthropic/claude-opus-4-6", + "api_key": "anthropic-key", + }, + }, + ] + ) + + +def _list_files_captured_kwargs( + mocker: MockerFixture, monkeypatch, router: Router, key_models: list +) -> dict: + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + proxy_logging_obj = setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + proxy_logging_obj.update_request_status = mocker.AsyncMock() + proxy_logging_obj.post_call_success_hook = mocker.AsyncMock(return_value=[]) + proxy_logging_obj.post_call_failure_hook = mocker.AsyncMock() + + captured_kwargs: dict = {} + + async def _mock_afile_list(**kwargs): + captured_kwargs.update(kwargs) + return [] + + monkeypatch.setattr(litellm, "afile_list", _mock_afile_list) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="test-user", + team_id="team-a", + team_models=["team-gpt", "claude-opus-4-6"], + models=key_models, + ) + + try: + response = client.get( + "/v1/files", + headers={ + "Authorization": "Bearer test-key", + "custom-llm-provider": "openai", + }, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + return captured_kwargs + + +def test_list_files_key_restricted_to_other_provider_does_not_leak_team_openai_credentials( + mocker: MockerFixture, monkeypatch +): + """ + Regression: a key restricted to an anthropic model on a team that also has + an openai deployment must not attach the team's openai credentials to a + provider-only openai files call; key-level model restrictions apply to + credential resolution, not just completions. + """ + captured_kwargs = _list_files_captured_kwargs( + mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["claude-opus-4-6"] + ) + assert captured_kwargs.get("api_key") != "team-openai-key" + + +def test_list_files_key_allowed_openai_model_still_resolves_team_credentials( + mocker: MockerFixture, monkeypatch +): + """ + A key whose allowlist includes the team's openai model keeps resolving that + deployment's credentials for provider-only openai files calls. + """ + captured_kwargs = _list_files_captured_kwargs( + mocker, monkeypatch, _team_openai_plus_global_anthropic_router(), ["team-gpt"] + ) + assert captured_kwargs.get("api_key") == "team-openai-key" From 6e8655762c3f06e44e1b916f8f230fb47da75bf6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 28 Jul 2026 21:12:57 -0700 Subject: [PATCH 23/54] test(router): directly cover team-ownership credential filter helpers --- tests/test_litellm/test_router.py | 57 +++++++++++++++++++++++++++++++ 1 file changed, 57 insertions(+) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 8dcefc0405a..4df48a6dc56 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3838,6 +3838,63 @@ def test_get_deployment_credentials_with_provider_no_fallback_to_other_team_only ) +def test_deployment_usable_by_team_helpers(): + """ + Direct coverage of the team-ownership filter: a team-scoped deployment is + usable only by its owning team, shared deployments by anyone, and the + model-group picker returns the first usable deployment or None. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "team-b-project", + }, + "model_info": { + "id": "team-b-vertex", + "team_id": "team-b", + "team_public_model_name": "gemini-2.5-pro", + }, + }, + { + "model_name": "gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + "vertex_project": "shared-project", + }, + }, + ], + ) + + team_owned, shared = router.model_list + assert router._deployment_usable_by_team(team_owned, "team-b") is True + assert router._deployment_usable_by_team(team_owned, "team-a") is False + assert router._deployment_usable_by_team(team_owned, None) is False + assert router._deployment_usable_by_team(shared, "team-a") is True + assert router._deployment_usable_by_team(shared, None) is True + + picked = router._get_model_group_deployment_usable_by_team( + model_group_name="gemini-2.5-pro", team_id="team-a" + ) + assert picked is not None + assert picked.litellm_params.vertex_project == "shared-project" + + owner_picked = router._get_model_group_deployment_usable_by_team( + model_group_name="gemini-2.5-pro", team_id="team-b" + ) + assert owner_picked is not None + assert owner_picked.litellm_params.vertex_project == "team-b-project" + + assert ( + router._get_model_group_deployment_usable_by_team( + model_group_name="unknown-model", team_id="team-a" + ) + is None + ) + + def test_get_deployment_credentials_with_provider_skips_other_team_wildcard(): """ Global wildcard resolution must skip a team-scoped wildcard deployment for From 2f7574d7c1767830c0163b331215010b38f0eed7 Mon Sep 17 00:00:00 2001 From: Napuh <55241721+Napuh@users.noreply.github.com> Date: Wed, 29 Jul 2026 06:51:59 +0200 Subject: [PATCH 24/54] fix(anthropic-adapter): open the first content block with the real upstream type so reasoning-first streams start with thinking (#34433) * fix(anthropic-adapter): open first content block with the real upstream type * fix(anthropic): defer blank leading stream deltas --- .../adapters/streaming_iterator.py | 75 +++++---- .../test_streaming_iterator_first_delta.py | 159 ++++++++++++++++++ .../messages/test_parallel_tool_calls.py | 21 --- 3 files changed, 205 insertions(+), 50 deletions(-) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index f02333c34c8..853bea636af 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if compaction_event is not None: return compaction_event - if self.sent_content_block_start is False: - self.sent_content_block_start = True - self.sent_content_block_finish = False - self.chunk_queue.append( - { - "type": "content_block_start", - "index": self.current_content_block_index, - "content_block": {"type": "text", "text": ""}, - } - ) - return self.chunk_queue.popleft() - for chunk in self.completion_stream: if chunk == "None" or chunk is None: raise Exception should_start_new_block = self._should_start_new_content_block(chunk) - if should_start_new_block: + is_opening_first_block = self.sent_content_block_start is False + if is_opening_first_block and self._is_blank_delta(chunk): + continue + if is_opening_first_block: + self.sent_content_block_start = True + self.sent_content_block_finish = False + self.chunk_queue.append( + { + "type": "content_block_start", + "index": self.current_content_block_index, + "content_block": self.current_content_block_start, + } + ) + elif should_start_new_block: self._increment_content_block_index() # applied_edits only needs to flow to the final message_delta @@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # ``not self.queued_usage_chunk``. continue - if should_start_new_block and not self.sent_content_block_finish: + if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start # -> (optionally) the trigger chunk's delta. # @@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): if compaction_event is not None: return compaction_event - if self.sent_content_block_start is False: - self.sent_content_block_start = True - self.sent_content_block_finish = False - self.chunk_queue.append( - { - "type": "content_block_start", - "index": self.current_content_block_index, - "content_block": {"type": "text", "text": ""}, - } - ) - return self.chunk_queue.popleft() - async for chunk in self.completion_stream: if chunk == "None" or chunk is None: raise Exception - # Check if we need to start a new content block should_start_new_block = self._should_start_new_content_block(chunk) - if should_start_new_block: + is_opening_first_block = self.sent_content_block_start is False + if is_opening_first_block and self._is_blank_delta(chunk): + continue + if is_opening_first_block: + self.sent_content_block_start = True + self.sent_content_block_finish = False + self.chunk_queue.append( + { + "type": "content_block_start", + "index": self.current_content_block_index, + "content_block": self.current_content_block_start, + } + ) + elif should_start_new_block: self._increment_content_block_index() # applied_edits only needs to flow to the final message_delta @@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): # Check if this processed chunk has a stop_reason - hold it for next chunk if not self.queued_usage_chunk: - if should_start_new_block and not self.sent_content_block_finish: + if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish: # Queue the sequence: content_block_stop -> content_block_start # -> (optionally) the trigger chunk's delta. # @@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): return False return bool(delta.get(_delta_payload_field(delta_type))) + @staticmethod + def _is_blank_delta(chunk: "ModelResponseStream") -> bool: + choice = chunk.choices[0] + if choice.finish_reason is not None: + return False + delta = choice.delta + if getattr(delta, "tool_calls", None): + return False + if getattr(delta, "content", None): + return False + if getattr(delta, "reasoning_content", None): + return False + if getattr(delta, "thinking_blocks", None): + return False + return True + def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool: """ Determine if we should start a new content block based on the processed chunk. diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py index 19ec1a04b45..17a57d974de 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_streaming_iterator_first_delta.py @@ -506,3 +506,162 @@ def test_empty_content_chunk_mid_text_block_is_suppressed_sync(): assert _text_deltas(events) == ["Hi", " there"] _assert_deltas_match_their_block_type(events) + + +def _thinking_first_chunks() -> List[MagicMock]: + return [ + _thinking_chunk("Let me think"), + _thinking_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def _assert_thinking_first_block_opens_at_index_zero(events: List[dict]) -> None: + starts = [ + (e["index"], e["content_block"]["type"]) + for e in events + if e.get("type") == "content_block_start" + ] + assert starts == [(0, "thinking"), (1, "text")], starts + assert "" not in _text_deltas(events) + assert _thinking_deltas(events) == ["Let me think", "about it."] + assert _text_deltas(events) == ["42"] + _assert_deltas_match_their_block_type(events) + + +def test_thinking_first_stream_opens_thinking_block_at_index_zero_sync(): + """Bug A regression: when the model's first output is reasoning the adapter + must open the first content block as ``thinking`` at index 0. The previous + code pre-emitted a hardcoded empty ``text`` block at index 0 before + inspecting any upstream chunk, then opened ``thinking`` at index 1; strict + Anthropic SDK clients with thinking enabled reject that stream with + "Content block is not a thinking block". + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_thinking_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +@pytest.mark.asyncio +async def test_thinking_first_stream_opens_thinking_block_at_index_zero_async(): + """Async twin of the Bug A regression; the proxy serves the async iterator, + so the first block must be ``thinking`` at index 0 on this path too. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_thinking_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +def _reasoning_content_chunk(reasoning: str) -> MagicMock: + return _make_chunk(Delta(content=None, reasoning_content=reasoning)) + + +def _reasoning_first_chunks() -> List[MagicMock]: + return [ + _reasoning_content_chunk("Let me think"), + _reasoning_content_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def test_reasoning_content_first_stream_opens_thinking_block_at_index_zero_sync(): + """The reported backend (hosted_vllm; vLLM and SGLang reasoning parsers) + surfaces reasoning as OpenAI ``reasoning_content`` with no + ``thinking_blocks``. Such a stream must also open the first content block as + ``thinking`` at index 0, exercising the reasoning_content branch of the + chunk translator rather than the thinking_blocks branch the other twins use. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_reasoning_first_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +def _blank_lead_chunks() -> List[MagicMock]: + return [ + _make_chunk(Delta(content=None)), + _thinking_chunk("Let me think"), + _thinking_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def _role_only_reasoning_content_lead_chunks() -> List[MagicMock]: + return [ + _make_chunk(Delta(role="assistant", content=None, tool_calls=[])), + _reasoning_content_chunk("Let me think"), + _reasoning_content_chunk("about it."), + _make_chunk(Delta(content="42")), + _make_chunk(Delta(content=None), finish_reason="stop"), + ] + + +def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_sync(): + """OpenAI-compatible streaming backends open the response with a contentless + priming chunk (an empty delta, e.g. the {role: assistant} lead-in) before the + first real token. Such a lead chunk must NOT commit index 0 to an empty text + block; the following thinking chunk must still open thinking at index 0, or + strict Anthropic SDK clients reject the stream. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_blank_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_sync(): + wrapper = AnthropicStreamWrapper( + completion_stream=iter(_role_only_reasoning_content_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(_drain_sync(wrapper)) + + +@pytest.mark.asyncio +async def test_contentless_lead_chunk_does_not_open_text_block_before_thinking_async(): + """Async twin; the proxy serves the async iterator, so the contentless lead + chunk must be skipped on this path too. + """ + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_blank_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +@pytest.mark.asyncio +async def test_role_only_lead_chunk_does_not_open_text_block_before_reasoning_content_async(): + wrapper = AnthropicStreamWrapper( + completion_stream=_AsyncStream(_role_only_reasoning_content_lead_chunks()), + model="claude-x", + ) + _assert_thinking_first_block_opens_at_index_zero(await _drain_async(wrapper)) + + +def test_finish_first_chunk_is_not_deferred_sync(): + """A stream whose first upstream chunk is already the finish event must not + be skipped by the blank-delta deferral. ``_is_blank_delta`` returns False + for a finish chunk so the message_delta still flows (with an empty text + block opened and closed first); without that guard the deferral would drop + the terminal event entirely. + """ + chunks = [_make_chunk(Delta(content=None), finish_reason="stop")] + wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="claude-x") + events = _drain_sync(wrapper) + + assert [e["type"] for e in events] == [ + "message_start", + "content_block_start", + "content_block_stop", + "message_delta", + "message_stop", + ] diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py index e44413cf837..6d7cd2f88be 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_parallel_tool_calls.py @@ -134,13 +134,6 @@ def test_anthropic_stream_wrapper_single_tool_call(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} @@ -196,13 +189,6 @@ def test_anthropic_stream_wrapper_back_to_back_tool_calls(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} @@ -267,13 +253,6 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): # Verify the expected sequence of chunk types expected_types = [ "message_start", # Initial message start - # TODO: for future contributors: if the initial content_block_start - # respects the upstream's starting chunk, the initial empty text block - # should be removed (and this test should be updated accordingly) - # --------------------------------------------------------------------- - "content_block_start", # Initial empty text block start - "content_block_stop", # End of empty text block - # --------------------------------------------------------------------- "content_block_start", # Start of first tool_use content block "content_block_delta", # {"city": "content_block_delta", # "NY"} From c274cf321c5c35c629220a89bb497d15b56f870f Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Tue, 28 Jul 2026 22:15:21 -0700 Subject: [PATCH 25/54] test(e2e): poll MCP tools across multi-worker lag (#35047) * fix(mcp): resolve call_tool by registry without requiring tool map Multi-worker reloads put MCP servers in the registry from the DB but do not re-run tools/list on every process. Gating call_tool on tool_name_to_mcp_server_name_mapping made cold workers 500 with Tool not found after another worker had already listed the tool. Treat a registry match on server id/name/alias as enough; upstream rejects unknown tools * test(e2e): poll MCP register, tools/list, and tools/call across multi-worker lag Stage multi-worker gateways only load MCP servers and tool maps on the process that handled the request. Poll until the server is listed, the tool appears on tools/list, and tools/call is not a cold-worker 500 so key-access and Datadog MCP e2e stop racing the LB * Revert "fix(mcp): resolve call_tool by registry without requiring tool map" This reverts commit 8b56e51e39b876d13d1112efa4130554ddf5f173. * test(e2e): tighten MCP multi-worker lag classifier Only retry tools/call on gateway shapes Tool not found and server_not_found, not any 500 that mentions tool/server not found, so upstream failures are not retried until the poll deadline * test(e2e): drop unit file for MCP lag classifier The live await_call_tool polls already cover multi-worker lag; a separate string-match unit module is not worth keeping --- tests/e2e/mcp/mcp_client.py | 91 +++++++++++++++++++++- tests/e2e/mcp/test_mcp_access_group_e2e.py | 1 + tests/e2e/mcp/test_mcp_datadog_e2e.py | 28 ++++--- tests/e2e/mcp/test_mcp_key_access_e2e.py | 15 ++-- 4 files changed, 111 insertions(+), 24 deletions(-) diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index 6d0f6ddc760..33ec557c339 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -11,13 +11,14 @@ request/response bodies are co-located here because only this suite speaks MCP. from __future__ import annotations +import re import time from collections.abc import Mapping from dataclasses import dataclass from pydantic import BaseModel, ConfigDict, Field, RootModel -from e2e_http import Headers, NoBody, Result, Success, unwrap +from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap from models import KeyGenerateBody, ObjectPermission from proxy_client import ProxyClient @@ -270,6 +271,60 @@ class McpClient: ) time.sleep(self.proxy.poll_interval) + def await_call_tool( + self, + key: str, + *, + server_id: str, + name: str, + arguments: McpToolArguments, + ) -> McpCallToolResponse: + """Poll tools/call until the result is not a multi-worker registry miss. + + Retries only on the gateway's own cold-worker 500 shapes (Tool + not found / server_not_found). Upstream tool errors and other 500s fail + immediately so non-idempotent calls are not repeated. + """ + deadline = time.monotonic() + self.proxy.poll_timeout + last: Result[McpCallToolResponse] | None = None + while True: + last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments) + if not _is_mcp_not_synced(last, tool_name=name): + return unwrap(last) + if time.monotonic() >= deadline: + raise AssertionError( + f"tools/call for {name!r} on server {server_id} still missing on the " + f"data plane after {self.proxy.poll_timeout}s (multi-worker registry lag); " + f"last result: {last}" + ) + time.sleep(self.proxy.poll_interval) + + def await_call_tool_denied( + self, + key: str, + *, + server_id: str, + name: str, + arguments: McpToolArguments, + ) -> UnknownApiError: + """Poll tools/call until a cold-worker miss clears and the call is 403 access_denied.""" + deadline = time.monotonic() + self.proxy.poll_timeout + last: Result[McpCallToolResponse] | None = None + while True: + last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments) + if isinstance(last, UnknownApiError) and last.status_code == 403: + return last + if not _is_mcp_not_synced(last, tool_name=name): + raise AssertionError( + f"ungranted key's tools/call was not 403 access_denied: {last}" + ) + if time.monotonic() >= deadline: + raise AssertionError( + f"ungranted key never got 403 for {name!r} within {self.proxy.poll_timeout}s; " + f"last result: {last}" + ) + time.sleep(self.proxy.poll_interval) + def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str: """Register a default-on content-filter guardrail that runs on the MCP tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is @@ -317,5 +372,39 @@ class McpClient: ) +def _is_mcp_not_synced( + result: Result[McpCallToolResponse], + *, + tool_name: str | None = None, +) -> bool: + """True only for gateway multi-worker registry misses, not upstream errors. + + Matches the proxy's own shapes: + - ValueError ``Tool not found`` wrapped as HTTP 500 (cold tool map / + unresolved server on this process) + - REST ``server_not_found`` when this worker has not loaded the MCP server row + + Does not treat arbitrary 500 bodies that merely mention "tool" and "not found" + (e.g. upstream MCP payload text) as lag, so await_call_tool does not retry + real failures or non-idempotent calls. + """ + if not isinstance(result, UnknownApiError) or result.status_code != 500: + return False + body = result.body + body_l = body.lower() + + if "server_not_found" in body_l: + return True + if re.search(r"mcp server ['\"][^'\"]+['\"] was not found", body_l): + return True + + # Gateway: "Tool search_datadog_logs not found" (optionally inside a longer message) + if tool_name is not None: + return ( + re.search(rf"\btool\s+{re.escape(tool_name)}\s+not found\b", body_l) is not None + ) + return re.search(r"\btool\s+\S+\s+not found\b", body_l) is not None + + def build_client(proxy: ProxyClient) -> McpClient: return McpClient(proxy=proxy) diff --git a/tests/e2e/mcp/test_mcp_access_group_e2e.py b/tests/e2e/mcp/test_mcp_access_group_e2e.py index 1b53d1ca0b4..f72b75fd43d 100644 --- a/tests/e2e/mcp/test_mcp_access_group_e2e.py +++ b/tests/e2e/mcp/test_mcp_access_group_e2e.py @@ -29,6 +29,7 @@ class TestMcpAccessGroupToolSelection: ) -> None: group = f"e2e-mcp-grp-{unique_marker()}" server_id = register_datadog_mcp(client, resources, mcp_access_groups=[group]) + client.await_registered(server_id) granted = client.generate_key( user_id=f"e2e-mcp-ag-granted-{unique_marker()}", diff --git a/tests/e2e/mcp/test_mcp_datadog_e2e.py b/tests/e2e/mcp/test_mcp_datadog_e2e.py index 8a539b86bff..d093e307f99 100644 --- a/tests/e2e/mcp/test_mcp_datadog_e2e.py +++ b/tests/e2e/mcp/test_mcp_datadog_e2e.py @@ -60,6 +60,7 @@ class TestDatadogMcpRoundTrip: _assert_datadog_logger_active(client.proxy) server_id = register_datadog_mcp(client, resources) + client.await_registered(server_id) marker = f"{MARKER_PREFIX}{unique_marker()}" key = client.generate_key( @@ -78,22 +79,19 @@ class TestDatadogMcpRoundTrip: ) tool_name = client.await_tool(key, server_id, SEARCH_LOGS_TOOL) - - call = unwrap( - client.call_tool( - key, - server_id=server_id, - name=tool_name, - arguments={ - "query": marker, - "from": DD_SEARCH_FROM, - "to": "now", - "max_tokens": 5000, - "telemetry": { - "intent": "e2e assert seeded litellm completion log is searchable via MCP" - }, + call = client.await_call_tool( + key, + server_id=server_id, + name=tool_name, + arguments={ + "query": marker, + "from": DD_SEARCH_FROM, + "to": "now", + "max_tokens": 5000, + "telemetry": { + "intent": "e2e assert seeded litellm completion log is searchable via MCP" }, - ) + }, ) assert call.is_error is not True, f"search_datadog_logs errored: {call}" body = call.all_text diff --git a/tests/e2e/mcp/test_mcp_key_access_e2e.py b/tests/e2e/mcp/test_mcp_key_access_e2e.py index 35c864c07d8..678424e36d1 100644 --- a/tests/e2e/mcp/test_mcp_key_access_e2e.py +++ b/tests/e2e/mcp/test_mcp_key_access_e2e.py @@ -16,7 +16,7 @@ import pytest from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp from e2e_config import DD_SEARCH_FROM, unique_marker -from e2e_http import UnknownApiError, unwrap +from e2e_http import unwrap from lifecycle import ResourceManager from mcp_client import McpClient @@ -72,13 +72,12 @@ class TestMcpKeyWithoutAccessIsDenied: "max_tokens": 1000, "telemetry": {"intent": "e2e control call proving granted key can invoke Datadog MCP"}, } - permitted_call = unwrap( - client.call_tool(permitted_key, server_id=server_id, name=tool_name, arguments=search_args) + permitted_call = client.await_call_tool( + permitted_key, server_id=server_id, name=tool_name, arguments=search_args ) assert permitted_call.is_error is not True, f"granted key's tool call errored: {permitted_call}" - match client.call_tool(denied_key, server_id=server_id, name=tool_name, arguments=search_args): - case UnknownApiError(status_code=403, body=body): - assert "access_denied" in body, f"403 was not an MCP access denial: {body}" - case other: - pytest.fail(f"ungranted key's tool call was not refused with 403 access_denied: {other}") + denied = client.await_call_tool_denied( + denied_key, server_id=server_id, name=tool_name, arguments=search_args + ) + assert "access_denied" in denied.body, f"403 was not an MCP access denial: {denied.body}" From 44e091aedb3f7877e1b058c0480963621d4d6e0f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 09:16:18 +0000 Subject: [PATCH 26/54] chore(typing): clear basedpyright Any errors in proxy management endpoints Convert pydantic table-model construction from Cls(**row.model_dump()) kwargs-unpacking to Cls.model_validate(...) across the management endpoint hotspot files (team, key, internal user, scim, model management, spend tracking, auth checks, proxy_server). Unpacking an untyped dict reports one Any-typed argument per matched model field, so each converted site clears 10-35 diagnostics while running the exact same pydantic validation. Conversions were limited to models verified to use pydantic's default __init__; UserAPIKeyAuth and LiteLLM_VerificationTokenView keep their custom kwargs-rewriting __init__ and are untouched. Two locally-verified helper params move from Any to object. Whole-tree basedpyright, measured against the branch point in the same environment: reportAny 24,431 -> 22,741 (-1,690), reportArgumentType 2,189 -> 2,136 (-53), reportUnknownArgumentType 34,370 -> 34,067 (-303), reportExplicitAny 7,285 -> 7,283 (-2); total 154,882 -> 152,834 (-2,048) with no rule increasing anywhere and no per-file increases. No casts, no suppressions, no behavior changes. Budgets ratcheted: basedpyright -2,048 across 4 rules, ruff ANN401 -2. --- basedpyright-code-budget.json | 8 +- litellm/proxy/auth/auth_checks.py | 34 ++++----- .../internal_user_endpoints.py | 24 +++--- .../key_management_endpoints.py | 36 +++++---- .../model_management_endpoints.py | 4 +- .../management_endpoints/scim/scim_v2.py | 8 +- .../management_endpoints/team_endpoints.py | 76 ++++++++++--------- litellm/proxy/proxy_server.py | 12 ++- .../spend_management_endpoints.py | 4 +- ruff-strict-budget.json | 2 +- tests/test_litellm/proxy/test_proxy_server.py | 24 +++--- 11 files changed, 125 insertions(+), 107 deletions(-) diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 28602fc235f..db3c2502e94 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,9 +1,9 @@ { "reportAny": { - "limit": 34906 + "limit": 33216 }, "reportArgumentType": { - "limit": 2701 + "limit": 2648 }, "reportAssignmentType": { "limit": 330 @@ -24,7 +24,7 @@ "limit": 42 }, "reportExplicitAny": { - "limit": 10230 + "limit": 10228 }, "reportFunctionMemberAccess": { "limit": 11 @@ -99,7 +99,7 @@ "limit": 0 }, "reportUnknownArgumentType": { - "limit": 45870 + "limit": 45567 }, "reportUnknownLambdaType": { "limit": 113 diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 07dfdc4fb43..d02fc02d4bf 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -953,7 +953,7 @@ async def get_default_end_user_budget( ) return None - _budget_obj = LiteLLM_BudgetTable(**budget_record.dict()) + _budget_obj = LiteLLM_BudgetTable.model_validate(budget_record.dict()) # Cache the budget for 60 seconds await user_api_key_cache.async_set_cache( key=cache_key, @@ -999,7 +999,7 @@ async def get_team_member_default_budget( if isinstance(cached_budget, LiteLLM_BudgetTable): return cached_budget if isinstance(cached_budget, dict): - return LiteLLM_BudgetTable(**cached_budget) + return LiteLLM_BudgetTable.model_validate(cached_budget) try: budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id}) @@ -1014,7 +1014,7 @@ async def get_team_member_default_budget( ttl=get_management_object_ttl(user_api_key_cache), ) - return LiteLLM_BudgetTable(**budget_record.dict()) + return LiteLLM_BudgetTable.model_validate(budget_record.dict()) except Exception: verbose_proxy_logger.exception(f"Error fetching team-default member budget {budget_id}") @@ -1168,7 +1168,7 @@ async def get_end_user_object( raise Exception # Convert to LiteLLM_EndUserTable object - _response = LiteLLM_EndUserTable(**response.dict()) + _response = LiteLLM_EndUserTable.model_validate(response.dict()) # Apply default budget if needed _response = await _apply_default_budget_to_end_user( @@ -1360,7 +1360,7 @@ async def get_tag_objects_batch( for db_tag in db_tags: tag_name = db_tag.tag_name cache_key = f"tag:{tag_name}" - _tag_obj = LiteLLM_TagTable(**db_tag.dict()) + _tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict()) await user_api_key_cache.async_set_cache( key=cache_key, value=_tag_obj, @@ -1453,7 +1453,7 @@ async def get_team_membership( if response is None: return None - _response = LiteLLM_TeamMembership(**response.dict()) + _response = LiteLLM_TeamMembership.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=_key, value=_response, @@ -1719,13 +1719,13 @@ async def get_user_object( if response.organization_memberships is not None and len(response.organization_memberships) > 0: # dump each organization membership to type LiteLLM_OrganizationMembershipTable _dumped_memberships = [ - LiteLLM_OrganizationMembershipTable(**membership.model_dump()) + LiteLLM_OrganizationMembershipTable.model_validate(membership.model_dump()) for membership in response.organization_memberships if membership is not None ] response.organization_memberships = _dumped_memberships - _response = LiteLLM_UserTable(**dict(response)) + _response = LiteLLM_UserTable.model_validate(dict(response)) response_dict = _response.model_dump() # save the user object to cache @@ -1862,7 +1862,7 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_ http_request=mock_request, user_api_key_dict=system_admin_user, ) - response = LiteLLM_TeamTable(**created_team_dict) + response = LiteLLM_TeamTable.model_validate(created_team_dict) return response @@ -1894,7 +1894,7 @@ async def _get_team_object_from_user_api_key_cache( if response is None: raise Exception - _response = LiteLLM_TeamTableCachedObj(**response.dict()) + _response = LiteLLM_TeamTableCachedObj.model_validate(response.dict()) # Load object_permission if object_permission_id exists but object_permission is not loaded if _response.object_permission_id and not _response.object_permission: @@ -2085,7 +2085,7 @@ async def get_access_object( detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."}, ) - _response = LiteLLM_AccessGroupTable(**response.dict()) + _response = LiteLLM_AccessGroupTable.model_validate(response.dict()) # Save to cache await _cache_access_object( @@ -2170,7 +2170,7 @@ async def get_team_object_by_alias( ) team = teams[0] - team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump()) + team_obj = LiteLLM_TeamTableCachedObj.model_validate(team.model_dump()) # Load object_permission if object_permission_id exists but object_permission is not loaded if team_obj.object_permission_id and not team_obj.object_permission: @@ -2272,7 +2272,7 @@ async def get_org_object_by_alias( ) org = orgs[0] - org_obj = LiteLLM_OrganizationTable(**org.model_dump()) + org_obj = LiteLLM_OrganizationTable.model_validate(org.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( @@ -2605,7 +2605,7 @@ async def get_object_permission( if response is None: return None - _perm_obj = LiteLLM_ObjectPermissionTable(**response.dict()) + _perm_obj = LiteLLM_ObjectPermissionTable.model_validate(response.dict()) await user_api_key_cache.async_set_cache( key=key, value=_perm_obj, @@ -2665,7 +2665,7 @@ async def get_managed_vector_store_rows_by_uuids( row_dict = dict(row) if hasattr(row, "__dict__") else {} if not row_dict: continue - cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict) + cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict) key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id) await user_api_key_cache.async_set_cache( key=key, @@ -2746,7 +2746,7 @@ async def get_org_object( f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call." ) - _org_obj = LiteLLM_OrganizationTable(**response.model_dump()) + _org_obj = LiteLLM_OrganizationTable.model_validate(response.model_dump()) # Cache the result await user_api_key_cache.async_set_cache( key=cache_key, @@ -4221,7 +4221,7 @@ async def get_project_object( if project_row is None: return None - project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump()) + project_obj = LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump()) # Cache with TTL following _cache_management_object pattern project_obj.last_refreshed_at = time.time() diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index a2c16e88839..1bd0a19bfb3 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -513,7 +513,7 @@ async def new_user( response_dict["key"] = response.get("token", "") - new_user_response = NewUserResponse(**response_dict) + new_user_response = NewUserResponse.model_validate(response_dict) ######################################################### ########## USER CREATED HOOK ################ @@ -879,7 +879,7 @@ async def _check_user_info_v2_access( # Get all teams the caller belongs to teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": caller_user.teams}}) for team in teams: - team_obj = LiteLLM_TeamTable(**team.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): # Check if target user is in this team if team.team_id in (target_user.teams or []): @@ -1013,11 +1013,11 @@ async def _get_user_info_for_proxy_admin(user_api_key_dict: UserAPIKeyAuth): for key in _keys_in_db: if key.get("models") is None: key["models"] = [] - keys_in_db.append(LiteLLM_VerificationToken(**key)) + keys_in_db.append(LiteLLM_VerificationToken.model_validate(key)) # cast all teams to LiteLLM_TeamTable _teams_in_db: list = results[0]["teams"] or [] - _teams_in_db = [LiteLLM_TeamTable(**team) for team in _teams_in_db] + _teams_in_db = [LiteLLM_TeamTable.model_validate(team) for team in _teams_in_db] _teams_in_db.sort(key=lambda x: getattr(x, "team_alias", "") or "") returned_keys = _process_keys_for_user_info(keys=keys_in_db, all_teams=_teams_in_db) @@ -1146,7 +1146,7 @@ async def _schedule_user_update_audit_log( try: updated_user_row = await UserRepository(prisma_client).table.find_first(where={"user_id": response["user_id"]}) if updated_user_row: - user_row_typed = LiteLLM_UserTable(**updated_user_row.model_dump(exclude_none=True)) + user_row_typed = LiteLLM_UserTable.model_validate(updated_user_row.model_dump(exclude_none=True)) asyncio.create_task( UserManagementEventHooks.create_internal_user_audit_log( user_id=user_row_typed.user_id, @@ -1172,7 +1172,7 @@ def _check_user_update_authz( raise HTTPException(status_code=403, detail="Only proxy admins can modify user roles.") if existing_user_row is not None: - typed_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + typed_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) if not can_user_call_user_update(user_api_key_dict=user_api_key_dict, user_info=typed_row): raise HTTPException( status_code=403, @@ -1248,7 +1248,7 @@ async def _update_single_user_helper( _check_user_update_authz(user_request, user_api_key_dict, existing_user_row) if existing_user_row is not None: - existing_user_row = LiteLLM_UserTable(**existing_user_row.model_dump(exclude_none=True)) + existing_user_row = LiteLLM_UserTable.model_validate(existing_user_row.model_dump(exclude_none=True)) # Prevent budget self-escalation (GHSA-wvg4-6222-3q4r): non-admin callers # must not be able to raise their own budget/spend fields. @@ -1998,7 +1998,11 @@ async def get_users( for user in users: user_dump = user.model_dump() user_dump["metadata"] = _redact_scim_enterprise_metadata(user_dump.get("metadata")) - user_list.append(LiteLLM_UserTableWithKeyCount(**user_dump, key_count=user_key_counts.get(user.user_id, 0))) + user_list.append( + LiteLLM_UserTableWithKeyCount.model_validate( + {**user_dump, "key_count": user_key_counts.get(user.user_id, 0)} + ) + ) else: user_list = [] @@ -2157,7 +2161,7 @@ async def delete_user( teams_to_update = [] for team in fetch_all_teams: is_member_in_team, new_team_members = _cleanup_members_with_roles( - existing_team_row=LiteLLM_TeamTable(**team.model_dump()), + existing_team_row=LiteLLM_TeamTable.model_validate(team.model_dump()), data=TeamMemberDeleteRequest( team_id=team.team_id, user_id=user_row.user_id, @@ -2438,7 +2442,7 @@ async def ui_view_users( if not users: return [] - return [LiteLLM_UserTableFiltered(**user.model_dump()) for user in users] + return [LiteLLM_UserTableFiltered.model_validate(user.model_dump()) for user in users] except HTTPException: raise diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index ac6a2a4a7db..e7ee5ffa849 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1078,7 +1078,7 @@ async def _common_key_generation_helper( response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response - response = GenerateKeyResponse(**response) + response = GenerateKeyResponse.model_validate(response) response.token = response.token_id # remap token to use the hash, and leave the key in the `key` field [TODO]: clean up generate_key_helper_fn to do this @@ -3047,10 +3047,12 @@ async def bulk_update_team_keys( ) # team_id from validated scope, never user payload — drives _check_team_key_limits. - update_key_request = UpdateKeyRequest( - key=token, - team_id=data.team_id, - **update_field_dict, + update_key_request = UpdateKeyRequest.model_validate( + { + "key": token, + "team_id": data.team_id, + **update_field_dict, + } ) updated_key_info = await _process_single_key_update( update_key_request=update_key_request, @@ -4048,12 +4050,14 @@ def _transform_verification_tokens_to_deleted_records( records = [] for key in keys: key_payload = key.model_dump() - deleted_record = LiteLLM_DeletedVerificationToken( - **key_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedVerificationToken.model_validate( + { + **key_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -4535,7 +4539,7 @@ async def _execute_virtual_key_regeneration( proxy_logging_obj=proxy_logging_obj, ) - response = GenerateKeyResponse(**updated_token_dict) + response = GenerateKeyResponse.model_validate(updated_token_dict) asyncio.create_task( KeyManagementEventHooks.async_key_rotated_hook( data=data, @@ -4853,7 +4857,7 @@ async def _check_proxy_or_team_admin_for_key( ) -def _validate_reset_spend_value(reset_to: Any, key_in_db: LiteLLM_VerificationToken) -> float: +def _validate_reset_spend_value(reset_to: object, key_in_db: LiteLLM_VerificationToken) -> float: if not isinstance(reset_to, (int, float)): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -5029,7 +5033,7 @@ async def validate_key_list_check( code=status.HTTP_403_FORBIDDEN, ) - complete_user_info = LiteLLM_UserTable(**complete_user_info_db_obj.model_dump()) + complete_user_info = LiteLLM_UserTable.model_validate(complete_user_info_db_obj.model_dump()) # internal user can only see their own keys if user_id: @@ -5102,7 +5106,7 @@ async def _fetch_user_team_objects( if teams is None: return [] - return [LiteLLM_TeamTable(**team.model_dump()) for team in teams] + return [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in teams] def _get_admin_team_ids_from_objects( @@ -5851,7 +5855,7 @@ async def _list_key_helper( if return_full_object is True or (expand and "user" in expand): if use_deleted_table: # Use deleted key type to preserve deleted_at, deleted_by, etc. - key_list.append(LiteLLM_DeletedVerificationToken(**key_dict)) + key_list.append(LiteLLM_DeletedVerificationToken.model_validate(key_dict)) else: key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object else: diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 1c0e7211493..b6422d7f5ae 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -1051,7 +1051,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, @@ -1089,7 +1089,7 @@ class ModelManagementAuthChecks: status_code=400, detail={"error": "Team id={} does not exist in db".format(model_params.model_info.team_id)}, ) - team_obj = LiteLLM_TeamTable(**team_obj_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_obj_row.model_dump()) return ModelManagementAuthChecks.can_user_make_team_model_call( team_id=model_params.model_info.team_id, diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 90eae5bbb21..582d34dcec8 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -1317,7 +1317,7 @@ async def delete_user( where={"team_id": team.team_id}, data={"members": new_members} ) - team_row = LiteLLM_TeamTable(**team.model_dump()) + team_row = LiteLLM_TeamTable.model_validate(team.model_dump()) if any(member.user_id == user_id for member in team_row.members_with_roles or []): await team_member_delete( data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id), @@ -2145,7 +2145,9 @@ async def patch_group( refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) refreshed_current = ( - set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump()))) + set( + await _get_team_member_user_ids_from_team(LiteLLM_TeamTable.model_validate(refreshed_team.model_dump())) + ) if refreshed_team else snapshot_members ) @@ -2173,7 +2175,7 @@ async def patch_group( # Convert to SCIM format and return scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( - LiteLLM_TeamTable(**updated_team.model_dump()) + LiteLLM_TeamTable.model_validate(updated_team.model_dump()) ) return scim_group diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 59b0cbc4ae7..c35c17aa359 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -141,7 +141,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import ( router = APIRouter() -def _sanitize_for_log(value: Any) -> str: +def _sanitize_for_log(value: object) -> str: """Strip CR/LF from user-controlled values to prevent log injection.""" try: text = str(value) @@ -171,7 +171,7 @@ async def _refresh_cached_team( """ await _cache_team_object( team_id=team_row.team_id, - team_table=LiteLLM_TeamTableCachedObj(**team_row.model_dump()), + team_table=LiteLLM_TeamTableCachedObj.model_validate(team_row.model_dump()), user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) @@ -510,7 +510,7 @@ async def get_all_team_memberships( returned_tm: List[LiteLLM_TeamMembership] = [] for tm in team_memberships: - returned_tm.append(LiteLLM_TeamMembership(**tm.model_dump())) + returned_tm.append(LiteLLM_TeamMembership.model_validate(tm.model_dump())) return returned_tm @@ -772,7 +772,7 @@ async def _check_org_team_limits( # Convert teams to LiteLLM_TeamTable objects team_objs: List[LiteLLM_TeamTable] = [] for team in teams: - team_objs.append(LiteLLM_TeamTable(**team.model_dump())) + team_objs.append(LiteLLM_TeamTable.model_validate(team.model_dump())) check_org_team_model_specific_limits( teams=team_objs, @@ -1467,9 +1467,9 @@ async def fetch_and_validate_organization( ) is_proxy_admin = user_api_key_dict is not None and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN - organization = LiteLLM_OrganizationTableWithMembers(**organization_row.model_dump()) + organization = LiteLLM_OrganizationTableWithMembers.model_validate(organization_row.model_dump()) validate_team_org_change( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, llm_router=llm_router, is_proxy_admin=is_proxy_admin, @@ -1477,7 +1477,7 @@ async def fetch_and_validate_organization( if is_proxy_admin: await _auto_add_team_members_to_organization( - team=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), organization=organization, prisma_client=prisma_client, ) @@ -1714,7 +1714,7 @@ async def update_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team_row.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -2013,7 +2013,7 @@ async def patch_team( existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"]) - update_request = UpdateTeamRequest(team_id=team_id, **patch_fields) + update_request = UpdateTeamRequest.model_validate({"team_id": team_id, **patch_fields}) result = await update_team( data=update_request, @@ -2591,7 +2591,7 @@ async def team_member_add( detail={"error": f"Team not found for team_id={getattr(data, 'team_id', None)}"}, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) team_member_add_duplication_check( data=data, @@ -2636,10 +2636,12 @@ async def team_member_add( _emit_team_members_metric(complete_team_data) - return TeamAddMemberResponse( - **updated_team.model_dump(), - updated_users=updated_users, - updated_team_memberships=updated_team_memberships, + return TeamAddMemberResponse.model_validate( + { + **updated_team.model_dump(), + "updated_users": updated_users, + "updated_team_memberships": updated_team_memberships, + } ) @@ -2711,7 +2713,7 @@ async def team_member_delete( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -2915,7 +2917,7 @@ async def team_member_update( status_code=400, detail={"error": "Team id={} does not exist in db".format(data.team_id)}, ) - existing_team_row = LiteLLM_TeamTable(**_existing_team_row.model_dump()) + existing_team_row = LiteLLM_TeamTable.model_validate(_existing_team_row.model_dump()) ## CHECK IF USER IS PROXY ADMIN OR TEAM ADMIN OR ORG ADMIN @@ -3261,7 +3263,7 @@ async def delete_team( status_code=404, detail={"error": f"Team not found, passed team_id={team_id}"}, ) - team_row_pydantic = LiteLLM_TeamTable(**team_row_base.model_dump()) + team_row_pydantic = LiteLLM_TeamTable.model_validate(team_row_base.model_dump()) # Verify caller has access to manage this team await _verify_team_access( @@ -3385,12 +3387,14 @@ def _transform_teams_to_deleted_records( records = [] for team in teams: team_payload = team.model_dump() - deleted_record = LiteLLM_DeletedTeamTable( - **team_payload, - deleted_at=deleted_at, - deleted_by=user_api_key_dict.user_id, - deleted_by_api_key=user_api_key_dict.api_key, - litellm_changed_by=litellm_changed_by, + deleted_record = LiteLLM_DeletedTeamTable.model_validate( + { + **team_payload, + "deleted_at": deleted_at, + "deleted_by": user_api_key_dict.user_id, + "deleted_by_api_key": user_api_key_dict.api_key, + "litellm_changed_by": litellm_changed_by, + } ) record = deleted_record.model_dump() @@ -3580,7 +3584,7 @@ async def team_info( ) await validate_membership( user_api_key_dict=user_api_key_dict, - team_table=LiteLLM_TeamTable(**team_info.model_dump()), + team_table=LiteLLM_TeamTable.model_validate(team_info.model_dump()), ) ## GET ALL KEYS ## @@ -3615,9 +3619,9 @@ async def team_info( returned_tm = await get_all_team_memberships(prisma_client, [team_id], user_id=None) if isinstance(team_info, dict): - _team_info = TeamInfoResponseObjectTeamTable(**team_info) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info) elif isinstance(team_info, BaseModel): - _team_info = TeamInfoResponseObjectTeamTable(**team_info.model_dump()) + _team_info = TeamInfoResponseObjectTeamTable.model_validate(team_info.model_dump()) else: _team_info = TeamInfoResponseObjectTeamTable() @@ -3823,7 +3827,7 @@ async def block_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3872,7 +3876,7 @@ async def unblock_team( # Verify caller has access to manage this team await _verify_team_access( - team_obj=LiteLLM_TeamTable(**existing_team.model_dump()), + team_obj=LiteLLM_TeamTable.model_validate(existing_team.model_dump()), user_api_key_dict=user_api_key_dict, ) @@ -3916,13 +3920,13 @@ async def list_available_teams( status_code=404, detail={"error": "User not found"}, ) - user_info_correct_type = LiteLLM_UserTable(**user_info.model_dump()) + user_info_correct_type = LiteLLM_UserTable.model_validate(user_info.model_dump()) available_teams = [team for team in available_teams if team not in user_info_correct_type.teams] available_teams_db = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": available_teams}}) - available_teams_correct_type = [LiteLLM_TeamTable(**team.model_dump()) for team in available_teams_db] + available_teams_correct_type = [LiteLLM_TeamTable.model_validate(team.model_dump()) for team in available_teams_db] return available_teams_correct_type @@ -4090,7 +4094,7 @@ def _convert_teams_to_response_models( team_dict = team.dict() if use_deleted_table: - team_list.append(LiteLLM_DeletedTeamTable(**team_dict)) + team_list.append(LiteLLM_DeletedTeamTable.model_validate(team_dict)) else: members_with_roles = team_dict.get("members_with_roles") if not isinstance(members_with_roles, list): @@ -4705,7 +4709,7 @@ async def team_model_add( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can add models if ( @@ -4805,7 +4809,7 @@ async def team_model_delete( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) # Authorization check - only proxy admin, team admin, or org admin can remove models if ( @@ -4873,7 +4877,7 @@ async def team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Admin Viewer follows the read-parity rule: see team permissions like # a Proxy Admin would. Team / org admins keep their existing scope. @@ -4940,7 +4944,7 @@ async def update_team_member_permissions( check_db_only=True, ) - complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + complete_team_data = LiteLLM_TeamTable.model_validate(existing_team_row.model_dump()) # Available-team self-join must NOT grant write access to team-wide # permission policies; only proxy/team/org admins can update them. @@ -5201,7 +5205,7 @@ async def get_team_daily_activity( if not _user_has_admin_view(user_api_key_dict) and team_ids_list and team_aliases: has_full_team_view = True for team_alias in team_aliases: - team_obj = LiteLLM_TeamTable(**team_alias.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_alias.model_dump()) is_admin = _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) has_perm = _team_member_has_permission( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4486cd7de59..70484eb1e4e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11214,11 +11214,15 @@ async def get_all_team_models( if user_teams == "*": team_db_objects = await TeamRepository(prisma_client).table.find_many() - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] else: team_db_objects = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_teams}}) - team_db_objects_typed = [LiteLLM_TeamTable(**team_db_object.model_dump()) for team_db_object in team_db_objects] + team_db_objects_typed = [ + LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) for team_db_object in team_db_objects + ] team_models = _add_team_models_to_all_models( team_db_objects_typed=team_db_objects_typed, @@ -11292,7 +11296,7 @@ async def _populate_team_access_on_models( where={"user_id": user_api_key_dict.user_id} ) if user_db_object is not None: - user_object = LiteLLM_UserTable(**user_db_object.model_dump()) + user_object = LiteLLM_UserTable.model_validate(user_db_object.model_dump()) user_teams = user_object.teams or [] direct_access_models = get_direct_access_models( user_db_object=user_object, @@ -11827,7 +11831,7 @@ async def _load_team_object_for_model_filter(team_id: str, prisma_client: Prisma if team_db_object is None: verbose_proxy_logger.warning(f"Team {team_id} not found in database") return None - return LiteLLM_TeamTable(**team_db_object.model_dump()) + return LiteLLM_TeamTable.model_validate(team_db_object.model_dump()) except Exception as e: verbose_proxy_logger.exception(f"Error fetching team {team_id}: {str(e)}") return None diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 0c525ee9466..9aae9ca2875 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3548,7 +3548,7 @@ async def _can_team_member_view_log( team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if team_row is None: return False - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): return True return _team_member_has_permission( @@ -3640,7 +3640,7 @@ async def _get_permitted_team_ids_for_spend_logs( permitted: List[str] = [] for team_row in team_rows: - team_obj = LiteLLM_TeamTable(**team_row.model_dump()) + team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): permitted.append(team_obj.team_id) elif _team_member_has_permission( diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index addee5fc68a..f3b4fce97d3 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -24,7 +24,7 @@ "limit": 130 }, "ANN401": { - "limit": 2015 + "limit": 2013 }, "ASYNC230": { "limit": 14 diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 536d24d4b4e..62f5ced7a39 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -1571,14 +1571,14 @@ async def test_get_all_team_models(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: # Configure the mock class to return proper instances - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams="*", @@ -1607,7 +1607,7 @@ async def test_get_all_team_models(): mock_litellm_teamtable.find_many.return_value = [mock_team1] with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -1658,7 +1658,7 @@ async def test_get_all_team_models(): mock_router.get_model_list.side_effect = mock_get_model_list_with_none with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class: - mock_team_table_class.side_effect = mock_team_table_constructor + mock_team_table_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], @@ -2373,14 +2373,14 @@ async def test_get_all_team_models_with_access_groups(): with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class: - def mock_team_table_constructor(**kwargs): + def mock_team_table_constructor(data): mock_instance = MagicMock() - mock_instance.team_id = kwargs["team_id"] - mock_instance.models = kwargs["models"] - mock_instance.access_group_ids = kwargs.get("access_group_ids") + mock_instance.team_id = data["team_id"] + mock_instance.models = data["models"] + mock_instance.access_group_ids = data.get("access_group_ids") return mock_instance - mock_tt_class.side_effect = mock_team_table_constructor + mock_tt_class.model_validate.side_effect = mock_team_table_constructor result = await get_all_team_models( user_teams=["team1"], From 095364fd046bc746cbb58a7808d8b249126d867f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 10:04:58 +0000 Subject: [PATCH 27/54] test: cover the model_validate conversion sites flagged by codecov Add regression tests for the db-fetch paths whose converted construction lines were uncovered: the auth_checks getters (default end user budget, end user, team membership, access group, team by alias, org by alias, object permission, managed vector stores, project), get_all_team_memberships and list_available_teams in team_endpoints, and the proxy admin user info helper. Each test feeds a mocked prisma row through the real function and asserts the validated model's fields, so a bad model_validate conversion on any of these paths now fails a test instead of only dropping coverage. --- .../proxy/auth/test_auth_checks.py | 251 ++++++++++++++++++ .../test_internal_user_endpoints.py | 35 +++ .../test_team_endpoints.py | 65 +++++ 3 files changed, 351 insertions(+) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index ccb20976df9..a5bf5e280d8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -4762,3 +4762,254 @@ async def test_skip_user_budget_on_team_key_flag_restores_old_behavior(): request=MagicMock(spec=Request), ) assert result is True + + +@pytest.mark.asyncio +async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(monkeypatch): + from litellm.proxy.auth.auth_checks import get_default_end_user_budget + + monkeypatch.setattr(litellm, "max_end_user_budget_id", "budget-default-1") + + budget_row = MagicMock() + budget_row.dict = lambda: {"budget_id": "budget-default-1", "max_budget": 12.5, "tpm_limit": 100} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_default_end_user_budget( + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_BudgetTable) + assert result.max_budget == 12.5 + assert result.tpm_limit == 100 + mock_cache.async_set_cache.assert_awaited_once() + assert mock_cache.async_set_cache.call_args.kwargs["value"] is result + + +@pytest.mark.asyncio +async def test_get_end_user_object_db_fetch_returns_validated_end_user(): + from litellm.proxy.auth.auth_checks import get_end_user_object + + end_user_row = MagicMock() + end_user_row.dict = lambda: {"user_id": "eu-1", "blocked": False, "spend": 3.0} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_end_user_object( + end_user_id="eu-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_EndUserTable) + assert result.user_id == "eu-1" + assert result.blocked is False + assert result.spend == 3.0 + + +@pytest.mark.asyncio +async def test_get_team_membership_db_fetch_returns_validated_membership(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.auth.auth_checks import get_team_membership + + membership_row = MagicMock() + membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 1.5} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_membership( + user_id="u-1", + team_id="t-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamMembership) + assert result.user_id == "u-1" + assert result.team_id == "t-1" + assert result.spend == 1.5 + + +@pytest.mark.asyncio +async def test_get_access_object_db_fetch_returns_validated_access_group(): + from litellm.proxy._types import LiteLLM_AccessGroupTable + from litellm.proxy.auth.auth_checks import get_access_object + + access_row = MagicMock() + access_row.dict = lambda: { + "access_group_id": "ag-1", + "access_group_name": "group one", + "access_model_names": ["gpt-4"], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_access_object( + access_group_id="ag-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + proxy_logging_obj=None, + ) + + assert isinstance(result, LiteLLM_AccessGroupTable) + assert result.access_group_id == "ag-1" + assert result.access_model_names == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_team_object_by_alias_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_TeamTableCachedObj + from litellm.proxy.auth.auth_checks import get_team_object_by_alias + + team_row = MagicMock() + team_row.model_dump = lambda: {"team_id": "t-9", "team_alias": "alias-9", "models": ["gpt-4"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_team_object_by_alias( + team_alias="alias-9", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_TeamTableCachedObj) + assert result.team_id == "t-9" + assert result.team_alias == "alias-9" + assert result.models == ["gpt-4"] + + +@pytest.mark.asyncio +async def test_get_org_object_by_alias_db_fetch_returns_validated_org(): + from litellm.proxy._types import LiteLLM_OrganizationTable + from litellm.proxy.auth.auth_checks import get_org_object_by_alias + + org_row = MagicMock() + org_row.model_dump = lambda: { + "organization_id": "org-1", + "organization_alias": "org-alias", + "budget_id": "b-1", + "created_by": "admin", + "updated_by": "admin", + "models": [], + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_org_object_by_alias( + org_alias="org-alias", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_OrganizationTable) + assert result.organization_id == "org-1" + assert result.budget_id == "b-1" + + +@pytest.mark.asyncio +async def test_get_object_permission_db_fetch_returns_validated_permission(): + from litellm.proxy.auth.auth_checks import get_object_permission + + perm_row = MagicMock() + perm_row.dict = lambda: {"object_permission_id": "op-1", "vector_stores": ["vs-1"]} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_object_permission( + object_permission_id="op-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ObjectPermissionTable) + assert result.object_permission_id == "op-1" + assert result.vector_stores == ["vs-1"] + + +@pytest.mark.asyncio +async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows(): + from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable + from litellm.proxy.auth.auth_checks import get_managed_vector_store_rows_by_uuids + + vs_row = MagicMock() + vs_row.model_dump = lambda: {"vector_store_id": "vs-7", "custom_llm_provider": "openai"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row]) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_managed_vector_store_rows_by_uuids( + uuids=["vs-7"], + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_ManagedVectorStoresTable) + assert result[0].vector_store_id == "vs-7" + assert result[0].custom_llm_provider == "openai" + + +@pytest.mark.asyncio +async def test_get_project_object_db_fetch_returns_cached_obj(): + from litellm.proxy._types import LiteLLM_ProjectTableCachedObj + from litellm.proxy.auth.auth_checks import get_project_object + + project_row = MagicMock() + project_row.model_dump = lambda: {"project_id": "p-1", "project_alias": "proj", "team_id": "t-1"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row) + + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_set_cache = AsyncMock() + + result = await get_project_object( + project_id="p-1", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + assert isinstance(result, LiteLLM_ProjectTableCachedObj) + assert result.project_id == "p-1" + assert result.project_alias == "proj" diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 5cbc3e72d83..8de42ca89da 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -3667,3 +3667,38 @@ async def test_add_user_to_team_keeps_already_a_member_quiet(mocker, caplog): ) assert [r.getMessage() for r in caplog.records if r.levelno >= logging.ERROR] == [] + + +@pytest.mark.asyncio +async def test_get_user_info_for_proxy_admin_validates_keys_and_teams(): + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _get_user_info_for_proxy_admin, + ) + + raw_rows = [ + { + "teams": [ + {"team_id": "team-b", "team_alias": "beta"}, + {"team_id": "team-a", "team_alias": "alpha"}, + ], + "keys": [ + {"token": "hashed-token-1", "team_id": "team-a", "models": None, "spend": 1.0}, + ], + } + ] + + mock_prisma_client = MagicMock() + mock_prisma_client.db.query_raw = AsyncMock(return_value=raw_rows) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await _get_user_info_for_proxy_admin(user_api_key_dict=UserAPIKeyAuth(user_id=None)) + + assert all(isinstance(team, LiteLLM_TeamTable) for team in result.teams) + assert [team.team_alias for team in result.teams] == ["alpha", "beta"] + assert len(result.keys) == 1 + returned_key = result.keys[0] + assert returned_key["team_id"] == "team-a" + assert returned_key["models"] == [] diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 5202c8cbfc0..1e4d1759062 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -10223,3 +10223,68 @@ def test_patch_team_route_publishes_its_request_body_schema(): assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"} properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"] assert "tpm_limit" in properties and "metadata" in properties + + +@pytest.mark.asyncio +async def test_get_all_team_memberships_validates_rows(): + from litellm.proxy._types import LiteLLM_TeamMembership + from litellm.proxy.management_endpoints.team_endpoints import ( + get_all_team_memberships, + ) + + membership_row = MagicMock() + membership_row.model_dump = lambda: { + "user_id": "member-1", + "team_id": "team-1", + "spend": 2.5, + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership_row]) + + result = await get_all_team_memberships(mock_prisma_client, ["team-1"], user_id="member-1") + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamMembership) + assert result[0].user_id == "member-1" + assert result[0].team_id == "team-1" + assert result[0].spend == 2.5 + find_many_kwargs = mock_prisma_client.db.litellm_teammembership.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-1"]}, "user_id": {"in": ["member-1"]}} + + +@pytest.mark.asyncio +async def test_list_available_teams_filters_joined_and_validates_rows(monkeypatch): + from fastapi import Request + + import litellm + from litellm.proxy.management_endpoints.team_endpoints import list_available_teams + + monkeypatch.setattr( + litellm, + "default_internal_user_params", + {"available_teams": ["team-open", "team-joined"]}, + ) + + user_row = MagicMock() + user_row.model_dump = lambda: {"user_id": "u-1", "teams": ["team-joined"]} + + open_team_row = MagicMock() + open_team_row.model_dump = lambda: {"team_id": "team-open", "team_alias": "open team"} + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[open_team_row]) + + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): + result = await list_available_teams( + http_request=MagicMock(spec=Request), + user_api_key_dict=UserAPIKeyAuth(user_id="u-1"), + ) + + assert len(result) == 1 + assert isinstance(result[0], LiteLLM_TeamTable) + assert result[0].team_id == "team-open" + assert result[0].team_alias == "open team" + find_many_kwargs = mock_prisma_client.db.litellm_teamtable.find_many.call_args.kwargs + assert find_many_kwargs["where"] == {"team_id": {"in": ["team-open"]}} From 25ebe5600a392d8bc75921d6a22fb4ca61273b5e Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 09:38:17 -0700 Subject: [PATCH 28/54] fix(scim): stop provisioning nested group ids as internal users (#34997) * fix(scim): stop provisioning nested group ids as internal users POST/PUT/PATCH /scim/v2/Groups treated every member.value as a user id, so with the default scim_upsert_user=true an unknown id was auto-created as an internal user. Entra sends nested groups as members carrying "type": "Group", which meant every nested group produced a phantom internal user whose id and email were the group GUID, and those users counted toward licensed seats. Group members are now classified before they are used: members typed "Group" are skipped without a database hit, an id that names an existing team is skipped too (Okta sends untyped ids through filtered paths, so the type alone is not enough), and only ids that resolve to a user, or that resolve to nothing at all, keep today's behavior. The user lookup runs before the team lookup so a user whose id collides with a team id keeps syncing. The type was previously dropped at parse time on POST/PUT because SCIMMember had no such field, and on PATCH because the raw member dicts were reduced to bare ids; both paths now share one resolver and one parser that preserves it. Member removals no longer upsert: a remove of an id we do not know is an idempotent no-op rather than a reason to create a user and immediately drop it, and strict mode (scim_upsert_user=false) no longer rejects it. Removal of an id that is on the roster but has no user row still cleans up membership. Responses now state members are of type "User" instead of emitting a null, and the advertised Group schema documents the members.type sub-attribute. * fix(scim): harden group member classification after adversarial review Removals now bypass classification and drop exactly the ids they name, restoring cleanup of roster entries the old bug left behind. The team-id fallback only applies to untyped members, so an explicit User type always provisions even when the id collides with a team. Member types are normalized before matching; a type other than User or Group only skips when the id is not an existing user. Non-string type values are tolerated as absent on every verb instead of failing validation. Admitted member ids are deduped order-preserving, which also closes a pre-existing duplicate-row hazard on group creation. * fix(scim): only treat scim-managed teams as nested groups A PR reviewer flagged that an untyped SCIM member whose id collides with an admin-created team was silently skipped, suppressing that user's provisioning. SCIM group writes (POST, PUT, and every PATCH) now stamp the team with scim_managed metadata, and the typeless team-id skip only applies to teams carrying that marker or the scim_data blob older PUTs already wrote. Admin-created teams stay unmarked, so a colliding untyped member provisions the user in permissive mode and returns the standard unknown-user 400 in strict mode. Teams SCIM touched before this change adopt the marker on their next group write. --- .../scim/scim_transformations.py | 1 + .../management_endpoints/scim/scim_v2.py | 425 ++++++--- .../proxy/management_endpoints/scim_v2.py | 12 + .../scim/test_scim_transformations.py | 15 + .../scim/test_scim_v2_discovery.py | 13 + .../scim/test_scim_v2_endpoints.py | 884 +++++++++++++++++- 6 files changed, 1230 insertions(+), 120 deletions(-) diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py index 65651752944..d572255fd32 100644 --- a/litellm/proxy/management_endpoints/scim/scim_transformations.py +++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py @@ -175,6 +175,7 @@ class ScimTransformations: SCIMMember( value=ScimTransformations._get_scim_member_value(member), display=ScimTransformations._get_scim_member_display(member), + type="User", ) ) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 90eae5bbb21..5f0ad1a2983 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -5,7 +5,9 @@ This is an enterprise feature and requires a premium license. """ import re -from typing import Any, Dict, Iterable, List, Optional, Set, Tuple +from collections.abc import Mapping, Sequence +from itertools import chain +from typing import Any, Dict, Iterable, List, NamedTuple, Optional, Set, Tuple from fastapi import ( APIRouter, @@ -17,8 +19,8 @@ from fastapi import ( Request, Response, ) -from pydantic import BaseModel, ValidationError -from typing_extensions import TypedDict +from pydantic import BaseModel, TypeAdapter, ValidationError +from typing_extensions import TypedDict, assert_never import litellm from litellm._logging import verbose_proxy_logger @@ -50,7 +52,11 @@ from litellm.proxy.management_endpoints.team_endpoints import ( team_member_add, team_member_delete, ) -from litellm.proxy.utils import _premium_user_check, handle_exception_on_proxy +from litellm.proxy.utils import ( + PrismaClient, + _premium_user_check, + handle_exception_on_proxy, +) from litellm.repositories.table_repositories import ( InvitationLinkRepository, OrganizationMembershipRepository, @@ -143,7 +149,11 @@ class ScimUserData(TypedDict): class GroupMemberExtractionResult(BaseModel): - """Result of extracting and processing group members.""" + """Result of extracting and processing group members. + + ``all_member_ids`` is deduped order-preserving; ``existing_member_ids`` is not, + so a repeated resolved id appears once in the former and twice in the latter. + """ existing_member_ids: List[str] created_users: List[NewUserResponse] @@ -371,6 +381,216 @@ async def _recompute_scim_member_roles(prisma_client: Any, user_ids: Iterable[st ) +class _ResolvedUserMember(NamedTuple): + user_id: str + + +class _SkippedGroupMember(NamedTuple): + value: str + reason: Literal["nested_group", "non_user_type", "existing_team"] + + +class _UnknownMember(NamedTuple): + value: str + + +_ClassifiedGroupMember = Union[_ResolvedUserMember, _SkippedGroupMember, _UnknownMember] + + +class _PartitionedMembers(NamedTuple): + resolved_ids: tuple[str, ...] + skipped: tuple[_SkippedGroupMember, ...] + unknown_ids: tuple[str, ...] + + +def _member_value(member: SCIMMember) -> str: + """A member id is opaque to us but has to be there; an empty one is a client error.""" + if not member.value or not member.value.strip(): + raise HTTPException( + status_code=400, + detail={"error": "Invalid member: user ID cannot be empty."}, + ) + return member.value + + +def _normalized_member_type(member: SCIMMember) -> str | None: + """The canonical ``type`` a member declares, lowercased; blank or absent means none.""" + normalized = (member.type or "").strip().lower() + return normalized or None + + +_JSON_OBJECT_ADAPTER = TypeAdapter(Dict[str, object]) + + +def _json_object_fields(raw: object) -> Mapping[str, object] | None: + """A typed, read-only view of a JSON object, or None when it is not one.""" + try: + return _JSON_OBJECT_ADAPTER.validate_python(raw) + except ValidationError: + return None + + +def _team_metadata_has_scim_provenance(team_metadata: object) -> bool: + """Whether a group write from the identity provider left its mark on this team. + + ``SCIM_TEAM_DATA_METADATA_KEY`` counts because PUT has been writing it since + long before the explicit marker, so a team the identity provider already + syncs is recognized without waiting to be written again. + """ + fields = _json_object_fields(team_metadata) + if fields is None: + return False + return bool(fields.get(SCIM_MANAGED_TEAM_METADATA_KEY)) or fields.get(SCIM_TEAM_DATA_METADATA_KEY) is not None + + +async def _classify_group_member(member: SCIMMember, prisma_client: PrismaClient) -> _ClassifiedGroupMember: + """ + Decide what a single SCIM group member refers to. + + A LiteLLM team only holds users, so a member is dropped when it declares a type + other than ``User`` or when its id names an existing team. Both of those checks + are placed around the user lookup rather than before it, because the id of a + real user is the one thing that outranks them: + + - ``"type": "Group"`` (what Entra sends for a nested group) is dropped without + a lookup. This bug provisioned nested group GUIDs as users, so those rows + exist in the wild and would otherwise resolve as members all over again. + - any other unrecognized type is dropped only after the user lookup misses. + Clients do send non-canonical types on real members (RFC 7643 defines + ``direct`` for ``User.groups``), and dropping a live user over one would + revoke that user's team access on the next full sync. + - an id that names an existing team is dropped only when the member arrives + untyped, which is how Okta sends nested groups, and only when that team is + one the identity provider writes. An id the IdP called a User is a user + even if some team happens to share the id, and a team created here rather + than through SCIM is not evidence of anything about the member. + """ + value = _member_value(member) + member_type = _normalized_member_type(member) + + if member_type == "group": + return _SkippedGroupMember(value=value, reason="nested_group") + + user = await UserRepository(prisma_client).table.find_unique(where={"user_id": value}) + if user is not None: + return _ResolvedUserMember(user_id=value) + + if member_type is not None and member_type != "user": + return _SkippedGroupMember(value=value, reason="non_user_type") + + if member_type is None: + team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": value}) + if team is not None and _team_metadata_has_scim_provenance(team.metadata): + return _SkippedGroupMember(value=value, reason="existing_team") + + return _UnknownMember(value=value) + + +def _bucketed_member(entry: _ClassifiedGroupMember) -> _PartitionedMembers: + """The single-member partition one classified entry contributes.""" + match entry: + case _ResolvedUserMember(user_id=user_id): + return _PartitionedMembers(resolved_ids=(user_id,), skipped=(), unknown_ids=()) + case _SkippedGroupMember(): + return _PartitionedMembers(resolved_ids=(), skipped=(entry,), unknown_ids=()) + case _UnknownMember(value=value): + return _PartitionedMembers(resolved_ids=(), skipped=(), unknown_ids=(value,)) + case _: + assert_never(entry) + + +def _partition_classified_members(classified: Iterable[_ClassifiedGroupMember]) -> _PartitionedMembers: + """Split classified members into the buckets the resolver acts on, keeping request order.""" + bucketed = tuple(_bucketed_member(entry) for entry in classified) + return _PartitionedMembers( + resolved_ids=tuple(chain.from_iterable(bucket.resolved_ids for bucket in bucketed)), + skipped=tuple(chain.from_iterable(bucket.skipped for bucket in bucketed)), + unknown_ids=tuple(chain.from_iterable(bucket.unknown_ids for bucket in bucketed)), + ) + + +def _admitted_member_id(entry: _ClassifiedGroupMember, created_ids: frozenset[str]) -> str | None: + match entry: + case _ResolvedUserMember(user_id=user_id): + return user_id + case _UnknownMember(value=value): + return value if value in created_ids else None + case _SkippedGroupMember(): + return None + case _: + assert_never(entry) + + +def _admitted_member_ids(classified: Iterable[_ClassifiedGroupMember], created_ids: frozenset[str]) -> tuple[str, ...]: + """Member ids that survive resolution, in the order the request listed them. + + An id the request repeats is one member: the roster these ids are written to + holds one row per member, and a second creation attempt for the same id fails + against the real unique constraint even though the first one succeeded. + """ + return tuple( + dict.fromkeys( + member_id for entry in classified if (member_id := _admitted_member_id(entry, created_ids)) is not None + ) + ) + + +async def _resolve_group_member_ids( + members: Sequence[SCIMMember], + created_via: str, + prisma_client: PrismaClient, +) -> GroupMemberExtractionResult: + """ + Resolve SCIM group members to LiteLLM user ids, dropping members that are not users. + + Only the operations that put ids onto a roster resolve their members: an id + that resolves to nothing is created when litellm_settings.scim_upsert_user is + True (default) and rejected per SCIM 2.0 otherwise. Removals do not come + through here; dropping an id is idempotent, so it needs neither a lookup nor a + user to drop. + + Raises: + HTTPException: 400 when a member id is empty, or when scim_upsert_user is + False and a member id is neither an existing user, an existing team, nor a + member declared to be something other than a user. + """ + classified = tuple([await _classify_group_member(member, prisma_client) for member in members]) + partition = _partition_classified_members(classified) + + for skipped in partition.skipped: + verbose_proxy_logger.info( + "SCIM: ignoring non-user group member '%s' (%s); LiteLLM teams contain users only", + skipped.value, + skipped.reason, + ) + + if partition.unknown_ids and not await _get_scim_upsert_user_setting(): + raise HTTPException( + status_code=400, + detail={ + "error": f"User with ID '{partition.unknown_ids[0]}' does not exist. " + "Please create the user first via POST /Users before adding to group." + }, + ) + + creations = tuple( + [ + (user_id, await _create_user_if_not_exists(user_id=user_id, created_via=created_via)) + for user_id in partition.unknown_ids + ] + ) + created_users = tuple(created for _, created in creations if created is not None) + + return GroupMemberExtractionResult( + existing_member_ids=partition.resolved_ids, + created_users=created_users, + all_member_ids=_admitted_member_ids( + classified, + frozenset(user_id for user_id, created in creations if created is not None), + ), + ) + + async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult: """ Extract member IDs from SCIMGroup, validating that all users exist. @@ -386,56 +606,10 @@ async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionRe HTTPException: If scim_upsert_user is False and any member user does not exist (400 Bad Request) """ prisma_client = await _get_prisma_client_or_raise_exception() - existing_member_ids = [] - created_users = [] - all_member_ids = [] - - # Check the feature flag - scim_upsert_user = await _get_scim_upsert_user_setting() - - if group.members: - for member in group.members: - user_id = member.value - - # Validate user_id is not empty - if not user_id or not user_id.strip(): - raise HTTPException( - status_code=400, - detail={"error": "Invalid member: user ID cannot be empty."}, - ) - - # Check if user exists - user = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - - if user: - existing_member_ids.append(user_id) - all_member_ids.append(user_id) - else: - if scim_upsert_user: - # Create the user if they don't exist (backward compatible behavior) - created_user = await _create_user_if_not_exists( - user_id=user_id, created_via="scim_group_membership" - ) - if created_user: - created_users.append(created_user) - all_member_ids.append(user_id) - # If creation failed, user is skipped (logged in helper) - else: - # User doesn't exist - reject per SCIM 2.0 protocol - # This prevents security issues where users not assigned to app - # get provisioned via group membership - raise HTTPException( - status_code=400, - detail={ - "error": f"User with ID '{user_id}' does not exist. " - "Please create the user first via POST /Users before adding to group." - }, - ) - - return GroupMemberExtractionResult( - existing_member_ids=existing_member_ids, - created_users=created_users, - all_member_ids=all_member_ids, + return await _resolve_group_member_ids( + members=group.members or [], + created_via="scim_group_membership", + prisma_client=prisma_client, ) @@ -448,7 +622,7 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id}) if user: display_name = user.user_email or user.user_id - members.append(SCIMMember(value=user.user_id, display=display_name)) + members.append(SCIMMember(value=user.user_id, display=display_name, type="User")) return members @@ -863,6 +1037,14 @@ def _get_schemas() -> list: type="string", description="Member display name.", ), + SCIMSchemaAttribute( + name="type", + type="string", + description=( + 'The type of member; canonical values are "User" and "Group". ' + "Only members of type User are honored, LiteLLM teams contain users only." + ), + ), ], ), ], @@ -1336,21 +1518,42 @@ async def delete_user( raise handle_exception_on_proxy(e) -def _extract_group_values(value: Any) -> List[str]: +def _parse_member_entry(entry: object) -> SCIMMember | None: + """Parse one entry of a SCIM patch value, or None when it carries no id.""" + if isinstance(entry, str): + return SCIMMember(value=entry) + + fields = _json_object_fields(entry) + if fields is None: + return None + + entry_value = fields.get("value") + if not entry_value: + return None + + entry_display = fields.get("display") + entry_type = fields.get("type") + return SCIMMember( + value=str(entry_value), + display=str(entry_display) if entry_display is not None else None, + type=entry_type if isinstance(entry_type, str) else None, + ) + + +def _parse_member_entries(value: object) -> tuple[SCIMMember, ...]: + """Parse a SCIM patch value into members, keeping each entry's ``type``. + + PATCH bodies bypass SCIMGroup parsing (SCIMPatchOperation.value is untyped), + so member objects arrive as raw dicts and the ``type`` that marks a nested + group would otherwise be lost. + """ + entries: tuple[object, ...] = tuple(value) if isinstance(value, list) else (value,) + return tuple(member for member in (_parse_member_entry(entry) for entry in entries) if member is not None) + + +def _extract_group_values(value: object) -> List[str]: """Return group ids from a SCIM patch value.""" - group_values: List[str] = [] - if isinstance(value, list): - for v in value: - if isinstance(v, dict) and v.get("value"): - group_values.append(str(v.get("value"))) - elif isinstance(v, str): - group_values.append(v) - elif isinstance(value, dict): - if value.get("value"): - group_values.append(str(value.get("value"))) - elif isinstance(value, str): - group_values.append(value) - return group_values + return [member.value for member in _parse_member_entries(value)] def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]: @@ -1833,6 +2036,7 @@ async def create_group( team_id=team_id, team_alias=group.displayName, members_with_roles=members_with_roles, + metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True}, ), http_request=Request(scope={"type": "http", "path": "/scim/v2/Groups"}), user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), @@ -1875,7 +2079,11 @@ async def update_group( # Prepare update data existing_metadata = existing_team.metadata if existing_team.metadata else {} - updated_metadata = {**existing_metadata, "scim_data": group.model_dump()} + updated_metadata = { + **existing_metadata, + SCIM_TEAM_DATA_METADATA_KEY: group.model_dump(), + SCIM_MANAGED_TEAM_METADATA_KEY: True, + } update_data = { "team_alias": group.displayName, @@ -1968,12 +2176,17 @@ async def _process_group_patch_operations( is absolute: it declares the roster is exactly this set, so the caller must reconcile against it as a set-to-target rather than rebasing it onto a concurrently-mutated roster. + + A ``remove`` drops the ids it names without resolving them first. Removal is + idempotent and cannot put anything on a roster, while resolving would make it + conditional on what the id turns out to be and leave members we should never + have admitted - the phantom users this endpoint used to create for nested + groups - impossible to clean up. """ update_data: Dict[str, Any] = {} # Create a fresh copy of existing metadata to avoid Prisma issues - existing_metadata = existing_team.metadata or {} - metadata = dict(existing_metadata) if existing_metadata else {} + metadata = {**(existing_team.metadata or {}), SCIM_MANAGED_TEAM_METADATA_KEY: True} # Track member changes. members_with_roles is the source of truth for team # membership; the legacy `members` column is not populated by team creation @@ -2001,50 +2214,26 @@ async def _process_group_patch_operations( metadata["externalId"] = str(value) elif path.startswith("members"): # Handle member operations - member_values = _extract_group_values(value) - if not member_values and value is None: - member_values = _extract_ids_from_path_filter(op.path, "members") - # Check the feature flag - scim_upsert_user = await _get_scim_upsert_user_setting() - # Validate all users exist or create them based on feature flag - valid_members = [] - for member_id in member_values: - # Validate member_id is not empty - if not member_id or not member_id.strip(): - raise HTTPException( - status_code=400, - detail={"error": "Invalid member: user ID cannot be empty."}, - ) + patched_members = ( + _parse_member_entries(value) + if value is not None + else tuple( + SCIMMember(value=member_id) for member_id in _extract_ids_from_path_filter(op.path, "members") + ) + ) - user = await UserRepository(prisma_client).table.find_unique(where={"user_id": member_id}) - if user: - valid_members.append(member_id) - else: - if scim_upsert_user: - # Create the user if they don't exist (backward compatible behavior) - created_user = await _create_user_if_not_exists( - user_id=member_id, created_via="scim_group_patch" - ) - if created_user: - valid_members.append(member_id) - # If creation failed, user is skipped (logged in helper) - else: - # User doesn't exist - reject per SCIM 2.0 protocol - raise HTTPException( - status_code=400, - detail={ - "error": f"User with ID '{member_id}' does not exist. " - "Please create the user first via POST /Users before adding to group." - }, - ) - - if op_type == "replace": - final_members = set(valid_members) - elif op_type == "add": - final_members.update(valid_members) - elif op_type == "remove": - for member_id in valid_members: - final_members.discard(member_id) + if op_type == "remove": + final_members = final_members - {_member_value(member) for member in patched_members} + else: + member_result = await _resolve_group_member_ids( + members=patched_members, + created_via="scim_group_patch", + prisma_client=prisma_client, + ) + if op_type == "replace": + final_members = set(member_result.all_member_ids) + elif op_type == "add": + final_members = final_members | set(member_result.all_member_ids) else: # Handle other generic metadata if op_type == "remove": @@ -2052,9 +2241,7 @@ async def _process_group_patch_operations( else: metadata[path] = value - # Include metadata in update data if it exists - if metadata: - update_data["metadata"] = metadata + update_data["metadata"] = metadata member_replace_present = any( op.op == "replace" and (op.path or "").lower().startswith("members") for op in patch_ops.Operations diff --git a/litellm/types/proxy/management_endpoints/scim_v2.py b/litellm/types/proxy/management_endpoints/scim_v2.py index 8c434481975..e7f5c85e6c6 100644 --- a/litellm/types/proxy/management_endpoints/scim_v2.py +++ b/litellm/types/proxy/management_endpoints/scim_v2.py @@ -18,6 +18,9 @@ SCIM_ENTERPRISE_METADATA_KEY = "scim_enterprise" SCIM_ENTITLEMENTS_METADATA_KEY = "scim_entitlements" SCIM_ROLES_METADATA_KEY = "scim_roles" +SCIM_MANAGED_TEAM_METADATA_KEY = "scim_managed" +SCIM_TEAM_DATA_METADATA_KEY = "scim_data" + class LiteLLM_UserScimMetadata(BaseModel): """ @@ -131,6 +134,15 @@ class SCIMUser(SCIMResource): class SCIMMember(BaseModel): value: str # User ID display: Optional[str] = None # Username or email + type: str | None = None + + @field_validator("type", mode="before") + @classmethod + def normalize_type(cls, v: object) -> str | None: + """Anything that is not a string carries no canonical type, and rejecting the + request over it would be a regression: before this field existed the value was + parsed away silently.""" + return v if isinstance(v, str) else None class SCIMGroup(SCIMResource): diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py index 458c7c42eb6..6970e34f759 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py @@ -332,6 +332,21 @@ class TestScimTransformations: assert scim_group.members[1].value == "test2@example.com" assert scim_group.members[1].display == "test2@example.com" + @pytest.mark.asyncio + async def test_transform_team_marks_members_as_users( + self, mock_team, mock_prisma_client + ): + """A LiteLLM team only holds users, and stating the member type keeps the + response from emitting a null ``type`` now that SCIMMember carries one.""" + mock_client, _ = mock_prisma_client + + with patch("litellm.proxy.proxy_server.prisma_client", mock_client): + scim_group = await ScimTransformations.transform_litellm_team_to_scim_group( + mock_team + ) + + assert [member.type for member in scim_group.members] == ["User", "User"] + def test_get_scim_user_name(self, mock_user, mock_user_minimal): # User with email result = ScimTransformations._get_scim_user_name(mock_user) diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py index 94ca0dc11f5..6ced5264267 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_discovery.py @@ -108,6 +108,19 @@ class TestGetSchemas: assert "displayName" in attr_names assert "members" in attr_names + def test_group_schema_advertises_member_type(self): + """IdPs read the schema to learn we understand ``members.type``, which is how + a nested group announces itself.""" + schemas = _get_schemas() + group_schema = next( + s for s in schemas if s.id == "urn:ietf:params:scim:schemas:core:2.0:Group" + ) + members = next(a for a in group_schema.attributes if a.name == "members") + member_type = next(a for a in members.subAttributes or [] if a.name == "type") + assert member_type.type == "string" + assert member_type.multiValued is False + assert "Group" in (member_type.description or "") + def test_schema_meta_fields(self): schemas = _get_schemas() user_schema = next( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index 7bb74285ac6..e333bf1e3fe 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -20,8 +20,10 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _extract_group_member_ids, _extract_ids_from_path_filter, _handle_team_membership_changes, + _parse_member_entries, _process_group_patch_operations, _recompute_scim_member_roles, + _resolve_group_member_ids, create_group, create_user, delete_group, @@ -36,6 +38,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( ) from litellm.types.proxy.management_endpoints.scim_v2 import ( SCIM_ENTERPRISE_USER_SCHEMA, + SCIM_MANAGED_TEAM_METADATA_KEY, + SCIM_TEAM_DATA_METADATA_KEY, SCIMGroup, SCIMMember, SCIMPatchOp, @@ -1611,7 +1615,10 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): mock_prisma_client.db.litellm_usertable = mocker.MagicMock() # Mock team operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) + def mock_team_lookup(where): + return mock_existing_team if where["team_id"] == group_id else None + + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=mock_team_lookup) # Mock updated team response mock_updated_team = mocker.MagicMock() @@ -1775,6 +1782,7 @@ async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, mon mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user lookup - only existing-user exists initially def mock_user_lookup(where): @@ -1842,6 +1850,7 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user lookup - only existing-user exists def mock_user_lookup(where): @@ -1902,6 +1911,7 @@ async def test_process_group_patch_operations_with_flag_true_creates_users(mocke # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Mock user creation created_user = NewUserResponse(user_id="new-user-1", key="test-key-1") @@ -1956,6 +1966,7 @@ async def test_process_group_patch_operations_with_flag_false_rejects(mocker, mo # Mock user lookup - new-user-1 doesn't exist mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) # Execute the function - should raise HTTPException with pytest.raises(HTTPException) as exc_info: @@ -3519,3 +3530,874 @@ async def test_process_group_patch_replace_empty_value_does_not_use_path_filter( ) assert final_members == set() + + +def _member_resolution_prisma(mocker, *, users: set, teams: set, unmanaged_teams: frozenset = frozenset()): + """Prisma mock where only the given ids resolve to a user row / team row. + + ``teams`` are teams a SCIM group write created, so they carry provenance; + ``unmanaged_teams`` resolve too but look like a team an admin created here. + """ + + def team_row(team_id: str) -> LiteLLM_TeamTable | None: + if team_id in teams: + return LiteLLM_TeamTable(team_id=team_id, metadata={SCIM_MANAGED_TEAM_METADATA_KEY: True}) + if team_id in unmanaged_teams: + return LiteLLM_TeamTable(team_id=team_id, metadata={}) + return None + + prisma_client = mocker.MagicMock() + prisma_client.db = mocker.MagicMock() + prisma_client.db.litellm_usertable = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=lambda where: ( + LiteLLM_UserTable(user_id=where["user_id"]) if where["user_id"] in users else None + ) + ) + prisma_client.db.litellm_teamtable = mocker.MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=lambda where: team_row(where["team_id"])) + return prisma_client + + +@pytest.fixture +def scim_upsert_user_enabled(monkeypatch): + from litellm.proxy.proxy_server import proxy_config + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": True}} + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + +@pytest.fixture +def scim_upsert_user_disabled(monkeypatch): + from litellm.proxy.proxy_server import proxy_config + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": False}} + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + +@pytest.mark.asyncio +async def test_create_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """Entra sends nested groups as members with ``type: "Group"``. Treating that + GUID as a user id provisioned a phantom internal user per nested group.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", display="Real User", type="User"), + SCIMMember(value=nested_group_id, display="Nested Group", type="Group"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams=set())), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + create_user_mock.assert_not_called() + assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")] + + +@pytest.mark.asyncio +async def test_update_group_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """PUT /Groups must drop nested-group members too, so a full sync from the IdP + neither provisions nor enrolls the nested group's GUID.""" + group_id = "parent-group" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + existing_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Parent Group", + members=[], + members_with_roles=[], + metadata={}, + ) + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", display="Real User", type="User"), + SCIMMember(value=nested_group_id, display="Nested Group", type="Group"), + ], + ) + + prisma_client = _member_resolution_prisma(mocker, users={"real-user"}, teams={group_id}) + prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists", + AsyncMock(return_value=existing_team), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + patch_membership_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + + await update_group(group_id=group_id, group=scim_group) + + create_user_mock.assert_not_called() + enrolled = {call.kwargs["user_id"] for call in patch_membership_mock.call_args_list} + assert enrolled == {"real-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_ignores_nested_group_members(mocker, scim_upsert_user_enabled): + """PATCH bodies bypass SCIMGroup parsing, so ``type`` must be read off the raw + member dicts; otherwise a nested group is indistinguishable from a user id.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation( + op="add", + path="members", + value=[ + {"value": "real-user", "display": "Real User", "type": "User"}, + {"value": nested_group_id, "display": "Nested Group", "type": "Group"}, + ], + ) + ], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="incumbent", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"incumbent", "real-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_ignores_lowercase_group_type(mocker, scim_upsert_user_enabled): + """The ``type`` comparison is case-insensitive; IdPs are not consistent about it.""" + nested_group_id = "8f1e9d70-0000-4a0e-9a1e-nested" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="add", path="members", value=[{"value": nested_group_id, "type": "group"}]) + ], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == set() + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_skips_member_matching_existing_team(mocker, scim_upsert_user_enabled): + """Okta sends filtered paths and untyped ids, so a nested group arrives with no + ``type`` at all; an id that names an existing team is still not a user.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="incumbent", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"child-team", "parent-group"}), + ) + + create_user_mock.assert_not_called() + assert final_members == {"incumbent"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_prefers_user_over_team_for_colliding_id( + mocker, scim_upsert_user_enabled +): + """Nothing stops a user id from also being a team id, so the user lookup has to + win; ordering the team check first would silently stop syncing that user.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "dual-id"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"dual-id"}, teams={"dual-id"}), + ) + + assert final_members == {"dual-id"} + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_accepts_group_and_team_members(mocker, scim_upsert_user_disabled): + """Strict mode (scim_upsert_user=False) rejects unknown *users*; a nested group + is not a user, so it must be dropped rather than 400 the whole sync.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="real-user", type="User"), + SCIMMember(value="nested-group-guid", type="Group"), + SCIMMember(value="child-team"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users={"real-user"}, teams={"child-team"})), + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + create_user_mock.assert_not_called() + assert new_team_mock.call_args.kwargs["data"].members_with_roles == [Member(user_id="real-user", role="user")] + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_still_rejects_unknown_user(mocker, scim_upsert_user_disabled): + """The strict-mode 400 must name the unknown *user* and stay quiet about the + nested group sharing the request.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[ + SCIMMember(value="nested-group-guid", type="Group"), + SCIMMember(value="unknown-user"), + ], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())), + ) + + with pytest.raises(ProxyException) as exc_info: + await create_group(group=scim_group) + + assert int(exc_info.value.code) == 400 + assert "unknown-user" in str(exc_info.value.message) + assert "nested-group-guid" not in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_process_group_patch_remove_unknown_member_does_not_create_user(mocker, scim_upsert_user_enabled): + """A ``remove`` of an id we don't know is an idempotent no-op. Upserting the id + first, only to drop it from the roster, made removals a phantom-user factory.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +@pytest.mark.parametrize( + "operation", + [ + SCIMPatchOperation(op="remove", path='members[value eq "long-gone"]', value=None), + SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone"}]), + SCIMPatchOperation(op="remove", path="members", value=[{"value": "long-gone", "type": "Group"}]), + ], + ids=["path-filter", "unknown-id", "nested-group"], +) +@pytest.mark.asyncio +async def test_process_group_patch_remove_unknown_member_does_not_reject_in_strict_mode( + mocker, scim_upsert_user_disabled, operation +): + """Strict mode must not 400 a removal: refusing to drop an id the IdP already + forgot leaves the roster permanently out of sync.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[operation], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + assert final_members == {"keep-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_remove_drops_member_without_user_row(mocker, scim_upsert_user_enabled): + """Phantom members already on a roster (their user row is gone) must still be + removable, so the removal id is honoured even though it resolves to nothing.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "phantom"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[ + Member(user_id="keep-user", role="user"), + Member(user_id="phantom", role="user"), + ], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +_NESTED_GROUP_ID = "8f1e9d70-0000-4a0e-9a1e-nested" + + +@pytest.mark.parametrize( + "member_entry, user_rows, team_rows", + [ + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user", _NESTED_GROUP_ID}, set()), + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, set()), + ({"value": _NESTED_GROUP_ID, "type": "Group"}, {"keep-user"}, {_NESTED_GROUP_ID}), + ({"value": _NESTED_GROUP_ID}, {"keep-user"}, {_NESTED_GROUP_ID}), + ], + ids=["phantom-user-row-exists", "user-row-already-deleted", "child-group-is-a-team", "untyped-team-id"], +) +@pytest.mark.asyncio +async def test_process_group_patch_remove_discards_non_user_member( + mocker, scim_upsert_user_enabled, member_entry, user_rows, team_rows +): + """Rosters written before nested groups were understood still carry those ids, + and the IdP removes them exactly as it added them; a removal that resolved its + ids first would classify them as non-users and leave them stuck on the team.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[member_entry])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[ + Member(user_id="keep-user", role="user"), + Member(user_id=_NESTED_GROUP_ID, role="user"), + ], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=user_rows, teams=team_rows), + ) + + create_user_mock.assert_not_called() + assert final_members == {"keep-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_add_keeps_member_typed_user_that_collides_with_team_id( + mocker, scim_upsert_user_enabled +): + """The team lookup only exists to catch nested groups that arrive untyped. An id + the IdP calls a User is a user, and IdP ids collide with team ids easily.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "123456", "type": "User"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="123456", key="new-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams={"123456", "parent-group"}), + ) + + assert create_user_mock.call_args.kwargs["user_id"] == "123456" + assert final_members == {"123456"} + + +@pytest.mark.parametrize("member_type", ["Device", " group ", "Machine"]) +@pytest.mark.asyncio +async def test_process_group_patch_operations_skips_non_user_member_types( + mocker, scim_upsert_user_enabled, member_type +): + """A team holds users, so a member that declares itself to be anything else is + dropped; enumerating the types worth skipping would leave the next one to leak.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "not-a-user", "type": member_type}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == set() + + +@pytest.mark.parametrize( + "team_metadata, expect_provisioned", + [ + ({SCIM_MANAGED_TEAM_METADATA_KEY: True}, False), + ({SCIM_TEAM_DATA_METADATA_KEY: {"displayName": "Child.Apps"}}, False), + ({}, True), + (None, True), + ({SCIM_MANAGED_TEAM_METADATA_KEY: False}, True), + ({SCIM_TEAM_DATA_METADATA_KEY: None}, True), + ], + ids=[ + "scim-managed", + "legacy-scim-data", + "admin-created", + "no-metadata", + "marker-unset", + "legacy-key-without-value", + ], +) +@pytest.mark.asyncio +async def test_process_group_patch_team_match_needs_scim_provenance( + mocker, scim_upsert_user_enabled, team_metadata, expect_provisioned +): + """A bare member id that names a team is only evidence of a nested group when the + identity provider is what wrote that team. Teams created here can share an id with + a real user, and skipping those members stops provisioning them entirely.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "child-team"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set()) + prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="child-team", metadata=team_metadata) + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="child-team", key="new-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=prisma_client, + ) + + assert create_user_mock.called is expect_provisioned + assert final_members == ({"child-team"} if expect_provisioned else set()) + + +@pytest.mark.asyncio +async def test_create_group_strict_mode_rejects_id_matching_admin_created_team(mocker, scim_upsert_user_disabled): + """Strict mode drops nested groups but reports unknown users. A team an admin + created here says nothing about the member, so the member is an unknown user.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="parent-group", + displayName="Parent Group", + members=[SCIMMember(value="admin-team")], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock( + return_value=_member_resolution_prisma( + mocker, users=set(), teams=set(), unmanaged_teams=frozenset({"admin-team"}) + ) + ), + ) + + with pytest.raises(ProxyException) as exc_info: + await create_group(group=scim_group) + + assert int(exc_info.value.code) == 400 + assert "admin-team" in str(exc_info.value.message) + + +@pytest.mark.asyncio +async def test_create_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """The provenance the classifier reads only exists if the group writes stamp it; + a SCIM-created team that carries no mark looks admin-created forever after.""" + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id="child-group", + displayName="Child.Apps", + members=[], + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=_member_resolution_prisma(mocker, users=set(), teams=set())), + ) + new_team_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.new_team", + AsyncMock(return_value=mocker.MagicMock()), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + + await create_group(group=scim_group) + + assert new_team_mock.call_args.kwargs["data"].metadata == {SCIM_MANAGED_TEAM_METADATA_KEY: True} + + +@pytest.mark.asyncio +async def test_update_group_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """A PUT full sync adopts a team the identity provider now owns, and the stamp has + to land alongside the existing metadata rather than replacing it.""" + import json + + group_id = "child-group" + existing_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Child.Apps", + members=[], + members_with_roles=[], + metadata={"existing_key": "kept"}, + ) + scim_group = SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Child.Apps", + members=[], + ) + + prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set()) + prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists", + AsyncMock(return_value=existing_team), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", + AsyncMock(return_value=scim_group), + ) + + await update_group(group_id=group_id, group=scim_group) + + written = json.loads(prisma_client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"]) + assert written[SCIM_MANAGED_TEAM_METADATA_KEY] is True + assert written["existing_key"] == "kept" + assert SCIM_TEAM_DATA_METADATA_KEY in written + + +@pytest.mark.asyncio +async def test_process_group_patch_stamps_scim_provenance(mocker, scim_upsert_user_enabled): + """PATCH is how Okta adopts a group, so a membership-only patch has to stamp the + team too; otherwise the group it manages never gains provenance.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user"}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + metadata={"existing_key": "kept"}, + ) + + update_data, _, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + assert update_data["metadata"][SCIM_MANAGED_TEAM_METADATA_KEY] is True + assert update_data["metadata"]["existing_key"] == "kept" + + +@pytest.mark.parametrize("member_type", ["direct", "Device"]) +@pytest.mark.asyncio +async def test_process_group_patch_keeps_existing_user_with_unrecognized_type( + mocker, scim_upsert_user_enabled, member_type +): + """Clients do stamp non-canonical types on real members (RFC 7643 defines + ``direct`` for ``User.groups``). Dropping a member whose id is a live user would + revoke that user's team access on the next full sync, so the type is only + grounds for skipping once the user lookup has missed.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "real-user", "type": member_type}])], + ) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[], + ) + create_user_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=NewUserResponse(user_id="phantom-user", key="phantom-key")), + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"real-user"}, teams=set()), + ) + + create_user_mock.assert_not_called() + assert final_members == {"real-user"} + + +@pytest.mark.parametrize( + "second_creation", + [None, NewUserResponse(user_id="dup-user", key="second-key")], + ids=["second-creation-fails", "both-creations-succeed"], +) +@pytest.mark.asyncio +async def test_resolve_group_member_ids_dedupes_repeated_member(mocker, scim_upsert_user_enabled, second_creation): + """An id the request lists twice is one member. Admitting it twice writes a + duplicate members_with_roles row, and the second creation of the same id fails + against the real unique constraint even when the first one succeeded.""" + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(side_effect=[NewUserResponse(user_id="dup-user", key="first-key"), second_creation]), + ) + + result = await _resolve_group_member_ids( + members=[SCIMMember(value="dup-user"), SCIMMember(value="dup-user")], + created_via="scim_group_membership", + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + assert result.all_member_ids == ["dup-user"] + + +@pytest.mark.parametrize( + "operation", + [ + SCIMPatchOperation(op="add", path="members", value=[{"value": " "}]), + SCIMPatchOperation(op="remove", path="members", value=[{"value": " "}]), + SCIMPatchOperation(op="remove", path='members[value eq " "]', value=None), + ], + ids=["add", "remove", "remove-path-filter"], +) +@pytest.mark.asyncio +async def test_process_group_patch_rejects_blank_member_id(mocker, scim_upsert_user_enabled, operation): + """A blank id names nobody. The removal path stopped resolving its members, so it + has to keep rejecting one on its own.""" + patch_ops = SCIMPatchOp(schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations=[operation]) + existing_team = LiteLLM_TeamTable( + team_id="parent-group", + team_alias="Parent Group", + members=[], + members_with_roles=[Member(user_id="keep-user", role="user")], + ) + + with pytest.raises(HTTPException) as exc_info: + await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=_member_resolution_prisma(mocker, users={"keep-user"}, teams=set()), + ) + + assert exc_info.value.status_code == 400 + + +def test_scim_member_round_trips_type(): + """``type`` has to survive parsing; dropping it is what made a nested group + look like a user id.""" + assert SCIMMember.model_validate({"value": "x", "type": "Group"}).type == "Group" + assert SCIMMember(value="x").type is None + + +@pytest.mark.parametrize("junk_type", [123, True, {}, [], 1.5]) +def test_scim_member_treats_non_string_type_as_absent(junk_type): + """Before ``type`` was a field, junk in it was parsed away; typing the field must + not start rejecting those requests, and both parsers have to agree it is typeless.""" + assert SCIMMember.model_validate({"value": "x", "type": junk_type}).type is None + assert _parse_member_entries([{"value": "x", "type": junk_type}])[0].type is None + + +@pytest.mark.asyncio +async def test_get_groups_members_are_typed_as_users(mocker): + """Group members we report back are always users, and saying so keeps the + response from emitting a null ``type``.""" + team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[Member(user_id="member-1", role="user")], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team]) + mock_prisma_client.db.litellm_teamtable.count = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mocker.MagicMock(user_id="member-1", user_email="member-1@example.com") + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + + response = await get_groups(startIndex=1, count=10, filter=None) + + assert [m.type for m in response.Resources[0].members] == ["User"] From 802ed1c74fa9770685d08baf919d0d0dd831806c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 09:46:49 -0700 Subject: [PATCH 29/54] fix(ui): show public model names in usage breakdowns --- .../common_daily_activity.py | 29 +++--- .../test_common_daily_activity.py | 91 ++++++++++++++++++- .../EntityUsage/EntityUsage.test.tsx | 38 ++++++-- .../components/EntityUsage/EntityUsage.tsx | 23 ++++- .../components/ModelViewToggle.tsx | 29 ++++++ .../components/UsagePageView.test.tsx | 44 +++++++-- .../_components/components/UsagePageView.tsx | 34 ++----- .../src/components/activity_metrics.test.tsx | 54 +++++++++++ .../src/components/activity_metrics.tsx | 2 +- 9 files changed, 285 insertions(+), 59 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 9b756d14815..9bd358289b3 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -233,27 +233,28 @@ def update_breakdown_metrics( ) # Update model group breakdown - if record.model_group and record.model_group not in breakdown.model_groups: - breakdown.model_groups[record.model_group] = MetricWithMetadata( + model_group_key = record.model_group or record.model + if model_group_key and model_group_key not in breakdown.model_groups: + breakdown.model_groups[model_group_key] = MetricWithMetadata( metrics=SpendMetrics(), - metadata=model_metadata.get(record.model_group, {}), + metadata=model_metadata.get(model_group_key, {}), ) - if record.model_group: - breakdown.model_groups[record.model_group].metrics = update_metrics( - breakdown.model_groups[record.model_group].metrics, record + if model_group_key: + breakdown.model_groups[model_group_key].metrics = update_metrics( + breakdown.model_groups[model_group_key].metrics, record ) # Update API key breakdown for this model - if record.api_key not in breakdown.model_groups[record.model_group].api_key_breakdown: - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key] = KeyMetricWithMetadata( + if record.api_key not in breakdown.model_groups[model_group_key].api_key_breakdown: + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key] = KeyMetricWithMetadata( metrics=SpendMetrics(), metadata=KeyMetadata( key_alias=api_key_metadata.get(record.api_key, {}).get("key_alias", None), team_id=api_key_metadata.get(record.api_key, {}).get("team_id", None), ), ) - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics = update_metrics( - breakdown.model_groups[record.model_group].api_key_breakdown[record.api_key].metrics, + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics = update_metrics( + breakdown.model_groups[model_group_key].api_key_breakdown[record.api_key].metrics, record, ) @@ -574,11 +575,11 @@ def _build_aggregated_sql_query( date, api_key, model, - model_group, + COALESCE(model_group, model) AS model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint, - GROUPING(date, api_key, model, model_group, + GROUPING(date, api_key, model, COALESCE(model_group, model), custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level, SUM(spend)::float AS spend, @@ -599,8 +600,8 @@ def _build_aggregated_sql_query( (date, api_key), (date, model), (date, model, api_key), - (date, model_group), - (date, model_group, api_key), + (date, COALESCE(model_group, model)), + (date, COALESCE(model_group, model), api_key), (date, custom_llm_provider), (date, custom_llm_provider, api_key), (date, mcp_namespaced_tool_name), diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 0769568d6cd..5aee7ff0236 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -650,14 +650,14 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): assert key_data.metrics.spend == 10.0 -def _daily_user_spend_record(*, user_id, api_key, spend): +def _daily_user_spend_record(*, user_id, api_key, spend, model="gpt-4", model_group="gpt-4"): """A LiteLLM_DailyUserSpend row as the per-user breakdown reads it.""" return SimpleNamespace( date="2024-01-01", user_id=user_id, api_key=api_key, - model="gpt-4", - model_group="gpt-4", + model=model, + model_group=model_group, custom_llm_provider="openai", mcp_namespaced_tool_name=None, endpoint="/chat/completions", @@ -731,6 +731,64 @@ async def test_get_daily_activity_applies_resolve_entity_metadata_to_breakdown() assert entities["user-no-email"].metadata == {} +@pytest.mark.asyncio +async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback(): + """The usage UI labels model traffic with the model_groups breakdown. + + Keys must be the requested public model name (model_group), and rows with a + NULL or empty model_group (pre-routing failures, rows written before the + column existed) must fall back to their model name instead of being dropped + from the breakdown. The models breakdown keeps the upstream litellm names. + """ + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + + records = [ + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=7.0, model="gpt-5.2", model_group="gpt-5.2-eu" + ), + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=3.0, model="gpt-5.2", model_group=None + ), + _daily_user_spend_record( + user_id="u1", api_key="key-1", spend=2.0, model="claude-x", model_group="" + ), + ] + + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=len(records)) + mock_table.find_many = AsyncMock(return_value=records) + mock_prisma.db.litellm_dailyuserspend = mock_table + mock_prisma.db.litellm_verificationtoken = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + result = await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + entity_metadata_field=None, + start_date="2024-01-01", + end_date="2024-01-01", + model=None, + api_key=None, + page=1, + page_size=1000, + ) + + breakdown = result.results[0].breakdown + + assert set(breakdown.model_groups.keys()) == {"gpt-5.2-eu", "gpt-5.2", "claude-x"} + assert breakdown.model_groups["gpt-5.2-eu"].metrics.spend == 7.0 + assert breakdown.model_groups["gpt-5.2"].metrics.spend == 3.0 + assert breakdown.model_groups["claude-x"].metrics.spend == 2.0 + assert breakdown.model_groups["gpt-5.2"].api_key_breakdown["key-1"].metrics.spend == 3.0 + + assert set(breakdown.models.keys()) == {"gpt-5.2", "claude-x"} + assert breakdown.models["gpt-5.2"].metrics.spend == 10.0 + assert breakdown.models["claude-x"].metrics.spend == 2.0 + + class TestAdjustDatesForTimezone: """ Regression tests for the timezone double-counting bug. @@ -852,6 +910,33 @@ class TestBuildAggregatedSqlQuery: assert "model = $4" in sql assert "api_key = $5" in sql + def test_model_group_rollups_fall_back_to_model_name(self): + """Aggregated model_groups rollups must coalesce NULL model_group to model. + + The (date, model_group) grouping level cannot recover the model column + after the fact (it is rolled up), so the fallback has to happen in SQL; + without it, group-less rows silently vanish from the model_groups + breakdown that the usage UI now renders by default. + """ + sql, _ = _build_aggregated_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id=None, + start_date="2026-07-01", + end_date="2026-07-01", + model=None, + api_key=None, + ) + + normalized = " ".join(sql.split()) + assert "COALESCE(model_group, model) AS model_group" in normalized + assert ( + "GROUPING(date, api_key, model, COALESCE(model_group, model), " + "custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized + ) + assert "(date, COALESCE(model_group, model)), (date, COALESCE(model_group, model), api_key)," in normalized + assert "(date, model_group)" not in normalized + @pytest.mark.asyncio async def test_get_daily_activity_aggregated_empty_result_set(): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 89c38c6274f..82ca66b10c0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -1,4 +1,4 @@ -import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { act, fireEvent, render, screen, waitFor, within } from "@testing-library/react"; import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import * as networking from "@/components/networking"; import EntityUsage from "./EntityUsage"; @@ -497,7 +497,7 @@ describe("EntityUsage", () => { it.each([ ["Cost", "Tag Spend Overview"], - ["Model Activity", "metrics-source:models"], + ["Model Activity", "metrics-source:model_groups"], ["Key Activity", "metrics-source:api_keys"], ["Endpoint Activity", "Endpoint Usage Panel"], ])("shows only the %s panel for a non-team entity type", async (tabLabel, marker) => { @@ -518,7 +518,7 @@ describe("EntityUsage", () => { it.each([ ["Cost", "Team Spend Overview"], - ["Model Activity", "metrics-source:models"], + ["Model Activity", "metrics-source:model_groups"], ["Agent Activity", "metrics-source:entities"], ["Key Activity", "metrics-source:api_keys"], ["Endpoint Activity", "Endpoint Usage Panel"], @@ -584,15 +584,41 @@ describe("EntityUsage", () => { expect(screen.getByText("Request / Token Consumption")).toBeInTheDocument(); }); - it("should display Top Models title for non-agent entity types", async () => { + it("should display Top Public Model Names title for non-agent entity types", async () => { render(); await waitFor(() => { expect(mockTagDailyActivityCall).toHaveBeenCalled(); }); - const topModelsElements = screen.getAllByText("Top Models"); - expect(topModelsElements.length).toBeGreaterThan(0); + expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); + }); + + it("defaults Model Activity to public model names and toggles to litellm models", async () => { + const { container } = render(); + + await waitFor(() => { + expect(mockTagDailyActivityCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getByText("Model Activity")); + }); + + const modelActivityPanel = () => selectedPanels(container)[0] as HTMLElement; + expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups"); + + act(() => { + fireEvent.click(within(modelActivityPanel()).getByText("Litellm Model Name")); + }); + + expect(modelActivityPanel().textContent).toContain("metrics-source:models"); + + act(() => { + fireEvent.click(within(modelActivityPanel()).getByText("Public Model Name")); + }); + + expect(modelActivityPanel().textContent).toContain("metrics-source:model_groups"); }); it("should display Top Agents title for agent entity type", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index e330983b6f9..4d44791d1a9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -49,6 +49,7 @@ import { } from "@/components/UsagePage/types"; import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; import EndpointUsage from "../EndpointUsage/EndpointUsage"; +import ModelViewToggle, { ModelViewType } from "../ModelViewToggle"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import TopModelView from "./TopModelView"; @@ -110,6 +111,7 @@ const ENTITY_FETCH_FNS: Record Promise> = { const EntityUsage: React.FC = ({ accessToken, entityType, entityId, entityList, dateValue }) => { const { teams } = useTeams(); const [selectedTags, setSelectedTags] = useState([]); + const [modelViewType, setModelViewType] = useState("groups"); const [topKeysLimit, setTopKeysLimit] = useState(5); const [topModelsLimit, setTopModelsLimit] = useState(5); const [topAgentsLimit, setTopAgentsLimit] = useState(5); @@ -153,14 +155,15 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData; - const modelMetrics = processActivityData(spendData, "models", teams || []); + const modelBreakdownKey = modelViewType === "groups" ? "model_groups" : "models"; + const modelMetrics = processActivityData(spendData, modelBreakdownKey, teams || []); const keyMetrics = processActivityData(spendData, "api_keys", teams || []); const agentMetrics = entityType === "team" ? processActivityData(agentSpendData, "entities", teams || []) : {}; const getTopModels = () => { const modelSpend: { [key: string]: any } = {}; spendData.results.forEach((day) => { - Object.entries(day.breakdown.models || {}).forEach(([model, metrics]) => { + Object.entries(day.breakdown[modelBreakdownKey] || {}).forEach(([model, metrics]) => { if (!modelSpend[model]) { modelSpend[model] = { spend: 0, @@ -406,6 +409,8 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti const capitalizedEntityLabel = entityType.charAt(0).toUpperCase() + entityType.slice(1); + const modelViewTitle = modelViewType === "groups" ? "Top Public Model Names" : "Top Litellm Models"; + const costPanel = ( {/* Total Spend Card */} @@ -604,7 +609,10 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti {/* Top Models */} - {entityType === "agent" ? "Top Agents" : "Top Models"} +
+ {entityType === "agent" ? "Top Agents" : modelViewTitle} + +
= ({ accessToken, entityType, enti { key: "models", label: entityType === "agent" ? "Request / Token Consumption" : "Model Activity", - content: , + content: ( + <> +
+ +
+ + + ), }, ...(entityType === "team" ? [{ key: "agents", label: "Agent Activity", content: }] diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx new file mode 100644 index 00000000000..0ee6dd3b19c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/ModelViewToggle.tsx @@ -0,0 +1,29 @@ +export type ModelViewType = "groups" | "individual"; + +const MODEL_VIEW_OPTIONS: readonly { value: ModelViewType; label: string }[] = [ + { value: "groups", label: "Public Model Name" }, + { value: "individual", label: "Litellm Model Name" }, +]; + +interface ModelViewToggleProps { + value: ModelViewType; + onChange: (value: ModelViewType) => void; +} + +export default function ModelViewToggle({ value, onChange }: ModelViewToggleProps) { + return ( +
+ {MODEL_VIEW_OPTIONS.map((option) => ( + + ))} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx index cf122137f91..98dae51fa37 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.test.tsx @@ -30,8 +30,10 @@ vi.mock("@/components/networking", () => ({ // Mock child components to simplify testing vi.mock("@/components/activity_metrics", () => ({ - ActivityMetrics: () =>
Activity Metrics
, - processActivityData: () => ({ data: [], metadata: {} }), + ActivityMetrics: ({ modelMetrics }: { modelMetrics?: { __source?: string } }) => ( +
{`activity-source:${modelMetrics?.__source ?? "none"}`}
+ ), + processActivityData: (_data: unknown, key: string) => ({ __source: key }), })); vi.mock("@/components/view_user_spend", () => ({ @@ -1043,8 +1045,8 @@ describe("UsagePage", () => { // Default should be "groups" view showing "Top Public Model Names" expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); - expect(screen.getByText("Public Model Name")).toBeInTheDocument(); - expect(screen.getByText("Litellm Model Name")).toBeInTheDocument(); + expect(screen.getAllByText("Public Model Name").length).toBeGreaterThan(0); + expect(screen.getAllByText("Litellm Model Name").length).toBeGreaterThan(0); }); it("should switch to Litellm Model Name view on toggle click", async () => { @@ -1055,7 +1057,7 @@ describe("UsagePage", () => { }); // Click the "Litellm Model Name" toggle - const litellmToggle = screen.getByText("Litellm Model Name"); + const litellmToggle = screen.getAllByText("Litellm Model Name")[0]; act(() => { fireEvent.click(litellmToggle); }); @@ -1074,7 +1076,7 @@ describe("UsagePage", () => { }); // Switch to individual first - const litellmToggle = screen.getByText("Litellm Model Name"); + const litellmToggle = screen.getAllByText("Litellm Model Name")[0]; act(() => { fireEvent.click(litellmToggle); }); @@ -1084,7 +1086,7 @@ describe("UsagePage", () => { }); // Switch back to groups - const publicToggle = screen.getByText("Public Model Name"); + const publicToggle = screen.getAllByText("Public Model Name")[0]; act(() => { fireEvent.click(publicToggle); }); @@ -1093,6 +1095,34 @@ describe("UsagePage", () => { expect(screen.getByText("Top Public Model Names")).toBeInTheDocument(); }); }); + + it("should feed the Model Activity tab from the model_groups breakdown by default", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + expect(screen.getByText("activity-source:model_groups")).toBeInTheDocument(); + expect(screen.queryByText("activity-source:models")).not.toBeInTheDocument(); + }); + + it("should switch the Model Activity tab to the litellm models breakdown on toggle click", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + act(() => { + fireEvent.click(screen.getAllByText("Litellm Model Name")[0]); + }); + + await waitFor(() => { + expect(screen.getByText("activity-source:models")).toBeInTheDocument(); + }); + expect(screen.queryByText("activity-source:model_groups")).not.toBeInTheDocument(); + }); }); describe("customer usage banner", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx index d2f75609d18..46a17017d39 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsagePageView.tsx @@ -55,6 +55,7 @@ import { DailyData, KeyMetricWithMetadata, MetricWithMetadata } from "@/componen import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; import EndpointUsage from "./EndpointUsage/EndpointUsage"; import EntityUsage, { EntityList } from "./EntityUsage/EntityUsage"; +import ModelViewToggle, { ModelViewType } from "./ModelViewToggle"; import SpendByProvider from "./EntityUsage/SpendByProvider"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import UsageAIChatPanel from "./UsageAIChatPanel"; @@ -143,7 +144,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { // For admins: null means global view (all users), a string means filter by that user // For non-admins: always set to their own user ID const [selectedUserId, setSelectedUserId] = useState(isAdmin ? null : userID || null); - const [modelViewType, setModelViewType] = useState<"groups" | "individual">("groups"); + const [modelViewType, setModelViewType] = useState("groups"); const [isCloudZeroModalOpen, setIsCloudZeroModalOpen] = useState(false); const [isGlobalExportModalOpen, setIsGlobalExportModalOpen] = useState(false); const [isAiChatOpen, setIsAiChatOpen] = useState(false); @@ -438,7 +439,10 @@ const UsagePage: React.FC = ({ teams, organizations }) => { () => [...userSpendData.results].sort((a, b) => new Date(a.date).getTime() - new Date(b.date).getTime()), [userSpendData.results], ); - const modelMetrics = useMemo(() => processActivityData(userSpendData, "models", teams), [userSpendData, teams]); + const modelMetrics = useMemo( + () => processActivityData(userSpendData, modelViewType === "groups" ? "model_groups" : "models", teams), + [userSpendData, modelViewType, teams], + ); const keyMetrics = useMemo(() => processActivityData(userSpendData, "api_keys", teams), [userSpendData, teams]); const mcpServerMetrics = useMemo( () => processActivityData(userSpendData, "mcp_servers", teams), @@ -753,28 +757,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { value={topModelsLimit} onChange={(value) => setTopModelsLimit(value as number)} /> -
- - -
+ {loading ? ( @@ -839,6 +822,9 @@ const UsagePage: React.FC = ({ teams, organizations }) => { {/* Activity Panel */} +
+ +
diff --git a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx index f37643f2366..21d66991bb5 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.test.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.test.tsx @@ -716,6 +716,60 @@ describe("processActivityData", () => { expect(result["gpt-4"].total_spend).toBe(100.5); }); + it("should process model_groups data keyed by public model name including fallback entries", () => { + const upstreamModelMetrics = { + ...EMPTY_SPEND_METRICS, + spend: 10, + api_requests: 10, + successful_requests: 10, + }; + const dailyActivityWithModelGroups: { results: DailyData[] } = { + results: [ + { + date: "2025-01-01", + metrics: upstreamModelMetrics, + breakdown: { + ...EMPTY_BREAKDOWN, + models: { + "gpt-5.2": { + metrics: upstreamModelMetrics, + metadata: {}, + api_key_breakdown: {}, + }, + }, + model_groups: { + "gpt-5.2-eu": { + metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, successful_requests: 7 }, + metadata: {}, + api_key_breakdown: { + "key-1": { + metrics: { ...EMPTY_SPEND_METRICS, spend: 7, api_requests: 7, total_tokens: 700 }, + metadata: { key_alias: "eu-key", team_id: "team1" }, + }, + }, + }, + "gpt-5.2": { + metrics: { ...EMPTY_SPEND_METRICS, spend: 3, api_requests: 3, successful_requests: 3 }, + metadata: {}, + api_key_breakdown: {}, + }, + }, + }, + }, + ], + }; + + const result = processActivityData(dailyActivityWithModelGroups, "model_groups"); + + expect(Object.keys(result).sort()).toEqual(["gpt-5.2", "gpt-5.2-eu"]); + expect(result["gpt-5.2-eu"].label).toBe("gpt-5.2-eu"); + expect(result["gpt-5.2-eu"].total_spend).toBe(7); + expect(result["gpt-5.2-eu"].top_api_keys).toHaveLength(1); + expect(result["gpt-5.2-eu"].top_api_keys[0].key_alias).toBe("eu-key"); + expect(result["gpt-5.2"].total_spend).toBe(3); + expect(result["gpt-5.2"].total_requests).toBe(3); + }); + it("should process data for mcp_servers key", () => { const dailyActivityWithMCP: { results: DailyData[] } = { results: [ diff --git a/ui/litellm-dashboard/src/components/activity_metrics.tsx b/ui/litellm-dashboard/src/components/activity_metrics.tsx index 54ac5ae0ee6..b8199ee29b6 100644 --- a/ui/litellm-dashboard/src/components/activity_metrics.tsx +++ b/ui/litellm-dashboard/src/components/activity_metrics.tsx @@ -362,7 +362,7 @@ export const formatKeyLabel = (modelData: KeyMetricWithMetadata, model: string, // Process data function export const processActivityData = ( dailyActivity: { results: DailyData[] }, - key: "models" | "api_keys" | "mcp_servers" | "entities", + key: "models" | "model_groups" | "api_keys" | "mcp_servers" | "entities", teams: Team[] = [], ): Record => { const modelMetrics: Record = {}; From fe1670fc068bbce3a370001f4103a03804f2d0df Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 09:47:54 -0700 Subject: [PATCH 30/54] fix(ui): size object permissions card grid by container width (#35019) The card variant used viewport breakpoints (md:grid-cols-2 lg:grid-cols-3) but every card usage sits in a one-third-width grid cell, so on desktop the narrow card still rendered three internal columns of roughly 100px each and the text spilled out of its boxes. Switch to Tailwind container queries so the internal column count follows the card's own width --- .../src/components/object_permissions_view.tsx | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/object_permissions_view.tsx b/ui/litellm-dashboard/src/components/object_permissions_view.tsx index 687d1a5a846..327e127e8fe 100644 --- a/ui/litellm-dashboard/src/components/object_permissions_view.tsx +++ b/ui/litellm-dashboard/src/components/object_permissions_view.tsx @@ -28,7 +28,7 @@ export function ObjectPermissionsView({ const searchTools = objectPermission?.search_tools || []; const content = ( -
+
+
Object Permissions From 40878a1ed52915d3e98ba019f14dd7f7bfdc4e98 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 09:48:08 -0700 Subject: [PATCH 31/54] fix(proxy): allow /key/update to identify the key by key_alias (#34851) * fix(proxy): allow /key/update to identify the key by key_alias * fix(ui): drop machine-dependent union-order churn from generated schema.d.ts --- litellm/proxy/_types.py | 7 +- .../key_management_endpoints.py | 71 +++++-- .../test_key_management_endpoints.py | 189 ++++++++++++++++++ tests/test_litellm/proxy/test_proxy_types.py | 20 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 +- 5 files changed, 277 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6ffd6816ca6..c6d4ee1120a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1169,7 +1169,6 @@ class GenerateKeyResponse(KeyRequestBase): class UpdateKeyRequest(KeyRequestBase): # Note: the defaults of all Params here MUST BE NONE # else they will get overwritten - key: str # type: ignore duration: Optional[str] = None spend: Optional[float] = None metadata: Optional[dict] = None @@ -1186,6 +1185,12 @@ class UpdateKeyRequest(KeyRequestBase): raise ValueError("temp_budget_increase and temp_budget_expiry must be set together") return self + @model_validator(mode="after") + def validate_key_identifier(self) -> "UpdateKeyRequest": + if self.key is None and self.key_alias is None: + raise ValueError("either key or key_alias must be provided") + return self + class RegenerateKeyRequest(GenerateKeyRequest): # This needs to be different from UpdateKeyRequest, because "key" is optional for this diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index e7ee5ffa849..a94a75fdfa3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2023,7 +2023,7 @@ def _validate_max_budget(max_budget: Optional[float]) -> None: async def _get_and_validate_existing_key( - token: str, prisma_client: Optional[PrismaClient] + token: str | None, prisma_client: Optional[PrismaClient], key_alias: str | None = None ) -> LiteLLM_VerificationToken: """ Get existing key from database and validate it exists. @@ -2031,12 +2031,13 @@ async def _get_and_validate_existing_key( Args: token: The key token to look up prisma_client: Prisma client instance + key_alias: Alias to look the key up by when token is not provided Returns: LiteLLM_VerificationToken: The existing key row Raises: - ProxyException: 404 if key is not found + ProxyException: 404 if key is not found, 400 if the alias matches multiple keys """ if prisma_client is None: raise HTTPException( @@ -2044,19 +2045,65 @@ async def _get_and_validate_existing_key( detail={"error": "Database not connected"}, ) - hashed_token = _hash_token_if_needed(token=token) + if token is not None: + hashed_token = _hash_token_if_needed(token=token) - existing_key_row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": hashed_token}) + existing_key_row: LiteLLM_VerificationToken | None = await VerificationTokenRepository( + prisma_client + ).table.find_unique(where={"token": hashed_token}) - if existing_key_row is None: + if existing_key_row is None: + raise ProxyException( + message="Key not found.", + type=ProxyErrorTypes.not_found_error, + param="key", + code=status.HTTP_404_NOT_FOUND, + ) + + return existing_key_row + + if key_alias is None: + raise ProxyException( + message="either key or key_alias must be provided", + type=ProxyErrorTypes.bad_request_error, + param="key", + code=status.HTTP_400_BAD_REQUEST, + ) + + rows: list[LiteLLM_VerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( + where={"key_alias": key_alias}, take=2 + ) + + if len(rows) == 0: + raise ProxyException( + message=f"Key not found. No key with key_alias='{key_alias}'.", + type=ProxyErrorTypes.not_found_error, + param="key_alias", + code=status.HTTP_404_NOT_FOUND, + ) + + if len(rows) > 1: + raise ProxyException( + message=f"Multiple keys share key_alias='{key_alias}', so it cannot be used as an identifier.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=status.HTTP_400_BAD_REQUEST, + ) + + return rows[0] + + +def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_VerificationToken) -> str: + if data.key is not None: + return data.key + if existing_key_row.token is None: raise ProxyException( message="Key not found.", type=ProxyErrorTypes.not_found_error, param="key", code=status.HTTP_404_NOT_FOUND, ) - - return existing_key_row + return existing_key_row.token async def _process_single_key_update( @@ -2508,8 +2555,8 @@ async def update_key_fn( Update an existing API key's parameters. Parameters: - - key: str - The key to update - - key_alias: Optional[str] - User-friendly key alias + - key: Optional[str] - The key to update. Either key or key_alias must be provided. + - key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases) - user_id: Optional[str] - User ID associated with key - team_id: Optional[str] - Team ID associated with key - agent_id: Optional[str] - The agent id associated with the key. @@ -2592,14 +2639,14 @@ async def update_key_fn( detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"}, ) - data_json: dict = data.model_dump(exclude_unset=True) - key = data_json.pop("key") - # get the row from db existing_key_row = await _get_and_validate_existing_key( token=data.key, prisma_client=prisma_client, + key_alias=data.key_alias, ) + key = _resolve_token_to_update(data=data, existing_key_row=existing_key_row) + data.key = key await _validate_update_key_data( data=data, diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 51f72f91dc3..867ef759fb3 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 @@ -2507,6 +2507,195 @@ async def test_update_key_nonexistent_key_returns_404(monkeypatch): assert "Authentication Error" not in str(exc_info.value.message) +def _setup_update_key_mocks(monkeypatch, mock_prisma_client): + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + + +@pytest.mark.asyncio +async def test_update_key_by_alias_only(monkeypatch): + """ + /key/update identified by key_alias alone resolves the key row via + find_many on the alias and updates using the resolved token. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken( + token=hashed_token, + key_alias="prod-alias", + user_id="test-user", + max_budget=200.0, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[key_in_db] + ) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=None + ) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"max_budget": 50.0}} + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + request_data = UpdateKeyRequest(key_alias="prod-alias", max_budget=50.0) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache: + mock_delete_cache.return_value = None + result = await update_key_fn( + request=MagicMock(), + data=request_data, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_many.assert_called_once_with( + where={"key_alias": "prod-alias"}, take=2 + ) + assert request_data.key == hashed_token + mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_not_called() + mock_prisma_client.update_data.assert_awaited_once() + assert mock_prisma_client.update_data.call_args.kwargs["token"] == hashed_token + assert ( + mock_prisma_client.update_data.call_args.kwargs["data"]["token"] == hashed_token + ) + assert result["key"] == hashed_token + + +@pytest.mark.asyncio +async def test_update_key_by_alias_not_found_returns_404(monkeypatch): + """ + /key/update with a key_alias matching no key returns 404. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[] + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key_alias="no-such-alias", max_budget=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "404" + assert "not found" in str(exc_info.value.message).lower() + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_by_duplicate_alias_returns_400(monkeypatch): + """ + /key/update with a key_alias shared by multiple keys returns 400 + instead of silently updating one of them. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + rows = [ + LiteLLM_VerificationToken(token="hashed-token-1", key_alias="dup-alias"), + LiteLLM_VerificationToken(token="hashed-token-2", key_alias="dup-alias"), + ] + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=rows + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with pytest.raises(ProxyException) as exc_info: + await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key_alias="dup-alias", max_budget=50.0), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert str(exc_info.value.code) == "400" + assert "multiple keys" in str(exc_info.value.message).lower() + mock_prisma_client.update_data.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_key_with_key_and_alias_selects_by_key(monkeypatch): + """ + Regression: passing both key and key_alias keeps today's behavior. The key + identifies the row (find_unique, never find_many) and key_alias is the new + alias to set; the response echoes the caller-passed key. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + hashed_token = "0d62f396c1317066f55a96086517047c737087c61eb2bf016b72e6298927b15b" + key_in_db = LiteLLM_VerificationToken( + token=hashed_token, + key_alias="old-name", + user_id="test-user", + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=key_in_db + ) + mock_prisma_client.db.litellm_verificationtoken.find_first = AsyncMock( + return_value=None + ) + mock_prisma_client.update_data = AsyncMock( + return_value={"data": {"key_alias": "new-name"}} + ) + _setup_update_key_mocks(monkeypatch, mock_prisma_client) + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin-user" + ) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object" + ) as mock_delete_cache: + mock_delete_cache.return_value = None + result = await update_key_fn( + request=MagicMock(), + data=UpdateKeyRequest(key="sk-test-key", key_alias="new-name"), + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + mock_prisma_client.db.litellm_verificationtoken.find_many.assert_not_called() + mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once() + assert mock_prisma_client.update_data.call_args.kwargs["token"] == "sk-test-key" + assert result["key"] == "sk-test-key" + + @pytest.mark.asyncio async def test_block_key_existing_key_succeeds(monkeypatch): """ diff --git a/tests/test_litellm/proxy/test_proxy_types.py b/tests/test_litellm/proxy/test_proxy_types.py index c8e0b3a730a..5354de182a0 100644 --- a/tests/test_litellm/proxy/test_proxy_types.py +++ b/tests/test_litellm/proxy/test_proxy_types.py @@ -139,3 +139,23 @@ def test_key_request_router_settings_keeps_enable_tag_filtering(): dumped = req.router_settings.model_dump(exclude_none=True) assert dumped["enable_tag_filtering"] is True assert dumped["num_retries"] == 2 + + +def test_update_key_request_requires_key_or_key_alias(): + """``/key/update`` can be addressed by ``key`` or by ``key_alias``; + a request with neither has no way to identify the target key and must + fail validation before hitting the endpoint.""" + import pydantic + + from litellm.proxy._types import UpdateKeyRequest + + with pytest.raises(pydantic.ValidationError, match="either key or key_alias must be provided"): + UpdateKeyRequest(max_budget=10.0) + + by_key = UpdateKeyRequest(key="sk-1234") + assert by_key.key == "sk-1234" + assert by_key.key_alias is None + + by_alias = UpdateKeyRequest(key_alias="my-alias") + assert by_alias.key is None + assert by_alias.key_alias == "my-alias" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 541d7a17ae1..79f03978b12 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -6935,8 +6935,8 @@ export interface paths { * @description Update an existing API key's parameters. * * Parameters: - * - key: str - The key to update - * - key_alias: Optional[str] - User-friendly key alias + * - key: Optional[str] - The key to update. Either key or key_alias must be provided. + * - key_alias: Optional[str] - User-friendly key alias. If key is omitted, also identifies the key to update (must match exactly one key, same as /key/delete's key_aliases) * - user_id: Optional[str] - User ID associated with key * - team_id: Optional[str] - Team ID associated with key * - agent_id: Optional[str] - The agent id associated with the key. @@ -32437,7 +32437,7 @@ export interface components { /** Guardrails */ guardrails?: string[] | null; /** Key */ - key: string; + key?: string | null; /** Key Alias */ key_alias?: string | null; /** Max Budget */ From 9b7a6b9b90c771bc430ee7e4dd6153757779fe7e Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 09:48:17 -0700 Subject: [PATCH 32/54] feat(ui): split failed requests into their own series on the cache dashboard (#34862) * feat(ui): chart failed requests as their own series on the cache dashboard Spend logs for failed requests are stored with an empty call_type, so the Cache Hits vs API Requests chart lumped them into an Unknown bar that read as normal LLM API traffic. The activity query now also returns a per-group failed_rows count (status = 'failure') and the dashboard charts it as a third stacked series, so failures are visibly separate from successful requests and cache hits. The chart data transform moves into a pure summarizeCacheActivity helper with unit tests; header stats keep their existing semantics (cache hit ratio still counts failures in the denominator). * refactor(ui): move cache dashboard aggregation server-side with a typed response The /global/activity/cache_hits endpoint previously returned raw per (key, call_type, model) spend-log aggregates typed as LiteLLM_SpendLogs (wrong), and the dashboard reduced them in the browser: grouping by call_type, relabeling empty call_type as Unknown, and computing the stat card totals. All of that now happens server-side. The SQL groups per call_type and splits cache hits vs successful vs failed requests, a new cache_activity module validates rows into Pydantic models and computes totals plus the key-alias/model filter options, and the endpoint declares a real response_model so schema.d.ts types it correctly. The dashboard consumes it through a typed $api react-query hook (filters ride the query key and are applied in SQL instead of the browser), the hand-rolled summarizeCacheActivity transform and the adminGlobalCacheActivity fetch helper are deleted, and the refresh button now actually refetches. The endpoint is UI-internal (hidden from the public swagger), so the response reshape is not a public API break. --- .gitignore | 1 + .../analytics_endpoints.py | 116 ++++------- .../analytics_endpoints/cache_activity.py | 137 +++++++++++++ .../proxy/analytics_endpoints/__init__.py | 0 .../test_analytics_endpoints.py | 133 +++++++++++++ ui/litellm-dashboard/eslint-suppressions.json | 2 +- .../_components/cache_dashboard.test.tsx | 92 ++++++--- .../caching/_components/cache_dashboard.tsx | 182 ++++-------------- .../hooks/caching/useCacheActivity.test.ts | 72 +++++++ .../hooks/caching/useCacheActivity.ts | 32 +++ .../src/components/networking.tsx | 36 ---- ui/litellm-dashboard/src/lib/http/schema.d.ts | 76 +++++--- 12 files changed, 570 insertions(+), 309 deletions(-) create mode 100644 litellm/proxy/analytics_endpoints/cache_activity.py create mode 100644 tests/test_litellm/proxy/analytics_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts diff --git a/.gitignore b/.gitignore index b812d45e349..13f2202305d 100644 --- a/.gitignore +++ b/.gitignore @@ -141,3 +141,4 @@ crash.*.log .coverage ui/litellm-dashboard/out/ +litellm.log diff --git a/litellm/proxy/analytics_endpoints/analytics_endpoints.py b/litellm/proxy/analytics_endpoints/analytics_endpoints.py index 4c1ff31e5a1..cb22c468c1e 100644 --- a/litellm/proxy/analytics_endpoints/analytics_endpoints.py +++ b/litellm/proxy/analytics_endpoints/analytics_endpoints.py @@ -1,105 +1,61 @@ #### Analytics Endpoints ##### from datetime import datetime, timezone -from typing import List, Optional +from typing import Annotated import fastapi from fastapi import APIRouter, Depends, HTTPException, status from litellm.proxy._types import * +from litellm.proxy.analytics_endpoints.cache_activity import CacheActivityResponse, get_cache_activity from litellm.proxy.auth.user_api_key_auth import user_api_key_auth router = APIRouter() +def _parse_date(value: str, param_name: str) -> datetime: + try: + return datetime.strptime(value, "%Y-%m-%d").replace(tzinfo=timezone.utc) + except ValueError: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"{param_name} must be a YYYY-MM-DD date, got {value!r}"}, + ) + + @router.get( "/global/activity/cache_hits", tags=["Budget & Spend Tracking"], dependencies=[Depends(user_api_key_auth)], - responses={ - 200: {"model": List[LiteLLM_SpendLogs]}, - }, + response_model=CacheActivityResponse, include_in_schema=False, ) async def get_global_activity( - start_date: Optional[str] = fastapi.Query( - default=None, - description="Time from which to start viewing spend", - ), - end_date: Optional[str] = fastapi.Query( - default=None, - description="Time till which to view spend", - ), -): + start_date: Annotated[str, fastapi.Query(description="Time from which to start viewing spend")], + end_date: Annotated[str, fastapi.Query(description="Time till which to view spend")], + key_aliases: Annotated[ + list[str] | None, fastapi.Query(description="Only include spend from these key aliases") + ] = None, + models: Annotated[list[str] | None, fastapi.Query(description="Only include spend for these models")] = None, +) -> CacheActivityResponse: """ - Get number of cache hits, vs misses - - { - "daily_data": [ - const chartdata = [ - { - date: 'Jan 22', - cache_hits: 10, - llm_api_calls: 2000 - }, - { - date: 'Jan 23', - cache_hits: 10, - llm_api_calls: 12 - }, - ], - "sum_cache_hits": 20, - "sum_llm_api_calls": 2012 - } + Cache activity for the Admin UI cache dashboard, aggregated per call_type: + cache hits vs successful LLM API requests vs failed requests, plus totals + for the stat cards and the available key-alias/model filter options. """ - - if start_date is None or end_date is None: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "Please provide start_date and end_date"}, - ) - - start_date_obj = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - end_date_obj = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - from litellm.proxy.proxy_server import prisma_client - try: - if prisma_client is None: - raise ValueError( - "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" - ) - - sql_query = """ - SELECT - CASE - WHEN vt."key_alias" IS NOT NULL THEN vt."key_alias" - ELSE 'Unnamed Key' - END AS api_key, - sl."call_type", - sl."model", - COUNT(*) AS total_rows, - SUM(CASE WHEN sl."cache_hit" = 'True' THEN 1 ELSE 0 END) AS cache_hit_true_rows, - SUM(CASE WHEN sl."cache_hit" = 'True' THEN sl."completion_tokens" ELSE 0 END) AS cached_completion_tokens, - SUM(CASE WHEN sl."cache_hit" != 'True' THEN sl."completion_tokens" ELSE 0 END) AS generated_completion_tokens - FROM "LiteLLM_SpendLogs" sl - LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" - WHERE - sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') - AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') - GROUP BY - vt."key_alias", - sl."call_type", - sl."model" - """ - db_response = await prisma_client.db.query_raw(sql_query, start_date_obj, end_date_obj) - - if db_response is None: - return [] - - return db_response - - except Exception as e: + if prisma_client is None: raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": str(e)}, + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail={ + "error": "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys" + }, ) + + return await get_cache_activity( + prisma_client=prisma_client, + start_date=_parse_date(start_date, "start_date"), + end_date=_parse_date(end_date, "end_date"), + key_aliases=key_aliases or [], + models=models or [], + ) diff --git a/litellm/proxy/analytics_endpoints/cache_activity.py b/litellm/proxy/analytics_endpoints/cache_activity.py new file mode 100644 index 00000000000..6e20382f56e --- /dev/null +++ b/litellm/proxy/analytics_endpoints/cache_activity.py @@ -0,0 +1,137 @@ +import asyncio +import json +from datetime import datetime +from typing import TYPE_CHECKING, Sequence + +from pydantic import BaseModel, TypeAdapter + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +UNKNOWN_CALL_TYPE = "Unknown" + + +class CacheActivityGroup(BaseModel): + call_type: str + api_requests: int + cache_hits: int + failed_requests: int + cached_completion_tokens: int + generated_completion_tokens: int + + +class CacheActivityTotals(BaseModel): + api_requests: int + cache_hits: int + failed_requests: int + cached_completion_tokens: int + cache_hit_ratio: float + + +class CacheActivityFilterOptions(BaseModel): + key_aliases: list[str] + models: list[str] + + +class CacheActivityResponse(BaseModel): + groups: list[CacheActivityGroup] + totals: CacheActivityTotals + filter_options: CacheActivityFilterOptions + + +GROUPS_SQL = """ + SELECT + CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type, + (COUNT(*) + - SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END) + - SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END))::int AS api_requests, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)::int AS cache_hits, + SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END)::int AS failed_requests, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN sl."completion_tokens" ELSE 0 END)::int + AS cached_completion_tokens, + SUM(CASE WHEN COALESCE(sl."cache_hit", '') != 'True' THEN sl."completion_tokens" ELSE 0 END)::int + AS generated_completion_tokens + FROM "LiteLLM_SpendLogs" sl + LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + AND ($3::jsonb = '[]'::jsonb + OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb))) + AND ($4::jsonb = '[]'::jsonb + OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb))) + GROUP BY 1 + ORDER BY (COUNT(*)) DESC +""" + +KEY_ALIAS_OPTIONS_SQL = """ + SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias + FROM "LiteLLM_SpendLogs" sl + LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token" + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + ORDER BY 1 +""" + +MODEL_OPTIONS_SQL = """ + SELECT DISTINCT sl."model" AS model + FROM "LiteLLM_SpendLogs" sl + WHERE + sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC') + AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC') + AND sl."model" != '' + ORDER BY 1 +""" + + +class _KeyAliasRow(BaseModel): + key_alias: str + + +class _ModelRow(BaseModel): + model: str + + +_groups_adapter = TypeAdapter(list[CacheActivityGroup]) +_key_alias_rows_adapter = TypeAdapter(list[_KeyAliasRow]) +_model_rows_adapter = TypeAdapter(list[_ModelRow]) + + +def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals: + api_requests = sum(group.api_requests for group in groups) + cache_hits = sum(group.cache_hits for group in groups) + failed_requests = sum(group.failed_requests for group in groups) + all_requests = api_requests + cache_hits + failed_requests + return CacheActivityTotals( + api_requests=api_requests, + cache_hits=cache_hits, + failed_requests=failed_requests, + cached_completion_tokens=sum(group.cached_completion_tokens for group in groups), + cache_hit_ratio=(cache_hits / all_requests) * 100 if all_requests > 0 else 0.0, + ) + + +async def get_cache_activity( + prisma_client: "PrismaClient", + start_date: datetime, + end_date: datetime, + key_aliases: Sequence[str], + models: Sequence[str], +) -> CacheActivityResponse: + group_rows, key_alias_rows, model_rows = await asyncio.gather( + prisma_client.db.query_raw( + GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models)) + ), + prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date), + prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date), + ) + groups = _groups_adapter.validate_python(group_rows or []) + return CacheActivityResponse( + groups=groups, + totals=compute_totals(groups), + filter_options=CacheActivityFilterOptions( + key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])], + models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])], + ), + ) diff --git a/tests/test_litellm/proxy/analytics_endpoints/__init__.py b/tests/test_litellm/proxy/analytics_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py new file mode 100644 index 00000000000..c48b8cfd5a5 --- /dev/null +++ b/tests/test_litellm/proxy/analytics_endpoints/test_analytics_endpoints.py @@ -0,0 +1,133 @@ +""" +The cache dashboard chart is fed by /global/activity/cache_hits. Aggregation +lives server-side: the SQL groups per call_type (splitting cache hits vs +successful vs failed requests; failed spend logs have call_type '' today and +must surface as 'Unknown'), and the endpoint returns chart-ready groups, +totals for the stat cards, and the filter options for the UI dropdowns. +""" + +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy.analytics_endpoints.analytics_endpoints import get_global_activity +from litellm.proxy.analytics_endpoints.cache_activity import ( + GROUPS_SQL, + CacheActivityGroup, + compute_totals, +) + +GROUP_ROWS = [ + { + "call_type": "acompletion", + "api_requests": 1000, + "cache_hits": 300, + "failed_requests": 200, + "cached_completion_tokens": 12000, + "generated_completion_tokens": 48000, + }, + { + "call_type": "Unknown", + "api_requests": 0, + "cache_hits": 0, + "failed_requests": 110, + "cached_completion_tokens": 0, + "generated_completion_tokens": 0, + }, +] +KEY_ALIAS_ROWS = [{"key_alias": "Unnamed Key"}, {"key_alias": "my-key"}] +MODEL_ROWS = [{"model": "gpt-5.1"}] + + +def build_prisma(query_raw: AsyncMock) -> MagicMock: + prisma = MagicMock() + prisma.db.query_raw = query_raw + return prisma + + +def dispatching_query_raw() -> AsyncMock: + async def dispatch(sql: str, *params: object) -> list[dict[str, object]]: + if "GROUP BY" in sql: + return GROUP_ROWS + if "key_alias" in sql: + return KEY_ALIAS_ROWS + return MODEL_ROWS + + return AsyncMock(side_effect=dispatch) + + +@pytest.fixture +def mock_prisma(monkeypatch: pytest.MonkeyPatch) -> MagicMock: + prisma = build_prisma(dispatching_query_raw()) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) + return prisma + + +@pytest.mark.asyncio +async def test_returns_groups_totals_and_filter_options(mock_prisma: MagicMock): + response = await get_global_activity( + start_date="2026-07-01", end_date="2026-07-27", key_aliases=[], models=[] + ) + + assert [group.call_type for group in response.groups] == ["acompletion", "Unknown"] + assert response.groups[0].api_requests == 1000 + assert response.groups[0].failed_requests == 200 + assert response.totals.api_requests == 1000 + assert response.totals.cache_hits == 300 + assert response.totals.failed_requests == 310 + assert response.totals.cached_completion_tokens == 12000 + assert response.totals.cache_hit_ratio == pytest.approx((300 / 1610) * 100) + assert response.filter_options.key_aliases == ["Unnamed Key", "my-key"] + assert response.filter_options.models == ["gpt-5.1"] + + +@pytest.mark.asyncio +async def test_filters_are_passed_to_sql_as_json_arrays(mock_prisma: MagicMock): + await get_global_activity( + start_date="2026-07-01", + end_date="2026-07-27", + key_aliases=["my-key"], + models=["gpt-5.1", "claude-opus-4-8"], + ) + + groups_call = next( + call for call in mock_prisma.db.query_raw.call_args_list if "GROUP BY" in call.args[0] + ) + assert groups_call.args[3] == json.dumps(["my-key"]) + assert groups_call.args[4] == json.dumps(["gpt-5.1", "claude-opus-4-8"]) + + +@pytest.mark.asyncio +async def test_rejects_malformed_dates_with_400(mock_prisma: MagicMock): + with pytest.raises(HTTPException) as exc_info: + await get_global_activity(start_date="07/01/2026", end_date="2026-07-27", key_aliases=[], models=[]) + + assert exc_info.value.status_code == 400 + mock_prisma.db.query_raw.assert_not_called() + + +def test_totals_ratio_is_zero_without_requests(): + totals = compute_totals([]) + + assert totals.cache_hit_ratio == 0.0 + assert totals.api_requests == 0 + + +def test_totals_denominator_includes_failed_requests(): + group = CacheActivityGroup( + call_type="acompletion", + api_requests=60, + cache_hits=20, + failed_requests=20, + cached_completion_tokens=0, + generated_completion_tokens=0, + ) + + assert compute_totals([group]).cache_hit_ratio == pytest.approx(20.0) + + +def test_groups_sql_splits_failures_and_labels_empty_call_type_unknown(): + assert "SUM(CASE WHEN sl.\"status\" = 'failure' THEN 1 ELSE 0 END)" in GROUPS_SQL + assert "CASE WHEN sl.\"call_type\" = '' THEN 'Unknown' ELSE sl.\"call_type\" END" in GROUPS_SQL diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 1c07f2d6247..7686cc05fa6 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -152,7 +152,7 @@ "count": 1 }, "prefer-const": { - "count": 3 + "count": 1 }, "react-hooks/purity": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx index e1d02e9352d..fd8dd011b05 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx @@ -4,36 +4,50 @@ import { screen, waitFor, within } from "@testing-library/react"; import { renderWithProviders } from "../../../../../tests/test-utils"; import CacheDashboard from "./cache_dashboard"; -const { adminGlobalCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({ - adminGlobalCacheActivity: vi.fn(), +const { useCacheActivity, cachingHealthCheckCall } = vi.hoisted(() => ({ + useCacheActivity: vi.fn(), cachingHealthCheckCall: vi.fn(), })); vi.mock("@/components/networking", () => ({ - adminGlobalCacheActivity, cachingHealthCheckCall, })); -const cacheActivity = [ - { - api_key: "sk-1", - model: "gpt-5.1", - call_type: "acompletion", - total_rows: 1500, - cache_hit_true_rows: 300, - cached_completion_tokens: 12000, - generated_completion_tokens: 48000, +vi.mock("@/app/(dashboard)/hooks/caching/useCacheActivity", () => ({ + useCacheActivity, +})); + +const cacheActivity = { + groups: [ + { + call_type: "acompletion", + api_requests: 1000, + cache_hits: 300, + failed_requests: 200, + cached_completion_tokens: 12000, + generated_completion_tokens: 48000, + }, + { + call_type: "aembedding", + api_requests: 550, + cache_hits: 100, + failed_requests: 50, + cached_completion_tokens: 2000, + generated_completion_tokens: 9000, + }, + ], + totals: { + api_requests: 1550, + cache_hits: 400, + failed_requests: 250, + cached_completion_tokens: 14000, + cache_hit_ratio: (400 / 2200) * 100, }, - { - api_key: "sk-2", - model: "text-embedding-3-large", - call_type: "aembedding", - total_rows: 700, - cache_hit_true_rows: 100, - cached_completion_tokens: 2000, - generated_completion_tokens: 9000, + filter_options: { + key_aliases: ["my-key", "Unnamed Key"], + models: ["gpt-5.1", "text-embedding-3-large"], }, -]; +}; const renderDashboard = () => renderWithProviders( @@ -75,7 +89,7 @@ const legendFillByCategory = (card: HTMLElement) => describe("CacheDashboard cache analytics charts", () => { beforeEach(() => { vi.clearAllMocks(); - adminGlobalCacheActivity.mockResolvedValue(cacheActivity); + useCacheActivity.mockReturnValue({ data: cacheActivity, refetch: vi.fn() }); }); it("renders both chart card titles", async () => { @@ -108,8 +122,13 @@ describe("CacheDashboard cache analytics charts", () => { expect(legendFillByCategory(requestsCard)).toEqual({ "LLM API requests": "var(--color-sky-500, #0ea5e9)", "Cache hit": "var(--color-teal-500, #14b8a6)", + "Failed requests": "var(--color-red-500, #ef4444)", }); - expect(barFills(requestsCard)).toEqual(["var(--color-sky-500, #0ea5e9)", "var(--color-teal-500, #14b8a6)"]); + expect(barFills(requestsCard)).toEqual([ + "var(--color-sky-500, #0ea5e9)", + "var(--color-teal-500, #14b8a6)", + "var(--color-red-500, #ef4444)", + ]); }); it("renders the tokens chart with each category legend-bound to its fill and stacked in order", async () => { @@ -133,18 +152,39 @@ describe("CacheDashboard cache analytics charts", () => { } }); - it("stacks the two categories into one column per call_type", async () => { + it("stacks all categories into one column per call_type", async () => { renderDashboard(); const { requestsCard, tokensCard } = await findChartCards(); - for (const card of [requestsCard, tokensCard]) { + const expectedRects = { requests: 6, tokens: 4 }; + for (const [card, rectCount] of [ + [requestsCard, expectedRects.requests], + [tokensCard, expectedRects.tokens], + ] as const) { const rects = Array.from(card.querySelectorAll("path.recharts-rectangle")); - expect(rects).toHaveLength(4); + expect(rects).toHaveLength(rectCount); const xPositions = rects.map((rect) => rect.getAttribute("d")?.split(",")[0]); expect(new Set(xPositions).size).toBe(2); } }); + it("renders the server-computed cache hit ratio", async () => { + renderDashboard(); + + expect(await screen.findByText("18.18%")).toBeInTheDocument(); + }); + + it("passes the date range and selected filters to the activity query", () => { + renderDashboard(); + + expect(useCacheActivity).toHaveBeenCalledWith({ + startDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/), + endDate: expect.stringMatching(/^\d{4}-\d{2}-\d{2}$/), + keyAliases: [], + models: [], + }); + }); + it("formats y-axis ticks with compact notation", async () => { renderDashboard(); const { requestsCard, tokensCard } = await findChartCards(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx index 95b73d1aacb..47c266ceac0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx @@ -19,13 +19,29 @@ import { import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { RefreshCw } from "lucide-react"; -import { adminGlobalCacheActivity, cachingHealthCheckCall } from "@/components/networking"; +import { cachingHealthCheckCall } from "@/components/networking"; +import { useCacheActivity, type CacheActivityGroup } from "@/app/(dashboard)/hooks/caching/useCacheActivity"; // Import the new component import { CacheHealthTab } from "./cache_health"; import CacheSettings from "./cache_settings"; import CoordinationRedisSettings from "./coordination_redis_settings"; +const REQUEST_SERIES = { + apiRequests: "LLM API requests", + cacheHits: "Cache hit", + failed: "Failed requests", +} as const; + +const toChartDatum = (group: CacheActivityGroup) => ({ + name: group.call_type, + [REQUEST_SERIES.apiRequests]: group.api_requests, + [REQUEST_SERIES.cacheHits]: group.cache_hits, + [REQUEST_SERIES.failed]: group.failed_requests, + "Cached Completion Tokens": group.cached_completion_tokens, + "Generated Completion Tokens": group.generated_completion_tokens, +}); + const formatDateWithoutTZ = (date: Date | undefined) => { if (!date) return undefined; return date.toISOString().split("T")[0]; @@ -49,26 +65,6 @@ interface CachePageProps { premiumUser: boolean; } -interface cacheDataItem { - api_key: string; - model: string; - cache_hit_true_rows: number; - cached_completion_tokens: number; - total_rows: number; - generated_completion_tokens: number; - call_type: string; - - // Add other properties as needed -} - -type uiData = { - name: string; - "LLM API requests": number; - "Cache hit": number; - "Cached Completion Tokens": number; - "Generated Completion Tokens": number; -}; - interface CacheHealthResponse { status?: string; cache_type?: string; @@ -97,13 +93,8 @@ const deepParse = (input: any) => { }; const CacheDashboard: React.FC = ({ accessToken, token, userRole, userID, premiumUser }) => { - const [filteredData, setFilteredData] = useState([]); const [selectedApiKeys, setSelectedApiKeys] = useState([]); const [selectedModels, setSelectedModels] = useState([]); - const [data, setData] = useState([]); - const [cachedResponses, setCachedResponses] = useState("0"); - const [cachedTokens, setCachedTokens] = useState("0"); - const [cacheHitRatio, setCacheHitRatio] = useState("0"); const [dateValue, setDateValue] = useState({ from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000), @@ -113,120 +104,24 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole const [lastRefreshed, setLastRefreshed] = useState(""); const [healthCheckResponse, setHealthCheckResponse] = useState(""); - useEffect(() => { - if (!accessToken || !dateValue) { - return; - } - const fetchData = async () => { - const response = await adminGlobalCacheActivity( - accessToken, - formatDateWithoutTZ(dateValue.from), - formatDateWithoutTZ(dateValue.to), - ); - setData(response); - }; - fetchData(); - - const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleString()); - }, [accessToken]); - - const uniqueApiKeys = Array.from(new Set(data.map((item) => item?.api_key ?? ""))); - const uniqueModels = Array.from(new Set(data.map((item) => item?.model ?? ""))); - const uniqueCallTypes = Array.from(new Set(data.map((item) => item?.call_type ?? ""))); - - const updateCachingData = async (startTime: Date | undefined, endTime: Date | undefined) => { - if (!startTime || !endTime || !accessToken) { - return; - } - - let new_cache_data = await adminGlobalCacheActivity( - accessToken, - formatDateWithoutTZ(startTime), - formatDateWithoutTZ(endTime), - ); - - setData(new_cache_data); - }; + const { data: activity, refetch } = useCacheActivity({ + startDate: formatDateWithoutTZ(dateValue.from), + endDate: formatDateWithoutTZ(dateValue.to), + keyAliases: selectedApiKeys, + models: selectedModels, + }); useEffect(() => { - let newData: cacheDataItem[] = data; - if (selectedApiKeys.length > 0) { - newData = newData.filter((item) => selectedApiKeys.includes(item.api_key)); - } + setLastRefreshed(new Date().toLocaleString()); + }, []); - if (selectedModels.length > 0) { - newData = newData.filter((item) => selectedModels.includes(item.model)); - } - - /* - Data looks like this - [{"api_key":"sk-test-mock-key-001","call_type":"acompletion","model":"llama3-8b-8192","total_rows":13,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-002","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-123","call_type":"acompletion","model":"gpt-3.5-turbo","total_rows":19,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-123","call_type":"aimage_generation","model":"","total_rows":3,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-003","call_type":"None","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-004","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - {"api_key":"sk-test-mock-key-005","call_type":"","model":"chatgpt-v-2","total_rows":1,"cache_hit_true_rows":0}, - */ - - // What data we need for bar chat - // ui_data = [ - // { - // name: "Call Type", - // Cache hit: 20, - // LLM API requests: 10, - // } - // ] - - let llm_api_requests = 0; - let cache_hits = 0; - let cached_tokens = 0; - const processedData = newData.reduce((acc: uiData[], item) => { - if (!item.call_type) { - item.call_type = "Unknown"; - } - - llm_api_requests += (item.total_rows || 0) - (item.cache_hit_true_rows || 0); - cache_hits += item.cache_hit_true_rows || 0; - cached_tokens += item.cached_completion_tokens || 0; - - const existingItem = acc.find((i) => i.name === item.call_type); - if (existingItem) { - existingItem["LLM API requests"] += (item.total_rows || 0) - (item.cache_hit_true_rows || 0); - existingItem["Cache hit"] += item.cache_hit_true_rows || 0; - existingItem["Cached Completion Tokens"] += item.cached_completion_tokens || 0; - existingItem["Generated Completion Tokens"] += item.generated_completion_tokens || 0; - } else { - acc.push({ - name: item.call_type, - "LLM API requests": (item.total_rows || 0) - (item.cache_hit_true_rows || 0), - "Cache hit": item.cache_hit_true_rows || 0, - "Cached Completion Tokens": item.cached_completion_tokens || 0, - "Generated Completion Tokens": item.generated_completion_tokens || 0, - }); - } - return acc; - }, []); - - // set header cache statistics - setCachedResponses(valueFormatterNumbers(cache_hits)); - setCachedTokens(valueFormatterNumbers(cached_tokens)); - let allRequests = cache_hits + llm_api_requests; - if (allRequests > 0) { - let cache_hit_ratio = ((cache_hits / allRequests) * 100).toFixed(2); - setCacheHitRatio(cache_hit_ratio); - } else { - setCacheHitRatio("0"); - } - - setFilteredData(processedData); - }, [selectedApiKeys, selectedModels, dateValue, data]); + const uniqueApiKeys = activity?.filter_options.key_aliases ?? []; + const uniqueModels = activity?.filter_options.models ?? []; + const chartData = (activity?.groups ?? []).map(toChartDatum); const handleRefreshClick = () => { - // Update the 'lastRefreshed' state to the current date and time - const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleString()); + refetch(); + setLastRefreshed(new Date().toLocaleString()); }; const runCachingHealthCheck = async () => { @@ -257,10 +152,12 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole } }; + const totals = activity?.totals; + const hasRequests = totals != null && totals.api_requests + totals.cache_hits + totals.failed_requests > 0; const statCards = [ - { label: "Cache Hit Ratio", value: `${cacheHitRatio}%` }, - { label: "Cache Hits", value: cachedResponses }, - { label: "Cached Completion Tokens", value: cachedTokens }, + { label: "Cache Hit Ratio", value: `${hasRequests ? totals.cache_hit_ratio.toFixed(2) : "0"}%` }, + { label: "Cache Hits", value: valueFormatterNumbers(totals?.cache_hits ?? 0) }, + { label: "Cached Completion Tokens", value: valueFormatterNumbers(totals?.cached_completion_tokens ?? 0) }, ]; return ( @@ -380,7 +277,6 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole value={dateValue} onValueChange={(value) => { setDateValue(value); - updateCachingData(value.from, value.to); }} />
@@ -404,12 +300,12 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole @@ -423,7 +319,7 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole ({ + $api: { useQuery: (...args: unknown[]) => useQueryMock(...args) }, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +const params: CacheActivityParams = { + startDate: "2026-07-20", + endDate: "2026-07-27", + keyAliases: ["my-key"], + models: ["gpt-5.1"], +}; + +const lastCallOptions = (): { enabled: boolean } => { + const calls = useQueryMock.mock.calls; + return calls[calls.length - 1][3] as { enabled: boolean }; +}; + +describe("useCacheActivity", () => { + beforeEach(() => { + vi.clearAllMocks(); + useQueryMock.mockReturnValue({ data: undefined }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-access-token" }); + }); + + it("queries GET /global/activity/cache_hits with dates and filters as query params", () => { + renderHook(() => useCacheActivity(params)); + + expect(useQueryMock).toHaveBeenCalledWith( + "get", + "/global/activity/cache_hits", + { + params: { + query: { + start_date: "2026-07-20", + end_date: "2026-07-27", + key_aliases: ["my-key"], + models: ["gpt-5.1"], + }, + }, + }, + expect.any(Object), + ); + }); + + it("enables the query when authorized and both dates are set", () => { + renderHook(() => useCacheActivity(params)); + + expect(lastCallOptions().enabled).toBe(true); + }); + + it("disables the query without an access token", () => { + mockUseAuthorized.mockReturnValue({ accessToken: null }); + renderHook(() => useCacheActivity(params)); + + expect(lastCallOptions().enabled).toBe(false); + }); + + it("disables the query while the date range is incomplete", () => { + renderHook(() => useCacheActivity({ ...params, endDate: undefined })); + + expect(lastCallOptions().enabled).toBe(false); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts new file mode 100644 index 00000000000..af4486ad33b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/caching/useCacheActivity.ts @@ -0,0 +1,32 @@ +import { $api } from "@/lib/http/api"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import type { components } from "@/lib/http/schema"; + +export type CacheActivityResponse = components["schemas"]["CacheActivityResponse"]; +export type CacheActivityGroup = components["schemas"]["CacheActivityGroup"]; + +export interface CacheActivityParams { + startDate: string | undefined; + endDate: string | undefined; + keyAliases: string[]; + models: string[]; +} + +export const useCacheActivity = ({ startDate, endDate, keyAliases, models }: CacheActivityParams) => { + const { accessToken } = useAuthorized(); + return $api.useQuery( + "get", + "/global/activity/cache_hits", + { + params: { + query: { + start_date: startDate ?? "", + end_date: endDate ?? "", + key_aliases: keyAliases, + models, + }, + }, + }, + { enabled: Boolean(accessToken && startDate && endDate) }, + ); +}; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a15767d26cb..1331018c259 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2094,42 +2094,6 @@ export const adminGlobalActivity = async ( } }; -export const adminGlobalCacheActivity = async ( - accessToken: string, - startTime: string | undefined, - endTime: string | undefined, -) => { - try { - let url = proxyBaseUrl ? `${proxyBaseUrl}/global/activity/cache_hits` : `/global/activity/cache_hits`; - - if (startTime && endTime) { - url += `?start_date=${startTime}&end_date=${endTime}`; - } - - const requestOptions = { - method: "GET", - headers: { - [globalLitellmHeaderName]: `Bearer ${accessToken}`, - }, - }; - - const response = await fetch(url, requestOptions); - - if (!response.ok) { - const errorData = await response.json(); - const errorMessage = deriveErrorMessage(errorData); - handleError(errorMessage); - throw new Error(errorMessage); - } - - const data = await response.json(); - return data; - } catch (error) { - console.error("Failed to fetch spend data:", error); - throw error; - } -}; - export const adminGlobalActivityPerModel = async ( accessToken: string, startTime: string | undefined, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 79f03978b12..ed975c6be0a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -4356,25 +4356,9 @@ export interface paths { }; /** * Get Global Activity - * @description Get number of cache hits, vs misses - * - * { - * "daily_data": [ - * const chartdata = [ - * { - * date: 'Jan 22', - * cache_hits: 10, - * llm_api_calls: 2000 - * }, - * { - * date: 'Jan 23', - * cache_hits: 10, - * llm_api_calls: 12 - * }, - * ], - * "sum_cache_hits": 20, - * "sum_llm_api_calls": 2012 - * } + * @description Cache activity for the Admin UI cache dashboard, aggregated per call_type: + * cache hits vs successful LLM API requests vs failed requests, plus totals + * for the stat cards and the available key-alias/model filter options. */ get: operations["get_global_activity_global_activity_cache_hits_get"]; put?: never; @@ -21742,6 +21726,48 @@ export interface components { /** Total Requested */ total_requested: number; }; + /** CacheActivityFilterOptions */ + CacheActivityFilterOptions: { + /** Key Aliases */ + key_aliases: string[]; + /** Models */ + models: string[]; + }; + /** CacheActivityGroup */ + CacheActivityGroup: { + /** Api Requests */ + api_requests: number; + /** Cache Hits */ + cache_hits: number; + /** Cached Completion Tokens */ + cached_completion_tokens: number; + /** Call Type */ + call_type: string; + /** Failed Requests */ + failed_requests: number; + /** Generated Completion Tokens */ + generated_completion_tokens: number; + }; + /** CacheActivityResponse */ + CacheActivityResponse: { + filter_options: components["schemas"]["CacheActivityFilterOptions"]; + /** Groups */ + groups: components["schemas"]["CacheActivityGroup"][]; + totals: components["schemas"]["CacheActivityTotals"]; + }; + /** CacheActivityTotals */ + CacheActivityTotals: { + /** Api Requests */ + api_requests: number; + /** Cache Hit Ratio */ + cache_hit_ratio: number; + /** Cache Hits */ + cache_hits: number; + /** Cached Completion Tokens */ + cached_completion_tokens: number; + /** Failed Requests */ + failed_requests: number; + }; /** CachePingResponse */ CachePingResponse: { /** Cache Type */ @@ -40613,11 +40639,15 @@ export interface operations { }; get_global_activity_global_activity_cache_hits_get: { parameters: { - query?: { + query: { /** @description Time from which to start viewing spend */ - start_date?: string | null; + start_date: string; /** @description Time till which to view spend */ - end_date?: string | null; + end_date: string; + /** @description Only include spend from these key aliases */ + key_aliases?: string[] | null; + /** @description Only include spend for these models */ + models?: string[] | null; }; header?: never; path?: never; @@ -40631,7 +40661,7 @@ export interface operations { [name: string]: unknown; }; content: { - "application/json": components["schemas"]["LiteLLM_SpendLogs"][]; + "application/json": components["schemas"]["CacheActivityResponse"]; }; }; /** @description Validation Error */ From d17387e2e105bcbb42cd40333d9834cec38c8a44 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 09:59:34 -0700 Subject: [PATCH 33/54] fix(proxy): fall back on empty-string model_group in aggregated usage SQL --- .../management_endpoints/common_daily_activity.py | 8 ++++---- .../test_common_daily_activity.py | 15 ++++++++++----- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 9bd358289b3..8a5a31710cf 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -575,11 +575,11 @@ def _build_aggregated_sql_query( date, api_key, model, - COALESCE(model_group, model) AS model_group, + COALESCE(NULLIF(model_group, ''), model) AS model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint, - GROUPING(date, api_key, model, COALESCE(model_group, model), + GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model), custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level, SUM(spend)::float AS spend, @@ -600,8 +600,8 @@ def _build_aggregated_sql_query( (date, api_key), (date, model), (date, model, api_key), - (date, COALESCE(model_group, model)), - (date, COALESCE(model_group, model), api_key), + (date, COALESCE(NULLIF(model_group, ''), model)), + (date, COALESCE(NULLIF(model_group, ''), model), api_key), (date, custom_llm_provider), (date, custom_llm_provider, api_key), (date, mcp_namespaced_tool_name), diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 5aee7ff0236..c2a0d34a915 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -911,12 +911,15 @@ class TestBuildAggregatedSqlQuery: assert "api_key = $5" in sql def test_model_group_rollups_fall_back_to_model_name(self): - """Aggregated model_groups rollups must coalesce NULL model_group to model. + """Aggregated model_groups rollups must fall back to model for group-less rows. The (date, model_group) grouping level cannot recover the model column after the fact (it is rolled up), so the fallback has to happen in SQL; without it, group-less rows silently vanish from the model_groups - breakdown that the usage UI now renders by default. + breakdown that the usage UI now renders by default. Group-less rows are + stored as empty strings, not NULL (spend_tracking_utils defaults + model_group to ""), so a plain COALESCE is not enough: the fallback must + be NULLIF-wrapped to catch both """ sql, _ = _build_aggregated_sql_query( table_name="litellm_dailyuserspend", @@ -929,13 +932,15 @@ class TestBuildAggregatedSqlQuery: ) normalized = " ".join(sql.split()) - assert "COALESCE(model_group, model) AS model_group" in normalized + fallback = "COALESCE(NULLIF(model_group, ''), model)" + assert f"{fallback} AS model_group" in normalized assert ( - "GROUPING(date, api_key, model, COALESCE(model_group, model), " + f"GROUPING(date, api_key, model, {fallback}, " "custom_llm_provider, mcp_namespaced_tool_name, endpoint) AS group_level" in normalized ) - assert "(date, COALESCE(model_group, model)), (date, COALESCE(model_group, model), api_key)," in normalized + assert f"(date, {fallback}), (date, {fallback}, api_key)," in normalized assert "(date, model_group)" not in normalized + assert "COALESCE(model_group, model)" not in normalized @pytest.mark.asyncio From fdea50daa24eb98b8b4f27d649b281d52e25eb81 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 10:01:09 -0700 Subject: [PATCH 34/54] feat(ui): shareable log links via log_id query param on the logs page (#34879) * feat(ui): shareable log links via log_id query param on the logs page Clicking a log row now writes ?log_id= to the URL, closing the drawer removes it, and loading the logs page with ?log_id= opens the drawer for that log. When the log is not in the loaded page, it is fetched by request_id (the backend already drops the date window for id lookups), so links keep working for logs of any age. Drawer open state derives from the URL, mirroring the models page ?model= pattern. * fix(ui): close the log drawer on browser back after opening via session id Session opens now write ?session_id= to the URL instead of holding local state, so back removes both params and the drawer closes (Greptile P1). Session views become shareable links as a side effect. In-drawer log switching now replaces the history entry instead of pushing, so back always closes the drawer in one step rather than replaying every viewed log. * fix(proxy): scope /spend/logs/session/ui to the requesting user's visible logs Non-admin callers now only receive session rows they could already see on /spend/logs/ui: their own logs plus logs of teams where they hold the spend-logs permission. Previously any authenticated user could read any session's log metadata by id, which shareable ?session_id= links made trivial to trigger. Admin views are unchanged. Also, clicking a log row now clears a lingering ?session_id= from the URL so the drawer shows the clicked log instead of a stale session (Greptile P1). --- .../spend_management_endpoints.py | 38 ++- .../test_spend_management_endpoints.py | 149 +++++++++++- .../src/app/(dashboard)/navigateWithParams.ts | 8 +- .../view_logs/RequestLogsPanel.test.tsx | 226 +++++++++++++++++- .../components/view_logs/RequestLogsPanel.tsx | 89 +++++-- .../components/view_logs/logDetailRouting.ts | 60 +++++ 6 files changed, 529 insertions(+), 41 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 9aae9ca2875..42788227acc 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3308,8 +3308,36 @@ async def ui_view_session_spend_logs( detail="Database not connected", ) - # Build query conditions - where_conditions = {"session_id": session_id} + if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): + scope_sql = "" + scope_params = () + where_conditions = {"session_id": session_id} + else: + try: + permitted_team_ids = ( + await _get_permitted_team_ids_for_spend_logs( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) + else [] + ) + except Exception: # noqa: BLE001 # mirror /spend/logs/ui: failed team lookup falls back to own-logs-only scope + permitted_team_ids = [] + if permitted_team_ids: + scope_sql = ' AND ("user" = $4 OR team_id = ANY($5::text[]))' + scope_params = (user_api_key_dict.user_id, permitted_team_ids) + where_conditions = { + "session_id": session_id, + "OR": [ + {"user": user_api_key_dict.user_id}, + {"team_id": {"in": permitted_team_ids}}, + ], + } + else: + scope_sql = ' AND "user" = $4' + scope_params = (user_api_key_dict.user_id,) + where_conditions = {"session_id": session_id, "user": user_api_key_dict.user_id} # Calculate pagination offsets skip = (page - 1) * page_size @@ -3318,7 +3346,7 @@ async def ui_view_session_spend_logs( total_records = await SpendLogsRepository(prisma_client).table.count(where=where_conditions) # Query with raw SQL to exclude heavy columns (messages, response, proxy_server_request) - sql_query = """ + sql_query = f""" SELECT request_id, call_type, api_key, spend, total_tokens, prompt_tokens, completion_tokens, "startTime", "endTime", @@ -3328,11 +3356,11 @@ async def ui_view_session_spend_logs( organization_id, end_user, requester_ip_address, session_id, status, mcp_namespaced_tool_name, agent_id FROM "LiteLLM_SpendLogs" - WHERE session_id = $1 + WHERE session_id = $1{scope_sql} ORDER BY "startTime" DESC LIMIT $2 OFFSET $3 """ - result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip) + result = await prisma_client.db.query_raw(sql_query, session_id, page_size, skip, *scope_params) total_pages = (total_records + page_size - 1) // page_size diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 67945436987..71206687b5c 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -1548,6 +1548,7 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): assert page_size == 1 assert skip == 1 # page=2, page_size=1 assert 'ORDER BY "startTime" DESC' in sql_query + assert '"user" = $4' not in sql_query return [mock_spend_logs[0]] class MockPrismaClient: @@ -1558,20 +1559,144 @@ async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): mock_prisma_client = MockPrismaClient() monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - response = client.get( - "/spend/logs/session/ui", - params={"session_id": "session-123", "page": 2, "page_size": 1}, - headers={"Authorization": "Bearer sk-test"}, + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" ) - assert response.status_code == 200 - data = response.json() - assert data["total"] == 2 - assert data["page"] == 2 - assert data["page_size"] == 1 - assert data["total_pages"] == 2 - assert len(data["data"]) == 1 - assert data["data"][0]["request_id"] == "req1" + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 2, "page_size": 1}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 2 + assert data["page"] == 2 + assert data["page_size"] == 1 + assert data["total_pages"] == 2 + assert len(data["data"]) == 1 + assert data["data"][0]["request_id"] == "req1" + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, monkeypatch): + own_log = { + "id": "log1", + "request_id": "req1", + "session_id": "session-123", + "user": "user-1", + "startTime": "2024-01-01T00:00:00Z", + } + + class MockDB: + async def count(self, *args, **kwargs): + assert kwargs.get("where") == {"session_id": "session-123", "user": "user-1"} + return 1 + + async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user): + assert session_id == "session-123" + assert scoped_user == "user-1" + assert '"user" = $4' in sql_query + return [own_log] + + class MockPrismaClient: + def __init__(self): + self.db = MockDB() + self.db.litellm_spendlogs = self.db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + + async def no_permitted_teams(*args, **kwargs): + return [] + + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", + no_permitted_teams, + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" + ) + + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 1, "page_size": 50}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert [row["request_id"] for row in data["data"]] == ["req1"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, monkeypatch): + class MockDB: + async def count(self, *args, **kwargs): + assert kwargs.get("where") == { + "session_id": "session-123", + "OR": [ + {"user": "user-1"}, + {"team_id": {"in": ["team-9"]}}, + ], + } + return 1 + + async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user, team_ids): + assert session_id == "session-123" + assert scoped_user == "user-1" + assert team_ids == ["team-9"] + assert '("user" = $4 OR team_id = ANY($5::text[]))' in sql_query + return [ + { + "id": "log2", + "request_id": "req2", + "session_id": "session-123", + "team_id": "team-9", + "startTime": "2024-01-02T00:00:00Z", + } + ] + + class MockPrismaClient: + def __init__(self): + self.db = MockDB() + self.db.litellm_spendlogs = self.db + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + + async def permitted_teams(*args, **kwargs): + return ["team-9"] + + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", + permitted_teams, + ) + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" + ) + + try: + response = client.get( + "/spend/logs/session/ui", + params={"session_id": "session-123", "page": 1, "page_size": 50}, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["total"] == 1 + assert [row["request_id"] for row in data["data"]] == ["req2"] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts index c49c73c3578..5acf444a359 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/navigateWithParams.ts @@ -1,7 +1,11 @@ -export function navigateWithParams(mutate: (params: URLSearchParams) => void): void { +export function navigateWithParams(mutate: (params: URLSearchParams) => void, mode: "push" | "replace" = "push"): void { const params = new URLSearchParams(window.location.search); mutate(params); const qs = params.toString(); const url = qs ? `${window.location.pathname}?${qs}` : window.location.pathname; - window.history.pushState(null, "", url); + if (mode === "replace") { + window.history.replaceState(null, "", url); + } else { + window.history.pushState(null, "", url); + } } diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx index e8e2fae3f4d..49870e51c2b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx @@ -1,7 +1,7 @@ import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import moment from "moment"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterAll, beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; import type { LogEntry } from "./columns"; @@ -22,11 +22,75 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ })); vi.mock("./LogDetailsDrawer", () => ({ - LogDetailsDrawer: function LogDetailsDrawerMock({ open }: { open: boolean }) { - return
{open ? "open" : "closed"}
; + LogDetailsDrawer: function LogDetailsDrawerMock({ + open, + logEntry, + sessionId, + onClose, + allLogs = [], + onSelectLog, + }: { + open: boolean; + logEntry?: { request_id: string } | null; + sessionId?: string | null; + onClose: () => void; + allLogs?: { request_id: string }[]; + onSelectLog?: (log: { request_id: string }) => void; + }) { + const nextLog = allLogs.find((log) => log.request_id !== logEntry?.request_id); + return ( +
+ {open ? "open" : "closed"} + + +
+ ); }, })); +vi.mock("next/navigation", async (importOriginal) => { + const actual = await importOriginal(); + const { useSyncExternalStore } = await import("react"); + return { + ...actual, + useSearchParams: () => { + const search = useSyncExternalStore( + (onChange: () => void) => { + window.addEventListener("test-locationchange", onChange); + window.addEventListener("popstate", onChange); + return () => { + window.removeEventListener("test-locationchange", onChange); + window.removeEventListener("popstate", onChange); + }; + }, + () => window.location.search, + ); + return new URLSearchParams(search); + }, + }; +}); + +const originalPushState = window.history.pushState.bind(window.history); +const originalReplaceState = window.history.replaceState.bind(window.history); +beforeAll(() => { + window.history.pushState = (data, unused, url) => { + originalPushState(data, unused, url); + window.dispatchEvent(new Event("test-locationchange")); + }; + window.history.replaceState = (data, unused, url) => { + originalReplaceState(data, unused, url); + window.dispatchEvent(new Event("test-locationchange")); + }; +}); +afterAll(() => { + window.history.pushState = originalPushState; + window.history.replaceState = originalReplaceState; +}); + import { uiSpendLogsCall } from "../networking"; const logEntry = (overrides: Partial): LogEntry => ({ @@ -73,6 +137,7 @@ describe("RequestLogsPanel", () => { vi.clearAllMocks(); sessionStorage.clear(); testQueryClient.clear(); + window.history.replaceState(null, "", "/logs/"); respondWith([]); }); @@ -185,6 +250,161 @@ describe("RequestLogsPanel", () => { }); }); + describe("shareable log links (?log_id=)", () => { + const drawer = () => screen.getByTestId("log-details-drawer"); + + it("clicking a row writes ?log_id= to the URL and opens the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + + expect(new URLSearchParams(window.location.search).get("log_id")).toBe("req-1"); + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-1"); + }); + }); + + it("opens the drawer on load when ?log_id= matches a log in the loaded page", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-2"); + respondWith([logEntry({ request_id: "req-1" }), logEntry({ request_id: "req-2" })]); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-2"); + }); + }); + + it("fetches the log by request_id and opens the drawer when it is not in the loaded page", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-old"); + vi.mocked(uiSpendLogsCall).mockImplementation(async ({ params }) => + params?.request_id === "req-old" + ? { data: [logEntry({ request_id: "req-old" })], total: 1, page: 1, page_size: 1, total_pages: 1 } + : { data: [], total: 0, page: 1, page_size: 50, total_pages: 0 }, + ); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-old"); + }); + + const byIdCall = vi + .mocked(uiSpendLogsCall) + .mock.calls.find(([options]) => options.params?.request_id === "req-old")?.[0]; + if (!byIdCall) throw new Error("expected a by-id uiSpendLogsCall"); + expect(byIdCall.page).toBe(1); + expect(byIdCall.page_size).toBe(1); + }); + + it("closing the drawer removes ?log_id= from the URL and closes the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + await waitFor(() => expect(drawer()).toHaveTextContent("open")); + + await user.click(screen.getByRole("button", { name: "close-drawer" })); + + expect(new URLSearchParams(window.location.search).get("log_id")).toBeNull(); + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + }); + + it("switching logs inside the drawer replaces the URL, so back closes the drawer in one step", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-1" }), logEntry({ request_id: "req-2" })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-1")).not.toBeNull()); + await user.click(row("req-1") as HTMLElement); + await waitFor(() => expect(drawer()).toHaveAttribute("data-log-id", "req-1")); + + await user.click(screen.getByRole("button", { name: "select-next-log" })); + await waitFor(() => expect(drawer()).toHaveAttribute("data-log-id", "req-2")); + expect(new URLSearchParams(window.location.search).get("log_id")).toBe("req-2"); + + window.history.back(); + + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + expect(new URLSearchParams(window.location.search).get("log_id")).toBeNull(); + }); + + it("clicking a session id writes ?session_id= and ?log_id= and opens the session drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-solo", session_id: "sess-solo", session_total_count: 1 })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-solo")).not.toBeNull()); + await user.click(within(row("req-solo") as HTMLElement).getByText("sess-solo")); + + const params = new URLSearchParams(window.location.search); + expect(params.get("session_id")).toBe("sess-solo"); + expect(params.get("log_id")).toBe("req-solo"); + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-session-id", "sess-solo"); + }); + }); + + it("clicking a log row clears a lingering ?session_id= so the drawer shows the clicked log", async () => { + const user = userEvent.setup(); + respondWith([ + logEntry({ request_id: "req-a", session_id: "sess-a", session_total_count: 1 }), + logEntry({ request_id: "req-b" }), + ]); + renderWithProviders(); + + await waitFor(() => expect(row("req-a")).not.toBeNull()); + await user.click(within(row("req-a") as HTMLElement).getByText("sess-a")); + await waitFor(() => expect(new URLSearchParams(window.location.search).get("session_id")).toBe("sess-a")); + + await user.click(row("req-b") as HTMLElement); + + const params = new URLSearchParams(window.location.search); + expect(params.get("log_id")).toBe("req-b"); + expect(params.get("session_id")).toBeNull(); + await waitFor(() => { + expect(drawer()).toHaveAttribute("data-log-id", "req-b"); + expect(drawer()).toHaveAttribute("data-session-id", ""); + }); + }); + + it("browser back after opening via a session id closes the drawer", async () => { + const user = userEvent.setup(); + respondWith([logEntry({ request_id: "req-solo", session_id: "sess-solo", session_total_count: 1 })]); + renderWithProviders(); + + await waitFor(() => expect(row("req-solo")).not.toBeNull()); + await user.click(within(row("req-solo") as HTMLElement).getByText("sess-solo")); + await waitFor(() => expect(drawer()).toHaveTextContent("open")); + + window.history.back(); + + await waitFor(() => expect(drawer()).toHaveTextContent("closed")); + expect(new URLSearchParams(window.location.search).get("session_id")).toBeNull(); + }); + + it("opens a deep-linked multi-call session log in session mode", async () => { + window.history.replaceState(null, "", "/logs/?log_id=req-llm"); + respondWith([ + logEntry({ request_id: "req-llm", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }), + ]); + renderWithProviders(); + + await waitFor(() => { + expect(drawer()).toHaveTextContent("open"); + expect(drawer()).toHaveAttribute("data-log-id", "req-llm"); + expect(drawer()).toHaveAttribute("data-session-id", "sess-1"); + }); + }); + }); + describe("live tail", () => { it("shows the auto-refresh banner on the first page and hides it once stopped", async () => { const user = userEvent.setup(); diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx index b3b1e8c0640..06c8ca26a7e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx @@ -8,7 +8,7 @@ import { useCallback, useEffect, useMemo, useState } from "react"; import { AutoRouterModelGroupsProvider } from "@/components/shared/table_cells"; import { internalUserRoles } from "../../utils/roles"; import type { KeyResponse } from "../key_team_helpers/key_list"; -import { keyInfoV1Call } from "../networking"; +import { keyInfoV1Call, uiSpendLogsCall } from "../networking"; import KeyInfoView from "../templates/key_info_view"; import type { LogEntry } from "./columns"; import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; @@ -17,8 +17,10 @@ import { formatLogsWindow, getLogsWindowEndBound, LOG_FILTER_IDS, + type PaginatedResponse, useLogFilterLogic, } from "./log_filter_logic"; +import { useLogDetailRouting } from "./logDetailRouting"; import { LogDetailsDrawer } from "./LogDetailsDrawer"; import { LiveTailBanner, LogsTableToolbar } from "./LogsTableToolbar"; import { RequestLogsTable } from "./RequestLogsTable"; @@ -52,8 +54,15 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, const [selectedKeyIdInfoView, setSelectedKeyIdInfoView] = useState(null); const [selectedLog, setSelectedLog] = useState(null); - const [isDrawerOpen, setIsDrawerOpen] = useState(false); - const [selectedSessionId, setSelectedSessionId] = useState(null); + + const { + logId: urlLogId, + sessionId: urlSessionId, + openLog, + openSession, + selectLog, + close: closeUrlLog, + } = useLogDetailRouting(); const [isLiveTail, setIsLiveTail] = useState(() => { const storedValue = sessionStorage.getItem("isLiveTail"); @@ -106,6 +115,43 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, const { data: selectedKeyInfo } = useQuery(keyInfoQueryOptions); + const urlLogQueryOptions: UseQueryOptions = { + queryKey: ["logs", "byId", urlLogId, accessToken], + queryFn: async () => { + if (urlLogId === null) return null; + const window = formatLogsWindow(startTime, endTime, isCustomDate); + const response: PaginatedResponse = await uiSpendLogsCall({ + accessToken, + start_date: window.start_date, + end_date: window.end_date, + page: 1, + page_size: 1, + params: { request_id: urlLogId }, + }); + return response.data.find((log) => log.request_id === urlLogId) ?? null; + }, + enabled: urlLogId !== null && selectedLog?.request_id !== urlLogId, + staleTime: Infinity, + }; + + const { data: urlLog } = useQuery(urlLogQueryOptions); + + const displayLog = useMemo(() => { + if (urlLogId === null) return null; + if (selectedLog?.request_id === urlLogId) return selectedLog; + return filteredLogs.data.find((log) => log.request_id === urlLogId) ?? urlLog ?? null; + }, [urlLogId, selectedLog, filteredLogs.data, urlLog]); + + const displaySessionId = useMemo(() => { + if (urlSessionId !== null) return urlSessionId; + if (displayLog?.session_id !== undefined && (displayLog.session_total_count || 1) > 1) { + return displayLog.session_id; + } + return null; + }, [urlSessionId, displayLog]); + + const isDrawerOpen = displayLog !== null || displaySessionId !== null; + const rows = useMemo(() => { const searchedLogs = filteredLogs.data; @@ -186,22 +232,30 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, resetToFirstPage(); }, [resetToFirstPage]); - const handleRowClick = useCallback((log: LogEntry) => { - const isMultiCallSession = log.session_id !== undefined && (log.session_total_count || 1) > 1; - setSelectedSessionId(isMultiCallSession ? log.session_id ?? null : null); - setSelectedLog(log); - setIsDrawerOpen(true); - }, []); + const handleRowClick = useCallback( + (log: LogEntry) => { + setSelectedLog(log); + openLog(log.request_id); + }, + [openLog], + ); const handleSessionClick = useCallback( (sessionId: string) => { if (!sessionId) return; const log = rows.find((candidate) => candidate.session_id === sessionId) ?? null; - setSelectedSessionId(sessionId); setSelectedLog(log); - setIsDrawerOpen(true); + openSession(sessionId, log?.request_id ?? null); }, - [rows], + [rows, openSession], + ); + + const handleSelectLog = useCallback( + (log: LogEntry) => { + setSelectedLog(log); + selectLog(log.request_id); + }, + [selectLog], ); const handleKeyHashClick = useCallback((keyHash: string) => { @@ -267,15 +321,12 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, { - setIsDrawerOpen(false); - setSelectedSessionId(null); - }} - logEntry={selectedLog} - sessionId={selectedSessionId} + onClose={closeUrlLog} + logEntry={displayLog} + sessionId={displaySessionId} accessToken={accessToken} allLogs={rows} - onSelectLog={setSelectedLog} + onSelectLog={handleSelectLog} startTime={moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss")} /> diff --git a/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts b/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts new file mode 100644 index 00000000000..5b311c94627 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/logDetailRouting.ts @@ -0,0 +1,60 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "@/app/(dashboard)/navigateWithParams"; + +export const LOG_ID_QUERY_PARAM = "log_id"; +export const SESSION_ID_QUERY_PARAM = "session_id"; + +export interface LogDetailRouting { + logId: string | null; + sessionId: string | null; + openLog: (requestId: string) => void; + openSession: (sessionId: string, requestId: string | null) => void; + selectLog: (requestId: string) => void; + close: () => void; +} + +export function useLogDetailRouting(): LogDetailRouting { + const searchParams = useSearchParams(); + + const openLog = useCallback((requestId: string) => { + navigateWithParams((params) => { + params.set(LOG_ID_QUERY_PARAM, requestId); + params.delete(SESSION_ID_QUERY_PARAM); + }); + }, []); + + const openSession = useCallback((sessionId: string, requestId: string | null) => { + navigateWithParams((params) => { + params.set(SESSION_ID_QUERY_PARAM, sessionId); + if (requestId === null) { + params.delete(LOG_ID_QUERY_PARAM); + } else { + params.set(LOG_ID_QUERY_PARAM, requestId); + } + }); + }, []); + + const selectLog = useCallback((requestId: string) => { + navigateWithParams((params) => { + params.set(LOG_ID_QUERY_PARAM, requestId); + }, "replace"); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete(LOG_ID_QUERY_PARAM); + params.delete(SESSION_ID_QUERY_PARAM); + }); + }, []); + + return { + logId: searchParams?.get(LOG_ID_QUERY_PARAM) ?? null, + sessionId: searchParams?.get(SESSION_ID_QUERY_PARAM) ?? null, + openLog, + openSession, + selectLog, + close, + }; +} From 74244ddd4579eb1e81e66d0c53b2ed53e63cc8da Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Wed, 29 Jul 2026 10:59:53 -0700 Subject: [PATCH 35/54] Merge pull request #35041 from BerriAI/litellm_/ui-perf-regression-d7888c fix(ui): point the navbar and sidebar logos at the dashboard home route --- ui/litellm-dashboard/src/components/leftnav.test.tsx | 6 ++++++ ui/litellm-dashboard/src/components/leftnav.tsx | 2 +- ui/litellm-dashboard/src/components/navbar.test.tsx | 6 ++++++ ui/litellm-dashboard/src/components/navbar.tsx | 3 ++- 4 files changed, 15 insertions(+), 2 deletions(-) diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index 2c8fe52f97f..87c830e69de 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -80,6 +80,12 @@ describe("Sidebar (leftnav)", () => { collapsed: false, }; + it("should link the logo to the UI home route rather than the proxy origin", () => { + renderWithProviders(); + + expect(screen.getByRole("link", { name: /litellm home/i })).toHaveAttribute("href", "/ui"); + }); + it("renders all top-level (non-nested) tabs for admin", () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 4aa15f5fc28..af76fccf9eb 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -582,7 +582,7 @@ const Sidebar_: React.FC = ({
- + LiteLLM { expect(screen.getByRole("button", { name: /open account menu/i })).toBeInTheDocument(); }); + it("should link the logo to the UI home route rather than the proxy origin", () => { + renderWithProviders(); + + expect(screen.getByRole("link", { name: /litellm brand/i })).toHaveAttribute("href", "/ui"); + }); + it("should display user information in dropdown", async () => { const user = userEvent.setup(); renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/navbar.tsx b/ui/litellm-dashboard/src/components/navbar.tsx index 40638e7b8ba..c999ee8035e 100644 --- a/ui/litellm-dashboard/src/components/navbar.tsx +++ b/ui/litellm-dashboard/src/components/navbar.tsx @@ -3,6 +3,7 @@ import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBounci import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts"; import { useWorker } from "@/hooks/useWorker"; import { getProxyBaseUrl } from "@/components/networking"; +import { migratedHref } from "@/utils/migratedPages"; import { useTheme } from "@/contexts/ThemeContext"; import { clearTokenCookies } from "@/utils/cookieUtils"; import { clearStoredReturnUrl, getLoginUrl } from "@/utils/returnUrlUtils"; @@ -75,7 +76,7 @@ const Navbar: React.FC = ({ )}
- +
Date: Wed, 29 Jul 2026 11:35:06 -0700 Subject: [PATCH 36/54] feat(ui): deep link team detail page via ?team= query param (#35112) The teams page kept the selected team in React state, so a team detail page had no URL: it could not be shared, bookmarked, or opened from another page, and the browser back button dropped you out of the page instead of closing the detail view Adds useTeamDetailRouting reading ?team= (same pattern as the api-keys, models, and logs deep links) and derives the open team in Teams.tsx from the URL. TeamInfo now also derives team-admin rights from the fetched team data, so team admins arriving via a deep link are not stuck with a read-only view --- .../teams/detailNavigation.test.ts | 53 +++++++++++++++ .../app/(dashboard)/teams/detailNavigation.ts | 32 +++++++++ .../src/components/Teams.test.tsx | 67 +++++++++++++++++++ ui/litellm-dashboard/src/components/Teams.tsx | 11 +-- .../src/components/team/TeamInfo.test.tsx | 23 +++++++ .../src/components/team/TeamInfo.tsx | 10 ++- 6 files changed, 190 insertions(+), 6 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts new file mode 100644 index 00000000000..e5d5b1a4073 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.test.ts @@ -0,0 +1,53 @@ +/* @vitest-environment jsdom */ +import { act, renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useTeamDetailRouting } from "./detailNavigation"; + +vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) })); + +describe("useTeamDetailRouting", () => { + beforeEach(() => { + window.history.pushState(null, "", "/teams/"); + }); + + it("openTeam sets ?team= via history.pushState (no full navigation)", () => { + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.openTeam("team-abc123")); + expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("team=team-abc123")); + spy.mockRestore(); + }); + + it("openTeam preserves unrelated query params", () => { + window.history.pushState(null, "", "/teams/?foo=bar"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.openTeam("team-abc123")); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).toContain("team=team-abc123"); + spy.mockRestore(); + }); + + it("close removes only the team param", () => { + window.history.pushState(null, "", "/teams/?foo=bar&team=team-abc123"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useTeamDetailRouting()); + act(() => result.current.close()); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).not.toContain("team="); + spy.mockRestore(); + }); + + it("exposes teamId from ?team=", () => { + window.history.pushState(null, "", "/teams/?team=team-abc123"); + const { result } = renderHook(() => useTeamDetailRouting()); + expect(result.current.teamId).toBe("team-abc123"); + }); + + it("teamId is null when no team param is present", () => { + const { result } = renderHook(() => useTeamDetailRouting()); + expect(result.current.teamId).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts new file mode 100644 index 00000000000..d5208f094cb --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/teams/detailNavigation.ts @@ -0,0 +1,32 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "../navigateWithParams"; + +export interface TeamDetailRouting { + teamId: string | null; + openTeam: (id: string) => void; + close: () => void; +} + +export function useTeamDetailRouting(): TeamDetailRouting { + const searchParams = useSearchParams(); + + const openTeam = useCallback((id: string) => { + navigateWithParams((params) => { + params.set("team", id); + }); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete("team"); + }); + }, []); + + return { + teamId: searchParams?.get("team") ?? null, + openTeam, + close, + }; +} diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index 7065b1a5fb6..742a88d864c 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -72,6 +72,31 @@ vi.mock("@/components/team/TeamInfo", () => ({ }, })); +// The selected team is URL-derived (?team=) via useTeamDetailRouting. Next's real useSearchParams +// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate. +vi.mock("next/navigation", async () => { + const { useSyncExternalStore } = await import("react"); + const LOCATION_CHANGE_EVENT = "test-locationchange"; + for (const method of ["pushState", "replaceState"] as const) { + const original = window.history[method].bind(window.history); + window.history[method] = (...args: Parameters) => { + original(...args); + window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT)); + }; + } + const subscribe = (onChange: () => void) => { + window.addEventListener(LOCATION_CHANGE_EVENT, onChange); + window.addEventListener("popstate", onChange); + return () => { + window.removeEventListener(LOCATION_CHANGE_EVENT, onChange); + window.removeEventListener("popstate", onChange); + }; + }; + return { + useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)), + }; +}); + vi.mock("./ModelSelect/ModelSelect", () => { const ModelSelect = React.forwardRef(({ value, onChange, dataTestId, id }: any, ref: any) => { return ( @@ -159,6 +184,7 @@ const renderWithQueryClient = (component: React.ReactElement) => { // Re-establish safe defaults before every test (clearAllMocks keeps return values, so restore them here). beforeEach(() => { mockTeamsTableProps = null; + window.history.replaceState(null, "", "/teams/"); }); describe("Teams - handleCreate organization handling", () => { @@ -436,6 +462,47 @@ describe("Teams - premium props", () => { }); }); +describe("Teams - team detail deep link (?team=)", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockTeamInfoView.mockClear(); + vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue([]); + vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]); + vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] }); + mockUseOrganizations.mockReturnValue({ data: [] }); + }); + + it("selecting a team pushes ?team= to the URL", async () => { + renderWithQueryClient(); + + await waitFor(() => expect(mockTeamsTableProps).not.toBeNull()); + act(() => mockTeamsTableProps.onSelectTeam({ ...baseTableTeam, team_id: "team-deep-link" })); + + expect(window.location.search).toContain("team=team-deep-link"); + await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled()); + expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-deep-link" })); + }); + + it("opens the team detail view directly from a ?team= deep link", async () => { + window.history.replaceState(null, "", "/teams/?team=team-from-url"); + renderWithQueryClient(); + + await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled()); + expect(mockTeamInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ teamId: "team-from-url" })); + }); + + it("closing the team detail view removes ?team= from the URL", async () => { + window.history.replaceState(null, "", "/teams/?team=team-from-url"); + renderWithQueryClient(); + + await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled()); + act(() => mockTeamInfoView.mock.calls.at(-1)?.[0].onClose()); + + expect(window.location.search).not.toContain("team="); + await waitFor(() => expect(screen.queryByTestId("team-info-view")).not.toBeInTheDocument()); + }); +}); + describe("Teams - Create Team CTA is grouped with the tabs on the left", () => { beforeEach(() => { vi.clearAllMocks(); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index 20e9e78e7e4..0a3c7fc736d 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -12,6 +12,7 @@ import { useQueryClient } from "@tanstack/react-query"; import { PageHeader } from "@/components/shared/PageHeader"; import { Button as UIButton } from "@/components/ui/button"; import { teamsTableKeys } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { useTeamDetailRouting } from "@/app/(dashboard)/teams/detailNavigation"; import { TeamsTable } from "./TeamsPage/TeamsTable"; import AccessGroupSelector from "./common_components/AccessGroupSelector"; import PassThroughRoutesSelector from "./common_components/PassThroughRoutesSelector"; @@ -135,7 +136,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser const [editModalVisible, setEditModalVisible] = useState(false); const [selectedTeam, setSelectedTeam] = useState(null); - const [selectedTeamId, setSelectedTeamId] = useState(null); + const { teamId: selectedTeamId, openTeam, close: closeTeamDetail } = useTeamDetailRouting(); const [editTeam, setEditTeam] = useState(false); const [isTeamModalVisible, setIsTeamModalVisible] = useState(false); @@ -482,12 +483,12 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser userID={userID} onSelectTeam={(team) => { setSelectedTeam(team); - setSelectedTeamId(team.team_id); + openTeam(team.team_id); setEditTeam(false); }} onEditTeam={(team) => { setSelectedTeam(team); - setSelectedTeamId(team.team_id); + openTeam(team.team_id); setEditTeam(true); }} onDeleteTeam={handleDelete} @@ -547,11 +548,11 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser }} onClose={() => { setSelectedTeam(null); - setSelectedTeamId(null); + closeTeamDetail(); setEditTeam(false); }} accessToken={accessToken} - is_team_admin={is_team_admin(selectedTeam)} + is_team_admin={is_team_admin(selectedTeam?.team_id === selectedTeamId ? selectedTeam : null)} is_proxy_admin={userRole == "Admin"} userModels={userModels} editTeam={editTeam} diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index dcc72ccac9c..25365d26a12 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -444,6 +444,29 @@ describe("TeamInfoView", () => { }); }); + it("shows edit tabs when the fetched team data marks the session user as team admin, even without the is_team_admin prop", async () => { + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ + members_with_roles: [ + { + user_id: "user-1", + user_email: "admin@test.com", + role: "admin", + spend: 0, + budget_id: "budget1", + }, + ], + }), + ); + + renderWithProviders(); + + await waitFor(() => { + expect(screen.getByRole("tab", { name: "Settings" })).toBeInTheDocument(); + }); + expect(screen.getByRole("tab", { name: "Members" })).toBeInTheDocument(); + }); + it("should navigate to settings tab when clicked", async () => { const user = userEvent.setup({ delay: null }); vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData()); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index ff881a68938..acd5a8966a7 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -225,7 +225,15 @@ const TeamInfoView: React.FC = ({ return unfurlWildcardModelsInList(selected, userModels); }, [selectedModelsInForm, teamData, userModels]); - const canEditTeam = is_team_admin || is_proxy_admin || is_org_admin || isOrgAdminForTeam; + const isTeamAdminFromTeamData = useMemo( + () => + teamData?.team_info?.members_with_roles?.some( + (member) => member.user_id != null && member.user_id === userId && member.role === "admin", + ) ?? false, + [teamData, userId], + ); + + const canEditTeam = is_team_admin || is_proxy_admin || is_org_admin || isOrgAdminForTeam || isTeamAdminFromTeamData; const visibleTabs = useMemo(() => getTeamInfoVisibleTabs(canEditTeam), [canEditTeam]); const defaultTabKey = useMemo(() => getTeamInfoDefaultTab(editTeam, canEditTeam), [editTeam, canEditTeam]); From ba7d8ae17fe45dd355f2fc8f629e8822eb6bda57 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 11:56:37 -0700 Subject: [PATCH 37/54] feat(ui): deep link organization detail page via ?org= query param (#35117) * feat(ui): deep link organization detail page via ?org= query param The organizations page kept the selected organization in React state, so an org detail page had no URL: it could not be shared, bookmarked, or opened from another page, and the browser back button dropped you out of the page instead of closing the detail view Adds useOrgDetailRouting reading ?org= (same pattern as the api-keys, models, logs, and teams deep links) and derives the open organization in OrganizationsPanel from the URL * fix(ui): reset org edit mode on plain row selection and type test mocks Greptile P1: with the selected org now URL-derived, browser Back leaves the detail view without running onClose, so a stale editOrg=true made the next plain row click open on the Settings tab. Reset the flag on row selection, matching the teams page Greptile P2: type the panel test's captured table and detail-view props from the real components instead of any --- .../_components/OrganizationsPanel.test.tsx | 113 +++++++++++++++++- .../_components/OrganizationsPanel.tsx | 12 +- .../organizations/detailNavigation.test.ts | 53 ++++++++ .../organizations/detailNavigation.ts | 32 +++++ 4 files changed, 201 insertions(+), 9 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx index d381e5e65ca..3f9de478069 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx @@ -1,7 +1,9 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render, screen } from "@testing-library/react"; +import { act, render, screen } from "@testing-library/react"; import React from "react"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import type OrganizationsTableComponent from "./OrganizationsTable"; +import type OrganizationInfoViewComponent from "@/components/organization/organization_view"; vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ __esModule: true, @@ -18,12 +20,50 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ userRole: null, }), })); +type OrganizationsTableProps = React.ComponentProps; +type OrganizationInfoViewProps = React.ComponentProps; + +let capturedTableProps: OrganizationsTableProps | null = null; vi.mock("./OrganizationsTable", () => ({ __esModule: true, - default: (props: { isLoading: boolean }) => ( -
isLoading:{String(props.isLoading)}
- ), + default: (props: OrganizationsTableProps) => { + capturedTableProps = props; + return
isLoading:{String(props.isLoading)}
; + }, })); +const mockOrgInfoView = vi.fn<(props: OrganizationInfoViewProps) => void>(); +vi.mock("@/components/organization/organization_view", () => ({ + __esModule: true, + default: (props: OrganizationInfoViewProps) => { + mockOrgInfoView(props); + return
; + }, +})); + +// The selected org is URL-derived (?org=) via useOrgDetailRouting. Next's real useSearchParams +// re-renders subscribers on history.pushState/replaceState; mirror that so URL changes propagate. +vi.mock("next/navigation", async () => { + const { useSyncExternalStore } = await import("react"); + const LOCATION_CHANGE_EVENT = "test-locationchange"; + for (const method of ["pushState", "replaceState"] as const) { + const original = window.history[method].bind(window.history); + window.history[method] = (...args: Parameters) => { + original(...args); + window.dispatchEvent(new Event(LOCATION_CHANGE_EVENT)); + }; + } + const subscribe = (onChange: () => void) => { + window.addEventListener(LOCATION_CHANGE_EVENT, onChange); + window.addEventListener("popstate", onChange); + return () => { + window.removeEventListener(LOCATION_CHANGE_EVENT, onChange); + window.removeEventListener("popstate", onChange); + }; + }; + return { + useSearchParams: () => new URLSearchParams(useSyncExternalStore(subscribe, () => window.location.search)), + }; +}); import OrganizationsPanel from "./OrganizationsPanel"; @@ -34,6 +74,12 @@ const renderWithQueryClient = (ui: React.ReactElement) => { return render({ui}); }; +beforeEach(() => { + capturedTableProps = null; + mockOrgInfoView.mockClear(); + window.history.replaceState(null, "", "/organizations/"); +}); + describe("OrganizationsPanel", () => { it("gates non-premium users behind the enterprise notice", () => { renderWithQueryClient(); @@ -55,3 +101,60 @@ describe("OrganizationsPanel", () => { expect(screen.getByTestId("organizations-table")).toHaveTextContent("isLoading:false"); }); }); + +describe("OrganizationsPanel - org detail deep link (?org=)", () => { + it("clicking an organization pushes ?org= and opens the detail view", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onOrganizationClick("org-deep-link")); + + expect(window.location.search).toContain("org=org-deep-link"); + expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ organizationId: "org-deep-link" })); + }); + + it("opens the org detail directly from a ?org= deep link", () => { + window.history.replaceState(null, "", "/organizations/?org=org-from-url"); + renderWithQueryClient(); + + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-from-url", editOrg: false }), + ); + expect(screen.queryByTestId("organizations-table")).not.toBeInTheDocument(); + }); + + it("closing the org detail removes ?org= and returns to the list", () => { + window.history.replaceState(null, "", "/organizations/?org=org-from-url"); + renderWithQueryClient(); + + act(() => mockOrgInfoView.mock.calls.at(-1)?.[0].onClose()); + + expect(window.location.search).not.toContain("org="); + expect(screen.queryByTestId("organization-info-view")).not.toBeInTheDocument(); + expect(screen.getByTestId("organizations-table")).toBeInTheDocument(); + }); + + it("the edit action opens the detail in edit mode with ?org= set", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onEditClick("org-edit")); + + expect(window.location.search).toContain("org=org-edit"); + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-edit", editOrg: true }), + ); + }); + + it("a plain row click after leaving an edit view via browser history does not reopen in edit mode", () => { + renderWithQueryClient(); + + act(() => capturedTableProps?.onEditClick("org-edit")); + expect(mockOrgInfoView).toHaveBeenLastCalledWith(expect.objectContaining({ editOrg: true })); + + act(() => window.history.pushState(null, "", "/organizations/")); + act(() => capturedTableProps?.onOrganizationClick("org-plain")); + + expect(mockOrgInfoView).toHaveBeenLastCalledWith( + expect.objectContaining({ organizationId: "org-plain", editOrg: false }), + ); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx index b1c026d3904..a21c0669677 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx @@ -1,5 +1,6 @@ import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels"; +import { useOrgDetailRouting } from "@/app/(dashboard)/organizations/detailNavigation"; import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters"; import { useQueryClient } from "@tanstack/react-query"; import React, { useState } from "react"; @@ -19,7 +20,7 @@ interface OrganizationsPanelProps { } const OrganizationsPanel: React.FC = ({ userRole, accessToken, premiumUser }) => { - const [selectedOrgId, setSelectedOrgId] = useState(null); + const { orgId: selectedOrgId, openOrg, close: closeOrgDetail } = useOrgDetailRouting(); const [editOrg, setEditOrg] = useState(false); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [orgToDelete, setOrgToDelete] = useState(null); @@ -108,7 +109,7 @@ const OrganizationsPanel: React.FC = ({ userRole, acces { - setSelectedOrgId(null); + closeOrgDetail(); setEditOrg(false); }} accessToken={accessToken} @@ -132,9 +133,12 @@ const OrganizationsPanel: React.FC = ({ userRole, acces isLoading={isLoading} userRole={userRole} searchActive={searchActive} - onOrganizationClick={setSelectedOrgId} + onOrganizationClick={(organizationId) => { + setEditOrg(false); + openOrg(organizationId); + }} onEditClick={(organizationId) => { - setSelectedOrgId(organizationId); + openOrg(organizationId); setEditOrg(true); }} onDeleteClick={handleDelete} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts new file mode 100644 index 00000000000..46b7c4313ea --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.test.ts @@ -0,0 +1,53 @@ +/* @vitest-environment jsdom */ +import { act, renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useOrgDetailRouting } from "./detailNavigation"; + +vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) })); + +describe("useOrgDetailRouting", () => { + beforeEach(() => { + window.history.pushState(null, "", "/organizations/"); + }); + + it("openOrg sets ?org= via history.pushState (no full navigation)", () => { + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.openOrg("org-abc123")); + expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("org=org-abc123")); + spy.mockRestore(); + }); + + it("openOrg preserves unrelated query params", () => { + window.history.pushState(null, "", "/organizations/?foo=bar"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.openOrg("org-abc123")); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).toContain("org=org-abc123"); + spy.mockRestore(); + }); + + it("close removes only the org param", () => { + window.history.pushState(null, "", "/organizations/?foo=bar&org=org-abc123"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useOrgDetailRouting()); + act(() => result.current.close()); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("foo=bar"); + expect(url).not.toContain("org="); + spy.mockRestore(); + }); + + it("exposes orgId from ?org=", () => { + window.history.pushState(null, "", "/organizations/?org=org-abc123"); + const { result } = renderHook(() => useOrgDetailRouting()); + expect(result.current.orgId).toBe("org-abc123"); + }); + + it("orgId is null when no org param is present", () => { + const { result } = renderHook(() => useOrgDetailRouting()); + expect(result.current.orgId).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts new file mode 100644 index 00000000000..8c55c7b750c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/detailNavigation.ts @@ -0,0 +1,32 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +import { navigateWithParams } from "../navigateWithParams"; + +export interface OrgDetailRouting { + orgId: string | null; + openOrg: (id: string) => void; + close: () => void; +} + +export function useOrgDetailRouting(): OrgDetailRouting { + const searchParams = useSearchParams(); + + const openOrg = useCallback((id: string) => { + navigateWithParams((params) => { + params.set("org", id); + }); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete("org"); + }); + }, []); + + return { + orgId: searchParams?.get("org") ?? null, + openOrg, + close, + }; +} From bf5334bc59faf9d7a35bc75fa3d0ed9d8e20f344 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 12:32:55 -0700 Subject: [PATCH 38/54] fix(router): drop duplicate Mapping import that fails ruff F811 (#35122) router.py imports Mapping from collections.abc and again from typing, which ruff flags as a redefinition and fails the lint CI job on every open PR. Keep the collections.abc import --- litellm/router.py | 1 - 1 file changed, 1 deletion(-) diff --git a/litellm/router.py b/litellm/router.py index 6336b12258b..69535d7c74a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -31,7 +31,6 @@ from typing import ( Generator, List, Literal, - Mapping, Optional, Set, Tuple, From ea783cc35c91d1d6d989421e8cc045f23156feae Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 29 Jul 2026 13:41:31 -0700 Subject: [PATCH 39/54] refactor(rust): make litellm-core the callable messages() SDK; drop the ai-gateway handler (#35044) 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-rust/ADDING_A_PROVIDER.md | 11 ++-- litellm-rust/AGENTS.md | 26 ++++++-- litellm-rust/CLAUDE.md | 61 +++++++++++++------ litellm-rust/README.md | 41 ++++++++----- .../PROVIDER_CODING_STANDARDS.md | 10 +-- litellm-rust/crates/ai-gateway/AGENTS.md | 9 ++- litellm-rust/crates/ai-gateway/README.md | 10 +-- .../crates/ai-gateway/src/constants.rs | 16 ----- .../crates/ai-gateway/src/io/messages.rs | 1 - litellm-rust/crates/ai-gateway/src/io/mod.rs | 1 - litellm-rust/crates/ai-gateway/src/lib.rs | 5 +- .../crates/ai-gateway/src/messages/mod.rs | 49 --------------- .../crates/ai-gateway/src/messages/types.rs | 24 -------- .../crates/ai-gateway/src/routes/AGENTS.md | 7 ++- .../ai-gateway/src/routes/messages/service.rs | 23 +++---- litellm-rust/crates/core/AGENTS.md | 8 ++- litellm-rust/crates/core/CLAUDE.md | 31 ++++++++-- litellm-rust/crates/core/Cargo.toml | 2 +- litellm-rust/crates/core/src/constants.rs | 16 +++++ .../src/messages/client.rs | 0 .../src/messages/common_utils.rs | 10 +-- .../src/messages/handler.rs | 17 ++---- litellm-rust/crates/core/src/messages/mod.rs | 30 +++++++++ .../src/messages/prepare.rs | 7 +-- .../src/messages/tests.rs | 14 +++-- .../crates/core/src/messages/types.rs | 24 ++++++++ litellm-rust/crates/python-bridge/AGENTS.md | 4 +- litellm-rust/crates/python-bridge/CLAUDE.md | 6 +- litellm-rust/crates/python-bridge/src/lib.rs | 18 ++++-- 29 files changed, 277 insertions(+), 204 deletions(-) delete mode 100644 litellm-rust/crates/ai-gateway/src/io/messages.rs delete mode 100644 litellm-rust/crates/ai-gateway/src/messages/mod.rs delete mode 100644 litellm-rust/crates/ai-gateway/src/messages/types.rs rename litellm-rust/crates/{ai-gateway => core}/src/messages/client.rs (100%) rename litellm-rust/crates/{ai-gateway => core}/src/messages/common_utils.rs (83%) rename litellm-rust/crates/{ai-gateway => core}/src/messages/handler.rs (84%) rename litellm-rust/crates/{ai-gateway => core}/src/messages/prepare.rs (92%) rename litellm-rust/crates/{ai-gateway => core}/src/messages/tests.rs (97%) diff --git a/litellm-rust/ADDING_A_PROVIDER.md b/litellm-rust/ADDING_A_PROVIDER.md index 5f933ec4fa8..857a744e014 100644 --- a/litellm-rust/ADDING_A_PROVIDER.md +++ b/litellm-rust/ADDING_A_PROVIDER.md @@ -1,10 +1,11 @@ # Adding a provider / route to litellm-rust -Three layers, same for every route (see `ocr` and `realtime` as references): +Everything for a route lives in `crates/core/src//`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint. -1. **Transform contract (pure)** — `crates/core/src//transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth. -2. **Provider config (pure)** — `crates/providers/src///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests. -3. **HTTP / transport (the host)** — `crates/providers/src/.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O. +1. **Entrypoint** — `mod.rs`: `pub async fn (request) -> CoreResult`, the Rust equivalent of `litellm.()`, plus a `_stream` variant when the route streams. It is the only thing a host touches. +2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`. +3. **Provider config** — `crates/core/src/providers///transformation.rs`: implement that trait as a `const __CONFIG`, mirroring the Python provider tree. Add parity unit tests. +4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response. ## Coding standards @@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider is a few declarative lines, not a new file of duplicated flow. Only diverge from the base when behavior is genuinely different, and say so explicitly in the PR. -**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`. +**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`. diff --git a/litellm-rust/AGENTS.md b/litellm-rust/AGENTS.md index 398eec4685c..36a5ad5a8f4 100644 --- a/litellm-rust/AGENTS.md +++ b/litellm-rust/AGENTS.md @@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes ( ## Crates -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| Crate | Role | +|-------|------| +| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. | +| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. | +| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. +## Where a route lives + +A top-level LiteLLM call is a module under `crates/core/src//`, shaped like `messages`: + +``` +core/src/messages/ + mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE) + types.rs # request/response types, MessagesRequest + transformation.rs # the provider template trait + prepare.rs # provider resolution, auth headers, URL + handler.rs # the provider call + client.rs # the shared reqwest client +``` + +Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched. + Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these. Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional. diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 0659e63df39..fe6ceedbb86 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR. ## Crates (exactly three — see AGENTS.md) -`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge` -exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates. +`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call. +`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and +`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not +a route — add modules, not crates. ## Core Boundary -`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work. +`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()` +is `litellm_core::messages::messages(request).await`: you call it, it does the +provider call, and you get a typed non-streaming response back. Route-level Rust structure mirrors LiteLLM's Python responsibilities: -- `core/src//` owns the route contract, shared types, and provider - template traits. For OCR, this means `core/src/ocr`. +- `core/src//` owns the route end to end: the public entrypoint fn named + after the route in `mod.rs`, the request/response types (`types.rs`), the + provider template trait (`transformation.rs`), the provider/auth/URL + resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that + performs the call (`handler.rs`). `core/src/messages` is the reference. - `core/src/providers///transformation.rs` owns the - provider-specific transform. For Mistral OCR, this means - `core/src/providers/mistral/ocr/transformation.rs`. -- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`), - never inside `core`. + provider-specific transform. For Anthropic Messages, this means + `core/src/providers/anthropic/messages/transformation.rs`. +- Handlers live in `core`, never in a host. `ai-gateway` must not contain a + route handler that talks to a provider; its axum route reads the HTTP request, + picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals + Python objects and calls the same entrypoint. + +Streaming keeps the same shape: the route entrypoint has a `_stream` +variant in `core` that returns the upstream response so a host can splice it to +its own caller; the host still owns no provider logic. Call-hook and lifecycle instrumentation, including phase timing, usage accumulation, and callback payload construction, always lives in `core`. @@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through their I/O logger; hosts must not own callback orchestration. Allowed in `core`: -- Pure request transforms -- Pure response transforms -- Pure stream chunk normalization +- The public entrypoint for a top-level LiteLLM call +- Request/response transforms and stream chunk normalization +- Provider resolution, auth header construction, and URL building +- The provider HTTP call itself, through a shared reused client with connect and + request timeouts - Shared data types and validation errors - Deterministic token/cost helper logic Not allowed in `core`: -- Network calls -- Environment variable or secret reads +- Serving HTTP: axum routes, extractors, and transport concerns stay in the host - Filesystem access -- Database or cache access -- Provider SDK signing or auth flows +- Database access +- Config file reading and rollout state - Logging callbacks, spend writes, or custom callbacks - Global mutable runtime state +Env reads in `core` are limited to credential fallback inside a route's +`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when +no key is passed. Everything else config-shaped is resolved by the host and +passed in. + +Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`) +predate this rule and are being moved into `core` route modules; do not add new +ones there, and prefer moving one when you touch it. + Python owns rollout state and fallback while Rust is being introduced. Rust paths must be off by default until parity tests prove equivalence with Python. A new provider/route may instead be implemented rust-only with no Python @@ -93,10 +116,10 @@ the first PR: - Preserve Python output shape intentionally. If a field is always serialized as `null` for Python parity, leave a short comment explaining that parity choice. -## Host I/O Rules +## Network I/O Rules -These rules apply when adding future crates or modules that execute network I/O, -such as `ai-gateway`, router hosts, or standalone servers: +These rules apply to every module that executes network I/O, whether it is a +`core` route handler or a host such as `ai-gateway`: - Set connect and full-request timeouts. No unbounded waits. - Reuse HTTP clients; do not construct clients per request. diff --git a/litellm-rust/README.md b/litellm-rust/README.md index 1646c90ad76..bcccf93300b 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -2,18 +2,31 @@ This workspace contains the staged Rust implementation for LiteLLM. -Rust starts as a pure transform core used by the existing Python host. Python -continues to own auth, configuration, network I/O, retries, routing, logging, +`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call +that makes the LLM call and hands back a typed response, the same shape as +`litellm.messages()` in Python. + +```rust +let response = litellm_core::messages::messages(MessagesRequest { + model: "claude-sonnet-4-5", + body, + api_key: Some(key), + .. +}) +.await?; +``` + +Python continues to own configuration, retries, routing policy, logging, callbacks, spend tracking, and customer plugins until each Rust path has parity coverage and production evidence. ## Crates -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| Crate | Role | +|-------|------| +| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. | +| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. | +| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. @@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm- ```text crates/ - core/ Route contracts, shared pure types, errors, and templates. - src/ocr/ - providers/ Provider-specific pure transforms. - src/mistral/ocr/transformation.rs + core/ The SDK: route modules + provider transforms. + src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client + src/providers/anthropic/messages/transformation.rs + ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints. python-bridge/ PyO3 bridge for Python LiteLLM. ``` -The folder shape should follow the Python provider tree: -`providers/src///transformation.rs`. The bridge should expose -one function per top-level route, starting with `ocr(payload)`. +The folder shape follows the Python provider tree: +`core/src/providers///transformation.rs`. The bridge exposes one +function per top-level route, mirroring the core entrypoints. ## Checks diff --git a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md index ed44dc4c729..a1860d8a9c9 100644 --- a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md +++ b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md @@ -1,6 +1,6 @@ # Provider coding standards (litellm-rust) -Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port. +Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response. ## Provider resolution @@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST ## Boundaries -7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter. +7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request. 8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers///`; a route is a module, never a new crate. -9. Route entry point stays thin: `()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them. -10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`. +9. Route entry point stays thin: `core::::()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them. +10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`. ## Types and errors @@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST 16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary. 17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer. -18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS. +18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS. ## Tests and rollout diff --git a/litellm-rust/crates/ai-gateway/AGENTS.md b/litellm-rust/crates/ai-gateway/AGENTS.md index d9e6e1adde5..92567091cd3 100644 --- a/litellm-rust/crates/ai-gateway/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/AGENTS.md @@ -1,7 +1,9 @@ # ai-gateway — folder architecture The Axum server that fronts the Rust gateway. It owns transport + config + auth -only; deployment selection lives in `core::router`, transforms in `core`/`providers`. +only; deployment selection lives in `core::router`, and the LLM call itself +(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint +such as `litellm_core::messages::messages`. No provider handler lives here. ``` src/ @@ -32,6 +34,11 @@ src/ args; it runs during extraction. Never re-implement the check per route. - **Handlers are thin.** A handler validates and delegates to its `service`. No business logic, no provider calls, no transforms in handlers. +- **Services call `core`, they don't reimplement it.** A `service` picks the + deployment and calls the `core` route entrypoint. Provider resolution, auth + headers, URL building, and the HTTP call are `core`'s job; a service that + builds a provider request itself is a bug (`routes/messages/service.rs` is + the reference). - **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in `state.rs`; read env/config only in `main.rs` when building state. diff --git a/litellm-rust/crates/ai-gateway/README.md b/litellm-rust/crates/ai-gateway/README.md index f913beff6d5..7a6c620ee84 100644 --- a/litellm-rust/crates/ai-gateway/README.md +++ b/litellm-rust/crates/ai-gateway/README.md @@ -8,11 +8,11 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame. `litellm-rust` is exactly three crates (a crate is a **layer**, not a route): -| Crate | Role | Pure / I/O | -|-------|------|------------| -| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. Builds requests/responses; no network. | Pure | -| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under `io/`) plus the Axum server binary (behind the `server` feature). | I/O | -| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding | +| Crate | Role | +|-------|------| +| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. | +| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. | +| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. | Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge. diff --git a/litellm-rust/crates/ai-gateway/src/constants.rs b/litellm-rust/crates/ai-gateway/src/constants.rs index 74808cf1ce6..78af374bf70 100644 --- a/litellm-rust/crates/ai-gateway/src/constants.rs +++ b/litellm-rust/crates/ai-gateway/src/constants.rs @@ -29,18 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500; #[cfg(feature = "server")] pub(crate) const DEFAULT_PROVIDER: &str = "openai"; -/// Full-request timeout ceiling for Anthropic Messages provider calls, in -/// seconds. Mirrors the Python Anthropic Messages default. The per-request -/// timeout from `litellm_params` still overrides this on the request builder. -pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; - -/// Connect timeout for Anthropic Messages provider calls, in seconds. -pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; - -/// Max characters of an upstream error body echoed across the host boundary -/// before truncation, so provider bodies are bounded and data-minimized. -pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; - pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10; pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; @@ -48,10 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300; #[cfg(feature = "server")] pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages"; -/// Provider name used by the Anthropic Messages route when a deployment's -/// provider model does not carry an explicit provider prefix. -pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; - /// Request headers owned by the gateway and never forwarded upstream. #[cfg(feature = "server")] pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] = diff --git a/litellm-rust/crates/ai-gateway/src/io/messages.rs b/litellm-rust/crates/ai-gateway/src/io/messages.rs deleted file mode 100644 index 86170e45678..00000000000 --- a/litellm-rust/crates/ai-gateway/src/io/messages.rs +++ /dev/null @@ -1 +0,0 @@ -pub use crate::messages::{MessagesRequest, messages}; diff --git a/litellm-rust/crates/ai-gateway/src/io/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index 6129a808965..cce56dd2121 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -1,5 +1,4 @@ pub mod audio_transcription; -pub mod messages; pub mod ocr; pub mod realtime; pub mod realtime_pool; diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index c44d661c29e..057db6457c4 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -4,7 +4,9 @@ //! without pulling in the HTTP server: //! //! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks, -//! and provider I/O. Always available — no feature required. +//! and provider I/O. Always available — no feature required. These predate the +//! rule that a route's entrypoint and handler live in `litellm-core` (see +//! `litellm_core::messages`) and move there as they are touched. //! - [`io`]: compatibility exports and realtime WebSocket splice helpers. //! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling //! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway` @@ -14,7 +16,6 @@ pub mod audio_transcription; mod client; pub mod io; -pub mod messages; pub mod ocr; /// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs deleted file mode 100644 index fd2dd546941..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ /dev/null @@ -1,49 +0,0 @@ -use litellm_core::CoreResult; -use serde_json::Value; - -mod client; -mod common_utils; -mod handler; -mod prepare; -mod types; - -pub use types::MessagesRequest; - -use handler::{execute_messages_provider_call, execute_messages_provider_stream}; -use prepare::prepare_messages_call; - -pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { - match execute_messages(request, false).await? { - MessagesResponse::Json(body) => Ok(body), - MessagesResponse::Stream(response) => { - drop(response); - Err(litellm_core::CoreError::InvalidResponse( - "non-streaming messages execution returned a stream".to_string(), - )) - } - } -} - -pub(crate) enum MessagesResponse { - Json(Value), - Stream(reqwest::Response), -} - -pub(crate) async fn execute_messages( - request: MessagesRequest<'_>, - stream: bool, -) -> CoreResult { - let prepared = prepare_messages_call(request)?; - if stream { - execute_messages_provider_stream(prepared) - .await - .map(MessagesResponse::Stream) - } else { - execute_messages_provider_call(prepared) - .await - .map(MessagesResponse::Json) - } -} - -#[cfg(test)] -mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs deleted file mode 100644 index 848fadb4b02..00000000000 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ /dev/null @@ -1,24 +0,0 @@ -use std::time::Duration; - -use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; -use serde_json::{Map, Value}; - -pub struct MessagesRequest<'a> { - pub model: &'a str, - pub body: Value, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, - pub extra_headers: Option>, - pub timeout: Option, -} - -pub(crate) struct ProviderMessagesRequest { - pub(crate) provider: String, - pub(crate) model: String, - pub(crate) config: &'static dyn AnthropicMessagesProviderConfig, - pub(crate) url: String, - pub(crate) body: Value, - pub(crate) upstream_headers: Vec<(String, String)>, - pub(crate) timeout: Option, -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md index 02c5f18c4f3..3eee43e7a2f 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md +++ b/litellm-rust/crates/ai-gateway/src/routes/AGENTS.md @@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... } When a route has business logic worth testing without axum, put it in a sibling `service` (a file, or a folder if the route grows). The route file stays the **axum surface** (router + handler + any socket/SSE adapter); `service` is plain -Rust with **no axum types**. `realtime/` is the example: +Rust with **no axum types**, and its job is to pick the deployment and call the +`core` route entrypoint (see `messages/service.rs` calling +`litellm_core::messages::messages`). Never build a provider request, resolve a +key, or perform the provider call here. `realtime/` is the older example: ``` realtime/ mod.rs # axum surface: router() + handler + the WS<->events adapter @@ -33,6 +36,8 @@ genuinely gets hard to read. `crate::auth::RequireMasterKey` to its arguments; it runs during extraction. Never re-implement the check per route. - **Handlers contain no business logic; `service` contains no axum types.** +- **No provider handlers in this crate.** Transforms, auth headers, and the + provider HTTP call live in `core/src//`. - A route owns its paths in its own `router()`; `mod.rs` only merges. - Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`, not duplicated in handlers. diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 75ed26e5be8..5f4c5fe8de4 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -1,12 +1,12 @@ use std::sync::Arc; +use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER; +use litellm_core::messages::types::MessagesRequest; +use litellm_core::messages::{messages, messages_stream}; use litellm_core::router::Router; use litellm_core::{CoreError, CoreResult}; use serde_json::{Map, Value}; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -use crate::messages::{MessagesRequest, execute_messages}; - pub(crate) enum MessagesResponse { Json(Value), Stream(reqwest::Response), @@ -52,13 +52,14 @@ pub async fn run( extra_headers, timeout: None, }; - let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true); - execute_messages(request, stream) - .await - .map(|response| match response { - crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body), - crate::messages::MessagesResponse::Stream(upstream) => { - MessagesResponse::Stream(upstream) - } + if request.body.get("stream").and_then(Value::as_bool) == Some(true) { + return messages_stream(request).await.map(MessagesResponse::Stream); + } + + let response = messages(request).await?; + serde_json::to_value(response) + .map(MessagesResponse::Json) + .map_err(|err| { + CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) }) } diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 8740dccaf01..aee8b4937ef 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -1,3 +1,7 @@ -litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads. +litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src//` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back. -Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates. +A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate. + +Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`. + +Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates. diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md index 20873878967..5d36305ded5 100644 --- a/litellm-rust/crates/core/CLAUDE.md +++ b/litellm-rust/crates/core/CLAUDE.md @@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`. ## Responsibility -`core` owns shared data types, typed errors, and deterministic helper contracts. -It must stay pure and host-independent. +`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level +LiteLLM call has a public entrypoint here, named after the route +(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and +calling it returns a typed non-streaming response. Allowed: +- The public entrypoint for a route, plus its `_stream` variant when the + route supports streaming. +- Provider resolution, auth header construction, URL building, and the provider + HTTP call (shared reused client, connect + request timeouts). - Shared request/response structs. - Typed errors with stable, non-sensitive messages. - Deterministic validation helpers. - Serialization helpers that intentionally mirror Python output shape. - Route templates that match Python base config responsibilities, such as - `ocr::transformation::OcrProviderConfig`. + `messages::transformation::AnthropicMessagesProviderConfig`. Not allowed: -- Network, filesystem, database, cache, or environment access. -- Secret reads or auth/header construction. +- Serving HTTP: axum routers, extractors, and other transport concerns. +- Filesystem, database, or cache access. +- Config file reading or rollout state; the host resolves those and passes them + in. Env reads are limited to credential fallback in a route's `prepare.rs`. - Logging callbacks, tracing spans, spend writes, or customer callbacks. - Provider-specific branching that belongs in `providers`. - Panics for user/provider-controlled input. @@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API. ## Structure -Use route names directly under `src/`: `ocr`, future `messages`, +Use route names directly under `src/`: `messages`, `ocr`, future `chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not invent broad names like `engine` for route contracts. +`src/messages` is the reference shape for a route module: + +``` +mod.rs pub async fn messages(..) (+ messages_stream) +types.rs request/response types +transformation.rs the provider template trait +prepare.rs provider resolution, auth headers, URL +handler.rs the provider call +client.rs the shared reqwest client +``` + ## Parity Rules - Every shared type used by a provider transform needs unit tests for diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 65c6db7412c..ab8050734f2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true [dependencies] rand.workspace = true +reqwest.workspace = true serde.workspace = true serde_json.workspace = true thiserror.workspace = true @@ -30,5 +31,4 @@ bedrock-auth = [ ] [dev-dependencies] -reqwest.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 5826a5bc9c1..caada1d98b0 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -1,3 +1,19 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1"; pub const OPENAI_RESPONSES_PATH: &str = "/responses"; + +/// Full-request timeout ceiling for Anthropic Messages provider calls, in +/// seconds. Mirrors the Python Anthropic Messages default. The per-request +/// timeout from the caller still overrides this on the request builder. +pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; + +/// Connect timeout for Anthropic Messages provider calls, in seconds. +pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; + +/// Max characters of an upstream error body echoed across the call boundary +/// before truncation, so provider bodies are bounded and data-minimized. +pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; + +/// Provider name used for Anthropic Messages when a deployment's provider model +/// does not carry an explicit provider prefix. +pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; diff --git a/litellm-rust/crates/ai-gateway/src/messages/client.rs b/litellm-rust/crates/core/src/messages/client.rs similarity index 100% rename from litellm-rust/crates/ai-gateway/src/messages/client.rs rename to litellm-rust/crates/core/src/messages/client.rs diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs similarity index 83% rename from litellm-rust/crates/ai-gateway/src/messages/common_utils.rs rename to litellm-rust/crates/core/src/messages/common_utils.rs index 68ecc3f17c1..9dcfcaa71e3 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,11 +1,11 @@ -use litellm_core::CoreResult; -use litellm_core::error::{CoreError, json_type_name}; -use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; -use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; -use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; use serde_json::{Map, Value}; use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; +use crate::error::{CoreError, CoreResult, json_type_name}; +use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; +use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; + +use super::transformation::AnthropicMessagesProviderConfig; pub(super) fn truncate_error_body(body: &str) -> String { if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS { diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs similarity index 84% rename from litellm-rust/crates/ai-gateway/src/messages/handler.rs rename to litellm-rust/crates/core/src/messages/handler.rs index 90c12367f50..1c895f66eba 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,15 +1,13 @@ -use litellm_core::CoreResult; -use litellm_core::error::CoreError; -use serde_json::Value; +use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +use crate::error::{CoreError, CoreResult}; use super::client::http_client; use super::common_utils::truncate_error_body; -use super::types::ProviderMessagesRequest; -use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest}; pub(super) async fn execute_messages_provider_call( request: ProviderMessagesRequest, -) -> CoreResult { +) -> CoreResult { let mut request_builder = http_client().post(&request.url).json(&request.body); for (key, value) in &request.upstream_headers { request_builder = request_builder.header(key, value); @@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call( let response = serde_json::from_str(&text).map_err(|err| { CoreError::InvalidResponse(format!("invalid messages response JSON: {err}")) })?; - let transformed = request - .config - .transform_response(&request.model, response)?; - serde_json::to_value(transformed).map_err(|err| { - CoreError::InvalidResponse(format!("failed to serialize messages response: {err}")) - }) + request.config.transform_response(&request.model, response) } pub(super) async fn execute_messages_provider_stream( diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index ec2fbb969a6..acb36d89daf 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -1,2 +1,32 @@ +//! The Anthropic Messages call, the Rust equivalent of Python's +//! `litellm.messages()`. +//! +//! [`messages`] is the top-level entrypoint: give it a model, a body, and +//! credentials, and it resolves the provider, transforms the request, calls the +//! provider, and returns a typed non-streaming response. [`messages_stream`] +//! is the streaming variant; it hands the raw upstream response back so a host +//! can splice the event stream to its own caller. + +mod client; +mod common_utils; +mod handler; +mod prepare; pub mod transformation; pub mod types; + +use crate::error::CoreResult; + +use handler::{execute_messages_provider_call, execute_messages_provider_stream}; +use prepare::prepare_messages_call; +use types::{AnthropicMessagesResponse, MessagesRequest}; + +pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { + execute_messages_provider_call(prepare_messages_call(request)?).await +} + +pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult { + execute_messages_provider_stream(prepare_messages_call(request)?).await +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs similarity index 92% rename from litellm-rust/crates/ai-gateway/src/messages/prepare.rs rename to litellm-rust/crates/core/src/messages/prepare.rs index 9a027490eb6..94b5b1eaed7 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,9 +1,8 @@ -use litellm_core::CoreError; -use litellm_core::CoreResult; -use litellm_core::messages::transformation::MessagesAuthStrategy; -use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; +use crate::error::{CoreError, CoreResult}; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; +use super::transformation::MessagesAuthStrategy; use super::types::{MessagesRequest, ProviderMessagesRequest}; pub(super) fn prepare_messages_call( diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs similarity index 97% rename from litellm-rust/crates/ai-gateway/src/messages/tests.rs rename to litellm-rust/crates/core/src/messages/tests.rs index 23a53e98045..9fc1763683b 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -1,14 +1,16 @@ use std::time::Duration; -use litellm_core::error::CoreError; use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; +use crate::error::CoreError; + use super::common_utils::{ has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; -use super::{MessagesRequest, messages}; +use super::messages; +use super::types::MessagesRequest; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); @@ -152,8 +154,8 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() .await .expect("messages request succeeds"); - assert_eq!(response["content"][0]["text"], "hi"); - assert_eq!(response["stop_reason"], "end_turn"); + assert_eq!(response.content[0]["text"], "hi"); + assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); let request = server.await.expect("server task completes"); let (head, body) = request.split_once("\r\n\r\n").expect("has body"); @@ -208,8 +210,8 @@ async fn messages_round_trip_builds_native_anthropic_request() { .await .expect("messages request succeeds"); - assert_eq!(response["content"][0]["text"], "hi"); - assert_eq!(response["stop_reason"], "end_turn"); + assert_eq!(response.content[0]["text"], "hi"); + assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); let request = server.await.expect("server task completes"); let (head, _) = request.split_once("\r\n\r\n").expect("has body"); diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 11fe17ea40f..b9f807c29fd 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,6 +1,30 @@ +use std::time::Duration; + use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use super::transformation::AnthropicMessagesProviderConfig; + +pub struct MessagesRequest<'a> { + pub model: &'a str, + pub body: Value, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, + pub timeout: Option, +} + +pub(super) struct ProviderMessagesRequest { + pub(super) provider: String, + pub(super) model: String, + pub(super) config: &'static dyn AnthropicMessagesProviderConfig, + pub(super) url: String, + pub(super) body: Value, + pub(super) upstream_headers: Vec<(String, String)>, + pub(super) timeout: Option, +} + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(untagged)] pub enum SystemPrompt { diff --git a/litellm-rust/crates/python-bridge/AGENTS.md b/litellm-rust/crates/python-bridge/AGENTS.md index d6d3d90e6ab..ad3cddfa5fd 100644 --- a/litellm-rust/crates/python-bridge/AGENTS.md +++ b/litellm-rust/crates/python-bridge/AGENTS.md @@ -1,3 +1,3 @@ -litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over litellm-ai-gateway. +litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`). -Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call into litellm-ai-gateway. +Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint. diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md index e5d021ec25b..3ce8b8c639a 100644 --- a/litellm-rust/crates/python-bridge/CLAUDE.md +++ b/litellm-rust/crates/python-bridge/CLAUDE.md @@ -11,11 +11,11 @@ Python-compatible dictionaries. ## Bridge Shape - Prefer one stable method per top-level LiteLLM route, for example - `ocr(payload)`. + `messages(...)`, calling the matching `litellm-core` entrypoint. - Do not add one exported PyO3 function per provider helper unless there is a measured reason. -- Provider dispatch belongs in Rust route modules such as - `litellm_providers::ocr`, not in this PyO3 crate. +- Provider dispatch belongs in the `litellm-core` route module (e.g. + `litellm_core::messages`), not in this PyO3 crate. - Python owns rollout state and fallback. Rust should return errors; Python decides whether to raise or fall back. For a rust-only provider/route (no Python reference), the Python side is a thin dispatch that calls Rust and diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index ee9bdd0b81f..f0cc26a0cca 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -4,10 +4,11 @@ use std::time::Duration; use litellm_ai_gateway::io::audio_transcription::{ AudioTranscriptionRequest, audio_transcription as run_audio_transcription, }; -use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages}; use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use litellm_core::error::CoreError; +use litellm_core::messages::messages as run_messages; +use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; use pyo3::exceptions::{PyRuntimeError, PyValueError}; use pyo3::prelude::*; use pyo3::types::{PyAny, PyDict}; @@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult> { Ok(json.call_method1("loads", (encoded,))?.unbind()) } +fn messages_response_to_py( + py: Python<'_>, + response: AnthropicMessagesResponse, +) -> PyResult> { + let value = + serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?; + json_to_py(py, value) +} + fn core_error_to_pyerr(err: CoreError) -> PyErr { match err { CoreError::Auth(message) => PyValueError::new_err(message), @@ -382,7 +392,7 @@ fn messages( }); match result { - Ok(value) => json_to_py(py, value), + Ok(response) => messages_response_to_py(py, response), Err(err) => Err(core_error_to_pyerr(err)), } } @@ -404,7 +414,7 @@ fn amessages( marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; pyo3_async_runtimes::tokio::future_into_py(py, async move { - let value = run_messages(MessagesRequest { + let response = run_messages(MessagesRequest { model: &model, body, api_key: api_key.as_deref(), @@ -416,7 +426,7 @@ fn amessages( .await .map_err(core_error_to_pyerr)?; - Python::attach(|py| json_to_py(py, value)) + Python::attach(|py| messages_response_to_py(py, response)) }) } From 440b1bcf654d637967d6036cca269a3db7e4cb9f Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 29 Jul 2026 13:43:33 -0700 Subject: [PATCH 40/54] fix(otel): make OTLP export work against Grafana Cloud (#35060) Three defects kept LiteLLM's OTel metrics from reaching an OTLP backend. OTEL_EXPORTER_OTLP_HEADERS is W3C Baggage encoded per the OTLP spec, so its values are percent-encoded. litellm split the string on "," and "=" and passed the raw value straight to the exporter, so a vendor that documents "Authorization=Basic%20" got a literal "%20" on the wire and the backend rejected the credential. Grafana Cloud documents exactly that shape, which made its OTLP gateway unreachable. Header parsing now delegates to the OTel SDK's own W3C Baggage parser in liberal mode, so percent-encoded values decode and values that were never encoded keep working. It moves from model/utils.py to plumbing/providers.py because model/ is deliberately free of opentelemetry imports; providers.parse_headers was already the entry point every caller used. The OTLP metric exporters then overrode histogram temporality to delta. Prometheus and Mimir, which back Grafana Cloud's OTLP gateway, reject delta histograms outright: the gateway answers 400 "invalid temporality and type combination" and drops the entire batch, so every GenAI metric was silently lost while traces kept flowing. Backends that prefer delta still accept cumulative, so the SDK default is the compatible choice in both directions, and the enterprise billing exporter already relies on it. Three GenAI instruments also carried names no convention or backend defines, so nothing downstream could chart them. Time to first token and time per output token take their semconv names, gen_ai.server.time_to_first_token and gen_ai.server.time_per_output_token; the gen_ai.client.response.* spellings litellm used are not conventions at all. Cost has no semconv instrument, so it takes gen_ai.usage.cost, the name backends already query for spend. All three are listed verbatim in Grafana Cloud's AI Observability integration reference, so its prebuilt panels find them. Both engines now read the names from the shared Metric constants rather than repeating string literals, so v1 and v2 cannot drift. The renames are breaking for anyone charting the former names; the docs and the release changelog carry the migration note. --- litellm/integrations/opentelemetry.py | 21 +++++----- litellm/integrations/otel/model/semconv.py | 22 ++++++++-- litellm/integrations/otel/model/utils.py | 22 +++------- .../integrations/otel/plumbing/providers.py | 35 +++++++++++----- .../otel/test_otel_v2_components.py | 42 +++++++++++++++++++ .../integrations/otel/test_otel_v2_logger.py | 6 +-- .../integrations/otel/test_otel_v2_metrics.py | 6 +-- .../integrations/test_opentelemetry.py | 34 +++++++++++++++ 8 files changed, 142 insertions(+), 46 deletions(-) diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 12465377b51..11e7ab062b6 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -27,6 +27,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import ( OTELSemconvCategory, parse_semconv_opt_in, ) +from litellm.integrations.otel.model.semconv import Metric from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.secret_managers.main import get_secret_bool, str_to_bool @@ -597,32 +598,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): meter = meter_provider.get_meter(__name__) self._operation_duration_histogram = meter.create_histogram( - name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38 + name=Metric.OPERATION_DURATION, description="GenAI operation duration", unit="s", ) self._token_usage_histogram = meter.create_histogram( - name="gen_ai.client.token.usage", # Replace with semconv constant in otel 1.38 + name=Metric.TOKEN_USAGE, description="GenAI token usage", unit="{token}", ) self._cost_histogram = meter.create_histogram( - name="gen_ai.client.token.cost", + name=Metric.TOKEN_COST, description="GenAI request cost", unit="USD", ) self._time_to_first_token_histogram = meter.create_histogram( - name="gen_ai.client.response.time_to_first_token", + name=Metric.TIME_TO_FIRST_TOKEN, description="Time to first token for streaming requests", unit="s", ) self._time_per_output_token_histogram = meter.create_histogram( - name="gen_ai.client.response.time_per_output_token", + name=Metric.TIME_PER_OUTPUT_TOKEN, description="Average time per output token (generation time / completion tokens)", unit="s", ) self._response_duration_histogram = meter.create_histogram( - name="gen_ai.client.response.duration", + name=Metric.RESPONSE_DURATION, description="Total LLM API generation time (excludes LiteLLM overhead)", unit="s", ) @@ -2980,10 +2981,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def _get_metric_reader(self): """ Get the appropriate metric reader based on the configuration. + + Histograms keep the SDK's default cumulative temporality: Prometheus-backed + OTLP receivers reject delta histograms and drop the whole batch, while + backends that prefer delta still accept cumulative. """ - from opentelemetry.sdk.metrics import Histogram from opentelemetry.sdk.metrics.export import ( - AggregationTemporality, ConsoleMetricExporter, PeriodicExportingMetricReader, ) @@ -3014,7 +3017,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): exporter = OTLPMetricExporter( endpoint=normalized_endpoint, headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) @@ -3032,7 +3034,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): exporter = OTLPMetricExporter( endpoint=normalized_endpoint, headers=_split_otel_headers, - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) return PeriodicExportingMetricReader(exporter, export_interval_millis=5000) diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 44b2f7e0488..1abe8ca33fa 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -257,13 +257,27 @@ class LiteLLM: class Metric: - """GenAI metric instrument names.""" + """GenAI metric instrument names. + + Every name here that a convention or a backend defines uses that name, so a + consumer charting GenAI telemetry finds litellm's series where it looks for + them. ``TOKEN_USAGE``, ``OPERATION_DURATION``, ``TIME_TO_FIRST_TOKEN`` and + ``TIME_PER_OUTPUT_TOKEN`` are semconv instruments, defined in the GenAI + conventions; the ``gen_ai.client.response.*`` spellings litellm used for the + latter two are not conventions at all, so nothing downstream could chart + them. Cost has no semconv instrument, so it takes ``gen_ai.usage.cost``, the + name backends already query for spend. + + ``RESPONSE_DURATION`` keeps its vendor spelling deliberately: the closest + convention, ``gen_ai.server.request.duration``, would collide in meaning with + ``OPERATION_DURATION``, which litellm already emits for the whole operation. + """ TOKEN_USAGE: Final = "gen_ai.client.token.usage" OPERATION_DURATION: Final = "gen_ai.client.operation.duration" - TOKEN_COST: Final = "gen_ai.client.token.cost" - TIME_TO_FIRST_TOKEN: Final = "gen_ai.client.response.time_to_first_token" - TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.client.response.time_per_output_token" + TOKEN_COST: Final = "gen_ai.usage.cost" + TIME_TO_FIRST_TOKEN: Final = "gen_ai.server.time_to_first_token" + TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.server.time_per_output_token" RESPONSE_DURATION: Final = "gen_ai.client.response.duration" diff --git a/litellm/integrations/otel/model/utils.py b/litellm/integrations/otel/model/utils.py index f37afc97879..ab54a558a9a 100644 --- a/litellm/integrations/otel/model/utils.py +++ b/litellm/integrations/otel/model/utils.py @@ -1,9 +1,11 @@ """Shared, OpenTelemetry-free helpers for the otel integration. -Generic value coercion (for reading heterogeneous logging-payload dicts), time -conversion, and header parsing — pulled out of the individual modules so they -live in one place. Deliberately free of any ``opentelemetry`` import so the -OTel-free sources of truth (payloads, semconv, spans, config) can use it too. +Generic value coercion (for reading heterogeneous logging-payload dicts) and +time conversion — pulled out of the individual modules so they live in one +place. Deliberately free of any ``opentelemetry`` import so the OTel-free +sources of truth (payloads, semconv, spans, config) can use it too. OTLP header +parsing lives in :mod:`litellm.integrations.otel.plumbing.providers` instead, +because it delegates to the OTel SDK's own W3C Baggage parser. """ from datetime import datetime @@ -89,15 +91,3 @@ def to_seconds(value: datetime | float | int | str | None) -> float | None: except ValueError: continue return None - - -def parse_headers(raw: str | None) -> dict[str, str]: - """Parse an OTLP ``"k=v,k=v"`` header string into a dict.""" - headers: dict[str, str] = {} - if not raw: - return headers - for pair in raw.split(","): - if "=" in pair: - key, _, value = pair.partition("=") - headers[key.strip()] = value.strip() - return headers diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index ced65aa1ec3..ede9acc6d8f 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( InMemorySpanExporter, ) from opentelemetry.trace import Span, SpanKind, Tracer +from opentelemetry.util.re import parse_env_headers from litellm._version import version as litellm_version from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config from litellm.integrations.otel.model.semconv import LiteLLM from litellm.integrations.otel.model.spans import LiteLLMSpanKind -# Re-exported so ``providers.parse_headers`` remains a stable entry point. -from litellm.integrations.otel.model.utils import parse_headers as parse_headers - if TYPE_CHECKING: from opentelemetry.metrics import Meter from opentelemetry.sdk.metrics.export import MetricReader @@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None: return endpoint + "/v1/traces" +def parse_headers(raw: str | None) -> dict[str, str]: + """Parse an OTLP ``"k=v,k=v"`` header string into a dict. + + ``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so + values are percent-decoded: a vendor that documents + ``Authorization=Basic%20`` (Grafana Cloud does, because a bare space + is not representable there) has to reach the exporter as ``Basic ``, + not with a literal ``%20`` that the backend rejects as malformed. The SDK's + own parser is used so litellm decodes exactly what the OTLP exporters do + when they read the env var themselves; ``liberal`` keeps values that are not + percent-encoded (``Authorization=Bearer ``) working unchanged. + """ + if not raw: + return {} + return dict(parse_env_headers(raw, liberal=True)) + + def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter: kind = (spec.kind or "console").lower() factory = _EXPORTER_FACTORIES.get(kind) @@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": ``console`` (and any unrecognized kind) exports to the console; ``otlp_http`` and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The reader exports on a 5s period, matching v1. + + Histograms keep the SDK's default cumulative temporality. Prometheus-backed + OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint) + reject delta histograms outright with ``invalid temporality and type + combination``, which drops the whole metric batch, while backends that + prefer delta still accept cumulative. The enterprise billing exporter + already relies on the same default. """ from opentelemetry.sdk.metrics.export import ( ConsoleMetricExporter, @@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( OTLPMetricExporter as HTTPMetricExporter, ) - from opentelemetry.sdk.metrics import Histogram - from opentelemetry.sdk.metrics.export import AggregationTemporality exporter: Any = HTTPMetricExporter( endpoint=_otlp_metrics_endpoint(config.endpoint), headers=parse_headers(config.headers), - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) elif kind in ("otlp_grpc", "grpc"): - from opentelemetry.sdk.metrics import Histogram - from opentelemetry.sdk.metrics.export import AggregationTemporality - try: from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import ( OTLPMetricExporter as GRPCMetricExporter, @@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader": exporter = GRPCMetricExporter( endpoint=config.endpoint, headers=parse_headers(config.headers), - preferred_temporality={Histogram: AggregationTemporality.DELTA}, ) else: exporter = ConsoleMetricExporter() 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 5191414edeb..d856d6871a3 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_components.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_components.py @@ -451,6 +451,30 @@ def test_parse_headers(): assert providers.parse_headers("no-equals") == {} +def test_parse_headers_percent_decodes_values(): + """A percent-encoded OTLP header value reaches the exporter decoded. + + ``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded, and Grafana Cloud + documents ``Authorization=Basic%20``. Forwarding the literal ``%20`` + makes the backend reject the export as a malformed credential. + """ + token = "MTMzNzc4MzpnbGNfZXlKdklqb2lNVEl6TkNJPQ==" + assert providers.parse_headers(f"Authorization=Basic%20{token}") == {"authorization": f"Basic {token}"} + assert providers.parse_headers("x-scope-orgid=team%20a") == {"x-scope-orgid": "team a"} + + +def test_parse_headers_keeps_unencoded_values_working(): + """Values that are not percent-encoded keep parsing unchanged. + + Vendors that document a bare space, and litellm's own presets, must survive + the switch to the spec-compliant parser. Base64 padding also means a value + can contain ``=``, so only the first one may split the pair. + """ + assert providers.parse_headers("Authorization=Bearer sk-123") == {"authorization": "Bearer sk-123"} + assert providers.parse_headers("api_key=abc,space_id=xyz") == {"api_key": "abc", "space_id": "xyz"} + assert providers.parse_headers("api_key=YWJjZA==") == {"api_key": "YWJjZA=="} + + def test_otlp_traces_endpoint_normalization(): norm = providers._otlp_traces_endpoint # A base endpoint gets the signal path appended (the common OTLP env shape). @@ -487,6 +511,24 @@ def test_build_span_exporter_variants(): assert "OTLPSpanExporter" in type(http_exporter).__name__ +def test_otlp_metric_exporter_uses_cumulative_histogram_temporality(): + """Histograms must export as cumulative, not delta. + + Prometheus-backed OTLP receivers (Grafana Cloud / Mimir) reject delta + histograms with ``invalid temporality and type combination`` and drop the + entire metric batch, so a delta default silently loses every GenAI metric. + """ + from opentelemetry.sdk.metrics import Histogram + from opentelemetry.sdk.metrics.export import AggregationTemporality + + reader = providers.build_metric_reader( + OpenTelemetryV2Config(exporter="otlp_http", endpoint="http://h:4318") + ) + temporality = reader._exporter._preferred_temporality # noqa: SLF001 # exporter exposes no public accessor + + assert temporality[Histogram] is AggregationTemporality.CUMULATIVE + + def test_otlp_logs_endpoint_normalization(): norm = providers._otlp_logs_endpoint # A base endpoint gets the signal path appended (the common OTLP env shape). diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 02954578644..41c02501acc 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -2120,9 +2120,9 @@ def test_valid_metric_filter_records_six_metrics(monkeypatch): assert _emitted_metric_names(reader) == { "gen_ai.client.operation.duration", "gen_ai.client.token.usage", - "gen_ai.client.token.cost", - "gen_ai.client.response.time_to_first_token", - "gen_ai.client.response.time_per_output_token", + "gen_ai.usage.cost", + "gen_ai.server.time_to_first_token", + "gen_ai.server.time_per_output_token", "gen_ai.client.response.duration", } diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py index 29067f91b5a..b2d89053ba3 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_metrics.py @@ -39,9 +39,9 @@ from litellm.integrations.otel.plumbing.providers import ( # noqa: E402 OPERATION_DURATION = "gen_ai.client.operation.duration" TOKEN_USAGE = "gen_ai.client.token.usage" -TOKEN_COST = "gen_ai.client.token.cost" -TIME_TO_FIRST_TOKEN = "gen_ai.client.response.time_to_first_token" -TIME_PER_OUTPUT_TOKEN = "gen_ai.client.response.time_per_output_token" +TOKEN_COST = "gen_ai.usage.cost" +TIME_TO_FIRST_TOKEN = "gen_ai.server.time_to_first_token" +TIME_PER_OUTPUT_TOKEN = "gen_ai.server.time_per_output_token" RESPONSE_DURATION = "gen_ai.client.response.duration" ALL_METRICS = frozenset( diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 7ffd09b931f..05205cb76f2 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -412,6 +412,40 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): current_provider is existing_provider ), "Existing TracerProvider should be respected and not overridden" + @patch.dict( + os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True + ) + def test_init_metrics_creates_instruments_under_their_published_names(self): + """ + The v1 engine's instrument names are a public contract. + + Every name here is what a backend queries: four are GenAI semantic + conventions and gen_ai.usage.cost is the name backends query for spend. + A rename is breaking for anyone charting them, so it has to be a + deliberate edit to the shared Metric constants and to this list, never + a silent drift between the v1 and v2 engines. + """ + from opentelemetry import metrics + + metrics.set_meter_provider(MeterProvider(metric_readers=[InMemoryMetricReader()])) + otel_integration = OpenTelemetry(config=OpenTelemetryConfig.from_env()) + + assert { + otel_integration._operation_duration_histogram.name, + otel_integration._token_usage_histogram.name, + otel_integration._cost_histogram.name, + otel_integration._time_to_first_token_histogram.name, + otel_integration._time_per_output_token_histogram.name, + otel_integration._response_duration_histogram.name, + } == { + "gen_ai.client.operation.duration", + "gen_ai.client.token.usage", + "gen_ai.usage.cost", + "gen_ai.server.time_to_first_token", + "gen_ai.server.time_per_output_token", + "gen_ai.client.response.duration", + } + @patch.dict( os.environ, {"LITELLM_OTEL_INTEGRATION_ENABLE_METRICS": "true"}, clear=True ) From 5182dfa66bd0c8c8afb3e56f7cf9ac2ec87c6ee5 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Wed, 29 Jul 2026 14:20:39 -0700 Subject: [PATCH 41/54] test(e2e): remove the Presidio guardrail suite (#35129) Drops tests/e2e/guardrails/test_presidio_guardrail_e2e.py and the PresidioParamsBody it was the only caller of. Both cases were red on most stage runs between 07-25 and 07-29: pre_call failed 6 of 11 runs, post_call 6 of 11, with post_call reporting the raw address reaching the caller while apply_to_output was set. The cause was propagation, not masking. GuardrailsClient.register() posts /guardrails and returns immediately with no readiness wait, unlike ProxyClient._await_model_servable or GuardrailsClient._await_team, and the data plane only picks a new guardrail up on its next periodic DB sync. Calls issued before that sync pass the raw value through. #34833 has since made both cases poll to the deadline, and on the current build each masks on the first attempt, so the suite is expected to be green now; it is being removed because it spends real provider money on every retry and because a pod replaced mid-poll still reproduces the old failure. The three guardrail.presidio.* rows stay in coverage_registry/guardrail.yaml and go uncovered on purpose, so Presidio reads as a tier-P0 gap in Grafana rather than dropping out of the denominator. --- tests/e2e/guardrails/guardrails_client.py | 11 -- .../guardrails/test_presidio_guardrail_e2e.py | 141 ------------------ 2 files changed, 152 deletions(-) delete mode 100644 tests/e2e/guardrails/test_presidio_guardrail_e2e.py diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index 5a54a4f0bbc..93861d19922 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -63,16 +63,6 @@ class OpenAIModerationParamsBody(GuardrailParamsBase): model: str | None = None -class PresidioParamsBody(GuardrailParamsBase): - guardrail: Literal["presidio"] = "presidio" - presidio_analyzer_api_base: str | None = None - presidio_anonymizer_api_base: str | None = None - # apply_to_output masks PII the model itself emitted, which also makes the - # guardrail run post_call. logging_only masks what the proxy logs. - apply_to_output: bool | None = None - logging_only: bool | None = None - - class BlockCodeExecutionParamsBody(GuardrailParamsBase): guardrail: Literal["block_code_execution"] = "block_code_execution" @@ -81,7 +71,6 @@ GuardrailParamsBody = ( ContentFilterParamsBody | BedrockGuardrailParamsBody | OpenAIModerationParamsBody - | PresidioParamsBody | BlockCodeExecutionParamsBody ) diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py deleted file mode 100644 index 9742dfc6ae7..00000000000 --- a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py +++ /dev/null @@ -1,141 +0,0 @@ -"""Live e2e: the built-in Presidio PII guardrail masks PII on the request and on -the model output. - -Presidio replaces detected PII with `` placeholders (e.g. -``) via a real analyzer + anonymizer. Two modes are checked -independently, each opted into per request (default_on=False) so it never touches -unrelated traffic: - -- pre_call: the prompt is anonymized before it reaches the model, so a - repeat-verbatim request comes back with the placeholder, never the raw email -- post_call (apply_to_output): PII the model itself emits is masked on the way - out, so the caller never receives the raw value the model produced - -A third mode, logging_only, is not covered here: the raw email stayed in the OTEL -span's `gen_ai.input.messages` on every attempt over a full poll deadline while -these two modes masked correctly, so that cell is tracked in LIT-4841 rather than -asserted against known-failing behavior. - -Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE / -PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at -locally published container ports for a host run). The chat backend is a gemini -deployment created for the test. -""" - -from __future__ import annotations - -import os -import time -from collections.abc import Callable - -import pytest - -from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker -from e2e_http import unwrap -from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody -from lifecycle import ResourceManager -from models import ChatResponse - -pytestmark = pytest.mark.e2e - -RAW_EMAIL = "alice.example.person@example.com" -PLACEHOLDER = "" - -ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}" -EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today" - - -def _content(response: ChatResponse) -> str: - if not response.choices: - return "" - message = response.choices[0].message - return (message.content if message else None) or "" - - -def _presidio_params( - mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False -) -> PresidioParamsBody: - analyzer = os.environ["PRESIDIO_ANALYZER_API_BASE"] - anonymizer = os.environ["PRESIDIO_ANONYMIZER_API_BASE"] - return PresidioParamsBody( - mode=mode, - default_on=False, - presidio_analyzer_api_base=analyzer, - presidio_anonymizer_api_base=anonymizer, - apply_to_output=apply_to_output, - logging_only=logging_only, - ) - - -def _poll_until_masked(call: Callable[[], str]) -> str: - """Retry a call until the guardrail masks its PII, returning the last content. - - Registering a guardrail is a control-plane write; the data-plane worker that - serves /chat/completions only picks it up on its next periodic DB sync (~30s - in proxy_server.py), so a call issued the instant after the create runs - against a worker that has no guardrail yet and passes the raw value through. - That is in-flight propagation, not a masking failure. Polling to the deadline - waits it out, so the assertions that follow judge the synced state; if the - mask never lands the last unmasked content is returned and they still fail. - """ - deadline = time.monotonic() + POLL_TIMEOUT - last = call() - while time.monotonic() < deadline: - if PLACEHOLDER in last and RAW_EMAIL not in last: - return last - time.sleep(POLL_INTERVAL) - last = call() - return last - - -class TestPresidioGuardrail: - @pytest.mark.covers( - "guardrail.presidio.pre_call.masks", - exercised_on=["chat_completions"], - ) - def test_pre_call_masks_pii_before_the_model_sees_it( - self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str - ) -> None: - model = client.create_backend_model(resources, prefix="e2e-presidio-pre") - name = f"e2e-presidio-pre-{unique_marker()}" - guardrail_id = client.register(name, _presidio_params("pre_call")) - resources.defer(lambda: client.delete_guardrail(guardrail_id)) - - echoed = _poll_until_masked( - lambda: _content( - unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128)) - ) - ) - assert RAW_EMAIL not in echoed, ( - "pre_call masking must strip the raw email before the model sees it, but the " - f"model echoed it back: {echoed[:300]!r}" - ) - assert PLACEHOLDER in echoed, ( - "the model should have echoed the masked placeholder the guardrail substituted, " - f"got: {echoed[:300]!r}" - ) - - @pytest.mark.covers( - "guardrail.presidio.post_call.masks", - exercised_on=["chat_completions"], - ) - def test_post_call_masks_pii_in_model_output( - self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str - ) -> None: - model = client.create_backend_model(resources, prefix="e2e-presidio-post") - name = f"e2e-presidio-post-{unique_marker()}" - guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True)) - resources.defer(lambda: client.delete_guardrail(guardrail_id)) - - out = _poll_until_masked( - lambda: _content( - unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128)) - ) - ) - assert RAW_EMAIL not in out, ( - "post_call masking must strip PII the model emitted, but the raw email reached the " - f"caller: {out[:300]!r}" - ) - assert PLACEHOLDER in out, ( - f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}" - ) From 551e5d097c11f08fd2400a25a651b1844fcf89c2 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 29 Jul 2026 14:22:22 -0700 Subject: [PATCH 42/54] feat(dashscope): add qwen3.7-plus and qwen3.7-max to the model cost map (#35123) * feat(dashscope): add qwen3.7-plus and qwen3.7-max to the model cost map Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: limit backup cost map diff to the new dashscope entries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(cost_calculator): adjust tier-only alias assertion for mapped qwen3.7-plus Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(dashscope): drop redundant cost map pinning tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(cost_calculator): point tier-only alias check at an unmapped model Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: shivam Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 50 +++++++++++++++++++ model_prices_and_context_window.json | 50 +++++++++++++++++++ tests/test_litellm/test_cost_calculator.py | 10 ++-- 3 files changed, 106 insertions(+), 4 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2cef600ea32..87d9b6afc18 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -13567,6 +13567,56 @@ } ] }, + "dashscope/qwen3.7-max": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 2.5e-06, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "dashscope/qwen3.7-plus": { + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tiered_pricing": [ + { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "output_cost_per_token": 1.6e-06, + "range": [ + 0, + 256000.0 + ] + }, + { + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.8e-06, + "range": [ + 256000.0, + 1000000.0 + ] + } + ] + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index db28118d52b..0edd3bd5f30 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -13567,6 +13567,56 @@ } ] }, + "dashscope/qwen3.7-max": { + "cache_read_input_token_cost": 5e-07, + "input_cost_per_token": 2.5e-06, + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "dashscope/qwen3.7-plus": { + "litellm_provider": "dashscope", + "max_input_tokens": 991808, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "source": "https://www.alibabacloud.com/help/en/model-studio/models", + "supports_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tiered_pricing": [ + { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "output_cost_per_token": 1.6e-06, + "range": [ + 0, + 256000.0 + ] + }, + { + "cache_read_input_token_cost": 2.4e-07, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.8e-06, + "range": [ + 256000.0, + 1000000.0 + ] + } + ] + }, "dashscope/qwq-plus": { "input_cost_per_token": 8e-07, "litellm_provider": "dashscope", diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 276ee96ed65..e6ae1f85cfd 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1014,9 +1014,9 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): router = Router( model_list=[ { - "model_name": "qwen-3.7-plus", + "model_name": "qwen-tier-only", "litellm_params": { - "model": "dashscope/qwen3.7-plus", + "model": "dashscope/qwen-tier-only-test", "api_key": "sk-fake", }, "model_info": { @@ -1037,10 +1037,12 @@ def test_tiered_pricing_only_deployment_selects_router_model_id(): assert entry.get("input_cost_per_token") is None assert entry.get("tiered_pricing") is not None # The stripped shared alias must not carry tiered pricing. - assert litellm.model_cost["dashscope/qwen3.7-plus"].get("tiered_pricing") is None + assert ( + litellm.model_cost["dashscope/qwen-tier-only-test"].get("tiered_pricing") is None + ) selected = _select_model_name_for_cost_calc( - model="dashscope/qwen3.7-plus", + model="dashscope/qwen-tier-only-test", completion_response=None, custom_pricing=True, custom_llm_provider="dashscope", From 0a6b372126e71a6a46ef4e96c9a83ecee7e2f8aa Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 17:20:59 -0700 Subject: [PATCH 43/54] feat(ui): link organization teams to their team detail pages (#35120) * feat(ui): link organization teams to their team detail pages On the organization info page the teams shown for an org were plain badges, so walking to a team meant copying its id and finding it by hand on the teams page Team badges now link to /teams?team=, which opens that team's detail page directly since #35112. Adds a shared BadgeLink (a badge rendered as a real anchor with modifier-aware client-side navigation, so cmd-click opens a new tab) and a teamDetailHref builder for reuse by future entity links * fix(ui): format BadgeLink, split its modifier-click chain, and size it up prettier wanted the Badge props wrapped, and local/no-long-condition-chain flagged the four-way modifier-click guard; the guard is now two named conditions. Linked badges also render slightly larger (text-sm, roomier padding) than plain badges so clickable entries stand out * feat(ui): size org model badges to match the linked team badges BadgeLink's href is now optional; without one it renders the same enlarged plain badge (no pointer, no hover), so the org page's model badges share the component and the size while staying non-clickable --- .../organization/organization_view.test.tsx | 62 ++++++++++++++++++- .../organization/organization_view.tsx | 19 +++--- .../src/components/shared/BadgeLink.test.tsx | 42 +++++++++++++ .../src/components/shared/BadgeLink.tsx | 46 ++++++++++++++ ui/litellm-dashboard/src/utils/entityLinks.ts | 5 ++ 5 files changed, 161 insertions(+), 13 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/BadgeLink.tsx create mode 100644 ui/litellm-dashboard/src/utils/entityLinks.ts diff --git a/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx b/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx index 7fed639c491..ec7baecad2d 100644 --- a/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx +++ b/ui/litellm-dashboard/src/components/organization/organization_view.test.tsx @@ -6,7 +6,14 @@ import { renderWithProviders } from "../../../tests/test-utils"; import OrganizationInfoView from "./organization_view"; import { useOrganization } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; -// Mock networking calls used by the component's mutation handlers +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: vi.fn() }), + usePathname: () => "/organizations", + useSearchParams: () => new URLSearchParams(window.location.search), +})); + +// Mock networking calls used by the component's mutation handlers. entityLinks -> migratedPages +// imports serverRootPath from the same module, so the mock must export it too. vi.mock("../networking", () => { return { __esModule: true, @@ -14,6 +21,7 @@ vi.mock("../networking", () => { organizationMemberUpdateCall: vi.fn(), organizationMemberDeleteCall: vi.fn(), organizationUpdateCall: vi.fn(), + serverRootPath: "", }; }); @@ -206,6 +214,58 @@ test("should display team ID as fallback when alias is not found", async () => { }); }); +test("links each team badge to that team's detail page", async () => { + const orgWithTeams = { + ...mockOrg, + teams: [{ team_id: "team_123" }, { team_id: "team_456" }], + }; + mockUseOrganization.mockReturnValue({ data: orgWithTeams, isLoading: false } as any); + + renderWithProviders( + {}} + accessToken="test-token" + is_org_admin={false} + is_proxy_admin={false} + userModels={[]} + editOrg={false} + />, + ); + + await waitFor(() => { + expect(screen.getByRole("link", { name: "Engineering Team" })).toHaveAttribute( + "href", + expect.stringContaining("/teams?team=team_123"), + ); + expect(screen.getByRole("link", { name: "Marketing Team" })).toHaveAttribute( + "href", + expect.stringContaining("/teams?team=team_456"), + ); + }); +}); + +test("model badges stay non-clickable", async () => { + mockUseOrganization.mockReturnValue({ data: mockOrg, isLoading: false } as any); + + renderWithProviders( + {}} + accessToken="test-token" + is_org_admin={false} + is_proxy_admin={false} + userModels={[]} + editOrg={false} + />, + ); + + await waitFor(() => { + expect(screen.getByText("gpt-4o-mini")).toBeInTheDocument(); + }); + expect(screen.queryByRole("link", { name: "gpt-4o-mini" })).not.toBeInTheDocument(); +}); + test("should keep unsaved settings edits when switching tabs and back", async () => { mockUseOrganization.mockReturnValue({ data: mockOrg, isLoading: false } as any); diff --git a/ui/litellm-dashboard/src/components/organization/organization_view.tsx b/ui/litellm-dashboard/src/components/organization/organization_view.tsx index 44b67765850..10af5bc1a07 100644 --- a/ui/litellm-dashboard/src/components/organization/organization_view.tsx +++ b/ui/litellm-dashboard/src/components/organization/organization_view.tsx @@ -4,12 +4,13 @@ import { useQueryClient } from "@tanstack/react-query"; import { useVisitedTabs } from "@/hooks/useVisitedTabs"; import { MoneyCell } from "@/components/shared/table_cells"; import CopyButton from "@/components/shared/CopyButton"; -import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Card, CardContent } from "@/components/ui/card"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { formatNumberWithCommas } from "@/utils/dataUtils"; +import { teamDetailHref } from "@/utils/entityLinks"; import { createTeamAliasMap } from "@/utils/teamUtils"; +import { BadgeLink } from "@/components/shared/BadgeLink"; import type { ColumnsType } from "antd/es/table"; import { ArrowLeft } from "lucide-react"; import React, { useMemo, useState } from "react"; @@ -220,13 +221,9 @@ const OrganizationInfoView: React.FC = ({

Models

{orgData.models.length === 0 ? ( - All proxy models + All proxy models ) : ( - orgData.models.map((model, index) => ( - - {model} - - )) + orgData.models.map((model, index) => {model}) )}
@@ -237,9 +234,9 @@ const OrganizationInfoView: React.FC = ({

Teams

{orgData.teams?.map((team, index) => ( - + {teamAliasMap[team.team_id] || team.team_id} - + ))}
@@ -309,9 +306,7 @@ const OrganizationInfoView: React.FC = ({

Models

{orgData.models.map((model, index) => ( - - {model} - + {model} ))}
diff --git a/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx b/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx new file mode 100644 index 00000000000..10a192c8af0 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/BadgeLink.test.tsx @@ -0,0 +1,42 @@ +/* @vitest-environment jsdom */ +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { BadgeLink } from "./BadgeLink"; + +const push = vi.fn(); +vi.mock("next/navigation", () => ({ useRouter: () => ({ push }) })); + +describe("BadgeLink", () => { + beforeEach(() => { + push.mockClear(); + }); + + it("renders an anchor pointing at the target href", () => { + render(My Team); + expect(screen.getByRole("link", { name: "My Team" })).toHaveAttribute("href", "/ui/teams?team=t1"); + }); + + it("navigates client-side on plain click", async () => { + const user = userEvent.setup(); + render(My Team); + await user.click(screen.getByRole("link", { name: "My Team" })); + expect(push).toHaveBeenCalledWith("/ui/teams?team=t1"); + }); + + it("leaves modified clicks to the browser so new-tab shortcuts keep working", async () => { + const user = userEvent.setup(); + render(My Team); + await user.keyboard("{Meta>}"); + await user.click(screen.getByRole("link", { name: "My Team" })); + await user.keyboard("{/Meta}"); + expect(push).not.toHaveBeenCalled(); + }); + + it("renders a plain same-sized badge when no href is given", () => { + render(all-proxy-models); + expect(screen.getByText("all-proxy-models")).toBeInTheDocument(); + expect(screen.queryByRole("link", { name: "all-proxy-models" })).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx b/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx new file mode 100644 index 00000000000..444d2acdcfe --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/BadgeLink.tsx @@ -0,0 +1,46 @@ +"use client"; + +import { useRouter } from "next/navigation"; +import * as React from "react"; + +import { Badge } from "@/components/ui/badge"; +import { cn } from "@/lib/cva.config"; + +const ENTITY_BADGE_SIZE = "px-2.5 py-1 text-sm"; + +interface BadgeLinkProps { + href?: string; + variant?: React.ComponentProps["variant"]; + className?: string; + children: React.ReactNode; +} + +export function BadgeLink({ href, variant = "secondary", className, children }: BadgeLinkProps) { + const router = useRouter(); + + if (!href) { + return ( + + {children} + + ); + } + + const handleClick = (e: React.MouseEvent) => { + const hasModifierKey = e.metaKey || e.ctrlKey || e.shiftKey; + const isNativeNewTabClick = hasModifierKey || e.button === 1; + if (isNativeNewTabClick) return; + e.preventDefault(); + router.push(href); + }; + + return ( + } + > + {children} + + ); +} diff --git a/ui/litellm-dashboard/src/utils/entityLinks.ts b/ui/litellm-dashboard/src/utils/entityLinks.ts new file mode 100644 index 00000000000..2659a307866 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/entityLinks.ts @@ -0,0 +1,5 @@ +import { migratedHref } from "@/utils/migratedPages"; + +export function teamDetailHref(teamId: string): string { + return `${migratedHref("teams")}?team=${encodeURIComponent(teamId)}`; +} From ae74b06ae3d9b63381c854c01a7e8b3f9f4b8dbc Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Wed, 29 Jul 2026 17:49:30 -0700 Subject: [PATCH 44/54] fix(tests): assert Content variants are identified by type, not by the discriminator keyword (#35161) `test_content_schema_uses_discriminator` fetches Google's live Interactions OpenAPI document and required an OpenAPI `discriminator` on the `Content` union. Google has since dropped that keyword and now pins `type` with a `const` on each variant instead, so the assertion fails on the current spec and the `misc` shard is red on every open PR against staging The information the transformation actually needs did not change: a content part is still routed by reading its `type`, and each variant still declares exactly one distinct value for it. So the test now asserts that property directly, and accepts either spelling, a `discriminator` on the union or a `const` (or single-value `enum`) on each member It stays a real check rather than a weakened one. Against the live spec it fails if TextContent loses its type property, if `text` is renamed, if two variants claim the same type value, if `Content` stops being a union of named variants, or if a discriminator appears on some property other than `type` --- .../interactions/test_openapi_compliance.py | 66 ++++++++++++++----- 1 file changed, 49 insertions(+), 17 deletions(-) diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 11b08fa45a8..1fe343ca6ee 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -37,6 +37,13 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: ) +def _declared_type_value(variant_schema: Dict[str, Any]) -> Any: + """The single `type` value a union variant pins, whether spelled as a const or a 1-item enum.""" + type_property = variant_schema.get("properties", {}).get("type", {}) + enum_values = type_property.get("enum") or [] + return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None) + + @pytest.fixture(scope="module") def spec_dict() -> Dict[str, Any]: """Load raw spec dict for manual validation.""" @@ -105,26 +112,51 @@ class TestRequestCompliance: assert "string" in input_types, "Input should support string" assert "array" in input_types, "Input should support array" - def test_content_schema_uses_discriminator(self, spec_dict): - """Verify Content uses type discriminator.""" + def test_content_variants_are_identified_by_their_type_field(self, spec_dict): + """Verify a Content part can be told apart by its `type`, however the spec spells that. + + Our transformation reads `type` off each content part to route it, so what has to hold is + that every variant of the union pins a distinct `type` value and that text is one of them. + A spec may express that with an OpenAPI `discriminator` on the union or with a `const` on + each member's own `type`; both are equivalent for us, so accepting only the first makes + this test fail on a stylistic change upstream that costs us nothing. + """ content_schema = spec_dict["components"]["schemas"]["Content"] - assert "discriminator" in content_schema - assert content_schema["discriminator"]["propertyName"] == "type" - - # Check TextContent is an option (via mapping if present, or via oneOf refs) - mapping = content_schema["discriminator"].get("mapping") - if mapping: - assert "text" in mapping - print(f"Content type discriminator mapping: {list(mapping.keys())}") - else: - # Discriminator without explicit mapping — verify via oneOf - one_of = content_schema.get("oneOf", []) - ref_names = [opt["$ref"].split("/")[-1] for opt in one_of if "$ref" in opt] + discriminator = content_schema.get("discriminator") + if discriminator is not None: assert ( - "TextContent" in ref_names - ), f"TextContent not found in oneOf refs: {ref_names}" - print(f"Content type discriminator (no mapping), oneOf refs: {ref_names}") + discriminator.get("propertyName") == "type" + ), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + + variant_names = [ + option["$ref"].split("/")[-1] + for option in content_schema.get("oneOf", []) + if "$ref" in option + ] + assert variant_names, f"Content is not a union of named variants: {content_schema}" + + mapping = (discriminator or {}).get("mapping") or {} + type_values = { + variant: mapping_value + for mapping_value, ref in mapping.items() + for variant in [ref.split("/")[-1]] + } or { + variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {})) + for variant in variant_names + } + + assert set(type_values) == set(variant_names) and all(type_values.values()), ( + f"every Content variant needs a discoverable type value, " + f"got {type_values} for variants {sorted(variant_names)}" + ) + assert len(set(type_values.values())) == len(type_values), ( + f"Content variants must pin DISTINCT type values, got {type_values}" + ) + assert type_values.get("TextContent") == "text", ( + f"TextContent must be reachable as type 'text', got {type_values}" + ) + print(f"Content variants by type: {type_values}") def test_text_content_schema(self, spec_dict): """Verify TextContent schema.""" From 7041f5768f4d5bbbe5d66789a6d9878f3f86cfce Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 28 Jul 2026 18:41:56 -0700 Subject: [PATCH 45/54] fix(mcp): never write discovery results to the row, heal rows a release already stamped, and retry failed discovery with backoff An interactive oauth2 MCP server created with explicit endpoint URLs and no issuer served 400 "authorization url is not configured" from /authorize about a minute after creation, with the admin's endpoints intact in the row the whole time (#34985). Discovery wrote its trust-on-first-use issuer into the same column an admin writes, so the next registry build read the gateway's own output back as an admin pin, anchored the server to RFC 8414 section 3.3, and discarded the stored endpoint columns; one transient metadata fetch failure then had nothing to serve, and the reload fast path pinned the broken entry until an unrelated config write The core of the fix is a deletion. The gateway no longer writes discovery results anywhere: the OAuth columns and credentials.scopes carry admin intent alone, and everything discovery learns lives on the in-memory registry entry, as the existing carry-forward already assumes. With no gateway write there is no value whose provenance a later build can misread, so the accidental anchoring cannot be expressed Deleting the write cannot fix a row a released version already stamped, which still reads as pinned, so a one-time startup heal clears those stamps. The signal is necessarily a heuristic: updated_by records only the most recent writer and no audit trail says which field it touched. A row is therefore healed only on the full signature of the defect, which is discovery as the last writer plus an issuer plus at least one configured endpoint column that anchoring is actively discarding; rows with an issuer but no configured endpoints are left alone, since for them both paths resolve from the same upstream document. Every heal logs the cleared value so an admin who pinned deliberately can re-pin, and the heal records its own actor, which makes it idempotent The reload fast path exempts servers missing an endpoint their flow needs, so failed discovery retries on the normal reload cadence rather than waiting for a config write. Flow requirements are read through effective_oauth2_flow, the column-first shape-fallback judge every flow decision uses, so a legacy null-flow M2M row is classified exactly as the request path classifies it instead of re-discovering forever; a dcr_bridge server with no configured client needs its registration endpoint for the relay arm, and an entra_obo server needs a scope, both of which discovery can supply. Retries back off per server, doubling from one reload cadence to a fifteen-minute cap, so a permanently unresolvable server cannot re-run the RFC 9728 to 8414 chain and re-log its warning every cycle forever Deployments with store_model_in_db unset or false loaded MCP servers exactly once at startup, leaving that retry with no driver, so they now refresh the registry on the same reload interval. That job deliberately calls a reload-only entry point rather than the startup composite, keeping the one-time oauth2_flow backfill and issuer heal out of a recurring path Losing the persisted trust-on-first-use issuer also means the issuer column no longer changes underneath the OAuth token identity, so user tokens are purged only when an admin actually edits the server Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 288 ++++---- .../mcp_server/oauth_issuer_stamp_backfill.py | 148 ++++ .../mcp_management_endpoints.py | 1 - litellm/proxy/proxy_server.py | 54 ++ .../mcp_server/test_mcp_partial_update.py | 28 +- .../mcp_server/test_mcp_server_manager.py | 654 ++++++++---------- .../test_oauth_issuer_stamp_backfill.py | 129 ++++ .../_components/OAuthFormFields.tsx | 2 +- 8 files changed, 768 insertions(+), 536 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 82b820d8cd9..3e0775ac09e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = ( ) +# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one +# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request +# amplification and log volume of a permanently broken configuration. +_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0 +_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0 + + def _blank_to_none(value: str | None) -> str | None: """Collapse an absent, empty, or whitespace-only string to ``None``. @@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer( authorization_url: str | None, token_url: str | None, registration_url: str | None, + server_ref: str, ) -> tuple[str | None, str | None, str | None]: """The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual @@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer( i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site so the invariant holds in one place instead of being re-derived per merge. """ - if issuer is not None and is_discovery_auth_type: - return None, None, None - return authorization_url, token_url, registration_url + if issuer is None or not is_discovery_auth_type: + return authorization_url, token_url, registration_url + discarded = sorted( + label + for label, value in ( + ("authorization_url", authorization_url), + ("token_url", token_url), + ("registration_url", registration_url), + ) + if value + ) + if discarded: + verbose_logger.warning( + "MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the " + "sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather " + "than falling back to them. To use manually configured endpoints instead, clear the Issuer " + "field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may " + "have been resolved under it), or clear the Issuer alone to re-discover from the server url.", + server_ref, + ", ".join(discarded), + "is" if len(discarded) == 1 else "are", + ) + return None, None, None def _normalized_authorize_endpoint(url: str) -> str: @@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool: return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer) +def _flow_endpoints_missing( + auth_type: MCPAuthType | None, + oauth2_flow: str | None, + authorization_url: str | None, + token_url: str | None, + token_exchange_endpoint: str | None = None, +) -> bool: + """Whether a built server is missing an endpoint its flow needs to run at all. + + Used by the reload fast-path exemption: discovery runs at build time only, and the fast path + reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty + (transient upstream failure, rate limiting) would stay broken until some unrelated config write + bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries + discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and + none for those with no discovery source, since the build skips discovery for both. + """ + if auth_type == MCPAuth.oauth2_token_exchange: + # A configured exchange endpoint replaces discovery entirely; only a server that must + # discover its token endpoint and still has none is unresolved. + return token_exchange_endpoint is None and token_url is None + if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: + return False + if oauth2_flow == "client_credentials": + return token_url is None + return authorization_url is None or token_url is None + + +def _oauth_endpoints_unresolved(server: MCPServer) -> bool: + """``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check. + + The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every + flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left + unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here + would classify it as interactive-missing-endpoints and re-run discovery on every reload. + """ + if ( + server.auth_type == MCPAuth.oauth2_token_exchange + and server.token_exchange_profile == "entra_obo" + and not server.scopes + ): + # entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes + # can come from resource discovery, so a server that resolved its endpoints but no scopes is + # still unresolved for its flow. + return True + if server.is_dcr_bridge and not server.client_id and server.registration_url is None: + # A DCR bridge with no admin-configured client can only register callers through the + # upstream's registration endpoint, so a build that resolved the authorize and token + # endpoints but not registration_endpoint (partial metadata) is still unresolved for its + # flow and must keep retrying; without this it silently degrades to the short-circuit arm + # until an unrelated config write. Scopes are deliberately NOT part of completeness: they + # are a request hint the authorization server bounds at consent (RFC 6749 section 3.3), + # and a server without them is fully functional. + return True + return _flow_endpoints_missing( + server.auth_type, + MCPServerManager.effective_oauth2_flow(server), + server.authorization_url, + server.token_url, + server.token_exchange_endpoint, + ) + + def _endpoints_corroborate_authorization_url( source_authorization_url: str | None, trusted_authorization_url: str | None, @@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv during re-discovery downgrades a working server (``authorization_url`` set) to a broken one (``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix`` carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous - endpoints may then belong to a different upstream. ``registration_url`` IS carried even though - ``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores - the same in-memory value the previous build already ran with, while persisting it would flip - ``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge - servers that never had one configured. + endpoints may then belong to a different upstream. Discovery results live only on the in-memory + registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent + alone, so this carry is the sole last-known-good mechanism and restores exactly the values the + previous build already ran with. Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous @@ -1182,6 +1271,40 @@ class MCPServerManager: # empty result, or failure). Used to throttle re-probes for servers that do # not return instructions, and to apply a short cooldown after failures. self._upstream_initialize_instructions_probed_at: dict[str, float] = {} + # Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a + # server whose endpoints never resolve backs off instead of re-running the full + # RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever. + self._oauth_discovery_retry_state: dict[ + str, tuple[int, float] + ] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success + + def _oauth_discovery_retry_due(self, server_id: str) -> bool: + """Whether an unresolved server is due for another discovery attempt. + + The reload fast-path exemption is what retries a failed discovery, so without a cooldown a + permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback + chain and re-emits its unresolved-endpoints warning on every reload, per server, forever. + Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to + ``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next + reload while a broken configuration settles to one attempt per cap. + """ + state = self._oauth_discovery_retry_state.get(server_id) + if state is None: + return True + failures, attempted_at = state + delay = min( + _OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)), + _OAUTH_DISCOVERY_RETRY_MAX_SECONDS, + ) + return (time.monotonic() - attempted_at) >= delay + + def _record_oauth_discovery_outcome(self, server: MCPServer) -> None: + """Advance or clear a server's retry cooldown after a rebuild resolved it or did not.""" + if not _oauth_endpoints_unresolved(server): + self._oauth_discovery_retry_state.pop(server.server_id, None) + return + failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0)) + self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic()) def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: raw = getattr(client, "_last_initialize_instructions", None) @@ -1357,6 +1480,7 @@ class MCPServerManager: manual_authorization_url, manual_token_url, manual_registration_url, + server_name or server_id, ) should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( is_discovery_auth_type or obo_needs_discovery @@ -1834,7 +1958,6 @@ class MCPServerManager: *, credentials_are_encrypted: bool = True, env_vars_are_encrypted: Optional[bool] = None, - persist_discovered_endpoints: bool = True, ) -> MCPServer: _mcp_info: MCPInfo = mcp_server.mcp_info or {} env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None)) @@ -1925,7 +2048,12 @@ class MCPServerManager: or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url), ) manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer( - manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url + manual_issuer, + is_discovery_auth_type, + manual_authorization_url, + manual_token_url, + manual_registration_url, + mcp_server.alias or mcp_server.server_name or mcp_server.server_id, ) gated_oauth_metadata = await self._resolve_table_oauth_metadata( mcp_server=mcp_server, @@ -2033,143 +2161,8 @@ class MCPServerManager: max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None), ) _warn_internal_delegate_pkce_if_applicable(new_server, source="database") - if persist_discovered_endpoints: - await self._persist_discovered_obo_token_url( - server_id=mcp_server.server_id, - auth_type=auth_type, - existing_token_url=manual_token_url, - discovered_token_url=new_server.token_url, - ) - await self._persist_discovered_oauth_endpoints( - server_id=mcp_server.server_id, - auth_type=auth_type, - existing_issuer=manual_issuer, - existing_authorization_url=manual_authorization_url, - existing_token_url=manual_token_url, - existing_scopes=scopes, - metadata=gated_oauth_metadata, - is_issuer_anchored=use_issuer_anchor, - ) return new_server - async def _persist_discovered_obo_token_url( - self, - *, - server_id: str, - auth_type: Optional[MCPAuthType], - existing_token_url: Optional[str], - discovered_token_url: Optional[str], - ) -> None: - """Write a freshly discovered OBO token endpoint back onto the DB row. - - ``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an - ``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise - lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild - re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no - endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery`` - return False on the next build. Fires at most once per server (skipped once the row has a - value), and is best-effort: a write failure just means discovery runs again next time. - """ - if auth_type != MCPAuth.oauth2_token_exchange: - return - if existing_token_url or not discovered_token_url: - return - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 - - if prisma_client is None: - return - try: - await MCPServerRepository(prisma_client).table.update( - where={"server_id": server_id}, - data={"token_url": discovered_token_url}, - ) - verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id) - except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build - verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc) - - async def _persist_discovered_oauth_endpoints( - self, - *, - server_id: str, - auth_type: MCPAuthType | None, - existing_issuer: str | None, - existing_authorization_url: str | None, - existing_token_url: str | None, - existing_scopes: list[str] | None, - metadata: MCPOAuthMetadata | None, - is_issuer_anchored: bool = False, - ) -> None: - """Write freshly discovered OAuth endpoints back onto the DB row. - - Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2 - family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on - the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path - calls ``update_server``) and on every post-write DB reload, so one failed re-discovery - serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds. - Only fills row fields that are currently empty, never persists origin-fallback guesses - (RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url`` - because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a - failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so - they merge into the credentials blob without touching the stored client credentials. - - For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the - §3.3-validated issuer document on every build, so they are NOT persisted into the endpoint - columns: persisting them would make the next build see populated endpoints and treat them as - authoritative stored values, defeating the "endpoints come solely from the issuer" invariant. - Only the resource-driven scopes are persisted for such servers. - """ - if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: - return - if metadata is None or metadata.from_origin_fallback: - return - issuer_update = ( - {"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {} - ) - authorization_url_update = ( - {"authorization_url": metadata.authorization_url} - if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored - else {} - ) - token_url_update = ( - {"token_url": metadata.token_url} - if metadata.token_url and not existing_token_url and not is_issuer_anchored - else {} - ) - scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {} - updates: dict[str, object] = { - **issuer_update, - **authorization_url_update, - **token_url_update, - **scopes_update, - } - if not updates: - return - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load - update_mcp_server, - ) - from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time - from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup - - if prisma_client is None: - return - try: - await update_mcp_server( - prisma_client=prisma_client, - data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}), - touched_by="mcp_oauth_discovery", - ) - verbose_logger.info( - "Persisted discovered OAuth endpoints for MCP server %s: %s", - server_id, - sorted(updates), - ) - except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build - verbose_logger.warning( - "Failed to persist discovered OAuth endpoints for MCP server %s: %s", - server_id, - exc, - ) - async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True): """Register OpenAPI tools if the server has a spec_path configured.""" if server.spec_path: @@ -5347,6 +5340,10 @@ class MCPServerManager: and existing_server.updated_at is not None and server.updated_at is not None and existing_server.updated_at == server.updated_at + and not ( + _oauth_endpoints_unresolved(existing_server) + and self._oauth_discovery_retry_due(server.server_id) + ) ): # Re-use existing server instance to avoid re-running build_mcp_server_from_table() # which can perform network discovery for OAuth2 servers. @@ -5364,6 +5361,7 @@ class MCPServerManager: # already-decrypted records add_server/update_server are handed. # Decrypt them while building the registry entry. new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True) + self._record_oauth_discovery_outcome(new_server) # Carry the cached short_prefix from the previous registry entry # (if any) so the prefix is stable across reloads. if existing_server is not None and existing_server.short_prefix: diff --git a/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py new file mode 100644 index 00000000000..874fcc64772 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/oauth_issuer_stamp_backfill.py @@ -0,0 +1,148 @@ +"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself. + +Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer`` +column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next +registry build read the gateway's own output back as admin intent: the server turned issuer-anchored +(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a +failed issuer-document fetch left it with no authorize endpoint (GH #34985). + +Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which +still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints +apply again. + +The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent +writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is +not directly knowable. Two independent clauses bound it, and each rules out a different way of +destroying a pin an admin meant. + +Configured endpoints must be present. A deliberately pinned row very often has none, both because the +Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an +issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never +be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the +anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the +unresolved-endpoint retry and the anchored-discard warning. + +The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the +one self-attested by the authorization-server document discovery reached from this very server, so +endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed +endpoints for a different authority is expressing an intent that clearing the issuer would discard, so +that row is warned about and never healed. + +What survives both clauses is a row whose configured endpoints and stamped issuer share an origin, +which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and +for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the +RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the +cleared value so it can be restored, and the clear is recorded under this module's actor so the heal +runs at most once per row. +""" + +from typing import Protocol +from urllib.parse import urlparse + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity +from litellm.proxy.utils import PrismaClient + +# The actor the removed discovery write-back stamped rows with. +_DISCOVERY_ACTOR = "mcp_oauth_discovery" + +# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it +# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered. +_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill" + +_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate") + + +def _origin(url: str) -> str | None: + """The scheme-and-authority identity of ``url``, or ``None`` when it has none. + + Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the + RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here. + """ + parsed = urlparse(canonicalize_url_identity(url)) + if not parsed.scheme or not parsed.netloc: + return None + return f"{parsed.scheme}://{parsed.netloc}" + + +class _MCPServerRow(Protocol): + """The MCP server row fields this heal reads, so the untyped DB record is narrowed once here.""" + + server_id: str + alias: str | None + server_name: str | None + auth_type: str | None + issuer: str | None + authorization_url: str | None + token_url: str | None + registration_url: str | None + updated_by: str | None + + +def _is_stamped_issuer_row(row: _MCPServerRow) -> bool: + """Whether this row carries the full signature of a gateway-written issuer stamp. + + The whole rule lives here, including the writer check the query also filters on, so the decision + to clear an admin-visible field is auditable in one place rather than split between a predicate + and a query. + """ + if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR: + return False + if not (getattr(row, "issuer", None) or "").strip(): + return False + if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING: + return False + configured = tuple( + value.strip() + for value in (row.authorization_url, row.token_url, row.registration_url) + if value and value.strip() + ) + if not configured: + return False + issuer_origin = _origin(row.issuer or "") + return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured) + + +async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int: + """Clear gateway-written issuer stamps, returning the number of rows healed.""" + candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many( + where={ + "updated_by": _DISCOVERY_ACTOR, + "auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)}, + }, + ) + stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row)) + if not stamped: + return 0 + + healed = 0 + for row in stamped: + try: + await prisma_client.db.litellm_mcpservertable.update( + where={"server_id": row.server_id}, + data={"issuer": None, "updated_by": _BACKFILL_ACTOR}, + ) + except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc + ) + continue + healed += 1 + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery " + "had written that value onto the Issuer column, which made the server issuer-anchored and " + "fail-closed, and its configured Authorization/Token/Registration URLs were being ignored " + "as a result; those now apply again. If you pinned this issuer deliberately, set it again " + "via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.", + row.issuer, + row.server_id, + row.alias or row.server_name, + ) + + if healed: + verbose_proxy_logger.warning( + "MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth " + "discovery rather than by an admin", + healed, + ) + return healed diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 282184d6495..1205d23ce02 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -1526,7 +1526,6 @@ if MCP_AVAILABLE: temporary_server = await global_mcp_server_manager.build_mcp_server_from_table( temp_record, credentials_are_encrypted=False, - persist_discovered_endpoints=False, ) _cache_temporary_mcp_server( temporary_server, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 70484eb1e4e..18a927e7a44 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6758,6 +6758,9 @@ class ProxyConfig: from litellm.proxy._experimental.mcp_server.oauth2_flow_backfill import ( backfill_null_oauth2_flows, ) + from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import ( + backfill_discovery_stamped_issuers, + ) try: if prisma_client is not None: @@ -6767,6 +6770,16 @@ class ProxyConfig: "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db backfill - {}".format(str(e)) ) + try: + if prisma_client is not None: + await backfill_discovery_stamped_issuers(prisma_client) + except Exception as e: # noqa: BLE001 + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db issuer stamp backfill - {}".format( + str(e) + ) + ) + try: await global_mcp_server_manager.reload_servers_from_database() except Exception as e: @@ -6778,6 +6791,31 @@ class ProxyConfig: if self._should_load_db_object(object_type="mcp"): await self._init_mcp_servers_in_db() + async def reload_mcp_servers_from_db(self) -> None: + """Registry refresh only, for the periodic job in store_model_in_db-off deployments. + + Deliberately narrower than ``init_mcp_servers_from_db``: the oauth2_flow backfill is a write + path that only needs to run once at startup, so the cadence here is purely the read-side + reload whose fast-path exemption retries failed OAuth discovery. Gated the same way, so an + admin who excluded mcp from supported_db_objects opts out of this too. + """ + if not self._should_load_db_object(object_type="mcp"): + return + from litellm.proxy._experimental.mcp_server.utils import is_mcp_available + + if not is_mcp_available(): + return + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + try: + await global_mcp_server_manager.reload_servers_from_database() + except Exception as e: # noqa: BLE001 # scheduled job: a reload failure must not kill the recurring retry + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:reload_mcp_servers_from_db - {}".format(str(e)) + ) + async def _init_agents_in_db(self, prisma_client: PrismaClient): from litellm.proxy.agent_endpoints.agent_registry import ( global_agent_registry as AGENT_REGISTRY, @@ -8099,6 +8137,22 @@ class ProxyStartupEvent: if store_model_in_db is not True: await proxy_config.init_mcp_servers_from_db() + if prisma_client is not None: + # DB-backed MCP servers are live objects in every mode, so the registry refresh that + # store_model_in_db=True deployments get via the add_deployment job must run here + # too; without it, a server whose OAuth discovery failed at startup is rebuilt only + # by a management write, since the reload fast path is the retry's only driver. + mcp_reload_interval_seconds = proxy_config_reload_interval_seconds + if not isinstance(mcp_reload_interval_seconds, int) or mcp_reload_interval_seconds <= 0: + mcp_reload_interval_seconds = 30 + scheduler.add_job( + proxy_config.reload_mcp_servers_from_db, + "interval", + seconds=mcp_reload_interval_seconds, + id="reload_mcp_servers_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) await cls._initialize_slack_alerting_jobs( scheduler=scheduler, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py index c063915e2e8..f6bd79c5d2d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_partial_update.py @@ -240,10 +240,10 @@ async def test_explicit_null_clears_upstream_resource_and_keeps_the_rest_of_the_ @pytest.mark.asyncio -async def test_url_change_clears_stale_discovered_oauth_fields(): - """Re-pointing the server url at a potentially different upstream must clear the discovered or - trust-on-first-use OAuth issuer and endpoints, so the new upstream re-discovers instead of - anchoring on the previous upstream's issuer (RFC 8414 §3.3 against a stale anchor).""" +async def test_url_change_clears_stale_oauth_fields(): + """Re-pointing the server url at a potentially different upstream must clear the OAuth issuer and + endpoints, so the new upstream re-discovers instead of anchoring on the previous upstream's issuer + (RFC 8414 §3.3 against a stale anchor).""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" @@ -350,11 +350,13 @@ async def test_repointing_pinned_issuer_clears_stale_endpoints_keeps_new_issuer( @pytest.mark.asyncio -async def test_establishing_issuer_first_time_preserves_discovered_fields(): - """Establishing an issuer for the first time (None -> X), which is exactly what the trust-on-first-use - discovery write-back does, must NOT clear the endpoints or oauth2_flow it discovered in the same - write. Only an issuer that was already pinned and is now changed or cleared invalidates its - endpoints, so the discovery persist cannot wipe the fields it just resolved.""" +async def test_establishing_issuer_first_time_preserves_endpoints_set_in_the_same_write(): + """Establishing an issuer for the first time (None -> X) must NOT clear endpoints or oauth2_flow + submitted in the same write. Only an issuer that was already pinned and is now changed or cleared + invalidates its endpoints, so an admin configuring an issuer and its endpoints together keeps + both. The write-back this once guarded (trust-on-first-use discovery stamping the issuer it had + just resolved) no longer exists; the db.py rule it relies on still governs admin writes, which is + what this now covers.""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" @@ -370,7 +372,7 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields(): token_url="https://discovered-idp.example.com/token", oauth2_flow="authorization_code", ) - await update_mcp_server(mock_prisma, data, "mcp_oauth_discovery") + await update_mcp_server(mock_prisma, data, "some-admin@example.com") data_dict = mock_prisma.db.litellm_mcpservertable.update.call_args[1]["data"] assert data_dict["issuer"] == "https://discovered-idp.example.com" @@ -380,9 +382,9 @@ async def test_establishing_issuer_first_time_preserves_discovered_fields(): @pytest.mark.asyncio -async def test_unchanged_url_does_not_clear_discovered_oauth_fields(): - """A partial update that resends the same url (or omits it) must not clear the discovered OAuth - fields, so a routine save does not force needless re-discovery.""" +async def test_unchanged_url_does_not_clear_oauth_fields(): + """A partial update that resends the same url (or omits it) must not clear the OAuth fields, so a + routine save does not force needless re-discovery.""" mock_prisma = _mock_prisma() existing = MagicMock() existing.auth_type = "oauth2" 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 5f7f2267fc7..8a8dea0ba28 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 @@ -2,6 +2,7 @@ import importlib import asyncio import json import logging +import time import os import sys from datetime import datetime @@ -35,6 +36,8 @@ from mcp.types import Tool as MCPTool from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _deserialize_json_dict, + _flow_endpoints_missing, + _oauth_endpoints_unresolved, _deserialize_json_list, _normalize_mcp_server_cost_info, _should_strip_caller_authorization, @@ -1594,21 +1597,15 @@ class TestMCPServerManager: token_url="https://idp.example.com/token", scopes=["read"], ) - with ( - patch.object( - manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved) - ) as anchored, - patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist, - ): + with patch.object( + manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=issuer_resolved) + ) as anchored: built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) anchored.assert_awaited_once_with("https://idp.example.com", "https://up.example.com/mcp") assert built.authorization_url == "https://idp.example.com/authorize" assert built.token_url == "https://idp.example.com/token" assert built.token_url != "https://attacker.example.com/steal" - # The issuer-anchored endpoints are never persisted into the endpoint columns, so a later - # build cannot treat them as authoritative stored values. - assert mock_persist.await_args.kwargs["is_issuer_anchored"] is True @pytest.mark.asyncio @pytest.mark.parametrize( @@ -1624,8 +1621,8 @@ class TestMCPServerManager: and PKCE verifier to the attacker (config-time RFC 9700 mix-up). The resource-driven scopes are kept, because scope selection is resource-driven (MCP Scope Selection Strategy) and scope inflation is bounded by the authorization server at consent (RFC 6749 §3.3), not by dropping - scopes on an endpoint mismatch. Both the in-memory merge and the persisted metadata drop only - the uncorroborated endpoints.""" + scopes on an endpoint mismatch. The gateway persists nothing, so the in-memory merge is the + entire behavior.""" manager = MCPServerManager() row = LiteLLM_MCPServerTable( server_id="manual-auth-url-3", @@ -1645,20 +1642,13 @@ class TestMCPServerManager: registration_url="https://attacker.example.com/register", scopes=["read", "admin"], ) - with ( - patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)), - patch.object(manager, "_persist_discovered_oauth_endpoints", new=AsyncMock()) as mock_persist, - ): + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=metadata)): built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) assert built.authorization_url == "https://idp.example.com/authorize" assert built.token_url is None assert built.registration_url is None assert built.scopes == ["read", "admin"] - persisted_metadata = mock_persist.await_args.kwargs["metadata"] - assert persisted_metadata.token_url is None - assert persisted_metadata.registration_url is None - assert persisted_metadata.scopes == ["read", "admin"] @pytest.mark.asyncio async def test_build_from_table_skips_discovery_when_all_upstream_oauth_fields_present(self): @@ -5586,388 +5576,300 @@ class TestMCPServerTimestamps: assert server.token_exchange_endpoint == "https://idp.example.com/token" @pytest.mark.asyncio - async def test_build_mcp_server_from_table_persists_discovered_obo_token_url(self): - """A DB-backed OBO server with no configured endpoint discovers token_url and must write it - back to the row, so the next rebuild skips discovery instead of re-running it every time.""" + async def test_discovery_never_writes_the_database(self): + """The #34985 regression, stated as the design invariant that fixes it: the gateway never + writes discovery results to the row. The OAuth columns and credentials.scopes carry admin + intent alone, so nothing the gateway learns can read back as an admin pin on a later build + (which is what anchored stamped servers fail-closed and 400ed /authorize). Discovery output + lives on the in-memory registry entry only, for oauth2 and OBO alike.""" manager = MCPServerManager() async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - assert server_url == "https://example.com/mcp" - assert allow_origin_fallback is False # OBO never guesses the origin return MCPOAuthMetadata( - scopes=None, - authorization_url=None, - token_url="https://discovered.example.com/token", - registration_url=None, - ) - - manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] - - record = LiteLLM_MCPServerTable( - server_id="obo-persist-1", - server_name="obo_persist", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"}, - ) - - update_mock = AsyncMock() - repo_instance = MagicMock() - repo_instance.table.update = update_mock - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) - - assert server.token_url == "https://discovered.example.com/token" - update_mock.assert_awaited_once() - assert update_mock.call_args.kwargs["where"] == {"server_id": "obo-persist-1"} - assert update_mock.call_args.kwargs["data"] == {"token_url": "https://discovered.example.com/token"} - - @pytest.mark.asyncio - async def test_persist_discovered_obo_token_url_skips_when_not_needed(self): - """The write-back fires only for an OBO server that discovered a new endpoint: a row that - already has token_url, a non-OBO auth_type, or a discovery that found nothing all no-op.""" - manager = MCPServerManager() - update_mock = AsyncMock() - repo_instance = MagicMock() - repo_instance.table.update = update_mock - - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - # already populated -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url="https://already.example.com/token", - discovered_token_url="https://new.example.com/token", - ) - # not an OBO server -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_token_url=None, - discovered_token_url="https://new.example.com/token", - ) - # discovery found nothing -> no write - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url=None, - discovered_token_url=None, - ) - - update_mock.assert_not_awaited() - - @pytest.mark.asyncio - async def test_persist_discovered_obo_token_url_is_best_effort(self): - """A write-back failure must not propagate; discovery just re-runs on the next build.""" - manager = MCPServerManager() - update_mock = AsyncMock(side_effect=Exception("db unavailable")) - repo_instance = MagicMock() - repo_instance.table.update = update_mock - - with ( - patch( - "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", - return_value=repo_instance, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_obo_token_url( - server_id="s", - auth_type=MCPAuth.oauth2_token_exchange, - existing_token_url=None, - discovered_token_url="https://new.example.com/token", - ) - - update_mock.assert_awaited_once() - - @pytest.mark.asyncio - async def test_build_mcp_server_from_table_persists_discovered_oauth_endpoints(self): - """A DB-backed oauth2 server with no configured endpoints discovers them and must write - authorization_url, token_url, and scopes back to the row; otherwise the resolved values - live only in memory and one failed re-discovery serves the 400 "authorization url is not configured" - from /authorize. registration_url must never be persisted because - _dcr_bridge_relays_client_registration keys off that column.""" - manager = MCPServerManager() - - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - assert allow_origin_fallback is True - return MCPOAuthMetadata( - scopes=["mcp.read", "mcp.write"], + scopes=["mcp.read"], authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", registration_url="https://idp.example.com/register", - ) - - manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] - - record = LiteLLM_MCPServerTable( - server_id="oauth-persist-1", - server_name="oauth_persist", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2, - oauth2_flow="authorization_code", - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - server = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) - - assert server.authorization_url == "https://idp.example.com/authorize" - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.server_id == "oauth-persist-1" - assert persisted.authorization_url == "https://idp.example.com/authorize" - assert persisted.token_url == "https://idp.example.com/token" - assert persisted.credentials == {"scopes": ["mcp.read", "mcp.write"]} - assert "registration_url" not in persisted.fields_set() - assert update_mcp_server_mock.call_args.kwargs["touched_by"] == "mcp_oauth_discovery" - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_guards(self): - """The write-back must no-op for non-discovery auth types, empty discovery, origin-fallback - guesses (never harden an inferred authorization server into configuration), and rows whose - fields are all already populated.""" - manager = MCPServerManager() - advertised = MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.api_key, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=advertised, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=None, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=advertised.model_copy(update={"from_origin_fallback": True}), - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url="https://configured.example.com/authorize", - existing_token_url="https://configured.example.com/token", - existing_scopes=["configured"], - metadata=advertised, - ) - - update_mcp_server_mock.assert_not_awaited() - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_only_fills_empty_fields(self): - """A row that already has token_url keeps it; only the missing authorization_url and - scopes are written, so admin-typed values always win over discovery.""" - manager = MCPServerManager() - - update_mcp_server_mock = AsyncMock() - with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url="https://configured.example.com/token", - existing_scopes=None, - metadata=MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - ), - ) - - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.authorization_url == "https://idp.example.com/authorize" - assert persisted.credentials == {"scopes": ["s1"]} - assert "token_url" not in persisted.fields_set() - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_writes_discovered_issuer_trust_on_first_use(self): - """A server with no configured issuer records the discovered issuer trust-on-first-use, so the - next rebuild anchors discovery on it (RFC 8414 §3.3) instead of re-trusting the resource. When - an issuer is already set (admin-typed or a prior discovery), it is never overwritten.""" - manager = MCPServerManager() - metadata = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - discovered_issuer="https://idp.example.com", - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer=None, - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=metadata, - ) - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer="https://admin-configured.example.com", - existing_authorization_url="https://admin-configured.example.com/authorize", - existing_token_url="https://admin-configured.example.com/token", - existing_scopes=["cfg"], - metadata=metadata, - ) - - assert update_mcp_server_mock.await_count == 1 - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert persisted.issuer == "https://idp.example.com" - - @pytest.mark.asyncio - async def test_persist_discovered_oauth_endpoints_does_not_persist_endpoints_for_issuer_anchored(self): - """For an issuer-anchored server the endpoints are re-derived from the §3.3-validated issuer - document every build, so they must NOT be written into the endpoint columns: persisting them - would make the next build see populated endpoints and treat them as authoritative stored - values, defeating the issuer-only invariant. Only the resource-driven scopes are persisted.""" - manager = MCPServerManager() - metadata = MCPOAuthMetadata( - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", - scopes=["read"], - ) - - update_mcp_server_mock = AsyncMock() - with ( - patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), - patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), - ): - await manager._persist_discovered_oauth_endpoints( - server_id="s", - auth_type=MCPAuth.oauth2, - existing_issuer="https://idp.example.com", - existing_authorization_url=None, - existing_token_url=None, - existing_scopes=None, - metadata=metadata, - is_issuer_anchored=True, - ) - - update_mcp_server_mock.assert_awaited_once() - persisted = update_mcp_server_mock.call_args.kwargs["data"] - assert "authorization_url" not in persisted.fields_set() - assert "token_url" not in persisted.fields_set() - assert persisted.credentials == {"scopes": ["read"]} - - @pytest.mark.asyncio - async def test_build_mcp_server_from_table_skips_persistence_for_temporary_servers(self): - """The session endpoint builds temporary servers whose server_id has no DB row; with - persist_discovered_endpoints=False neither the oauth2 nor the OBO write-back may fire.""" - manager = MCPServerManager() - - async def fake_discovery(server_url: str, *, allow_origin_fallback: bool = True, warn_when_no_metadata: bool = False): - return MCPOAuthMetadata( - scopes=["s1"], - authorization_url="https://idp.example.com/authorize", - token_url="https://idp.example.com/token", + discovered_issuer="https://idp.example.com", ) manager._descovery_metadata = fake_discovery # type: ignore[attr-defined] update_mcp_server_mock = AsyncMock() - obo_update_mock = AsyncMock() repo_instance = MagicMock() - repo_instance.table.update = obo_update_mock + repo_instance.table.update = AsyncMock() with ( - patch( - "litellm.proxy._experimental.mcp_server.db.update_mcp_server", - new=update_mcp_server_mock, - ), + patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=update_mcp_server_mock), patch( "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", return_value=repo_instance, ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), ): - oauth2_record = LiteLLM_MCPServerTable( - server_id="temp-oauth-1", - server_name="temp_oauth", - url="https://example.com/mcp", + for auth_type, flow in ((MCPAuth.oauth2, "authorization_code"), (MCPAuth.oauth2_token_exchange, None)): + record = LiteLLM_MCPServerTable( + server_id=f"no-write-{auth_type}", + server_name=f"no_write_{auth_type}", + url="https://example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + oauth2_flow=flow, + credentials={"client_id": "cid", "client_secret": "csec", "audience": "aud"}, + ) + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + assert built.token_url == "https://idp.example.com/token" + + update_mcp_server_mock.assert_not_awaited() + repo_instance.table.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_declared_endpoints_survive_a_failed_discovery(self): + """The reporter's configuration: explicit authorization_url/token_url/registration_url, + issuer left empty. With the gateway never stamping the issuer column, the server never turns + anchored, so the declared endpoints resolve on every build, including one whose discovery + fails entirely; /authorize keeps redirecting instead of serving the 400.""" + manager = MCPServerManager() + record = LiteLLM_MCPServerTable( + server_id="declared-1", + alias="declared", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + with patch.object(manager, "_descovery_metadata", new=AsyncMock(return_value=None)): + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + + assert built.issuer_is_anchored is False + assert built.authorization_url == "https://idp.example.com/authorize" + assert built.token_url == "https://idp.example.com/token" + assert built.registration_url == "https://idp.example.com/register" + + def test_flow_endpoints_missing_arms(self): + """The reload fast-path exemption's completeness rule. Interactive needs authorize+token, + client_credentials and OBO need token only, an OBO server with a configured exchange + endpoint never discovers and must not be sent into a rebuild loop, and non-OAuth auth types + are never unresolved.""" + assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", None, "https://idp/token") is True + assert ( + _flow_endpoints_missing(MCPAuth.oauth2, "authorization_code", "https://idp/auth", "https://idp/token") + is False + ) + assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, "https://idp/token") is False + assert _flow_endpoints_missing(MCPAuth.oauth2, "client_credentials", None, None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None) is True + assert _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, "https://idp/token") is False + assert ( + _flow_endpoints_missing(MCPAuth.oauth2_token_exchange, None, None, None, "https://idp/exchange") is False + ) + assert _flow_endpoints_missing(MCPAuth.api_key, None, None, None) is False + + def test_unresolved_check_uses_the_flow_judge_not_the_raw_column(self): + """A legacy row the startup backfill deliberately left unstamped (token_url plus client + credentials, no authorization_url: the ambiguous M2M shape) serves client_credentials at + request time via effective_oauth2_flow. The reload check must reach the same verdict, or the + row is classified as interactive-missing-endpoints and re-runs discovery on every reload + forever. A null-flow row without the M2M shape stays interactive and genuinely unresolved.""" + m2m_shaped = MCPServer( + server_id="null-flow-m2m", + name="null_flow_m2m", + server_name="null_flow_m2m", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow=None, + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + assert _oauth_endpoints_unresolved(m2m_shaped) is False + + interactive_unresolved = m2m_shaped.model_copy(update={"client_id": None, "client_secret": None}) + assert _oauth_endpoints_unresolved(interactive_unresolved) is True + + def test_dcr_bridge_relay_arm_needs_its_registration_endpoint(self): + """A dcr_bridge server with no admin-configured client can only register callers through the + upstream registration endpoint, so a partial discovery that resolved authorize and token but + not registration_endpoint leaves it silently degraded to the short-circuit arm. That counts as + unresolved so it keeps retrying. A bridge with a configured client_id uses the short-circuit + arm by design and is unaffected.""" + relay_arm = MCPServer( + server_id="bridge-partial", + name="bridge_partial", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + # dcr_bridge is only valid on the client-forwarded modes (see MCPServer.is_dcr_bridge) + auth_type=MCPAuth.oauth_delegate, + dcr_bridge=True, + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + registration_url=None, + ) + assert _oauth_endpoints_unresolved(relay_arm) is True + assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"registration_url": "https://idp/reg"})) is False + assert _oauth_endpoints_unresolved(relay_arm.model_copy(update={"client_id": "admin-client"})) is False + + def test_entra_obo_without_scopes_is_unresolved(self): + """entra_obo token exchange fails closed without a scope, and scopes can come from resource + discovery, so an entra_obo server that resolved its token endpoint but no scopes is still + unresolved for its flow. The default rfc8693 profile has no such requirement.""" + entra = MCPServer( + server_id="entra-noscope", + name="entra_noscope", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_profile="entra_obo", + token_url="https://idp.example.com/token", + scopes=None, + ) + assert _oauth_endpoints_unresolved(entra) is True + assert _oauth_endpoints_unresolved(entra.model_copy(update={"scopes": ["api://app/.default"]})) is False + assert _oauth_endpoints_unresolved(entra.model_copy(update={"token_exchange_profile": "rfc8693"})) is False + + def test_oauth_discovery_retry_backs_off_per_server(self): + """Without a cooldown the fast-path exemption re-runs the full discovery chain, and re-emits + the unresolved warning, on every reload forever for a server that can never resolve. Delay + doubles per consecutive failure up to the cap, a success clears the state so the next failure + starts from the base delay again, and the cooldown is per server.""" + manager = MCPServerManager() + + def unresolved(server_id): + return MCPServer( + server_id=server_id, + name=server_id, + url="https://up.example.com/mcp", transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - obo_record = LiteLLM_MCPServerTable( - server_id="temp-obo-1", - server_name="temp_obo", - url="https://example.com/mcp", - transport=MCPTransport.http, - auth_type=MCPAuth.oauth2_token_exchange, - credentials={"client_id": "cid", "client_secret": "csec"}, - ) - built_oauth2 = await manager.build_mcp_server_from_table( - oauth2_record, credentials_are_encrypted=False, persist_discovered_endpoints=False - ) - await manager.build_mcp_server_from_table( - obo_record, credentials_are_encrypted=False, persist_discovered_endpoints=False ) - assert built_oauth2.authorization_url == "https://idp.example.com/authorize" - update_mcp_server_mock.assert_not_awaited() - obo_update_mock.assert_not_awaited() + assert manager._oauth_discovery_retry_due("a") is True + + manager._record_oauth_discovery_outcome(unresolved("a")) + assert manager._oauth_discovery_retry_due("a") is False + assert manager._oauth_discovery_retry_due("b") is True, "cooldown must be per server" + + failures_before, _ = manager._oauth_discovery_retry_state["a"] + manager._record_oauth_discovery_outcome(unresolved("a")) + failures_after, _ = manager._oauth_discovery_retry_state["a"] + assert failures_after == failures_before + 1 + + # An elapsed cooldown lets the retry through, and the delay grows with the failure count + manager._oauth_discovery_retry_state["a"] = (1, time.monotonic() - 31.0) + assert manager._oauth_discovery_retry_due("a") is True + manager._oauth_discovery_retry_state["a"] = (5, time.monotonic() - 31.0) + assert manager._oauth_discovery_retry_due("a") is False + + resolved = unresolved("a").model_copy( + update={ + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + } + ) + manager._record_oauth_discovery_outcome(resolved) + assert "a" not in manager._oauth_discovery_retry_state + assert manager._oauth_discovery_retry_due("a") is True + + @pytest.mark.asyncio + async def test_reload_fast_path_retries_unresolved_oauth_servers(self): + """A server whose discovery failed must not be pinned broken by the updated_at fast path: + the next reload rebuilds it, retrying discovery on the normal cadence instead of waiting for + an unrelated config write. A resolved server with an unchanged row still takes the fast path, + so the exemption costs nothing in the steady state.""" + manager = MCPServerManager() + stamp = datetime.now() + row = LiteLLM_MCPServerTable( + server_id="retry-1", + server_name="retry_server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + created_at=stamp, + updated_at=stamp, + ) + + def entry(authorization_url, token_url): + return MCPServer( + server_id="retry-1", + name="retry_server", + server_name="retry_server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + authorization_url=authorization_url, + token_url=token_url, + updated_at=stamp, + ) + + raw_row = MagicMock() + raw_row.model_dump.return_value = row.model_dump() + repo_instance = MagicMock() + repo_instance.table.find_many = AsyncMock(return_value=[raw_row]) + + async def run_reload(previous_entry): + manager.registry = {"retry-1": previous_entry} + build_mock = AsyncMock(return_value=previous_entry) + with ( + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPServerRepository", + return_value=repo_instance, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch.object(manager, "build_mcp_server_from_table", new=build_mock), + ): + await manager.reload_servers_from_database() + return build_mock + + unresolved_build = await run_reload(entry(None, None)) + unresolved_build.assert_awaited_once() + + resolved_build = await run_reload(entry("https://idp.example.com/authorize", "https://idp.example.com/token")) + resolved_build.assert_not_awaited() + + @pytest.mark.asyncio + async def test_anchored_issuer_discarding_stored_endpoints_warns(self, caplog): + """An anchored server ignoring stored endpoint columns must say so: that state is exactly + what a row stamped by an earlier release looks like after upgrade, and the warning names the + remedy (clear the Issuer field) instead of leaving the 400 undiagnosable.""" + manager = MCPServerManager() + record = LiteLLM_MCPServerTable( + server_id="stamped-1", + alias="stamped_row", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="authorization_code", + issuer="https://idp.example.com", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + created_at=datetime.now(), + updated_at=datetime.now(), + ) + + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=None)), + caplog.at_level(logging.WARNING, logger="LiteLLM"), + ): + built = await manager.build_mcp_server_from_table(record, credentials_are_encrypted=False) + + assert built.issuer_is_anchored is True + assert built.authorization_url is None + assert "stamped_row" in caplog.text + assert "authorization_url, token_url" in caplog.text + assert "clear the Issuer" in caplog.text @pytest.mark.asyncio async def test_update_server_carries_forward_last_known_good_oauth_endpoints(self): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py new file mode 100644 index 00000000000..b6c946b95fa --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -0,0 +1,129 @@ +"""Tests for the one-time heal of issuer values a released version's discovery write-back stamped.""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy._experimental.mcp_server.oauth_issuer_stamp_backfill import ( + backfill_discovery_stamped_issuers, +) + + +def _row(**overrides): + fields = { + "server_id": "srv-1", + "alias": "srv_one", + "server_name": "srv_one", + "auth_type": "oauth2", + "issuer": "https://idp.example.com", + "authorization_url": "https://idp.example.com/authorize", + "token_url": "https://idp.example.com/token", + "registration_url": None, + "updated_by": "mcp_oauth_discovery", + } + fields.update(overrides) + return SimpleNamespace(**fields) + + +def _prisma(rows): + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows) + prisma_client.db.litellm_mcpservertable.update = AsyncMock() + return prisma_client + + +@pytest.mark.asyncio +async def test_clears_the_stamp_and_records_its_own_actor(): + """The GH #34985 row: discovery wrote the issuer, so the server reads as issuer-anchored and its + configured endpoints are ignored. Clearing the stamp makes them apply again. The heal records its + own actor, which is also what makes it idempotent: the row no longer matches the discovery-actor + filter, so it is never reconsidered on a later boot.""" + prisma_client = _prisma([_row()]) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + + call = prisma_client.db.litellm_mcpservertable.update.call_args + assert call.kwargs["where"] == {"server_id": "srv-1"} + assert call.kwargs["data"]["issuer"] is None + assert call.kwargs["data"]["updated_by"] == "mcp_oauth_issuer_stamp_backfill" + + where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"] + assert where["updated_by"] == "mcp_oauth_discovery" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "overrides, reason", + [ + ({"updated_by": "some-admin@example.com"}, "an admin was the last writer, so the pin is theirs"), + ({"issuer": None}, "nothing to heal"), + ({"issuer": " "}, "blank issuer is not a pin"), + ( + {"authorization_url": None, "token_url": None, "registration_url": None}, + "issuer set with no configured endpoints is the canonical shape of a deliberate pin, and " + "there is nothing configured for anchoring to discard anyway", + ), + ( + {"authorization_url": "https://other-idp.example.com/authorize", "token_url": None}, + "endpoints addressing a different authority than the issuer are an intent a clear would " + "discard, so the row is warned about rather than healed", + ), + ( + {"issuer": "https://pinned.example.com"}, + "same shape from the other side: a pinned issuer whose origin differs from the configured " + "endpoints cannot have been derived from them by discovery", + ), + ], +) +async def test_leaves_rows_alone_that_do_not_carry_the_defect_signature(overrides, reason): + """updated_by records only the most recent writer and no audit trail says which field it touched, + so the heal is deliberately narrow: it fires only on the full signature of the defect. Every + exclusion here protects a row whose issuer may be a deliberate admin pin.""" + prisma_client = _prisma([_row(**overrides)]) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 0, reason + prisma_client.db.litellm_mcpservertable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_heals_across_url_forms_that_denote_the_same_origin(): + """Origin comparison runs through the shared canonicalizer, so a default port or host casing + difference between the stamped issuer and the endpoints an admin typed does not make a #34985 row + look like a deliberate pin at a different authority.""" + prisma_client = _prisma( + [ + _row( + issuer="https://IDP.example.com:443", + authorization_url="https://idp.example.com/authorize", + token_url="https://idp.example.com/token", + ) + ] + ) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + + +@pytest.mark.asyncio +async def test_query_is_scoped_to_auth_types_where_an_issuer_anchors(): + """Only the discovery auth types read an issuer as a trust anchor; clearing it elsewhere would be + an unrelated mutation.""" + prisma_client = _prisma([]) + + await backfill_discovery_stamped_issuers(prisma_client) + + where = prisma_client.db.litellm_mcpservertable.find_many.call_args.kwargs["where"] + assert set(where["auth_type"]["in"]) == {"oauth2", "true_passthrough", "oauth_delegate"} + + +@pytest.mark.asyncio +async def test_a_failed_row_does_not_abort_the_rest(): + """Per-row best effort: one write failure must not leave later rows unhealed, and the next boot + retries the failed one since its updated_by is unchanged.""" + prisma_client = _prisma([_row(server_id="bad"), _row(server_id="good")]) + prisma_client.db.litellm_mcpservertable.update = AsyncMock( + side_effect=[Exception("write failed"), MagicMock()] + ) + + assert await backfill_discovery_stamped_issuers(prisma_client) == 1 + assert prisma_client.db.litellm_mcpservertable.update.await_count == 2 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx index 8dffc80a70e..5650bd1d7e4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx @@ -190,7 +190,7 @@ const OAuthFormFields: React.FC = ({ label={ } name="issuer" From 581f5c319e30c23179d12cfdb6d765484b98f83d Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 29 Jul 2026 18:25:05 -0700 Subject: [PATCH 46/54] feat(cli): read base_url from persistent config file (#35015) * feat(cli): read base_url from persistent config file Adds a lite config command group (set/get/unset) backed by ~/.litellm/config.json so users no longer need to export LITELLM_PROXY_URL in every shell session. Resolution order is --base-url flag, then LITELLM_PROXY_URL, then the config file, then the localhost default. A config-file base_url counts as an explicit server choice for lite auth print-token, matching the env var semantics it replaces. * fix(cli): harden config persistence after review feedback Rejects base_url values containing a query string or fragment, including bare trailing ? or # which parse as empty but still corrupt every joined request URL. Writes config.json and token.json atomically through a shared write_private_json helper (0600 at creation, fsync, os.replace) so an interrupted save can no longer truncate the file or leave it world-readable. Warns on stderr when an existing config file is invalid instead of silently ignoring it, including invalid UTF-8. Resolves the eager --version flag through the same env, config file, default chain as every other command, and reads the config file once per invocation so base_url and base_url_explicit always come from the same snapshot. * fix(cli): resolve --version after option parsing The eager --version callback ran before --base-url and --api-key were parsed, so it could not see an explicitly named server. Combined with the env fallback added for config-file support, that sent the resolved API key to whichever server the config file pointed at even when the user named a different one on the command line. Making the flag a normal option and handling it in the group callback gives the version request the same flag, env, config, default precedence as every other command, and lets the stored-token lookup stay origin-checked. --- litellm/proxy/client/cli/README.md | 25 +- litellm/proxy/client/cli/commands/auth.py | 8 +- litellm/proxy/client/cli/commands/config.py | 108 +++++++ .../proxy/client/cli/commands/private_json.py | 20 ++ litellm/proxy/client/cli/main.py | 35 ++- .../proxy/client/cli/test_auth_commands.py | 133 +++++++- .../proxy/client/cli/test_config_commands.py | 284 ++++++++++++++++++ .../proxy/client/cli/test_global_options.py | 162 +++++++++- 8 files changed, 726 insertions(+), 49 deletions(-) create mode 100644 litellm/proxy/client/cli/commands/config.py create mode 100644 litellm/proxy/client/cli/commands/private_json.py create mode 100644 tests/test_litellm/proxy/client/cli/test_config_commands.py diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 2ad8a08b8c3..de9d38963c1 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -10,11 +10,32 @@ uv tool install 'litellm[proxy]' ## Configuration -The CLI can be configured using environment variables or command-line options: +The CLI can be configured using environment variables, command-line options, or a persistent config file: - `LITELLM_PROXY_URL`: Base URL of the LiteLLM proxy server (default: http://localhost:4000) - `LITELLM_PROXY_API_KEY`: API key for authentication +To stop exporting `LITELLM_PROXY_URL` in every shell session, store the proxy URL once in `~/.litellm/config.json`: + +```bash +lite config set base_url https://your-proxy.example.com +``` + +Manage the stored config with: + +```bash +lite config get base_url # print the stored value +lite config get # print all stored config +lite config unset base_url # remove the stored value +``` + +The base URL is resolved in this order of precedence: + +1. `--base-url` command-line option +2. `LITELLM_PROXY_URL` environment variable +3. `base_url` from `~/.litellm/config.json` +4. `http://localhost:4000` + ## Global Options - `--version`, `-v`: Print the LiteLLM Proxy client and server version and exit. @@ -581,6 +602,8 @@ The CLI respects the following environment variables: - `LITELLM_PROXY_URL`: Base URL of the proxy server - `LITELLM_PROXY_API_KEY`: API key for authentication +`LITELLM_PROXY_URL` takes precedence over a `base_url` stored via `lite config set`, and the `--base-url` option overrides both. See the Configuration section for the full precedence order. + ## Examples 1. List all models in table format: diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 61495403407..970d801dc6d 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -15,6 +15,8 @@ from rich.table import Table from litellm.constants import CLI_JWT_EXPIRATION_HOURS from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh +from .private_json import write_private_json + # Token storage utilities def get_token_file_path() -> str: @@ -27,11 +29,7 @@ def get_token_file_path() -> str: def save_token(token_data: Dict[str, Any]) -> None: """Save token data to file""" - token_file = get_token_file_path() - with open(token_file, "w") as f: - json.dump(token_data, f, indent=2) - # Set file permissions to be readable only by owner - os.chmod(token_file, 0o600) + write_private_json(get_token_file_path(), token_data) def load_token() -> Optional[Dict[str, Any]]: diff --git a/litellm/proxy/client/cli/commands/config.py b/litellm/proxy/client/cli/commands/config.py new file mode 100644 index 00000000000..851a6c11529 --- /dev/null +++ b/litellm/proxy/client/cli/commands/config.py @@ -0,0 +1,108 @@ +import json +import os +import sys +from collections.abc import Mapping +from pathlib import Path +from urllib.parse import urlparse + +import click +from pydantic import TypeAdapter + +from .private_json import write_private_json + +ALLOWED_CONFIG_KEYS: tuple[str, ...] = ("base_url",) + +_config_adapter: TypeAdapter[Mapping[str, str]] = TypeAdapter(Mapping[str, str]) + + +def get_config_file_path() -> str: + """Get the path to the persistent CLI config file""" + home_dir = Path.home() + config_dir = home_dir / ".litellm" + return str(config_dir / "config.json") + + +def load_config() -> Mapping[str, str]: + """Load CLI config from file; returns {} if missing or unreadable""" + try: + config_file = get_config_file_path() + except RuntimeError: + return {} + if not os.path.exists(config_file): + return {} + try: + with open(config_file, "r") as f: + return _config_adapter.validate_python(json.load(f)) + except (OSError, ValueError) as e: + click.echo(f"Warning: ignoring invalid config file {config_file}: {e}", err=True) + return {} + + +def save_config(config: Mapping[str, str]) -> None: + """Save CLI config to file""" + write_private_json(get_config_file_path(), config) + + +def get_config_value(key: str) -> str | None: + """Get a single value from the persistent CLI config""" + return load_config().get(key) + + +@click.group(name="config") +def config_commands() -> None: + """Manage persistent CLI configuration (~/.litellm/config.json)""" + + +@config_commands.command(name="set") +@click.argument("key") +@click.argument("value") +def set_config(key: str, value: str) -> None: + """Set a config KEY to VALUE (e.g. `lite config set base_url https://your-proxy.example.com`)""" + if key not in ALLOWED_CONFIG_KEYS: + raise click.UsageError(f"Unknown config key '{key}'. Allowed keys: {', '.join(ALLOWED_CONFIG_KEYS)}") + + if key == "base_url": + parsed = urlparse(value) + if parsed.scheme not in ("http", "https") or not parsed.netloc: + raise click.UsageError("base_url must be a full http:// or https:// URL including a host") + if "?" in value or "#" in value: + raise click.UsageError("base_url must not include a query string or fragment") + + normalized_value = value.rstrip("/") + save_config({**load_config(), key: normalized_value}) + click.echo(f"Set {key} = {normalized_value} in {get_config_file_path()}") + + +@config_commands.command(name="get") +@click.argument("key", required=False) +def get_config(key: str | None) -> None: + """Print the value of KEY, or all stored config when KEY is omitted""" + config = load_config() + + if key is not None: + value = config.get(key) + if value is None: + click.echo(f"{key} is not set", err=True) + sys.exit(1) + click.echo(value) + return + + if not config: + click.echo("(no config set)") + return + + for entry_key, entry_value in config.items(): + click.echo(f"{entry_key} = {entry_value}") + + +@config_commands.command(name="unset") +@click.argument("key") +def unset_config(key: str) -> None: + """Remove KEY from the config file""" + config = load_config() + if key not in config: + click.echo(f"{key} was not set") + return + + save_config({k: v for k, v in config.items() if k != key}) + click.echo(f"Removed {key} from {get_config_file_path()}") diff --git a/litellm/proxy/client/cli/commands/private_json.py b/litellm/proxy/client/cli/commands/private_json.py new file mode 100644 index 00000000000..70aac0c6de0 --- /dev/null +++ b/litellm/proxy/client/cli/commands/private_json.py @@ -0,0 +1,20 @@ +import json +import os +import tempfile +from collections.abc import Mapping +from pathlib import Path + + +def write_private_json(path: str, data: Mapping[str, object]) -> None: + """Atomically write JSON to path with owner-only permissions (0600)""" + parent = Path(path).parent + parent.mkdir(parents=True, exist_ok=True) + fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-", suffix=".json") + try: + with os.fdopen(fd, "w") as f: + json.dump(data, f, indent=2) + f.flush() + os.fsync(f.fileno()) + os.replace(tmp_path, path) + finally: + Path(tmp_path).unlink(missing_ok=True) diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index e641956b2c5..24e5cdf747b 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -11,6 +11,7 @@ from .commands.agents import agent_commands from .commands.auth import auth_group, get_stored_api_key, login, logout, whoami from .commands.autoroute.commands import autoroute_group from .commands.chat import chat +from .commands.config import config_commands, get_config_value from .commands.credentials import credentials from .commands.encryption import encryption from .commands.http import http @@ -45,27 +46,16 @@ def print_version(base_url: str, api_key: Optional[str]): @click.option( "--version", "-v", + "show_version", is_flag=True, - is_eager=True, - expose_value=False, help="Show the LiteLLM Proxy CLI and server version and exit.", - callback=lambda ctx, param, value: ( - ( - print_version( - ctx.params.get("base_url") or "http://localhost:4000", - ctx.params.get("api_key"), - ) - or ctx.exit() - ) - if value and not ctx.resilient_parsing - else None - ), ) @click.option( "--base-url", envvar="LITELLM_PROXY_URL", show_envvar=True, - default="http://localhost:4000", + default=None, + show_default="base_url from `lite config`, else http://localhost:4000", help="Base URL of the LiteLLM proxy server", ) @click.option( @@ -75,13 +65,16 @@ def print_version(base_url: str, api_key: Optional[str]): help="API key for authentication", ) @click.pass_context -def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: +def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: Optional[str]) -> None: """LiteLLM Proxy CLI - Manage your LiteLLM proxy server""" ctx.ensure_object(dict) + stored_base_url = get_config_value("base_url") + base_url_provided = base_url is not None + # Normalize once here so every downstream command (login, agents, http, ...) can safely # do f"{base_url}/some/path" without producing a double slash. - base_url = base_url.rstrip("/") + base_url = ((stored_base_url or "http://localhost:4000") if base_url is None else base_url).rstrip("/") # If no API key provided via flag or environment variable, try to load from saved token. # Pass base_url so we only use the stored key when it was issued for this server. @@ -94,8 +87,13 @@ def cli(ctx: click.Context, base_url: str, api_key: Optional[str]) -> None: # apiKeyHelper is invoked bare (no flags) -- commands that must work # unattended (print-token) need to tell "user didn't say" apart from # "user said localhost:4000 on purpose" so they can fall back to - # whatever server the stored token was actually issued for. - ctx.obj["base_url_explicit"] = ctx.get_parameter_source("base_url") != click.core.ParameterSource.DEFAULT + # whatever server the stored token was actually issued for. A base_url + # saved via `lite config set` counts as the user saying it. + ctx.obj["base_url_explicit"] = base_url_provided or bool(stored_base_url) + + if show_version: + print_version(base_url, api_key) + ctx.exit() # If no subcommand was invoked, start interactive mode if ctx.invoked_subcommand is None: @@ -141,6 +139,7 @@ cli.add_command(down) cli.add_command(model_groups) # Add the autoroute command group (QA auto-routing against your real proxy) cli.add_command(autoroute_group, name="autoroute") +cli.add_command(config_commands) if __name__ == "__main__": diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 2fbc9c5c82f..f0aa49ff123 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -1,5 +1,6 @@ import json import os +import stat import sys import time from pathlib import Path @@ -12,6 +13,7 @@ import pytest from click.testing import CliRunner from litellm.constants import CLI_JWT_EXPIRATION_HOURS +from litellm.proxy.client.cli import cli from litellm.proxy.client.cli.commands.auth import ( clear_token, get_stored_api_key, @@ -201,31 +203,22 @@ class TestTokenUtilities: mock_mkdir.assert_called_once_with(exist_ok=True) - def test_save_token(self): + def test_save_token(self, tmp_path): """Test saving token data to file""" token_data = { "key": "test-key", "user_id": "test-user", "timestamp": 1234567890, } + token_file = tmp_path / "token.json" - with ( - patch("builtins.open", mock_open()) as mock_file, - patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path, - patch("os.chmod") as mock_chmod, - ): - mock_path.return_value = "/test/path/token.json" + with patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path: + mock_path.return_value = str(token_file) save_token(token_data) - mock_file.assert_called_once_with("/test/path/token.json", "w") - mock_file().write.assert_called() - mock_chmod.assert_called_once_with("/test/path/token.json", 0o600) - - # Verify JSON content was written correctly - written_content = "".join(call[0][0] for call in mock_file().write.call_args_list) - parsed_content = json.loads(written_content) - assert parsed_content == token_data + assert json.loads(token_file.read_text()) == token_data + assert stat.S_IMODE(token_file.stat().st_mode) == 0o600 def test_load_token_success(self): """Test loading token data from file successfully""" @@ -808,7 +801,8 @@ class TestPrintTokenCommand: since there is no explicit target to check it against. `--base-url`/ `LITELLM_PROXY_URL` only enforces the match when a caller explicitly passes it (tracked via ctx.obj["base_url_explicit"], set by the `cli` - group from click's ParameterSource). + group from click's ParameterSource); a base_url saved via + `lite config set` counts as explicit too. """ def setup_method(self): @@ -928,3 +922,110 @@ class TestPrintTokenCommand: assert "sk-stale-key" not in result.output assert "lite login" in result.output mock_post.assert_not_called() + + +def _write_home_json(home: Path, filename: str, payload: dict[str, object]) -> None: + litellm_dir = home / ".litellm" + litellm_dir.mkdir(exist_ok=True) + (litellm_dir / filename).write_text(json.dumps(payload)) + + +class TestPrintTokenWithConfigFile: + """A config-file base_url is a drop-in replacement for exporting + LITELLM_PROXY_URL, so print-token must treat it as an explicit server + choice: a token minted for a different proxy is never handed out.""" + + @pytest.fixture + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + def test_config_base_url_mismatch_fails_closed(self, isolated_home): + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": "https://server-b.example.com"}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 1 + assert "sk-issued-for-a" not in result.output + assert "Not authenticated for this server" in result.output + + def test_config_base_url_match_prints_token(self, isolated_home): + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": "https://server-a.example.com"}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + def test_empty_config_base_url_treated_as_unset(self, isolated_home): + """A hand-edited config.json with base_url "" must behave like no config at all: + base_url falls back to the default AND explicitness stays False.""" + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + _write_home_json(isolated_home, "config.json", {"base_url": ""}) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + def test_bare_invocation_without_config_file_unchanged(self, isolated_home): + """No config file means base_url_explicit stays False, so the stored + token's own server is trusted (pre-config behavior must not regress).""" + _write_home_json( + isolated_home, + "token.json", + {"base_url": "https://server-a.example.com", "key": "sk-issued-for-a", "timestamp": time.time()}, + ) + + result = CliRunner().invoke(cli, ["auth", "print-token"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "sk-issued-for-a" + + +class TestSaveTokenPrivateWrite: + """token.json holds the real API key: it must never be world-readable at any + instant, and a failed write must not destroy the previously stored token.""" + + @pytest.fixture + def isolated_home(self, monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + def test_save_token_owner_only_permissions_and_no_temp_leftovers(self, isolated_home): + save_token({"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890}) + + token_file = isolated_home / ".litellm" / "token.json" + assert json.loads(token_file.read_text()) == {"key": "sk-secret", "user_id": "u-1", "timestamp": 1234567890} + assert stat.S_IMODE(token_file.stat().st_mode) == 0o600 + assert list(token_file.parent.glob(".tmp-*")) == [] + + def test_save_token_failure_mid_write_preserves_existing_token(self, isolated_home): + _write_home_json(isolated_home, "token.json", {"key": "sk-original", "timestamp": 1234567890}) + token_file = isolated_home / ".litellm" / "token.json" + + with pytest.raises(TypeError): + save_token({"key": object()}) + + assert json.loads(token_file.read_text()) == {"key": "sk-original", "timestamp": 1234567890} + assert list(token_file.parent.glob(".tmp-*")) == [] diff --git a/tests/test_litellm/proxy/client/cli/test_config_commands.py b/tests/test_litellm/proxy/client/cli/test_config_commands.py new file mode 100644 index 00000000000..698d6188768 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_config_commands.py @@ -0,0 +1,284 @@ +import json +import os +import stat +import sys +from pathlib import Path + +import pytest +from click.testing import CliRunner + +sys.path.insert(0, os.path.abspath("../../..")) + + +from litellm.proxy.client.cli import cli +from litellm.proxy.client.cli.commands.config import ( + get_config_file_path, + get_config_value, + load_config, + save_config, +) +from litellm.proxy.client.cli.commands.private_json import write_private_json + + +@pytest.fixture +def cli_runner(): + return CliRunner() + + +@pytest.fixture +def isolated_home(monkeypatch, tmp_path): + """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm.""" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + +def _config_path(home: Path) -> Path: + return home / ".litellm" / "config.json" + + +def _raise_home_unresolvable() -> str: + raise RuntimeError("Could not determine home directory.") + + +class TestConfigSet: + @pytest.mark.parametrize( + "value", + ["https://your-proxy.example.com", "http://your-proxy.example.com:8080"], + ) + def test_set_stores_value_with_owner_only_permissions(self, cli_runner, isolated_home, value): + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code == 0 + config_file = _config_path(isolated_home) + assert json.loads(config_file.read_text()) == {"base_url": value} + assert stat.S_IMODE(config_file.stat().st_mode) == 0o600 + assert str(config_file) in result.output + + def test_set_strips_trailing_slash(self, cli_runner, isolated_home): + """Downstream commands join paths onto base_url; a stored trailing + slash would produce double slashes in every request URL.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com/"]) + + assert result.exit_code == 0 + assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://your-proxy.example.com"} + + def test_set_unknown_key_rejected_and_names_allowed_keys(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "set", "api_key", "sk-secret"]) + + assert result.exit_code != 0 + assert "base_url" in result.output + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize("value", ["your-proxy.example.com", "ftp://your-proxy.example.com"]) + def test_set_base_url_without_http_scheme_rejected(self, cli_runner, isolated_home, value): + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert "http" in result.output + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize("value", ["https://", "http://", "https:///some-path"]) + def test_set_base_url_without_host_rejected(self, cli_runner, isolated_home, value): + """rstrip("/") would otherwise persist a bare "https:" that breaks every later request.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert not _config_path(isolated_home).exists() + + @pytest.mark.parametrize( + "value", + [ + "https://proxy.example.com?env=prod", + "https://proxy.example.com#prod", + "https://proxy.example.com/?", + "https://proxy.example.com/#", + ], + ) + def test_set_base_url_with_query_or_fragment_rejected(self, cli_runner, isolated_home, value): + """Downstream commands join paths onto base_url; a stored query string or + fragment would silently corrupt every request URL built from it. Bare + trailing '?' / '#' parse as EMPTY query/fragment yet still break every + joined path, so rejection must key off the raw characters.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", value]) + + assert result.exit_code != 0 + assert "query" in result.output or "fragment" in result.output + assert not _config_path(isolated_home).exists() + + def test_set_base_url_with_path_prefix_accepted(self, cli_runner, isolated_home): + """Proxies are commonly served under a path prefix; the query/fragment + rejection must not over-reach into legitimate paths.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://proxy.example.com/litellm"]) + + assert result.exit_code == 0 + assert json.loads(_config_path(isolated_home).read_text()) == {"base_url": "https://proxy.example.com/litellm"} + + def test_set_leaves_no_temp_files_behind(self, cli_runner, isolated_home): + """The atomic write goes through a .tmp-* sibling; it must be renamed away, + never abandoned next to the config.""" + result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + + assert result.exit_code == 0 + config_file = _config_path(isolated_home) + assert stat.S_IMODE(config_file.stat().st_mode) == 0o600 + assert list(config_file.parent.glob(".tmp-*")) == [] + + +class TestConfigGet: + def test_get_prints_only_the_value(self, cli_runner, isolated_home): + """stdout must be exactly the value so scripts can do URL=$(lite config get base_url).""" + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "get", "base_url"]) + + assert result.exit_code == 0 + assert result.stdout.strip() == "https://your-proxy.example.com" + + def test_get_unset_key_exits_one_with_stderr_message(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "get", "base_url"]) + + assert result.exit_code == 1 + assert result.stdout.strip() == "" + assert result.stderr != "" + + def test_get_without_key_lists_entries(self, cli_runner, isolated_home): + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "base_url = https://your-proxy.example.com" in result.output + + def test_get_without_key_when_nothing_set(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "no config" in result.output.lower() + + +class TestConfigUnset: + def test_unset_removes_key_from_file(self, cli_runner, isolated_home): + set_result = cli_runner.invoke(cli, ["config", "set", "base_url", "https://your-proxy.example.com"]) + assert set_result.exit_code == 0 + + result = cli_runner.invoke(cli, ["config", "unset", "base_url"]) + + assert result.exit_code == 0 + assert "base_url" not in load_config() + assert cli_runner.invoke(cli, ["config", "get", "base_url"]).exit_code == 1 + + def test_unset_missing_key_is_idempotent(self, cli_runner, isolated_home): + result = cli_runner.invoke(cli, ["config", "unset", "base_url"]) + + assert result.exit_code == 0 + assert "not set" in result.output.lower() + + +class TestConfigHelpers: + def test_get_config_file_path_under_home(self, isolated_home): + assert get_config_file_path() == str(isolated_home / ".litellm" / "config.json") + + def test_load_config_missing_file_returns_empty(self, isolated_home): + assert load_config() == {} + + def test_home_unresolvable_does_not_crash_cli(self, cli_runner, isolated_home, monkeypatch): + """Path.home() raises RuntimeError in HOME-less containers; invocations that + never needed the home dir (--api-key supplied) must keep working.""" + monkeypatch.setattr( + "litellm.proxy.client.cli.commands.config.get_config_file_path", + _raise_home_unresolvable, + ) + + assert load_config() == {} + + result = cli_runner.invoke(cli, ["--api-key", "sk-test", "config", "get"]) + assert result.exit_code == 0 + assert "(no config set)" in result.output + + @pytest.mark.parametrize( + "content", + [ + "{not json", + '{"base_url": 123}', + '["https://your-proxy.example.com"]', + '"https://your-proxy.example.com"', + ], + ) + def test_load_config_invalid_content_returns_empty(self, isolated_home, content): + """A corrupt or wrongly-shaped config file must degrade to defaults, never crash the CLI.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_text(content) + + assert load_config() == {} + + def test_load_config_invalid_utf8_returns_empty(self, isolated_home): + """json.load raises UnicodeDecodeError (a ValueError but not a JSONDecodeError) + on undecodable bytes; before catching ValueError this crashed every CLI + invocation, including the `config set` needed to repair the file.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_bytes(b"\xff\xfe{}") + + assert load_config() == {} + + def test_save_config_round_trip_creates_dir_and_restricts_permissions(self, isolated_home): + save_config({"base_url": "https://your-proxy.example.com"}) + + assert load_config() == {"base_url": "https://your-proxy.example.com"} + assert stat.S_IMODE(_config_path(isolated_home).stat().st_mode) == 0o600 + + def test_get_config_value_unset_then_set(self, isolated_home): + assert get_config_value("base_url") is None + + save_config({"base_url": "https://your-proxy.example.com"}) + + assert get_config_value("base_url") == "https://your-proxy.example.com" + + def test_corrupt_config_file_warns_on_stderr_but_command_succeeds(self, cli_runner, isolated_home): + """Silently ignoring a broken config file leaves users debugging why their + stored base_url stopped applying; the CLI must keep working but say why.""" + config_file = _config_path(isolated_home) + config_file.parent.mkdir(parents=True, exist_ok=True) + config_file.write_text("{not json") + + result = cli_runner.invoke(cli, ["config", "get"]) + + assert result.exit_code == 0 + assert "Warning: ignoring invalid config file" in result.stderr + + +class TestWritePrivateJson: + def test_failed_write_preserves_previous_file_and_removes_temp(self, tmp_path): + """json.dump can fail partway through serializing; writing to a temp file + and renaming keeps the previous file intact through a crash mid-write.""" + target = tmp_path / "config.json" + original = '{"base_url": "https://original.example.com"}' + target.write_text(original) + + with pytest.raises(TypeError): + write_private_json(str(target), {"bad": object()}) + + assert target.read_text() == original + assert list(tmp_path.glob(".tmp-*")) == [] + + def test_interrupted_write_removes_temp_file(self, tmp_path, monkeypatch): + """Ctrl-C is BaseException, which `except Exception` misses; an interrupt + mid-write must not abandon a .tmp-* file next to the config forever.""" + + def _interrupt(*args: object, **kwargs: object) -> None: + raise KeyboardInterrupt() + + monkeypatch.setattr("litellm.proxy.client.cli.commands.private_json.json.dump", _interrupt) + target = tmp_path / "config.json" + + with pytest.raises(KeyboardInterrupt): + write_private_json(str(target), {"base_url": "https://your-proxy.example.com"}) + + assert not target.exists() + assert list(tmp_path.glob(".tmp-*")) == [] diff --git a/tests/test_litellm/proxy/client/cli/test_global_options.py b/tests/test_litellm/proxy/client/cli/test_global_options.py index 8df763d35c2..9995cb1bca5 100644 --- a/tests/test_litellm/proxy/client/cli/test_global_options.py +++ b/tests/test_litellm/proxy/client/cli/test_global_options.py @@ -1,4 +1,5 @@ # stdlib imports +import json import os import sys from pathlib import Path @@ -7,9 +8,7 @@ from unittest.mock import Mock, patch import pytest from click.testing import CliRunner -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path import litellm.proxy.client.cli @@ -71,13 +70,9 @@ def test_base_url_trailing_slash_normalized(cli_runner): ) as mock_post, patch("requests.get", side_effect=ValueError("stop after start request")), ): - cli_runner.invoke( - cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"] - ) + cli_runner.invoke(cli, ["--base-url", "https://gateway.litellm-sandbox.ai/", "login"]) - mock_post.assert_called_once_with( - "https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10 - ) + mock_post.assert_called_once_with("https://gateway.litellm-sandbox.ai/sso/cli/start", timeout=10) def test_cli_version_command(cli_runner): @@ -94,3 +89,152 @@ def test_cli_version_command(cli_runner): assert f"LiteLLM Proxy CLI Version: {litellm_version}" in result.output assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output assert "LiteLLM Proxy Server Version: 1.2.3" in result.output + + +@pytest.fixture +def isolated_home(monkeypatch, tmp_path): + """Point HOME at tmp_path so tests never touch the developer's real ~/.litellm.""" + monkeypatch.setenv("HOME", str(tmp_path)) + monkeypatch.setenv("USERPROFILE", str(tmp_path)) + monkeypatch.delenv("LITELLM_PROXY_URL", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + return tmp_path + + +def _write_config_file(home: Path, config: dict[str, str]) -> None: + config_dir = home / ".litellm" + config_dir.mkdir(exist_ok=True) + (config_dir / "config.json").write_text(json.dumps(config)) + + +def _invoke_version(cli_runner: CliRunner, *args: str): + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + return cli_runner.invoke(cli, [*args, "version"]) + + +def test_base_url_read_from_config_file(cli_runner, isolated_home): + """base_url precedence: flag > env > config file > default.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output + + +def test_env_var_beats_config_file_base_url(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output + + +def test_base_url_flag_beats_env_var_and_config_file(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + result = _invoke_version(cli_runner, "--base-url", "http://flag-proxy.example.com:9000") + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://flag-proxy.example.com:9000" in result.output + + +def test_default_base_url_unchanged_without_config_file(cli_runner, isolated_home): + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output + + +def test_corrupt_config_file_falls_back_to_default(cli_runner, isolated_home): + """A corrupt config file must never crash the CLI. Exactly one warning proves + the config file is read once per invocation, not once per lookup.""" + config_dir = isolated_home / ".litellm" + config_dir.mkdir(exist_ok=True) + (config_dir / "config.json").write_text("{not json") + + result = _invoke_version(cli_runner) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://localhost:4000" in result.output + assert result.stderr.count("Warning: ignoring invalid config file") == 1 + + +def test_empty_base_url_flag_is_not_treated_as_unset(cli_runner, isolated_home): + """`--base-url ""` explicitly provided an (empty) value; falling back to the + config file or localhost would silently redirect auth-sensitive commands.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + result = _invoke_version(cli_runner, "--base-url", "") + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL:" not in result.output + + +def test_version_flag_reads_config_file_base_url(cli_runner, isolated_home): + """--version resolves through the same precedence chain as every other command.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://config-proxy.example.com" in result.output + + +def test_version_flag_prefers_env_var_over_config_file(cli_runner, isolated_home, monkeypatch): + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_URL", "http://env-proxy.example.com:5000") + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: http://env-proxy.example.com:5000" in result.output + + +def test_version_flag_prefers_explicit_base_url_over_config_file(cli_runner, isolated_home): + """An eager --version could not see the flag and silently queried the config + server instead of the one the user named.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + + with patch( + "litellm.proxy.client.health.HealthManagementClient.get_server_version", + return_value="1.2.3", + ): + result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"]) + + assert result.exit_code == 0 + assert "LiteLLM Proxy Server URL: https://flag-proxy.example.com" in result.output + assert "config-proxy.example.com" not in result.output + + +def test_version_flag_never_sends_api_key_to_unnamed_server(cli_runner, isolated_home, monkeypatch): + """The version request carries a bearer token; it must reach only the server the + user named, never whichever host happens to sit in the config file.""" + _write_config_file(isolated_home, {"base_url": "https://config-proxy.example.com"}) + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-intended-for-flag-proxy") + + with patch("litellm.proxy.client.http_client.requests.request") as mock_request: + mock_request.return_value.json.return_value = {"litellm_version": "1.2.3"} + mock_request.return_value.raise_for_status.return_value = None + result = cli_runner.invoke(cli, ["--base-url", "https://flag-proxy.example.com", "--version"]) + + assert result.exit_code == 0 + requested_urls = [call.kwargs["url"] for call in mock_request.call_args_list] + assert requested_urls + assert all(url.startswith("https://flag-proxy.example.com") for url in requested_urls) + sent_keys = [call.kwargs["headers"].get("Authorization") for call in mock_request.call_args_list] + assert sent_keys == ["Bearer sk-intended-for-flag-proxy"] * len(requested_urls) From 2a84c397620d0eea027440950ed33925a1dce41e Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 29 Jul 2026 17:13:50 -0700 Subject: [PATCH 47/54] fix(complexity_router): capture the classifier request body in spend logs --- .../complexity_router/complexity_router.py | 17 +++++- .../router_strategy/test_complexity_router.py | 56 +++++++++++++++++++ 2 files changed, 72 insertions(+), 1 deletion(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index e5268b5107b..836c2e9d4c0 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -429,12 +429,21 @@ class ComplexityRouter(CustomLogger): # internal classifier call) is responsible for reconciling. metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata")) + proxy_server_request = { + "body": { + "model": llm_config.model, + "messages": [{"role": "user", "content": classification_prompt}], + "response_format": TierClassification.model_json_schema(), + } + } + response: ModelResponse = await self.litellm_router_instance.acompletion( model=llm_config.model, messages=[{"role": "user", "content": classification_prompt}], response_format=TierClassification, timeout=llm_config.timeout_ms / 1000, metadata=metadata, + proxy_server_request=proxy_server_request, ) content = response.choices[0].message.content if not content: @@ -821,8 +830,14 @@ class ComplexityRouter(CustomLogger): # key/team budget. Key/team attribution fields are preserved for spend logging. metadata = _classifier_call_metadata(request_kwargs.get("metadata")) litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata")) + proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}} query_vector = ( - await encoder.aencode_queries([user_message], metadata=metadata, litellm_metadata=litellm_metadata) + await encoder.aencode_queries( + [user_message], + metadata=metadata, + litellm_metadata=litellm_metadata, + proxy_server_request=proxy_server_request, + ) )[0] route_choice = await routelayer.acall(vector=query_vector) diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index ef70687bd97..a9b86c16c33 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -1417,6 +1417,29 @@ class TestLLMClassifier: call_kwargs = mock_router_instance.acompletion.call_args.kwargs assert call_kwargs["metadata"] == request_metadata + @pytest.mark.asyncio + async def test_aclassify_captures_request_body_in_proxy_server_request( + self, llm_complexity_router, mock_router_instance + ): + """The classifier call must supply proxy_server_request so its request body is logged. + + proxy_server_request["body"] is populated only by the proxy's HTTP ingress + middleware, which never runs for this internally-initiated router.acompletion + call. Without it _get_proxy_server_request_for_spend_logs_payload reads nothing + and stores "{}" for the request, so the classifier's spend-log row shows a + populated response but an empty request and the log cannot show which prompt + drove the tier decision. The captured body must carry the classification prompt + actually sent, so the classifier model, the classification prompt, and the user + text are all asserted here. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "COMPLEX"}')) + await llm_complexity_router.aclassify("explain quantum tunneling in depth") + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + body = call_kwargs["proxy_server_request"]["body"] + assert body["model"] == "haiku-classifier" + assert body["messages"] == call_kwargs["messages"] + assert "explain quantum tunneling in depth" in body["messages"][0]["content"] + @pytest.mark.asyncio async def test_aclassify_strips_budget_reservation_from_classifier_metadata( self, llm_complexity_router, mock_router_instance @@ -2169,6 +2192,39 @@ class TestSemanticKeywordTierRules: assert fake_router.async_embedding_kwargs[0]["metadata"] == caller_metadata assert fake_router.async_embedding_kwargs[0]["litellm_metadata"] == caller_litellm_metadata + @pytest.mark.asyncio + async def test_semantic_embedding_call_captures_request_body_in_proxy_server_request(self, basic_config): + """The query embedding call must supply proxy_server_request so its request is logged. + + Like the LLM classifier, this embedding is fired internally and never passes + through the proxy's HTTP ingress middleware, so proxy_server_request is unset and + the embedding's spend-log row stores "{}" for the request while its response is + captured. The captured body must carry the embedded input so the log shows what + was classified. + """ + fake_router = FakeEmbeddingRouter() + config = { + **basic_config, + "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}], + "semantic_keyword_matching": True, + "embedding_model": "fake-embed", + "match_threshold": 0.5, + } + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=fake_router, + complexity_router_config=config, + ) + await router.async_pre_routing_hook( + model="test-model", + request_kwargs={}, + messages=[{"role": "user", "content": "roll out my k8s cluster"}], + ) + assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt" + body = fake_router.async_embedding_kwargs[0]["proxy_server_request"]["body"] + assert body["model"] == "fake-embed" + assert body["input"] == ["roll out my k8s cluster"] + @pytest.mark.asyncio async def test_semantic_embedding_call_strips_budget_reservation(self, basic_config): """The embedding call must not carry the parent request's budget reservation. From 78207064d8d426b2e408c6572f780ccda613eaf6 Mon Sep 17 00:00:00 2001 From: tin Date: Thu, 30 Jul 2026 01:09:37 +0000 Subject: [PATCH 48/54] fix(complexity_router): log the classifier request on chat completions too The classifier read its metadata only from litellm_metadata, which the proxy populates just for LITELLM_METADATA_ROUTES (/v1/messages, /v1/responses, ...); /v1/chat/completions puts it under metadata, so the classifier call arrived unattributed and _should_track_cost_callback dropped it, leaving no spend-log row at all for the captured request body to show up in. Also log response_format in the wire shape litellm actually sends (type_to_response_format_param) instead of the bare pydantic JSON schema Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../complexity_router/complexity_router.py | 6 +++-- .../router_strategy/test_complexity_router.py | 25 +++++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 836c2e9d4c0..4bc847c923e 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -25,6 +25,7 @@ from pydantic import BaseModel from litellm._logging import verbose_router_logger from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY from litellm.integrations.custom_logger import CustomLogger +from litellm.llms.base_llm.base_utils import type_to_response_format_param from litellm.types.utils import ModelResponse from .config import ( @@ -427,13 +428,14 @@ class ComplexityRouter(CustomLogger): # attributed to the calling key/team instead of being dropped. Excludes the # parent request's budget reservation, which the routed completion (not this # internal classifier call) is responsible for reconciling. - metadata = _classifier_call_metadata((request_kwargs or {}).get("litellm_metadata")) + request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata") + metadata = _classifier_call_metadata(request_metadata) proxy_server_request = { "body": { "model": llm_config.model, "messages": [{"role": "user", "content": classification_prompt}], - "response_format": TierClassification.model_json_schema(), + "response_format": type_to_response_format_param(TierClassification), } } diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index a9b86c16c33..1feeb150c87 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -1417,6 +1417,24 @@ class TestLLMClassifier: call_kwargs = mock_router_instance.acompletion.call_args.kwargs assert call_kwargs["metadata"] == request_metadata + @pytest.mark.asyncio + async def test_aclassify_forwards_metadata_key_used_by_chat_completions( + self, llm_complexity_router, mock_router_instance + ): + """/v1/chat/completions puts the request metadata under "metadata", not "litellm_metadata". + + Only the routes in LITELLM_METADATA_ROUTES (/v1/messages, /v1/responses, ...) get a + "litellm_metadata" bucket; chat completions gets "metadata". Reading only + "litellm_metadata" leaves the classifier call unattributed on the most common route, + so _should_track_cost_callback drops it and no spend-log row is written at all, + which also makes the captured request body unreachable in the Logs UI. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + request_metadata = {"user_api_key": "sk-abc", "user_api_key_team_id": "team-1"} + await llm_complexity_router.aclassify("hi", request_kwargs={"metadata": request_metadata}) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["metadata"] == request_metadata + @pytest.mark.asyncio async def test_aclassify_captures_request_body_in_proxy_server_request( self, llm_complexity_router, mock_router_instance @@ -1439,6 +1457,13 @@ class TestLLMClassifier: assert body["model"] == "haiku-classifier" assert body["messages"] == call_kwargs["messages"] assert "explain quantum tunneling in depth" in body["messages"][0]["content"] + assert body["response_format"]["type"] == "json_schema" + assert body["response_format"]["json_schema"]["schema"]["properties"]["tier"]["enum"] == [ + "SIMPLE", + "MEDIUM", + "COMPLEX", + "REASONING", + ] @pytest.mark.asyncio async def test_aclassify_strips_budget_reservation_from_classifier_metadata( From 3d5b8e5960bbb3c02a63207aa189e4618dc8b38a Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 29 Jul 2026 18:17:34 -0700 Subject: [PATCH 49/54] fix(complexity_router): propagate turn_off_message_logging to internal sub-calls The classifier and semantic-embedding sub-calls now capture proxy_server_request, but neither forwarded the caller's turn_off_message_logging opt-out. A caller who disabled message logging still had their prompt stored in the clear in these internal sub-calls' spend-log rows, since should_redact_message_logging reads the flag per-call and this internal call never inherited it. --- .../complexity_router/complexity_router.py | 12 +++ .../router_strategy/test_complexity_router.py | 78 +++++++++++++++++++ 2 files changed, 90 insertions(+) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 4bc847c923e..bbeeeb0be64 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -113,6 +113,14 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] } +def _effective_turn_off_message_logging(request_kwargs: dict[str, Any] | None) -> bool | None: + from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + initialize_standard_callback_dynamic_params, + ) + + return initialize_standard_callback_dynamic_params(request_kwargs or {}).get("turn_off_message_logging") + + class DimensionScore: """Represents a score for a single dimension with optional signal.""" @@ -430,6 +438,7 @@ class ComplexityRouter(CustomLogger): # internal classifier call) is responsible for reconciling. request_metadata = (request_kwargs or {}).get("litellm_metadata") or (request_kwargs or {}).get("metadata") metadata = _classifier_call_metadata(request_metadata) + turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs) proxy_server_request = { "body": { @@ -446,6 +455,7 @@ class ComplexityRouter(CustomLogger): timeout=llm_config.timeout_ms / 1000, metadata=metadata, proxy_server_request=proxy_server_request, + turn_off_message_logging=turn_off_message_logging, ) content = response.choices[0].message.content if not content: @@ -832,6 +842,7 @@ class ComplexityRouter(CustomLogger): # key/team budget. Key/team attribution fields are preserved for spend logging. metadata = _classifier_call_metadata(request_kwargs.get("metadata")) litellm_metadata = _classifier_call_metadata(request_kwargs.get("litellm_metadata")) + turn_off_message_logging = _effective_turn_off_message_logging(request_kwargs) proxy_server_request = {"body": {"model": self.config.embedding_model, "input": [user_message]}} query_vector = ( await encoder.aencode_queries( @@ -839,6 +850,7 @@ class ComplexityRouter(CustomLogger): metadata=metadata, litellm_metadata=litellm_metadata, proxy_server_request=proxy_server_request, + turn_off_message_logging=turn_off_message_logging, ) )[0] route_choice = await routelayer.acall(vector=query_vector) diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 1feeb150c87..f31ca32f4c5 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -1465,6 +1465,54 @@ class TestLLMClassifier: "REASONING", ] + @pytest.mark.asyncio + async def test_aclassify_propagates_top_level_turn_off_message_logging( + self, llm_complexity_router, mock_router_instance + ): + """A caller's top-level turn_off_message_logging must reach the classifier call. + + Without this, a caller who opts a request out of message logging still has their + prompt captured in full by the classifier's proxy_server_request: the spend-log + redaction gate (should_redact_message_logging) reads turn_off_message_logging off + the classifier call's own kwargs, and this internal call is not the caller's + request, so it never inherits the opt-out unless it's forwarded explicitly. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify("secret prompt", request_kwargs={"turn_off_message_logging": True}) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is True + + @pytest.mark.asyncio + async def test_aclassify_propagates_metadata_slot_turn_off_message_logging( + self, llm_complexity_router, mock_router_instance + ): + """turn_off_message_logging set inside metadata/litellm_metadata must also propagate. + + initialize_standard_callback_dynamic_params reads this flag from either the + top-level request kwargs or the metadata/litellm_metadata dicts (the same slots a + real HTTP request populates), so the classifier call must resolve it from there too. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify( + "secret prompt", request_kwargs={"litellm_metadata": {"turn_off_message_logging": True}} + ) + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is True + + @pytest.mark.asyncio + async def test_aclassify_defaults_turn_off_message_logging_to_none( + self, llm_complexity_router, mock_router_instance + ): + """With no caller opt-out, the classifier call must not force redaction on or off. + + Passing None (rather than omitting the kwarg or defaulting to False) preserves the + existing header- and global-setting fallbacks in should_redact_message_logging. + """ + mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) + await llm_complexity_router.aclassify("hi") + call_kwargs = mock_router_instance.acompletion.call_args.kwargs + assert call_kwargs["turn_off_message_logging"] is None + @pytest.mark.asyncio async def test_aclassify_strips_budget_reservation_from_classifier_metadata( self, llm_complexity_router, mock_router_instance @@ -2250,6 +2298,36 @@ class TestSemanticKeywordTierRules: assert body["model"] == "fake-embed" assert body["input"] == ["roll out my k8s cluster"] + @pytest.mark.asyncio + async def test_semantic_embedding_call_propagates_turn_off_message_logging(self, basic_config): + """A caller's turn_off_message_logging must reach the query embedding call. + + The embedding now captures the user's prompt in proxy_server_request, so a caller + who opts out of message logging must have that opt-out forwarded; otherwise the + embedding's spend-log row stores the prompt in the clear despite the parent request + being redacted, exposing it to anyone authorized to read the team's spend logs. + """ + fake_router = FakeEmbeddingRouter() + config = { + **basic_config, + "keyword_tier_rules": [{"keywords": ["kubernetes deployment"], "tier": "REASONING"}], + "semantic_keyword_matching": True, + "embedding_model": "fake-embed", + "match_threshold": 0.5, + } + router = ComplexityRouter( + model_name="test-router", + litellm_router_instance=fake_router, + complexity_router_config=config, + ) + await router.async_pre_routing_hook( + model="test-model", + request_kwargs={"turn_off_message_logging": True}, + messages=[{"role": "user", "content": "roll out my k8s cluster"}], + ) + assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt" + assert fake_router.async_embedding_kwargs[0]["turn_off_message_logging"] is True + @pytest.mark.asyncio async def test_semantic_embedding_call_strips_budget_reservation(self, basic_config): """The embedding call must not carry the parent request's budget reservation. From dbc0d23c1ecea6450e248b3d2b846c28a25b2869 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 18:38:29 -0700 Subject: [PATCH 50/54] fix(vertex_ai): skip context caching when the cached block ends on a model turn --- .../context_caching/transformation.py | 13 +- .../vertex_ai_context_caching.py | 15 ++ .../test_vertex_ai_context_caching.py | 129 ++++++++++++++++++ 3 files changed, 156 insertions(+), 1 deletion(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index f0ce3323ef6..a74e0c97abc 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works """ import re -from typing import List, Optional, Tuple, Literal +from typing import List, Optional, Sequence, Tuple, Literal from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import CachedContentRequestBody @@ -152,6 +152,17 @@ def separate_cached_messages( return cached_messages, non_cached_messages +def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool: + """ + The cachedContents API rejects contents ending on a model turn, which is how it + classifies both assistant messages and tool results, with HTTP 400 + "Requests ending with a model turn are not supported". + """ + if not cached_messages: + return False + return cached_messages[-1].get("role") not in ("assistant", "tool", "function") + + def transform_openai_messages_to_gemini_context_caching( model: str, messages: List[AllMessageValues], diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index 0bf3715f798..fe4cd4ec451 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import ( from ..common_utils import VertexAIError, get_vertex_base_url from ..vertex_llm_base import VertexBase from .transformation import ( + cached_messages_end_on_supported_turn, separate_cached_messages, transform_openai_messages_to_gemini_context_caching, ) @@ -308,6 +309,13 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on an assistant or " + "tool turn, which the cachedContents API rejects. Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( @@ -459,6 +467,13 @@ class ContextCachingEndpoints(VertexBase): if len(cached_messages) == 0: return messages, optional_params, None + if not cached_messages_end_on_supported_turn(cached_messages): + verbose_logger.debug( + "Vertex AI context caching: cached message block ends on an assistant or " + "tool turn, which the cachedContents API rejects. Skipping context caching." + ) + return messages, optional_params, None + # Gemini requires a minimum of 1024 tokens for context caching. # Skip caching if the cached content is too small to avoid API errors. if not is_prompt_caching_valid_prompt( diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index cf75964ddb7..1aa724e551e 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1452,6 +1452,135 @@ class TestContextCachingEndpoints: # Restart the patcher so teardown_method can stop it cleanly self._token_check_patcher.start() + def _model_turn_final_messages(self, final_cached_role): + tool_call = { + "id": "call_abc123", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'}, + } + cached_tail = ( + [ + { + "role": "tool", + "tool_call_id": "call_abc123", + "content": "72F and sunny", + "cache_control": {"type": "ephemeral"}, + } + ] + if final_cached_role == "tool" + else [] + ) + return [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Use the weather tool for every answer.", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + { + "role": "assistant", + "content": "", + "tool_calls": [tool_call], + "cache_control": {"type": "ephemeral"}, + }, + *cached_tail, + {"role": "user", "content": "What is the weather in Boston?"}, + ] + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """The cachedContents API rejects contents ending on an assistant or tool turn + with HTTP 400 "Requests ending with a model turn are not supported", so the + request must proceed uncached instead of failing. + """ + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = self.context_caching.check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_client.get.assert_not_called() + self.mock_client.post.assert_not_called() + + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.asyncio + async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( + self, final_cached_role + ): + """Async variant: an unsupported terminal turn skips caching instead of failing.""" + all_messages = self._model_turn_final_messages(final_cached_role) + optional_params = self.sample_optional_params.copy() + + result = await self.context_caching.async_check_and_create_cache( + messages=all_messages, + optional_params=optional_params, + api_key="test_key", + api_base=None, + model="gemini-3.6-flash", + client=self.mock_async_client, + timeout=30.0, + logging_obj=self.mock_logging, + cached_content=None, + custom_llm_provider="vertex_ai", + vertex_project="test_project", + vertex_location="us-central1", + vertex_auth_header="test_token", + ) + + messages, returned_params, returned_cache = result + assert messages == all_messages + assert returned_cache is None + assert "tools" in returned_params + self.mock_async_client.get.assert_not_called() + self.mock_async_client.post.assert_not_called() + + +def test_cached_messages_end_on_supported_turn(): + from litellm.llms.vertex_ai.context_caching.transformation import ( + cached_messages_end_on_supported_turn, + ) + + assert ( + cached_messages_end_on_supported_turn( + [{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}] + ) + is True + ) + assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True + assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False + assert ( + cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}]) + is False + ) + assert ( + cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}]) + is False + ) + assert cached_messages_end_on_supported_turn([]) is False + class TestCheckCachePagination: """Test pagination logic in check_cache and async_check_cache methods.""" From 074eda52222da35bd43a6f9e6666aa4203eedeb1 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Wed, 29 Jul 2026 18:45:51 -0700 Subject: [PATCH 51/54] fix(complexity_router): use Mapping instead of dict in turn_off_message_logging helper parameter Accept read-only Mapping[str, Any] instead of mutable dict[str, Any] in _effective_turn_off_message_logging's parameter to satisfy LIT001 (mutable collections in type annotations). Convert to dict for the function that expects Dict. --- .../router_strategy/complexity_router/complexity_router.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index bbeeeb0be64..1da8ee68c6e 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -18,6 +18,7 @@ from __future__ import annotations import asyncio import random import re +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Literal, Union, cast from pydantic import BaseModel @@ -113,12 +114,14 @@ def _classifier_call_metadata(metadata: dict[str, Any] | None) -> dict[str, Any] } -def _effective_turn_off_message_logging(request_kwargs: dict[str, Any] | None) -> bool | None: +def _effective_turn_off_message_logging(request_kwargs: Mapping[str, Any] | None) -> bool | None: from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( initialize_standard_callback_dynamic_params, ) - return initialize_standard_callback_dynamic_params(request_kwargs or {}).get("turn_off_message_logging") + return initialize_standard_callback_dynamic_params(dict(request_kwargs) if request_kwargs else {}).get( + "turn_off_message_logging" + ) class DimensionScore: From 04d702c46ae028d5a1375173a9c9209cb0d6fa39 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 18:51:25 -0700 Subject: [PATCH 52/54] test(managed-files): call store_unified_file_id twice and assert upsert payloads --- .../proxy/test_managed_files_hook.py | 31 ++++++++++++------- 1 file changed, 19 insertions(+), 12 deletions(-) diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 1526aad7a24..2580197d6d2 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -5,6 +5,8 @@ Regression test for afile_retrieve called without credentials in async_post_call_success_hook when processing completed batch responses. """ +import json + import pytest from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch @@ -421,18 +423,23 @@ async def test_store_unified_file_id_is_idempotent_via_upsert(): unified_file_id, never do an unconditional create that raises on conflict.""" managed_files, mock_prisma = _make_real_managed_files_instance() file_id = "litellm_proxy_unified_output_id_abc" + model_mappings = {"model-deploy-xyz": "file-output-abc"} - await managed_files.store_unified_file_id( - file_id=file_id, - file_object=_make_file_object(), - litellm_parent_otel_span=None, - model_mappings={"model-deploy-xyz": "file-output-abc"}, - user_api_key_dict=_make_user_api_key_dict(), - ) + for _ in range(2): + await managed_files.store_unified_file_id( + file_id=file_id, + file_object=_make_file_object(), + litellm_parent_otel_span=None, + model_mappings=model_mappings, + user_api_key_dict=_make_user_api_key_dict(), + ) mock_prisma.db.litellm_managedfiletable.create.assert_not_awaited() - mock_prisma.db.litellm_managedfiletable.upsert.assert_awaited_once() - assert ( - mock_prisma.db.litellm_managedfiletable.upsert.await_args.kwargs["where"] - == {"unified_file_id": file_id} - ) + upsert_mock = mock_prisma.db.litellm_managedfiletable.upsert + assert upsert_mock.await_count == 2 + for upsert_call in upsert_mock.await_args_list: + assert upsert_call.kwargs["where"] == {"unified_file_id": file_id} + upsert_data = upsert_call.kwargs["data"] + assert upsert_data["create"]["unified_file_id"] == file_id + assert json.loads(upsert_data["create"]["model_mappings"]) == model_mappings + assert json.loads(upsert_data["update"]["model_mappings"]) == model_mappings From 56d51bc32edda4bbd8019fbef45dabaca0879b88 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:00:31 -0700 Subject: [PATCH 53/54] build(makefile): give local basedpyright runs the node heap CI uses (#35173) --- Makefile | 2 ++ 1 file changed, 2 insertions(+) diff --git a/Makefile b/Makefile index 8b657dcb465..e9b2fb9d8f1 100644 --- a/Makefile +++ b/Makefile @@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \ else echo "No changed .py files to check."; fi +lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288 + lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging From 819dc7812af842b7b7844dc7c794c302d8e1075a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:47:26 -0700 Subject: [PATCH 54/54] fix(vertex_ai): evaluate cached-block terminal turn after system extraction --- .../context_caching/transformation.py | 11 ++++-- .../vertex_ai_context_caching.py | 10 +++-- .../test_vertex_ai_context_caching.py | 38 +++++++++++++++---- 3 files changed, 43 insertions(+), 16 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index a74e0c97abc..36c78974aca 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -156,11 +156,14 @@ def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageVa """ The cachedContents API rejects contents ending on a model turn, which is how it classifies both assistant messages and tool results, with HTTP 400 - "Requests ending with a model turn are not supported". + "Requests ending with a model turn are not supported". System messages are + extracted into system_instruction before contents are built, so the terminal + turn is the last non-system message. """ - if not cached_messages: - return False - return cached_messages[-1].get("role") not in ("assistant", "tool", "function") + non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system") + if not non_system_messages: + return bool(cached_messages) + return non_system_messages[-1].get("role") not in ("assistant", "tool", "function") def transform_openai_messages_to_gemini_context_caching( diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index fe4cd4ec451..f8774e33ca4 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -311,8 +311,9 @@ class ContextCachingEndpoints(VertexBase): if not cached_messages_end_on_supported_turn(cached_messages): verbose_logger.debug( - "Vertex AI context caching: cached message block ends on an assistant or " - "tool turn, which the cachedContents API rejects. Skipping context caching." + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." ) return messages, optional_params, None @@ -469,8 +470,9 @@ class ContextCachingEndpoints(VertexBase): if not cached_messages_end_on_supported_turn(cached_messages): verbose_logger.debug( - "Vertex AI context caching: cached message block ends on an assistant or " - "tool turn, which the cachedContents API rejects. Skipping context caching." + "Vertex AI context caching: cached message block ends on a model turn once " + "system messages are extracted, which the cachedContents API rejects. " + "Skipping context caching." ) return messages, optional_params, None diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 1aa724e551e..ad890d0c7ea 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1458,18 +1458,24 @@ class TestContextCachingEndpoints: "type": "function", "function": {"name": "get_weather", "arguments": '{"location": "Boston"}'}, } - cached_tail = ( - [ + cached_tail = { + "assistant": [], + "tool": [ { "role": "tool", "tool_call_id": "call_abc123", "content": "72F and sunny", "cache_control": {"type": "ephemeral"}, } - ] - if final_cached_role == "tool" - else [] - ) + ], + "system": [ + { + "role": "system", + "content": "Tool results are authoritative.", + "cache_control": {"type": "ephemeral"}, + } + ], + }[final_cached_role] return [ { "role": "user", @@ -1491,7 +1497,7 @@ class TestContextCachingEndpoints: {"role": "user", "content": "What is the weather in Boston?"}, ] - @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( self, final_cached_role ): @@ -1525,7 +1531,7 @@ class TestContextCachingEndpoints: self.mock_client.get.assert_not_called() self.mock_client.post.assert_not_called() - @pytest.mark.parametrize("final_cached_role", ["assistant", "tool"]) + @pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"]) @pytest.mark.asyncio async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn( self, final_cached_role @@ -1571,6 +1577,22 @@ def test_cached_messages_end_on_supported_turn(): ) assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False + assert ( + cached_messages_end_on_supported_turn( + [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + {"role": "system", "content": "be brief"}, + ] + ) + is False + ) + assert ( + cached_messages_end_on_supported_turn( + [{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}] + ) + is True + ) assert ( cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}]) is False