Merge pull request #4352 from pipecat-ai/mb/pyright-fixes-1-per-file

More pyright fixes
This commit is contained in:
Mark Backman
2026-04-23 14:14:36 -04:00
committed by GitHub
55 changed files with 242 additions and 163 deletions

View File

@@ -0,0 +1 @@
- `ToolsSchema(standard_tools=...)` now accepts any `Sequence[FunctionSchema | DirectFunction]` rather than requiring an exact `list` of the union. Callers can pass a narrower `list[FunctionSchema]` (or any other `Sequence`) without the type checker complaining about list invariance.

View File

@@ -0,0 +1 @@
- Fixed direct-function registration crashing for functions without a docstring. `DirectFunctionWrapper` passed `inspect.getdoc()`'s result to `docstring_parser.parse()`, which raises when the docstring is `None`. Functions now register cleanly whether or not they have a docstring; an empty docstring produces empty description and parameter metadata as expected.

1
changelog/4352.fixed.md Normal file
View File

@@ -0,0 +1 @@
- Fixed `AssemblyAISTTService`, `CartesiaSTTService`, `GradiumSTTService`, and `SonioxSTTService` crashing the pipeline on transient WebSocket send failures. Each `run_stt` sent audio directly without catching errors, so a single network hiccup mid-stream raised an uncaught exception through `process_frame`. The guards now log a warning and let the connection-state check on the next call handle recovery, matching the pattern used by Deepgram, xAI, Azure, and other push-based STTs.

View File

@@ -6,9 +6,6 @@
"exclude": ["**/*_pb2.py", "**/__pycache__"], "exclude": ["**/*_pb2.py", "**/__pycache__"],
"ignore": [ "ignore": [
"tests", "tests",
"src/pipecat/adapters/base_llm_adapter.py",
"src/pipecat/adapters/schemas/direct_function.py",
"src/pipecat/adapters/schemas/tools_schema.py",
"src/pipecat/adapters/services/anthropic_adapter.py", "src/pipecat/adapters/services/anthropic_adapter.py",
"src/pipecat/adapters/services/aws_nova_sonic_adapter.py", "src/pipecat/adapters/services/aws_nova_sonic_adapter.py",
"src/pipecat/adapters/services/bedrock_adapter.py", "src/pipecat/adapters/services/bedrock_adapter.py",
@@ -23,7 +20,6 @@
"src/pipecat/audio/filters/aic_filter.py", "src/pipecat/audio/filters/aic_filter.py",
"src/pipecat/audio/filters/krisp_viva_filter.py", "src/pipecat/audio/filters/krisp_viva_filter.py",
"src/pipecat/audio/filters/rnnoise_filter.py", "src/pipecat/audio/filters/rnnoise_filter.py",
"src/pipecat/audio/resamplers/soxr_stream_resampler.py",
"src/pipecat/audio/turn/smart_turn/local_smart_turn_v2.py", "src/pipecat/audio/turn/smart_turn/local_smart_turn_v2.py",
"src/pipecat/audio/turn/smart_turn/local_smart_turn_v3.py", "src/pipecat/audio/turn/smart_turn/local_smart_turn_v3.py",
"src/pipecat/audio/vad/silero.py", "src/pipecat/audio/vad/silero.py",
@@ -35,8 +31,6 @@
"src/pipecat/processors/frameworks/rtvi/processor.py", "src/pipecat/processors/frameworks/rtvi/processor.py",
"src/pipecat/processors/frameworks/strands_agents.py", "src/pipecat/processors/frameworks/strands_agents.py",
"src/pipecat/processors/gstreamer/pipeline_source.py", "src/pipecat/processors/gstreamer/pipeline_source.py",
"src/pipecat/processors/metrics/frame_processor_metrics.py",
"src/pipecat/services/ai_service.py",
"src/pipecat/services/anthropic/llm.py", "src/pipecat/services/anthropic/llm.py",
"src/pipecat/services/assemblyai/stt.py", "src/pipecat/services/assemblyai/stt.py",
"src/pipecat/services/asyncai/tts.py", "src/pipecat/services/asyncai/tts.py",
@@ -57,7 +51,6 @@
"src/pipecat/services/deepgram/flux/stt.py", "src/pipecat/services/deepgram/flux/stt.py",
"src/pipecat/services/deepgram/sagemaker/stt.py", "src/pipecat/services/deepgram/sagemaker/stt.py",
"src/pipecat/services/deepgram/sagemaker/tts.py", "src/pipecat/services/deepgram/sagemaker/tts.py",
"src/pipecat/services/deepgram/stt.py",
"src/pipecat/services/deepgram/tts.py", "src/pipecat/services/deepgram/tts.py",
"src/pipecat/services/elevenlabs/stt.py", "src/pipecat/services/elevenlabs/stt.py",
"src/pipecat/services/elevenlabs/tts.py", "src/pipecat/services/elevenlabs/tts.py",
@@ -71,22 +64,18 @@
"src/pipecat/services/google/tts.py", "src/pipecat/services/google/tts.py",
"src/pipecat/services/google/vertex/llm.py", "src/pipecat/services/google/vertex/llm.py",
"src/pipecat/services/gradium/stt.py", "src/pipecat/services/gradium/stt.py",
"src/pipecat/services/gradium/tts.py",
"src/pipecat/services/groq/tts.py", "src/pipecat/services/groq/tts.py",
"src/pipecat/services/heygen/api_interactive_avatar.py", "src/pipecat/services/heygen/api_interactive_avatar.py",
"src/pipecat/services/heygen/base_api.py", "src/pipecat/services/heygen/base_api.py",
"src/pipecat/services/heygen/client.py", "src/pipecat/services/heygen/client.py",
"src/pipecat/services/heygen/video.py", "src/pipecat/services/heygen/video.py",
"src/pipecat/services/hume/tts.py", "src/pipecat/services/hume/tts.py",
"src/pipecat/services/image_service.py",
"src/pipecat/services/inworld/realtime/llm.py", "src/pipecat/services/inworld/realtime/llm.py",
"src/pipecat/services/inworld/tts.py", "src/pipecat/services/inworld/tts.py",
"src/pipecat/services/kokoro/tts.py", "src/pipecat/services/kokoro/tts.py",
"src/pipecat/services/llm_service.py", "src/pipecat/services/llm_service.py",
"src/pipecat/services/lmnt/tts.py", "src/pipecat/services/lmnt/tts.py",
"src/pipecat/services/mcp_service.py",
"src/pipecat/services/mem0/memory.py", "src/pipecat/services/mem0/memory.py",
"src/pipecat/services/mistral/llm.py",
"src/pipecat/services/mistral/stt.py", "src/pipecat/services/mistral/stt.py",
"src/pipecat/services/mistral/tts.py", "src/pipecat/services/mistral/tts.py",
"src/pipecat/services/moondream/vision.py", "src/pipecat/services/moondream/vision.py",
@@ -105,11 +94,9 @@
"src/pipecat/services/resembleai/tts.py", "src/pipecat/services/resembleai/tts.py",
"src/pipecat/services/rime/tts.py", "src/pipecat/services/rime/tts.py",
"src/pipecat/services/sambanova/llm.py", "src/pipecat/services/sambanova/llm.py",
"src/pipecat/services/sarvam/llm.py",
"src/pipecat/services/sarvam/stt.py", "src/pipecat/services/sarvam/stt.py",
"src/pipecat/services/sarvam/tts.py", "src/pipecat/services/sarvam/tts.py",
"src/pipecat/services/simli/video.py", "src/pipecat/services/simli/video.py",
"src/pipecat/services/smallest/stt.py",
"src/pipecat/services/smallest/tts.py", "src/pipecat/services/smallest/tts.py",
"src/pipecat/services/soniox/stt.py", "src/pipecat/services/soniox/stt.py",
"src/pipecat/services/speechmatics/stt.py", "src/pipecat/services/speechmatics/stt.py",
@@ -118,26 +105,20 @@
"src/pipecat/services/tavus/video.py", "src/pipecat/services/tavus/video.py",
"src/pipecat/services/tts_service.py", "src/pipecat/services/tts_service.py",
"src/pipecat/services/ultravox/llm.py", "src/pipecat/services/ultravox/llm.py",
"src/pipecat/services/vision_service.py",
"src/pipecat/services/websocket_service.py", "src/pipecat/services/websocket_service.py",
"src/pipecat/services/whisper/stt.py", "src/pipecat/services/whisper/stt.py",
"src/pipecat/services/xai/realtime/events.py",
"src/pipecat/services/xai/realtime/llm.py", "src/pipecat/services/xai/realtime/llm.py",
"src/pipecat/services/xai/stt.py",
"src/pipecat/services/xai/tts.py",
"src/pipecat/services/xtts/tts.py", "src/pipecat/services/xtts/tts.py",
"src/pipecat/transports/base_output.py", "src/pipecat/transports/base_output.py",
"src/pipecat/transports/daily/transport.py", "src/pipecat/transports/daily/transport.py",
"src/pipecat/transports/heygen/transport.py", "src/pipecat/transports/heygen/transport.py",
"src/pipecat/transports/lemonslice/transport.py", "src/pipecat/transports/lemonslice/transport.py",
"src/pipecat/transports/livekit/transport.py", "src/pipecat/transports/livekit/transport.py",
"src/pipecat/transports/local/tk.py",
"src/pipecat/transports/smallwebrtc/connection.py", "src/pipecat/transports/smallwebrtc/connection.py",
"src/pipecat/transports/smallwebrtc/request_handler.py", "src/pipecat/transports/smallwebrtc/request_handler.py",
"src/pipecat/transports/smallwebrtc/transport.py", "src/pipecat/transports/smallwebrtc/transport.py",
"src/pipecat/transports/tavus/transport.py", "src/pipecat/transports/tavus/transport.py",
"src/pipecat/transports/websocket/client.py", "src/pipecat/transports/websocket/client.py",
"src/pipecat/transports/websocket/fastapi.py",
"src/pipecat/transports/websocket/server.py", "src/pipecat/transports/websocket/server.py",
"src/pipecat/transports/whatsapp/client.py" "src/pipecat/transports/whatsapp/client.py"
], ],

View File

@@ -127,7 +127,7 @@ class BaseDirectFunctionWrapper:
self.name = self.function.__name__ self.name = self.function.__name__
# Parse docstring for description and parameters # Parse docstring for description and parameters
docstring = docstring_parser.parse(inspect.getdoc(self.function)) docstring = docstring_parser.parse(inspect.getdoc(self.function) or "")
# Get function description # Get function description
self.description = (docstring.description or "").strip() self.description = (docstring.description or "").strip()

View File

