mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #38414 from BerriAI/litellm_fix_speech_metadata_spend_tracking
fix(speech): keep proxy metadata and completion cost through the TTS completion bridge
This commit is contained in:
commit
d77eef3d11
8 changed files with 176 additions and 10 deletions
|
|
@ -99,7 +99,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44530
|
||||
"limit": 44528
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
|
|
@ -138,9 +138,9 @@
|
|||
"limit": 139
|
||||
},
|
||||
"reportUnusedImport": {
|
||||
"limit": 545
|
||||
"limit": 544
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 146
|
||||
"limit": 145
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,6 +8,14 @@ if TYPE_CHECKING:
|
|||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
def _completion_response_cost(model_response: "ModelResponse") -> float | None:
|
||||
hidden_params: Final = getattr(model_response, "_hidden_params", None)
|
||||
if not isinstance(hidden_params, dict):
|
||||
return None
|
||||
response_cost: Final = hidden_params.get("response_cost")
|
||||
return response_cost if isinstance(response_cost, float) else None
|
||||
|
||||
|
||||
class SpeechToCompletionBridgeTransformationHandler:
|
||||
def transform_request(
|
||||
self,
|
||||
|
|
@ -123,4 +131,6 @@ class SpeechToCompletionBridgeTransformationHandler:
|
|||
|
||||
# Create an httpx.Response object
|
||||
response: Final = httpx.Response(status_code=200, content=binary_data, headers=headers)
|
||||
return HttpxBinaryResponseContent(response)
|
||||
binary_response: Final = HttpxBinaryResponseContent(response)
|
||||
binary_response.set_response_cost(_completion_response_cost(model_response))
|
||||
return binary_response
|
||||
|
|
|
|||
|
|
@ -1590,7 +1590,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if transformed_result is not None:
|
||||
result = transformed_result
|
||||
|
||||
if isinstance(result, BaseModel) and hasattr(result, "_hidden_params"):
|
||||
if isinstance(result, (BaseModel, HttpxBinaryResponseContent)) and hasattr(result, "_hidden_params"):
|
||||
hidden_params: Final = getattr(result, "_hidden_params", {})
|
||||
if (
|
||||
"response_cost" in hidden_params and hidden_params["response_cost"] is not None
|
||||
|
|
|
|||
|
|
@ -8013,7 +8013,7 @@ def speech(
|
|||
|
||||
if max_retries is None:
|
||||
max_retries = litellm.num_retries or openai.DEFAULT_MAX_RETRIES
|
||||
litellm_params_dict: Final = get_litellm_params(**kwargs)
|
||||
litellm_params_dict: Final = get_litellm_params(metadata=metadata, api_key=api_key or dynamic_api_key, **kwargs)
|
||||
|
||||
# Get provider-specific text-to-speech config and map parameters
|
||||
text_to_speech_provider_config = ProviderConfigManager.get_provider_text_to_speech_config(
|
||||
|
|
|
|||
|
|
@ -107,7 +107,17 @@ EmbeddingInput = str | list[str]
|
|||
|
||||
|
||||
class HttpxBinaryResponseContent(_HttpxBinaryResponseContent):
|
||||
_hidden_params: dict = {}
|
||||
_hidden_params: dict
|
||||
|
||||
def __init__(self, response: httpx.Response) -> None:
|
||||
super().__init__(response)
|
||||
self._hidden_params = {} # mutable-ok: mutable-dict contract shared with ModelResponse logging consumers
|
||||
|
||||
def set_response_cost(self, response_cost: float | None) -> None:
|
||||
if response_cost is None:
|
||||
self._hidden_params.pop("response_cost", None)
|
||||
return
|
||||
self._hidden_params["response_cost"] = response_cost
|
||||
|
||||
|
||||
class NotGiven:
|
||||
|
|
|
|||
|
|
@ -108,7 +108,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"F401": {
|
||||
"limit": 14
|
||||
"limit": 13
|
||||
},
|
||||
"LOG015": {
|
||||
"limit": 5
|
||||
|
|
@ -171,7 +171,7 @@
|
|||
"limit": 175
|
||||
},
|
||||
"RUF012": {
|
||||
"limit": 240
|
||||
"limit": 239
|
||||
},
|
||||
"RUF015": {
|
||||
"limit": 8
|
||||
|
|
@ -183,7 +183,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"RUF059": {
|
||||
"limit": 67
|
||||
"limit": 66
|
||||
},
|
||||
"RUF100": {
|
||||
"limit": 0
|
||||
|
|
|
|||
|
|
@ -1,7 +1,12 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import contextlib
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -14,6 +19,9 @@ from unittest.mock import MagicMock, patch
|
|||
|
||||
import litellm
|
||||
from litellm import main as litellm_main
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
async def _async_fake_bedrock_image_details(image_url):
|
||||
|
|
@ -2957,3 +2965,109 @@ async def test_acompletion_resolves_provider_from_api_base():
|
|||
)
|
||||
|
||||
assert response.choices[0].message.content == "resolved"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RecordedSpeechSuccess:
|
||||
call_type: str | None
|
||||
spend_metadata: Mapping[str, object]
|
||||
response_cost: float | None
|
||||
logged_response_cost: float | None
|
||||
|
||||
|
||||
def _record_speech_success(payload: dict[str, object]) -> _RecordedSpeechSuccess:
|
||||
call_type: Final = payload.get("call_type")
|
||||
response_cost: Final = payload.get("response_cost")
|
||||
logging_payload: Final = payload.get("standard_logging_object")
|
||||
logged_cost: Final = logging_payload.get("response_cost") if isinstance(logging_payload, dict) else None
|
||||
return _RecordedSpeechSuccess(
|
||||
call_type=call_type if isinstance(call_type, str) else None,
|
||||
spend_metadata=get_litellm_metadata_from_kwargs(payload),
|
||||
response_cost=response_cost if isinstance(response_cost, float) else None,
|
||||
logged_response_cost=logged_cost if isinstance(logged_cost, float) else None,
|
||||
)
|
||||
|
||||
|
||||
class _SuccessEventRecorder(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.events: list[_RecordedSpeechSuccess] = [] # mutable-ok: test recorder of success-callback events
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self.events.append(_record_speech_success(kwargs))
|
||||
|
||||
|
||||
async def _wait_for_success_event(recorder: _SuccessEventRecorder, call_type: str) -> _RecordedSpeechSuccess:
|
||||
for _ in range(100):
|
||||
if (event := next((e for e in recorder.events if e.call_type == call_type), None)) is not None:
|
||||
return event
|
||||
await asyncio.sleep(0.05)
|
||||
pytest.fail(f"no {call_type} success event; got {[e.call_type for e in recorder.events]}")
|
||||
|
||||
|
||||
def _gemini_tts_generate_content_response() -> dict[str, object]:
|
||||
return {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {
|
||||
"parts": [
|
||||
{
|
||||
"inlineData": {
|
||||
"mimeType": "audio/L16;codec=pcm;rate=24000",
|
||||
"data": base64.b64encode(b"pcm-audio-bytes").decode(),
|
||||
}
|
||||
}
|
||||
],
|
||||
"role": "model",
|
||||
},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 5,
|
||||
"candidatesTokenCount": 60,
|
||||
"totalTokenCount": 65,
|
||||
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 5}],
|
||||
"candidatesTokensDetails": [{"modality": "AUDIO", "tokenCount": 60}],
|
||||
},
|
||||
"modelVersion": "gemini-2.5-flash-preview-tts",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aspeech_gemini_bridge_keeps_proxy_metadata_for_spend_tracking(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("GEMINI_API_KEY", raising=False)
|
||||
monkeypatch.delenv("GOOGLE_API_KEY", raising=False)
|
||||
recorder: Final = _SuccessEventRecorder()
|
||||
monkeypatch.setattr(litellm, "callbacks", [recorder])
|
||||
mock_route: Final = respx_mock.post(
|
||||
url__regex=r"https://generativelanguage\.googleapis\.com/v1beta/models/gemini-2\.5-flash-preview-tts:generateContent.*"
|
||||
).mock(return_value=httpx.Response(200, json=_gemini_tts_generate_content_response()))
|
||||
|
||||
await litellm.aspeech(
|
||||
model="gemini/gemini-2.5-flash-preview-tts",
|
||||
input="spend tracking check",
|
||||
voice="Kore",
|
||||
api_key="fake-gemini-key",
|
||||
metadata={"user_api_key": "hashed-virtual-key", "user_api_key_user_id": "user-1"},
|
||||
)
|
||||
|
||||
assert mock_route.called
|
||||
assert mock_route.calls.last.request.headers["x-goog-api-key"] == "fake-gemini-key"
|
||||
speech_event: Final = await _wait_for_success_event(recorder, call_type="aspeech")
|
||||
assert speech_event.spend_metadata["user_api_key"] == "hashed-virtual-key"
|
||||
assert speech_event.spend_metadata["user_api_key_user_id"] == "user-1"
|
||||
expected_prompt_cost, expected_completion_cost = litellm.cost_per_token(
|
||||
model="gemini/gemini-2.5-flash-preview-tts",
|
||||
usage_object=Usage(prompt_tokens=5, completion_tokens=60, total_tokens=65),
|
||||
)
|
||||
expected_cost: Final = expected_prompt_cost + expected_completion_cost
|
||||
assert expected_cost > 0
|
||||
assert speech_event.response_cost == pytest.approx(expected_cost)
|
||||
assert speech_event.logged_response_cost == pytest.approx(expected_cost)
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ import pytest
|
|||
import json
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import HttpxBinaryResponseContent
|
||||
|
||||
|
||||
def test_generic_event():
|
||||
|
|
@ -522,3 +523,34 @@ class TestOpenAIFileObjectBatchGuardrailSerialization:
|
|||
|
||||
page = FileListPage(object="list", data=[self._file_object()], has_more=False)
|
||||
assert "litellm_batch_guardrail" not in page.model_dump(mode="json")["data"][0]
|
||||
|
||||
|
||||
def _binary_content(payload: bytes) -> HttpxBinaryResponseContent:
|
||||
import httpx
|
||||
|
||||
return HttpxBinaryResponseContent(httpx.Response(200, content=payload))
|
||||
|
||||
|
||||
def test_httpx_binary_response_content_hidden_params_are_per_instance():
|
||||
first = _binary_content(b"first")
|
||||
second = _binary_content(b"second")
|
||||
|
||||
first._hidden_params["response_cost"] = 0.5
|
||||
|
||||
assert second._hidden_params == {}
|
||||
|
||||
|
||||
def test_set_response_cost_none_leaves_hidden_params_empty():
|
||||
binary_response = _binary_content(b"audio")
|
||||
|
||||
binary_response.set_response_cost(None)
|
||||
|
||||
assert "response_cost" not in binary_response._hidden_params
|
||||
|
||||
binary_response.set_response_cost(0.25)
|
||||
|
||||
assert binary_response._hidden_params["response_cost"] == 0.25
|
||||
|
||||
binary_response.set_response_cost(None)
|
||||
|
||||
assert "response_cost" not in binary_response._hidden_params
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue