Update GladiaSTTService to use the Gladia V2 API
This commit is contained in:
@@ -15,6 +15,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
- Added a foundational example for Gladia transcription:
|
- Added a foundational example for Gladia transcription:
|
||||||
`13c-gladia-transcription.py`
|
`13c-gladia-transcription.py`
|
||||||
|
|
||||||
|
### Changed
|
||||||
|
|
||||||
|
- Updated `GladiaSTTService` to use the V2 API.
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
- Fixed `enable_usage_metrics` to control LLM/TTS usage metrics separately
|
- Fixed `enable_usage_metrics` to control LLM/TTS usage metrics separately
|
||||||
|
|||||||
@@ -5,12 +5,16 @@
|
|||||||
#
|
#
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import aiohttp
|
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from loguru import logger
|
||||||
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import LLMMessagesFrame
|
from pipecat.frames.frames import EndFrame, LLMMessagesFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
@@ -20,12 +24,6 @@ from pipecat.services.gladia import GladiaSTTService
|
|||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
|
||||||
from runner import configure
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
|
|
||||||
load_dotenv(override=True)
|
load_dotenv(override=True)
|
||||||
|
|
||||||
logger.remove(0)
|
logger.remove(0)
|
||||||
@@ -90,6 +88,11 @@ async def main():
|
|||||||
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
await task.queue_frames([LLMMessagesFrame(messages)])
|
await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
|
# Register an event handler to exit the application when the user leaves.
|
||||||
|
@transport.event_handler("on_participant_left")
|
||||||
|
async def on_participant_left(transport, participant, reason):
|
||||||
|
await task.queue_frame(EndFrame())
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import base64
|
|||||||
import json
|
import json
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
|
|
||||||
@@ -23,7 +24,6 @@ from pipecat.services.ai_services import STTService
|
|||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
# See .env.example for Gladia configuration needed
|
|
||||||
try:
|
try:
|
||||||
import websockets
|
import websockets
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
@@ -38,15 +38,16 @@ class GladiaSTTService(STTService):
|
|||||||
class InputParams(BaseModel):
|
class InputParams(BaseModel):
|
||||||
sample_rate: Optional[int] = 16000
|
sample_rate: Optional[int] = 16000
|
||||||
language: Optional[Language] = Language.EN
|
language: Optional[Language] = Language.EN
|
||||||
transcription_hint: Optional[str] = None
|
endpointing: Optional[float] = 0.2
|
||||||
endpointing: Optional[int] = 200
|
maximum_duration_without_endpointing: Optional[int] = 10
|
||||||
prosody: Optional[bool] = None
|
audio_enhancer: Optional[bool] = None
|
||||||
|
words_accurate_timestamps: Optional[bool] = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
url: str = "wss://api.gladia.io/audio/text/audio-transcription",
|
url: str = "https://api.gladia.io/v2/live",
|
||||||
confidence: float = 0.5,
|
confidence: float = 0.5,
|
||||||
params: InputParams = InputParams(),
|
params: InputParams = InputParams(),
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -56,101 +57,82 @@ class GladiaSTTService(STTService):
|
|||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._url = url
|
self._url = url
|
||||||
self._settings = {
|
self._settings = {
|
||||||
|
"encoding": "wav/pcm",
|
||||||
|
"bit_depth": 16,
|
||||||
"sample_rate": params.sample_rate,
|
"sample_rate": params.sample_rate,
|
||||||
"language": self.language_to_service_language(params.language)
|
"channels": 1,
|
||||||
if params.language
|
"language_config": {
|
||||||
else Language.EN,
|
"languages": [self.language_to_service_language(params.language)]
|
||||||
"transcription_hint": params.transcription_hint,
|
if params.language
|
||||||
|
else [],
|
||||||
|
"code_switching": False,
|
||||||
|
},
|
||||||
"endpointing": params.endpointing,
|
"endpointing": params.endpointing,
|
||||||
"prosody": params.prosody,
|
"maximum_duration_without_endpointing": params.maximum_duration_without_endpointing,
|
||||||
|
"pre_processing": {
|
||||||
|
"audio_enhancer": params.audio_enhancer,
|
||||||
|
},
|
||||||
|
"realtime_processing": {
|
||||||
|
"words_accurate_timestamps": params.words_accurate_timestamps,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
self._confidence = confidence
|
self._confidence = confidence
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> str | None:
|
def language_to_service_language(self, language: Language) -> str | None:
|
||||||
match language:
|
language_map = {
|
||||||
case Language.BG:
|
Language.BG: "bg",
|
||||||
return "bulgarian"
|
Language.CA: "ca",
|
||||||
case Language.CA:
|
Language.ZH: "zh",
|
||||||
return "catalan"
|
Language.CS: "cs",
|
||||||
case Language.ZH:
|
Language.DA: "da",
|
||||||
return "chinese"
|
Language.NL: "nl",
|
||||||
case Language.CS:
|
Language.EN: "en",
|
||||||
return "czech"
|
Language.EN_US: "en",
|
||||||
case Language.DA:
|
Language.EN_AU: "en",
|
||||||
return "danish"
|
Language.EN_GB: "en",
|
||||||
case Language.NL:
|
Language.EN_NZ: "en",
|
||||||
return "dutch"
|
Language.EN_IN: "en",
|
||||||
case (
|
Language.ET: "et",
|
||||||
Language.EN
|
Language.FI: "fi",
|
||||||
| Language.EN_US
|
Language.FR: "fr",
|
||||||
| Language.EN_AU
|
Language.FR_CA: "fr",
|
||||||
| Language.EN_GB
|
Language.DE: "de",
|
||||||
| Language.EN_NZ
|
Language.DE_CH: "de",
|
||||||
| Language.EN_IN
|
Language.EL: "el",
|
||||||
):
|
Language.HI: "hi",
|
||||||
return "english"
|
Language.HU: "hu",
|
||||||
case Language.ET:
|
Language.ID: "id",
|
||||||
return "estonian"
|
Language.IT: "it",
|
||||||
case Language.FI:
|
Language.JA: "ja",
|
||||||
return "finnish"
|
Language.KO: "ko",
|
||||||
case Language.FR | Language.FR_CA:
|
Language.LV: "lv",
|
||||||
return "french"
|
Language.LT: "lt",
|
||||||
case Language.DE | Language.DE_CH:
|
Language.MS: "ms",
|
||||||
return "german"
|
Language.NO: "no",
|
||||||
case Language.EL:
|
Language.PL: "pl",
|
||||||
return "greek"
|
Language.PT: "pt",
|
||||||
case Language.HI:
|
Language.PT_BR: "pt",
|
||||||
return "hindi"
|
Language.RO: "ro",
|
||||||
case Language.HU:
|
Language.RU: "ru",
|
||||||
return "hungarian"
|
Language.SK: "sk",
|
||||||
case Language.ID:
|
Language.ES: "es",
|
||||||
return "indonesian"
|
Language.SV: "sv",
|
||||||
case Language.IT:
|
Language.TH: "th",
|
||||||
return "italian"
|
Language.TR: "tr",
|
||||||
case Language.JA:
|
Language.UK: "uk",
|
||||||
return "japanese"
|
Language.VI: "vi",
|
||||||
case Language.KO:
|
}
|
||||||
return "korean"
|
return language_map.get(language)
|
||||||
case Language.LV:
|
|
||||||
return "latvian"
|
|
||||||
case Language.LT:
|
|
||||||
return "lithuanian"
|
|
||||||
case Language.MS:
|
|
||||||
return "malay"
|
|
||||||
case Language.NO:
|
|
||||||
return "norwegian"
|
|
||||||
case Language.PL:
|
|
||||||
return "polish"
|
|
||||||
case Language.PT | Language.PT_BR:
|
|
||||||
return "portuguese"
|
|
||||||
case Language.RO:
|
|
||||||
return "romanian"
|
|
||||||
case Language.RU:
|
|
||||||
return "russian"
|
|
||||||
case Language.SK:
|
|
||||||
return "slovak"
|
|
||||||
case Language.ES:
|
|
||||||
return "spanish"
|
|
||||||
case Language.SV:
|
|
||||||
return "slovenian"
|
|
||||||
case Language.TH:
|
|
||||||
return "thai"
|
|
||||||
case Language.TR:
|
|
||||||
return "turkish"
|
|
||||||
case Language.UK:
|
|
||||||
return "ukrainian"
|
|
||||||
case Language.VI:
|
|
||||||
return "vietnamese"
|
|
||||||
return None
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
self._websocket = await websockets.connect(self._url)
|
response = await self._setup_gladia()
|
||||||
|
self._websocket = await websockets.connect(response["url"])
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
||||||
await self._setup_gladia()
|
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
|
await self._send_stop_recording()
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
@@ -164,39 +146,37 @@ class GladiaSTTService(STTService):
|
|||||||
yield None
|
yield None
|
||||||
|
|
||||||
async def _setup_gladia(self):
|
async def _setup_gladia(self):
|
||||||
configuration = {
|
async with aiohttp.ClientSession() as session:
|
||||||
"x_gladia_key": self._api_key,
|
async with session.post(
|
||||||
"encoding": "WAV/PCM",
|
self._url,
|
||||||
"model_type": "fast",
|
headers={"X-Gladia-Key": self._api_key, "Content-Type": "application/json"},
|
||||||
"language_behaviour": "manual",
|
json=self._settings,
|
||||||
"sample_rate": self._settings["sample_rate"],
|
) as response:
|
||||||
"language": self._settings["language"],
|
if response.ok:
|
||||||
"transcription_hint": self._settings["transcription_hint"],
|
return await response.json()
|
||||||
"endpointing": self._settings["endpointing"],
|
else:
|
||||||
"prosody": self._settings["prosody"],
|
logger.error(
|
||||||
}
|
f"Gladia error: {response.status}: {response.text or response.reason}"
|
||||||
|
)
|
||||||
await self._websocket.send(json.dumps(configuration))
|
raise Exception(f"Failed to initialize Gladia session: {response.status}")
|
||||||
|
|
||||||
async def _send_audio(self, audio: bytes):
|
async def _send_audio(self, audio: bytes):
|
||||||
message = {"frames": base64.b64encode(audio).decode("utf-8")}
|
data = base64.b64encode(audio).decode("utf-8")
|
||||||
|
message = {"type": "audio_chunk", "data": {"chunk": data}}
|
||||||
await self._websocket.send(json.dumps(message))
|
await self._websocket.send(json.dumps(message))
|
||||||
|
|
||||||
|
async def _send_stop_recording(self):
|
||||||
|
await self._websocket.send(json.dumps({"type": "stop_recording"}))
|
||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
async for message in self._websocket:
|
async for message in self._websocket:
|
||||||
utterance = json.loads(message)
|
content = json.loads(message)
|
||||||
if not utterance:
|
if content["type"] == "transcript":
|
||||||
continue
|
utterance = content["data"]["utterance"]
|
||||||
|
confidence = utterance.get("confidence", 0)
|
||||||
if "error" in utterance:
|
transcript = utterance["text"]
|
||||||
message = utterance["message"]
|
|
||||||
logger.error(f"Gladia error: {message}")
|
|
||||||
elif "confidence" in utterance:
|
|
||||||
type = utterance["type"]
|
|
||||||
confidence = utterance["confidence"]
|
|
||||||
transcript = utterance["transcription"]
|
|
||||||
if confidence >= self._confidence:
|
if confidence >= self._confidence:
|
||||||
if type == "final":
|
if content["data"]["is_final"]:
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(transcript, "", time_now_iso8601())
|
TranscriptionFrame(transcript, "", time_now_iso8601())
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user