@@ -10,6 +10,7 @@ This module provides schemas for managing both standardized function tools
and custom adapter-specific tools in the Pipecat framework. and custom adapter-specific tools in the Pipecat framework.
""" """
from collections.abc import Sequence
from enum import Enum from enum import Enum
from typing import Any from typing import Any
@@ -39,7 +40,7 @@ class ToolsSchema:
def __init__( def __init__(
self, self,
standard_tools: list[FunctionSchema | DirectFunction], standard_tools: Sequence[FunctionSchema | DirectFunction],
custom_tools: dict[AdapterType, list[dict[str, Any]]] | None = None, custom_tools: dict[AdapterType, list[dict[str, Any]]] | None = None,
) -> None: ) -> None:
"""Initialize the tools schema. """Initialize the tools schema.
@@ -75,7 +76,7 @@ class ToolsSchema:
return self._standard_tools return self._standard_tools
@property @property
def custom_tools(self) -> dict[AdapterType, list[dict[str, Any]]]: def custom_tools(self) -> dict[AdapterType, list[dict[str, Any]]] | None:
"""Get the custom tools dictionary. """Get the custom tools dictionary.
Returns: Returns:

View File

@@ -0,0 +1,135 @@
#
# Copyright (c) 2024-2026, Daily
#
# SPDX-License-Identifier: BSD 2-Clause License
#
"""Mistral LLM adapter for Pipecat.
Mistral's API uses an OpenAI-compatible interface but imposes three
conversation-history constraints that OpenAI does not:
1. **Tool messages must be followed by an assistant message.** A ``"tool"``
role message that isn't followed by an ``"assistant"`` message is
rejected.
2. **Only the initial contiguous system block is permitted.** A
``"system"`` message appearing after any non-system message must be
converted to ``"user"``.
3. **A trailing assistant message requires ``prefix=True``.** When the
conversation ends on an assistant message, Mistral expects the
``prefix`` flag set so it can continue from that partial reply.
This adapter extends ``OpenAILLMAdapter`` and applies those three fixups
before the messages reach ``build_chat_completion_params``.
"""
import copy
from typing import Any, cast
from openai.types.chat import ChatCompletionMessageParam
from pipecat.adapters.services.open_ai_adapter import OpenAILLMAdapter, OpenAILLMInvocationParams
from pipecat.processors.aggregators.llm_context import LLMContext
class MistralLLMAdapter(OpenAILLMAdapter):
"""Adapter that transforms messages to satisfy Mistral's API constraints.
Mistral accepts the OpenAI chat-completions schema but enforces extra
rules on conversation history. This adapter extends ``OpenAILLMAdapter``
and rewrites the messages produced by the parent to comply with those
rules before the request is built.
"""
def get_llm_invocation_params(
self,
context: LLMContext,
*,
system_instruction: str | None = None,
convert_developer_to_user: bool,
) -> OpenAILLMInvocationParams:
"""Get OpenAI-compatible invocation parameters with Mistral message fixes applied.
Args:
context: The LLM context containing messages, tools, etc.
system_instruction: Optional system instruction from service settings
or ``run_inference``. Forwarded to the parent adapter.
convert_developer_to_user: If True, convert "developer"-role messages
to "user"-role messages. Forwarded to the parent adapter.
Returns:
Dictionary of parameters for Mistral's ChatCompletion API, with
messages transformed to satisfy Mistral's constraints.
"""
params = super().get_llm_invocation_params(
context,
system_instruction=system_instruction,
convert_developer_to_user=convert_developer_to_user,
)
params["messages"] = self._transform_messages(list(params["messages"]))
return params
def _transform_messages(
self, messages: list[ChatCompletionMessageParam]
) -> list[ChatCompletionMessageParam]:
"""Transform messages to satisfy Mistral's API constraints.
Applies three transformation steps in order:
1. **Insert assistant messages after tool messages** — Any ``"tool"``
message not followed by an ``"assistant"`` message gets a minimal
``{"role": "assistant", "content": " "}`` inserted after it.
2. **Convert non-initial system messages to user** — System messages
after the initial contiguous system block are converted to
``"user"``, since Mistral only accepts system messages at the
start of a conversation.
3. **Set prefix on trailing assistant message** — If the final message
is an assistant message without a ``prefix`` field, set
``prefix=True`` so Mistral will continue the partial reply.
Args:
messages: List of OpenAI-shaped message dicts.
Returns:
Transformed list of messages satisfying Mistral's constraints.
"""
if not messages:
return messages
# Work on plain dicts: we need to mutate "role" (which OpenAI TypedDict
# variants tag with fixed Literals) and to attach Mistral's non-standard
# "prefix" field. Cast back on return — the outgoing list is valid for
# Mistral's extended schema even though it doesn't fit OpenAI's.
msgs: list[dict[str, Any]] = copy.deepcopy([dict(m) for m in messages])
# Step 1: ensure every "tool" message is followed by an "assistant".
insert_at: list[int] = []
for i, msg in enumerate(msgs):
if msg.get("role") == "tool":
is_last = i == len(msgs) - 1
if is_last or msgs[i + 1].get("role") != "assistant":
insert_at.append(i + 1)
for idx in reversed(insert_at):
msgs.insert(idx, {"role": "assistant", "content": " "})
# Step 2: convert non-initial system messages to "user".
# Mistral rejects system messages after any non-system message.
first_non_system = next(
(i for i, m in enumerate(msgs) if m.get("role") != "system"),
len(msgs),
)
for i in range(first_non_system, len(msgs)):
if msgs[i].get("role") == "system":
msgs[i]["role"] = "user"
# Step 3: set prefix on a trailing assistant message so Mistral will
# continue it rather than rejecting the turn.
last = msgs[-1]
if last.get("role") == "assistant" and "prefix" not in last:
last["prefix"] = True
return cast(list[ChatCompletionMessageParam], msgs)

View File

@@ -68,7 +68,7 @@ class SOXRStreamAudioResampler(BaseAudioResampler):
self._soxr_stream.clear() self._soxr_stream.clear()
self._last_resample_time = current_time self._last_resample_time = current_time
def _maybe_initialize_sox_stream(self, in_rate: int, out_rate: int): def _maybe_initialize_sox_stream(self, in_rate: int, out_rate: int) -> "soxr.ResampleStream":
if self._soxr_stream is None: if self._soxr_stream is None:
self._initialize(in_rate, out_rate) self._initialize(in_rate, out_rate)
else: else:
@@ -80,6 +80,9 @@ class SOXRStreamAudioResampler(BaseAudioResampler):
f"expected {self._in_rate}->{self._out_rate}, got {in_rate}->{out_rate}" f"expected {self._in_rate}->{self._out_rate}, got {in_rate}->{out_rate}"
) )
assert self._soxr_stream is not None
return self._soxr_stream
async def resample(self, audio: bytes, in_rate: int, out_rate: int) -> bytes: async def resample(self, audio: bytes, in_rate: int, out_rate: int) -> bytes:
"""Resample audio data using soxr.ResampleStream resampler library. """Resample audio data using soxr.ResampleStream resampler library.
@@ -94,8 +97,8 @@ class SOXRStreamAudioResampler(BaseAudioResampler):
if in_rate == out_rate: if in_rate == out_rate:
return audio return audio
self._maybe_initialize_sox_stream(in_rate, out_rate) stream = self._maybe_initialize_sox_stream(in_rate, out_rate)
audio_data = np.frombuffer(audio, dtype=np.int16) audio_data = np.frombuffer(audio, dtype=np.int16)
resampled_audio = self._soxr_stream.resample_chunk(audio_data) resampled_audio = stream.resample_chunk(audio_data)
result = resampled_audio.astype(np.int16).tobytes() result = resampled_audio.astype(np.int16).tobytes()
return result return result

View File

@@ -66,6 +66,8 @@ class FrameProcessorMetrics(BaseObject):
Returns: Returns:
The task manager instance for async operations. The task manager instance for async operations.
""" """
if self._task_manager is None:
raise RuntimeError("task_manager not set; call setup() first")
return self._task_manager return self._task_manager
@property @property

View File

@@ -66,8 +66,9 @@ class AIService(FrameProcessor):
Args: Args:
model: The name of the AI model to use. model: The name of the AI model to use.
""" """
model = self._settings.model
self.set_core_metrics_data( self.set_core_metrics_data(
MetricsData(processor=self.name, model=self._settings.model or "") MetricsData(processor=self.name, model=model if isinstance(model, str) else "")
) )
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):

View File

@@ -418,7 +418,7 @@ class AssemblyAISTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text conversion. """Process audio data for speech-to-text conversion.
Args: Args:
@@ -433,7 +433,11 @@ class AssemblyAISTTService(WebsocketSTTService):
while len(self._audio_buffer) >= self._chunk_size_bytes: while len(self._audio_buffer) >= self._chunk_size_bytes:
chunk = bytes(self._audio_buffer[: self._chunk_size_bytes]) chunk = bytes(self._audio_buffer[: self._chunk_size_bytes])
self._audio_buffer = self._audio_buffer[self._chunk_size_bytes :] self._audio_buffer = self._audio_buffer[self._chunk_size_bytes :]
await self._websocket.send(chunk) try:
await self._websocket.send(chunk)
except Exception as e:
logger.warning(f"{self}: send failed: {e}")
break
yield None yield None

View File

@@ -449,7 +449,7 @@ class AsyncAITTSService(WebsocketTTSService):
await super().on_audio_context_completed(context_id) await super().on_audio_context_completed(context_id)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Async API websocket endpoint. """Generate speech from text using Async API websocket endpoint.
Args: Args:
@@ -620,7 +620,7 @@ class AsyncAIHttpTTSService(TTSService):
self._output_sample_rate = self.sample_rate self._output_sample_rate = self.sample_rate
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Async's HTTP streaming API. """Generate speech from text using Async's HTTP streaming API.
Args: Args:

View File

@@ -196,7 +196,7 @@ class AWSTranscribeSTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data and send to AWS Transcribe. """Process audio data and send to AWS Transcribe.
Args: Args:

View File

@@ -191,7 +191,7 @@ class AzureSTTService(STTService):
return changed return changed
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text conversion. """Process audio data for speech-to-text conversion.
Feeds audio data to the Azure speech recognizer for processing. Feeds audio data to the Azure speech recognizer for processing.

View File

@@ -277,7 +277,7 @@ class CartesiaSTTService(WebsocketSTTService):
if self._websocket and self._websocket.state is State.OPEN: if self._websocket and self._websocket.state is State.OPEN:
await self._websocket.send("finalize") await self._websocket.send("finalize")
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text transcription. """Process audio data for speech-to-text transcription.
Args: Args:
@@ -290,7 +290,10 @@ class CartesiaSTTService(WebsocketSTTService):
if not self._websocket or self._websocket.state is not State.OPEN: if not self._websocket or self._websocket.state is not State.OPEN:
await self._connect() await self._connect()
await self._websocket.send(audio) try:
await self._websocket.send(audio)
except Exception as e:
logger.warning(f"{self}: send failed: {e}")
yield None yield None
async def _connect(self): async def _connect(self):

View File

@@ -660,7 +660,7 @@ class CartesiaTTSService(WebsocketTTSService):
await self._connect_websocket() await self._connect_websocket()
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Cartesia's streaming API. """Generate speech from text using Cartesia's streaming API.
Args: Args:
@@ -873,7 +873,7 @@ class CartesiaHttpTTSService(TTSService):
await self._close_session() await self._close_session()
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Cartesia's HTTP API. """Generate speech from text using Cartesia's HTTP API.
Args: Args:

