diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 37d267fcd6e..f1b4c6b5b17 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -54,6 +54,9 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( normalize_mime_type_for_provider, resolve_managed_output_file_model_name, ) +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import ( + request_tags_from_metadata, +) from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue] AllMessageValues, AsyncCursorPage, @@ -1146,6 +1149,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ## Check if unified_file_id is in the response unified_file_id = response._hidden_params.get("unified_file_id") # managed file id unified_batch_id = response._hidden_params.get("unified_batch_id") # managed batch id + is_batch_create: Final = unified_file_id is not None model_id = cast(Optional[str], response._hidden_params.get("model_id")) model_name = cast(Optional[str], response._hidden_params.get("model_name")) @@ -1216,6 +1220,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): model_mappings={model_id: provider_file_id}, user_api_key_dict=user_api_key_dict, ) + request_metadata: Final = data.get("litellm_metadata") await self.store_unified_object_id( unified_object_id=response.id, file_object=response, @@ -1223,6 +1228,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): model_object_id=original_response_id, file_purpose="batch", user_api_key_dict=user_api_key_dict, + request_tags=request_tags_from_metadata(request_metadata if isinstance(request_metadata, dict) else {}), + persist_attribution=is_batch_create, ) # Only record batch creation metric on actual create (not retrieve/cancel). 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 d260f79a09a..6cc31f991a3 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -828,3 +828,49 @@ async def test_cost_job_and_retrieve_paths_mint_identical_unified_output_file_id model_id="model-deploy-xyz", model_name=cost_job_model_name, ) + + +@pytest.mark.asyncio +async def test_batch_create_hook_persists_creating_key_and_tags(): + """Regression: the /v1/batches create hook must persist the creating key and the + request's tags on the managed object row. CheckBatchCost, which owns the batch's + accounting once the retrieve path defers to it, bills whatever the row carries, and + without these columns the cost lands on the user alone and the key's spend and + budget never see it.""" + managed_files = _make_managed_files_instance() + creator = UserAPIKeyAuth(api_key="sk-the-creator", user_id="alice", parent_otel_span=None) + create_response = _make_batch_response(status="validating", output_file_id=None) + + await managed_files.async_post_call_success_hook( + data={"litellm_metadata": {"tags": ["env:prod", "team:ml"], "user_api_key": creator.api_key}}, + user_api_key_dict=creator, + response=create_response, + ) + + managed_files.store_unified_object_id.assert_awaited_once() + stored = managed_files.store_unified_object_id.await_args.kwargs + assert stored["persist_attribution"] is True + assert stored["request_tags"] == ("env:prod", "team:ml") + assert stored["user_api_key_dict"] is creator + + +@pytest.mark.asyncio +async def test_batch_retrieve_hook_does_not_claim_attribution(): + """A retrieve carries unified_batch_id but no unified_file_id, so it must not rewrite + the row's paying key to whoever happens to poll the batch.""" + managed_files = _make_managed_files_instance() + retrieve_response = _make_batch_response(status="in_progress", output_file_id=None) + retrieve_response._hidden_params = { + "unified_batch_id": "some-unified-batch-id", + "model_id": "model-deploy-xyz", + "model_name": "azure/gpt-4", + } + + await managed_files.async_post_call_success_hook( + data={"litellm_metadata": {"tags": ["poller:tag"]}}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-the-poller", user_id="bob", parent_otel_span=None), + response=retrieve_response, + ) + + managed_files.store_unified_object_id.assert_awaited_once() + assert managed_files.store_unified_object_id.await_args.kwargs["persist_attribution"] is False