mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[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:
parent
37ecb03d4f
commit
860cdc81d3
3 changed files with 131 additions and 18 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue