From 860cdc81d3a540c64d17cc6112ac577f1f9dd926 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 1 Dec 2025 18:26:56 -0800 Subject: [PATCH] [Fix] Fix Watsonx Audio Transcription API (#17326) * """ add * fix transform_audio_transcription_request * fix tests * test_watsonx_transcription_request_body --- .../audio_transcription/transformation.py | 78 ++++++++++++++++--- litellm/types/llms/watsonx.py | 36 ++++++++- ...sonx_audio_transcription_transformation.py | 35 ++++++++- 3 files changed, 131 insertions(+), 18 deletions(-) diff --git a/litellm/llms/watsonx/audio_transcription/transformation.py b/litellm/llms/watsonx/audio_transcription/transformation.py index 8c8324cb72d..8fe8b4a4248 100644 --- a/litellm/llms/watsonx/audio_transcription/transformation.py +++ b/litellm/llms/watsonx/audio_transcription/transformation.py @@ -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 diff --git a/litellm/types/llms/watsonx.py b/litellm/types/llms/watsonx.py index 4eb2f2531a0..6c42c3ecea0 100644 --- a/litellm/types/llms/watsonx.py +++ b/litellm/types/llms/watsonx.py @@ -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" diff --git a/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py b/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py index 84a9d25d98e..1286c2d4fe6 100644 --- a/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/watsonx/audio_transcription/test_watsonx_audio_transcription_transformation.py @@ -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)