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:
kerry 2026-09-19 00:49:34 +00:00
parent a6e3a72ed8
commit 6b0ad3bed3
8 changed files with 434 additions and 1 deletions

View file

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

View file

@ -0,0 +1,3 @@
from .transformation import XAIAudioTranscriptionConfig
__all__ = ["XAIAudioTranscriptionConfig"]

View 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

View file

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

View file

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

View file

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

View file

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

View file

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