mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-22 00:31:44 +00:00
feat(proxy): price Azure Speech short audio pass-through from the recognized duration
Short audio responses carry Offset and Duration in 100ns ticks; convert their sum to seconds and price it with the existing azure/speech/azure-stt entry through transcription_cost. Batch calls and responses without an integer duration stay at zero cost Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2e8dc0a627
commit
f2305879d0
5 changed files with 171 additions and 24 deletions
|
|
@ -1579,6 +1579,8 @@ AZURE_SPEECH_COGNITIVE_SERVICES_DOMAIN: Final = "api.cognitive.microsoft.com"
|
|||
AZURE_SPEECH_SUBSCRIPTION_KEY_HEADER: Final = "Ocp-Apim-Subscription-Key"
|
||||
AZURE_SPEECH_SHORT_AUDIO_MODEL: Final = "short-audio"
|
||||
AZURE_SPEECH_BATCH_MODEL: Final = "batch-transcription"
|
||||
AZURE_SPEECH_PRICING_MODEL: Final = "azure/speech/azure-stt"
|
||||
AZURE_SPEECH_TICKS_PER_SECOND: Final = 10_000_000
|
||||
|
||||
BASE_MCP_ROUTE: Final = "/mcp"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from urllib.parse import urlparse
|
||||
|
|
@ -9,9 +9,12 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.constants import (
|
||||
AZURE_SPEECH_BATCH_MODEL,
|
||||
AZURE_SPEECH_CUSTOM_LLM_PROVIDER,
|
||||
AZURE_SPEECH_PRICING_MODEL,
|
||||
AZURE_SPEECH_SHORT_AUDIO_MODEL,
|
||||
AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX,
|
||||
AZURE_SPEECH_TICKS_PER_SECOND,
|
||||
)
|
||||
from litellm.cost_calculator import transcription_cost
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
|
|
@ -21,16 +24,50 @@ from litellm.types.utils import StandardPassThroughResponseObject
|
|||
|
||||
|
||||
class AzureSpeechPassthroughLoggingHandler:
|
||||
@staticmethod
|
||||
def _is_short_audio_route(url_route: str) -> bool:
|
||||
return urlparse(url_route).path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX)
|
||||
|
||||
@staticmethod
|
||||
def _model_from_url_route(url_route: str) -> str:
|
||||
path: Final = urlparse(url_route).path
|
||||
if path.startswith(AZURE_SPEECH_SHORT_AUDIO_PATH_PREFIX):
|
||||
if AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route):
|
||||
return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_SHORT_AUDIO_MODEL}"
|
||||
return f"{AZURE_SPEECH_CUSTOM_LLM_PROVIDER}/{AZURE_SPEECH_BATCH_MODEL}"
|
||||
|
||||
@staticmethod
|
||||
def _recognized_audio_seconds(response_body: Mapping[str, object] | Sequence[object] | None) -> float:
|
||||
if not isinstance(response_body, Mapping):
|
||||
return 0.0
|
||||
offset: Final = response_body.get("Offset")
|
||||
duration: Final = response_body.get("Duration")
|
||||
if not isinstance(offset, int) or not isinstance(duration, int):
|
||||
return 0.0
|
||||
return (offset + duration) / AZURE_SPEECH_TICKS_PER_SECOND
|
||||
|
||||
@staticmethod
|
||||
def _response_cost(url_route: str, response_body: Mapping[str, object] | Sequence[object] | None) -> float:
|
||||
if not AzureSpeechPassthroughLoggingHandler._is_short_audio_route(url_route):
|
||||
return 0.0
|
||||
audio_seconds: Final = AzureSpeechPassthroughLoggingHandler._recognized_audio_seconds(response_body)
|
||||
if audio_seconds <= 0.0:
|
||||
return 0.0
|
||||
try:
|
||||
prompt_cost, completion_cost = transcription_cost(
|
||||
model=AZURE_SPEECH_PRICING_MODEL,
|
||||
custom_llm_provider="azure",
|
||||
duration=audio_seconds,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a missing price entry must not drop the spend log row
|
||||
verbose_proxy_logger.warning(
|
||||
"No price for %s, logging Azure Speech call at zero cost: %s", AZURE_SPEECH_PRICING_MODEL, e
|
||||
)
|
||||
return 0.0
|
||||
return prompt_cost + completion_cost
|
||||
|
||||
@staticmethod
|
||||
def azure_speech_passthrough_handler(
|
||||
httpx_response: httpx.Response,
|
||||
response_body: Mapping[str, object] | Sequence[object] | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
url_route: str,
|
||||
result: str,
|
||||
|
|
@ -40,25 +77,20 @@ class AzureSpeechPassthroughLoggingHandler:
|
|||
request_body: Mapping[str, object],
|
||||
**kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
"""
|
||||
Records model and provider for an Azure AI Speech REST call. Azure bills per audio
|
||||
hour after the fact and neither the short-audio response nor the batch job carries
|
||||
a billable duration this path can trust, so response_cost is recorded as 0.0 rather
|
||||
than estimated.
|
||||
"""
|
||||
try:
|
||||
model_name: Final = AzureSpeechPassthroughLoggingHandler._model_from_url_route(url_route)
|
||||
response_cost: Final = AzureSpeechPassthroughLoggingHandler._response_cost(url_route, response_body)
|
||||
|
||||
updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict
|
||||
**kwargs,
|
||||
"model": model_name,
|
||||
"custom_llm_provider": AZURE_SPEECH_CUSTOM_LLM_PROVIDER,
|
||||
"response_cost": 0.0,
|
||||
"response_cost": response_cost,
|
||||
}
|
||||
logging_obj.model_call_details.update(
|
||||
model=model_name,
|
||||
custom_llm_provider=AZURE_SPEECH_CUSTOM_LLM_PROVIDER,
|
||||
response_cost=0.0,
|
||||
response_cost=response_cost,
|
||||
)
|
||||
|
||||
standard_logging_object: Final = get_standard_logging_object_payload(
|
||||
|
|
|
|||
|
|
@ -264,6 +264,7 @@ class PassThroughEndpointLogging:
|
|||
|
||||
azure_speech_handler_result: Final = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler(
|
||||
httpx_response=httpx_response,
|
||||
response_body=response_body,
|
||||
logging_obj=logging_obj,
|
||||
url_route=url_route,
|
||||
result=result,
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.azure_speech_passthrough_logging_handler import (
|
||||
AzureSpeechPassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -13,7 +15,24 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
|
|||
|
||||
SHORT_AUDIO_URL = "https://eastus.stt.speech.microsoft.com/speech/recognition/conversation/cognitiveservices/v1?language=en-US"
|
||||
BATCH_URL = "https://eastus.api.cognitive.microsoft.com/speechtotext/v3.2/transcriptions"
|
||||
TRANSCRIPT = '{"RecognitionStatus":"Success","DisplayText":"Hello world."}'
|
||||
TRANSCRIPT_BODY = {"RecognitionStatus": "Success", "Offset": 5000000, "Duration": 25000000, "DisplayText": "Hello world."}
|
||||
TRANSCRIPT = json.dumps(TRANSCRIPT_BODY)
|
||||
TRANSCRIPT_AUDIO_SECONDS = 3.0
|
||||
PRICE_PER_SECOND = 0.5
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def azure_stt_price(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"azure/speech/azure-stt",
|
||||
{
|
||||
"litellm_provider": "azure",
|
||||
"mode": "audio_transcription",
|
||||
"input_cost_per_second": PRICE_PER_SECOND,
|
||||
"output_cost_per_second": 0.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _make_response(url: str) -> httpx.Response:
|
||||
|
|
@ -30,18 +49,19 @@ def _make_logging_obj() -> MagicMock:
|
|||
|
||||
class TestAzureSpeechPassthroughHandler:
|
||||
@pytest.mark.parametrize(
|
||||
"url_route,expected_model",
|
||||
"url_route,expected_model,expected_cost",
|
||||
[
|
||||
(SHORT_AUDIO_URL, "azure_speech/short-audio"),
|
||||
(BATCH_URL, "azure_speech/batch-transcription"),
|
||||
(f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription"),
|
||||
(SHORT_AUDIO_URL, "azure_speech/short-audio", TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND),
|
||||
(BATCH_URL, "azure_speech/batch-transcription", 0.0),
|
||||
(f"{BATCH_URL}/8a5d3f2c-0b1e-4c7d-9e6f-1234567890ab/files", "azure_speech/batch-transcription", 0.0),
|
||||
],
|
||||
)
|
||||
def test_records_model_provider_and_zero_cost(self, url_route: str, expected_model: str):
|
||||
def test_records_model_provider_and_cost(self, url_route: str, expected_model: str, expected_cost: float):
|
||||
logging_obj = _make_logging_obj()
|
||||
|
||||
handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler(
|
||||
httpx_response=_make_response(url_route),
|
||||
response_body=TRANSCRIPT_BODY,
|
||||
logging_obj=logging_obj,
|
||||
url_route=url_route,
|
||||
result=TRANSCRIPT,
|
||||
|
|
@ -54,16 +74,65 @@ class TestAzureSpeechPassthroughHandler:
|
|||
assert handler_result["result"] == {"response": TRANSCRIPT}
|
||||
assert handler_result["kwargs"]["model"] == expected_model
|
||||
assert handler_result["kwargs"]["custom_llm_provider"] == "azure_speech"
|
||||
assert handler_result["kwargs"]["response_cost"] == 0.0
|
||||
assert handler_result["kwargs"]["standard_logging_object"]["response_cost"] == 0.0
|
||||
assert handler_result["kwargs"]["response_cost"] == pytest.approx(expected_cost)
|
||||
assert handler_result["kwargs"]["standard_logging_object"]["response_cost"] == pytest.approx(expected_cost)
|
||||
assert handler_result["kwargs"]["standard_logging_object"]["model"] == expected_model
|
||||
assert logging_obj.model_call_details["model"] == expected_model
|
||||
assert logging_obj.model_call_details["custom_llm_provider"] == "azure_speech"
|
||||
assert logging_obj.model_call_details["response_cost"] == 0.0
|
||||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(expected_cost)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"response_body",
|
||||
[
|
||||
{"RecognitionStatus": "NoMatch", "Offset": 0, "Duration": 0},
|
||||
{"RecognitionStatus": "InitialSilenceTimeout"},
|
||||
{"Offset": "5000000", "Duration": "25000000"},
|
||||
{},
|
||||
[],
|
||||
None,
|
||||
],
|
||||
)
|
||||
def test_short_audio_without_recognized_duration_logs_zero_cost(
|
||||
self, response_body: dict[str, object] | list[dict[str, object]] | None
|
||||
):
|
||||
handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler(
|
||||
httpx_response=_make_response(SHORT_AUDIO_URL),
|
||||
response_body=response_body,
|
||||
logging_obj=_make_logging_obj(),
|
||||
url_route=SHORT_AUDIO_URL,
|
||||
result="",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={},
|
||||
)
|
||||
|
||||
assert handler_result["kwargs"]["model"] == "azure_speech/short-audio"
|
||||
assert handler_result["kwargs"]["response_cost"] == 0.0
|
||||
|
||||
def test_missing_price_entry_still_logs_the_row_at_zero_cost(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delitem(litellm.model_cost, "azure/speech/azure-stt")
|
||||
|
||||
handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler(
|
||||
httpx_response=_make_response(SHORT_AUDIO_URL),
|
||||
response_body=TRANSCRIPT_BODY,
|
||||
logging_obj=_make_logging_obj(),
|
||||
url_route=SHORT_AUDIO_URL,
|
||||
result=TRANSCRIPT,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={},
|
||||
)
|
||||
|
||||
assert handler_result["kwargs"]["model"] == "azure_speech/short-audio"
|
||||
assert handler_result["kwargs"]["custom_llm_provider"] == "azure_speech"
|
||||
assert handler_result["kwargs"]["response_cost"] == 0.0
|
||||
|
||||
def test_subscription_key_never_reaches_the_logging_payload(self):
|
||||
handler_result = AzureSpeechPassthroughLoggingHandler.azure_speech_passthrough_handler(
|
||||
httpx_response=_make_response(SHORT_AUDIO_URL),
|
||||
response_body=TRANSCRIPT_BODY,
|
||||
logging_obj=_make_logging_obj(),
|
||||
url_route=SHORT_AUDIO_URL,
|
||||
result=TRANSCRIPT,
|
||||
|
|
@ -106,18 +175,18 @@ class TestNormalizeDispatch:
|
|||
def test_normalize_routes_to_azure_speech_handler(self):
|
||||
normalized = PassThroughEndpointLogging().normalize_llm_passthrough_logging_payload(
|
||||
httpx_response=_make_response(SHORT_AUDIO_URL),
|
||||
response_body={"RecognitionStatus": "Success"},
|
||||
response_body=TRANSCRIPT_BODY,
|
||||
request_body={},
|
||||
logging_obj=_make_logging_obj(),
|
||||
url_route=SHORT_AUDIO_URL,
|
||||
result=TRANSCRIPT,
|
||||
result="",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
custom_llm_provider="azure_speech",
|
||||
)
|
||||
|
||||
assert normalized["standard_logging_response_object"] == {"response": TRANSCRIPT}
|
||||
assert normalized["standard_logging_response_object"] == {"response": ""}
|
||||
assert normalized["kwargs"]["model"] == "azure_speech/short-audio"
|
||||
assert normalized["kwargs"]["custom_llm_provider"] == "azure_speech"
|
||||
assert normalized["kwargs"]["response_cost"] == 0.0
|
||||
assert normalized["kwargs"]["response_cost"] == pytest.approx(TRANSCRIPT_AUDIO_SECONDS * PRICE_PER_SECOND)
|
||||
|
|
|
|||
|
|
@ -6278,6 +6278,49 @@ class TestAzureSpeechProxyRoute:
|
|||
("azure_speech/batch-transcription", "azure_speech", 0.0)
|
||||
]
|
||||
|
||||
def test_short_audio_spend_is_priced_from_the_recognized_duration(
|
||||
self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class _Recorder(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.payloads: list[dict[str, object]] = [] # mutable-ok: test recorder accumulates callback payloads
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
|
||||
self.payloads.append(kwargs["standard_logging_object"])
|
||||
|
||||
recorder: Final = _Recorder()
|
||||
monkeypatch.setattr(litellm, "_async_success_callback", [*litellm._async_success_callback, recorder])
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"azure/speech/azure-stt",
|
||||
{
|
||||
"litellm_provider": "azure",
|
||||
"mode": "audio_transcription",
|
||||
"input_cost_per_second": 0.25,
|
||||
"output_cost_per_second": 0.0,
|
||||
},
|
||||
)
|
||||
transcript: Final = {**AZURE_SPEECH_TRANSCRIPT, "Offset": 10_000_000, "Duration": 30_000_000}
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
upstream.post(f"https://eastus.stt.speech.microsoft.com{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}").mock(
|
||||
return_value=httpx.Response(200, json=transcript)
|
||||
)
|
||||
|
||||
response = azure_speech_client.post(
|
||||
f"/azure_speech{AZURE_SPEECH_SHORT_AUDIO_ENDPOINT}",
|
||||
content=AZURE_SPEECH_WAV_BYTES,
|
||||
headers={"Content-Type": "audio/wav", "Authorization": "Bearer sk-virtual"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [(p["model"], p["custom_llm_provider"]) for p in recorder.payloads] == [
|
||||
("azure_speech/short-audio", "azure_speech")
|
||||
]
|
||||
assert recorder.payloads[0]["response_cost"] == pytest.approx(4.0 * 0.25)
|
||||
|
||||
def test_api_base_wins_over_region_for_both_families(
|
||||
self, azure_speech_client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue