fix(responses): bill background responses against the creating org and tags
Some checks failed
Terraform Provider / gofmt, vet, build, test (push) Has been cancelled
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
jesus 2026-09-17 22:30:47 +00:00
parent a537523b5e
commit 350ed95411
4 changed files with 111 additions and 5 deletions

View file

@ -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,
}

View file

@ -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)

View file

@ -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

View file

@ -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.