Merge pull request #1085 from pipecat-ai/aleix/task-creation-and-cancellation
improve task creation and cancellation
This commit is contained in:
16
CHANGELOG.md
16
CHANGELOG.md
@@ -9,6 +9,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
|
- In order to create tasks in Pipecat frame processors it is now recommended to
|
||||||
|
use `FrameProcessor.create_task()` (which uses the new
|
||||||
|
`utils.asyncio.create_task()`). It takes care of uncaught exceptions, task
|
||||||
|
cancellation handling and task management. To cancel or wait for a task there
|
||||||
|
is `FrameProcessor.cancel_task()` and `FrameProcessor.wait_for_task()`. All of
|
||||||
|
Pipecat processors have been updated accordingly. Also, when a pipeline runner
|
||||||
|
finishes, a warning about dangling tasks might appear, which indicates if any
|
||||||
|
of the created tasks was never cancelled or awaited for (using these new
|
||||||
|
functions).
|
||||||
|
|
||||||
- It is now possible to specify the period of the `PipelineTask` heartbeat
|
- It is now possible to specify the period of the `PipelineTask` heartbeat
|
||||||
frames with `heartbeats_period_secs`.
|
frames with `heartbeats_period_secs`.
|
||||||
|
|
||||||
@@ -41,6 +51,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed an `GeminiMultimodalLiveLLMService` issue that was preventing the user
|
||||||
|
to push initial LLM assistant messages (using `LLMMessagesAppendFrame`).
|
||||||
|
|
||||||
|
- Added missing `FrameProcessor.cleanup()` calls to `Pipeline`,
|
||||||
|
`ParallelPipeline` and `UserIdleProcessor`.
|
||||||
|
|
||||||
- Fixed a type error when using `voice_settings` in `ElevenLabsHttpTTSService`.
|
- Fixed a type error when using `voice_settings` in `ElevenLabsHttpTTSService`.
|
||||||
|
|
||||||
- Fixed an issue where `OpenAIRealtimeBetaLLMService` function calling resulted
|
- Fixed an issue where `OpenAIRealtimeBetaLLMService` function calling resulted
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -31,7 +31,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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 PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
@@ -93,7 +93,7 @@ async def main():
|
|||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": """You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way.
|
"content": """You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way.
|
||||||
|
|
||||||
You have one functions available:
|
You have one functions available:
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
@@ -93,7 +93,7 @@ async def main():
|
|||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": """You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way.
|
"content": """You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way.
|
||||||
|
|
||||||
You have one functions available:
|
You have one functions available:
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from openai.types.chat import ChatCompletionToolParam
|
|||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import TTSSpeakFrame
|
||||||
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
|
||||||
@@ -32,7 +32,7 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
async def start_fetch_weather(function_name, llm, context):
|
async def start_fetch_weather(function_name, llm, context):
|
||||||
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
"""Push a frame to the LLM; this is handy when the LLM response might take a while."""
|
||||||
await llm.push_frame(TextFrame("Let me check on that."))
|
await llm.push_frame(TTSSpeakFrame("Let me check on that."))
|
||||||
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
logger.debug(f"Starting fetch_weather_from_api with function_name: {function_name}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -169,8 +169,7 @@ class OutputGate(FrameProcessor):
|
|||||||
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
||||||
|
|
||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
self._gate_task.cancel()
|
await self.cancel_task(self._gate_task)
|
||||||
await self._gate_task
|
|
||||||
|
|
||||||
async def _gate_task_handler(self):
|
async def _gate_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
|
|||||||
@@ -101,12 +101,12 @@ HIGH PRIORITY SIGNALS:
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Complete Wh-question
|
# Complete Wh-question
|
||||||
[{"role": "assistant", "content": "I can help you learn."},
|
[{"role": "assistant", "content": "I can help you learn."},
|
||||||
{"role": "user", "content": "What's the fastest way to learn Spanish"}]
|
{"role": "user", "content": "What's the fastest way to learn Spanish"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Complete Yes/No question despite STT error
|
# Complete Yes/No question despite STT error
|
||||||
[{"role": "assistant", "content": "I know about planets."},
|
[{"role": "assistant", "content": "I know about planets."},
|
||||||
{"role": "user", "content": "Is is Jupiter the biggest planet"}]
|
{"role": "user", "content": "Is is Jupiter the biggest planet"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -118,12 +118,12 @@ Output: YES
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Direct instruction
|
# Direct instruction
|
||||||
[{"role": "assistant", "content": "I can explain many topics."},
|
[{"role": "assistant", "content": "I can explain many topics."},
|
||||||
{"role": "user", "content": "Tell me about black holes"}]
|
{"role": "user", "content": "Tell me about black holes"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Action demand
|
# Action demand
|
||||||
[{"role": "assistant", "content": "I can help with math."},
|
[{"role": "assistant", "content": "I can help with math."},
|
||||||
{"role": "user", "content": "Solve this equation x plus 5 equals 12"}]
|
{"role": "user", "content": "Solve this equation x plus 5 equals 12"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -134,12 +134,12 @@ Output: YES
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Specific answer
|
# Specific answer
|
||||||
[{"role": "assistant", "content": "What's your favorite color?"},
|
[{"role": "assistant", "content": "What's your favorite color?"},
|
||||||
{"role": "user", "content": "I really like blue"}]
|
{"role": "user", "content": "I really like blue"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Option selection
|
# Option selection
|
||||||
[{"role": "assistant", "content": "Would you prefer morning or evening?"},
|
[{"role": "assistant", "content": "Would you prefer morning or evening?"},
|
||||||
{"role": "user", "content": "Morning"}]
|
{"role": "user", "content": "Morning"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -153,17 +153,17 @@ MEDIUM PRIORITY SIGNALS:
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Self-correction reaching completion
|
# Self-correction reaching completion
|
||||||
[{"role": "assistant", "content": "What would you like to know?"},
|
[{"role": "assistant", "content": "What would you like to know?"},
|
||||||
{"role": "user", "content": "Tell me about... no wait, explain how rainbows form"}]
|
{"role": "user", "content": "Tell me about... no wait, explain how rainbows form"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Topic change with complete thought
|
# Topic change with complete thought
|
||||||
[{"role": "assistant", "content": "The weather is nice today."},
|
[{"role": "assistant", "content": "The weather is nice today."},
|
||||||
{"role": "user", "content": "Actually can you tell me who invented the telephone"}]
|
{"role": "user", "content": "Actually can you tell me who invented the telephone"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Mid-sentence completion
|
# Mid-sentence completion
|
||||||
[{"role": "assistant", "content": "Hello I'm ready."},
|
[{"role": "assistant", "content": "Hello I'm ready."},
|
||||||
{"role": "user", "content": "What's the capital of? France"}]
|
{"role": "user", "content": "What's the capital of? France"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -175,12 +175,12 @@ Output: YES
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Acknowledgment
|
# Acknowledgment
|
||||||
[{"role": "assistant", "content": "Should we talk about history?"},
|
[{"role": "assistant", "content": "Should we talk about history?"},
|
||||||
{"role": "user", "content": "Sure"}]
|
{"role": "user", "content": "Sure"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Disagreement with completion
|
# Disagreement with completion
|
||||||
[{"role": "assistant", "content": "Is that what you meant?"},
|
[{"role": "assistant", "content": "Is that what you meant?"},
|
||||||
{"role": "user", "content": "No not really"}]
|
{"role": "user", "content": "No not really"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -194,12 +194,12 @@ LOW PRIORITY SIGNALS:
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Word repetition but complete
|
# Word repetition but complete
|
||||||
[{"role": "assistant", "content": "I can help with that."},
|
[{"role": "assistant", "content": "I can help with that."},
|
||||||
{"role": "user", "content": "What what is the time right now"}]
|
{"role": "user", "content": "What what is the time right now"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Missing punctuation but complete
|
# Missing punctuation but complete
|
||||||
[{"role": "assistant", "content": "I can explain that."},
|
[{"role": "assistant", "content": "I can explain that."},
|
||||||
{"role": "user", "content": "Please tell me how computers work"}]
|
{"role": "user", "content": "Please tell me how computers work"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
@@ -211,12 +211,12 @@ Output: YES
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Filler words but complete
|
# Filler words but complete
|
||||||
[{"role": "assistant", "content": "What would you like to know?"},
|
[{"role": "assistant", "content": "What would you like to know?"},
|
||||||
{"role": "user", "content": "Um uh how do airplanes fly"}]
|
{"role": "user", "content": "Um uh how do airplanes fly"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Thinking pause but incomplete
|
# Thinking pause but incomplete
|
||||||
[{"role": "assistant", "content": "I can explain anything."},
|
[{"role": "assistant", "content": "I can explain anything."},
|
||||||
{"role": "user", "content": "Well um I want to know about the"}]
|
{"role": "user", "content": "Well um I want to know about the"}]
|
||||||
Output: NO
|
Output: NO
|
||||||
|
|
||||||
@@ -241,17 +241,17 @@ DECISION RULES:
|
|||||||
|
|
||||||
Examples:
|
Examples:
|
||||||
# Incomplete despite corrections
|
# Incomplete despite corrections
|
||||||
[{"role": "assistant", "content": "What would you like to know about?"},
|
[{"role": "assistant", "content": "What would you like to know about?"},
|
||||||
{"role": "user", "content": "Can you tell me about"}]
|
{"role": "user", "content": "Can you tell me about"}]
|
||||||
Output: NO
|
Output: NO
|
||||||
|
|
||||||
# Complete despite multiple artifacts
|
# Complete despite multiple artifacts
|
||||||
[{"role": "assistant", "content": "I can help you learn."},
|
[{"role": "assistant", "content": "I can help you learn."},
|
||||||
{"role": "user", "content": "How do you I mean what's the best way to learn programming"}]
|
{"role": "user", "content": "How do you I mean what's the best way to learn programming"}]
|
||||||
Output: YES
|
Output: YES
|
||||||
|
|
||||||
# Trailing off incomplete
|
# Trailing off incomplete
|
||||||
[{"role": "assistant", "content": "I can explain anything."},
|
[{"role": "assistant", "content": "I can explain anything."},
|
||||||
{"role": "user", "content": "I was wondering if you could tell me why"}]
|
{"role": "user", "content": "I was wondering if you could tell me why"}]
|
||||||
Output: NO
|
Output: NO
|
||||||
"""
|
"""
|
||||||
@@ -374,8 +374,7 @@ class OutputGate(FrameProcessor):
|
|||||||
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
||||||
|
|
||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
self._gate_task.cancel()
|
await self.cancel_task(self._gate_task)
|
||||||
await self._gate_task
|
|
||||||
|
|
||||||
async def _gate_task_handler(self):
|
async def _gate_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
|
|||||||
@@ -44,9 +44,7 @@ from pipecat.processors.aggregators.openai_llm_context import (
|
|||||||
)
|
)
|
||||||
from pipecat.processors.filters.function_filter import FunctionFilter
|
from pipecat.processors.filters.function_filter import FunctionFilter
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.user_idle_processor import UserIdleProcessor
|
|
||||||
from pipecat.services.cartesia import CartesiaTTSService
|
from pipecat.services.cartesia import CartesiaTTSService
|
||||||
from pipecat.services.deepgram import DeepgramSTTService
|
|
||||||
from pipecat.services.google import GoogleLLMContext, GoogleLLMService
|
from pipecat.services.google import GoogleLLMContext, GoogleLLMService
|
||||||
from pipecat.sync.base_notifier import BaseNotifier
|
from pipecat.sync.base_notifier import BaseNotifier
|
||||||
from pipecat.sync.event_notifier import EventNotifier
|
from pipecat.sync.event_notifier import EventNotifier
|
||||||
@@ -440,11 +438,11 @@ class CompletenessCheck(FrameProcessor):
|
|||||||
|
|
||||||
if isinstance(frame, UserStartedSpeakingFrame):
|
if isinstance(frame, UserStartedSpeakingFrame):
|
||||||
if self._idle_task:
|
if self._idle_task:
|
||||||
self._idle_task.cancel()
|
await self.cancel_task(self._idle_task)
|
||||||
elif isinstance(frame, TextFrame) and frame.text.startswith("YES"):
|
elif isinstance(frame, TextFrame) and frame.text.startswith("YES"):
|
||||||
logger.debug("Completeness check YES")
|
logger.debug("Completeness check YES")
|
||||||
if self._idle_task:
|
if self._idle_task:
|
||||||
self._idle_task.cancel()
|
await self.cancel_task(self._idle_task)
|
||||||
await self.push_frame(UserStoppedSpeakingFrame())
|
await self.push_frame(UserStoppedSpeakingFrame())
|
||||||
await self._audio_accumulator.reset()
|
await self._audio_accumulator.reset()
|
||||||
await self._notifier.notify()
|
await self._notifier.notify()
|
||||||
@@ -602,8 +600,7 @@ class OutputGate(FrameProcessor):
|
|||||||
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
||||||
|
|
||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
self._gate_task.cancel()
|
await self.cancel_task(self._gate_task)
|
||||||
await self._gate_task
|
|
||||||
|
|
||||||
async def _gate_task_handler(self):
|
async def _gate_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from runner import configure
|
|||||||
|
|
||||||
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
from pipecat.audio.vad.silero import SileroVADAnalyzer
|
||||||
from pipecat.audio.vad.vad_analyzer import VADParams
|
from pipecat.audio.vad.vad_analyzer import VADParams
|
||||||
|
from pipecat.frames.frames import LLMMessagesAppendFrame
|
||||||
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
|
||||||
@@ -71,6 +72,21 @@ async def main():
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@transport.event_handler("on_first_participant_joined")
|
||||||
|
async def on_first_participant_joined(transport, participant):
|
||||||
|
await task.queue_frames(
|
||||||
|
[
|
||||||
|
LLMMessagesAppendFrame(
|
||||||
|
messages=[
|
||||||
|
{
|
||||||
|
"role": "assistant",
|
||||||
|
"content": "Greet the user.",
|
||||||
|
}
|
||||||
|
]
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|||||||
@@ -159,5 +159,5 @@ class SileroVADAnalyzer(VADAnalyzer):
|
|||||||
return new_confidence
|
return new_confidence
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# This comes from an empty audio array
|
# This comes from an empty audio array
|
||||||
logger.exception(f"Error analyzing audio with Silero VAD: {e}")
|
logger.error(f"Error analyzing audio with Silero VAD: {e}")
|
||||||
return 0
|
return 0
|
||||||
|
|||||||
@@ -11,6 +11,18 @@ from pipecat.frames.frames import Frame
|
|||||||
|
|
||||||
|
|
||||||
class BaseTask(ABC):
|
class BaseTask(ABC):
|
||||||
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def id(self) -> int:
|
||||||
|
"""Returns the unique indetifier for this task."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def name(self) -> str:
|
||||||
|
"""Returns the name of this task."""
|
||||||
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def has_finished(self) -> bool:
|
def has_finished(self) -> bool:
|
||||||
"""Indicates whether the tasks has finished. That is, all processors
|
"""Indicates whether the tasks has finished. That is, all processors
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from pipecat.pipeline.pipeline import Pipeline
|
|||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
|
||||||
|
|
||||||
class Source(FrameProcessor):
|
class ParallelPipelineSource(FrameProcessor):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
upstream_queue: asyncio.Queue,
|
upstream_queue: asyncio.Queue,
|
||||||
@@ -46,7 +46,7 @@ class Source(FrameProcessor):
|
|||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
class Sink(FrameProcessor):
|
class ParallelPipelineSink(FrameProcessor):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
downstream_queue: asyncio.Queue,
|
downstream_queue: asyncio.Queue,
|
||||||
@@ -92,8 +92,8 @@ class ParallelPipeline(BasePipeline):
|
|||||||
raise TypeError(f"ParallelPipeline argument {processors} is not a list")
|
raise TypeError(f"ParallelPipeline argument {processors} is not a list")
|
||||||
|
|
||||||
# We will add a source before the pipeline and a sink after.
|
# We will add a source before the pipeline and a sink after.
|
||||||
source = Source(self._up_queue, self._parallel_push_frame)
|
source = ParallelPipelineSource(self._up_queue, self._parallel_push_frame)
|
||||||
sink = Sink(self._down_queue, self._parallel_push_frame)
|
sink = ParallelPipelineSink(self._down_queue, self._parallel_push_frame)
|
||||||
self._sources.append(source)
|
self._sources.append(source)
|
||||||
self._sinks.append(sink)
|
self._sinks.append(sink)
|
||||||
|
|
||||||
@@ -117,6 +117,7 @@ class ParallelPipeline(BasePipeline):
|
|||||||
#
|
#
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
|
await super().cleanup()
|
||||||
await asyncio.gather(*[s.cleanup() for s in self._sources])
|
await asyncio.gather(*[s.cleanup() for s in self._sources])
|
||||||
await asyncio.gather(*[p.cleanup() for p in self._pipelines])
|
await asyncio.gather(*[p.cleanup() for p in self._pipelines])
|
||||||
await asyncio.gather(*[s.cleanup() for s in self._sinks])
|
await asyncio.gather(*[s.cleanup() for s in self._sinks])
|
||||||
@@ -150,22 +151,18 @@ class ParallelPipeline(BasePipeline):
|
|||||||
|
|
||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
# The up task doesn't receive an EndFrame, so we just cancel it.
|
# The up task doesn't receive an EndFrame, so we just cancel it.
|
||||||
self._up_task.cancel()
|
await self.cancel_task(self._up_task)
|
||||||
await self._up_task
|
# The down tasks waits for the last EndFrame sent by the internal
|
||||||
# The down tasks waits for the last EndFrame send by the internal
|
|
||||||
# pipelines.
|
# pipelines.
|
||||||
await self._down_task
|
await self._down_task
|
||||||
|
|
||||||
async def _cancel(self):
|
async def _cancel(self):
|
||||||
self._up_task.cancel()
|
await self.cancel_task(self._up_task)
|
||||||
await self._up_task
|
await self.cancel_task(self._down_task)
|
||||||
self._down_task.cancel()
|
|
||||||
await self._down_task
|
|
||||||
|
|
||||||
async def _create_tasks(self):
|
async def _create_tasks(self):
|
||||||
loop = self.get_event_loop()
|
self._up_task = self.create_task(self._process_up_queue())
|
||||||
self._up_task = loop.create_task(self._process_up_queue())
|
self._down_task = self.create_task(self._process_down_queue())
|
||||||
self._down_task = loop.create_task(self._process_down_queue())
|
|
||||||
|
|
||||||
async def _drain_queues(self):
|
async def _drain_queues(self):
|
||||||
while not self._up_queue.empty:
|
while not self._up_queue.empty:
|
||||||
@@ -185,32 +182,26 @@ class ParallelPipeline(BasePipeline):
|
|||||||
|
|
||||||
async def _process_up_queue(self):
|
async def _process_up_queue(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame = await self._up_queue.get()
|
||||||
frame = await self._up_queue.get()
|
await self._parallel_push_frame(frame, FrameDirection.UPSTREAM)
|
||||||
await self._parallel_push_frame(frame, FrameDirection.UPSTREAM)
|
self._up_queue.task_done()
|
||||||
self._up_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _process_down_queue(self):
|
async def _process_down_queue(self):
|
||||||
running = True
|
running = True
|
||||||
while running:
|
while running:
|
||||||
try:
|
frame = await self._down_queue.get()
|
||||||
frame = await self._down_queue.get()
|
|
||||||
|
|
||||||
endframe_counter = self._endframe_counter.get(frame.id, 0)
|
endframe_counter = self._endframe_counter.get(frame.id, 0)
|
||||||
|
|
||||||
# If we have a counter, decrement it.
|
# If we have a counter, decrement it.
|
||||||
if endframe_counter > 0:
|
if endframe_counter > 0:
|
||||||
self._endframe_counter[frame.id] -= 1
|
self._endframe_counter[frame.id] -= 1
|
||||||
endframe_counter = self._endframe_counter[frame.id]
|
endframe_counter = self._endframe_counter[frame.id]
|
||||||
|
|
||||||
# If we don't have a counter or we reached 0, push the frame.
|
# If we don't have a counter or we reached 0, push the frame.
|
||||||
if endframe_counter == 0:
|
if endframe_counter == 0:
|
||||||
await self._parallel_push_frame(frame, FrameDirection.DOWNSTREAM)
|
await self._parallel_push_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
running = not (endframe_counter == 0 and isinstance(frame, EndFrame))
|
running = not (endframe_counter == 0 and isinstance(frame, EndFrame))
|
||||||
|
|
||||||
self._down_queue.task_done()
|
self._down_queue.task_done()
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ class Pipeline(BasePipeline):
|
|||||||
#
|
#
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
|
await super().cleanup()
|
||||||
await self._cleanup_processors()
|
await self._cleanup_processors()
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import signal
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
|
from pipecat.utils.asyncio import current_tasks
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
|
|
||||||
@@ -19,6 +20,7 @@ class PipelineRunner:
|
|||||||
self.name: str = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
self.name: str = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
|
|
||||||
self._tasks = {}
|
self._tasks = {}
|
||||||
|
self._sig_task = None
|
||||||
|
|
||||||
if handle_sigint:
|
if handle_sigint:
|
||||||
self._setup_sigint()
|
self._setup_sigint()
|
||||||
@@ -28,6 +30,11 @@ class PipelineRunner:
|
|||||||
self._tasks[task.name] = task
|
self._tasks[task.name] = task
|
||||||
await task.run()
|
await task.run()
|
||||||
del self._tasks[task.name]
|
del self._tasks[task.name]
|
||||||
|
# If we are cancelling through a signal, make sure we wait for it so
|
||||||
|
# everything gets cleaned up nicely.
|
||||||
|
if self._sig_task:
|
||||||
|
await self._sig_task
|
||||||
|
self._print_dangling_tasks()
|
||||||
logger.debug(f"Runner {self} finished running {task}")
|
logger.debug(f"Runner {self} finished running {task}")
|
||||||
|
|
||||||
async def stop_when_done(self):
|
async def stop_when_done(self):
|
||||||
@@ -40,16 +47,21 @@ class PipelineRunner:
|
|||||||
|
|
||||||
def _setup_sigint(self):
|
def _setup_sigint(self):
|
||||||
loop = asyncio.get_running_loop()
|
loop = asyncio.get_running_loop()
|
||||||
loop.add_signal_handler(
|
loop.add_signal_handler(signal.SIGINT, lambda *args: self._sig_handler())
|
||||||
signal.SIGINT, lambda *args: asyncio.create_task(self._sig_handler())
|
loop.add_signal_handler(signal.SIGTERM, lambda *args: self._sig_handler())
|
||||||
)
|
|
||||||
loop.add_signal_handler(
|
|
||||||
signal.SIGTERM, lambda *args: asyncio.create_task(self._sig_handler())
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _sig_handler(self):
|
def _sig_handler(self):
|
||||||
|
if not self._sig_task:
|
||||||
|
self._sig_task = asyncio.create_task(self._sig_cancel())
|
||||||
|
|
||||||
|
async def _sig_cancel(self):
|
||||||
logger.warning(f"Interruption detected. Canceling runner {self}")
|
logger.warning(f"Interruption detected. Canceling runner {self}")
|
||||||
await self.cancel()
|
await self.cancel()
|
||||||
|
|
||||||
|
def _print_dangling_tasks(self):
|
||||||
|
tasks = [t.get_name() for t in current_tasks()]
|
||||||
|
if tasks:
|
||||||
|
logger.warning(f"Dangling tasks detected: {tasks}")
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return self.name
|
return self.name
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ class SyncFrame(ControlFrame):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
class Source(FrameProcessor):
|
class SyncParallelPipelineSource(FrameProcessor):
|
||||||
def __init__(self, upstream_queue: asyncio.Queue):
|
def __init__(self, upstream_queue: asyncio.Queue):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._up_queue = upstream_queue
|
self._up_queue = upstream_queue
|
||||||
@@ -39,7 +39,7 @@ class Source(FrameProcessor):
|
|||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
class Sink(FrameProcessor):
|
class SyncParallelPipelineSink(FrameProcessor):
|
||||||
def __init__(self, downstream_queue: asyncio.Queue):
|
def __init__(self, downstream_queue: asyncio.Queue):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._down_queue = downstream_queue
|
self._down_queue = downstream_queue
|
||||||
@@ -76,8 +76,8 @@ class SyncParallelPipeline(BasePipeline):
|
|||||||
# We add a source at the beginning of the pipeline and a sink at the end.
|
# We add a source at the beginning of the pipeline and a sink at the end.
|
||||||
up_queue = asyncio.Queue()
|
up_queue = asyncio.Queue()
|
||||||
down_queue = asyncio.Queue()
|
down_queue = asyncio.Queue()
|
||||||
source = Source(up_queue)
|
source = SyncParallelPipelineSource(up_queue)
|
||||||
sink = Sink(down_queue)
|
sink = SyncParallelPipelineSink(down_queue)
|
||||||
processors: List[FrameProcessor] = [source] + processors + [sink]
|
processors: List[FrameProcessor] = [source] + processors + [sink]
|
||||||
|
|
||||||
# Keep track of sources and sinks. We also keep the output queue of
|
# Keep track of sources and sinks. We also keep the output queue of
|
||||||
@@ -101,6 +101,10 @@ class SyncParallelPipeline(BasePipeline):
|
|||||||
# Frame processor
|
# Frame processor
|
||||||
#
|
#
|
||||||
|
|
||||||
|
async def cleanup(self):
|
||||||
|
await super().cleanup()
|
||||||
|
await asyncio.gather(*[p.cleanup() for p in self._pipelines])
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ from pipecat.pipeline.base_pipeline import BasePipeline
|
|||||||
from pipecat.pipeline.base_task import BaseTask
|
from pipecat.pipeline.base_task import BaseTask
|
||||||
from pipecat.pipeline.task_observer import TaskObserver
|
from pipecat.pipeline.task_observer import TaskObserver
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
from pipecat.utils.asyncio import cancel_task, create_task, wait_for_task
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
HEARTBEAT_SECONDS = 1.0
|
HEARTBEAT_SECONDS = 1.0
|
||||||
@@ -49,7 +50,7 @@ class PipelineParams(BaseModel):
|
|||||||
heartbeats_period_secs: float = HEARTBEAT_SECONDS
|
heartbeats_period_secs: float = HEARTBEAT_SECONDS
|
||||||
|
|
||||||
|
|
||||||
class Source(FrameProcessor):
|
class PipelineTaskSource(FrameProcessor):
|
||||||
"""This is the source processor that is linked at the beginning of the
|
"""This is the source processor that is linked at the beginning of the
|
||||||
pipeline given to the pipeline task. It allows us to easily push frames
|
pipeline given to the pipeline task. It allows us to easily push frames
|
||||||
downstream to the pipeline and also receive upstream frames coming from the
|
downstream to the pipeline and also receive upstream frames coming from the
|
||||||
@@ -57,8 +58,8 @@ class Source(FrameProcessor):
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, up_queue: asyncio.Queue):
|
def __init__(self, up_queue: asyncio.Queue, **kwargs):
|
||||||
super().__init__()
|
super().__init__(**kwargs)
|
||||||
self._up_queue = up_queue
|
self._up_queue = up_queue
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
@@ -71,15 +72,15 @@ class Source(FrameProcessor):
|
|||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
class Sink(FrameProcessor):
|
class PipelineTaskSink(FrameProcessor):
|
||||||
"""This is the sink processor that is linked at the end of the pipeline
|
"""This is the sink processor that is linked at the end of the pipeline
|
||||||
given to the pipeline task. It allows us to receive downstream frames and
|
given to the pipeline task. It allows us to receive downstream frames and
|
||||||
act on them, for example, waiting to receive an EndFrame.
|
act on them, for example, waiting to receive an EndFrame.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, down_queue: asyncio.Queue):
|
def __init__(self, down_queue: asyncio.Queue, **kwargs):
|
||||||
super().__init__()
|
super().__init__(**kwargs)
|
||||||
self._down_queue = down_queue
|
self._down_queue = down_queue
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
@@ -94,8 +95,8 @@ class PipelineTask(BaseTask):
|
|||||||
params: PipelineParams = PipelineParams(),
|
params: PipelineParams = PipelineParams(),
|
||||||
clock: BaseClock = SystemClock(),
|
clock: BaseClock = SystemClock(),
|
||||||
):
|
):
|
||||||
self.id: int = obj_id()
|
self._id: int = obj_id()
|
||||||
self.name: str = f"{self.__class__.__name__}#{obj_count(self)}"
|
self._name: str = f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
|
|
||||||
self._pipeline = pipeline
|
self._pipeline = pipeline
|
||||||
self._clock = clock
|
self._clock = clock
|
||||||
@@ -115,14 +116,24 @@ class PipelineTask(BaseTask):
|
|||||||
# down queue.
|
# down queue.
|
||||||
self._endframe_event = asyncio.Event()
|
self._endframe_event = asyncio.Event()
|
||||||
|
|
||||||
self._source = Source(self._up_queue)
|
self._source = PipelineTaskSource(self._up_queue)
|
||||||
self._source.link(pipeline)
|
self._source.link(pipeline)
|
||||||
|
|
||||||
self._sink = Sink(self._down_queue)
|
self._sink = PipelineTaskSink(self._down_queue)
|
||||||
pipeline.link(self._sink)
|
pipeline.link(self._sink)
|
||||||
|
|
||||||
self._observer = TaskObserver(params.observers)
|
self._observer = TaskObserver(params.observers)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> int:
|
||||||
|
"""Returns the unique indetifier for this task."""
|
||||||
|
return self._id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
"""Returns the name of this task."""
|
||||||
|
return self._name
|
||||||
|
|
||||||
def has_finished(self) -> bool:
|
def has_finished(self) -> bool:
|
||||||
"""Indicates whether the tasks has finished. That is, all processors
|
"""Indicates whether the tasks has finished. That is, all processors
|
||||||
have stopped.
|
have stopped.
|
||||||
@@ -147,14 +158,24 @@ class PipelineTask(BaseTask):
|
|||||||
# out-of-band from the main streaming task which is what we want since
|
# out-of-band from the main streaming task which is what we want since
|
||||||
# we want to cancel right away.
|
# we want to cancel right away.
|
||||||
await self._source.push_frame(CancelFrame())
|
await self._source.push_frame(CancelFrame())
|
||||||
await self._cancel_tasks(True)
|
# Only cancel the push task. Everything else will be cancelled in run().
|
||||||
|
await cancel_task(self._process_push_task)
|
||||||
|
await self._cleanup()
|
||||||
|
|
||||||
async def run(self):
|
async def run(self):
|
||||||
"""
|
"""
|
||||||
Starts running the given pipeline.
|
Starts running the given pipeline.
|
||||||
"""
|
"""
|
||||||
tasks = self._create_tasks()
|
try:
|
||||||
await asyncio.gather(*tasks)
|
push_task = self._create_tasks()
|
||||||
|
await wait_for_task(push_task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# We are awaiting on the push task and it might be cancelled
|
||||||
|
# (e.g. Ctrl-C). This means we will get a CancelledError here as
|
||||||
|
# well, because you get a CancelledError in every place you are
|
||||||
|
# awaiting a task.
|
||||||
|
pass
|
||||||
|
await self._cancel_tasks()
|
||||||
self._finished = True
|
self._finished = True
|
||||||
|
|
||||||
async def queue_frame(self, frame: Frame):
|
async def queue_frame(self, frame: Frame):
|
||||||
@@ -175,41 +196,41 @@ class PipelineTask(BaseTask):
|
|||||||
await self.queue_frame(frame)
|
await self.queue_frame(frame)
|
||||||
|
|
||||||
def _create_tasks(self):
|
def _create_tasks(self):
|
||||||
tasks = []
|
loop = asyncio.get_running_loop()
|
||||||
self._process_up_task = asyncio.create_task(self._process_up_queue())
|
self._process_up_task = create_task(
|
||||||
self._process_down_task = asyncio.create_task(self._process_down_queue())
|
loop, self._process_up_queue(), f"{self}::_process_up_queue"
|
||||||
self._process_push_task = asyncio.create_task(self._process_push_queue())
|
)
|
||||||
|
self._process_down_task = create_task(
|
||||||
|
loop, self._process_down_queue(), f"{self}::_process_down_queue"
|
||||||
|
)
|
||||||
|
self._process_push_task = create_task(
|
||||||
|
loop, self._process_push_queue(), f"{self}::_process_push_queue"
|
||||||
|
)
|
||||||
|
|
||||||
tasks = [self._process_up_task, self._process_down_task, self._process_push_task]
|
return self._process_push_task
|
||||||
|
|
||||||
return tasks
|
|
||||||
|
|
||||||
def _maybe_start_heartbeat_tasks(self):
|
def _maybe_start_heartbeat_tasks(self):
|
||||||
if self._params.enable_heartbeats:
|
if self._params.enable_heartbeats:
|
||||||
self._heartbeat_push_task = asyncio.create_task(self._heartbeat_push_handler())
|
loop = asyncio.get_running_loop()
|
||||||
self._heartbeat_monitor_task = asyncio.create_task(self._heartbeat_monitor_handler())
|
self._heartbeat_push_task = create_task(
|
||||||
|
loop, self._heartbeat_push_handler(), f"{self}::_heartbeat_push_handler"
|
||||||
|
)
|
||||||
|
self._heartbeat_monitor_task = create_task(
|
||||||
|
loop, self._heartbeat_monitor_handler(), f"{self}::_heartbeat_monitor_handler"
|
||||||
|
)
|
||||||
|
|
||||||
async def _cancel_tasks(self, cancel_push: bool):
|
async def _cancel_tasks(self):
|
||||||
await self._maybe_cancel_heartbeat_tasks()
|
await self._maybe_cancel_heartbeat_tasks()
|
||||||
|
|
||||||
if cancel_push:
|
await cancel_task(self._process_up_task)
|
||||||
self._process_push_task.cancel()
|
await cancel_task(self._process_down_task)
|
||||||
await self._process_push_task
|
|
||||||
|
|
||||||
self._process_up_task.cancel()
|
|
||||||
await self._process_up_task
|
|
||||||
|
|
||||||
self._process_down_task.cancel()
|
|
||||||
await self._process_down_task
|
|
||||||
|
|
||||||
await self._observer.stop()
|
await self._observer.stop()
|
||||||
|
|
||||||
async def _maybe_cancel_heartbeat_tasks(self):
|
async def _maybe_cancel_heartbeat_tasks(self):
|
||||||
if self._params.enable_heartbeats:
|
if self._params.enable_heartbeats:
|
||||||
self._heartbeat_push_task.cancel()
|
await cancel_task(self._heartbeat_push_task)
|
||||||
await self._heartbeat_push_task
|
await cancel_task(self._heartbeat_monitor_task)
|
||||||
self._heartbeat_monitor_task.cancel()
|
|
||||||
await self._heartbeat_monitor_task
|
|
||||||
|
|
||||||
def _initial_metrics_frame(self) -> MetricsFrame:
|
def _initial_metrics_frame(self) -> MetricsFrame:
|
||||||
processors = self._pipeline.processors_with_metrics()
|
processors = self._pipeline.processors_with_metrics()
|
||||||
@@ -223,6 +244,11 @@ class PipelineTask(BaseTask):
|
|||||||
await self._endframe_event.wait()
|
await self._endframe_event.wait()
|
||||||
self._endframe_event.clear()
|
self._endframe_event.clear()
|
||||||
|
|
||||||
|
async def _cleanup(self):
|
||||||
|
await self._source.cleanup()
|
||||||
|
await self._pipeline.cleanup()
|
||||||
|
await self._sink.cleanup()
|
||||||
|
|
||||||
async def _process_push_queue(self):
|
async def _process_push_queue(self):
|
||||||
"""This is the task that runs the pipeline for the first time by sending
|
"""This is the task that runs the pipeline for the first time by sending
|
||||||
a StartFrame and by pushing any other frames queued by the user. It runs
|
a StartFrame and by pushing any other frames queued by the user. It runs
|
||||||
@@ -249,24 +275,16 @@ class PipelineTask(BaseTask):
|
|||||||
running = True
|
running = True
|
||||||
should_cleanup = True
|
should_cleanup = True
|
||||||
while running:
|
while running:
|
||||||
try:
|
frame = await self._push_queue.get()
|
||||||
frame = await self._push_queue.get()
|
await self._source.queue_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
await self._source.queue_frame(frame, FrameDirection.DOWNSTREAM)
|
if isinstance(frame, EndFrame):
|
||||||
if isinstance(frame, EndFrame):
|
await self._wait_for_endframe()
|
||||||
await self._wait_for_endframe()
|
running = not isinstance(frame, (StopTaskFrame, EndFrame))
|
||||||
running = not isinstance(frame, (StopTaskFrame, EndFrame))
|
should_cleanup = not isinstance(frame, StopTaskFrame)
|
||||||
should_cleanup = not isinstance(frame, StopTaskFrame)
|
self._push_queue.task_done()
|
||||||
self._push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
# Cleanup only if we need to.
|
# Cleanup only if we need to.
|
||||||
if should_cleanup:
|
if should_cleanup:
|
||||||
await self._source.cleanup()
|
await self._cleanup()
|
||||||
await self._pipeline.cleanup()
|
|
||||||
await self._sink.cleanup()
|
|
||||||
# Finally, cancel internal tasks. We don't cancel the push tasks because
|
|
||||||
# that's us.
|
|
||||||
await self._cancel_tasks(False)
|
|
||||||
|
|
||||||
async def _process_up_queue(self):
|
async def _process_up_queue(self):
|
||||||
"""This is the task that processes frames coming upstream from the
|
"""This is the task that processes frames coming upstream from the
|
||||||
@@ -276,26 +294,23 @@ class PipelineTask(BaseTask):
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame = await self._up_queue.get()
|
||||||
frame = await self._up_queue.get()
|
if isinstance(frame, EndTaskFrame):
|
||||||
if isinstance(frame, EndTaskFrame):
|
# Tell the task we should end nicely.
|
||||||
# Tell the task we should end nicely.
|
await self.queue_frame(EndFrame())
|
||||||
await self.queue_frame(EndFrame())
|
elif isinstance(frame, CancelTaskFrame):
|
||||||
elif isinstance(frame, CancelTaskFrame):
|
# Tell the task we should end right away.
|
||||||
# Tell the task we should end right away.
|
await self.queue_frame(CancelFrame())
|
||||||
|
elif isinstance(frame, StopTaskFrame):
|
||||||
|
await self.queue_frame(StopTaskFrame())
|
||||||
|
elif isinstance(frame, ErrorFrame):
|
||||||
|
logger.error(f"Error running app: {frame}")
|
||||||
|
if frame.fatal:
|
||||||
|
# Cancel all tasks downstream.
|
||||||
await self.queue_frame(CancelFrame())
|
await self.queue_frame(CancelFrame())
|
||||||
elif isinstance(frame, StopTaskFrame):
|
# Tell the task we should stop.
|
||||||
await self.queue_frame(StopTaskFrame())
|
await self.queue_frame(StopTaskFrame())
|
||||||
elif isinstance(frame, ErrorFrame):
|
self._up_queue.task_done()
|
||||||
logger.error(f"Error running app: {frame}")
|
|
||||||
if frame.fatal:
|
|
||||||
# Cancel all tasks downstream.
|
|
||||||
await self.queue_frame(CancelFrame())
|
|
||||||
# Tell the task we should stop.
|
|
||||||
await self.queue_frame(StopTaskFrame())
|
|
||||||
self._up_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _process_down_queue(self):
|
async def _process_down_queue(self):
|
||||||
"""This tasks process frames coming downstream from the pipeline. For
|
"""This tasks process frames coming downstream from the pipeline. For
|
||||||
@@ -305,29 +320,23 @@ class PipelineTask(BaseTask):
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame = await self._down_queue.get()
|
||||||
frame = await self._down_queue.get()
|
if isinstance(frame, EndFrame):
|
||||||
if isinstance(frame, EndFrame):
|
self._endframe_event.set()
|
||||||
self._endframe_event.set()
|
elif isinstance(frame, HeartbeatFrame):
|
||||||
elif isinstance(frame, HeartbeatFrame):
|
await self._heartbeat_queue.put(frame)
|
||||||
await self._heartbeat_queue.put(frame)
|
self._down_queue.task_done()
|
||||||
self._down_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _heartbeat_push_handler(self):
|
async def _heartbeat_push_handler(self):
|
||||||
"""
|
"""
|
||||||
This tasks pushes a heartbeat frame every heartbeat period.
|
This tasks pushes a heartbeat frame every heartbeat period.
|
||||||
"""
|
"""
|
||||||
while True:
|
while True:
|
||||||
try:
|
# Don't use `queue_frame()` because if an EndFrame is queued the
|
||||||
# Don't use `queue_frame()` because if an EndFrame is queued the
|
# task will just stop waiting for the pipeline to finish not
|
||||||
# task will just stop waiting for the pipeline to finish not
|
# allowing more frames to be pushed.
|
||||||
# allowing more frames to be pushed.
|
await self._source.queue_frame(HeartbeatFrame(timestamp=self._clock.get_time()))
|
||||||
await self._source.queue_frame(HeartbeatFrame(timestamp=self._clock.get_time()))
|
await asyncio.sleep(self._params.heartbeats_period_secs)
|
||||||
await asyncio.sleep(self._params.heartbeats_period_secs)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _heartbeat_monitor_handler(self):
|
async def _heartbeat_monitor_handler(self):
|
||||||
"""This tasks monitors heartbeat frames. If a heartbeat frame has not
|
"""This tasks monitors heartbeat frames. If a heartbeat frame has not
|
||||||
@@ -347,8 +356,6 @@ class PipelineTask(BaseTask):
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
f"{self}: heartbeat frame not received for more than {wait_time} seconds"
|
f"{self}: heartbeat frame not received for more than {wait_time} seconds"
|
||||||
)
|
)
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return self.name
|
return self.name
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ from attr import dataclass
|
|||||||
from pipecat.frames.frames import Frame
|
from pipecat.frames.frames import Frame
|
||||||
from pipecat.observers.base_observer import BaseObserver
|
from pipecat.observers.base_observer import BaseObserver
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
from pipecat.utils.asyncio import cancel_task, create_task
|
||||||
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -54,13 +56,22 @@ class TaskObserver(BaseObserver):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, observers: List[BaseObserver] = []):
|
def __init__(self, observers: List[BaseObserver] = []):
|
||||||
|
self._id: int = obj_id()
|
||||||
|
self._name: str = f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
self._proxies: List[Proxy] = self._create_proxies(observers)
|
self._proxies: List[Proxy] = self._create_proxies(observers)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> int:
|
||||||
|
return self._id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
return self._name
|
||||||
|
|
||||||
async def stop(self):
|
async def stop(self):
|
||||||
"""Stops all proxy observer tasks."""
|
"""Stops all proxy observer tasks."""
|
||||||
for proxy in self._proxies:
|
for proxy in self._proxies:
|
||||||
proxy.task.cancel()
|
await cancel_task(proxy.task)
|
||||||
await proxy.task
|
|
||||||
|
|
||||||
async def on_push_frame(
|
async def on_push_frame(
|
||||||
self,
|
self,
|
||||||
@@ -79,19 +90,24 @@ class TaskObserver(BaseObserver):
|
|||||||
|
|
||||||
def _create_proxies(self, observers) -> List[Proxy]:
|
def _create_proxies(self, observers) -> List[Proxy]:
|
||||||
proxies = []
|
proxies = []
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
for observer in observers:
|
for observer in observers:
|
||||||
queue = asyncio.Queue()
|
queue = asyncio.Queue()
|
||||||
task = asyncio.create_task(self._proxy_task_handler(queue, observer))
|
task = create_task(
|
||||||
|
loop,
|
||||||
|
self._proxy_task_handler(queue, observer),
|
||||||
|
f"{self}::{observer.__class__.__name__}",
|
||||||
|
)
|
||||||
proxy = Proxy(queue=queue, task=task, observer=observer)
|
proxy = Proxy(queue=queue, task=task, observer=observer)
|
||||||
proxies.append(proxy)
|
proxies.append(proxy)
|
||||||
return proxies
|
return proxies
|
||||||
|
|
||||||
async def _proxy_task_handler(self, queue: asyncio.Queue, observer: BaseObserver):
|
async def _proxy_task_handler(self, queue: asyncio.Queue, observer: BaseObserver):
|
||||||
while True:
|
while True:
|
||||||
try:
|
data = await queue.get()
|
||||||
data = await queue.get()
|
await observer.on_push_frame(
|
||||||
await observer.on_push_frame(
|
data.src, data.dst, data.frame, data.direction, data.timestamp
|
||||||
data.src, data.dst, data.frame, data.direction, data.timestamp
|
)
|
||||||
)
|
|
||||||
except asyncio.CancelledError:
|
def __str__(self):
|
||||||
break
|
return self.name
|
||||||
|
|||||||
@@ -4,8 +4,6 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from pipecat.frames.frames import CancelFrame, EndFrame, Frame, StartFrame
|
from pipecat.frames.frames import CancelFrame, EndFrame, Frame, StartFrame
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
@@ -38,18 +36,14 @@ class GatedOpenAILLMContextAggregator(FrameProcessor):
|
|||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def _start(self):
|
async def _start(self):
|
||||||
self._gate_task = self.get_event_loop().create_task(self._gate_task_handler())
|
self._gate_task = self.create_task(self._gate_task_handler())
|
||||||
|
|
||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
self._gate_task.cancel()
|
await self.cancel_task(self._gate_task)
|
||||||
await self._gate_task
|
|
||||||
|
|
||||||
async def _gate_task_handler(self):
|
async def _gate_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
await self._notifier.wait()
|
||||||
await self._notifier.wait()
|
if self._last_context_frame:
|
||||||
if self._last_context_frame:
|
await self.push_frame(self._last_context_frame)
|
||||||
await self.push_frame(self._last_context_frame)
|
self._last_context_frame = None
|
||||||
self._last_context_frame = None
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|||||||
@@ -6,15 +6,15 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import inspect
|
import inspect
|
||||||
|
import sys
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Awaitable, Callable, Optional
|
from typing import Awaitable, Callable, Coroutine, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
from pipecat.clocks.base_clock import BaseClock
|
from pipecat.clocks.base_clock import BaseClock
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
@@ -24,6 +24,7 @@ from pipecat.frames.frames import (
|
|||||||
)
|
)
|
||||||
from pipecat.metrics.metrics import LLMTokenUsage, MetricsData
|
from pipecat.metrics.metrics import LLMTokenUsage, MetricsData
|
||||||
from pipecat.processors.metrics.frame_processor_metrics import FrameProcessorMetrics
|
from pipecat.processors.metrics.frame_processor_metrics import FrameProcessorMetrics
|
||||||
|
from pipecat.utils.asyncio import cancel_task, create_task, wait_for_task
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
|
|
||||||
@@ -41,8 +42,8 @@ class FrameProcessor:
|
|||||||
loop: asyncio.AbstractEventLoop | None = None,
|
loop: asyncio.AbstractEventLoop | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
self.id: int = obj_id()
|
self._id: int = obj_id()
|
||||||
self.name = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
self._name = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
self._parent: "FrameProcessor" | None = None
|
self._parent: "FrameProcessor" | None = None
|
||||||
self._prev: "FrameProcessor" | None = None
|
self._prev: "FrameProcessor" | None = None
|
||||||
self._next: "FrameProcessor" | None = None
|
self._next: "FrameProcessor" | None = None
|
||||||
@@ -83,6 +84,14 @@ class FrameProcessor:
|
|||||||
# the exception to this rule. This create this task.
|
# the exception to this rule. This create this task.
|
||||||
self.__create_push_task()
|
self.__create_push_task()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> int:
|
||||||
|
return self._id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
return self._name
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def interruptions_allowed(self):
|
def interruptions_allowed(self):
|
||||||
return self._allow_interruptions
|
return self._allow_interruptions
|
||||||
@@ -141,6 +150,16 @@ class FrameProcessor:
|
|||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
|
|
||||||
|
def create_task(self, coroutine: Coroutine) -> asyncio.Task:
|
||||||
|
name = f"{self}::{coroutine.cr_code.co_name}"
|
||||||
|
return create_task(self.get_event_loop(), coroutine, name)
|
||||||
|
|
||||||
|
async def cancel_task(self, task: asyncio.Task, timeout: Optional[float] = None):
|
||||||
|
await cancel_task(task, timeout)
|
||||||
|
|
||||||
|
async def wait_for_task(self, task: asyncio.Task, timeout: Optional[float] = None):
|
||||||
|
await wait_for_task(task, timeout)
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
await self.__cancel_input_task()
|
await self.__cancel_input_task()
|
||||||
await self.__cancel_push_task()
|
await self.__cancel_push_task()
|
||||||
@@ -188,7 +207,6 @@ class FrameProcessor:
|
|||||||
async def resume_processing_frames(self):
|
async def resume_processing_frames(self):
|
||||||
logger.trace(f"{self}: resuming frame processing")
|
logger.trace(f"{self}: resuming frame processing")
|
||||||
self.__input_event.set()
|
self.__input_event.set()
|
||||||
self.__should_block_frames = False
|
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
if isinstance(frame, StartFrame):
|
if isinstance(frame, StartFrame):
|
||||||
@@ -283,61 +301,44 @@ class FrameProcessor:
|
|||||||
def __create_input_task(self):
|
def __create_input_task(self):
|
||||||
self.__should_block_frames = False
|
self.__should_block_frames = False
|
||||||
self.__input_queue = asyncio.Queue()
|
self.__input_queue = asyncio.Queue()
|
||||||
self.__input_frame_task = self.get_event_loop().create_task(
|
|
||||||
self.__input_frame_task_handler()
|
|
||||||
)
|
|
||||||
self.__input_event = asyncio.Event()
|
self.__input_event = asyncio.Event()
|
||||||
|
self.__input_frame_task = self.create_task(self.__input_frame_task_handler())
|
||||||
|
|
||||||
async def __cancel_input_task(self):
|
async def __cancel_input_task(self):
|
||||||
self.__input_frame_task.cancel()
|
await self.cancel_task(self.__input_frame_task)
|
||||||
await self.__input_frame_task
|
|
||||||
|
|
||||||
async def __input_frame_task_handler(self):
|
async def __input_frame_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
if self.__should_block_frames:
|
||||||
if self.__should_block_frames:
|
logger.trace(f"{self}: frame processing paused")
|
||||||
logger.trace(f"{self}: frame processing paused")
|
await self.__input_event.wait()
|
||||||
await self.__input_event.wait()
|
self.__input_event.clear()
|
||||||
self.__input_event.clear()
|
self.__should_block_frames = False
|
||||||
logger.trace(f"{self}: frame processing resumed")
|
logger.trace(f"{self}: frame processing resumed")
|
||||||
|
|
||||||
(frame, direction, callback) = await self.__input_queue.get()
|
(frame, direction, callback) = await self.__input_queue.get()
|
||||||
|
|
||||||
# Process the frame.
|
# Process the frame.
|
||||||
await self.process_frame(frame, direction)
|
await self.process_frame(frame, direction)
|
||||||
|
|
||||||
# If this frame has an associated callback, call it now.
|
# If this frame has an associated callback, call it now.
|
||||||
if callback:
|
if callback:
|
||||||
await callback(self, frame, direction)
|
await callback(self, frame, direction)
|
||||||
|
|
||||||
self.__input_queue.task_done()
|
self.__input_queue.task_done()
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.trace(f"{self}: cancelled input task")
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self}: Uncaught exception {e}")
|
|
||||||
await self.push_error(ErrorFrame(str(e)))
|
|
||||||
|
|
||||||
def __create_push_task(self):
|
def __create_push_task(self):
|
||||||
self.__push_queue = asyncio.Queue()
|
self.__push_queue = asyncio.Queue()
|
||||||
self.__push_frame_task = self.get_event_loop().create_task(self.__push_frame_task_handler())
|
self.__push_frame_task = self.create_task(self.__push_frame_task_handler())
|
||||||
|
|
||||||
async def __cancel_push_task(self):
|
async def __cancel_push_task(self):
|
||||||
self.__push_frame_task.cancel()
|
await self.cancel_task(self.__push_frame_task)
|
||||||
await self.__push_frame_task
|
|
||||||
|
|
||||||
async def __push_frame_task_handler(self):
|
async def __push_frame_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
(frame, direction) = await self.__push_queue.get()
|
||||||
(frame, direction) = await self.__push_queue.get()
|
await self.__internal_push_frame(frame, direction)
|
||||||
await self.__internal_push_frame(frame, direction)
|
self.__push_queue.task_done()
|
||||||
self.__push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.trace(f"{self}: cancelled push task")
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self}: Uncaught exception {e}")
|
|
||||||
await self.push_error(ErrorFrame(str(e)))
|
|
||||||
|
|
||||||
async def _call_event_handler(self, event_name: str, *args, **kwargs):
|
async def _call_event_handler(self, event_name: str, *args, **kwargs):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -764,11 +764,11 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
|
|
||||||
# A task to process incoming action frames.
|
# A task to process incoming action frames.
|
||||||
self._action_queue = asyncio.Queue()
|
self._action_queue = asyncio.Queue()
|
||||||
self._action_task = self.get_event_loop().create_task(self._action_task_handler())
|
self._action_task = self.create_task(self._action_task_handler())
|
||||||
|
|
||||||
# A task to process incoming transport messages.
|
# A task to process incoming transport messages.
|
||||||
self._message_queue = asyncio.Queue()
|
self._message_queue = asyncio.Queue()
|
||||||
self._message_task = self.get_event_loop().create_task(self._message_task_handler())
|
self._message_task = self.create_task(self._message_task_handler())
|
||||||
|
|
||||||
self._register_event_handler("on_bot_started")
|
self._register_event_handler("on_bot_started")
|
||||||
self._register_event_handler("on_client_ready")
|
self._register_event_handler("on_client_ready")
|
||||||
@@ -873,13 +873,11 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
|
|
||||||
async def _cancel_tasks(self):
|
async def _cancel_tasks(self):
|
||||||
if self._action_task:
|
if self._action_task:
|
||||||
self._action_task.cancel()
|
await self.cancel_task(self._action_task)
|
||||||
await self._action_task
|
|
||||||
self._action_task = None
|
self._action_task = None
|
||||||
|
|
||||||
if self._message_task:
|
if self._message_task:
|
||||||
self._message_task.cancel()
|
await self.cancel_task(self._message_task)
|
||||||
await self._message_task
|
|
||||||
self._message_task = None
|
self._message_task = None
|
||||||
|
|
||||||
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
|
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
|
||||||
@@ -888,21 +886,15 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
|
|
||||||
async def _action_task_handler(self):
|
async def _action_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame = await self._action_queue.get()
|
||||||
frame = await self._action_queue.get()
|
await self._handle_action(frame.message_id, frame.rtvi_action_run)
|
||||||
await self._handle_action(frame.message_id, frame.rtvi_action_run)
|
self._action_queue.task_done()
|
||||||
self._action_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _message_task_handler(self):
|
async def _message_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
message = await self._message_queue.get()
|
||||||
message = await self._message_queue.get()
|
await self._handle_message(message)
|
||||||
await self._handle_message(message)
|
self._message_queue.task_done()
|
||||||
self._message_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _handle_transport_message(self, frame: TransportMessageUrgentFrame):
|
async def _handle_transport_message(self, frame: TransportMessageUrgentFrame):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -49,12 +49,11 @@ class IdleFrameProcessor(FrameProcessor):
|
|||||||
self._idle_event.set()
|
self._idle_event.set()
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
self._idle_task.cancel()
|
await self.cancel_task(self._idle_task)
|
||||||
await self._idle_task
|
|
||||||
|
|
||||||
def _create_idle_task(self):
|
def _create_idle_task(self):
|
||||||
self._idle_event = asyncio.Event()
|
self._idle_event = asyncio.Event()
|
||||||
self._idle_task = self.get_event_loop().create_task(self._idle_task_handler())
|
self._idle_task = self.create_task(self._idle_task_handler())
|
||||||
|
|
||||||
async def _idle_task_handler(self):
|
async def _idle_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
@@ -62,7 +61,5 @@ class IdleFrameProcessor(FrameProcessor):
|
|||||||
await asyncio.wait_for(self._idle_event.wait(), timeout=self._timeout)
|
await asyncio.wait_for(self._idle_event.wait(), timeout=self._timeout)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
await self._callback(self)
|
await self._callback(self)
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
finally:
|
finally:
|
||||||
self._idle_event.clear()
|
self._idle_event.clear()
|
||||||
|
|||||||
@@ -103,7 +103,7 @@ class UserIdleProcessor(FrameProcessor):
|
|||||||
def _create_idle_task(self) -> None:
|
def _create_idle_task(self) -> None:
|
||||||
"""Creates the idle task if it hasn't been created yet."""
|
"""Creates the idle task if it hasn't been created yet."""
|
||||||
if self._idle_task is None:
|
if self._idle_task is None:
|
||||||
self._idle_task = self.get_event_loop().create_task(self._idle_task_handler())
|
self._idle_task = self.create_task(self._idle_task_handler())
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def retry_count(self) -> int:
|
def retry_count(self) -> int:
|
||||||
@@ -113,11 +113,7 @@ class UserIdleProcessor(FrameProcessor):
|
|||||||
async def _stop(self) -> None:
|
async def _stop(self) -> None:
|
||||||
"""Stops and cleans up the idle monitoring task."""
|
"""Stops and cleans up the idle monitoring task."""
|
||||||
if self._idle_task is not None:
|
if self._idle_task is not None:
|
||||||
self._idle_task.cancel()
|
await self.cancel_task(self._idle_task)
|
||||||
try:
|
|
||||||
await self._idle_task
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass # Expected when task is cancelled
|
|
||||||
self._idle_task = None
|
self._idle_task = None
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
async def process_frame(self, frame: Frame, direction: FrameDirection) -> None:
|
||||||
@@ -160,6 +156,7 @@ class UserIdleProcessor(FrameProcessor):
|
|||||||
|
|
||||||
async def cleanup(self) -> None:
|
async def cleanup(self) -> None:
|
||||||
"""Cleans up resources when processor is shutting down."""
|
"""Cleans up resources when processor is shutting down."""
|
||||||
|
await super().cleanup()
|
||||||
if self._idle_task: # Only stop if task exists
|
if self._idle_task: # Only stop if task exists
|
||||||
await self._stop()
|
await self._stop()
|
||||||
|
|
||||||
@@ -178,7 +175,5 @@ class UserIdleProcessor(FrameProcessor):
|
|||||||
if not should_continue:
|
if not should_continue:
|
||||||
await self._stop()
|
await self._stop()
|
||||||
break
|
break
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
finally:
|
finally:
|
||||||
self._idle_event.clear()
|
self._idle_event.clear()
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import asyncio
|
|||||||
import io
|
import io
|
||||||
import wave
|
import wave
|
||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from typing import Any, AsyncGenerator, Dict, List, Optional, Tuple
|
from typing import Any, AsyncGenerator, Dict, List, Mapping, Optional, Tuple
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -69,7 +69,7 @@ class AIService(FrameProcessor):
|
|||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
async def _update_settings(self, settings: Dict[str, Any]):
|
async def _update_settings(self, settings: Mapping[str, Any]):
|
||||||
from pipecat.services.openai_realtime_beta.events import (
|
from pipecat.services.openai_realtime_beta.events import (
|
||||||
SessionProperties,
|
SessionProperties,
|
||||||
)
|
)
|
||||||
@@ -253,23 +253,21 @@ class TTSService(AIService):
|
|||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
if self._push_stop_frames:
|
if self._push_stop_frames:
|
||||||
self._stop_frame_task = self.get_event_loop().create_task(self._stop_frame_handler())
|
self._stop_frame_task = self.create_task(self._stop_frame_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
if self._stop_frame_task:
|
if self._stop_frame_task:
|
||||||
self._stop_frame_task.cancel()
|
await self.cancel_task(self._stop_frame_task)
|
||||||
await self._stop_frame_task
|
|
||||||
self._stop_frame_task = None
|
self._stop_frame_task = None
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
if self._stop_frame_task:
|
if self._stop_frame_task:
|
||||||
self._stop_frame_task.cancel()
|
await self.cancel_task(self._stop_frame_task)
|
||||||
await self._stop_frame_task
|
|
||||||
self._stop_frame_task = None
|
self._stop_frame_task = None
|
||||||
|
|
||||||
async def _update_settings(self, settings: Dict[str, Any]):
|
async def _update_settings(self, settings: Mapping[str, Any]):
|
||||||
for key, value in settings.items():
|
for key, value in settings.items():
|
||||||
if key in self._settings:
|
if key in self._settings:
|
||||||
logger.info(f"Updating TTS setting {key} to: [{value}]")
|
logger.info(f"Updating TTS setting {key} to: [{value}]")
|
||||||
@@ -364,23 +362,20 @@ class TTSService(AIService):
|
|||||||
await self.push_frame(TTSTextFrame(text))
|
await self.push_frame(TTSTextFrame(text))
|
||||||
|
|
||||||
async def _stop_frame_handler(self):
|
async def _stop_frame_handler(self):
|
||||||
try:
|
has_started = False
|
||||||
has_started = False
|
while True:
|
||||||
while True:
|
try:
|
||||||
try:
|
frame = await asyncio.wait_for(
|
||||||
frame = await asyncio.wait_for(
|
self._stop_frame_queue.get(), self._stop_frame_timeout_s
|
||||||
self._stop_frame_queue.get(), self._stop_frame_timeout_s
|
)
|
||||||
)
|
if isinstance(frame, TTSStartedFrame):
|
||||||
if isinstance(frame, TTSStartedFrame):
|
has_started = True
|
||||||
has_started = True
|
elif isinstance(frame, (TTSStoppedFrame, StartInterruptionFrame)):
|
||||||
elif isinstance(frame, (TTSStoppedFrame, StartInterruptionFrame)):
|
has_started = False
|
||||||
has_started = False
|
except asyncio.TimeoutError:
|
||||||
except asyncio.TimeoutError:
|
if has_started:
|
||||||
if has_started:
|
await self.push_frame(TTSStoppedFrame())
|
||||||
await self.push_frame(TTSStoppedFrame())
|
has_started = False
|
||||||
has_started = False
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class WordTTSService(TTSService):
|
class WordTTSService(TTSService):
|
||||||
@@ -388,7 +383,7 @@ class WordTTSService(TTSService):
|
|||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._initial_word_timestamp = -1
|
self._initial_word_timestamp = -1
|
||||||
self._words_queue = asyncio.Queue()
|
self._words_queue = asyncio.Queue()
|
||||||
self._words_task = self.get_event_loop().create_task(self._words_task_handler())
|
self._words_task = self.create_task(self._words_task_handler())
|
||||||
|
|
||||||
def start_word_timestamps(self):
|
def start_word_timestamps(self):
|
||||||
if self._initial_word_timestamp == -1:
|
if self._initial_word_timestamp == -1:
|
||||||
@@ -421,35 +416,29 @@ class WordTTSService(TTSService):
|
|||||||
|
|
||||||
async def _stop_words_task(self):
|
async def _stop_words_task(self):
|
||||||
if self._words_task:
|
if self._words_task:
|
||||||
self._words_task.cancel()
|
await self.cancel_task(self._words_task)
|
||||||
await self._words_task
|
|
||||||
self._words_task = None
|
self._words_task = None
|
||||||
|
|
||||||
async def _words_task_handler(self):
|
async def _words_task_handler(self):
|
||||||
last_pts = 0
|
last_pts = 0
|
||||||
while True:
|
while True:
|
||||||
try:
|
(word, timestamp) = await self._words_queue.get()
|
||||||
(word, timestamp) = await self._words_queue.get()
|
if word == "Reset" and timestamp == 0:
|
||||||
if word == "Reset" and timestamp == 0:
|
self.reset_word_timestamps()
|
||||||
self.reset_word_timestamps()
|
frame = None
|
||||||
frame = None
|
elif word == "LLMFullResponseEndFrame" and timestamp == 0:
|
||||||
elif word == "LLMFullResponseEndFrame" and timestamp == 0:
|
frame = LLMFullResponseEndFrame()
|
||||||
frame = LLMFullResponseEndFrame()
|
frame.pts = last_pts
|
||||||
frame.pts = last_pts
|
elif word == "TTSStoppedFrame" and timestamp == 0:
|
||||||
elif word == "TTSStoppedFrame" and timestamp == 0:
|
frame = TTSStoppedFrame()
|
||||||
frame = TTSStoppedFrame()
|
frame.pts = last_pts
|
||||||
frame.pts = last_pts
|
else:
|
||||||
else:
|
frame = TTSTextFrame(word)
|
||||||
frame = TTSTextFrame(word)
|
frame.pts = self._initial_word_timestamp + timestamp
|
||||||
frame.pts = self._initial_word_timestamp + timestamp
|
if frame:
|
||||||
if frame:
|
last_pts = frame.pts
|
||||||
last_pts = frame.pts
|
await self.push_frame(frame)
|
||||||
await self.push_frame(frame)
|
self._words_queue.task_done()
|
||||||
self._words_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} exception: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
class STTService(AIService):
|
class STTService(AIService):
|
||||||
@@ -479,7 +468,7 @@ class STTService(AIService):
|
|||||||
"""Returns transcript as a string"""
|
"""Returns transcript as a string"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
async def _update_settings(self, settings: Dict[str, Any]):
|
async def _update_settings(self, settings: Mapping[str, Any]):
|
||||||
logger.info(f"Updating STT settings: {self._settings}")
|
logger.info(f"Updating STT settings: {self._settings}")
|
||||||
for key, value in settings.items():
|
for key, value in settings.items():
|
||||||
if key in self._settings:
|
if key in self._settings:
|
||||||
|
|||||||
@@ -187,16 +187,13 @@ class CartesiaTTSService(WordTTSService, WebsocketService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
self._receive_task = self.get_event_loop().create_task(
|
self._receive_task = self.create_task(self._receive_task_handler(self.push_error))
|
||||||
self._receive_task_handler(self.push_error)
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
|
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task)
|
||||||
await self._receive_task
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
|
|||||||
@@ -299,20 +299,16 @@ class ElevenLabsTTSService(WordTTSService, WebsocketService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
self._receive_task = self.get_event_loop().create_task(
|
self._receive_task = self.create_task(self._receive_task_handler(self.push_error))
|
||||||
self._receive_task_handler(self.push_error)
|
self._keepalive_task = self.create_task(self._keepalive_task_handler())
|
||||||
)
|
|
||||||
self._keepalive_task = self.get_event_loop().create_task(self._keepalive_task_handler())
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task)
|
||||||
await self._receive_task
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
if self._keepalive_task:
|
if self._keepalive_task:
|
||||||
self._keepalive_task.cancel()
|
await self.cancel_task(self._keepalive_task)
|
||||||
await self._keepalive_task
|
|
||||||
self._keepalive_task = None
|
self._keepalive_task = None
|
||||||
|
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
@@ -383,13 +379,8 @@ class ElevenLabsTTSService(WordTTSService, WebsocketService):
|
|||||||
|
|
||||||
async def _keepalive_task_handler(self):
|
async def _keepalive_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
await asyncio.sleep(10)
|
||||||
await asyncio.sleep(10)
|
await self._send_text("")
|
||||||
await self._send_text("")
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} exception: {e}")
|
|
||||||
|
|
||||||
async def _send_text(self, text: str):
|
async def _send_text(self, text: str):
|
||||||
if self._websocket:
|
if self._websocket:
|
||||||
|
|||||||
@@ -104,15 +104,12 @@ class FishAudioTTSService(TTSService, WebsocketService):
|
|||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
self._receive_task = self.get_event_loop().create_task(
|
self._receive_task = self.create_task(self._receive_task_handler(self.push_error))
|
||||||
self._receive_task_handler(self.push_error)
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task)
|
||||||
await self._receive_task
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
|
|||||||
@@ -182,9 +182,13 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
|
|
||||||
self._audio_input_paused = start_audio_paused
|
self._audio_input_paused = start_audio_paused
|
||||||
self._video_input_paused = start_video_paused
|
self._video_input_paused = start_video_paused
|
||||||
|
self._context = None
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
self._context = None
|
self._transcribe_audio_task = None
|
||||||
|
self._transcribe_model_audio_task = None
|
||||||
|
self._transcribe_audio_queue = asyncio.Queue()
|
||||||
|
self._transcribe_model_audio_queue = asyncio.Queue()
|
||||||
|
|
||||||
self._disconnecting = False
|
self._disconnecting = False
|
||||||
self._api_session_ready = False
|
self._api_session_ready = False
|
||||||
@@ -244,6 +248,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
@@ -275,7 +280,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
)
|
)
|
||||||
await self.send_client_event(evt)
|
await self.send_client_event(evt)
|
||||||
if self._transcribe_user_audio and self._context:
|
if self._transcribe_user_audio and self._context:
|
||||||
asyncio.create_task(self._handle_transcribe_user_audio(audio, self._context))
|
await self._transcribe_audio_queue.put(audio)
|
||||||
|
|
||||||
async def _handle_transcribe_user_audio(self, audio, context):
|
async def _handle_transcribe_user_audio(self, audio, context):
|
||||||
text = await self._transcribe_audio(audio, context)
|
text = await self._transcribe_audio(audio, context)
|
||||||
@@ -381,17 +386,21 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
await self._ws_send(event.model_dump(exclude_none=True))
|
await self._ws_send(event.model_dump(exclude_none=True))
|
||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
|
if self._websocket:
|
||||||
|
# Here we assume that if we have a websocket, we are connected. We
|
||||||
|
# handle disconnections in the send/recv code paths.
|
||||||
|
return
|
||||||
|
|
||||||
logger.info("Connecting to Gemini service")
|
logger.info("Connecting to Gemini service")
|
||||||
try:
|
try:
|
||||||
if self._websocket:
|
|
||||||
# Here we assume that if we have a websocket, we are connected. We
|
|
||||||
# handle disconnections in the send/recv code paths.
|
|
||||||
return
|
|
||||||
|
|
||||||
uri = f"wss://{self.base_url}/ws/google.ai.generativelanguage.v1alpha.GenerativeService.BidiGenerateContent?key={self.api_key}"
|
uri = f"wss://{self.base_url}/ws/google.ai.generativelanguage.v1alpha.GenerativeService.BidiGenerateContent?key={self.api_key}"
|
||||||
logger.info(f"Connecting to {uri}")
|
logger.info(f"Connecting to {uri}")
|
||||||
self._websocket = await websockets.connect(uri=uri)
|
self._websocket = await websockets.connect(uri=uri)
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
self._receive_task = self.create_task(self._receive_task_handler())
|
||||||
|
self._transcribe_audio_task = self.create_task(self._transcribe_audio_handler())
|
||||||
|
self._transcribe_model_audio_task = self.create_task(
|
||||||
|
self._transcribe_model_audio_handler()
|
||||||
|
)
|
||||||
config = events.Config.model_validate(
|
config = events.Config.model_validate(
|
||||||
{
|
{
|
||||||
"setup": {
|
"setup": {
|
||||||
@@ -441,12 +450,14 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task, timeout=1.0)
|
||||||
try:
|
|
||||||
await asyncio.wait_for(self._receive_task, timeout=1.0)
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.warning("Timed out waiting for receive task to finish")
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
if self._transcribe_audio_task:
|
||||||
|
await self.cancel_task(self._transcribe_audio_task)
|
||||||
|
self._transcribe_audio_task = None
|
||||||
|
if self._transcribe_model_audio_task:
|
||||||
|
await self.cancel_task(self._transcribe_model_audio_task)
|
||||||
|
self._transcribe_model_audio_task = None
|
||||||
self._disconnecting = False
|
self._disconnecting = False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error disconnecting: {e}")
|
logger.error(f"{self} error disconnecting: {e}")
|
||||||
@@ -454,9 +465,8 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
async def _ws_send(self, message):
|
async def _ws_send(self, message):
|
||||||
# logger.debug(f"Sending message to websocket: {message}")
|
# logger.debug(f"Sending message to websocket: {message}")
|
||||||
try:
|
try:
|
||||||
if not self._websocket:
|
if self._websocket:
|
||||||
await self._connect()
|
await self._websocket.send(json.dumps(message))
|
||||||
await self._websocket.send(json.dumps(message))
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
if self._disconnecting:
|
if self._disconnecting:
|
||||||
return
|
return
|
||||||
@@ -473,32 +483,35 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
#
|
#
|
||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
try:
|
async for message in self._websocket:
|
||||||
async for message in self._websocket:
|
evt = events.parse_server_event(message)
|
||||||
evt = events.parse_server_event(message)
|
# logger.debug(f"Received event: {message[:500]}")
|
||||||
# logger.debug(f"Received event: {message[:500]}")
|
# logger.debug(f"Received event: {evt}")
|
||||||
# logger.debug(f"Received event: {evt}")
|
|
||||||
|
|
||||||
if evt.setupComplete:
|
if evt.setupComplete:
|
||||||
await self._handle_evt_setup_complete(evt)
|
await self._handle_evt_setup_complete(evt)
|
||||||
elif evt.serverContent and evt.serverContent.modelTurn:
|
elif evt.serverContent and evt.serverContent.modelTurn:
|
||||||
await self._handle_evt_model_turn(evt)
|
await self._handle_evt_model_turn(evt)
|
||||||
elif evt.serverContent and evt.serverContent.turnComplete:
|
elif evt.serverContent and evt.serverContent.turnComplete:
|
||||||
await self._handle_evt_turn_complete(evt)
|
await self._handle_evt_turn_complete(evt)
|
||||||
elif evt.toolCall:
|
elif evt.toolCall:
|
||||||
await self._handle_evt_tool_call(evt)
|
await self._handle_evt_tool_call(evt)
|
||||||
|
elif False: # !!! todo: error events?
|
||||||
|
await self._handle_evt_error(evt)
|
||||||
|
# errors are fatal, so exit the receive loop
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
pass
|
||||||
|
|
||||||
elif False: # !!! todo: error events?
|
async def _transcribe_audio_handler(self):
|
||||||
await self._handle_evt_error(evt)
|
while True:
|
||||||
# errors are fatal, so exit the receive loop
|
audio = await self._transcribe_audio_queue.get()
|
||||||
return
|
await self._handle_transcribe_user_audio(audio, self._context)
|
||||||
|
|
||||||
else:
|
async def _transcribe_model_audio_handler(self):
|
||||||
pass
|
while True:
|
||||||
except asyncio.CancelledError:
|
audio = await self._transcribe_model_audio_queue.get()
|
||||||
logger.debug("websocket receive task cancelled")
|
await self._handle_transcribe_model_audio(audio, self._context)
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} exception: {e}")
|
|
||||||
|
|
||||||
#
|
#
|
||||||
#
|
#
|
||||||
@@ -679,7 +692,7 @@ class GeminiMultimodalLiveLLMService(LLMService):
|
|||||||
self._bot_text_buffer = ""
|
self._bot_text_buffer = ""
|
||||||
|
|
||||||
if audio and self._transcribe_model_audio and self._context:
|
if audio and self._transcribe_model_audio and self._context:
|
||||||
asyncio.create_task(self._handle_transcribe_model_audio(audio, self._context))
|
await self._transcribe_model_audio_queue.put(audio)
|
||||||
elif text:
|
elif text:
|
||||||
await self.push_frame(LLMFullResponseEndFrame())
|
await self.push_frame(LLMFullResponseEndFrame())
|
||||||
|
|
||||||
|
|||||||
@@ -180,7 +180,7 @@ class GladiaSTTService(STTService):
|
|||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
response = await self._setup_gladia()
|
response = await self._setup_gladia()
|
||||||
self._websocket = await websockets.connect(response["url"])
|
self._websocket = await websockets.connect(response["url"])
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
self._receive_task = self.create_task(self._receive_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
|
|||||||
@@ -113,16 +113,13 @@ class LmntTTSService(TTSService, WebsocketService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
self._receive_task = self.get_event_loop().create_task(
|
self._receive_task = self.create_task(self._receive_task_handler(self.push_error))
|
||||||
self._receive_task_handler(self.push_error)
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
|
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task)
|
||||||
await self._receive_task
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
|
|||||||
@@ -277,7 +277,7 @@ class OpenAIRealtimeBetaLLMService(LLMService):
|
|||||||
"OpenAI-Beta": "realtime=v1",
|
"OpenAI-Beta": "realtime=v1",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
self._receive_task = self.create_task(self._receive_task_handler())
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} initialization error: {e}")
|
logger.error(f"{self} initialization error: {e}")
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
@@ -291,11 +291,7 @@ class OpenAIRealtimeBetaLLMService(LLMService):
|
|||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task, timeout=1.0)
|
||||||
try:
|
|
||||||
await asyncio.wait_for(self._receive_task, timeout=1.0)
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.warning("Timed out waiting for receive task to finish")
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
self._disconnecting = False
|
self._disconnecting = False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -332,40 +328,32 @@ class OpenAIRealtimeBetaLLMService(LLMService):
|
|||||||
#
|
#
|
||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
try:
|
async for message in self._websocket:
|
||||||
async for message in self._websocket:
|
evt = events.parse_server_event(message)
|
||||||
evt = events.parse_server_event(message)
|
if evt.type == "session.created":
|
||||||
if evt.type == "session.created":
|
await self._handle_evt_session_created(evt)
|
||||||
await self._handle_evt_session_created(evt)
|
elif evt.type == "session.updated":
|
||||||
elif evt.type == "session.updated":
|
await self._handle_evt_session_updated(evt)
|
||||||
await self._handle_evt_session_updated(evt)
|
elif evt.type == "response.audio.delta":
|
||||||
elif evt.type == "response.audio.delta":
|
await self._handle_evt_audio_delta(evt)
|
||||||
await self._handle_evt_audio_delta(evt)
|
elif evt.type == "response.audio.done":
|
||||||
elif evt.type == "response.audio.done":
|
await self._handle_evt_audio_done(evt)
|
||||||
await self._handle_evt_audio_done(evt)
|
elif evt.type == "conversation.item.created":
|
||||||
elif evt.type == "conversation.item.created":
|
await self._handle_evt_conversation_item_created(evt)
|
||||||
await self._handle_evt_conversation_item_created(evt)
|
elif evt.type == "conversation.item.input_audio_transcription.completed":
|
||||||
elif evt.type == "conversation.item.input_audio_transcription.completed":
|
await self.handle_evt_input_audio_transcription_completed(evt)
|
||||||
await self.handle_evt_input_audio_transcription_completed(evt)
|
elif evt.type == "response.done":
|
||||||
elif evt.type == "response.done":
|
await self._handle_evt_response_done(evt)
|
||||||
await self._handle_evt_response_done(evt)
|
elif evt.type == "input_audio_buffer.speech_started":
|
||||||
elif evt.type == "input_audio_buffer.speech_started":
|
await self._handle_evt_speech_started(evt)
|
||||||
await self._handle_evt_speech_started(evt)
|
elif evt.type == "input_audio_buffer.speech_stopped":
|
||||||
elif evt.type == "input_audio_buffer.speech_stopped":
|
await self._handle_evt_speech_stopped(evt)
|
||||||
await self._handle_evt_speech_stopped(evt)
|
elif evt.type == "response.audio_transcript.delta":
|
||||||
elif evt.type == "response.audio_transcript.delta":
|
await self._handle_evt_audio_transcript_delta(evt)
|
||||||
await self._handle_evt_audio_transcript_delta(evt)
|
elif evt.type == "error":
|
||||||
elif evt.type == "error":
|
await self._handle_evt_error(evt)
|
||||||
await self._handle_evt_error(evt)
|
# errors are fatal, so exit the receive loop
|
||||||
# errors are fatal, so exit the receive loop
|
return
|
||||||
return
|
|
||||||
|
|
||||||
else:
|
|
||||||
pass
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.debug("websocket receive task cancelled")
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} exception: {e}")
|
|
||||||
|
|
||||||
async def _handle_evt_session_created(self, evt):
|
async def _handle_evt_session_created(self, evt):
|
||||||
# session.created is received right after connecting. Send a message
|
# session.created is received right after connecting. Send a message
|
||||||
|
|||||||
@@ -165,16 +165,13 @@ class PlayHTTTSService(TTSService, WebsocketService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
await self._connect_websocket()
|
await self._connect_websocket()
|
||||||
|
|
||||||
self._receive_task = self.get_event_loop().create_task(
|
self._receive_task = self.create_task(self._receive_task_handler(self.push_error))
|
||||||
self._receive_task_handler(self.push_error)
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _disconnect(self):
|
async def _disconnect(self):
|
||||||
await self._disconnect_websocket()
|
await self._disconnect_websocket()
|
||||||
|
|
||||||
if self._receive_task:
|
if self._receive_task:
|
||||||
self._receive_task.cancel()
|
await self.cancel_task(self._receive_task)
|
||||||
await self._receive_task
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
async def _connect_websocket(self):
|
async def _connect_websocket(self):
|
||||||
|
|||||||
@@ -202,8 +202,8 @@ class ParakeetSTTService(STTService):
|
|||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
self._thread_task = self.get_event_loop().create_task(self._thread_task_handler())
|
self._thread_task = self.create_task(self._thread_task_handler())
|
||||||
self._response_task = self.get_event_loop().create_task(self._response_task_handler())
|
self._response_task = self.create_task(self._response_task_handler())
|
||||||
self._response_queue = asyncio.Queue()
|
self._response_queue = asyncio.Queue()
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
@@ -215,10 +215,8 @@ class ParakeetSTTService(STTService):
|
|||||||
await self._stop_tasks()
|
await self._stop_tasks()
|
||||||
|
|
||||||
async def _stop_tasks(self):
|
async def _stop_tasks(self):
|
||||||
self._thread_task.cancel()
|
await self.cancel_task(self._thread_task)
|
||||||
await self._thread_task
|
await self.cancel_task(self._response_task)
|
||||||
self._response_task.cancel()
|
|
||||||
await self._response_task
|
|
||||||
|
|
||||||
def _response_handler(self):
|
def _response_handler(self):
|
||||||
responses = self._asr_service.streaming_response_generator(
|
responses = self._asr_service.streaming_response_generator(
|
||||||
@@ -238,7 +236,7 @@ class ParakeetSTTService(STTService):
|
|||||||
await asyncio.to_thread(self._response_handler)
|
await asyncio.to_thread(self._response_handler)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
self._thread_running = False
|
self._thread_running = False
|
||||||
pass
|
raise
|
||||||
|
|
||||||
async def _handle_response(self, response):
|
async def _handle_response(self, response):
|
||||||
for result in response.results:
|
for result in response.results:
|
||||||
@@ -260,11 +258,8 @@ class ParakeetSTTService(STTService):
|
|||||||
|
|
||||||
async def _response_task_handler(self):
|
async def _response_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
response = await self._response_queue.get()
|
||||||
response = await self._response_queue.get()
|
await self._handle_response(response)
|
||||||
await self._handle_response(response)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
await self._queue.put(audio)
|
await self._queue.put(audio)
|
||||||
|
|||||||
@@ -49,45 +49,33 @@ class SimliVideoService(FrameProcessor):
|
|||||||
async def _start_connection(self):
|
async def _start_connection(self):
|
||||||
await self._simli_client.Initialize()
|
await self._simli_client.Initialize()
|
||||||
# Create task to consume and process audio and video
|
# Create task to consume and process audio and video
|
||||||
self._audio_task = asyncio.create_task(self._consume_and_process_audio())
|
self._audio_task = self.create_task(self._consume_and_process_audio())
|
||||||
self._video_task = asyncio.create_task(self._consume_and_process_video())
|
self._video_task = self.create_task(self._consume_and_process_video())
|
||||||
|
|
||||||
async def _consume_and_process_audio(self):
|
async def _consume_and_process_audio(self):
|
||||||
try:
|
await self._pipecat_resampler_event.wait()
|
||||||
await self._pipecat_resampler_event.wait()
|
async for audio_frame in self._simli_client.getAudioStreamIterator():
|
||||||
async for audio_frame in self._simli_client.getAudioStreamIterator():
|
resampled_frames = self._pipecat_resampler.resample(audio_frame)
|
||||||
resampled_frames = self._pipecat_resampler.resample(audio_frame)
|
for resampled_frame in resampled_frames:
|
||||||
for resampled_frame in resampled_frames:
|
await self.push_frame(
|
||||||
await self.push_frame(
|
TTSAudioRawFrame(
|
||||||
TTSAudioRawFrame(
|
audio=resampled_frame.to_ndarray().tobytes(),
|
||||||
audio=resampled_frame.to_ndarray().tobytes(),
|
sample_rate=self._pipecat_resampler.rate,
|
||||||
sample_rate=self._pipecat_resampler.rate,
|
num_channels=1,
|
||||||
num_channels=1,
|
),
|
||||||
),
|
)
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} exception: {e}")
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def _consume_and_process_video(self):
|
async def _consume_and_process_video(self):
|
||||||
try:
|
await self._pipecat_resampler_event.wait()
|
||||||
await self._pipecat_resampler_event.wait()
|
async for video_frame in self._simli_client.getVideoStreamIterator(targetFormat="rgb24"):
|
||||||
async for video_frame in self._simli_client.getVideoStreamIterator(
|
# Process the video frame
|
||||||
targetFormat="rgb24"
|
convertedFrame: OutputImageRawFrame = OutputImageRawFrame(
|
||||||
):
|
image=video_frame.to_rgb().to_image().tobytes(),
|
||||||
# Process the video frame
|
size=(video_frame.width, video_frame.height),
|
||||||
convertedFrame: OutputImageRawFrame = OutputImageRawFrame(
|
format="RGB",
|
||||||
image=video_frame.to_rgb().to_image().tobytes(),
|
)
|
||||||
size=(video_frame.width, video_frame.height),
|
convertedFrame.pts = video_frame.pts
|
||||||
format="RGB",
|
await self.push_frame(convertedFrame)
|
||||||
)
|
|
||||||
convertedFrame.pts = video_frame.pts
|
|
||||||
await self.push_frame(convertedFrame)
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} exception: {e}")
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
@@ -128,8 +116,6 @@ class SimliVideoService(FrameProcessor):
|
|||||||
async def _stop(self):
|
async def _stop(self):
|
||||||
await self._simli_client.stop()
|
await self._simli_client.stop()
|
||||||
if self._audio_task:
|
if self._audio_task:
|
||||||
self._audio_task.cancel()
|
await self.cancel_task(self._audio_task)
|
||||||
await self._audio_task
|
|
||||||
if self._video_task:
|
if self._video_task:
|
||||||
self._video_task.cancel()
|
await self.cancel_task(self._video_task)
|
||||||
await self._video_task
|
|
||||||
|
|||||||
@@ -75,6 +75,14 @@ class TavusVideoService(AIService):
|
|||||||
logger.debug(f"TavusVideoService persona grabbed {response_json}")
|
logger.debug(f"TavusVideoService persona grabbed {response_json}")
|
||||||
return response_json["persona_name"]
|
return response_json["persona_name"]
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
await super().stop(frame)
|
||||||
|
await self._end_conversation()
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
await super().cancel(frame)
|
||||||
|
await self._end_conversation()
|
||||||
|
|
||||||
async def _end_conversation(self) -> None:
|
async def _end_conversation(self) -> None:
|
||||||
url = f"https://tavusapi.com/v2/conversations/{self._conversation_id}/end"
|
url = f"https://tavusapi.com/v2/conversations/{self._conversation_id}/end"
|
||||||
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
headers = {"Content-Type": "application/json", "x-api-key": self._api_key}
|
||||||
@@ -105,8 +113,6 @@ class TavusVideoService(AIService):
|
|||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
elif isinstance(frame, StartInterruptionFrame):
|
||||||
await self._send_interrupt_message()
|
await self._send_interrupt_message()
|
||||||
elif isinstance(frame, (EndFrame, CancelFrame)):
|
|
||||||
await self._end_conversation()
|
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|||||||
@@ -85,10 +85,6 @@ class WebsocketService(ABC):
|
|||||||
await self._receive_messages()
|
await self._receive_messages()
|
||||||
logger.debug(f"{self} connection established successfully")
|
logger.debug(f"{self} connection established successfully")
|
||||||
retry_count = 0 # Reset counter on successful message receive
|
retry_count = 0 # Reset counter on successful message receive
|
||||||
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
retry_count += 1
|
retry_count += 1
|
||||||
if retry_count >= MAX_RETRIES:
|
if retry_count >= MAX_RETRIES:
|
||||||
|
|||||||
@@ -50,13 +50,12 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
# Create audio input queue and task if needed.
|
# Create audio input queue and task if needed.
|
||||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||||
self._audio_in_queue = asyncio.Queue()
|
self._audio_in_queue = asyncio.Queue()
|
||||||
self._audio_task = self.get_event_loop().create_task(self._audio_task_handler())
|
self._audio_task = self.create_task(self._audio_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
# Cancel and wait for the audio input task to finish.
|
# Cancel and wait for the audio input task to finish.
|
||||||
if self._audio_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
if self._audio_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||||
self._audio_task.cancel()
|
await self.cancel_task(self._audio_task)
|
||||||
await self._audio_task
|
|
||||||
self._audio_task = None
|
self._audio_task = None
|
||||||
# Stop audio filter.
|
# Stop audio filter.
|
||||||
if self._params.audio_in_filter:
|
if self._params.audio_in_filter:
|
||||||
@@ -65,8 +64,7 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
# Cancel and wait for the audio input task to finish.
|
# Cancel and wait for the audio input task to finish.
|
||||||
if self._audio_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
if self._audio_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||||
self._audio_task.cancel()
|
await self.cancel_task(self._audio_task)
|
||||||
await self._audio_task
|
|
||||||
self._audio_task = None
|
self._audio_task = None
|
||||||
|
|
||||||
def vad_analyzer(self) -> VADAnalyzer | None:
|
def vad_analyzer(self) -> VADAnalyzer | None:
|
||||||
@@ -173,27 +171,22 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
async def _audio_task_handler(self):
|
async def _audio_task_handler(self):
|
||||||
vad_state: VADState = VADState.QUIET
|
vad_state: VADState = VADState.QUIET
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame: InputAudioRawFrame = await self._audio_in_queue.get()
|
||||||
frame: InputAudioRawFrame = await self._audio_in_queue.get()
|
|
||||||
|
|
||||||
audio_passthrough = True
|
audio_passthrough = True
|
||||||
|
|
||||||
# If an audio filter is available, run it before VAD.
|
# If an audio filter is available, run it before VAD.
|
||||||
if self._params.audio_in_filter:
|
if self._params.audio_in_filter:
|
||||||
frame.audio = await self._params.audio_in_filter.filter(frame.audio)
|
frame.audio = await self._params.audio_in_filter.filter(frame.audio)
|
||||||
|
|
||||||
# Check VAD and push event if necessary. We just care about
|
# Check VAD and push event if necessary. We just care about
|
||||||
# changes from QUIET to SPEAKING and vice versa.
|
# changes from QUIET to SPEAKING and vice versa.
|
||||||
if self._params.vad_enabled:
|
if self._params.vad_enabled:
|
||||||
vad_state = await self._handle_vad(frame, vad_state)
|
vad_state = await self._handle_vad(frame, vad_state)
|
||||||
audio_passthrough = self._params.vad_audio_passthrough
|
audio_passthrough = self._params.vad_audio_passthrough
|
||||||
|
|
||||||
# Push audio downstream if passthrough.
|
# Push audio downstream if passthrough.
|
||||||
if audio_passthrough:
|
if audio_passthrough:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
self._audio_in_queue.task_done()
|
self._audio_in_queue.task_done()
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} error reading audio frames: {e}")
|
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ from pipecat.frames.frames import (
|
|||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.transports.base_transport import TransportParams
|
from pipecat.transports.base_transport import TransportParams
|
||||||
|
from pipecat.utils.asyncio import wait_for_task
|
||||||
from pipecat.utils.time import nanoseconds_to_seconds
|
from pipecat.utils.time import nanoseconds_to_seconds
|
||||||
|
|
||||||
|
|
||||||
@@ -87,9 +88,9 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
# for these tasks before cancelling the camera and audio tasks below
|
# for these tasks before cancelling the camera and audio tasks below
|
||||||
# because they might be still rendering.
|
# because they might be still rendering.
|
||||||
if self._sink_task:
|
if self._sink_task:
|
||||||
await self._sink_task
|
await wait_for_task(self._sink_task)
|
||||||
if self._sink_clock_task:
|
if self._sink_clock_task:
|
||||||
await self._sink_clock_task
|
await wait_for_task(self._sink_clock_task)
|
||||||
|
|
||||||
# We can now cancel the camera task.
|
# We can now cancel the camera task.
|
||||||
await self._cancel_camera_task()
|
await self._cancel_camera_task()
|
||||||
@@ -217,22 +218,19 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
#
|
#
|
||||||
|
|
||||||
def _create_sink_tasks(self):
|
def _create_sink_tasks(self):
|
||||||
loop = self.get_event_loop()
|
|
||||||
self._sink_queue = asyncio.Queue()
|
self._sink_queue = asyncio.Queue()
|
||||||
self._sink_task = loop.create_task(self._sink_task_handler())
|
|
||||||
self._sink_clock_queue = asyncio.PriorityQueue()
|
self._sink_clock_queue = asyncio.PriorityQueue()
|
||||||
self._sink_clock_task = loop.create_task(self._sink_clock_task_handler())
|
self._sink_task = self.create_task(self._sink_task_handler())
|
||||||
|
self._sink_clock_task = self.create_task(self._sink_clock_task_handler())
|
||||||
|
|
||||||
async def _cancel_sink_tasks(self):
|
async def _cancel_sink_tasks(self):
|
||||||
# Stop sink tasks.
|
# Stop sink tasks.
|
||||||
if self._sink_task:
|
if self._sink_task:
|
||||||
self._sink_task.cancel()
|
await self.cancel_task(self._sink_task)
|
||||||
await self._sink_task
|
|
||||||
self._sink_task = None
|
self._sink_task = None
|
||||||
# Stop sink clock tasks.
|
# Stop sink clock tasks.
|
||||||
if self._sink_clock_task:
|
if self._sink_clock_task:
|
||||||
self._sink_clock_task.cancel()
|
await self.cancel_task(self._sink_clock_task)
|
||||||
await self._sink_clock_task
|
|
||||||
self._sink_clock_task = None
|
self._sink_clock_task = None
|
||||||
|
|
||||||
async def _sink_frame_handler(self, frame: Frame):
|
async def _sink_frame_handler(self, frame: Frame):
|
||||||
@@ -269,7 +267,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
|
|
||||||
self._sink_clock_queue.task_done()
|
self._sink_clock_queue.task_done()
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"{self} error processing sink clock queue: {e}")
|
logger.exception(f"{self} error processing sink clock queue: {e}")
|
||||||
|
|
||||||
@@ -317,49 +315,42 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
return without_mixer(vad_stop_secs)
|
return without_mixer(vad_stop_secs)
|
||||||
|
|
||||||
async def _sink_task_handler(self):
|
async def _sink_task_handler(self):
|
||||||
try:
|
async for frame in self._next_frame():
|
||||||
async for frame in self._next_frame():
|
# Notify the bot started speaking upstream if necessary and that
|
||||||
# Notify the bot started speaking upstream if necessary and that
|
# it's actually speaking.
|
||||||
# it's actually speaking.
|
if isinstance(frame, TTSAudioRawFrame):
|
||||||
if isinstance(frame, TTSAudioRawFrame):
|
await self._bot_started_speaking()
|
||||||
await self._bot_started_speaking()
|
await self.push_frame(BotSpeakingFrame())
|
||||||
await self.push_frame(BotSpeakingFrame())
|
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
||||||
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
|
||||||
|
|
||||||
# No need to push EndFrame, it's pushed from process_frame().
|
# No need to push EndFrame, it's pushed from process_frame().
|
||||||
if isinstance(frame, EndFrame):
|
if isinstance(frame, EndFrame):
|
||||||
break
|
break
|
||||||
|
|
||||||
# Handle frame.
|
# Handle frame.
|
||||||
await self._sink_frame_handler(frame)
|
await self._sink_frame_handler(frame)
|
||||||
|
|
||||||
# Also, push frame downstream in case anyone else needs it.
|
# Also, push frame downstream in case anyone else needs it.
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
# Send audio.
|
# Send audio.
|
||||||
if isinstance(frame, OutputAudioRawFrame):
|
if isinstance(frame, OutputAudioRawFrame):
|
||||||
await self.write_raw_audio_frames(frame.audio)
|
await self.write_raw_audio_frames(frame.audio)
|
||||||
except asyncio.CancelledError:
|
|
||||||
pass
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} error writing to microphone: {e}")
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Camera task
|
# Camera task
|
||||||
#
|
#
|
||||||
|
|
||||||
def _create_camera_task(self):
|
def _create_camera_task(self):
|
||||||
loop = self.get_event_loop()
|
|
||||||
# Create camera output queue and task if needed.
|
# Create camera output queue and task if needed.
|
||||||
if self._params.camera_out_enabled:
|
if self._params.camera_out_enabled:
|
||||||
self._camera_out_queue = asyncio.Queue()
|
self._camera_out_queue = asyncio.Queue()
|
||||||
self._camera_out_task = loop.create_task(self._camera_out_task_handler())
|
self._camera_out_task = self.create_task(self._camera_out_task_handler())
|
||||||
|
|
||||||
async def _cancel_camera_task(self):
|
async def _cancel_camera_task(self):
|
||||||
# Stop camera output task.
|
# Stop camera output task.
|
||||||
if self._camera_out_task and self._params.camera_out_enabled:
|
if self._camera_out_task and self._params.camera_out_enabled:
|
||||||
self._camera_out_task.cancel()
|
await self.cancel_task(self._camera_out_task)
|
||||||
await self._camera_out_task
|
|
||||||
self._camera_out_task = None
|
self._camera_out_task = None
|
||||||
|
|
||||||
async def _draw_image(self, frame: OutputImageRawFrame):
|
async def _draw_image(self, frame: OutputImageRawFrame):
|
||||||
@@ -387,19 +378,14 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
self._camera_out_frame_duration = 1 / self._params.camera_out_framerate
|
self._camera_out_frame_duration = 1 / self._params.camera_out_framerate
|
||||||
self._camera_out_frame_reset = self._camera_out_frame_duration * 5
|
self._camera_out_frame_reset = self._camera_out_frame_duration * 5
|
||||||
while True:
|
while True:
|
||||||
try:
|
if self._params.camera_out_is_live:
|
||||||
if self._params.camera_out_is_live:
|
await self._camera_out_is_live_handler()
|
||||||
await self._camera_out_is_live_handler()
|
elif self._camera_images:
|
||||||
elif self._camera_images:
|
image = next(self._camera_images)
|
||||||
image = next(self._camera_images)
|
await self._draw_image(image)
|
||||||
await self._draw_image(image)
|
await asyncio.sleep(self._camera_out_frame_duration)
|
||||||
await asyncio.sleep(self._camera_out_frame_duration)
|
else:
|
||||||
else:
|
await asyncio.sleep(self._camera_out_frame_duration)
|
||||||
await asyncio.sleep(self._camera_out_frame_duration)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.exception(f"{self} error writing to camera: {e}")
|
|
||||||
|
|
||||||
async def _camera_out_is_live_handler(self):
|
async def _camera_out_is_live_handler(self):
|
||||||
image = await self._camera_out_queue.get()
|
image = await self._camera_out_queue.get()
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from pipecat.audio.filters.base_audio_filter import BaseAudioFilter
|
|||||||
from pipecat.audio.mixers.base_audio_mixer import BaseAudioMixer
|
from pipecat.audio.mixers.base_audio_mixer import BaseAudioMixer
|
||||||
from pipecat.audio.vad.vad_analyzer import VADAnalyzer
|
from pipecat.audio.vad.vad_analyzer import VADAnalyzer
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
|
|
||||||
class TransportParams(BaseModel):
|
class TransportParams(BaseModel):
|
||||||
@@ -46,15 +47,27 @@ class TransportParams(BaseModel):
|
|||||||
class BaseTransport(ABC):
|
class BaseTransport(ABC):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
input_name: str | None = None,
|
*,
|
||||||
output_name: str | None = None,
|
name: Optional[str] = None,
|
||||||
loop: asyncio.AbstractEventLoop | None = None,
|
input_name: Optional[str] = None,
|
||||||
|
output_name: Optional[str] = None,
|
||||||
|
loop: Optional[asyncio.AbstractEventLoop] = None,
|
||||||
):
|
):
|
||||||
|
self._id: int = obj_id()
|
||||||
|
self._name = name or f"{self.__class__.__name__}#{obj_count(self)}"
|
||||||
self._input_name = input_name
|
self._input_name = input_name
|
||||||
self._output_name = output_name
|
self._output_name = output_name
|
||||||
self._loop = loop or asyncio.get_running_loop()
|
self._loop = loop or asyncio.get_running_loop()
|
||||||
self._event_handlers: dict = {}
|
self._event_handlers: dict = {}
|
||||||
|
|
||||||
|
@property
|
||||||
|
def id(self) -> int:
|
||||||
|
return self._id
|
||||||
|
|
||||||
|
@property
|
||||||
|
def name(self) -> str:
|
||||||
|
return self._name
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> FrameProcessor:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
@@ -89,3 +102,6 @@ class BaseTransport(ABC):
|
|||||||
handler(self, *args, **kwargs)
|
handler(self, *args, **kwargs)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Exception in event handler {event_name}: {e}")
|
logger.exception(f"Exception in event handler {event_name}: {e}")
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return self.name
|
||||||
|
|||||||
@@ -98,6 +98,7 @@ class LocalAudioOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
class LocalAudioTransport(BaseTransport):
|
class LocalAudioTransport(BaseTransport):
|
||||||
def __init__(self, params: TransportParams):
|
def __init__(self, params: TransportParams):
|
||||||
|
super().__init__()
|
||||||
self._params = params
|
self._params = params
|
||||||
self._pyaudio = pyaudio.PyAudio()
|
self._pyaudio = pyaudio.PyAudio()
|
||||||
|
|
||||||
|
|||||||
@@ -127,6 +127,7 @@ class TkOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
class TkLocalTransport(BaseTransport):
|
class TkLocalTransport(BaseTransport):
|
||||||
def __init__(self, tk_root: tk.Tk, params: TransportParams):
|
def __init__(self, tk_root: tk.Tk, params: TransportParams):
|
||||||
|
super().__init__()
|
||||||
self._tk_root = tk_root
|
self._tk_root = tk_root
|
||||||
self._params = params
|
self._params = params
|
||||||
self._pyaudio = pyaudio.PyAudio()
|
self._pyaudio = pyaudio.PyAudio()
|
||||||
|
|||||||
@@ -16,6 +16,8 @@ from loguru import logger
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
|
CancelFrame,
|
||||||
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InputAudioRawFrame,
|
InputAudioRawFrame,
|
||||||
OutputAudioRawFrame,
|
OutputAudioRawFrame,
|
||||||
@@ -27,6 +29,7 @@ from pipecat.serializers.base_serializer import FrameSerializer, FrameSerializer
|
|||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
from pipecat.transports.base_output import BaseOutputTransport
|
from pipecat.transports.base_output import BaseOutputTransport
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
|
from pipecat.utils.asyncio import cancel_task
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from fastapi import WebSocket
|
from fastapi import WebSocket
|
||||||
@@ -68,11 +71,17 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
|
|||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
if self._params.session_timeout:
|
if self._params.session_timeout:
|
||||||
self._monitor_websocket_task = self.get_event_loop().create_task(
|
self._monitor_websocket_task = self.create_task(self._monitor_websocket())
|
||||||
self._monitor_websocket()
|
|
||||||
)
|
|
||||||
await self._callbacks.on_client_connected(self._websocket)
|
await self._callbacks.on_client_connected(self._websocket)
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_messages())
|
self._receive_task = self.create_task(self._receive_messages())
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
await super().stop(frame)
|
||||||
|
await cancel_task(self._receive_task)
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
await super().cancel(frame)
|
||||||
|
await cancel_task(self._receive_task)
|
||||||
|
|
||||||
def _iter_data(self) -> typing.AsyncIterator[bytes | str]:
|
def _iter_data(self) -> typing.AsyncIterator[bytes | str]:
|
||||||
if self._params.serializer.type == FrameSerializerType.BINARY:
|
if self._params.serializer.type == FrameSerializerType.BINARY:
|
||||||
@@ -96,11 +105,8 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
async def _monitor_websocket(self):
|
async def _monitor_websocket(self):
|
||||||
"""Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event."""
|
"""Wait for self._params.session_timeout seconds, if the websocket is still open, trigger timeout event."""
|
||||||
try:
|
await asyncio.sleep(self._params.session_timeout)
|
||||||
await asyncio.sleep(self._params.session_timeout)
|
await self._callbacks.on_session_timeout(self._websocket)
|
||||||
await self._callbacks.on_session_timeout(self._websocket)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.info(f"Monitoring task cancelled for: {self._websocket}")
|
|
||||||
|
|
||||||
|
|
||||||
class FastAPIWebsocketOutputTransport(BaseOutputTransport):
|
class FastAPIWebsocketOutputTransport(BaseOutputTransport):
|
||||||
|
|||||||
@@ -71,17 +71,16 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
self._server_task = self.get_event_loop().create_task(self._server_task_handler())
|
self._server_task = self.create_task(self._server_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
self._stop_server_event.set()
|
self._stop_server_event.set()
|
||||||
await self._server_task
|
await self.wait_for_task(self._server_task)
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
self._stop_server_event.set()
|
await self.cancel_task(self._server_task)
|
||||||
await self._server_task
|
|
||||||
|
|
||||||
async def _server_task_handler(self):
|
async def _server_task_handler(self):
|
||||||
logger.info(f"Starting websocket server on {self._host}:{self._port}")
|
logger.info(f"Starting websocket server on {self._host}:{self._port}")
|
||||||
@@ -131,6 +130,7 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
await self._callbacks.on_session_timeout(websocket)
|
await self._callbacks.on_session_timeout(websocket)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
logger.info(f"Monitoring task cancelled for: {websocket.remote_address}")
|
logger.info(f"Monitoring task cancelled for: {websocket.remote_address}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
class WebsocketServerOutputTransport(BaseOutputTransport):
|
class WebsocketServerOutputTransport(BaseOutputTransport):
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ from pipecat.transcriptions.language import Language
|
|||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
from pipecat.transports.base_output import BaseOutputTransport
|
from pipecat.transports.base_output import BaseOutputTransport
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
|
from pipecat.utils.asyncio import cancel_task, create_task
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from daily import CallClient, Daily, EventHandler
|
from daily import CallClient, Daily, EventHandler
|
||||||
@@ -180,6 +181,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
params: DailyParams,
|
params: DailyParams,
|
||||||
callbacks: DailyCallbacks,
|
callbacks: DailyCallbacks,
|
||||||
loop: asyncio.AbstractEventLoop,
|
loop: asyncio.AbstractEventLoop,
|
||||||
|
transport_name: str,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
@@ -193,6 +195,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
self._params: DailyParams = params
|
self._params: DailyParams = params
|
||||||
self._callbacks = callbacks
|
self._callbacks = callbacks
|
||||||
self._loop = loop
|
self._loop = loop
|
||||||
|
self._transport_name = transport_name
|
||||||
|
|
||||||
self._participant_id: str = ""
|
self._participant_id: str = ""
|
||||||
self._video_renderers = {}
|
self._video_renderers = {}
|
||||||
@@ -218,7 +221,11 @@ class DailyTransportClient(EventHandler):
|
|||||||
# future) we will deadlock because completions use event handlers (which
|
# future) we will deadlock because completions use event handlers (which
|
||||||
# are holding the GIL).
|
# are holding the GIL).
|
||||||
self._callback_queue = asyncio.Queue()
|
self._callback_queue = asyncio.Queue()
|
||||||
self._callback_task = self._loop.create_task(self._callback_task_handler())
|
self._callback_task = create_task(
|
||||||
|
self._loop,
|
||||||
|
self._callback_task_handler(),
|
||||||
|
f"{self._transport_name}::DailyTransportClient::callback_task",
|
||||||
|
)
|
||||||
|
|
||||||
self._camera: VirtualCameraDevice | None = None
|
self._camera: VirtualCameraDevice | None = None
|
||||||
if self._params.camera_out_enabled:
|
if self._params.camera_out_enabled:
|
||||||
@@ -469,8 +476,9 @@ class DailyTransportClient(EventHandler):
|
|||||||
return await asyncio.wait_for(future, timeout=10)
|
return await asyncio.wait_for(future, timeout=10)
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
self._callback_task.cancel()
|
if self._callback_task:
|
||||||
await self._callback_task
|
await cancel_task(self._callback_task)
|
||||||
|
self._callback_task = None
|
||||||
# Make sure we don't block the event loop in case `client.release()`
|
# Make sure we don't block the event loop in case `client.release()`
|
||||||
# takes extra time.
|
# takes extra time.
|
||||||
await self._loop.run_in_executor(self._executor, self._cleanup)
|
await self._loop.run_in_executor(self._executor, self._cleanup)
|
||||||
@@ -687,11 +695,8 @@ class DailyTransportClient(EventHandler):
|
|||||||
|
|
||||||
async def _callback_task_handler(self):
|
async def _callback_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
(callback, *args) = await self._callback_queue.get()
|
||||||
(callback, *args) = await self._callback_queue.get()
|
await callback(*args)
|
||||||
await callback(*args)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
|
|
||||||
class DailyInputTransport(BaseInputTransport):
|
class DailyInputTransport(BaseInputTransport):
|
||||||
@@ -721,7 +726,7 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
# Create audio task. It reads audio frames from Daily and push them
|
# Create audio task. It reads audio frames from Daily and push them
|
||||||
# internally for VAD processing.
|
# internally for VAD processing.
|
||||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||||
self._audio_in_task = self.get_event_loop().create_task(self._audio_in_task_handler())
|
self._audio_in_task = self.create_task(self._audio_in_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
# Parent stop.
|
# Parent stop.
|
||||||
@@ -730,8 +735,7 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
await self._client.leave()
|
await self._client.leave()
|
||||||
# Stop audio thread.
|
# Stop audio thread.
|
||||||
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||||
self._audio_in_task.cancel()
|
await self.cancel_task(self._audio_in_task)
|
||||||
await self._audio_in_task
|
|
||||||
self._audio_in_task = None
|
self._audio_in_task = None
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
@@ -741,8 +745,7 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
await self._client.leave()
|
await self._client.leave()
|
||||||
# Stop audio thread.
|
# Stop audio thread.
|
||||||
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||||
self._audio_in_task.cancel()
|
await self.cancel_task(self._audio_in_task)
|
||||||
await self._audio_in_task
|
|
||||||
self._audio_in_task = None
|
self._audio_in_task = None
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
@@ -779,12 +782,9 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
|
|
||||||
async def _audio_in_task_handler(self):
|
async def _audio_in_task_handler(self):
|
||||||
while True:
|
while True:
|
||||||
try:
|
frame = await self._client.read_next_audio_frame()
|
||||||
frame = await self._client.read_next_audio_frame()
|
if frame:
|
||||||
if frame:
|
await self.push_audio_frame(frame)
|
||||||
await self.push_audio_frame(frame)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Camera in
|
# Camera in
|
||||||
@@ -913,7 +913,7 @@ class DailyTransport(BaseTransport):
|
|||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
self._client = DailyTransportClient(
|
self._client = DailyTransportClient(
|
||||||
room_url, token, bot_name, params, callbacks, self._loop
|
room_url, token, bot_name, params, callbacks, self._loop, self.name
|
||||||
)
|
)
|
||||||
self._input: DailyInputTransport | None = None
|
self._input: DailyInputTransport | None = None
|
||||||
self._output: DailyOutputTransport | None = None
|
self._output: DailyOutputTransport | None = None
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from pipecat.processors.frame_processor import FrameDirection
|
|||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
from pipecat.transports.base_output import BaseOutputTransport
|
from pipecat.transports.base_output import BaseOutputTransport
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
|
from pipecat.utils.asyncio import create_task
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from livekit import rtc
|
from livekit import rtc
|
||||||
@@ -72,6 +73,7 @@ class LiveKitTransportClient:
|
|||||||
params: LiveKitParams,
|
params: LiveKitParams,
|
||||||
callbacks: LiveKitCallbacks,
|
callbacks: LiveKitCallbacks,
|
||||||
loop: asyncio.AbstractEventLoop,
|
loop: asyncio.AbstractEventLoop,
|
||||||
|
transport_name: str,
|
||||||
):
|
):
|
||||||
self._url = url
|
self._url = url
|
||||||
self._token = token
|
self._token = token
|
||||||
@@ -79,6 +81,7 @@ class LiveKitTransportClient:
|
|||||||
self._params = params
|
self._params = params
|
||||||
self._callbacks = callbacks
|
self._callbacks = callbacks
|
||||||
self._loop = loop
|
self._loop = loop
|
||||||
|
self._transport_name = transport_name
|
||||||
self._room = rtc.Room(loop=loop)
|
self._room = rtc.Room(loop=loop)
|
||||||
self._participant_id: str = ""
|
self._participant_id: str = ""
|
||||||
self._connected = False
|
self._connected = False
|
||||||
@@ -215,10 +218,18 @@ class LiveKitTransportClient:
|
|||||||
|
|
||||||
# Wrapper methods for event handlers
|
# Wrapper methods for event handlers
|
||||||
def _on_participant_connected_wrapper(self, participant: rtc.RemoteParticipant):
|
def _on_participant_connected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||||
asyncio.create_task(self._async_on_participant_connected(participant))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_participant_connected(participant),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_participant_connected",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_participant_disconnected_wrapper(self, participant: rtc.RemoteParticipant):
|
def _on_participant_disconnected_wrapper(self, participant: rtc.RemoteParticipant):
|
||||||
asyncio.create_task(self._async_on_participant_disconnected(participant))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_participant_disconnected(participant),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_participant_disconnected",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_track_subscribed_wrapper(
|
def _on_track_subscribed_wrapper(
|
||||||
self,
|
self,
|
||||||
@@ -226,7 +237,11 @@ class LiveKitTransportClient:
|
|||||||
publication: rtc.RemoteTrackPublication,
|
publication: rtc.RemoteTrackPublication,
|
||||||
participant: rtc.RemoteParticipant,
|
participant: rtc.RemoteParticipant,
|
||||||
):
|
):
|
||||||
asyncio.create_task(self._async_on_track_subscribed(track, publication, participant))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_track_subscribed(track, publication, participant),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_track_subscribed",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_track_unsubscribed_wrapper(
|
def _on_track_unsubscribed_wrapper(
|
||||||
self,
|
self,
|
||||||
@@ -234,16 +249,32 @@ class LiveKitTransportClient:
|
|||||||
publication: rtc.RemoteTrackPublication,
|
publication: rtc.RemoteTrackPublication,
|
||||||
participant: rtc.RemoteParticipant,
|
participant: rtc.RemoteParticipant,
|
||||||
):
|
):
|
||||||
asyncio.create_task(self._async_on_track_unsubscribed(track, publication, participant))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_track_unsubscribed(track, publication, participant),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_track_unsubscribed",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_data_received_wrapper(self, data: rtc.DataPacket):
|
def _on_data_received_wrapper(self, data: rtc.DataPacket):
|
||||||
asyncio.create_task(self._async_on_data_received(data))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_data_received(data),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_data_received",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_connected_wrapper(self):
|
def _on_connected_wrapper(self):
|
||||||
asyncio.create_task(self._async_on_connected())
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_connected(),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_connected",
|
||||||
|
)
|
||||||
|
|
||||||
def _on_disconnected_wrapper(self):
|
def _on_disconnected_wrapper(self):
|
||||||
asyncio.create_task(self._async_on_disconnected())
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._async_on_disconnected(),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_async_on_disconnected",
|
||||||
|
)
|
||||||
|
|
||||||
# Async methods for event handling
|
# Async methods for event handling
|
||||||
async def _async_on_participant_connected(self, participant: rtc.RemoteParticipant):
|
async def _async_on_participant_connected(self, participant: rtc.RemoteParticipant):
|
||||||
@@ -269,7 +300,11 @@ class LiveKitTransportClient:
|
|||||||
logger.info(f"Audio track subscribed: {track.sid} from participant {participant.sid}")
|
logger.info(f"Audio track subscribed: {track.sid} from participant {participant.sid}")
|
||||||
self._audio_tracks[participant.sid] = track
|
self._audio_tracks[participant.sid] = track
|
||||||
audio_stream = rtc.AudioStream(track)
|
audio_stream = rtc.AudioStream(track)
|
||||||
asyncio.create_task(self._process_audio_stream(audio_stream, participant.sid))
|
create_task(
|
||||||
|
self._loop,
|
||||||
|
self._process_audio_stream(audio_stream, participant.sid),
|
||||||
|
f"{self._transport_name}::LiveKitTransportClient::_process_audio_stream",
|
||||||
|
)
|
||||||
|
|
||||||
async def _async_on_track_unsubscribed(
|
async def _async_on_track_unsubscribed(
|
||||||
self,
|
self,
|
||||||
@@ -319,23 +354,21 @@ class LiveKitInputTransport(BaseInputTransport):
|
|||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
await self._client.connect()
|
await self._client.connect()
|
||||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||||
self._audio_in_task = asyncio.create_task(self._audio_in_task_handler())
|
self._audio_in_task = self.create_task(self._audio_in_task_handler())
|
||||||
logger.info("LiveKitInputTransport started")
|
logger.info("LiveKitInputTransport started")
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
await super().stop(frame)
|
await super().stop(frame)
|
||||||
await self._client.disconnect()
|
await self._client.disconnect()
|
||||||
if self._audio_in_task:
|
if self._audio_in_task:
|
||||||
self._audio_in_task.cancel()
|
await self.cancel_task(self._audio_in_task)
|
||||||
await self._audio_in_task
|
|
||||||
logger.info("LiveKitInputTransport stopped")
|
logger.info("LiveKitInputTransport stopped")
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
await self._client.disconnect()
|
await self._client.disconnect()
|
||||||
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
if self._audio_in_task and (self._params.audio_in_enabled or self._params.vad_enabled):
|
||||||
self._audio_in_task.cancel()
|
await self.cancel_task(self._audio_in_task)
|
||||||
await self._audio_in_task
|
|
||||||
|
|
||||||
def vad_analyzer(self) -> VADAnalyzer | None:
|
def vad_analyzer(self) -> VADAnalyzer | None:
|
||||||
return self._vad_analyzer
|
return self._vad_analyzer
|
||||||
@@ -347,22 +380,16 @@ class LiveKitInputTransport(BaseInputTransport):
|
|||||||
async def _audio_in_task_handler(self):
|
async def _audio_in_task_handler(self):
|
||||||
logger.info("Audio input task started")
|
logger.info("Audio input task started")
|
||||||
while True:
|
while True:
|
||||||
try:
|
audio_data = await self._client.get_next_audio_frame()
|
||||||
audio_data = await self._client.get_next_audio_frame()
|
if audio_data:
|
||||||
if audio_data:
|
audio_frame_event, participant_id = audio_data
|
||||||
audio_frame_event, participant_id = audio_data
|
pipecat_audio_frame = self._convert_livekit_audio_to_pipecat(audio_frame_event)
|
||||||
pipecat_audio_frame = self._convert_livekit_audio_to_pipecat(audio_frame_event)
|
input_audio_frame = InputAudioRawFrame(
|
||||||
input_audio_frame = InputAudioRawFrame(
|
audio=pipecat_audio_frame.audio,
|
||||||
audio=pipecat_audio_frame.audio,
|
sample_rate=pipecat_audio_frame.sample_rate,
|
||||||
sample_rate=pipecat_audio_frame.sample_rate,
|
num_channels=pipecat_audio_frame.num_channels,
|
||||||
num_channels=pipecat_audio_frame.num_channels,
|
)
|
||||||
)
|
await self.push_audio_frame(input_audio_frame)
|
||||||
await self.push_audio_frame(input_audio_frame)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
logger.info("Audio input task cancelled")
|
|
||||||
break
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error in audio input task: {e}")
|
|
||||||
|
|
||||||
def _convert_livekit_audio_to_pipecat(
|
def _convert_livekit_audio_to_pipecat(
|
||||||
self, audio_frame_event: rtc.AudioFrameEvent
|
self, audio_frame_event: rtc.AudioFrameEvent
|
||||||
@@ -451,7 +478,7 @@ class LiveKitTransport(BaseTransport):
|
|||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
self._client = LiveKitTransportClient(
|
self._client = LiveKitTransportClient(
|
||||||
url, token, room_name, self._params, callbacks, self._loop
|
url, token, room_name, self._params, callbacks, self._loop, self.name
|
||||||
)
|
)
|
||||||
self._input: LiveKitInputTransport | None = None
|
self._input: LiveKitInputTransport | None = None
|
||||||
self._output: LiveKitOutputTransport | None = None
|
self._output: LiveKitOutputTransport | None = None
|
||||||
|
|||||||
114
src/pipecat/utils/asyncio.py
Normal file
114
src/pipecat/utils/asyncio.py
Normal file
@@ -0,0 +1,114 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024–2025, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from typing import Coroutine, Optional, Set
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
_TASKS: Set[asyncio.Task] = set()
|
||||||
|
|
||||||
|
|
||||||
|
def create_task(loop: asyncio.AbstractEventLoop, coroutine: Coroutine, name: str) -> asyncio.Task:
|
||||||
|
"""
|
||||||
|
Creates and schedules a new asyncio Task that runs the given coroutine.
|
||||||
|
|
||||||
|
The task is added to a global set of created tasks.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
loop (asyncio.AbstractEventLoop): The event loop to use for creating the task.
|
||||||
|
coroutine (Coroutine): The coroutine to be executed within the task.
|
||||||
|
name (str): The name to assign to the task for identification.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
asyncio.Task: The created task object.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def run_coroutine():
|
||||||
|
try:
|
||||||
|
await coroutine
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.trace(f"{name}: task cancelled")
|
||||||
|
# Re-raise the exception to ensure the task is cancelled.
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"{name}: unexpected exception: {e}")
|
||||||
|
|
||||||
|
task = loop.create_task(run_coroutine())
|
||||||
|
task.set_name(name)
|
||||||
|
_TASKS.add(task)
|
||||||
|
logger.trace(f"{name}: task created")
|
||||||
|
return task
|
||||||
|
|
||||||
|
|
||||||
|
async def wait_for_task(task: asyncio.Task, timeout: Optional[float] = None):
|
||||||
|
"""Wait for an asyncio.Task to complete with optional timeout handling.
|
||||||
|
|
||||||
|
This function awaits the specified asyncio.Task and handles scenarios for
|
||||||
|
timeouts, cancellations, and other exceptions. It also ensures that the task
|
||||||
|
is removed from the set of registered tasks upon completion or failure.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
task (asyncio.Task): The asyncio Task to wait for.
|
||||||
|
timeout (Optional[float], optional): The maximum number of seconds
|
||||||
|
to wait for the task to complete. If None, waits indefinitely.
|
||||||
|
Defaults to None.
|
||||||
|
"""
|
||||||
|
name = task.get_name()
|
||||||
|
try:
|
||||||
|
if timeout:
|
||||||
|
await asyncio.wait_for(task, timeout=timeout)
|
||||||
|
else:
|
||||||
|
await task
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning(f"{name}: timed out waiting for task to finish")
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.trace(f"{name}: unexpected task cancellation (maybe Ctrl-C?)")
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"{name}: unexpected exception while stopping task: {e}")
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
_TASKS.remove(task)
|
||||||
|
except KeyError as e:
|
||||||
|
logger.trace(f"{name}: unable to remove task (already removed?): {e}")
|
||||||
|
|
||||||
|
|
||||||
|
async def cancel_task(task: asyncio.Task, timeout: Optional[float] = None):
|
||||||
|
"""Cancels the given asyncio Task and awaits its completion with an
|
||||||
|
optional timeout.
|
||||||
|
|
||||||
|
This function removes the task from the set of registered tasks upon
|
||||||
|
completion or failure.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
task (asyncio.Task): The task to be cancelled.
|
||||||
|
timeout (Optional[float]): The optional timeout in seconds to wait for the task to cancel.
|
||||||
|
|
||||||
|
"""
|
||||||
|
name = task.get_name()
|
||||||
|
task.cancel()
|
||||||
|
try:
|
||||||
|
if timeout:
|
||||||
|
await asyncio.wait_for(task, timeout=timeout)
|
||||||
|
else:
|
||||||
|
await task
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning(f"{name}: timed out waiting for task to cancel")
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# Here are sure the task is cancelled properly.
|
||||||
|
pass
|
||||||
|
except Exception as e:
|
||||||
|
logger.exception(f"{name}: unexpected exception while cancelling task: {e}")
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
_TASKS.remove(task)
|
||||||
|
except KeyError as e:
|
||||||
|
logger.trace(f"{name}: unable to remove task (already removed?): {e}")
|
||||||
|
|
||||||
|
|
||||||
|
def current_tasks() -> Set[asyncio.Task]:
|
||||||
|
"""Returns the list of currently created/registered tasks."""
|
||||||
|
return _TASKS
|
||||||
@@ -3,6 +3,7 @@
|
|||||||
#
|
#
|
||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import collections
|
import collections
|
||||||
import itertools
|
import itertools
|
||||||
|
|
||||||
|
|||||||
@@ -32,8 +32,7 @@ from pipecat.processors.frameworks.langchain import LangchainProcessor
|
|||||||
class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
class TestLangchain(unittest.IsolatedAsyncioTestCase):
|
||||||
class MockProcessor(FrameProcessor):
|
class MockProcessor(FrameProcessor):
|
||||||
def __init__(self, name):
|
def __init__(self, name):
|
||||||
super().__init__()
|
super().__init__(name=name)
|
||||||
self.name = name
|
|
||||||
self.token: list[str] = []
|
self.token: list[str] = []
|
||||||
# Start collecting tokens when we see the start frame
|
# Start collecting tokens when we see the start frame
|
||||||
self.start_collecting = False
|
self.start_collecting = False
|
||||||
|
|||||||
Reference in New Issue
Block a user