mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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
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:
parent
a537523b5e
commit
350ed95411
4 changed files with 111 additions and 5 deletions
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue