Skip to content
Open
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
6 changes: 6 additions & 0 deletions .env.example
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,12 @@ AZURE_OPENAI_API_VERSION_CHAT=2023-03-15-preview
# --- Groq ---
GROQ_API_KEY=

# --- OpenAI ---
OPENAI_API_KEY=
OPENAI_TRANSCRIPTION_MODEL=gpt-4o-transcribe
OPENAI_TRANSCRIPTION_ENDPOINT=https://api.openai.com/v1/audio/transcriptions
# OPENAI_MODEL and OPENAI_API_ENDPOINT are also supported for compatibility.

# --- DeepInfra ---
DEEPINFRA_API_KEY=

Expand Down
140 changes: 140 additions & 0 deletions sapat/providers/openai.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
# ABOUTME: OpenAI transcription provider
# ABOUTME: Uses OpenAI's Audio API for file-based speech-to-text

import os
from typing import Optional

import requests

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


DEFAULT_ENDPOINT = "https://api.openai.com/v1/audio/transcriptions"
DIARIZATION_MODEL = "gpt-4o-transcribe-diarize"


@register
class OpenAIProvider(OpenAICompatProvider):
name = "openai"
config = ProviderConfig(
required_env_vars=["OPENAI_API_KEY"],
max_file_size_mb=25.0,
preferred_format=AudioFormat.MP3,
supports_correction=False,
default_model=os.getenv(
"OPENAI_TRANSCRIPTION_MODEL",
os.getenv("OPENAI_MODEL", "gpt-4o-transcribe"),
),
)
_env_key_for_auth = "OPENAI_API_KEY"

@property
def base_url(self) -> str:
return (
os.getenv("OPENAI_TRANSCRIPTION_ENDPOINT")
or os.getenv("OPENAI_API_ENDPOINT")
or DEFAULT_ENDPOINT
)

def resolve_model(self, model_alias: str) -> str:
aliases = {
"4o": "gpt-4o-transcribe",
"gpt4o": "gpt-4o-transcribe",
"mini": "gpt-4o-mini-transcribe",
"4o-mini": "gpt-4o-mini-transcribe",
"gpt4o-mini": "gpt-4o-mini-transcribe",
"w": "whisper-1",
"whisper": "whisper-1",
"diarize": DIARIZATION_MODEL,
"diarization": DIARIZATION_MODEL,
}
return aliases.get(model_alias, model_alias)

def _build_data(
self, model: str, language: str, prompt, temperature: float, **kwargs
) -> dict:
data = super()._build_data(model, language, prompt, temperature, **kwargs)

response_format = kwargs.get("response_format")
if response_format:
data["response_format"] = response_format

if model == DIARIZATION_MODEL:
data["response_format"] = response_format or "diarized_json"
data["chunking_strategy"] = kwargs.get("chunking_strategy", "auto")

speaker_names = kwargs.get("known_speaker_names")
if speaker_names:
data["known_speaker_names[]"] = speaker_names

speaker_references = kwargs.get("known_speaker_references")
if speaker_references:
data["known_speaker_references[]"] = speaker_references

include = kwargs.get("include")
if include:
data["include[]"] = include

timestamp_granularities = kwargs.get("timestamp_granularities")
if timestamp_granularities:
data["timestamp_granularities[]"] = timestamp_granularities

return data

def transcribe(
self,
audio_file: str,
model: str,
language: str = "en",
prompt: Optional[str] = None,
temperature: float = 0,
**kwargs,
) -> TranscriptionResult:
model = self.resolve_model(model)
if model == DIARIZATION_MODEL and prompt:
raise ValueError("OpenAI diarization transcriptions do not support prompts")
if model == DIARIZATION_MODEL:
unsupported = [
name
for name in ("include", "timestamp_granularities")
if kwargs.get(name)
]
if unsupported:
unsupported_list = ", ".join(unsupported)
raise ValueError(
f"OpenAI diarization transcriptions do not support: {unsupported_list}"
)

headers = {self._auth_header_name: self._get_auth_value()}
data = self._build_data(model, language, prompt, temperature, **kwargs)

with open(audio_file, "rb") as f:
response = requests.post(
self.base_url,
headers=headers,
data=data,
files={"file": f},
)

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

if data.get("response_format") == "text":
return TranscriptionResult(text=response.text, raw_response=response.text)

