From 5a5bb8c9d844870c25684e169960f1571d08e5ce Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:29:39 +0000 Subject: [PATCH 01/78] fix(proxy): stop /{provider}/v1/files from capturing /openai_passthrough The native files and batches routes declare /{provider}/v1/... and their routers are mounted before the passthrough router, so /openai_passthrough/v1/files and /openai_passthrough/v1/batches matched them with provider="openai_passthrough" and 500'd on the LlmProviders lookup instead of reaching openai_proxy_route. Move the dedicated /openai_passthrough prefix onto its own router mounted ahead of the batches and files routers. /openai/... and every other provider prefix keep their current behavior. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../llm_passthrough_endpoints.py | 3 +- litellm/proxy/proxy_server.py | 2 + .../test_llm_pass_through_endpoints.py | 57 +++++++++++++++++++ 3 files changed, 61 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 38da00a3bb9..baa74c19182 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -60,6 +60,7 @@ from .passthrough_endpoint_router import PassthroughEndpointRouter vertex_llm_base: Final = VertexBase() router: Final = APIRouter() +openai_passthrough_router: Final = APIRouter() default_vertex_config: Final = None passthrough_endpoint_router: Final = PassthroughEndpointRouter() @@ -1875,7 +1876,7 @@ async def vertex_proxy_route( ) -@router.api_route( +@openai_passthrough_router.api_route( "/openai_passthrough/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], tags=["OpenAI Pass-through", "pass-through"], diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index fb9c4e67aad..e75277e7f0a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -522,6 +522,7 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import ( set_files_config, ) from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + openai_passthrough_router, passthrough_endpoint_router, vertex_ai_live_websocket_passthrough, ) @@ -16433,6 +16434,7 @@ app.include_router(search_router) app.include_router(image_router) app.include_router(fine_tuning_router) app.include_router(credential_router) +app.include_router(openai_passthrough_router) app.include_router(batches_router) app.include_router(openai_files_router) app.include_router(llm_passthrough_router) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 181846fe289..27d6e4c8585 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -2814,6 +2814,63 @@ class TestOpenAIPassthroughRoute: assert result == {"id": "asst_123", "object": "assistant"} +def _resolve_route_name(method: str, path: str) -> str | None: + from starlette.routing import Match + + from litellm.proxy.proxy_server import app + + scope = { + "type": "http", + "method": method, + "path": path, + "headers": [], + "query_string": b"", + "root_path": "", + } + for route in app.router.routes: + if route.matches(scope)[0] == Match.FULL: + return getattr(route, "name", None) + return None + + +@pytest.mark.parametrize( + "method, path", + [ + ("POST", "/openai_passthrough/v1/files"), + ("GET", "/openai_passthrough/v1/files"), + ("GET", "/openai_passthrough/v1/files/file-abc123"), + ("DELETE", "/openai_passthrough/v1/files/file-abc123"), + ("GET", "/openai_passthrough/v1/files/file-abc123/content"), + ("POST", "/openai_passthrough/v1/batches"), + ("GET", "/openai_passthrough/v1/batches"), + ("GET", "/openai_passthrough/v1/batches/batch_abc123"), + ("POST", "/openai_passthrough/v1/batches/batch_abc123/cancel"), + ("POST", "/openai_passthrough/v1/responses"), + ], +) +def test_openai_passthrough_prefix_wins_over_native_provider_routes(method, path): + """ + /openai_passthrough exists to guarantee passthrough, so the native + /{provider}/v1/files and /{provider}/v1/batches routes must never capture it + with provider="openai_passthrough" (which 500s on the LlmProviders lookup). + """ + assert _resolve_route_name(method, path) == "openai_proxy_route" + + +@pytest.mark.parametrize( + "method, path, expected_name", + [ + ("POST", "/openai/v1/files", "create_file"), + ("GET", "/azure/v1/files", "list_files"), + ("POST", "/v1/files", "create_file"), + ("POST", "/v1/batches", "create_batch"), + ("POST", "/openai/v1/chat/completions", "openai_proxy_route"), + ], +) +def test_native_provider_routes_are_unchanged(method, path, expected_name): + assert _resolve_route_name(method, path) == expected_name + + class TestCursorProxyRoute: """Tests for the Cursor Cloud Agents pass-through route.""" From 357f90fa39d18c9a158a978ebd1ed0fecac6044d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 15:33:40 +0000 Subject: [PATCH 02/78] fix(proxy): scope file list pagination cursors to the caller GET /v1/files filters data down to the caller's own managed files but left first_id and last_id as the upstream page's, so a non-owner got back file ids belonging to other users even with an empty data array --- .../proxy/hooks/managed_files.py | 15 +++ .../proxy/hooks/test_managed_files.py | 100 ++++++++++++++++++ 2 files changed, 115 insertions(+) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 0036603bcd1..851e202e2fb 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1270,10 +1270,25 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) ## Filter the response to only include the files created by the user response.data = user_created_file_ids # type: ignore + self._scope_list_page_cursors(response, user_created_file_ids) return response return response return response + @staticmethod + def _scope_list_page_cursors(response: AsyncCursorPage, data: List[OpenAIFileObject]) -> None: + """Rebuild ``first_id`` / ``last_id`` from the caller-scoped page. + + The upstream cursors point at rows that were just filtered out, so + leaving them in place discloses other callers' file ids. + """ + if hasattr(response, "first_id"): + response.first_id = data[0].id if data else None + if hasattr(response, "last_id"): + response.last_id = data[-1].id if data else None + if not data and hasattr(response, "has_more"): + response.has_more = False + async def afile_retrieve( self, file_id: str, litellm_parent_otel_span: Optional[Span], llm_router=None ) -> OpenAIFileObject: diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 50af6465d06..3384c553740 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -2861,3 +2861,103 @@ async def test_same_user_different_keys_can_access_batch(): assert "batch_id" in result2 # Both keys should get the same result assert result1["batch_id"] == result2["batch_id"] + + +@pytest.mark.asyncio +async def test_file_list_cursors_are_scoped_to_the_caller(): + """A non-owner must not learn other callers' file ids through the page cursors.""" + from openai.pagination import AsyncCursorPage + from openai.types import FileObject + + from litellm.proxy._types import UserAPIKeyAuth + + owner_file = FileObject( + id="file-owner-1", + bytes=100, + created_at=1, + filename="owner.jsonl", + object="file", + purpose="batch", + status="processed", + ) + upstream_page = AsyncCursorPage[FileObject].construct( + data=[owner_file], + has_more=True, + first_id=owner_file.id, + last_id=owner_file.id, + object="list", + ) + + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_many.return_value = [] + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + response = await proxy_managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth( + user_id="other-user", team_id="other-team", parent_otel_span=MagicMock() + ), + response=upstream_page, + ) + + assert response.data == [] + assert response.first_id is None + assert response.last_id is None + assert response.has_more is False + + +@pytest.mark.asyncio +async def test_file_list_cursors_follow_the_owner_scoped_page(): + from openai.pagination import AsyncCursorPage + from openai.types import FileObject + + from litellm.proxy._types import UserAPIKeyAuth + + def _raw_file(file_id: str) -> FileObject: + return FileObject( + id=file_id, + bytes=100, + created_at=1, + filename=f"{file_id}.jsonl", + object="file", + purpose="batch", + status="processed", + ) + + upstream_page = AsyncCursorPage[FileObject].construct( + data=[_raw_file("file-someone-else"), _raw_file("file-mine")], + has_more=False, + first_id="file-someone-else", + last_id="file-mine", + object="list", + ) + + managed_row = MagicMock() + managed_row.file_object = { + "id": "litellm_proxy:mine", + "bytes": 100, + "created_at": 1, + "filename": "mine.jsonl", + "object": "file", + "purpose": "batch", + "status": "processed", + } + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_many.return_value = [managed_row] + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=prisma_client + ) + + response = await proxy_managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth( + user_id="mine-user", parent_otel_span=MagicMock() + ), + response=upstream_page, + ) + + assert [file_object.id for file_object in response.data] == ["litellm_proxy:mine"] + assert response.first_id == "litellm_proxy:mine" + assert response.last_id == "litellm_proxy:mine" From f9b86b253a3fb87d003bb5ccc80c7d89aa91dd62 Mon Sep 17 00:00:00 2001 From: Harry Qian Date: Tue, 4 Aug 2026 17:14:26 +0800 Subject: [PATCH 03/78] fix(proxy): restore query-param validation under fastapi>=0.140.7 fastapi 0.140.7 removed get_flat_dependant(), which broke the import in management_v1/common.py and took down every /management/v1 route. Switch to get_flat_params() and filter to ParamTypes.query so unknown-query-param rejection keeps matching the old behavior. --- .../management_endpoints/management_v1/common.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/management_v1/common.py b/litellm/proxy/management_endpoints/management_v1/common.py index 8525d67a041..ec79820465a 100644 --- a/litellm/proxy/management_endpoints/management_v1/common.py +++ b/litellm/proxy/management_endpoints/management_v1/common.py @@ -4,7 +4,8 @@ from typing import Final from urllib.parse import urlencode from fastapi import Request -from fastapi.dependencies.utils import get_flat_dependant +from fastapi.dependencies.utils import get_flat_params +from fastapi.params import ParamTypes from fastapi.responses import JSONResponse from litellm.types.proxy.management_endpoints.management_v1 import ( @@ -42,7 +43,13 @@ def _declared_query_params(request: Request) -> frozenset[str]: dependant: Final = getattr(route, "dependant", None) if dependant is None: return frozenset() - return frozenset(field.alias for field in get_flat_dependant(dependant, skip_repeats=True).query_params) + # fastapi>=0.140.7 removed get_flat_dependant(); get_flat_params() returns the + # flattened (deduped) param list. Filter to query params to match the old behavior. + return frozenset( + field.alias + for field in get_flat_params(dependant) + if getattr(field.field_info, "in_", None) == ParamTypes.query + ) def escape_like(value: str) -> str: From da443d1266615507f52a101b461c80e0265069ae Mon Sep 17 00:00:00 2001 From: Harry Qian Date: Tue, 4 Aug 2026 18:21:22 +0800 Subject: [PATCH 04/78] test(proxy): lock in query-param validation across fastapi param types Guards _declared_query_params against a regression in the get_flat_params migration: the flatten step returns path, query, header and cookie params together, so a dropped ParamTypes.query filter would wrongly treat path or header names as declared query params and accept unknown ones. Removing the filter fails these tests. --- .../management_v1/test_common.py | 95 +++++++++++++++++++ 1 file changed, 95 insertions(+) create mode 100644 tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py diff --git a/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py new file mode 100644 index 00000000000..167a06ed551 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/management_v1/test_common.py @@ -0,0 +1,95 @@ +from typing import Annotated + +from fastapi import Depends, FastAPI, Header, Query, Request +from fastapi.testclient import TestClient + +from litellm.proxy.management_endpoints.management_v1.common import ( + ManagementProblem, + PROBLEM_CONTENT_TYPE, + _declared_query_params, + problem_response, + reject_unknown_query_params, +) + + +def _client() -> TestClient: + app = FastAPI() + + @app.exception_handler(ManagementProblem) + async def _handle(_request: Request, exc: ManagementProblem): + return problem_response(exc.problem) + + @app.get("/things/{thing_id}", dependencies=[Depends(reject_unknown_query_params)]) + def _handler( + thing_id: str, + request: Request, + status: Annotated[str | None, Query(alias="filter[status]")] = None, + page: Annotated[int, Query(ge=1)] = 1, + x_trace: Annotated[str | None, Header()] = None, + ) -> dict[str, bool]: + return {"ok": True} + + return TestClient(app, raise_server_exceptions=False) + + +def test_a_declared_query_param_is_accepted_by_its_alias(): + response = _client().get("/things/abc", params={"filter[status]": "active", "page": "2"}) + assert response.status_code == 200, response.text + + +def test_an_unknown_query_param_is_rejected_as_a_problem(): + response = _client().get("/things/abc", params={"bogus": "x"}) + assert response.status_code == 400 + assert response.headers["content-type"].startswith(PROBLEM_CONTENT_TYPE) + assert "bogus" in response.json()["detail"] + + +def test_a_path_param_name_is_not_a_declared_query_param(): + """The flatten step returns path+query+header together; only query names count as declared. + + If the ParamTypes.query filter were dropped, `thing_id` (a path param) would leak + into the declared set and this request would be wrongly accepted. + """ + response = _client().get("/things/abc", params={"thing_id": "x"}) + assert response.status_code == 400 + assert "thing_id" in response.json()["detail"] + + +def test_a_header_param_name_is_not_a_declared_query_param(): + response = _client().get("/things/abc", params={"x-trace": "x"}) + assert response.status_code == 400 + assert "x-trace" in response.json()["detail"] + + +def test_declared_query_params_isolates_query_aliases_from_other_param_types(): + captured: dict[str, frozenset[str]] = {} + app = FastAPI() + + @app.get("/things/{thing_id}") + def _handler( + thing_id: str, + request: Request, + status: Annotated[str | None, Query(alias="filter[status]")] = None, + page: Annotated[int, Query(ge=1)] = 1, + x_trace: Annotated[str | None, Header()] = None, + ) -> dict[str, bool]: + captured["declared"] = _declared_query_params(request) + return {"ok": True} + + TestClient(app).get("/things/abc") + assert captured["declared"] == frozenset({"filter[status]", "page"}) + + +def test_declared_query_params_is_empty_when_the_route_has_no_dependant(): + request = Request( + { + "type": "http", + "method": "GET", + "scheme": "http", + "root_path": "", + "path": "/things/abc", + "query_string": b"", + "headers": [(b"host", b"testserver")], + } + ) + assert _declared_query_params(request) == frozenset() From 5883aa354d42a3225fef485034e86a52a275cdd7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 04:45:39 -0700 Subject: [PATCH 05/78] fix(router): keep batch fallbacks inside the model group that owns the file A batch or fine-tuning job is created from a file the caller already uploaded, and that file only exists under the credentials of the deployment that stored it. When the router fell back to a different model group it handed that file id to a provider that has never seen it, so the caller got the second provider's complaint about the file id instead of the error that explains what was actually wrong with their request. run_async_fallback now skips fallback targets outside the original model group whenever the request carries input_file_id or training_file. Order-based fallbacks stay inside the group, so retrying across deployments still works. The same handler also crashed with "'NoneType' object has no attribute 'update'" whenever a fallback fired on a request with metadata set to None, which /v1/batches always does when the caller sends no metadata, turning the provider's 400 into a 500. Record the model group with a merge instead of setdefault, and write it to litellm_metadata on the endpoints that use it so the router's bookkeeping no longer lands in the metadata stored on the provider's batch. --- .../router_utils/fallback_event_handlers.py | 41 ++++- .../test_fallback_event_handlers.py | 141 ++++++++++++++++++ tests/test_litellm/test_router.py | 62 ++++++++ 3 files changed, 241 insertions(+), 3 deletions(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 1c6bb52ccb8..c4a84a1d61e 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -9,6 +9,7 @@ from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, get_fallback_error_info, ) +from litellm.router_utils.batch_utils import _get_router_metadata_variable_name from litellm.types.router import LiteLLMParamsTypedDict if TYPE_CHECKING: @@ -82,6 +83,28 @@ def get_fallback_model_group(fallbacks: list[Any], model_group: str) -> tuple[li return fallback_model_group, generic_fallback_idx +PROVIDER_SCOPED_RESOURCE_KEYS: Final = ("input_file_id", "training_file") + + +def _get_fallback_target_model_group(fallback_entry: str | dict[str, object]) -> str | None: + if isinstance(fallback_entry, str): + return fallback_entry + target: Final = fallback_entry.get("model") + return target if isinstance(target, str) else None + + +def references_provider_scoped_resource(kwargs: dict[str, object]) -> bool: + """ + True when the request names a file that only exists under one provider's credentials. + + Batch and fine-tuning jobs are created from a file the caller already uploaded, and + that file lives in the account of the deployment that stored it. Handing the id to a + different model group can only fail, and the second provider's error replaces the + error the caller actually needs to see. + """ + return any(kwargs.get(key) for key in PROVIDER_SCOPED_RESOURCE_KEYS) + + async def run_async_fallback( *args: tuple[Any], litellm_router: LitellmRouter, @@ -120,10 +143,21 @@ async def run_async_fallback( error_from_fallbacks = original_exception fallback_errors = (get_fallback_error_info(original_exception),) + metadata_variable_name: Final = _get_router_metadata_variable_name( + function_name=getattr(kwargs.get("original_function"), "__name__", None) + ) + same_model_group_only: Final = references_provider_scoped_resource(kwargs) for mg in fallback_model_group: if mg == original_model_group: continue + if same_model_group_only and _get_fallback_target_model_group(mg) != original_model_group: + verbose_router_logger.info( + "Skipping fallback to model_group = %s: request is pinned to model_group = %s by its uploaded file", + mask_sensitive_structure(mg), + original_model_group, + ) + continue try: # LOGGING kwargs = litellm_router.log_retry(kwargs=kwargs, e=original_exception) @@ -132,9 +166,10 @@ async def run_async_fallback( kwargs["model"] = mg elif isinstance(mg, dict): kwargs.update(mg) - kwargs.setdefault("metadata", {}).update( - {"model_group": kwargs.get("model", None)} - ) # update model_group used, if fallbacks are done + kwargs[metadata_variable_name] = { + **(kwargs.get(metadata_variable_name) or {}), + "model_group": kwargs.get("model", None), + } # update model_group used, if fallbacks are done fallback_depth = fallback_depth + 1 kwargs["fallback_depth"] = fallback_depth kwargs["max_fallbacks"] = max_fallbacks diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 98a34de295c..d93aa4ab023 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -142,6 +142,147 @@ async def test_run_async_fallback_skips_original_model_group(): assert response._hidden_params["additional_headers"]["x-litellm-attempted-fallbacks"] == 1 +class AttemptRecordingRouter: + def __init__(self): + self.attempted_model_groups = [] + self.received_kwargs = None + + def log_retry(self, kwargs, e): + return kwargs + + async def async_function_with_fallbacks(self, *args, **kwargs): + self.attempted_model_groups.append(kwargs.get("model")) + self.received_kwargs = kwargs + return StreamingWrapper() + + +async def _acreate_batch(*args, **kwargs): + raise AssertionError("only used for its __name__") + + +@pytest.mark.asyncio +async def test_run_async_fallback_keeps_uploaded_file_requests_in_their_model_group(): + """An input_file_id only exists under the credentials of the group it was uploaded + to, so a cross-group fallback can only fail with the wrong provider's error.""" + router = AttemptRecordingRouter() + owning_provider_error = RuntimeError("openai connection error") + + with pytest.raises(RuntimeError, match="openai connection error"): + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=owning_provider_error, + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + original_function=_acreate_batch, + ) + + assert router.attempted_model_groups == [] + + +@pytest.mark.asyncio +async def test_run_async_fallback_keeps_fine_tuning_requests_in_their_model_group(): + router = AttemptRecordingRouter() + + with pytest.raises(RuntimeError, match="openai connection error"): + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + training_file="file-owned-by-openai", + ) + + assert router.attempted_model_groups == [] + + +@pytest.mark.asyncio +async def test_run_async_fallback_allows_same_model_group_retry_for_uploaded_file_requests(): + """Order-based fallbacks stay inside the owning group, so they must still run.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=[{"model": "openai-group", "_target_order": 2}], + original_model_group="openai-group", + original_exception=RuntimeError("first deployment failed"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + original_function=_acreate_batch, + ) + + assert router.attempted_model_groups == ["openai-group"] + + +@pytest.mark.asyncio +async def test_run_async_fallback_still_crosses_model_groups_without_an_uploaded_file(): + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + ) + + assert router.attempted_model_groups == ["azure-group"] + + +@pytest.mark.asyncio +async def test_run_async_fallback_handles_explicitly_none_metadata(): + """/v1/batches always sets `metadata`, and sets it to None when the caller sent + none, so setdefault() on it hands back None instead of a dict.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=["azure-group"], + original_model_group="openai-group", + original_exception=RuntimeError("openai connection error"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + metadata=None, + ) + + assert router.received_kwargs["metadata"] == {"model_group": "azure-group"} + + +@pytest.mark.asyncio +async def test_run_async_fallback_records_batch_model_group_outside_provider_metadata(): + """`metadata` on a batch request is forwarded to the provider and stored on the + batch, so the router's own model_group belongs in litellm_metadata.""" + router = AttemptRecordingRouter() + + await run_async_fallback( + litellm_router=router, + fallback_model_group=[{"model": "openai-group", "_target_order": 2}], + original_model_group="openai-group", + original_exception=RuntimeError("first deployment failed"), + max_fallbacks=3, + fallback_depth=0, + model="openai-group", + input_file_id="file-owned-by-openai", + metadata={"caller": "nightly-job"}, + litellm_metadata={"model_group": "openai-group"}, + original_function=_acreate_batch, + ) + + assert router.received_kwargs["metadata"] == {"caller": "nightly-job"} + assert router.received_kwargs["litellm_metadata"]["model_group"] == "openai-group" + + def test_get_fallback_model_group_does_not_mutate_fallbacks(): """A string fallback must be resolved without mutating the caller's fallbacks list, which is the live router config shared across requests.""" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 4a3395a7d3f..b76e69bc978 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -6755,6 +6755,68 @@ async def test_acreate_batch_disable_fallbacks_surfaces_owning_provider_error(): assert mock_create.call_args.kwargs["model"] == "owning-model" +@pytest.mark.asyncio +async def test_acreate_batch_surfaces_owning_provider_error_without_disable_fallbacks(): + """The router itself has to keep a batch inside the group that owns the input file: + the proxy only sets disable_fallbacks on the managed-files route, so the caller + otherwise gets the fallback provider's error for a file it never received.""" + from litellm.types.utils import LiteLLMBatch + + router = litellm.Router( + model_list=[ + { + "model_name": "owning-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "sk-owning", + }, + }, + { + "model_name": "fallback-model", + "litellm_params": { + "model": "azure/gpt-4o-mini", + "api_key": "sk-fallback", + "api_base": "https://fallback.openai.azure.com", + "api_version": "2024-08-01-preview", + }, + }, + ], + fallbacks=[{"owning-model": ["fallback-model"]}], + num_retries=0, + ) + attempted_models = [] + + async def _acreate_batch(model, **kwargs): + attempted_models.append(model) + if model == "owning-model": + raise litellm.APIConnectionError( + message="Connection error - openai is unreachable", + model="openai/gpt-4o-mini", + llm_provider="openai", + ) + return LiteLLMBatch( + id="batch-created-on-the-wrong-provider", + completion_window="24h", + created_at=0, + endpoint="/v1/chat/completions", + input_file_id="file-owned-by-openai", + object="batch", + status="validating", + ) + + with patch.object(router, "_acreate_batch", _acreate_batch): + with pytest.raises(litellm.APIConnectionError, match="openai is unreachable"): + await router.acreate_batch( + model="owning-model", + input_file_id="file-owned-by-openai", + endpoint="/v1/chat/completions", + completion_window="24h", + metadata={"team": "batch-jobs"}, + ) + + assert attempted_models == ["owning-model"] + + @pytest.mark.asyncio async def test_acreate_batch_request_bedrock_tags_override_deployment_tags(): import httpx From d7bc63da5c5eb44daeaaa2a876f757257bb5d68a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 7 Aug 2026 05:25:48 -0700 Subject: [PATCH 06/78] style(router): drop the inline comment on the fallback metadata merge --- litellm/router_utils/fallback_event_handlers.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index c4a84a1d61e..00df20be845 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -169,7 +169,7 @@ async def run_async_fallback( kwargs[metadata_variable_name] = { **(kwargs.get(metadata_variable_name) or {}), "model_group": kwargs.get("model", None), - } # update model_group used, if fallbacks are done + } fallback_depth = fallback_depth + 1 kwargs["fallback_depth"] = fallback_depth kwargs["max_fallbacks"] = max_fallbacks From 855c49d0ef01a09161f8d2ec195be01447669f1a Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 01:44:56 -0700 Subject: [PATCH 07/78] fix(proxy): skip prisma-dependent hooks when no database is attached --- .../storage_backend_service.py | 10 ++ litellm/proxy/utils.py | 5 + .../test_storage_backend_service.py | 127 ++++++++++++++++++ .../utils/proxy_logging/test_lifecycle.py | 94 +++++++++++-- 4 files changed, 228 insertions(+), 8 deletions(-) create mode 100644 tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index 4c301c96f30..e766f335071 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -68,6 +68,16 @@ class StorageBackendFileService: code=400, ) + if target_model_names: + managed_files_hook: Final = proxy_logging_obj.get_proxy_hook("managed_files") + if not isinstance(managed_files_hook, BaseFileEndpoints): + raise ProxyException( + message="Uploading with target_model_names requires a database-connected proxy, and this proxy has no database configured", + type="invalid_request_error", + param="target_model_names", + code=400, + ) + # Extract file information file_content: Final = file_data["content"] filename: Final = file_data.get("filename", "file") diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e59c6adaf22..5f22ca021ac 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -544,6 +544,11 @@ class ProxyLogging: for hook in PROXY_HOOKS: proxy_hook = get_proxy_hook(hook) expected_args = inspect.getfullargspec(proxy_hook).args + if "prisma_client" in expected_args and prisma_client is None: + verbose_proxy_logger.debug( + "Skipping proxy hook %s: it requires a database and no prisma client is configured", hook + ) + continue passed_in_args: dict[str, Any] = {} if "internal_usage_cache" in expected_args: passed_in_args["internal_usage_cache"] = self.internal_usage_cache diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py new file mode 100644 index 00000000000..07a85a70815 --- /dev/null +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_storage_backend_service.py @@ -0,0 +1,127 @@ +import pytest + +from litellm.llms.base_llm.files.transformation import BaseFileEndpoints +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.openai_files_endpoints import storage_backend_service +from litellm.proxy.openai_files_endpoints.storage_backend_service import ( + StorageBackendFileService, +) + + +class _RecordingStorageBackend: + def __init__(self): + self.upload_calls = [] + + async def upload_file(self, **kwargs): + self.upload_calls.append(kwargs) + return "https://storage.example/blob-1" + + +class _FakeManagedFilesHook(BaseFileEndpoints): + def __init__(self): + self.stored = [] + + async def acreate_file( + self, create_file_request, llm_router, target_model_names_list, litellm_parent_otel_span, user_api_key_dict + ): + raise NotImplementedError + + async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router=None): + raise NotImplementedError + + async def afile_list(self, purpose, litellm_parent_otel_span, **data): + raise NotImplementedError + + async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError + + async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data): + raise NotImplementedError + + async def store_unified_file_id(self, **kwargs): + self.stored.append(kwargs) + + +class _FakeProxyLogging: + def __init__(self, hook): + self._hook = hook + + def get_proxy_hook(self, hook_name): + return self._hook if hook_name == "managed_files" else None + + +def _file_data(): + return {"content": b"x", "filename": "input.jsonl", "content_type": "application/jsonl"} + + +@pytest.mark.asyncio +async def test_upload_with_target_model_names_but_no_hook_raises_before_uploading(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + + with pytest.raises(ProxyException) as exc_info: + await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=["gpt-x"], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=None), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "code": exc_info.value.code, + "message_names_requirement": "requires a database-connected proxy" in exc_info.value.message, + "upload_calls": backend.upload_calls, + } + assert snapshot == {"code": "400", "message_names_requirement": True, "upload_calls": []} + + +@pytest.mark.asyncio +async def test_upload_without_target_model_names_skips_hook_requirement(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + + file_object = await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=[], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=None), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "upload_count": len(backend.upload_calls), + "id_prefix": file_object.id.split("-")[0], + } + assert snapshot == {"upload_count": 1, "id_prefix": "file"} + + +@pytest.mark.asyncio +async def test_upload_with_target_model_names_and_hook_stores_unified_id(monkeypatch): + backend = _RecordingStorageBackend() + monkeypatch.setattr(storage_backend_service, "get_storage_backend", lambda name: backend) + hook = _FakeManagedFilesHook() + + file_object = await StorageBackendFileService.upload_file_to_storage_backend( + file_data=_file_data(), + target_storage="azure_storage", + target_model_names=["gpt-x"], + purpose="batch", + proxy_logging_obj=_FakeProxyLogging(hook=hook), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + + snapshot = { + "upload_count": len(backend.upload_calls), + "store_count": len(hook.stored), + "stored_id_matches_response": hook.stored[0]["file_id"] == file_object.id, + "model_mappings": hook.stored[0]["model_mappings"], + } + assert snapshot == { + "upload_count": 1, + "store_count": 1, + "stored_id_matches_response": True, + "model_mappings": {"gpt-x": "https://storage.example/blob-1"}, + } diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py index e33da672599..cf906259246 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_lifecycle.py @@ -8,7 +8,6 @@ because they are direct dependents on the lifecycle state. from __future__ import annotations -import asyncio from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -17,7 +16,6 @@ import pytest import litellm from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.utils import ( - InternalUsageCache, ProxyLogging, ) @@ -102,9 +100,7 @@ def test_update_values_with_no_args_is_noop(proxy_logging): def test_update_values_invalid_type_for_alerting_raises(proxy_logging): - proxy_logging.slack_alerting_instance = MagicMock( - update_values=MagicMock(side_effect=TypeError("bad type")) - ) + proxy_logging.slack_alerting_instance = MagicMock(update_values=MagicMock(side_effect=TypeError("bad type"))) with pytest.raises(TypeError): proxy_logging.update_values(alerting={"not": "a list"}) # type: ignore[arg-type] @@ -190,6 +186,90 @@ def test_add_proxy_hooks_registers_callbacks(proxy_logging, monkeypatch): } +def _stub_hook_classes(): + class _PrismaFreeHook: + def __init__(self, internal_usage_cache): + self.internal_usage_cache = internal_usage_cache + + class _PrismaRequiringHook: + def __init__(self, internal_usage_cache, prisma_client): + self.internal_usage_cache = internal_usage_cache + self.prisma_client = prisma_client + + class _PrismaOnlyHook: + def __init__(self, prisma_client): + self.prisma_client = prisma_client + + return { + "cache_control_check": _PrismaFreeHook, + "needs_db_hook": _PrismaRequiringHook, + "db_only_hook": _PrismaOnlyHook, + } + + +def test_add_proxy_hooks_skips_prisma_requiring_hook_when_no_db(proxy_logging, monkeypatch): + hook_classes = _stub_hook_classes() + registered: List[Any] = [] + + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", list(hook_classes.keys())) + monkeypatch.setattr(utils_mod, "get_proxy_hook", hook_classes.__getitem__) + monkeypatch.setattr( + litellm.logging_callback_manager, + "add_litellm_callback", + lambda cb: registered.append(cb), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", None): + proxy_logging._add_proxy_hooks(llm_router=None) + + snapshot = { + "mapping_keys": list(proxy_logging.proxy_hook_mapping.keys()), + "registered_types": [type(r).__name__ for r in registered], + "needs_db_hook_lookup": proxy_logging.get_proxy_hook("needs_db_hook"), + "db_only_hook_lookup": proxy_logging.get_proxy_hook("db_only_hook"), + } + assert snapshot == { + "mapping_keys": ["cache_control_check"], + "registered_types": ["_PrismaFreeHook"], + "needs_db_hook_lookup": None, + "db_only_hook_lookup": None, + } + + +def test_add_proxy_hooks_registers_prisma_requiring_hook_with_db(proxy_logging, monkeypatch): + hook_classes = _stub_hook_classes() + registered: List[Any] = [] + fake_prisma = MagicMock() + + from litellm.proxy import utils as utils_mod + + monkeypatch.setattr(utils_mod, "PROXY_HOOKS", list(hook_classes.keys())) + monkeypatch.setattr(utils_mod, "get_proxy_hook", hook_classes.__getitem__) + monkeypatch.setattr( + litellm.logging_callback_manager, + "add_litellm_callback", + lambda cb: registered.append(cb), + ) + + with patch("litellm.proxy.proxy_server.prisma_client", fake_prisma): + proxy_logging._add_proxy_hooks(llm_router=None) + + snapshot = { + "mapping_keys": list(proxy_logging.proxy_hook_mapping.keys()), + "registered_count": len(registered), + "needs_db_hook_got_prisma": proxy_logging.proxy_hook_mapping["needs_db_hook"].prisma_client is fake_prisma, + "db_only_hook_got_prisma": proxy_logging.proxy_hook_mapping["db_only_hook"].prisma_client is fake_prisma, + } + assert snapshot == { + "mapping_keys": ["cache_control_check", "needs_db_hook", "db_only_hook"], + "registered_count": 3, + "needs_db_hook_got_prisma": True, + "db_only_hook_got_prisma": True, + } + + def test_add_proxy_hooks_unknown_hook_raises(proxy_logging, monkeypatch): from litellm.proxy import utils as utils_mod @@ -267,9 +347,7 @@ def test_init_litellm_callbacks_replaces_string_with_instance(proxy_logging, mon snapshot = { "replaced_first_item": litellm.callbacks[0] is sentinel_instance, "callbacks_grew_with_service": len(litellm.callbacks) >= 2, - "service_logging_appended": any( - "ServiceLogging" in type(c).__name__ for c in litellm.callbacks - ), + "service_logging_appended": any("ServiceLogging" in type(c).__name__ for c in litellm.callbacks), } assert snapshot == { "replaced_first_item": True, From 8812debeff43b077c64c2cbf2f2d2dd4e3187639 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 20:58:30 +0000 Subject: [PATCH 08/78] docs: allow functional comments as an exception in CLAUDE.md Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index f1bb46c1fd3..046932274cd 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. Exceptions are granted for comments that do something rather than document something for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From abbad8ad528fb6a36c037ad90e9a210fd3435753 Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 21:04:43 +0000 Subject: [PATCH 09/78] docs: limit the comment exception to tool-read directives Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index 046932274cd..4901346bc6e 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. Exceptions are granted for comments that do something rather than document something for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The one exception is a comment a tool reads and acts on, as opposed to one documenting code for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed. Human-readable annotations like TODO, FIXME, and section headers don't qualify Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 6e6e0d662b4db1d46552ca011b19b5775f0db2ad Mon Sep 17 00:00:00 2001 From: mateo Date: Sat, 8 Aug 2026 21:07:40 +0000 Subject: [PATCH 10/78] docs: frame the comment rule around AI slop and allow TODO/FIXME Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- CLAUDE.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index 4901346bc6e..b249e2d01fc 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,4 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The one exception is a comment a tool reads and acts on, as opposed to one documenting code for humans and agents: an entry in `.git-blame-ignore-revs` needs its comment to say which commit is being excluded from git blame, and a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` is the only way to silence a violation when introducing one is truly unavoidable. Write those, and the reasons they require, wherever they're needed. Human-readable annotations like TODO, FIXME, and section headers don't qualify +Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The point of that rule is to keep out AI slop comments that just restate what the code already says, so comments carrying information the code can't are fine. That covers comments a tool reads and acts on, such as an entry in `.git-blame-ignore-revs` saying which commit is excluded from git blame, or a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a violation is truly unavoidable, and it covers a real TODO or FIXME flagging known unfinished work. Write those, and the reasons they require, wherever they're needed Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 1f8964a4c4ab2d77eae31ee2a9bf8de2ae6c0f14 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 14:20:13 -0700 Subject: [PATCH 11/78] chore: handwrite the rule --- CLAUDE.md | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/CLAUDE.md b/CLAUDE.md index b249e2d01fc..6990408e686 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,4 +1,12 @@ -Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt. The point of that rule is to keep out AI slop comments that just restate what the code already says, so comments carrying information the code can't are fine. That covers comments a tool reads and acts on, such as an entry in `.git-blame-ignore-revs` saying which commit is excluded from git blame, or a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a violation is truly unavoidable, and it covers a real TODO or FIXME flagging known unfinished work. Write those, and the reasons they require, wherever they're needed +Do not write comments unless they are: +- absolutely necessary to explain some very complex business logic +- used as an input for tools to read and act on. For example: + - entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame + - a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a truly unavoidable violation +- a TODO or FIXME + - Not great to have those, but if it's unavoidable, make sure to include a strong reason for why it's there or, better yet, link to a GitHub issue for the follow-up work + +Explanation: code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From 805fc497761a47c82ceb062b10945434e37cc53b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 14:29:36 -0700 Subject: [PATCH 12/78] chore: mention AI slop reason --- CLAUDE.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/CLAUDE.md b/CLAUDE.md index 6990408e686..dcdfe2d15b9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,12 +1,12 @@ Do not write comments unless they are: -- absolutely necessary to explain some very complex business logic +- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear) - used as an input for tools to read and act on. For example: - entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame - a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # ` when introducing a truly unavoidable violation - a TODO or FIXME - - Not great to have those, but if it's unavoidable, make sure to include a strong reason for why it's there or, better yet, link to a GitHub issue for the follow-up work + - Not great to have those, but if it's unavoidable, make sure to include a strong, concise reason for why it's there or, better yet, link to a GitHub issue for the follow-up work -Explanation: code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance +Explanation: The point of this rule is to keep out AI slop comments. AI writes way too many and way too verbose comments. Code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in: From efc4e6f28c0951dc20fc61e8cb9be53834fdb7b7 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Sat, 8 Aug 2026 16:01:47 -0700 Subject: [PATCH 13/78] fix(batches): keep batch state in sync on a poll without claiming attribution (#34456) A poll of a Vertex passthrough batch wrote nothing to the managed-object row, so status and file_object stayed frozen at the create-time snapshot and GET /v1/batches served a stale status and an empty output file id for the life of the batch. Only the create may claim a batch, but every observation of one may refresh its state. store_unified_object_id takes create_if_missing, which the poll clears: it refreshes status and file_object through update_many, and leaves a row that is absent absent rather than creating one owned by the observer, since created_by and team_id are written by whoever reaches the create branch. The update payload is now shared with the upsert so it cannot drift into writing api_key, request_tags, created_by or team_id. The passthrough identity re-assertion that was previously part of this PR ships separately in #36121, so this PR keeps only the batch attribution work. The creating key owns user_api_key_alias only when it actually has one. Guarding the overwrite on the presence of a key rather than on a resolved alias nulled the field out for every key generated without key_alias, and for any key rotated or deleted before its batch finished, losing the creating user's alias that the spend row previously carried. The guard now matches the team-alias line below it. --- .../proxy/common_utils/check_batch_cost.py | 77 +++++++-- .../proxy/hooks/managed_files.py | 49 +++++- .../migration.sql | 5 + .../litellm_proxy_extras/schema.prisma | 2 + .../proxy/hooks/proxy_track_cost_callback.py | 23 ++- .../vertex_passthrough_logging_handler.py | 76 ++++++++- litellm/proxy/schema.prisma | 2 + schema.prisma | 2 + .../proxy_unit_tests/test_check_batch_cost.py | 150 ++++++++++++++++ ..._batch_update_db_managed_output_file_id.py | 139 +++++++++++++++ .../proxy/test_managed_files_hook.py | 111 ++++++++++++ .../hooks/test_proxy_track_cost_callback.py | 161 ++++++++++++++++++ .../test_vertex_ai_batch_passthrough.py | 136 +++++++++++++++ 13 files changed, 906 insertions(+), 27 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 7acdd5dbdaf..dc8f17fb665 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -3,7 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -43,11 +43,15 @@ class CheckBatchCost: # the guaranteed-failing primary query on every subsequent cycle. self._has_batch_processed_column: bool = True - async def _get_user_info(self, batch_id, user_id) -> dict: + async def _get_user_info(self, batch_id: str, user_id: Optional[str]) -> Dict[str, Any]: """ Look up user email and key alias by user_id for enriching the S3 callback metadata. Returns a dict with user_api_key_user_email and user_api_key_alias (both may be None). + Returns an empty dict when user_id is None: batches created by a team or service + account key carry no user id, and find_unique(where={"user_id": None}) raises. """ + if not user_id: + return {} try: user_row = await self.prisma_client.db.litellm_usertable.find_unique( where={"user_id": user_id} @@ -62,6 +66,66 @@ class CheckBatchCost: verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}") return {} + async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None: + """Resolve the creating virtual key's alias from its hashed token.""" + if not api_key: + return None + try: + key_row = await self.prisma_client.db.litellm_verificationtoken.find_unique( + where={"token": api_key} + ) + return getattr(key_row, "key_alias", None) if key_row is not None else None + except Exception as e: + verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}") + return None + + async def _get_team_alias(self, team_id: str | None) -> str | None: + """Resolve a team's alias from its id.""" + if not team_id: + return None + try: + team_row = await self.prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) + return getattr(team_row, "team_alias", None) if team_row is not None else None + except Exception as e: + verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}") + return None + + async def _build_creator_attribution_metadata( + self, job: "LiteLLM_ManagedObjectTable", batch_id: str + ) -> Dict[str, Any]: + """ + Rebuild the spend-tracking metadata for the key, team, and tags that created the + batch so the batch-cost spend log is attributed the same way a non-batch request + is. Rows created before api_key and request_tags were persisted carry only + created_by and team_id, and fall back to those. A named creating key owns + user_api_key_alias; when it has no alias, or the key has since been rotated or + deleted, the field keeps the creating user's alias that _get_user_info filled in, + because a resolvable name is more useful on the spend row than a null. + """ + api_key = getattr(job, "api_key", None) + team_id = getattr(job, "team_id", None) + request_tags = getattr(job, "request_tags", None) + + metadata: Dict[str, Any] = { + "user_api_key_user_id": job.created_by, + "user_api_key": api_key, + "user_api_key_team_id": team_id, + **(await self._get_user_info(batch_id, job.created_by)), + } + + key_alias = await self._get_key_alias(batch_id, api_key) + if key_alias is not None: + metadata["user_api_key_alias"] = key_alias + team_alias = await self._get_team_alias(team_id) + if team_alias is not None: + metadata["user_api_key_team_alias"] = team_alias + if isinstance(request_tags, list) and request_tags: + metadata["tags"] = [tag for tag in request_tags if isinstance(tag, str)] + + return metadata + async def _cleanup_stale_managed_objects(self) -> None: """ Mark managed objects older than MANAGED_OBJECT_STALENESS_CUTOFF_DAYS days @@ -485,9 +549,6 @@ class CheckBatchCost: function_id=str(uuid.uuid4()), ) - creator_user_id = job.created_by - user_info = await self._get_user_info(batch_id, job.created_by) - logging_obj.update_environment_variables( litellm_params={ # set the user-agent header so that S3 callback consumers can easily identify CheckBatchCost callbacks @@ -496,11 +557,7 @@ class CheckBatchCost: "user-agent": CHECK_BATCH_COST_USER_AGENT, } }, - "metadata": { - "user_api_key_user_id": creator_user_id, - "user_api_key_team_id": getattr(job, "team_id", None), - **user_info, - }, + "metadata": await self._build_creator_attribution_metadata(job, batch_id), }, optional_params={}, ) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index f0914240f79..a3994dccfd6 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -163,6 +163,8 @@ class _ManagedObjectTableActions(Protocol): self, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]] ) -> "PrismaManagedObjectRow": ... + async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int: ... + class _CursorPageArgs(TypedDict, total=False): cursor: Mapping[str, str] @@ -263,7 +265,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): model_object_id: str, file_purpose: Literal["batch", "fine-tune", "response"], user_api_key_dict: UserAPIKeyAuth, + request_tags: Sequence[str] | None = None, + persist_attribution: bool = False, + create_if_missing: bool = True, ) -> None: + """Persist a managed object row, caching it and upserting it in the DB. + + persist_attribution is set only by the batch create, which is the one caller + that can speak for the creator; it gates the api_key and request_tags columns + that CheckBatchCost bills against, so a later poll or retrieve of the same + batch cannot record itself as the paying key. Like created_by and team_id, + both are written only in the upsert create branch, never on update. + + create_if_missing is cleared by callers that observe a batch they did not + create, such as a poll. They still refresh status and file_object, but a + row absent from the table is left absent rather than created with the + observer as its creator, because created_by and team_id are written from + whoever calls the create branch. + """ verbose_logger.info(f"Storing LiteLLM Managed {file_purpose} object with id={unified_object_id} in cache") litellm_managed_object = LiteLLM_ManagedObjectTable( unified_object_id=unified_object_id, @@ -277,6 +296,29 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): litellm_parent_otel_span=litellm_parent_otel_span, ) + from prisma import Json + + api_key = user_api_key_dict.api_key or None + attribution_columns = ( + { + **({"api_key": api_key} if api_key is not None else {}), + **({"request_tags": Json(list(request_tags))} if request_tags else {}), + } + if persist_attribution + else {} + ) + # FIX: Update status and file_object on every operation to keep state in sync + update_columns: Final = { + "file_object": file_object.model_dump_json(), + "status": file_object.status, + "updated_by": user_api_key_dict.user_id, + } + if not create_if_missing: + await _managed_object_table(self.prisma_client).update_many( + where={"unified_object_id": unified_object_id}, + data=update_columns, + ) + return await _managed_object_table(self.prisma_client).upsert( where={"unified_object_id": unified_object_id}, data={ @@ -289,12 +331,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "team_id": user_api_key_dict.team_id, "updated_by": user_api_key_dict.user_id, "status": file_object.status, + **attribution_columns, }, - "update": { - "file_object": file_object.model_dump_json(), - "status": file_object.status, - "updated_by": user_api_key_dict.user_id, - }, # FIX: Update status and file_object on every operation to keep state in sync + "update": update_columns, }, ) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql new file mode 100644 index 00000000000..79bc6b24de8 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260730000000_add_api_key_and_request_tags_to_managed_object_table/migration.sql @@ -0,0 +1,5 @@ +-- Add api_key and request_tags columns to LiteLLM_ManagedObjectTable +-- Captured at batch-create time so CheckBatchCost can attribute batch-cost spend +-- back to the creating virtual key (and its tags) even when created_by is null. +ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "api_key" TEXT; +ALTER TABLE "LiteLLM_ManagedObjectTable" ADD COLUMN IF NOT EXISTS "request_tags" JSONB DEFAULT '[]'; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 3346f9d7e3b..0e22b5324c1 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -48,6 +48,15 @@ _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset( } ) +# Both spellings, because call_type reaches the callback as str(...) of either the +# enum member or its value. +_CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( + ( + CallTypes.aretrieve_batch.value, + str(CallTypes.aretrieve_batch), + ) +) + class _ProxyDBLogger(CustomLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -212,7 +221,10 @@ class _ProxyDBLogger(CustomLogger): # Only fetch key details when user_id wasn't already populated (e.g. direct MCP REST calls). # Avoids a cache/DB lookup on every normal LLM request. if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): - metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata=metadata) + metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original + metadata=metadata, + resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES, + ) _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata) user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None)) @@ -337,7 +349,7 @@ class _ProxyDBLogger(CustomLogger): spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) @staticmethod - async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: + async def _enrich_failure_metadata_with_key_info(metadata: dict, resolve_missing_key_identity: bool = True) -> dict: """ Enriches failure spend log metadata by looking up the key object (and team object) from cache/DB when key fields are missing. @@ -349,6 +361,11 @@ class _ProxyDBLogger(CustomLogger): 2. Post-auth failures (provider errors, rate limits): key fields are populated but team_alias is missing because LiteLLM_VerificationTokenView SQL view doesn't include it. We look up the team object to fill in team_alias. + + Scenario 1 reads the key's identity as it stands right now, so it is only correct + for a log emitted within the request it describes. Callers that log after a delay, + against an identity captured earlier, pass resolve_missing_key_identity=False and + keep their own user_id, team_id and org_id. """ api_key_hash: Final = metadata.get("user_api_key") if not api_key_hash: @@ -361,7 +378,7 @@ class _ProxyDBLogger(CustomLogger): ) # Step 1: If key fields are missing, look up the full key object - if metadata.get("user_api_key_alias") is None: + if resolve_missing_key_identity and metadata.get("user_api_key_alias") is None: try: key_obj: Final = await get_key_object( hashed_token=api_key_hash, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 2e3f7bb9aa6..9c5b7dc563e 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -1,4 +1,6 @@ +import asyncio import re +from collections.abc import Mapping from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast from urllib.parse import urlparse @@ -39,6 +41,32 @@ else: EndpointType = Any +def _optional_str(value: object) -> str | None: + return value if isinstance(value, str) else None + + +def _optional_str_tuple(value: object) -> tuple[str, ...] | None: + if not isinstance(value, list): + return None + items: Final = cast(list[object], value) # cast-ok: isinstance-narrowed; element type unknown + return tuple(tag for tag in items if isinstance(tag, str)) + + +def _request_tags(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None: + """Tags for the batch-cost spend row: the request's own tags when it sent any, + otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a + tagged key does not put its tags in the top-level metadata "tags" on the + passthrough path) + """ + tags: Final = _optional_str_tuple(request_metadata.get("tags")) + if tags: + return tags + key_auth_metadata: Final = request_metadata.get("user_api_key_auth_metadata") + if isinstance(key_auth_metadata, dict): + return _optional_str_tuple(key_auth_metadata.get("tags")) + return None + + class VertexPassthroughLoggingHandler: @staticmethod def vertex_passthrough_handler( @@ -657,11 +685,13 @@ class VertexPassthroughLoggingHandler: # Store the managed object for cost tracking # This will be picked up by check_batch_cost polling mechanism + is_batch_create: Final = url_route.split("?")[0].rstrip("/").endswith("batchPredictionJobs") VertexPassthroughLoggingHandler._store_batch_managed_object( unified_object_id=unified_object_id, batch_object=litellm_batch_response, model_object_id=batch_id, logging_obj=logging_obj, + is_batch_create=is_batch_create, **kwargs, ) @@ -779,17 +809,45 @@ class VertexPassthroughLoggingHandler: "kwargs": kwargs, } + @staticmethod + def _log_batch_registration_result( + finished: asyncio.Task, unified_object_id: str, model_object_id: str, is_batch_create: bool + ) -> None: + error: Final = finished.exception() if not finished.cancelled() else None + if finished.cancelled() or error is not None: + consequence: Final = ( + "its cost will not be tracked" if is_batch_create else "its status and output file may be stale" + ) + verbose_proxy_logger.error( + "Failed to store batch managed object with unified_object_id=%s, batch_id=%s; %s: %s", + unified_object_id, + model_object_id, + consequence, + error, + ) + return + verbose_proxy_logger.info( + "Stored batch managed object with unified_object_id=%s, batch_id=%s", + unified_object_id, + model_object_id, + ) + @staticmethod def _store_batch_managed_object( unified_object_id: str, batch_object: LiteLLMBatch, model_object_id: str, logging_obj: LiteLLMLoggingObj, + is_batch_create: bool, **kwargs, ) -> None: """ Store batch managed object for cost tracking. This will be picked up by the check_batch_cost polling mechanism. + + A poll refreshes the batch status and file object but neither creates the row + nor writes attribution, so the creating key and its tags are persisted from + the create alone. """ try: # Get the managed files hook from the logging object @@ -805,7 +863,7 @@ class VertexPassthroughLoggingHandler: user_api_key_dict: Final = UserAPIKeyAuth( user_id=_request_metadata.get("user_api_key_user_id", "default-user"), - api_key="", + api_key=_optional_str(_request_metadata.get("user_api_key")), team_id=_request_metadata.get("user_api_key_team_id"), team_alias=None, user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value @@ -827,9 +885,7 @@ class VertexPassthroughLoggingHandler: ) # Store the unified object for batch cost tracking - import asyncio - - asyncio.create_task( + task: Final = asyncio.create_task( managed_files_hook.store_unified_object_id( unified_object_id=unified_object_id, file_object=batch_object, @@ -837,13 +893,15 @@ class VertexPassthroughLoggingHandler: model_object_id=model_object_id, file_purpose="batch", user_api_key_dict=user_api_key_dict, + request_tags=_request_tags(_request_metadata), + persist_attribution=is_batch_create, + create_if_missing=is_batch_create, ) ) - - verbose_proxy_logger.info( - "Stored batch managed object with unified_object_id=%s, batch_id=%s", - unified_object_id, - model_object_id, + task.add_done_callback( + lambda finished: VertexPassthroughLoggingHandler._log_batch_registration_result( + finished, unified_object_id, model_object_id, is_batch_create + ) ) else: verbose_proxy_logger.warning( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/schema.prisma b/schema.prisma index 9c871b65f40..cabddf6f1a1 100644 --- a/schema.prisma +++ b/schema.prisma @@ -985,6 +985,8 @@ model LiteLLM_ManagedObjectTable { // for batches or finetuning jobs which use t created_at DateTime @default(now()) created_by String? team_id String? + api_key String? + request_tags Json? @default("[]") updated_at DateTime @updatedAt updated_by String? diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 800c97d7ba1..6d7ada17ec5 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -1641,3 +1641,153 @@ class TestManagedOutputFileIdEncodesPublicModelGroup: decoded = _is_base64_encoded_unified_file_id(output_file_id) assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP] +class TestBatchCostAttribution: + """CheckBatchCost rebuilds the creator's spend metadata from the managed-object row so + the batch-cost log is attributed like a non-batch request.""" + + def _instance(self, key_row=None, team_row=None, user_row=None): + from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost + + prisma = MagicMock() + prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_row) + prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + return CheckBatchCost( + proxy_logging_obj=MagicMock(), + prisma_client=prisma, + llm_router=MagicMock(), + ) + + def _job(self, **overrides): + from types import SimpleNamespace + + fields = { + "created_by": "alice", + "team_id": "team-alpha", + "api_key": "hash-alice", + "request_tags": ["env:prod"], + } + fields.update(overrides) + return SimpleNamespace(unified_object_id="uoi", **fields) + + @pytest.mark.asyncio + async def test_metadata_carries_key_team_and_tags(self): + """The spend row names the creating key, its team, both aliases, and the tags.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias="prod-key"), + team_row=SimpleNamespace(team_alias="Team Alpha"), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias=None), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata["user_api_key_user_id"] == "alice" + assert metadata["user_api_key_team_id"] == "team-alpha" + assert metadata["user_api_key_alias"] == "prod-key" + assert metadata["user_api_key_team_alias"] == "Team Alpha" + assert metadata["tags"] == ["env:prod"] + + @pytest.mark.asyncio + async def test_metadata_tolerates_legacy_row_without_columns(self): + """Rows created before the columns existed carry only created_by/team_id and must + still produce an attributed row rather than raising.""" + instance = self._instance() + job = self._job(api_key=None, request_tags=None) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["user_api_key"] is None + assert metadata["user_api_key_user_id"] == "alice" + assert metadata["user_api_key_team_id"] == "team-alpha" + assert "tags" not in metadata + + @pytest.mark.asyncio + async def test_metadata_keeps_key_when_team_key_has_no_user(self): + """A team-scoped key carries no user id. The user lookup is skipped (prisma rejects + a None user_id) and the key hash still drives key-level attribution.""" + from types import SimpleNamespace + + instance = self._instance(key_row=SimpleNamespace(key_alias="svc-key")) + job = self._job(created_by=None) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata["user_api_key_user_id"] is None + assert metadata["user_api_key_alias"] == "svc-key" + instance.prisma_client.db.litellm_usertable.find_unique.assert_not_called() + + @pytest.mark.asyncio + async def test_metadata_drops_non_string_tags(self): + """Non-string tags are dropped so a malformed stored value cannot slip past the + tag-budget checks that consume this metadata.""" + instance = self._instance() + job = self._job(request_tags=["env:prod", 7, None, "team:ml"]) + + metadata = await instance._build_creator_attribution_metadata(job, "batch-1") + + assert metadata["tags"] == ["env:prod", "team:ml"] + + @pytest.mark.asyncio + async def test_key_alias_lookup_failure_does_not_break_attribution(self): + """An alias lookup failure must not lose the spend row; the key hash and team still + attribute it.""" + instance = self._instance() + instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + side_effect=Exception("db down") + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key"] == "hash-alice" + assert metadata.get("user_api_key_alias") is None + + @pytest.mark.asyncio + async def test_unnamed_key_keeps_the_creating_user_alias(self): + """Regression: a key generated without key_alias resolves to no alias, and the + overwrite must not null out the creating user's alias that _get_user_info supplied. + Most keys carry no alias, so this is the common batch, not an edge case.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias=None), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "Alice Chen" + assert metadata["user_api_key"] == "hash-alice" + + @pytest.mark.asyncio + async def test_rotated_key_keeps_the_creating_user_alias(self): + """Batches outlive keys. When the creating key has been rotated or deleted the + lookup returns no row, and the spend log keeps a resolvable name instead of null.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=None, + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "Alice Chen" + + @pytest.mark.asyncio + async def test_named_key_still_owns_the_alias(self): + """The fallback must not weaken the intended precedence: a key that has its own + alias still overrides the creating user's.""" + from types import SimpleNamespace + + instance = self._instance( + key_row=SimpleNamespace(key_alias="prod-key"), + user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"), + ) + + metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1") + + assert metadata["user_api_key_alias"] == "prod-key" diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index 1b60c97b510..ebd33aa2e53 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -356,3 +356,142 @@ async def test_ensure_batch_response_returns_early_without_auth(): assert response.output_file_id == "file-raw-output" mock_managed_files.get_unified_output_file_id.assert_not_called() + + +def _in_memory_managed_files(): + """Build a real _PROXY_LiteLLMManagedFiles whose prisma upsert hits an in-memory row.""" + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + + store: dict = {} + + async def _upsert(where, data): + key = where["unified_object_id"] + if key in store: + store[key].update(data["update"]) + else: + store[key] = dict(data["create"]) + + table = MagicMock() + table.upsert = AsyncMock(side_effect=_upsert) + prisma = MagicMock() + prisma.db.litellm_managedobjecttable = table + + cache = MagicMock() + cache.async_set_cache = AsyncMock() + + return ( + _PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma), + store, + ) + + +@pytest.mark.asyncio +async def test_store_unified_object_id_persists_key_and_tags_on_create(): + """Regression (spend loss): the batch create persists the creating key hash and tags so + CheckBatchCost can write an attributed spend row instead of a blank one the DB drops.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key="hash-alice") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=["env:prod"], + persist_attribution=True, + ) + + row = store["unified-b"] + assert row["api_key"] == "hash-alice" + assert row["created_by"] == "alice" + assert row["team_id"] == "team-alpha" + assert row["request_tags"].data == ["env:prod"] + + +@pytest.mark.asyncio +async def test_store_unified_object_id_omits_key_and_tags_without_persist_attribution(): + """Regression (spend redirect): a caller that is not the batch create (a poll, or the + generic post-call hook on a retrieve) carries a real hashed key, but must never have it + recorded as the batch's paying key. created_by/team_id keep their existing behavior.""" + instance, store = _in_memory_managed_files() + poller = UserAPIKeyAuth(user_id="bob", team_id="team-bravo", api_key="hash-bob") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="in_progress"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=["env:dev"], + ) + + row = store["unified-b"] + assert "api_key" not in row + assert "request_tags" not in row + assert row["created_by"] == "bob" + + +@pytest.mark.asyncio +async def test_store_unified_object_id_attribution_columns_are_write_once(): + """Identity is written only in the upsert create branch, so a later store for the same + batch (a status update, a poll) can neither reassign the paying key nor clear it.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key="hash-alice") + poller = UserAPIKeyAuth(user_id="bob", team_id="team-bravo", api_key="hash-bob") + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=["env:prod"], + persist_attribution=True, + ) + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="completed"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=["poller-tag"], + persist_attribution=True, + ) + + row = store["unified-b"] + assert row["api_key"] == "hash-alice" + assert row["created_by"] == "alice" + assert row["status"] == "completed" + + upsert_data = instance.prisma_client.db.litellm_managedobjecttable.upsert.call_args.kwargs["data"] + assert "api_key" not in upsert_data["update"] + assert "request_tags" not in upsert_data["update"] + + +@pytest.mark.asyncio +async def test_store_unified_object_id_omits_unset_columns(): + """A batch created with no tags (the common case) still registers: the optional columns + are omitted rather than passed as None, which prisma rejects for the Json column.""" + instance, store = _in_memory_managed_files() + creator = UserAPIKeyAuth(user_id="alice", team_id="team-alpha", api_key=None) + + await instance.store_unified_object_id( + unified_object_id="unified-b", + file_object=_build_batch_response(batch_id="b", status="validating"), + litellm_parent_otel_span=None, + model_object_id="b", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=None, + persist_attribution=True, + ) + + create_data = instance.prisma_client.db.litellm_managedobjecttable.upsert.call_args.kwargs["data"]["create"] + assert "api_key" not in create_data + assert "request_tags" not in create_data + assert "unified-b" in store 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 63646ae53f8..34cd0cabc2c 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -510,6 +510,117 @@ def _make_real_managed_files_instance(): ) +def _make_object_store_instance(): + """A real store_unified_object_id over an AsyncMock prisma client, so both the + upsert and the update-only write path 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_managedobjecttable.upsert = AsyncMock() + mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock() + + return ( + _PROXY_LiteLLMManagedFiles( + internal_usage_cache=mock_cache, + prisma_client=mock_prisma, + ), + mock_prisma, + ) + + +@pytest.mark.asyncio +async def test_poll_refreshes_batch_state_without_claiming_the_row(): + """Regression (stale batch state): a poll observes a batch it did not create, so it + must still refresh status and file_object -- otherwise GET /v1/batches serves the + create-time snapshot forever -- while writing none of the attribution columns and + never creating a row it would then own.""" + managed_files, mock_prisma = _make_object_store_instance() + poller = UserAPIKeyAuth( + api_key="sk-the-poller", user_id="bob", team_id="team-bravo", parent_otel_span=None + ) + + await managed_files.store_unified_object_id( + unified_object_id="uoi-1", + file_object=_make_batch_response(status="completed"), + litellm_parent_otel_span=None, + model_object_id="batch-123", + file_purpose="batch", + user_api_key_dict=poller, + request_tags=("poller:tag",), + persist_attribution=False, + create_if_missing=False, + ) + + # the row is refreshed in place, and cannot be conjured by a poll + mock_prisma.db.litellm_managedobjecttable.upsert.assert_not_awaited() + update_many = mock_prisma.db.litellm_managedobjecttable.update_many + update_many.assert_awaited_once() + call = update_many.await_args + assert call.kwargs["where"] == {"unified_object_id": "uoi-1"} + + written = call.kwargs["data"] + assert written["status"] == "completed" + assert json.loads(written["file_object"])["output_file_id"] == "file-output-abc" + # nothing the poller could be billed for + for owned in ("api_key", "request_tags", "created_by", "team_id"): + assert owned not in written + + +@pytest.mark.asyncio +async def test_create_still_upserts_and_claims_attribution(): + """The create is the one caller that can speak for the batch, so it keeps the upsert + (creating the row when absent) and writes the attribution columns.""" + managed_files, mock_prisma = _make_object_store_instance() + creator = UserAPIKeyAuth( + api_key="sk-the-creator", user_id="alice", team_id="team-alpha", parent_otel_span=None + ) + + await managed_files.store_unified_object_id( + unified_object_id="uoi-2", + file_object=_make_batch_response(status="validating"), + litellm_parent_otel_span=None, + model_object_id="batch-456", + file_purpose="batch", + user_api_key_dict=creator, + request_tags=("env:prod",), + persist_attribution=True, + ) + + mock_prisma.db.litellm_managedobjecttable.update_many.assert_not_awaited() + upsert = mock_prisma.db.litellm_managedobjecttable.upsert + upsert.assert_awaited_once() + created = upsert.await_args.kwargs["data"]["create"] + # UserAPIKeyAuth hashes an sk- token on construction; the hash is what is billed + assert created["api_key"] == creator.api_key + assert created["api_key"] != "sk-the-creator" + assert created["created_by"] == "alice" + assert created["team_id"] == "team-alpha" + + +@pytest.mark.asyncio +async def test_default_callers_still_create_their_rows(): + """create_if_missing defaults to True, so the fine-tune, Responses and Anthropic + callers, none of which pass it, keep upserting exactly as before.""" + managed_files, mock_prisma = _make_object_store_instance() + + await managed_files.store_unified_object_id( + unified_object_id="uoi-3", + file_object=_make_batch_response(), + litellm_parent_otel_span=None, + model_object_id="ft-789", + file_purpose="fine-tune", + user_api_key_dict=_make_user_api_key_dict(), + ) + + mock_prisma.db.litellm_managedobjecttable.upsert.assert_awaited_once() + mock_prisma.db.litellm_managedobjecttable.update_many.assert_not_awaited() + + @pytest.mark.asyncio async def test_store_unified_file_id_is_idempotent_via_upsert(): """Regression test for the managed-batch retrieve 500 (UniqueViolationError on diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 69f04ce2bbe..2b162774aea 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,6 +17,7 @@ from litellm.proxy.hooks.proxy_track_cost_callback import ( _should_track_cost_callback, _update_database_and_spend_counters, ) +from litellm.types.utils import CallTypes @pytest.mark.asyncio @@ -783,6 +784,166 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): mock_get_key.assert_not_called() +@pytest.mark.asyncio +async def test_enrich_failure_metadata_keeps_captured_identity_when_not_resolving(): + """ + With resolve_missing_key_identity=False the key is not read, so a null user_id, + team_id and org_id captured earlier stay null instead of being refilled from the + key as it stands now. The team_alias lookup still runs off the captured team_id. + """ + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "alias-assigned-later" + mock_key_obj.user_id = "user-assigned-later" + mock_key_obj.team_id = "team-assigned-later" + mock_key_obj.org_id = "org-assigned-later" + + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "captured-team-alias" + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ), + ): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": None, + "user_api_key_user_id": None, + "user_api_key_team_id": "captured-team-id", + "user_api_key_team_alias": None, + "user_api_key_org_id": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata, resolve_missing_key_identity=False + ) + + mock_get_key.assert_not_called() + assert result["user_api_key_user_id"] is None + assert result["user_api_key_team_id"] == "captured-team-id" + assert result["user_api_key_org_id"] is None + assert result["user_api_key_alias"] is None + assert result["user_api_key_team_alias"] == "captured-team-alias" + + +@pytest.mark.asyncio +async def test_enrich_failure_metadata_ignores_flag_when_alias_present(): + """ + A captured alias already closes the key lookup, so resolve_missing_key_identity + changes nothing for a key that has one; only the alias-less key depends on it. + """ + mock_team_obj = MagicMock() + mock_team_obj.team_alias = "captured-team-alias" + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=mock_team_obj, + ), + ): + for resolve in (True, False): + metadata = { + "user_api_key": "hashed_key", + "user_api_key_alias": "captured-alias", + "user_api_key_user_id": None, + "user_api_key_team_id": "captured-team-id", + "user_api_key_team_alias": None, + "user_api_key_org_id": None, + } + result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( + metadata, resolve_missing_key_identity=resolve + ) + mock_get_key.assert_not_called() + assert result["user_api_key_user_id"] is None + assert result["user_api_key_alias"] == "captured-alias" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "call_type, expect_key_read", + [ + (CallTypes.aretrieve_batch.value, False), + (CallTypes.aretrieve_batch, False), + (CallTypes.acompletion.value, True), + ], +) +async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, expect_key_read): + """ + The batch cost row is logged long after the batch was created, so it keeps the + identity persisted at create time. Every other call type still backfills from + the key. + """ + logger = _ProxyDBLogger() + + mock_key_obj = MagicMock() + mock_key_obj.key_alias = "alias-assigned-later" + mock_key_obj.user_id = "user-assigned-later" + mock_key_obj.team_id = "team-assigned-later" + mock_key_obj.org_id = "org-assigned-later" + + kwargs = { + "call_type": call_type, + "model": None, + "litellm_call_id": "test-call-id", + "stream": False, + "litellm_params": { + "metadata": { + "user_api_key": "hashed_key", + "user_api_key_alias": None, + "user_api_key_user_id": None, + "user_api_key_team_id": None, + "user_api_key_org_id": None, + } + }, + } + + with ( + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=mock_key_obj, + ) as mock_get_key, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.get_team_object", + new_callable=AsyncMock, + return_value=MagicMock(team_alias=None), + ), + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + mock_proxy_logging.db_spend_update_writer = MagicMock() + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert mock_get_key.called is expect_key_read + + written = kwargs["litellm_params"]["metadata"] + if expect_key_read: + assert written["user_api_key_user_id"] == "user-assigned-later" + assert written["user_api_key_team_id"] == "team-assigned-later" + else: + assert written["user_api_key_user_id"] is None + assert written["user_api_key_team_id"] is None + assert written["user_api_key_org_id"] is None + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_enriches_auth_error_metadata(): """ diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py index 52da7a4a81d..d53e6dedf0b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py @@ -258,6 +258,7 @@ class TestVertexAIBatchPassthroughHandler: batch_object=batch_object, model_object_id=model_object_id, logging_obj=mock_logging_obj, + is_batch_create=True, user_api_key_dict={"user_id": "test-user"}, ) @@ -307,6 +308,7 @@ class TestVertexAIBatchPassthroughHandler: batch_object={"id": "b1", "object": "batch", "status": "validating"}, model_object_id="b1", logging_obj=mock_logging_obj, + is_batch_create=True, **kwargs, ) @@ -315,6 +317,140 @@ class TestVertexAIBatchPassthroughHandler: assert call_kwargs["user_api_key_dict"].user_id == expected_user_id assert call_kwargs["user_api_key_dict"].team_id == expected_team_id + def _store_with_metadata(self, mock_logging_obj, mock_managed_files_hook, metadata): + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_pl, + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.verbose_proxy_logger" + ), + ): + mock_pl.get_proxy_hook.return_value = mock_managed_files_hook + VertexPassthroughLoggingHandler._store_batch_managed_object( + unified_object_id="uoi", + batch_object={"id": "b1", "object": "batch", "status": "validating"}, + model_object_id="b1", + logging_obj=mock_logging_obj, + is_batch_create=True, + litellm_params={"metadata": metadata}, + ) + mock_managed_files_hook.store_unified_object_id.assert_called_once() + return mock_managed_files_hook.store_unified_object_id.call_args[1] + + def test_create_persists_key_hash_and_tags( + self, mock_logging_obj, mock_managed_files_hook + ): + """Regression (spend loss): the batch create must persist the creating key's hashed + token and its tags so CheckBatchCost can attribute the batch-cost spend row. Before + this fix the stored api_key was always "" and the row was dropped as unattributed.""" + call_kwargs = self._store_with_metadata( + mock_logging_obj, + mock_managed_files_hook, + { + "user_api_key": "hashed-key-a", + "user_api_key_user_id": "alice", + "user_api_key_team_id": "team-alpha", + "user_api_key_auth_metadata": {"tags": ["env:prod", 7, "team:ml"]}, + }, + ) + + assert call_kwargs["user_api_key_dict"].api_key == "hashed-key-a" + # non-string tags are dropped so downstream tag budgets cannot be bypassed + assert call_kwargs["request_tags"] == ("env:prod", "team:ml") + assert call_kwargs["persist_attribution"] is True + + @pytest.mark.parametrize( + "metadata, expected", + [ + # a request that sent its own tags (x-litellm-tags header or body metadata) + ({"tags": ["req:a", "req:b"]}, ("req:a", "req:b")), + # request tags win over the key's own tags + ( + {"tags": ["req:a"], "user_api_key_auth_metadata": {"tags": ["key:b"]}}, + ("req:a",), + ), + # no request tags: fall back to the tags the key itself carries + ({"user_api_key_auth_metadata": {"tags": ["key:b"]}}, ("key:b",)), + # neither: no tags on the spend row + ({}, None), + ], + ) + def test_request_tags_precedence( + self, mock_logging_obj, mock_managed_files_hook, metadata, expected + ): + """Request tags take precedence over the key's tags, and the key's tags are the + fallback because a tagged key does not put its tags in the top-level metadata.""" + call_kwargs = self._store_with_metadata( + mock_logging_obj, + mock_managed_files_hook, + {"user_api_key": "hashed-key-a", **metadata}, + ) + + assert call_kwargs["request_tags"] == expected + + @pytest.mark.parametrize( + "url_route, expected", + [ + ("/v1/projects/p/locations/us-central1/batchPredictionJobs", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs?alt=json", True), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/123456", False), + ("/v1/projects/p/locations/us-central1/batchPredictionJobs/123456?alt=json", False), + ], + ) + def test_batch_is_registered_from_the_create_route_only( + self, mock_logging_obj, url_route, expected + ): + """Only a POST to the collection route is the create, and only the create claims + attribution. Every id-scoped route is a poll or retrieve, which still reports the + batch so its status and file object stay in sync, but carries is_batch_create=False + so it neither claims the batch nor creates a row it would then own.""" + response = MagicMock() + response.status_code = 200 + response.json.return_value = { + "name": "projects/p/locations/us-central1/batchPredictionJobs/123456", + "model": "publishers/google/models/gemini-2.5-flash", + } + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.verbose_proxy_logger" + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.VertexPassthroughLoggingHandler._store_batch_managed_object" + ) as mock_store, + patch( + "litellm.llms.vertex_ai.batches.transformation.VertexAIBatchTransformation" + ) as mock_transformation, + patch( + "litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler.VertexPassthroughLoggingHandler.get_actual_model_id_from_router", + return_value="gemini-2.5-flash", + ), + ): + mock_transformation.transform_vertex_ai_batch_response_to_openai_batch_response.return_value = { + "id": "123456", + "object": "batch", + "status": "validating", + "created_at": 1704067200, + "input_file_id": "gs://bucket/in.jsonl", + "completion_window": "24h", + } + mock_transformation._get_batch_id_from_vertex_ai_batch_response.return_value = "123456" + + VertexPassthroughLoggingHandler.batch_prediction_jobs_handler( + httpx_response=response, + logging_obj=mock_logging_obj, + url_route=url_route, + result="", + start_time=datetime.now(), + end_time=datetime.now(), + cache_hit=False, + ) + + # every route reports the batch; only the create claims it + mock_store.assert_called_once() + assert mock_store.call_args[1]["unified_object_id"] + assert mock_store.call_args[1]["is_batch_create"] is expected + def test_batch_cost_calculation_integration(self): """Single Vertex AI response → non-zero cost with correct token counts.""" from litellm.batches.batch_utils import calculate_vertex_ai_batch_cost_and_usage From 2112422c713cf3e497da3351c8317992f77d975f Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 16:34:26 -0700 Subject: [PATCH 14/78] test(managed-files): read the scoped page id from the row's unified_file_id --- .../litellm_enterprise/proxy/hooks/test_managed_files.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 9e94b9f2a0f..d3efcf2e7a0 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -3119,8 +3119,9 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): ) managed_row = MagicMock() + managed_row.unified_file_id = "litellm_proxy:mine" managed_row.file_object = { - "id": "litellm_proxy:mine", + "id": "file-mine", "bytes": 100, "created_at": 1, "filename": "mine.jsonl", From 82662dc104db0b26e214c90346627801a7999da3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 8 Aug 2026 17:28:43 -0700 Subject: [PATCH 15/78] fix(proxy): report has_more false on caller-scoped file list pages --- enterprise/litellm_enterprise/proxy/hooks/managed_files.py | 6 ++++-- .../litellm_enterprise/proxy/hooks/test_managed_files.py | 3 ++- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a6ba1e0a791..37d267fcd6e 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1297,13 +1297,15 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """Rebuild ``first_id`` / ``last_id`` from the caller-scoped page. The upstream cursors point at rows that were just filtered out, so - leaving them in place discloses other callers' file ids. + leaving them in place discloses other callers' file ids. ``has_more`` + is always cleared because ``after`` is never forwarded upstream, so + no further page is reachable through the proxy. """ if hasattr(response, "first_id"): response.first_id = data[0].id if data else None if hasattr(response, "last_id"): response.last_id = data[-1].id if data else None - if not data and hasattr(response, "has_more"): + if hasattr(response, "has_more"): response.has_more = False async def afile_retrieve( diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index d3efcf2e7a0..fde1feb80e2 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -3112,7 +3112,7 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): upstream_page = AsyncCursorPage[FileObject].construct( data=[_raw_file("file-someone-else"), _raw_file("file-mine")], - has_more=False, + has_more=True, first_id="file-someone-else", last_id="file-mine", object="list", @@ -3146,3 +3146,4 @@ async def test_file_list_cursors_follow_the_owner_scoped_page(): assert [file_object.id for file_object in response.data] == ["litellm_proxy:mine"] assert response.first_id == "litellm_proxy:mine" assert response.last_id == "litellm_proxy:mine" + assert response.has_more is False From 30c4898de9cea90e777a43a5260fe77011acdb5b Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:16:22 -0700 Subject: [PATCH 16/78] fix(ui): hide admin-only Logs tabs from roles that cannot call their endpoints The Logs nav entry is open to internal users so they can read their own request logs, but the page rendered all four tabs unconditionally. Audit Logs calls GET /audit and Deleted Teams calls GET /v2/team/list?status=deleted, neither of which an internal user is permitted to call, so the page fired requests that came back 401. Gate both tabs on new viewAuditLogs / viewDeletedTeams capabilities, using the same CAPABILITY_ROLES map and useCan hook introduced for Tool Policies. Hiding a tab drops its panel from the tree entirely, so the request is never issued rather than issued and rejected. Selecting a tab also mapped index 0 to "request logs" and every other index to "audit logs", which activated the audit panel whenever a user opened Deleted Keys or Deleted Teams. Derive the active tab from the visible tab list instead, so the mapping survives tabs being filtered out. --- .../view_logs/index.integration.test.tsx | 104 ++++++++++++++++++ .../src/components/view_logs/index.test.tsx | 79 ++++++++++++- .../src/components/view_logs/index.tsx | 91 +++++++++------ .../src/utils/capabilities.test.ts | 13 +++ .../src/utils/capabilities.ts | 2 + 5 files changed, 254 insertions(+), 35 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx new file mode 100644 index 00000000000..b86ad015b91 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/index.integration.test.tsx @@ -0,0 +1,104 @@ +import { screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import SpendLogsTable from "./index"; +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; + +const { useAuthorizedMock } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn() })); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: useAuthorizedMock, +})); + +vi.mock("./RequestLogsPanel", () => ({ + default: function RequestLogsPanelMock() { + return
; + }, +})); + +const fetchMock = vi.fn(); + +const jsonResponse = (body: unknown) => ({ + ok: true, + status: 200, + statusText: "OK", + json: async () => body, +}); + +const requestedUrls = () => fetchMock.mock.calls.map(([url]) => String(url)); + +const emptyAuditLogs = { audit_logs: [], total: 0, page: 1, page_size: 50, total_pages: 0 }; + +const defaultProps = { + accessToken: "sk-test", + token: "jwt-test", + userRole: "Admin", + userID: "user-1", + premiumUser: true, +}; + +const renderAs = (sessionRole: string) => { + useAuthorizedMock.mockReturnValue({ accessToken: "sk-test", userRole: sessionRole, premiumUser: true }); + return renderWithProviders(); +}; + +describe("SpendLogsTable network access by role", () => { + beforeEach(() => { + testQueryClient.clear(); + vi.clearAllMocks(); + fetchMock.mockImplementation(async (url: string) => { + if (String(url).includes("/audit")) { + return jsonResponse(emptyAuditLogs); + } + if (String(url).includes("/v2/team/list")) { + return jsonResponse({ teams: [] }); + } + return jsonResponse({ keys: [], total_count: 0 }); + }); + vi.stubGlobal("fetch", fetchMock); + }); + + it("fires neither the audit nor the deleted-teams request for an internal user", async () => { + const user = userEvent.setup(); + renderAs("Internal User"); + + // Liveness gate: the sibling Deleted Keys panel does reach the network, so a + // silent absence below means the gate worked, not that nothing rendered. + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/key/list"))).toBe(true)); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + await user.click(screen.getByRole("tab", { name: "Request Logs" })); + + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + expect(requestedUrls().filter((url) => url.includes("/v2/team/list"))).toEqual([]); + }); + + it("fetches deleted teams and audit logs for an admin", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await waitFor(() => + expect(requestedUrls().some((url) => url.includes("/v2/team/list") && url.includes("status=deleted"))).toBe(true), + ); + + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + + await user.click(screen.getByRole("tab", { name: "Audit Logs" })); + + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/audit"))).toBe(true)); + }); + + it("leaves the audit request unsent when an admin selects a tab after Audit Logs", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Teams" })); + + expect(screen.getByRole("tab", { name: "Deleted Teams" })).toHaveAttribute("aria-selected", "true"); + expect(requestedUrls().filter((url) => url.includes("/audit"))).toEqual([]); + + await user.click(screen.getByRole("tab", { name: "Audit Logs" })); + + await waitFor(() => expect(requestedUrls().some((url) => url.includes("/audit"))).toBe(true)); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index b2e77ec7fd5..785fa0cc6f8 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,9 +1,15 @@ import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import SpendLogsTable from "./index"; import { renderWithProviders } from "../../../tests/test-utils"; +const { useAuthorizedMock } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn() })); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: useAuthorizedMock, +})); + vi.mock("./RequestLogsPanel", () => ({ default: function RequestLogsPanelMock({ isActive }: { isActive: boolean }) { return
{isActive ? "active" : "inactive"}
; @@ -36,9 +42,18 @@ const defaultProps = { premiumUser: false, }; +const renderAs = (sessionRole: string) => { + useAuthorizedMock.mockReturnValue({ userRole: sessionRole }); + return renderWithProviders(); +}; + describe("SpendLogsTable", () => { + beforeEach(() => { + useAuthorizedMock.mockReturnValue({ userRole: "Admin" }); + }); + it("renders the four log tabs", () => { - renderWithProviders(); + renderAs("Admin"); for (const label of ["Request Logs", "Audit Logs", "Deleted Keys", "Deleted Teams"]) { expect(screen.getByRole("tab", { name: label })).toBeInTheDocument(); @@ -47,7 +62,7 @@ describe("SpendLogsTable", () => { it("marks only the visible tab's panel active so background tabs do not query", async () => { const user = userEvent.setup(); - renderWithProviders(); + renderAs("Admin"); expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("active"); @@ -57,8 +72,64 @@ describe("SpendLogsTable", () => { expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); }); + describe("admin-only tabs", () => { + it.each(["Internal User", "Internal Viewer"])("hides Audit Logs and Deleted Teams from %s", (role) => { + renderAs(role); + + expect(screen.getByRole("tab", { name: "Request Logs" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Deleted Keys" })).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Audit Logs" })).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Deleted Teams" })).not.toBeInTheDocument(); + }); + + it("never mounts the panels that call the admin-only endpoints for an internal user", () => { + renderAs("Internal User"); + + expect(screen.queryByTestId("audit-logs-panel")).not.toBeInTheDocument(); + expect(screen.queryByTestId("deleted-teams-page")).not.toBeInTheDocument(); + expect(screen.getByTestId("deleted-keys-page")).toBeInTheDocument(); + }); + }); + + describe("tab index mapping", () => { + it("activates the panel the admin selected, not the one at the old hardcoded index", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + + expect(screen.getByTestId("audit-logs-panel")).toHaveTextContent("inactive"); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); + }); + + it("keeps the audit panel inert when an admin selects the last tab", async () => { + const user = userEvent.setup(); + renderAs("Admin"); + + await user.click(screen.getByRole("tab", { name: "Deleted Teams" })); + + expect(screen.getByTestId("audit-logs-panel")).toHaveTextContent("inactive"); + expect(screen.getByTestId("deleted-teams-page")).toBeInTheDocument(); + }); + + it("selects the last visible tab for an internal user and returns to Request Logs", async () => { + const user = userEvent.setup(); + renderAs("Internal User"); + + await user.click(screen.getByRole("tab", { name: "Deleted Keys" })); + + expect(screen.getByTestId("deleted-keys-page")).toBeInTheDocument(); + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("inactive"); + + await user.click(screen.getByRole("tab", { name: "Request Logs" })); + + expect(screen.getByTestId("request-logs-panel")).toHaveTextContent("active"); + }); + }); + describe("auth-not-ready guard", () => { it("shows a loading spinner when credentials are not yet resolved", () => { + useAuthorizedMock.mockReturnValue({ userRole: "Admin" }); renderWithProviders(); expect(document.querySelector(".ant-spin")).toBeInTheDocument(); @@ -66,7 +137,7 @@ describe("SpendLogsTable", () => { }); it("renders the tabs (no spinner) once all credentials are present", () => { - renderWithProviders(); + renderAs("Admin"); expect(document.querySelector(".ant-spin")).not.toBeInTheDocument(); expect(screen.getByRole("tab", { name: "Request Logs" })).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index 8e7423e3fae..7269564dcec 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -1,5 +1,6 @@ import { useState } from "react"; import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage"; import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage"; import AuditLogsPanel from "./AuditLogsPanel"; @@ -14,8 +15,22 @@ interface SpendLogsTableProps { premiumUser: boolean; } +type LogsTabId = "request logs" | "audit logs" | "deleted keys" | "deleted teams"; + +interface LogsTab { + id: LogsTabId; + label: string; +} + +const REQUEST_LOGS_TAB: LogsTab = { id: "request logs", label: "Request Logs" }; +const AUDIT_LOGS_TAB: LogsTab = { id: "audit logs", label: "Audit Logs" }; +const DELETED_KEYS_TAB: LogsTab = { id: "deleted keys", label: "Deleted Keys" }; +const DELETED_TEAMS_TAB: LogsTab = { id: "deleted teams", label: "Deleted Teams" }; + export default function SpendLogsTable({ accessToken, token, userRole, userID, premiumUser }: SpendLogsTableProps) { - const [activeTab, setActiveTab] = useState("request logs"); + const [activeTab, setActiveTab] = useState(REQUEST_LOGS_TAB.id); + const canViewAuditLogs = useCan("viewAuditLogs"); + const canViewDeletedTeams = useCan("viewDeletedTeams"); if (!accessToken || !token || !userRole || !userID) { return ( @@ -25,41 +40,55 @@ export default function SpendLogsTable({ accessToken, token, userRole, userID, p ); } + const tabs: LogsTab[] = [ + REQUEST_LOGS_TAB, + ...(canViewAuditLogs ? [AUDIT_LOGS_TAB] : []), + DELETED_KEYS_TAB, + ...(canViewDeletedTeams ? [DELETED_TEAMS_TAB] : []), + ]; + + const renderPanel = (tabId: LogsTabId) => { + switch (tabId) { + case "request logs": + return ( + + ); + case "audit logs": + return ( + + ); + case "deleted keys": + return ; + case "deleted teams": + return ; + } + }; + return (
- setActiveTab(index === 0 ? "request logs" : "audit logs")}> + setActiveTab(tabs[index].id)}> - Request Logs - Audit Logs - Deleted Keys - Deleted Teams + {tabs.map((tab) => ( + {tab.label} + ))} - - - - - - - - - - - - + {tabs.map((tab) => ( + {renderPanel(tab.id)} + ))}
diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index f48609b0b9d..611c9626065 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -18,6 +18,19 @@ describe("hasCapability", () => { ); }); +describe.each(["viewAuditLogs", "viewDeletedTeams"] as const)("hasCapability - %s", (capability) => { + it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(true); + }); + + it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])( + "should deny it to %s", + (role) => { + expect(hasCapability(role, capability)).toBe(false); + }, + ); +}); + describe("rolesWithCapability", () => { it("should return a copy so callers cannot mutate the capability map", () => { const roles = rolesWithCapability("viewToolPolicies"); diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index 77ead2568fb..f0847cc3400 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -2,6 +2,8 @@ import { all_admin_roles } from "./roles"; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, + viewAuditLogs: all_admin_roles, + viewDeletedTeams: all_admin_roles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; From 6a540a1bf848129dc16228d1b23f92120d7a7f03 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:17:12 -0700 Subject: [PATCH 17/78] fix(ui): gate organization and agent usage views behind capabilities The Usage page admits internal users because their own usage view works, but the entity breakdown selector inside it also offered Organization Usage, so picking it fired /organization/daily/activity and collected a 401. Neither that route nor /agent/daily/activity appears in any non-admin route list, so both are default-deny. The team breakdown leaked the second one too: it fetches agent activity unconditionally to fill its Top Agents card, which 401s for the same roles. Adds viewOrganizationUsage and viewAgentUsage to the existing capability map and points the selector option, the page section, and the fetch's enabled flag at the same capability, so a role that cannot call the endpoint never sees the breakdown and never issues the request. The team and tag breakdowns, which internal users can read, are untouched, and the default Usage view was already one of those. --- .../EntityUsage/EntityUsage.test.tsx | 43 ++++++++++++++++++- .../components/EntityUsage/EntityUsage.tsx | 34 +++++++++++---- .../components/UsagePageView.test.tsx | 25 +++++++++++ .../_components/components/UsagePageView.tsx | 9 ++-- .../UsageViewSelect/UsageViewSelect.test.tsx | 29 +++++++++++-- .../UsageViewSelect/UsageViewSelect.tsx | 20 +++++---- .../src/utils/capabilities.test.ts | 36 ++++++++++------ .../src/utils/capabilities.ts | 2 + 8 files changed, 160 insertions(+), 38 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 82ca66b10c0..11528117f1e 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, within } from "@testing-library/react"; +import { act, cleanup, 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"; @@ -856,6 +856,47 @@ describe("EntityUsage", () => { expect(logo.getAttribute("src")).toContain("openai_small"); }); + describe("capability gating", () => { + it.each([ + ["organization", () => mockOrganizationDailyActivityCall, "Organization Spend Overview"], + ["agent", () => mockAgentDailyActivityCall, "Agent Spend Overview"], + ] as const)("fetches %s activity for an admin but not for an internal user", async (entityType, call, heading) => { + render(); + await waitFor(() => { + expect(call()).toHaveBeenCalled(); + }); + + cleanup(); + call().mockClear(); + + render(); + expect(await screen.findByText(heading)).toBeInTheDocument(); + expect(call()).not.toHaveBeenCalled(); + }); + + it("keeps the team breakdown but drops its agent sub-fetch for an internal user", async () => { + render(); + + await waitFor(() => { + expect(mockTeamDailyActivityCall).toHaveBeenCalled(); + }); + expect(screen.getByText("Team Spend Overview")).toBeInTheDocument(); + + expect(mockAgentDailyActivityCall).not.toHaveBeenCalled(); + expect(screen.queryByText("Agent Activity")).not.toBeInTheDocument(); + expect(screen.queryByText("Top Agents Driving Spend")).not.toBeInTheDocument(); + }); + + it("keeps the tag breakdown for an internal user", async () => { + render(); + + await waitFor(() => { + expect(mockTagDailyActivityCall).toHaveBeenCalled(); + }); + expect(screen.getByText("Tag Spend Overview")).toBeInTheDocument(); + }); + }); + it("renders a letter avatar instead of an img for an unknown provider slug", async () => { const spendDataUnknownProvider = { ...mockSpendData, 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 4d44791d1a9..5a0f2abf15b 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 @@ -2,6 +2,7 @@ import useTeams from "@/app/(dashboard)/hooks/useTeams"; import { BarChart, DonutChart } from "@/components/shared/charts"; import { MoneyCell } from "@/components/shared/table_cells"; import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { hasCapability, type Capability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { Card, @@ -108,7 +109,19 @@ const ENTITY_FETCH_FNS: Record Promise> = { user: userDailyActivityCall, }; -const EntityUsage: React.FC = ({ accessToken, entityType, entityId, entityList, dateValue }) => { +const ENTITY_CAPABILITIES: Partial> = { + organization: "viewOrganizationUsage", + agent: "viewAgentUsage", +}; + +const EntityUsage: React.FC = ({ + accessToken, + entityType, + entityId, + entityList, + userRole, + dateValue, +}) => { const { teams } = useTeams(); const [selectedTags, setSelectedTags] = useState([]); const [modelViewType, setModelViewType] = useState("groups"); @@ -125,7 +138,11 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti }, [entityType, selectedTags]); const fetchFn = ENTITY_FETCH_FNS[entityType]; - const enabled = !!accessToken && !!startTime && !!endTime; + const entityCapability = ENTITY_CAPABILITIES[entityType]; + const canViewEntity = entityCapability === undefined || hasCapability(userRole, entityCapability); + const showAgentBreakdown = entityType === "team" && hasCapability(userRole, "viewAgentUsage"); + const hasRequestWindow = !!accessToken && !!startTime && !!endTime; + const enabled = hasRequestWindow && canViewEntity; const { data: spendDataRaw, @@ -150,7 +167,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti } = usePaginatedDailyActivity({ fetchFn: agentDailyActivityCall, args: [accessToken, startTime, endTime, null], - enabled: enabled && entityType === "team", + enabled: enabled && showAgentBreakdown, }); const agentSpendData = agentSpendDataRaw as unknown as EntitySpendData; @@ -158,7 +175,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti 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 agentMetrics = showAgentBreakdown ? processActivityData(agentSpendData, "entities", teams || []) : {}; const getTopModels = () => { const modelSpend: { [key: string]: any } = {}; @@ -621,8 +638,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti - {/* Top Agents - only for team entity type */} - {entityType === "team" && ( + {showAgentBreakdown && ( Top Agents Driving Spend @@ -708,7 +724,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti ), }, - ...(entityType === "team" + ...(showAgentBreakdown ? [{ key: "agents", label: "Agent Activity", content: }] : []), { @@ -757,7 +773,7 @@ const EntityUsage: React.FC = ({ accessToken, entityType, enti } /> )} - {agentIsFetchingMore && entityType === "team" && ( + {agentIsFetchingMore && showAgentBreakdown && ( = ({ accessToken, entityType, enti } /> )} - {agentCancelled && entityType === "team" && ( + {agentCancelled && showAgentBreakdown && ( { userId: "user-123", userEmail: "test@example.com", userRole: "Internal User", + userRoleLabel: "Internal User", + isViewOnly: false, premiumUser: true, disabledPersonalKeyCreation: false, showSSOBanner: false, @@ -861,6 +863,29 @@ describe("UsagePage", () => { }); }); + // The select hides both views from a non-admin, so this drives the section + // gate directly through the mocked select, which always offers every option. + it.each(["organization", "agent"])("should not render the %s usage view for an internal user", async (usageView) => { + mockUseAuthorized.mockReturnValue(nonAdminSession); + + renderWithProviders(); + + await waitFor(() => { + expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalled(); + }); + + const usageSelect = screen.getByTestId("usage-view-select"); + act(() => { + fireEvent.change(usageSelect, { target: { value: "team" } }); + }); + expect(screen.getAllByText("Entity Usage").length).toBeGreaterThan(0); + + act(() => { + fireEvent.change(usageSelect, { target: { value: usageView } }); + }); + expect(screen.queryByText("Entity Usage")).not.toBeInTheDocument(); + }); + describe("admin user selector", () => { it("should render user selector for admin users in global view", async () => { renderWithProviders(); 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 c3645d6371e..494df313ac0 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 @@ -33,6 +33,7 @@ import { useCustomers } from "@/app/(dashboard)/hooks/customers/useCustomers"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useCurrentUser } from "@/app/(dashboard)/hooks/users/useCurrentUser"; import { useInfiniteUsers } from "@/app/(dashboard)/hooks/users/useUsers"; +import { hasCapability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { all_admin_roles, internalUserRoles } from "@/utils/roles"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; @@ -109,6 +110,8 @@ const UsagePage: React.FC = ({ teams, organizations }) => { const { data: currentUser } = useCurrentUser(); const isAdmin = all_admin_roles.includes(userRole || ""); const canViewTagUsage = isAdmin || internalUserRoles.includes(userRole || ""); + const canViewOrganizationUsage = hasCapability(userRole, "viewOrganizationUsage"); + const canViewAgentUsage = hasCapability(userRole, "viewAgentUsage"); // Debounced search for user selector const [userSearchInput, setUserSearchInput] = useState(""); @@ -513,7 +516,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { setUsageView(value)} - isAdmin={isAdmin} + userRole={userRole} canViewTagUsage={canViewTagUsage} /> @@ -950,7 +953,7 @@ const UsagePage: React.FC = ({ teams, organizations }) => { )} {/* Organization Usage Panel */} - {usageView === "organization" && ( + {usageView === "organization" && canViewOrganizationUsage && ( = ({ teams, organizations }) => { /> )} - {usageView === "agent" && ( + {usageView === "agent" && canViewAgentUsage && ( { }); it("should render", () => { - render(); + render(); expect(screen.getByText("Usage View")).toBeInTheDocument(); expect(screen.getByText("Select the usage data you want to view")).toBeInTheDocument(); expect(screen.getByRole("combobox")).toBeInTheDocument(); + expect(screen.getByRole("option", { name: "Your Usage" })).toBeInTheDocument(); }); it("should call onChange when value changes", () => { - render(); + render(); const select = screen.getByRole("combobox"); act(() => { @@ -109,14 +110,34 @@ describe("UsageViewSelect", () => { }); it("should show Tag Usage for non-admin users with tag usage permission", () => { - render(); + render(); expect(screen.getByRole("option", { name: "Tag Usage" })).toBeInTheDocument(); }); it("should hide Tag Usage for non-admin users without tag usage permission", () => { - render(); + render(); expect(screen.queryByRole("option", { name: "Tag Usage" })).not.toBeInTheDocument(); }); + + it.each(["Organization Usage", "Agent Usage (A2A)"])("should show %s to an admin", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); + + // Neither /organization/daily/activity nor /agent/daily/activity admits an + // internal user, so the option that fires them must not be selectable. + it.each(["Organization Usage", "Agent Usage (A2A)"])("should hide %s from an internal user", (optionName) => { + render(); + + expect(screen.queryByRole("option", { name: optionName })).not.toBeInTheDocument(); + }); + + it.each(["Team Usage", "Tag Usage"])("should keep %s available to an internal user", (optionName) => { + render(); + + expect(screen.getByRole("option", { name: optionName })).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx index 94b483cb539..54c1d5ab7cc 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx @@ -11,6 +11,8 @@ import { } from "@ant-design/icons"; import { Badge, Select } from "antd"; import React from "react"; +import { hasCapability, type Capability } from "@/utils/capabilities"; +import { all_admin_roles } from "@/utils/roles"; export type UsageOption = | "global" | "my-usage" @@ -24,7 +26,7 @@ export type UsageOption = export interface UsageViewSelectProps { value: UsageOption; onChange: (value: UsageOption) => void; - isAdmin: boolean; + userRole: string | null; canViewTagUsage?: boolean; title?: string; description?: string; @@ -35,6 +37,7 @@ interface OptionConfig { label: string; description: string; icon: React.ReactNode; + capability?: Capability; adminOnly?: boolean; showForAdmin?: string; showForNonAdmin?: string; @@ -63,12 +66,9 @@ const OPTIONS: OptionConfig[] = [ { value: "organization", label: "Organization Usage", - showForAdmin: "Organization Usage", - showForNonAdmin: "Your Organization Usage", - description: "View organization-level usage", - descriptionForAdmin: "View usage across all organizations", - descriptionForNonAdmin: "View your organization's usage", + description: "View usage across all organizations", icon: , + capability: "viewOrganizationUsage", }, { value: "team", @@ -95,7 +95,7 @@ const OPTIONS: OptionConfig[] = [ label: "Agent Usage (A2A)", description: "View usage by AI agents", icon: , - adminOnly: true, + capability: "viewAgentUsage", }, { value: "user", @@ -115,14 +115,18 @@ const OPTIONS: OptionConfig[] = [ export const UsageViewSelect: React.FC = ({ value, onChange, - isAdmin, + userRole, canViewTagUsage = false, title = "Usage View", description = "Select the usage data you want to view", "data-id": dataId, }) => { + const isAdmin = all_admin_roles.includes(userRole ?? ""); const getFilteredOptions = () => { return OPTIONS.filter((option) => { + if (option.capability) { + return hasCapability(userRole, option.capability); + } if (option.value === "tag" && canViewTagUsage) { return true; } diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index f48609b0b9d..3f5a8ac81fb 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -1,21 +1,31 @@ import { describe, expect, it } from "vitest"; -import { hasCapability, rolesWithCapability } from "./capabilities"; +import { hasCapability, rolesWithCapability, type Capability } from "./capabilities"; + +const ADMIN_ROLES = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"]; +const NON_ADMIN_ROLES = [ + "Internal User", + "Internal Viewer", + "App User", + "Org Admin", + "Unknown Role", + "", + null, + undefined, +]; + +const ADMIN_ONLY_CAPABILITIES: Capability[] = ["viewToolPolicies", "viewOrganizationUsage", "viewAgentUsage"]; describe("hasCapability", () => { - it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])( - "should grant viewToolPolicies to %s", - (role) => { - expect(hasCapability(role, "viewToolPolicies")).toBe(true); - }, - ); + describe.each(ADMIN_ONLY_CAPABILITIES)("%s", (capability) => { + it.each(ADMIN_ROLES)("should grant it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(true); + }); - it.each(["Internal User", "Internal Viewer", "App User", "Org Admin", "Unknown Role", "", null, undefined])( - "should deny viewToolPolicies to %s", - (role) => { - expect(hasCapability(role, "viewToolPolicies")).toBe(false); - }, - ); + it.each(NON_ADMIN_ROLES)("should deny it to %s", (role) => { + expect(hasCapability(role, capability)).toBe(false); + }); + }); }); describe("rolesWithCapability", () => { diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index 77ead2568fb..c4d878b81a5 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -2,6 +2,8 @@ import { all_admin_roles } from "./roles"; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, + viewOrganizationUsage: all_admin_roles, + viewAgentUsage: all_admin_roles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; From 2502ee4a2ace88ac10dd8ed30e2a03fff075ce26 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 8 Aug 2026 20:34:56 -0700 Subject: [PATCH 18/78] fix(ui): gate policy and prompt lookups on an admin capability /policies/list and /prompts/list are default-deny for internal_user, but the Virtual Keys create/edit flow, the Teams forms and the Playground called them on mount, so every internal user landing on the dashboard fired two requests that 401. Add viewPolicies and viewPrompts to the capability map and use them to gate the nav entry, the form field and the fetch together, following the pattern from the Tool Policies migration. Non-admins now see no policy or prompt selector at all rather than an empty dropdown. --- .../playground/components/chat_ui/ChatUI.tsx | 56 +++---- .../components/complianceUI/ComplianceUI.tsx | 48 +++--- .../src/components/Teams.test.tsx | 55 ++++++- ui/litellm-dashboard/src/components/Teams.tsx | 68 ++++---- .../src/components/leftnav.test.tsx | 22 +++ .../src/components/leftnav.tsx | 10 +- .../organisms/create_key_button.test.tsx | 49 +++++- .../organisms/create_key_button.tsx | 145 +++++++++--------- .../policies/PolicySelector.test.tsx | 20 +++ .../components/policies/PolicySelector.tsx | 10 +- .../src/components/team/TeamInfo.test.tsx | 46 ++++++ .../src/components/team/TeamInfo.tsx | 56 +++---- .../templates/key_edit_view.test.tsx | 50 +++++- .../components/templates/key_edit_view.tsx | 87 ++++++----- .../src/utils/capabilities.test.ts | 28 ++++ .../src/utils/capabilities.ts | 2 + 16 files changed, 533 insertions(+), 219 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx index 57ff7906eda..0241ef8a77e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx @@ -25,6 +25,7 @@ import React, { useEffect, useRef, useState } from "react"; import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; import { v4 as uuidv4 } from "uuid"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import GuardrailSelector from "@/components/guardrails/GuardrailSelector"; import PolicySelector from "@/components/policies/PolicySelector"; import MCPToolArgumentsForm, { MCPToolArgumentsFormRef } from "@/components/mcp_tools/MCPToolArgumentsForm"; @@ -106,6 +107,7 @@ const ChatUI: React.FC = ({ simplified = false, fixedModel, }) => { + const canViewPolicies = useCan("viewPolicies"); const [mcpServers, setMCPServers] = useState([]); const [mcpToolsets, setMCPToolsets] = useState([]); const [isToolsetsInfoModalVisible, setIsToolsetsInfoModalVisible] = useState(false); @@ -1652,32 +1654,34 @@ const ChatUI: React.FC = ({ />
-
- - Policies - - Select policy/policies to apply to this LLM API call. Policies define which guardrails are - applied based on conditions. You can set up your policies{" "} - - here - - . - - } - > - - - - -
+ {canViewPolicies && ( +
+ + Policies + + Select policy/policies to apply to this LLM API call. Policies define which guardrails are + applied based on conditions. You can set up your policies{" "} + + here + + . + + } + > + + + + +
+ )} {/* Code Interpreter Toggle - Only for Responses endpoint */} {endpointType === EndpointType.RESPONSES && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx index 39346105f2a..c3b417987e6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx @@ -6,6 +6,7 @@ import { type ComplianceFramework, type CompliancePrompt, } from "@/data/compliancePrompts"; +import useCan from "@/app/(dashboard)/hooks/useCan"; import { getGuardrailsList, testPoliciesAndGuardrails } from "@/components/networking"; import PolicySelector, { getPolicyOptionEntries } from "@/components/policies/PolicySelector"; import { Policy } from "@/components/policies/types"; @@ -123,6 +124,7 @@ export default function ComplianceUI({ fixedModel, proxySettings, }: ComplianceUIProps) { + const canViewPolicies = useCan("viewPolicies"); const frameworks = getFrameworks(); const [policyValueToLabel, setPolicyValueToLabel] = useState>(new Map()); @@ -701,29 +703,37 @@ export default function ComplianceUI({

Test Configuration

-

Select policies, guardrails, or both to test against.

+

+ {canViewPolicies + ? "Select policies, guardrails, or both to test against." + : "Select guardrails to test against."} +

-
- - {accessToken && ( - - )} -
+ {canViewPolicies && ( + <> +
+ + {accessToken && ( + + )} +
-
-
- or -
-
+
+
+ or +
+
+ + )}