[Fix] Fix Watsonx Audio Transcription API (#17326)

* """
add

* fix transform_audio_transcription_request

* fix tests

* test_watsonx_transcription_request_body
This commit is contained in:
Ishaan Jaff 2025-12-01 18:26:56 -08:00 • committed by GitHub
parent 37ecb03d4f
commit 860cdc81d3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 131 additions and 18 deletions

View file

@ -4,11 +4,17 @@ Translates from OpenAI's `/v1/audio/transcriptions` to IBM WatsonX's `/ml/v1/aud
WatsonX follows the OpenAI spec for audio transcription.
"""
from typing import List, Optional
from typing import Any, Dict, List, Optional
import litellm
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.types.llms.openai import OpenAIAudioTranscriptionOptionalParams
from litellm.types.llms.watsonx import WatsonXAudioTranscriptionRequestBody
from litellm.types.utils import FileTypes
from ...base_llm.audio_transcription.transformation import (
AudioTranscriptionRequestData,
)
from ...openai.transcriptions.whisper_transformation import (
OpenAIWhisperAudioTranscriptionConfig,
)
@ -40,6 +46,60 @@ class IBMWatsonXAudioTranscriptionConfig(
"timestamp_granularities",
]
def transform_audio_transcription_request(
self,
model: str,
audio_file: FileTypes,
optional_params: dict,
litellm_params: dict,
) -> AudioTranscriptionRequestData:
"""
Transform the audio transcription request for WatsonX.
WatsonX expects multipart/form-data with:
- file: the audio file
- model: the model name (without watsonx/ prefix)
- project_id: the project ID (as form field, not query param)
- other optional params
"""
# Use common utility to process the audio file
processed_audio = process_audio_file(audio_file)
# Get API params to extract project_id
api_params = _get_api_params(params=optional_params.copy())
# Initialize form data with required fields
form_data: WatsonXAudioTranscriptionRequestBody = {
"model": model,
"project_id": api_params.get("project_id", ""),
}
# Add supported OpenAI params to form data
supported_params = self.get_supported_openai_params(model)
for key, value in optional_params.items():
if key in supported_params and value is not None:
form_data[key] = value # type: ignore
# Set default response_format for cost calculation
if "response_format" not in form_data or (
form_data.get("response_format") in ["text", "json"]
):
form_data["response_format"] = "verbose_json"
# Prepare files dict with the audio file
files = {
"file": (
processed_audio.filename,
processed_audio.file_content,
processed_audio.content_type,
)
}
# Convert TypedDict to regular dict for AudioTranscriptionRequestData
form_data_dict: Dict[str, Any] = dict(form_data)
return AudioTranscriptionRequestData(data=form_data_dict, files=files)
def get_complete_url(
self,
api_base: Optional[str],
@ -52,7 +112,9 @@ class IBMWatsonXAudioTranscriptionConfig(
"""
Construct the complete URL for WatsonX audio transcription.
URL format: {api_base}/ml/v1/audio/transcriptions?version={version}&project_id={project_id}
URL format: {api_base}/ml/v1/audio/transcriptions?version={version}
Note: project_id is sent as form data, not as a query parameter
"""
# Get base URL
url = self._get_base_url(api_base=api_base)
@ -61,18 +123,10 @@ class IBMWatsonXAudioTranscriptionConfig(
# Add the audio transcription endpoint
url = f"{url}/ml/v1/audio/transcriptions"
# Get API params for project_id
api_params = _get_api_params(params=optional_params.copy())
# Add version parameter
api_version = optional_params.pop(
# Add version parameter (only version in query string, not project_id)
api_version = optional_params.get(
"api_version", None
) or litellm.WATSONX_DEFAULT_API_VERSION
url = f"{url}?version={api_version}"
# Add project_id parameter
project_id = api_params.get("project_id")
if project_id:
url = f"{url}&project_id={project_id}"
return url

View file

@ -1,9 +1,7 @@
import json
from enum import Enum
from typing import Any, List, Optional, Union
from typing import List, Optional
from pydantic import BaseModel
from typing_extensions import TypedDict
from typing_extensions import NotRequired, TypedDict
class WatsonXAPIParams(TypedDict):
@ -18,6 +16,36 @@ class WatsonXCredentials(TypedDict):
token: Optional[str]
class WatsonXAudioTranscriptionRequestBody(TypedDict):
"""
WatsonX Audio Transcription API request body.
Follows multipart/form-data format for WatsonX Whisper models.
See: https://cloud.ibm.com/apidocs/watsonx-ai
"""
model: str
"""Model name (e.g., 'whisper-large-v3-turbo')"""
project_id: str
"""WatsonX project ID (required)"""
language: NotRequired[str]
"""Language code (e.g., 'en', 'es')"""
prompt: NotRequired[str]
"""Optional prompt to guide transcription"""
response_format: NotRequired[str]
"""Response format: 'json', 'text', 'srt', 'verbose_json', 'vtt'"""
temperature: NotRequired[float]
"""Sampling temperature (0-1)"""
timestamp_granularities: NotRequired[List[str]]
"""Timestamp granularities: ['word', 'segment']"""
class WatsonXAIEndpoint(str, Enum):
TEXT_GENERATION = "/ml/v1/text/generation"
TEXT_GENERATION_STREAM = "/ml/v1/text/generation_stream"

View file

@ -4,6 +4,7 @@ Tests for IBM WatsonX Audio Transcription.
Validates that litellm.transcription transforms requests correctly for WatsonX.
"""
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -29,6 +30,7 @@ class TestWatsonXAudioTranscription:
captured_request["url"] = str(kwargs.get("url", args[0] if args else None))
captured_request["headers"] = kwargs.get("headers", {})
captured_request["data"] = kwargs.get("data", {})
captured_request["files"] = kwargs.get("files", {})
mock_response = MagicMock()
mock_response.json.return_value = {
@ -54,16 +56,30 @@ class TestWatsonXAudioTranscription:
# Validate URL contains WatsonX audio transcription endpoint
assert "/ml/v1/audio/transcriptions" in captured_request["url"]
assert "version=" in captured_request["url"]
assert "project_id=test-project-123" in captured_request["url"]
# project_id should NOT be in URL (it should be in form data instead)
assert "project_id=test-project-123" not in captured_request["url"]
# Validate headers contain WatsonX auth
assert "Authorization" in captured_request["headers"]
assert "Bearer test-bearer-token" in captured_request["headers"]["Authorization"]
# Validate project_id is in form data, not URL
assert captured_request["data"].get("project_id") == "test-project-123"
# Validate file is in files dict
assert "file" in captured_request["files"]
@pytest.mark.asyncio
async def test_watsonx_transcription_request_body(self):
"""
Test that litellm.transcription sends correct request body for WatsonX.
Validates that:
- Request uses multipart/form-data (data + files)
- Model name has watsonx/ prefix removed
- project_id is in form data, not URL
- Audio file is in files dict
- OpenAI params are included in form data
"""
captured_request = {}
@ -94,9 +110,24 @@ class TestWatsonXAudioTranscription:
except Exception:
pass # We just want to capture the request
# Validate request body contains expected fields
# Validate form data contains expected fields
data = captured_request.get("data", {})
print("JSON DUMPS captured_request:")
print(json.dumps(captured_request, indent=4, default=str))
# Model name should NOT have watsonx/ prefix
assert data.get("model") == "whisper-large-v3-turbo"
# project_id should be in form data
assert data.get("project_id") == "test-project-123"
# OpenAI params should be in form data
assert data.get("language") == "en"
assert data.get("temperature") == 0.5
assert data.get("response_format") == "verbose_json" # Default for cost calculation
# Validate file is in files dict (multipart/form-data)
files = captured_request.get("files", {})
assert "file" in files
assert isinstance(files["file"], tuple) # Should be (filename, content, content_type)