diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 349fa873792..3f6cce87876 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -48,7 +48,7 @@ class CheckResponsesCost: async def _get_response( self, response_id: str, - litellm_metadata: dict[str, str], + litellm_metadata: dict[str, object], ) -> ResponsesAPIResponse: """Fetch the upstream response through the deployment that served it. @@ -214,11 +214,17 @@ class CheckResponsesCost: # Decrypts rows written before model_object_id held the provider's own id. responses_id_security = ResponsesIDSecurity().provider_response_id(job.model_object_id) - probe_metadata: dict[str, str] = { + probe_metadata: dict[str, object] = { "user_api_key_user_id": job.created_by or "default-user-id", **({"user_api_key_team_id": job.team_id} if job.team_id else {}), **({"user_api_key": job.api_key, "user_api_key_hash": job.api_key} if job.api_key else {}), **({"model": model_name, "model_group": model_name} if model_name else {}), + **({"user_api_key_org_id": job.org_id} if job.org_id else {}), + **( + {"tags": [tag for tag in job.request_tags if isinstance(tag, str)]} + if isinstance(job.request_tags, list) and job.request_tags + else {} + ), } except Exception as e: @@ -263,7 +269,7 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Response {unified_object_id} has terminal status {response.status}, marked as complete" ) - billing_metadata: dict[str, str] = { + billing_metadata: dict[str, object] = { **probe_metadata, INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN, } diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 415d630e015..02024dc14f4 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,7 +1,7 @@ import asyncio import json import time -from collections.abc import AsyncIterator, Awaitable, Mapping +from collections.abc import AsyncIterator, Awaitable, Mapping, Sequence from enum import Enum from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Protocol, cast, get_args @@ -30,6 +30,9 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_set_request_parsed_body, ) +from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import ( + request_tags_from_metadata, +) from litellm.types.llms.openai import ( REASONING_EFFORT, ResponsesAPIOptionalRequestParams, @@ -58,6 +61,7 @@ class BackgroundResponseStore(Protocol): model_object_id: str, file_purpose: Literal["response"], user_api_key_dict: UserAPIKeyAuth, + request_tags: Sequence[str] | None = None, persist_attribution: bool = False, ) -> None: ... @@ -80,6 +84,7 @@ async def store_background_response_object( response: ResponsesAPIResponse, managed_files_obj: BackgroundResponseStore, user_api_key_dict: UserAPIKeyAuth, + data: Mapping[str, object], ) -> None: """Record a queued background response so the cost poller can find and bill it. @@ -98,6 +103,7 @@ async def store_background_response_object( return provider_response_id: Final = ResponsesIDSecurity().provider_response_id(response.id) + litellm_metadata: Final = data.get("litellm_metadata") await managed_files_obj.store_unified_object_id( unified_object_id=response.id, file_object=response, @@ -105,6 +111,7 @@ async def store_background_response_object( model_object_id=provider_response_id, file_purpose="response", user_api_key_dict=user_api_key_dict, + request_tags=request_tags_from_metadata(litellm_metadata if isinstance(litellm_metadata, dict) else {}), persist_attribution=True, ) verbose_proxy_logger.info("Stored background response %s in managed objects table", response.id) @@ -444,6 +451,7 @@ async def responses_api( response=response, managed_files_obj=managed_files_obj, user_api_key_dict=user_api_key_dict, + data=data, ) except Exception as e: verbose_proxy_logger.error("Failed to store background response in managed objects table: %s", e) diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py index e6475c8b452..e6136f5491d 100644 --- a/tests/proxy_unit_tests/test_check_responses_cost.py +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -1049,6 +1049,73 @@ class TestCheckResponsesCost: } assert _release_calls(mock_prisma_client) == [] + @pytest.mark.asyncio + async def test_billing_read_includes_managed_row_attribution( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY + from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + + attributed_job = MagicMock() + attributed_job.unified_object_id = "resp_attributed" + attributed_job.model_object_id = _routed_response_id("resp_attributed") + attributed_job.created_by = "test-user" + attributed_job.org_id = "org-1" + attributed_job.request_tags = ["tag-a", "tag-b"] + attributed_job.id = "job-attributed" + attributed_job.file_object = {"model": "gpt-5", "id": "resp_attributed"} + + unattributed_job = MagicMock() + unattributed_job.unified_object_id = "resp_unattributed" + unattributed_job.model_object_id = _routed_response_id("resp_unattributed") + unattributed_job.created_by = "test-user" + unattributed_job.org_id = None + unattributed_job.request_tags = [] + unattributed_job.id = "job-unattributed" + unattributed_job.file_object = {"model": "gpt-5", "id": "resp_unattributed"} + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[attributed_job, unattributed_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_llm_router.aget_responses = AsyncMock( + side_effect=[ + ResponsesAPIResponse( + id="resp_attributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + + [ + ResponsesAPIResponse( + id="resp_unattributed", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + ] + * 2 + ) + + await check_responses_cost_instance.check_responses_cost() + + billing_metadata = [ + call.kwargs["litellm_metadata"] + for call in mock_llm_router.aget_responses.await_args_list + if call.kwargs["litellm_metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) + == BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN + ] + assert billing_metadata[0]["user_api_key_org_id"] == "org-1" + assert billing_metadata[0]["tags"] == ["tag-a", "tag-b"] + assert "user_api_key_org_id" not in billing_metadata[1] + assert "tags" not in billing_metadata[1] + @pytest.mark.asyncio async def test_non_terminal_probe_does_not_finalize_or_bill( self, check_responses_cost_instance, mock_prisma_client, mock_llm_router diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 907feb971dd..aa5287cf7ab 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -1997,7 +1997,12 @@ class TestBackgroundResponseManagedObjectId: response._hidden_params = {"model_id": model_id} if model_id else {} return response - async def _stored_kwargs(self, advertised_id: str, model_id: str | None = "deployment-1"): + async def _stored_kwargs( + self, + advertised_id: str, + model_id: str | None = "deployment-1", + data: dict[str, object] | None = None, + ): from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.response_api_endpoints.endpoints import ( store_background_response_object, @@ -2010,6 +2015,7 @@ class TestBackgroundResponseManagedObjectId: response=self._queued_response(advertised_id, model_id), managed_files_obj=managed_files_obj, user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"), + data=data if data is not None else {"background": True}, ) return managed_files_obj.store_unified_object_id @@ -2070,6 +2076,25 @@ class TestBackgroundResponseManagedObjectId: store.assert_not_awaited() + @pytest.mark.asyncio + async def test_request_tags_are_forwarded_from_litellm_metadata(self, monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt") + + tagged_store = await self._stored_kwargs( + self._encrypted_id("resp_tagged"), + data={ + "background": True, + "litellm_metadata": {"tags": ["tag-a", "tag-b"]}, + }, + ) + untagged_store = await self._stored_kwargs( + self._encrypted_id("resp_untagged"), + data={"background": True}, + ) + + assert tagged_store.await_args.kwargs["request_tags"] == ("tag-a", "tag-b") + assert untagged_store.await_args.kwargs["request_tags"] is None + class TestShouldStoreBackgroundResponse: """The gate `responses_api` applies before it writes a managed row.