View File

@@ -222,7 +222,7 @@ class DeepgramFluxSageMakerSTTService(DeepgramFluxSTTBase):
# Audio sending and response receiving # Audio sending and response receiving
# ------------------------------------------------------------------ # ------------------------------------------------------------------
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Deepgram Flux for transcription. """Send audio data to Deepgram Flux for transcription.
Args: Args:

View File

@@ -354,7 +354,7 @@ class DeepgramFluxSTTService(DeepgramFluxSTTBase, WebsocketService):
# Audio sending and receiving # Audio sending and receiving
# ------------------------------------------------------------------ # ------------------------------------------------------------------
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Deepgram Flux for transcription. """Send audio data to Deepgram Flux for transcription.
Transmits raw audio bytes to the Deepgram Flux API for real-time speech Transmits raw audio bytes to the Deepgram Flux API for real-time speech

View File

@@ -256,7 +256,7 @@ class DeepgramSageMakerSTTService(STTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Deepgram for transcription. """Send audio data to Deepgram for transcription.
Args: Args:

View File

@@ -325,7 +325,7 @@ class DeepgramSageMakerTTSService(TTSService):
logger.error(f"{self} error sending Flush message: {e}") logger.error(f"{self} error sending Flush message: {e}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Deepgram TTS on SageMaker. """Generate speech from text using Deepgram TTS on SageMaker.
Args: Args:

View File

@@ -514,7 +514,7 @@ class DeepgramSTTService(STTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Deepgram for transcription. """Send audio data to Deepgram for transcription.
Args: Args:

View File

@@ -330,7 +330,7 @@ class DeepgramTTSService(WebsocketTTSService):
logger.error(f"{self} error sending Flush message: {e}") logger.error(f"{self} error sending Flush message: {e}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Deepgram's WebSocket TTS API. """Generate speech from text using Deepgram's WebSocket TTS API.
Args: Args:
@@ -441,7 +441,7 @@ class DeepgramHttpTTSService(TTSService):
return True return True
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Deepgram's TTS API. """Generate speech from text using Deepgram's TTS API.
Args: Args:

View File

@@ -370,7 +370,7 @@ class ElevenLabsSTTService(SegmentedSTTService):
"""Handle a transcription result with tracing.""" """Handle a transcription result with tracing."""
await self.stop_processing_metrics() await self.stop_processing_metrics()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Transcribe an audio segment using ElevenLabs' STT API. """Transcribe an audio segment using ElevenLabs' STT API.
Args: Args:
@@ -674,7 +674,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
except Exception as e: except Exception as e:
logger.warning(f"Failed to send commit: {e}") logger.warning(f"Failed to send commit: {e}")
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text transcription. """Process audio data for speech-to-text transcription.
Args: Args:

View File

@@ -889,7 +889,7 @@ class ElevenLabsTTSService(WebsocketTTSService):
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using ElevenLabs' streaming WebSocket API. """Generate speech from text using ElevenLabs' streaming WebSocket API.
Args: Args:
@@ -1240,7 +1240,7 @@ class ElevenLabsHttpTTSService(TTSService):
return word_times return word_times
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using ElevenLabs streaming API with timestamps. """Generate speech from text using ElevenLabs streaming API with timestamps.
Makes a request to the ElevenLabs API to generate audio and timing data. Makes a request to the ElevenLabs API to generate audio and timing data.

View File

@@ -373,7 +373,7 @@ class FishAudioTTSService(InterruptibleTTSService):
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e) await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Fish Audio's streaming API. """Generate speech from text using Fish Audio's streaming API.
Args: Args:

View File

@@ -461,7 +461,7 @@ class GladiaSTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Run speech-to-text on audio data. """Run speech-to-text on audio data.
Args: Args:

View File

@@ -931,7 +931,7 @@ class GoogleSTTService(STTService):
except Exception as e: except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e) await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process an audio chunk for STT transcription. """Process an audio chunk for STT transcription.
Args: Args:

View File

@@ -333,7 +333,7 @@ class GradiumSTTService(WebsocketSTTService):
except Exception as e: except Exception as e:
logger.warning(f"Failed to send flush: {e}") logger.warning(f"Failed to send flush: {e}")
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text conversion. """Process audio data for speech-to-text conversion.
Args: Args:
@@ -350,7 +350,11 @@ class GradiumSTTService(WebsocketSTTService):
chunk = base64.b64encode(chunk).decode("utf-8") chunk = base64.b64encode(chunk).decode("utf-8")
msg = {"type": "audio", "audio": chunk} msg = {"type": "audio", "audio": chunk}
if self._websocket and self._websocket.state is State.OPEN: if self._websocket and self._websocket.state is State.OPEN:
await self._websocket.send(json.dumps(msg)) try:
await self._websocket.send(json.dumps(msg))
except Exception as e:
logger.warning(f"{self}: send failed: {e}")
break
yield None yield None

View File

@@ -356,7 +356,7 @@ class GradiumTTSService(WebsocketTTSService):
await self.push_error(error_msg=f"Error: {msg.get('message', msg)}") await self.push_error(error_msg=f"Error: {msg.get('message', msg)}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Gradium's streaming API. """Generate speech from text using Gradium's streaming API.
Args: Args:

View File

@@ -57,7 +57,8 @@ class ImageGenService(AIService):
Frame: Frames containing the generated image (typically ImageRawFrame Frame: Frames containing the generated image (typically ImageRawFrame
or URLImageRawFrame). or URLImageRawFrame).
""" """
pass raise NotImplementedError
yield # pragma: no cover
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
"""Process frames for image generation. """Process frames for image generation.

View File

@@ -283,7 +283,7 @@ class InworldHttpTTSService(TTSService):
return (word_times, chunk_end_time) return (word_times, chunk_end_time)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate TTS audio for the given text. """Generate TTS audio for the given text.
Args: Args:
@@ -1134,7 +1134,7 @@ class InworldTTSService(WebsocketTTSService):
await self.send_with_retry(json.dumps(msg), self._report_error) await self.send_with_retry(json.dumps(msg), self._report_error)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate TTS audio for the given text using the Inworld WebSocket TTS service. """Generate TTS audio for the given text using the Inworld WebSocket TTS service.
Args: Args:

View File

@@ -336,7 +336,7 @@ class LmntTTSService(InterruptibleTTSService):
logger.error(f"Invalid JSON message: {message}") logger.error(f"Invalid JSON message: {message}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate TTS audio from text using LMNT's streaming API. """Generate TTS audio from text using LMNT's streaming API.
Args: Args:

View File

@@ -10,8 +10,8 @@ from collections.abc import Sequence
from dataclasses import dataclass from dataclasses import dataclass
from loguru import logger from loguru import logger
from openai.types.chat import ChatCompletionMessageParam
from pipecat.adapters.services.mistral_adapter import MistralLLMAdapter
from pipecat.adapters.services.open_ai_adapter import OpenAILLMInvocationParams from pipecat.adapters.services.open_ai_adapter import OpenAILLMInvocationParams
from pipecat.frames.frames import FunctionCallFromLLM from pipecat.frames.frames import FunctionCallFromLLM
from pipecat.services.openai.base_llm import BaseOpenAILLMService from pipecat.services.openai.base_llm import BaseOpenAILLMService
@@ -36,6 +36,8 @@ class MistralLLMService(OpenAILLMService):
# This value is used by BaseOpenAILLMService when calling the adapter. # This value is used by BaseOpenAILLMService when calling the adapter.
supports_developer_role = False supports_developer_role = False
adapter_class = MistralLLMAdapter
Settings = MistralLLMSettings Settings = MistralLLMSettings
_settings: Settings _settings: Settings
@@ -92,60 +94,6 @@ class MistralLLMService(OpenAILLMService):
logger.debug(f"Creating Mistral client with api {base_url}") logger.debug(f"Creating Mistral client with api {base_url}")
return super().create_client(api_key, base_url, **kwargs) return super().create_client(api_key, base_url, **kwargs)
def _apply_mistral_fixups(
self, messages: list[ChatCompletionMessageParam]
) -> list[ChatCompletionMessageParam]:
"""Apply fixups to messages to meet Mistral-specific requirements.
1. A "tool"-role message must be followed by an assistant message.
2. "system"-role messages must only appear at the start of a
conversation.
3. Assistant messages must have prefix=True when they are the final
message in a conversation (but at no other point).
Args:
messages: The original list of messages.
Returns:
Messages with Mistral prefix requirement applied to final assistant message.
"""
if not messages:
return messages
# Create a copy to avoid modifying the original
fixed_messages = [dict(msg) for msg in messages]
# Ensure all tool responses are followed by an assistant message
assistant_insert_indices = []
for i, msg in enumerate(fixed_messages):
if msg.get("role") == "tool":
# If this is the last message or the next message is not assistant
if i == len(fixed_messages) - 1 or fixed_messages[i + 1].get("role") != "assistant":
assistant_insert_indices.append(i + 1)
for idx in reversed(assistant_insert_indices):
fixed_messages.insert(idx, {"role": "assistant", "content": " "})
# Convert any "system" messages that aren't at the start (i.e., after the initial contiguous block) to "user"
first_non_system_idx = next(
(i for i, msg in enumerate(fixed_messages) if msg.get("role") != "system"),
len(fixed_messages),
)
for i, msg in enumerate(fixed_messages):
if msg.get("role") == "system" and i >= first_non_system_idx:
msg["role"] = "user"
# Get the last message
last_message = fixed_messages[-1]
# Only add prefix=True to the last message if it's an assistant message
# and Mistral would otherwise reject it
if last_message.get("role") == "assistant" and "prefix" not in last_message:
last_message["prefix"] = True
return fixed_messages
async def run_function_calls(self, function_calls: Sequence[FunctionCallFromLLM]): async def run_function_calls(self, function_calls: Sequence[FunctionCallFromLLM]):
"""Execute function calls, filtering out already-completed ones. """Execute function calls, filtering out already-completed ones.
@@ -208,18 +156,14 @@ class MistralLLMService(OpenAILLMService):
def build_chat_completion_params(self, params_from_context: OpenAILLMInvocationParams) -> dict: def build_chat_completion_params(self, params_from_context: OpenAILLMInvocationParams) -> dict:
"""Build parameters for Mistral chat completion request. """Build parameters for Mistral chat completion request.
Handles Mistral-specific requirements including: Handles Mistral-specific parameter mapping (``random_seed`` in place
- Assistant message prefix requirement for API compatibility of ``seed``). Message-shape fixups required by Mistral are applied
- Parameter mapping (random_seed instead of seed) by :class:`MistralLLMAdapter` upstream.
- Core completion settings
""" """
# Apply Mistral's assistant prefix requirement for API compatibility
fixed_messages = self._apply_mistral_fixups(params_from_context["messages"])
params = { params = {
"model": self._settings.model, "model": self._settings.model,
"stream": True, "stream": True,
"messages": fixed_messages, "messages": params_from_context["messages"],
"tools": params_from_context["tools"], "tools": params_from_context["tools"],
"tool_choice": params_from_context["tool_choice"], "tool_choice": params_from_context["tool_choice"],
"frequency_penalty": self._settings.frequency_penalty, "frequency_penalty": self._settings.frequency_penalty,

View File

@@ -185,7 +185,7 @@ class MistralSTTService(STTService):
if self._connection and not self._connection.is_closed: if self._connection and not self._connection.is_closed:
await self._connection.flush_audio() await self._connection.flush_audio()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Mistral for transcription. """Send audio data to Mistral for transcription.
Args: Args:

View File

@@ -366,7 +366,7 @@ class NeuphonicTTSService(InterruptibleTTSService):
await self._websocket.send(json.dumps(msg)) await self._websocket.send(json.dumps(msg))
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Neuphonic's streaming API. """Generate speech from text using Neuphonic's streaming API.
Args: Args:
@@ -565,7 +565,7 @@ class NeuphonicHttpTTSService(TTSService):
return None return None
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Neuphonic streaming API. """Generate speech from text using Neuphonic streaming API.
Args: Args:

View File

@@ -498,7 +498,7 @@ class NvidiaSTTService(STTService):
) )
logger.trace(f"Interim Transcription: [{transcript}]") logger.trace(f"Interim Transcription: [{transcript}]")
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Process audio data for speech-to-text transcription. """Process audio data for speech-to-text transcription.
Args: Args:
@@ -756,7 +756,7 @@ class NvidiaSegmentedSTTService(SegmentedSTTService):
"""Handle a transcription result with tracing.""" """Handle a transcription result with tracing."""
pass pass
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Transcribe an audio segment. """Transcribe an audio segment.
Args: Args:

View File

@@ -526,7 +526,7 @@ class NvidiaTTSService(TTSService):
) )
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using NVIDIA Nemotron Speech TTS. """Generate speech from text using NVIDIA Nemotron Speech TTS.
On the first call for a turn, starts a persistent ``synthesize_online`` On the first call for a turn, starts a persistent ``synthesize_online``

View File

@@ -415,7 +415,7 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to the transcription session. """Send audio data to the transcription session.
Audio is streamed over the WebSocket. Transcription results arrive Audio is streamed over the WebSocket. Transcription results arrive

View File

@@ -431,7 +431,7 @@ class ResembleAITTSService(WebsocketTTSService):
await self._connect_websocket() await self._connect_websocket()
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Resemble AI's streaming API. """Generate speech from text using Resemble AI's streaming API.
Args: Args:

View File

@@ -603,7 +603,7 @@ class RimeTTSService(WebsocketTTSService):
self.reset_active_audio_context() self.reset_active_audio_context()
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Rime's streaming API. """Generate speech from text using Rime's streaming API.
Args: Args:
@@ -786,7 +786,7 @@ class RimeHttpTTSService(TTSService):
return language_to_rime_language(language) return language_to_rime_language(language)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Rime's HTTP API. """Generate speech from text using Rime's HTTP API.
Args: Args:
@@ -1142,7 +1142,7 @@ class RimeNonJsonTTSService(InterruptibleTTSService):
await self.push_error(error_msg=f"Error: {e}", exception=e) await self.push_error(error_msg=f"Error: {e}", exception=e)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Rime's streaming API. """Generate speech from text using Rime's streaming API.
Args: Args:

View File

@@ -83,7 +83,10 @@ class SarvamLLMService(OpenAILLMService):
if settings is not None: if settings is not None:
default_settings.apply_update(settings) default_settings.apply_update(settings)
self._validate_model(default_settings.model) model = default_settings.model
if not isinstance(model, str):
raise ValueError("Sarvam LLM requires a non-empty model string.")
self._validate_model(model)
super().__init__( super().__init__(
api_key=api_key, api_key=api_key,

View File

@@ -570,7 +570,7 @@ class SarvamSTTService(STTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Sarvam for transcription. """Send audio data to Sarvam for transcription.
Args: Args:

View File

@@ -569,7 +569,7 @@ class SarvamHttpTTSService(TTSService):
await super().start(frame) await super().start(frame)
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Sarvam AI's API. """Generate speech from text using Sarvam AI's API.
Args: Args:
@@ -1192,7 +1192,7 @@ class SarvamTTSService(InterruptibleTTSService):
logger.warning("WebSocket not ready, cannot send text") logger.warning("WebSocket not ready, cannot send text")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech audio frames from input text using Sarvam TTS. """Generate speech audio frames from input text using Sarvam TTS.
Sends text over WebSocket for synthesis and yields corresponding audio or status frames. Sends text over WebSocket for synthesis and yields corresponding audio or status frames.

View File

@@ -247,7 +247,7 @@ class SmallestSTTService(WebsocketSTTService):
except Exception as e: except Exception as e:
logger.warning(f"{self} failed to send finalize: {e}") logger.warning(f"{self} failed to send finalize: {e}")
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio to the Smallest Pulse WebSocket for transcription. """Send audio to the Smallest Pulse WebSocket for transcription.
Args: Args:

View File

@@ -390,7 +390,7 @@ class SmallestTTSService(InterruptibleTTSService):
logger.warning(f"{self} unknown message status: {msg}") logger.warning(f"{self} unknown message status: {msg}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using Smallest's WebSocket streaming API. """Generate speech from text using Smallest's WebSocket streaming API.
Args: Args:

View File

@@ -402,7 +402,7 @@ class SonioxSTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Send audio data to Soniox STT Service. """Send audio data to Soniox STT Service.
Args: Args:
@@ -412,7 +412,10 @@ class SonioxSTTService(WebsocketSTTService):
Frame: None (transcription results come via WebSocket callbacks). Frame: None (transcription results come via WebSocket callbacks).
""" """
if self._websocket and self._websocket.state is State.OPEN: if self._websocket and self._websocket.state is State.OPEN:
await self._websocket.send(audio) try:
await self._websocket.send(audio)
except Exception as e:
logger.warning(f"{self}: send failed: {e}")
yield None yield None

View File

@@ -1059,7 +1059,7 @@ class SpeechmaticsSTTService(STTService):
"""Record transcription event for tracing.""" """Record transcription event for tracing."""
pass pass
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Adds audio to the audio buffer and yields None.""" """Adds audio to the audio buffer and yields None."""
try: try:
if self._client: if self._client:

View File

@@ -274,7 +274,7 @@ class STTService(AIService):
return Language(language) return Language(language)
@abstractmethod @abstractmethod
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Run speech-to-text on the provided audio data. """Run speech-to-text on the provided audio data.
This method must be implemented by subclasses to provide actual speech This method must be implemented by subclasses to provide actual speech

View File

@@ -446,7 +446,7 @@ class TTSService(AIService):
# Converts the text to audio. # Converts the text to audio.
@abstractmethod @abstractmethod
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Run text-to-speech synthesis on the provided text. """Run text-to-speech synthesis on the provided text.
This method must be implemented by subclasses to provide actual TTS functionality. This method must be implemented by subclasses to provide actual TTS functionality.

View File

@@ -59,7 +59,8 @@ class VisionService(AIService):
Frame: Frames containing the vision analysis results, typically TextFrame Frame: Frames containing the vision analysis results, typically TextFrame
objects with descriptions or answers. objects with descriptions or answers.
""" """
pass raise NotImplementedError
yield # pragma: no cover
async def process_frame(self, frame: Frame, direction: FrameDirection): async def process_frame(self, frame: Frame, direction: FrameDirection):
"""Process frames, handling vision image frames for analysis. """Process frames, handling vision image frames for analysis.

View File

@@ -541,18 +541,6 @@ class InputAudioBufferCleared(ServerEvent):
type: Literal["input_audio_buffer.cleared"] type: Literal["input_audio_buffer.cleared"]
class ResponseCreated(ServerEvent):
"""Event indicating an assistant response has been created.
Parameters:
type: Event type, always "response.created".
response: The created response object.
"""
type: Literal["response.created"]
response: "Response"
class ResponseOutputItemAdded(ServerEvent): class ResponseOutputItemAdded(ServerEvent):
"""Event indicating an output item has been added to a response. """Event indicating an output item has been added to a response.

View File

@@ -209,7 +209,7 @@ class XAISTTService(WebsocketSTTService):
await super().cancel(frame) await super().cancel(frame)
await self._disconnect() await self._disconnect()
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]: async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame | None, None]:
"""Forward raw audio bytes to the xAI STT WebSocket. """Forward raw audio bytes to the xAI STT WebSocket.
Transcription frames are pushed from the receive task, not yielded Transcription frames are pushed from the receive task, not yielded

View File

@@ -188,7 +188,7 @@ class XAIHttpTTSService(TTSService):
self._session = None self._session = None
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate speech from text using xAI's TTS API.""" """Generate speech from text using xAI's TTS API."""
logger.debug(f"{self}: Generating TTS [{text}]") logger.debug(f"{self}: Generating TTS [{text}]")
@@ -466,7 +466,7 @@ class XAITTSService(InterruptibleTTSService):
logger.debug(f"{self}: unhandled xAI message type: {msg_type}") logger.debug(f"{self}: unhandled xAI message type: {msg_type}")
@traced_tts @traced_tts
async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame | None, None]:
"""Generate TTS audio from text using xAI's streaming WebSocket API.""" """Generate TTS audio from text using xAI's streaming WebSocket API."""
logger.debug(f"{self}: Generating TTS [{text}]") logger.debug(f"{self}: Generating TTS [{text}]")

View File

@@ -228,7 +228,7 @@ class TkOutputTransport(BaseOutputTransport):
# This holds a reference to the photo, preventing it from being garbage # This holds a reference to the photo, preventing it from being garbage
# collected. # collected.
self._image_label.image = photo self._image_label.image = photo # type: ignore[attr-defined]
class TkLocalTransport(BaseTransport): class TkLocalTransport(BaseTransport):

View File

@@ -255,7 +255,9 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
if self._params.serializer: if self._params.serializer:
await self._params.serializer.setup(frame) await self._params.serializer.setup(frame)
if not self._monitor_websocket_task and self._params.session_timeout: if not self._monitor_websocket_task and self._params.session_timeout:
self._monitor_websocket_task = self.create_task(self._monitor_websocket()) self._monitor_websocket_task = self.create_task(
self._monitor_websocket(self._params.session_timeout)
)
await self._client.trigger_client_connected() await self._client.trigger_client_connected()
await self.push_frame(ClientConnectedFrame()) await self.push_frame(ClientConnectedFrame())
if not self._receive_task: if not self._receive_task:
@@ -322,9 +324,9 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
if not self._client.is_closing: if not self._client.is_closing:
await self._client.trigger_client_disconnected() await self._client.trigger_client_disconnected()
async def _monitor_websocket(self): async def _monitor_websocket(self, timeout: int):
"""Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event.""" """Wait for ``timeout`` seconds, then trigger the client-timeout event if still open."""
await asyncio.sleep(self._params.session_timeout) await asyncio.sleep(timeout)
await self._client.trigger_client_timeout() await self._client.trigger_client_timeout()