Update STT service settings
This commit is contained in:
@@ -16,7 +16,7 @@ Provides two STT services:
|
||||
|
||||
import base64
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncGenerator, Literal, Optional, Union
|
||||
|
||||
from loguru import logger
|
||||
@@ -35,7 +35,7 @@ from pipecat.frames.frames import (
|
||||
VADUserStoppedSpeakingFrame,
|
||||
)
|
||||
from pipecat.processors.frame_processor import FrameDirection
|
||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
||||
from pipecat.services.settings import STTSettings, _NotGiven, _warn_deprecated_param
|
||||
from pipecat.services.stt_latency import OPENAI_REALTIME_TTFS_P99, OPENAI_TTFS_P99
|
||||
from pipecat.services.stt_service import WebsocketSTTService
|
||||
from pipecat.services.whisper.base_stt import (
|
||||
@@ -55,6 +55,13 @@ except ModuleNotFoundError:
|
||||
State = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class OpenAISTTSettings(BaseWhisperSTTSettings):
|
||||
"""Settings for the OpenAI STT service."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class OpenAISTTService(BaseWhisperSTTService):
|
||||
"""OpenAI Speech-to-Text service that generates text from audio.
|
||||
|
||||
@@ -62,6 +69,8 @@ class OpenAISTTService(BaseWhisperSTTService):
|
||||
set via the api_key parameter or OPENAI_API_KEY environment variable.
|
||||
"""
|
||||
|
||||
_settings: OpenAISTTSettings
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
@@ -71,7 +80,7 @@ class OpenAISTTService(BaseWhisperSTTService):
|
||||
language: Optional[Language] = Language.EN,
|
||||
prompt: Optional[str] = None,
|
||||
temperature: Optional[float] = None,
|
||||
settings: Optional[BaseWhisperSTTSettings] = None,
|
||||
settings: Optional[OpenAISTTSettings] = None,
|
||||
ttfs_p99_latency: Optional[float] = OPENAI_TTFS_P99,
|
||||
**kwargs,
|
||||
):
|
||||
@@ -81,13 +90,25 @@ class OpenAISTTService(BaseWhisperSTTService):
|
||||
model: Model to use — either gpt-4o or Whisper.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=BaseWhisperSTTSettings(model=...)`` instead.
|
||||
Use ``settings=OpenAISTTSettings(model=...)`` instead.
|
||||
|
||||
api_key: OpenAI API key. Defaults to None.
|
||||
base_url: API base URL. Defaults to None.
|
||||
language: Language of the audio input. Defaults to English.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=OpenAISTTSettings(language=...)`` instead.
|
||||
|
||||
prompt: Optional text to guide the model's style or continue a previous segment.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=OpenAISTTSettings(prompt=...)`` instead.
|
||||
|
||||
temperature: Optional sampling temperature between 0 and 1. Defaults to 0.0.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=OpenAISTTSettings(temperature=...)`` instead.
|
||||
|
||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||
parameters, ``settings`` values take precedence.
|
||||
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
||||
@@ -96,18 +117,23 @@ class OpenAISTTService(BaseWhisperSTTService):
|
||||
"""
|
||||
# --- 1. Hardcoded defaults ---
|
||||
_language = language or Language.EN
|
||||
default_settings = BaseWhisperSTTSettings(
|
||||
default_settings = OpenAISTTSettings(
|
||||
model="gpt-4o-transcribe",
|
||||
language=self.language_to_service_language(_language),
|
||||
base_url=base_url,
|
||||
prompt=prompt,
|
||||
temperature=temperature,
|
||||
prompt=None,
|
||||
temperature=None,
|
||||
)
|
||||
|
||||
# --- 2. Deprecated direct-arg overrides ---
|
||||
if model is not None:
|
||||
_warn_deprecated_param("model", BaseWhisperSTTSettings, "model")
|
||||
_warn_deprecated_param("model", OpenAISTTSettings, "model")
|
||||
default_settings.model = model
|
||||
if prompt is not None:
|
||||
_warn_deprecated_param("prompt", OpenAISTTSettings, "prompt")
|
||||
default_settings.prompt = prompt
|
||||
if temperature is not None:
|
||||
_warn_deprecated_param("temperature", OpenAISTTSettings, "temperature")
|
||||
default_settings.temperature = temperature
|
||||
|
||||
# --- 3. (no params object for this service) ---
|
||||
|
||||
@@ -124,7 +150,7 @@ class OpenAISTTService(BaseWhisperSTTService):
|
||||
)
|
||||
|
||||
async def _transcribe(self, audio: bytes) -> Transcription:
|
||||
assert self._language is not None # Assigned in the BaseWhisperSTTService class
|
||||
assert self._settings.language is not None
|
||||
|
||||
# Build kwargs dict with only set parameters
|
||||
kwargs = {
|
||||
@@ -162,7 +188,7 @@ class OpenAIRealtimeSTTSettings(STTSettings):
|
||||
prompt: Optional prompt text to guide transcription style.
|
||||
"""
|
||||
|
||||
prompt: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||
prompt: str | None | _NotGiven = None
|
||||
|
||||
|
||||
class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
@@ -228,8 +254,16 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
base_url: WebSocket base URL for the Realtime API.
|
||||
Defaults to ``"wss://api.openai.com/v1/realtime"``.
|
||||
language: Language of the audio input. Defaults to English.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=OpenAIRealtimeSTTSettings(language=...)`` instead.
|
||||
|
||||
prompt: Optional prompt text to guide transcription style
|
||||
or provide keyword hints.
|
||||
|
||||
.. deprecated:: 0.0.105
|
||||
Use ``settings=OpenAIRealtimeSTTSettings(prompt=...)`` instead.
|
||||
|
||||
turn_detection: Server-side VAD configuration. Defaults to
|
||||
``False`` (disabled), which relies on a local VAD
|
||||
processor in the pipeline. Pass ``None`` to use server
|
||||
@@ -257,14 +291,20 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
# --- 1. Hardcoded defaults ---
|
||||
default_settings = OpenAIRealtimeSTTSettings(
|
||||
model="gpt-4o-transcribe",
|
||||
language=language,
|
||||
prompt=prompt,
|
||||
language=Language.EN,
|
||||
prompt=None,
|
||||
)
|
||||
|
||||
# --- 2. Deprecated direct-arg overrides ---
|
||||
if model is not None:
|
||||
_warn_deprecated_param("model", OpenAIRealtimeSTTSettings, "model")
|
||||
default_settings.model = model
|
||||
if language is not None and language != Language.EN:
|
||||
_warn_deprecated_param("language", OpenAIRealtimeSTTSettings, "language")
|
||||
default_settings.language = language
|
||||
if prompt is not None:
|
||||
_warn_deprecated_param("prompt", OpenAIRealtimeSTTSettings, "prompt")
|
||||
default_settings.prompt = prompt
|
||||
|
||||
# --- 3. (no params object for this service) ---
|
||||
|
||||
@@ -281,7 +321,6 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
self._api_key = api_key
|
||||
self._base_url = base_url
|
||||
|
||||
self._prompt = self._settings.prompt
|
||||
self._turn_detection = turn_detection
|
||||
self._noise_reduction = noise_reduction
|
||||
self._should_interrupt = should_interrupt
|
||||
@@ -318,8 +357,7 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
||||
"""Apply a settings delta and send session update if needed.
|
||||
|
||||
Keeps ``_language_code`` and ``_prompt`` in sync with settings
|
||||
and sends a ``session.update`` to the server when the session is active.
|
||||
Sends a ``session.update`` to the server when the session is active.
|
||||
|
||||
Args:
|
||||
delta: A :class:`STTSettings` (or ``OpenAIRealtimeSTTSettings``) delta.
|
||||
@@ -329,13 +367,7 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
"""
|
||||
changed = await super()._update_settings(delta)
|
||||
|
||||
if not changed:
|
||||
return changed
|
||||
|
||||
if "prompt" in changed and isinstance(self._settings, OpenAIRealtimeSTTSettings):
|
||||
self._prompt = self._settings.prompt
|
||||
|
||||
if self._session_ready:
|
||||
if changed and self._session_ready:
|
||||
await self._send_session_update()
|
||||
|
||||
return changed
|
||||
@@ -492,8 +524,8 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||
if language_code:
|
||||
transcription["language"] = language_code
|
||||
|
||||
if self._prompt:
|
||||
transcription["prompt"] = self._prompt
|
||||
if self._settings.prompt:
|
||||
transcription["prompt"] = self._settings.prompt
|
||||
|
||||
input_audio: dict = {
|
||||
"format": {
|
||||
|
||||
Reference in New Issue
Block a user