diff --git a/.env.example b/.env.example index 3e70778..6544648 100644 --- a/.env.example +++ b/.env.example @@ -85,6 +85,10 @@ FAL_KEY= # --- Soniox --- SONIOX_API_KEY= +# --- CAMB.AI --- +CAMB_API_KEY= +# CAMB_WORD_LEVEL_TIMESTAMPS=false + # --- Picovoice Leopard --- PICOVOICE_ACCESS_KEY= diff --git a/sapat/providers/cambai.py b/sapat/providers/cambai.py new file mode 100644 index 0000000..3982d22 --- /dev/null +++ b/sapat/providers/cambai.py @@ -0,0 +1,188 @@ +# ABOUTME: CAMB.AI async transcription provider +# ABOUTME: Uses POST /transcribe, polls task status, then fetches run transcript + +import os +from typing import Any, Dict, List, Optional + +import requests + +from sapat.providers import register +from sapat.providers.async_poll import AsyncPollProvider +from sapat.providers.base import ( + AudioFormat, + ProviderConfig, + TranscriptionResult, +) + + +LANGUAGE_ALIASES = { + "ar": "ar-sa", + "de": "de-de", + "en": "en-us", + "es": "es-es", + "fr": "fr-fr", + "ja": "ja-jp", + "zh": "zh-cn", +} + + +@register +class CambAIProvider(AsyncPollProvider): + """CAMB.AI speech-to-text provider.""" + + name = "cambai" + config = ProviderConfig( + required_env_vars=["CAMB_API_KEY"], + max_file_size_mb=20.0, + preferred_format=AudioFormat.MP3, + supports_correction=False, + default_model="default", + ) + + poll_interval: float = float(os.getenv("CAMB_POLL_INTERVAL_SECONDS", "5")) + max_poll_time: float = float(os.getenv("CAMB_TIMEOUT_SECONDS", "600")) + + def __init__(self): + super().__init__() + self.api_key = os.getenv("CAMB_API_KEY") + self.base_url = os.getenv( + "CAMB_API_BASE_URL", "https://client.camb.ai/apis" + ).rstrip("/") + self.word_level_timestamps = os.getenv( + "CAMB_WORD_LEVEL_TIMESTAMPS", "false" + ).lower() in {"1", "true", "yes", "on"} + self._run_ids: Dict[str, int] = {} + + def _headers(self) -> Dict[str, str]: + return {"x-api-key": self.api_key or ""} + + def _upload(self, audio_file: str, model: str, language: str, **kwargs) -> str: + data: Dict[str, Any] = { + "language": self._normalize_language(language), + } + self._add_optional_metadata(data) + + with open(audio_file, "rb") as media_file: + response = requests.post( + f"{self.base_url}/transcribe", + headers=self._headers(), + data=data, + files={"media_file": (os.path.basename(audio_file), media_file)}, + timeout=120, + ) + + if response.status_code not in (200, 201): + raise RuntimeError(f"CAMB.AI transcription task failed: {response.text}") + + task = response.json() + task_id = task.get("task_id") + if not task_id: + raise RuntimeError( + f"CAMB.AI task response did not include task_id: {task}" + ) + return str(task_id) + + def _poll(self, job_id: str) -> str: + response = requests.get( + f"{self.base_url}/transcribe/{job_id}", + headers=self._headers(), + timeout=60, + ) + if response.status_code != 200: + raise RuntimeError(f"CAMB.AI status check failed: {response.text}") + + status_payload = response.json() + status = str(status_payload.get("status", "")).upper() + if status == "SUCCESS": + run_id = status_payload.get("run_id") + if run_id is None: + raise RuntimeError( + f"CAMB.AI status response did not include run_id: {status_payload}" + ) + self._run_ids[job_id] = int(run_id) + return "completed" + if status == "PENDING": + return "pending" + if status in {"ERROR", "TIMEOUT", "PAYMENT_REQUIRED"}: + return "failed" + return "pending" + + def _fetch_result(self, job_id: str) -> TranscriptionResult: + run_id = self._run_ids.get(job_id) + if run_id is None: + raise RuntimeError(f"CAMB.AI run_id missing for task {job_id}") + + response = requests.get( + f"{self.base_url}/transcription-result/{run_id}", + headers=self._headers(), + params={"word_level_timestamps": self.word_level_timestamps}, + timeout=60, + ) + if response.status_code != 200: + raise RuntimeError(f"CAMB.AI result fetch failed: {response.text}") + + payload = response.json() + segments = self._extract_segments(payload) + text = self._extract_text(payload, segments) + return TranscriptionResult( + text=text, + segments=segments or None, + raw_response=payload, + ) + + def _add_optional_metadata(self, data: Dict[str, Any]) -> None: + project_name = os.getenv("CAMB_PROJECT_NAME") + project_description = os.getenv("CAMB_PROJECT_DESCRIPTION") + folder_id = os.getenv("CAMB_FOLDER_ID") + + if project_name: + data["project_name"] = project_name + if project_description: + data["project_description"] = project_description + if folder_id: + data["folder_id"] = folder_id + + def _normalize_language(self, language: Optional[str]) -> str: + if not language: + return "en-us" + + normalized = language.strip().lower().replace("_", "-") + return LANGUAGE_ALIASES.get(normalized, normalized) + + def _extract_segments(self, payload: Any) -> List[Dict[str, Any]]: + if isinstance(payload, list): + return [item for item in payload if isinstance(item, dict)] + + if not isinstance(payload, dict): + return [] + + for key in ("transcript", "segments", "dialogue"): + value = payload.get(key) + if isinstance(value, list): + return [item for item in value if isinstance(item, dict)] + + result = payload.get("result") + if isinstance(result, list): + return [item for item in result if isinstance(item, dict)] + if isinstance(result, dict): + return self._extract_segments(result) + + return [] + + def _extract_text(self, payload: Any, segments: List[Dict[str, Any]]) -> str: + if segments: + return "\n".join( + str(segment.get("text", "")).strip() + for segment in segments + if segment.get("text") + ) + + if isinstance(payload, dict): + text = payload.get("text") + if isinstance(text, str): + return text + + if isinstance(payload, list): + return "\n".join(str(item).strip() for item in payload if item) + + return "" diff --git a/tests/providers/test_group_b.py b/tests/providers/test_group_b.py index 5dd344a..e4ae367 100644 --- a/tests/providers/test_group_b.py +++ b/tests/providers/test_group_b.py @@ -325,7 +325,128 @@ def test_rejected_job_raises(self, mock_post, mock_get, audio_file): # =========================================================================== -# 4. Yandex SpeechKit +# 4. CAMB.AI +# =========================================================================== + + +class TestCambAIProvider: + @patch.dict( + os.environ, + { + "CAMB_API_KEY": "test-key", + "CAMB_API_BASE_URL": "https://client.camb.test/apis", + }, + clear=False, + ) + @patch("sapat.providers.cambai.requests.get") + @patch("sapat.providers.cambai.requests.post") + def test_transcribe_uploads_polls_and_fetches_segments( + self, mock_post, mock_get, audio_file + ): + mock_post.return_value = FakeResponse( + status_code=200, + payload={"task_id": "task-123"}, + ) + mock_get.side_effect = [ + FakeResponse(status_code=200, payload={"status": "PENDING"}), + FakeResponse( + status_code=200, + payload={"status": "SUCCESS", "run_id": 456}, + ), + FakeResponse( + status_code=200, + payload={ + "transcript": [ + { + "start": 0.0, + "end": 1.4, + "speaker": "Speaker 1", + "text": "First sentence.", + }, + { + "start": 1.5, + "end": 3.0, + "speaker": "Speaker 2", + "text": "Second sentence.", + }, + ] + }, + ), + ] + + from sapat.providers.cambai import CambAIProvider + + provider = CambAIProvider() + provider.poll_interval = 0.01 + result = provider.transcribe(audio_file, model="default", language="en") + + assert isinstance(result, TranscriptionResult) + assert result.text == "First sentence.\nSecond sentence." + assert result.segments and len(result.segments) == 2 + + mock_post.assert_called_once() + assert mock_post.call_args.args[0] == "https://client.camb.test/apis/transcribe" + assert mock_post.call_args.kwargs["headers"]["x-api-key"] == "test-key" + assert mock_post.call_args.kwargs["data"]["language"] == "en-us" + assert "media_file" in mock_post.call_args.kwargs["files"] + + assert mock_get.call_args_list[0].args[0] == ( + "https://client.camb.test/apis/transcribe/task-123" + ) + assert mock_get.call_args_list[-1].args[0] == ( + "https://client.camb.test/apis/transcription-result/456" + ) + + @patch.dict(os.environ, {"CAMB_API_KEY": "test-key"}, clear=False) + def test_available_with_key(self): + from sapat.providers.cambai import CambAIProvider + + assert CambAIProvider.is_available() is True + + @patch.dict(os.environ, {}, clear=True) + def test_not_available_without_key(self): + from sapat.providers.cambai import CambAIProvider + + assert CambAIProvider.is_available() is False + + @patch.dict(os.environ, {"CAMB_API_KEY": "test-key"}, clear=False) + def test_default_model(self): + from sapat.providers.cambai import CambAIProvider + + assert CambAIProvider.config.default_model == "default" + + @patch.dict(os.environ, {"CAMB_API_KEY": "test-key"}, clear=False) + def test_language_aliases(self): + from sapat.providers.cambai import CambAIProvider + + provider = CambAIProvider() + assert provider._normalize_language("en") == "en-us" + assert provider._normalize_language("pt-br") == "pt-br" + assert provider._normalize_language(None) == "en-us" + + @patch.dict(os.environ, {"CAMB_API_KEY": "test-key"}, clear=False) + @patch("sapat.providers.cambai.requests.get") + @patch("sapat.providers.cambai.requests.post") + def test_failed_task_raises(self, mock_post, mock_get, audio_file): + mock_post.return_value = FakeResponse( + status_code=200, + payload={"task_id": "task-123"}, + ) + mock_get.return_value = FakeResponse( + status_code=200, + payload={"status": "ERROR"}, + ) + + from sapat.providers.cambai import CambAIProvider + + provider = CambAIProvider() + provider.poll_interval = 0.01 + with pytest.raises(RuntimeError, match="failed"): + provider.transcribe(audio_file, model="default", language="en") + + +# =========================================================================== +# 5. Yandex SpeechKit # ===========================================================================