result = response.json()
return TranscriptionResult(
text=result.get("text", ""),
language=result.get("language"),
duration=result.get("duration"),
segments=result.get("segments"),
raw_response=result,
)
161 changes: 148 additions & 13 deletions tests/providers/test_group_a.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# ABOUTME: Mock-based tests for all 8 transcription providers in group A
# ABOUTME: Mock-based tests for transcription providers in group A
# ABOUTME: Tests verify each provider sends correct auth, URL, and payload

import os
Expand Down Expand Up @@ -35,6 +35,117 @@ def json(self):
return self.payload


# ===========================================================================
# OpenAI
# ===========================================================================


class TestOpenAIProvider:
@patch.dict(
os.environ,
{
"OPENAI_API_KEY": "test-key",
"OPENAI_TRANSCRIPTION_ENDPOINT": "https://example.test/v1/audio/transcriptions",
},
clear=False,
)
@patch("sapat.providers.openai.requests.post")
def test_transcribe_sends_correct_request(self, mock_post, audio_file):
mock_post.return_value = FakeResponse(payload={"text": "hello openai"})

from sapat.providers.openai import OpenAIProvider

provider = OpenAIProvider()
result = provider.transcribe(
audio_file,
model="gpt4o",
language="en",
prompt="Product names include SAPAT.",
temperature=0,
)

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

_, kwargs = mock_post.call_args
assert (
mock_post.call_args.args[0]
== "https://example.test/v1/audio/transcriptions"
)
assert kwargs["headers"]["Authorization"] == "Bearer test-key"
assert kwargs["data"]["model"] == "gpt-4o-transcribe"
assert kwargs["data"]["language"] == "en"
assert kwargs["data"]["prompt"] == "Product names include SAPAT."
assert kwargs["data"]["response_format"] == "json"
assert "file" in kwargs["files"]

@patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}, clear=False)
@patch("sapat.providers.openai.requests.post")
def test_diarization_defaults_to_speaker_segments(self, mock_post, audio_file):
segments = [{"speaker": "speaker_0", "text": "hello", "start": 0, "end": 1}]
mock_post.return_value = FakeResponse(
payload={"text": "hello", "segments": segments, "duration": 1.0}
)

from sapat.providers.openai import OpenAIProvider

provider = OpenAIProvider()
result = provider.transcribe(audio_file, model="diarize")

_, kwargs = mock_post.call_args
assert kwargs["data"]["model"] == "gpt-4o-transcribe-diarize"
assert kwargs["data"]["response_format"] == "diarized_json"
assert kwargs["data"]["chunking_strategy"] == "auto"
assert result.segments == segments
assert result.raw_response["duration"] == 1.0

@patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}, clear=False)
def test_resolve_model_aliases(self):
from sapat.providers.openai import OpenAIProvider

provider = OpenAIProvider()
assert provider.resolve_model("whisper") == "whisper-1"
assert provider.resolve_model("mini") == "gpt-4o-mini-transcribe"
assert provider.resolve_model("diarization") == "gpt-4o-transcribe-diarize"

@patch.dict(os.environ, {}, clear=True)
def test_not_available_without_key(self):
from sapat.providers.openai import OpenAIProvider

assert OpenAIProvider.is_available() is False

@patch.dict(os.environ, {"OPENAI_API_KEY": "bad-key"}, clear=False)
@patch("sapat.providers.openai.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.openai import OpenAIProvider

provider = OpenAIProvider()
with pytest.raises(RuntimeError, match="401"):
provider.transcribe(audio_file, model="gpt-4o-transcribe")

@patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}, clear=False)
def test_diarization_rejects_prompt(self, audio_file):
from sapat.providers.openai import OpenAIProvider

provider = OpenAIProvider()
with pytest.raises(ValueError, match="do not support prompts"):
provider.transcribe(audio_file, model="diarize", prompt="Prefer names")

@patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}, clear=False)
def test_diarization_rejects_unsupported_options(self, audio_file):
from sapat.providers.openai import OpenAIProvider

provider = OpenAIProvider()
with pytest.raises(ValueError, match="timestamp_granularities"):
provider.transcribe(
audio_file,
model="diarize",
timestamp_granularities=["word"],
)


# ===========================================================================
# 1. DeepInfra
# ===========================================================================
Expand All @@ -56,7 +167,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 +219,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 +260,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 +339,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 +356,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 +378,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 +423,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 +469,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 +537,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
Loading