Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 11 additions & 5 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,10 @@ This tool automates the process of transcribing video files using multiple trans
OPENAI_MODEL=whisper-1
OPENAI_API_ENDPOINT=https://api.openai.com/v1/audio/transcriptions
OPENAI_MODEL_NAME_CHAT=gpt-4o

# IBM Watson Speech to Text
IBM_WATSON_STT_API_KEY=your_ibm_watson_stt_api_key_here
IBM_WATSON_STT_URL=https://api.us-south.speech-to-text.watson.cloud.ibm.com/instances/your-instance-id
```

## Building and Installing the Package
Expand Down Expand Up @@ -111,15 +115,17 @@ sapat <video_file_or_directory> [--language <language>] [--prompt <prompt>] [--t
- `--prompt`: Optional prompt to guide the model's transcription.
- `--temperature`: The sampling temperature, between 0 and 1 (default: 0).
- `--quality`: Quality of the MP3 audio: 'L' for low, 'M' for medium, and 'H' for high (default: 'M').
- `--api`: Specify the API to use for transcription.
- `--api azure` for Azure OpenAI API
- `--api groq` for Groq Cloud API
- `--api openai` for OpenAI API
- `--provider`: Specify the provider to use for transcription.
- `--provider azure` for Azure OpenAI API
- `--provider groq` for Groq Cloud API
- `--provider openai` for OpenAI API
- `--provider ibm_watson` for IBM Watson Speech to Text
- `--model`: Specify the provider-specific model name.

Example:

```
sapat my_video.mp4 --quality H --language es --prompt "This is a test prompt" --temperature 0.5 --api groq
sapat my_video.mp4 --quality H --language es --prompt "This is a test prompt" --temperature 0.5 --provider groq
```

- If a file is provided, it will process that single file.
Expand Down
107 changes: 107 additions & 0 deletions sapat/providers/ibm_watson.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
# ABOUTME: IBM Watson Speech to Text transcription provider
# ABOUTME: Uses the Watson Speech to Text REST recognize endpoint

import os
from pathlib import Path
from typing import Optional

import requests

from sapat.providers import register
from sapat.providers.base import (
AudioFormat,
ProviderConfig,
TranscriptionProvider,
TranscriptionResult,
)


@register
class IBMWatsonProvider(TranscriptionProvider):
name = "ibm_watson"
config = ProviderConfig(
required_env_vars=["IBM_WATSON_STT_API_KEY", "IBM_WATSON_STT_URL"],
max_file_size_mb=100.0,
preferred_format=AudioFormat.MP3,
supports_correction=False,
default_model="en-US_BroadbandModel",
)

def resolve_model(self, model_alias: str) -> str:
aliases = {
"default": "en-US_BroadbandModel",
"en": "en-US_BroadbandModel",
"en-us": "en-US_BroadbandModel",
"en-gb": "en-GB_BroadbandModel",
"es": "es-ES_BroadbandModel",
"fr": "fr-FR_BroadbandModel",
"de": "de-DE_BroadbandModel",
"ja": "ja-JP_BroadbandModel",
"pt": "pt-BR_BroadbandModel",
}
return aliases.get(model_alias.lower(), model_alias)

def transcribe(
self,
audio_file: str,
model: str,
language: str = "en",
prompt: Optional[str] = None,
temperature: float = 0,
**kwargs,
) -> TranscriptionResult:
api_key = os.getenv("IBM_WATSON_STT_API_KEY", "")
service_url = os.getenv("IBM_WATSON_STT_URL", "").rstrip("/")
url = f"{service_url}/v1/recognize"

params = {
"model": self.resolve_model(model or self.config.default_model),
}
content_type = kwargs.get("content_type") or self._content_type(audio_file)
headers = {"Content-Type": content_type}

with open(audio_file, "rb") as f:
response = requests.post(
url,
auth=("apikey", api_key),
headers=headers,
params=params,
data=f.read(),
)

if response.status_code != 200:
raise RuntimeError(
f"IBM Watson transcription failed ({response.status_code}): "
f"{response.text}"
)

payload = response.json()
return TranscriptionResult(
text=self._extract_text(payload),
language=language,
raw_response=payload,
)

@staticmethod
def _extract_text(payload: dict) -> str:
transcripts = []
for result in payload.get("results", []):
alternatives = result.get("alternatives") or []
if alternatives:
transcript = alternatives[0].get("transcript", "").strip()
if transcript:
transcripts.append(transcript)
return " ".join(transcripts)

@staticmethod
def _content_type(audio_file: str) -> str:
extension = Path(audio_file).suffix.lower()
content_types = {
".flac": "audio/flac",
".mp3": "audio/mp3",
".mpeg": "audio/mpeg",
".ogg": "audio/ogg",
".wav": "audio/wav",
".webm": "audio/webm",
}
return content_types.get(extension, "application/octet-stream")
153 changes: 140 additions & 13 deletions tests/providers/test_group_a.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@

from sapat.providers.base import TranscriptionResult


# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -56,7 +55,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert mock_post.call_args.args[0] == "https://api.deepinfra.com/v1/openai/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "https://api.deepinfra.com/v1/openai/audio/transcriptions"
)
assert kwargs["data"]["model"] == "openai/whisper-large-v3"
assert "file" in kwargs["files"]

Expand Down Expand Up @@ -105,7 +107,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert mock_post.call_args.args[0] == "https://api.venice.ai/api/v1/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "https://api.venice.ai/api/v1/audio/transcriptions"
)
assert kwargs["data"]["model"] == "whisper-large-v3"
assert "file" in kwargs["files"]

Expand Down Expand Up @@ -143,7 +148,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert mock_post.call_args.args[0] == "https://api.together.xyz/v1/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "https://api.together.xyz/v1/audio/transcriptions"
)
assert kwargs["data"]["model"] == "openai/whisper-large-v3"
assert "file" in kwargs["files"]

Expand Down Expand Up @@ -219,7 +227,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert mock_post.call_args.args[0] == "https://api.mistral.ai/v1/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "https://api.mistral.ai/v1/audio/transcriptions"
)
assert kwargs["data"]["model"] == "voxtral-mini-latest"
# Mistral uses context_bias instead of prompt
assert "prompt" not in kwargs["data"]
Expand All @@ -233,17 +244,19 @@ def test_transcribe_sends_context_bias_for_prompt(self, mock_post, audio_file):
from sapat.providers.mistral import MistralProvider

provider = MistralProvider()
provider.transcribe(audio_file, model="voxtral-mini-latest", prompt="Product: Sapat")
provider.transcribe(
audio_file, model="voxtral-mini-latest", prompt="Product: Sapat"
)

_, kwargs = mock_post.call_args
assert kwargs["data"]["context_bias"] == "Product: Sapat"

@patch.dict(os.environ, {"MISTRAL_API_KEY": "test-key"}, clear=False)
@patch("sapat.providers.mistral.requests.post")
def test_correct_transcript_uses_chat_endpoint(self, mock_post):
chat_response = FakeResponse(payload={
"choices": [{"message": {"content": "corrected text"}}]
})
chat_response = FakeResponse(
payload={"choices": [{"message": {"content": "corrected text"}}]}
)
mock_post.return_value = chat_response

from sapat.providers.mistral import MistralProvider
Expand All @@ -253,7 +266,9 @@ def test_correct_transcript_uses_chat_endpoint(self, mock_post):

assert result == "corrected text"
_, kwargs = mock_post.call_args
assert mock_post.call_args.args[0] == "https://api.mistral.ai/v1/chat/completions"
assert (
mock_post.call_args.args[0] == "https://api.mistral.ai/v1/chat/completions"
)
assert kwargs["json"]["messages"][1]["content"] == "raw text"

@patch.dict(os.environ, {"MISTRAL_API_KEY": "test-key"}, clear=False)
Expand Down Expand Up @@ -296,7 +311,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert mock_post.call_args.args[0] == "https://api.lemonfox.ai/v1/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "https://api.lemonfox.ai/v1/audio/transcriptions"
)
assert kwargs["data"]["model"] == "whisper-1"
assert "file" in kwargs["files"]

Expand Down Expand Up @@ -339,7 +357,10 @@ def test_transcribe_sends_correct_request(self, mock_post, audio_file):
_, kwargs = mock_post.call_args
# No auth header when LOCALAI_API_KEY is not set
assert "Authorization" not in kwargs["headers"]
assert mock_post.call_args.args[0] == "http://localhost:8080/v1/audio/transcriptions"
assert (
mock_post.call_args.args[0]
== "http://localhost:8080/v1/audio/transcriptions"
)
assert kwargs["data"]["model"] == "whisper-1"
assert "file" in kwargs["files"]

Expand Down Expand Up @@ -404,7 +425,9 @@ def test_transcribe_sends_xi_api_key_header(self, mock_post, audio_file):
assert "xi-api-key" in kwargs["headers"]
assert kwargs["headers"]["xi-api-key"] == "test-key"
assert "Authorization" not in kwargs["headers"]
assert mock_post.call_args.args[0] == "https://api.elevenlabs.io/v1/speech-to-text"
assert (
mock_post.call_args.args[0] == "https://api.elevenlabs.io/v1/speech-to-text"
)
# ElevenLabs uses model_id, not model
assert kwargs["data"]["model_id"] == "scribe_v2"
assert "file" in kwargs["files"]
Expand Down Expand Up @@ -445,3 +468,107 @@ def test_raises_on_api_error(self, mock_post, audio_file):
provider = ElevenLabsProvider()
with pytest.raises(RuntimeError, match="401"):
provider.transcribe(audio_file, model="scribe_v2")


# ===========================================================================
# 9. IBM Watson
# ===========================================================================


class TestIBMWatsonProvider:
@patch.dict(
os.environ,
{
"IBM_WATSON_STT_API_KEY": "test-key",
"IBM_WATSON_STT_URL": "https://api.us-south.speech-to-text.watson.cloud.ibm.com/instances/abc",
},
clear=False,
)
@patch("sapat.providers.ibm_watson.requests.post")
def test_transcribe_sends_correct_request(self, mock_post, audio_file):
mock_post.return_value = FakeResponse(
payload={
"results": [
{"alternatives": [{"transcript": "hello "}]},
{"alternatives": [{"transcript": "watson"}]},
]
}
)

from sapat.providers.ibm_watson import IBMWatsonProvider

provider = IBMWatsonProvider()
result = provider.transcribe(audio_file, model="en")

assert isinstance(result, TranscriptionResult)
assert result.text == "hello watson"

assert mock_post.call_args.args[0].endswith("/v1/recognize")
_, kwargs = mock_post.call_args
assert kwargs["auth"] == ("apikey", "test-key")
assert kwargs["headers"]["Content-Type"] == "audio/mp3"
assert kwargs["params"]["model"] == "en-US_BroadbandModel"
assert kwargs["data"] == b"fake audio data"

@patch.dict(
os.environ,
{
"IBM_WATSON_STT_API_KEY": "test-key",
"IBM_WATSON_STT_URL": "https://watson.example",
},
clear=False,
)
@patch("sapat.providers.ibm_watson.requests.post")
def test_allows_content_type_override(self, mock_post, audio_file):
mock_post.return_value = FakeResponse(payload={"results": []})

from sapat.providers.ibm_watson import IBMWatsonProvider

provider = IBMWatsonProvider()
provider.transcribe(
audio_file,
model="default",
content_type="audio/mpeg",
)

_, kwargs = mock_post.call_args
assert kwargs["headers"]["Content-Type"] == "audio/mpeg"

@patch.dict(
os.environ,
{
"IBM_WATSON_STT_API_KEY": "test-key",
"IBM_WATSON_STT_URL": "https://watson.example",
},
clear=False,
)
def test_resolve_model_aliases(self):
from sapat.providers.ibm_watson import IBMWatsonProvider

provider = IBMWatsonProvider()
assert provider.resolve_model("en-gb") == "en-GB_BroadbandModel"
assert provider.resolve_model("custom-model") == "custom-model"

@patch.dict(os.environ, {}, clear=True)
def test_not_available_without_env(self):
from sapat.providers.ibm_watson import IBMWatsonProvider

assert IBMWatsonProvider.is_available() is False

@patch.dict(
os.environ,
{
"IBM_WATSON_STT_API_KEY": "bad-key",
"IBM_WATSON_STT_URL": "https://watson.example",
},
clear=False,
)
@patch("sapat.providers.ibm_watson.requests.post")
def test_raises_on_api_error(self, mock_post, audio_file):
mock_post.return_value = FakeResponse(status_code=401, text="unauthorized")

from sapat.providers.ibm_watson import IBMWatsonProvider

provider = IBMWatsonProvider()
with pytest.raises(RuntimeError, match="401"):
provider.transcribe(audio_file, model="default")