mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(xai): add speech-to-text via /v1/audio/transcriptions
Route xai audio transcription through a provider config hitting POST https://api.x.ai/v1/stt instead of the openai-compatible chat handler which targets /audio/transcriptions. Supports language, diarize, keyterm, filler_words and other provider fields as passthrough kwargs Resolves LIT-8153 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
a6e3a72ed8
commit
6b0ad3bed3
8 changed files with 434 additions and 1 deletions
|
|
@ -998,6 +998,11 @@ openai_compatible_providers: Final[list] = [
|
|||
"cognition",
|
||||
"scx-ai",
|
||||
]
|
||||
|
||||
# Providers that are openai-compatible for chat but have their own audio
|
||||
# transcription endpoint, so litellm.transcription must route them through
|
||||
# their provider config instead of the OpenAI SDK handler.
|
||||
OPENAI_COMPATIBLE_PROVIDERS_WITH_NATIVE_AUDIO_TRANSCRIPTION: Final = frozenset({"xai"})
|
||||
openai_text_completion_compatible_providers: Final[list] = [ # providers that support `/v1/completions`
|
||||
"together_ai",
|
||||
"fireworks_ai",
|
||||
|
|
|
|||
3
litellm/llms/xai/audio_transcription/__init__.py
Normal file
3
litellm/llms/xai/audio_transcription/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import XAIAudioTranscriptionConfig
|
||||
|
||||
__all__ = ["XAIAudioTranscriptionConfig"]
|
||||
187
litellm/llms/xai/audio_transcription/transformation.py
Normal file
187
litellm/llms/xai/audio_transcription/transformation.py
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
"""
|
||||
Translates from OpenAI's `/v1/audio/transcriptions` to xAI's `/v1/stt`
|
||||
"""
|
||||
|
||||
from collections.abc import Iterable, Mapping
|
||||
from typing import Final, cast
|
||||
|
||||
from httpx import Headers, Response
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIAudioTranscriptionOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
from ...base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from ..common_utils import XAIModelInfo
|
||||
|
||||
|
||||
class XAIAudioTranscriptionError(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
class _XAISttWord(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
text: str = ""
|
||||
start: float = 0.0
|
||||
end: float = 0.0
|
||||
speaker: str | None = None
|
||||
|
||||
|
||||
class _XAISttResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
text: str = ""
|
||||
language: str = "unknown"
|
||||
duration: float | None = None
|
||||
words: list[_XAISttWord] | None = None
|
||||
|
||||
|
||||
def _serialize_form_value(value: object) -> str | list[str]:
|
||||
if isinstance(value, bool):
|
||||
return "true" if value else "false"
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [str(item) for item in cast(Iterable[object], value)]
|
||||
return str(value)
|
||||
|
||||
|
||||
class XAIAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> str:
|
||||
return litellm.LlmProviders.XAI.value
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list[OpenAIAudioTranscriptionOptionalParams]:
|
||||
return ["language"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict[str, object],
|
||||
optional_params: dict[str, object],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, object]:
|
||||
supported_params: Final = self.get_supported_openai_params(model)
|
||||
for k, v in non_default_params.items():
|
||||
if k in supported_params:
|
||||
optional_params[k] = v
|
||||
return optional_params
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: dict[str, object] | Headers
|
||||
) -> BaseLLMException:
|
||||
return XAIAudioTranscriptionError(message=error_message, status_code=status_code, headers=headers)
|
||||
|
||||
def transform_audio_transcription_request(
|
||||
self,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
) -> AudioTranscriptionRequestData:
|
||||
processed_audio: Final = process_audio_file(audio_file)
|
||||
|
||||
# Provider kwargs land in `extra_body` for openai_compatible_providers
|
||||
extra_body: Final = optional_params.get("extra_body")
|
||||
flat_params: Final[dict[str, object]] = {
|
||||
**(dict(cast(Mapping[str, object], extra_body)) if isinstance(extra_body, Mapping) else {}),
|
||||
**{k: v for k, v in optional_params.items() if k != "extra_body"},
|
||||
}
|
||||
|
||||
openai_params: Final = self.get_supported_openai_params(model)
|
||||
excluded_params: Final = frozenset({"model", "OPENAI_TRANSCRIPTION_PARAMS", *openai_params})
|
||||
provider_specific_params: Final[dict[str, object]] = {
|
||||
k: v for k, v in flat_params.items() if v is not None and k not in excluded_params
|
||||
}
|
||||
|
||||
form_data: Final[dict[str, str | list[str]]] = {"model": model}
|
||||
for key, value in provider_specific_params.items():
|
||||
form_data[key] = _serialize_form_value(value)
|
||||
for key in openai_params:
|
||||
value = flat_params.get(key)
|
||||
if value is not None:
|
||||
form_data[key] = _serialize_form_value(value)
|
||||
|
||||
files: Final = {
|
||||
"file": (
|
||||
processed_audio.filename,
|
||||
processed_audio.file_content,
|
||||
processed_audio.content_type,
|
||||
)
|
||||
}
|
||||
|
||||
return AudioTranscriptionRequestData(data=form_data, files=files)
|
||||
|
||||
def transform_audio_transcription_response(
|
||||
self,
|
||||
raw_response: Response,
|
||||
) -> TranscriptionResponse:
|
||||
try:
|
||||
payload: Final = _XAISttResponse.model_validate_json(raw_response.content)
|
||||
except Exception as e:
|
||||
raise XAIAudioTranscriptionError(
|
||||
message=f"Error parsing xAI response: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=dict(raw_response.headers),
|
||||
)
|
||||
|
||||
response: Final = TranscriptionResponse(text=payload.text)
|
||||
response["task"] = "transcribe"
|
||||
response["language"] = payload.language
|
||||
|
||||
if payload.duration is not None:
|
||||
response["duration"] = payload.duration
|
||||
|
||||
if payload.words is not None:
|
||||
response["words"] = [
|
||||
{
|
||||
"word": word.text,
|
||||
"start": word.start,
|
||||
"end": word.end,
|
||||
**({"speaker": word.speaker} if word.speaker is not None else {}),
|
||||
}
|
||||
for word in payload.words
|
||||
]
|
||||
|
||||
hidden_params: Final[dict[str, object]] = dict(payload.model_dump(mode="json"))
|
||||
if payload.duration is not None:
|
||||
hidden_params["audio_transcription_duration"] = payload.duration
|
||||
response._hidden_params = hidden_params # pyright: ignore[reportPrivateUsage] # TranscriptionResponse exposes no public hidden-params setter
|
||||
|
||||
return response
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
base: Final = (XAIModelInfo.get_api_base(api_base) or "").rstrip("/")
|
||||
normalized: Final = base.removesuffix("/v1")
|
||||
return f"{normalized}/v1/stt"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, object],
|
||||
model: str,
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict[str, object],
|
||||
litellm_params: dict[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict[str, object]:
|
||||
resolved_key: Final = XAIModelInfo.get_api_key(api_key)
|
||||
if resolved_key is None:
|
||||
raise ValueError("xAI API key is required. Set XAI_API_KEY environment variable.")
|
||||
|
||||
headers["Authorization"] = f"Bearer {resolved_key}"
|
||||
return headers
|
||||
|
|
@ -64,6 +64,7 @@ from litellm.constants import (
|
|||
AZURE_OPENAI_AUDIO_PROVIDERS,
|
||||
DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT,
|
||||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
OPENAI_COMPATIBLE_PROVIDERS_WITH_NATIVE_AUDIO_TRANSCRIPTION,
|
||||
)
|
||||
from litellm.exceptions import LiteLLMUnknownProvider
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -7859,7 +7860,10 @@ def transcription(
|
|||
litellm_params=litellm_params_dict,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
elif custom_llm_provider == "openai" or (custom_llm_provider in litellm.openai_compatible_providers):
|
||||
elif custom_llm_provider == "openai" or (
|
||||
custom_llm_provider in litellm.openai_compatible_providers
|
||||
and custom_llm_provider not in OPENAI_COMPATIBLE_PROVIDERS_WITH_NATIVE_AUDIO_TRANSCRIPTION
|
||||
):
|
||||
api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
|
|
|
|||
|
|
@ -63375,6 +63375,34 @@
|
|||
"video"
|
||||
]
|
||||
},
|
||||
"xai/grok-voice-transcribe-1.0": {
|
||||
"input_cost_per_second": 2.778e-05,
|
||||
"litellm_provider": "xai",
|
||||
"metadata": {
|
||||
"calculation": "$0.10/3600 seconds = $0.00002778 per second",
|
||||
"original_pricing_per_hour": 0.1
|
||||
},
|
||||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0,
|
||||
"source": "https://docs.x.ai/developers/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
]
|
||||
},
|
||||
"xai/grok-voice-transcribe-2.0": {
|
||||
"input_cost_per_second": 2.778e-05,
|
||||
"litellm_provider": "xai",
|
||||
"metadata": {
|
||||
"calculation": "$0.10/3600 seconds = $0.00002778 per second",
|
||||
"original_pricing_per_hour": 0.1
|
||||
},
|
||||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0,
|
||||
"source": "https://docs.x.ai/developers/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1024/grok-imagine-image-2.0": {
|
||||
"input_cost_per_image": 0.04,
|
||||
"litellm_provider": "xai",
|
||||
|
|
|
|||
|
|
@ -8730,6 +8730,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return ElevenLabsAudioTranscriptionConfig()
|
||||
elif litellm.LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.audio_transcription.transformation import (
|
||||
XAIAudioTranscriptionConfig,
|
||||
)
|
||||
|
||||
return XAIAudioTranscriptionConfig()
|
||||
elif litellm.LlmProviders.OPENAI == provider:
|
||||
if "gpt-4o" in model:
|
||||
return litellm.OpenAIGPTAudioTranscriptionConfig()
|
||||
|
|
|
|||
|
|
@ -63375,6 +63375,34 @@
|
|||
"video"
|
||||
]
|
||||
},
|
||||
"xai/grok-voice-transcribe-1.0": {
|
||||
"input_cost_per_second": 2.778e-05,
|
||||
"litellm_provider": "xai",
|
||||
"metadata": {
|
||||
"calculation": "$0.10/3600 seconds = $0.00002778 per second",
|
||||
"original_pricing_per_hour": 0.1
|
||||
},
|
||||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0,
|
||||
"source": "https://docs.x.ai/developers/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
]
|
||||
},
|
||||
"xai/grok-voice-transcribe-2.0": {
|
||||
"input_cost_per_second": 2.778e-05,
|
||||
"litellm_provider": "xai",
|
||||
"metadata": {
|
||||
"calculation": "$0.10/3600 seconds = $0.00002778 per second",
|
||||
"original_pricing_per_hour": 0.1
|
||||
},
|
||||
"mode": "audio_transcription",
|
||||
"output_cost_per_second": 0.0,
|
||||
"source": "https://docs.x.ai/developers/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
]
|
||||
},
|
||||
"low/1024-x-1024/grok-imagine-image-2.0": {
|
||||
"input_cost_per_image": 0.04,
|
||||
"litellm_provider": "xai",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,172 @@
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.xai.audio_transcription.transformation import (
|
||||
XAIAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
CONFIG = XAIAudioTranscriptionConfig()
|
||||
|
||||
WAV_BYTES = b"RIFF" + b"\x00" * 64
|
||||
|
||||
|
||||
def test_transform_request_serializes_provider_params():
|
||||
result = CONFIG.transform_audio_transcription_request(
|
||||
model="grok-voice-transcribe-2.0",
|
||||
audio_file=WAV_BYTES,
|
||||
optional_params={
|
||||
"language": "en",
|
||||
"diarize": True,
|
||||
"keyterm": ["LiteLLM", "Grok"],
|
||||
},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert isinstance(result, AudioTranscriptionRequestData)
|
||||
data = result.data
|
||||
assert data["model"] == "grok-voice-transcribe-2.0"
|
||||
assert data["language"] == "en"
|
||||
assert data["diarize"] == "true"
|
||||
assert data["keyterm"] == ["LiteLLM", "Grok"]
|
||||
filename, content, content_type = result.files["file"]
|
||||
assert content == WAV_BYTES
|
||||
assert isinstance(filename, str)
|
||||
assert isinstance(content_type, str)
|
||||
|
||||
|
||||
def test_transform_request_flattens_extra_body():
|
||||
result = CONFIG.transform_audio_transcription_request(
|
||||
model="grok-voice-transcribe-1.0",
|
||||
audio_file=WAV_BYTES,
|
||||
optional_params={
|
||||
"language": "en",
|
||||
"extra_body": {"diarize": False, "channels": 2},
|
||||
},
|
||||
litellm_params={},
|
||||
)
|
||||
assert result.data["diarize"] == "false"
|
||||
assert result.data["channels"] == "2"
|
||||
assert "extra_body" not in result.data
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base,expected",
|
||||
[
|
||||
(None, "https://api.x.ai/v1/stt"),
|
||||
("https://api.x.ai/v1", "https://api.x.ai/v1/stt"),
|
||||
("https://api.x.ai/v1/", "https://api.x.ai/v1/stt"),
|
||||
("https://proxy.example/", "https://proxy.example/v1/stt"),
|
||||
],
|
||||
)
|
||||
def test_get_complete_url(api_base, expected):
|
||||
url = CONFIG.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="grok-voice-transcribe-2.0",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url == expected
|
||||
|
||||
|
||||
def test_validate_environment_sets_bearer_header():
|
||||
headers = CONFIG.validate_environment(
|
||||
headers={},
|
||||
model="grok-voice-transcribe-2.0",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key="sk-test",
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer sk-test"
|
||||
assert "Content-Type" not in headers
|
||||
|
||||
|
||||
def test_validate_environment_requires_key(monkeypatch):
|
||||
monkeypatch.delenv("XAI_API_KEY", raising=False)
|
||||
monkeypatch.setattr(litellm, "xai_key", None)
|
||||
with pytest.raises(ValueError):
|
||||
CONFIG.validate_environment(
|
||||
headers={},
|
||||
model="grok-voice-transcribe-2.0",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
def test_transform_response_maps_xai_shape():
|
||||
raw = httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"text": "hello world",
|
||||
"language": "en",
|
||||
"duration": 3.2,
|
||||
"words": [
|
||||
{"text": "hello", "start": 0.0, "end": 0.5, "speaker": "1"},
|
||||
{"text": "world", "start": 0.5, "end": 1.0},
|
||||
],
|
||||
},
|
||||
request=httpx.Request("POST", "https://api.x.ai/v1/stt"),
|
||||
)
|
||||
response = CONFIG.transform_audio_transcription_response(raw_response=raw)
|
||||
|
||||
assert response.text == "hello world"
|
||||
assert response["language"] == "en"
|
||||
assert response["duration"] == 3.2
|
||||
assert response["task"] == "transcribe"
|
||||
assert response["words"] == [
|
||||
{"word": "hello", "start": 0.0, "end": 0.5, "speaker": "1"},
|
||||
{"word": "world", "start": 0.5, "end": 1.0},
|
||||
]
|
||||
assert response._hidden_params["audio_transcription_duration"] == 3.2
|
||||
|
||||
|
||||
def test_transcription_routes_to_xai_stt(monkeypatch):
|
||||
monkeypatch.delenv("XAI_API_KEY", raising=False)
|
||||
monkeypatch.setattr(litellm, "xai_key", None)
|
||||
captured: dict = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
captured["request"] = request
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={"text": "transcribed text", "language": "en", "duration": 1.5},
|
||||
request=request,
|
||||
)
|
||||
|
||||
http_handler = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
|
||||
response = litellm.transcription(
|
||||
model="xai/grok-voice-transcribe-2.0",
|
||||
file=("sample.wav", WAV_BYTES, "audio/wav"),
|
||||
api_key="sk-test",
|
||||
diarize=True,
|
||||
keyterm=["LiteLLM"],
|
||||
client=http_handler,
|
||||
)
|
||||
|
||||
request = captured["request"]
|
||||
assert str(request.url) == "https://api.x.ai/v1/stt"
|
||||
assert request.headers["Authorization"] == "Bearer sk-test"
|
||||
body = request.content.decode("utf-8", errors="replace")
|
||||
assert 'name="model"' in body and "grok-voice-transcribe-2.0" in body
|
||||
assert 'name="diarize"' in body and "true" in body
|
||||
assert 'name="keyterm"' in body and "LiteLLM" in body
|
||||
assert 'name="file"' in body
|
||||
assert response.text == "transcribed text"
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_xai_config():
|
||||
config = ProviderConfigManager.get_provider_audio_transcription_config(
|
||||
model="grok-voice-transcribe-2.0",
|
||||
provider=LlmProviders.XAI,
|
||||
)
|
||||
assert isinstance(config, XAIAudioTranscriptionConfig)
|
||||
Loading…
Add table
Reference in a new issue