Merge remote-tracking branch 'upstream/main'
This commit is contained in:
11
.github/workflows/tests.yaml
vendored
11
.github/workflows/tests.yaml
vendored
@@ -20,14 +20,17 @@ jobs:
|
|||||||
name: "Unit and Integration Tests"
|
name: "Unit and Integration Tests"
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v4
|
- name: Checkout repo
|
||||||
|
uses: actions/checkout@v4
|
||||||
- name: Set up Python
|
- name: Set up Python
|
||||||
id: setup_python
|
id: setup_python
|
||||||
uses: actions/setup-python@v4
|
uses: actions/setup-python@v4
|
||||||
with:
|
with:
|
||||||
python-version: "3.10"
|
python-version: "3.10"
|
||||||
- name: Install system packages
|
- name: Install system packages
|
||||||
run: sudo apt-get install -y portaudio19-dev
|
id: install_system_packages
|
||||||
|
run: |
|
||||||
|
sudo apt-get install -y portaudio19-dev
|
||||||
- name: Setup virtual environment
|
- name: Setup virtual environment
|
||||||
run: |
|
run: |
|
||||||
python -m venv .venv
|
python -m venv .venv
|
||||||
@@ -35,8 +38,8 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
source .venv/bin/activate
|
source .venv/bin/activate
|
||||||
python -m pip install --upgrade pip
|
python -m pip install --upgrade pip
|
||||||
pip install -r dev-requirements.txt
|
pip install -r test-requirements.txt
|
||||||
- name: Test with pytest
|
- name: Test with pytest
|
||||||
run: |
|
run: |
|
||||||
source .venv/bin/activate
|
source .venv/bin/activate
|
||||||
pytest --doctest-modules --ignore-glob="*to_be_updated*" src tests
|
pytest --ignore-glob="*to_be_updated*" --ignore-glob=*pipeline_source* src tests
|
||||||
|
|||||||
70
CHANGELOG.md
70
CHANGELOG.md
@@ -9,9 +9,30 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Added
|
### Added
|
||||||
|
|
||||||
- A clock can now be specified to `PipelineTask` (defaults to
|
- Added configurable LLM parameters (e.g., temperature, top_p, max_tokens, seed)
|
||||||
`SystemClock`). This clock will be passed to each frame processor via the
|
for OpenAI, Anthropic, and Together AI services along with corresponding
|
||||||
`StartFrame`.
|
setter functions.
|
||||||
|
|
||||||
|
- Added `sample_rate` as a constructor parameter for TTS services.
|
||||||
|
|
||||||
|
- Pipecat has a pipeline-based architecture. The pipeline consists of frame
|
||||||
|
processors linked to each other. The elements traveling across the pipeline
|
||||||
|
are called frames.
|
||||||
|
|
||||||
|
To have a deterministic behavior the frames traveling through the pipeline
|
||||||
|
should always be ordered, except system frames which are out-of-band
|
||||||
|
frames. To achieve that, each frame processor should only output frames from a
|
||||||
|
single task.
|
||||||
|
|
||||||
|
In this version we introduce synchronous and asynchronous frame
|
||||||
|
processors. The synchronous processors push output frames from the same task
|
||||||
|
that they receive input frames, and therefore only pushing frames from one
|
||||||
|
task. Asynchronous frame processors can have internal tasks to perform things
|
||||||
|
asynchronously (e.g. receiving data from a websocket) but they also have a
|
||||||
|
single task where they push frames from.
|
||||||
|
|
||||||
|
By default, frame processors are synchronous. To change a frame processor to
|
||||||
|
asynchronous you only need to pass `sync=False` to the base class constructor.
|
||||||
|
|
||||||
- Added pipeline clocks. A pipeline clock is used by the output transport to
|
- Added pipeline clocks. A pipeline clock is used by the output transport to
|
||||||
know when a frame needs to be presented. For that, all frames now have an
|
know when a frame needs to be presented. For that, all frames now have an
|
||||||
@@ -19,6 +40,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
clock implementation `SystemClock` and the `pts` field is currently only used
|
clock implementation `SystemClock` and the `pts` field is currently only used
|
||||||
for `TextFrame`s (audio and image frames will be next).
|
for `TextFrame`s (audio and image frames will be next).
|
||||||
|
|
||||||
|
- A clock can now be specified to `PipelineTask` (defaults to
|
||||||
|
`SystemClock`). This clock will be passed to each frame processor via the
|
||||||
|
`StartFrame`.
|
||||||
|
|
||||||
|
- Added `CartesiaHttpTTSService`. This is a synchronous frame processor
|
||||||
|
(i.e. given an input text frame it will wait for the whole output before
|
||||||
|
returning).
|
||||||
|
|
||||||
- `DailyTransport` now supports setting the audio bitrate to improve audio
|
- `DailyTransport` now supports setting the audio bitrate to improve audio
|
||||||
quality through the `DailyParams.audio_out_bitrate` parameter. The new
|
quality through the `DailyParams.audio_out_bitrate` parameter. The new
|
||||||
default is 96kbps.
|
default is 96kbps.
|
||||||
@@ -40,6 +69,29 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Changed
|
### Changed
|
||||||
|
|
||||||
|
- We now distinguish between input and output audio and image frames. We
|
||||||
|
introduce `InputAudioRawFrame`, `OutputAudioRawFrame`, `InputImageRawFrame`
|
||||||
|
and `OutputImageRawFrame` (and other subclasses of those). The input frames
|
||||||
|
usually come from an input transport and are meant to be processed inside the
|
||||||
|
pipeline to generate new frames. However, the input frames will not be sent
|
||||||
|
through an output transport. The output frames can also be processed by any
|
||||||
|
frame processor in the pipeline and they are allowed to be sent by the output
|
||||||
|
transport.
|
||||||
|
|
||||||
|
- `ParallelTask` has been renamed to `SyncParallelPipeline`. A
|
||||||
|
`SyncParallelPipeline` is a frame processor that contains a list of different
|
||||||
|
pipelines to be executed concurrently. The difference between a
|
||||||
|
`SyncParallelPipeline` and a `ParallelPipeline` is that, given an input frame,
|
||||||
|
the `SyncParallelPipeline` will wait for all the internal pipelines to
|
||||||
|
complete. This is achieved by ensuring all the processors in each of the
|
||||||
|
internal pipelines are synchronous.
|
||||||
|
|
||||||
|
- `StartFrame` is back a system frame so we make sure it's processed immediately
|
||||||
|
by all processors. `EndFrame` stays a control frame since it needs to be
|
||||||
|
ordered allowing the frames in the pipeline to be processed.
|
||||||
|
|
||||||
|
- Updated `MoondreamService` revision to `2024-08-26`.
|
||||||
|
|
||||||
- `CartesiaTTSService` and `ElevenLabsTTSService` now add presentation
|
- `CartesiaTTSService` and `ElevenLabsTTSService` now add presentation
|
||||||
timestamps to their text output. This allows the output transport to push the
|
timestamps to their text output. This allows the output transport to push the
|
||||||
text frames downstream at almost the same time the words are spoken. We say
|
text frames downstream at almost the same time the words are spoken. We say
|
||||||
@@ -60,6 +112,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
|
- Fixed a `BaseOutputTransport` issue that would stop audio and video rendering
|
||||||
|
tasks (after receiving and `EndFrame`) before the internal queue was emptied,
|
||||||
|
causing the pipeline to finish prematurely.
|
||||||
|
|
||||||
- `StartFrame` should be the first frame every processor receives to avoid
|
- `StartFrame` should be the first frame every processor receives to avoid
|
||||||
situations where things are not initialized (because initialization happens on
|
situations where things are not initialized (because initialization happens on
|
||||||
`StartFrame`) and other frames come in resulting in undesired behavior.
|
`StartFrame`) and other frames come in resulting in undesired behavior.
|
||||||
@@ -293,7 +349,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
- It is now possible to specify a Silero VAD version when using `SileroVADAnalyzer`
|
- It is now possible to specify a Silero VAD version when using `SileroVADAnalyzer`
|
||||||
or `SileroVAD`.
|
or `SileroVAD`.
|
||||||
|
|
||||||
- Added `AysncFrameProcessor` and `AsyncAIService`. Some services like
|
- Added `AysncFrameProcessor` and `AsyncAIService`. Some services like
|
||||||
`DeepgramSTTService` need to process things asynchronously. For example, audio
|
`DeepgramSTTService` need to process things asynchronously. For example, audio
|
||||||
is sent to Deepgram but transcriptions are not returned immediately. In these
|
is sent to Deepgram but transcriptions are not returned immediately. In these
|
||||||
cases we still require all frames (except system frames) to be pushed
|
cases we still require all frames (except system frames) to be pushed
|
||||||
@@ -310,7 +366,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
- `WhisperSTTService` model can now also be a string.
|
- `WhisperSTTService` model can now also be a string.
|
||||||
|
|
||||||
- Added missing * keyword separators in services.
|
- Added missing \* keyword separators in services.
|
||||||
|
|
||||||
### Fixed
|
### Fixed
|
||||||
|
|
||||||
@@ -387,7 +443,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
- Added new `TwilioFrameSerializer`. This is a new serializer that knows how to
|
- Added new `TwilioFrameSerializer`. This is a new serializer that knows how to
|
||||||
serialize and deserialize audio frames from Twilio.
|
serialize and deserialize audio frames from Twilio.
|
||||||
|
|
||||||
- Added Daily transport event: `on_dialout_answered`. See
|
- Added Daily transport event: `on_dialout_answered`. See
|
||||||
https://reference-python.daily.co/api_reference.html#daily.EventHandler
|
https://reference-python.daily.co/api_reference.html#daily.EventHandler
|
||||||
|
|
||||||
- Added new `AzureSTTService`. This allows you to use Azure Speech-To-Text.
|
- Added new `AzureSTTService`. This allows you to use Azure Speech-To-Text.
|
||||||
@@ -627,7 +683,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
- Added Daily transport support for dial-in use cases.
|
- Added Daily transport support for dial-in use cases.
|
||||||
|
|
||||||
- Added Daily transport events: `on_dialout_connected`, `on_dialout_stopped`,
|
- Added Daily transport events: `on_dialout_connected`, `on_dialout_stopped`,
|
||||||
`on_dialout_error` and `on_dialout_warning`. See
|
`on_dialout_error` and `on_dialout_warning`. See
|
||||||
https://reference-python.daily.co/api_reference.html#daily.EventHandler
|
https://reference-python.daily.co/api_reference.html#daily.EventHandler
|
||||||
|
|
||||||
## [0.0.21] - 2024-05-22
|
## [0.0.21] - 2024-05-22
|
||||||
|
|||||||
@@ -165,7 +165,7 @@ pip install "path_to_this_repo[option,...]"
|
|||||||
From the root directory, run:
|
From the root directory, run:
|
||||||
|
|
||||||
```shell
|
```shell
|
||||||
pytest --doctest-modules --ignore-glob="*to_be_updated*" src tests
|
pytest --doctest-modules --ignore-glob="*to_be_updated*" --ignore-glob=*pipeline_source* src tests
|
||||||
```
|
```
|
||||||
|
|
||||||
## Setting up your editor
|
## Setting up your editor
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
pipecat-ai[daily,openai,silero]
|
pipecat-ai[daily,elevenlabs,openai,silero]
|
||||||
fastapi
|
fastapi
|
||||||
uvicorn
|
uvicorn
|
||||||
python-dotenv
|
python-dotenv
|
||||||
|
|||||||
@@ -9,11 +9,11 @@ import aiohttp
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
from pipecat.frames.frames import TextFrame
|
from pipecat.frames.frames import EndFrame, TextFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.services.cartesia import CartesiaTTSService
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
|
||||||
from runner import configure
|
from runner import configure
|
||||||
@@ -34,7 +34,7 @@ async def main():
|
|||||||
transport = DailyTransport(
|
transport = DailyTransport(
|
||||||
room_url, None, "Say One Thing", DailyParams(audio_out_enabled=True))
|
room_url, None, "Say One Thing", DailyParams(audio_out_enabled=True))
|
||||||
|
|
||||||
tts = CartesiaTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("CARTESIA_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
@@ -48,7 +48,7 @@ async def main():
|
|||||||
@transport.event_handler("on_participant_joined")
|
@transport.event_handler("on_participant_joined")
|
||||||
async def on_new_participant_joined(transport, participant):
|
async def on_new_participant_joined(transport, participant):
|
||||||
participant_name = participant["info"]["userName"] or ''
|
participant_name = participant["info"]["userName"] or ''
|
||||||
await task.queue_frame(TextFrame(f"Hello there, {participant_name}!"))
|
await task.queue_frames([TextFrame(f"Hello there, {participant_name}!"), EndFrame()])
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|
||||||
|
|||||||
@@ -9,11 +9,11 @@ import aiohttp
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
from pipecat.frames.frames import LLMMessagesFrame
|
from pipecat.frames.frames import EndFrame, LLMMessagesFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
from pipecat.services.cartesia import CartesiaTTSService
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
|
||||||
@@ -38,7 +38,7 @@ async def main():
|
|||||||
"Say One Thing From an LLM",
|
"Say One Thing From an LLM",
|
||||||
DailyParams(audio_out_enabled=True))
|
DailyParams(audio_out_enabled=True))
|
||||||
|
|
||||||
tts = CartesiaTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("CARTESIA_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
@@ -59,7 +59,7 @@ async def main():
|
|||||||
|
|
||||||
@transport.event_handler("on_first_participant_joined")
|
@transport.event_handler("on_first_participant_joined")
|
||||||
async def on_first_participant_joined(transport, participant):
|
async def on_first_participant_joined(transport, participant):
|
||||||
await task.queue_frame(LLMMessagesFrame(messages))
|
await task.queue_frames([LLMMessagesFrame(messages), EndFrame()])
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,10 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
|
#
|
||||||
|
# This example broken on latest pipecat and needs updating.
|
||||||
|
#
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
|
|||||||
@@ -14,21 +14,18 @@ from dataclasses import dataclass
|
|||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AppFrame,
|
AppFrame,
|
||||||
Frame,
|
Frame,
|
||||||
ImageRawFrame,
|
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
TextFrame
|
TextFrame
|
||||||
)
|
)
|
||||||
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.sync_parallel_pipeline import SyncParallelPipeline
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
from pipecat.pipeline.parallel_task import ParallelTask
|
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.aggregators.gated import GatedAggregator
|
|
||||||
from pipecat.processors.aggregators.llm_response import LLMFullResponseAggregator
|
|
||||||
from pipecat.processors.aggregators.sentence import SentenceAggregator
|
from pipecat.processors.aggregators.sentence import SentenceAggregator
|
||||||
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
|
||||||
from pipecat.services.fal import FalImageGenService
|
from pipecat.services.fal import FalImageGenService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
|
||||||
@@ -88,9 +85,9 @@ async def main():
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
@@ -105,24 +102,23 @@ async def main():
|
|||||||
key=os.getenv("FAL_KEY"),
|
key=os.getenv("FAL_KEY"),
|
||||||
)
|
)
|
||||||
|
|
||||||
gated_aggregator = GatedAggregator(
|
|
||||||
gate_open_fn=lambda frame: isinstance(frame, ImageRawFrame),
|
|
||||||
gate_close_fn=lambda frame: isinstance(frame, LLMFullResponseStartFrame),
|
|
||||||
start_open=False
|
|
||||||
)
|
|
||||||
|
|
||||||
sentence_aggregator = SentenceAggregator()
|
sentence_aggregator = SentenceAggregator()
|
||||||
month_prepender = MonthPrepender()
|
month_prepender = MonthPrepender()
|
||||||
llm_full_response_aggregator = LLMFullResponseAggregator()
|
|
||||||
|
|
||||||
|
# With `SyncParallelPipeline` we synchronize audio and images by pushing
|
||||||
|
# them basically in order (e.g. I1 A1 A1 A1 I2 A2 A2 A2 A2 I3 A3). To do
|
||||||
|
# that, each pipeline runs concurrently and `SyncParallelPipeline` will
|
||||||
|
# wait for the input frame to be processed.
|
||||||
|
#
|
||||||
|
# Note that `SyncParallelPipeline` requires all processors in it to be
|
||||||
|
# synchronous (which is the default for most processors).
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline([
|
||||||
llm, # LLM
|
llm, # LLM
|
||||||
sentence_aggregator, # Aggregates LLM output into full sentences
|
sentence_aggregator, # Aggregates LLM output into full sentences
|
||||||
ParallelTask( # Run pipelines in parallel aggregating the result
|
SyncParallelPipeline( # Run pipelines in parallel aggregating the result
|
||||||
[month_prepender, tts], # Create "Month: sentence" and output audio
|
[month_prepender, tts], # Create "Month: sentence" and output audio
|
||||||
[llm_full_response_aggregator, imagegen] # Aggregate full LLM response
|
[imagegen] # Generate image
|
||||||
),
|
),
|
||||||
gated_aggregator, # Queues everything until an image is available
|
|
||||||
transport.output() # Transport output
|
transport.output() # Transport output
|
||||||
])
|
])
|
||||||
|
|
||||||
|
|||||||
@@ -11,18 +11,24 @@ import sys
|
|||||||
|
|
||||||
import tkinter as tk
|
import tkinter as tk
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, Frame, URLImageRawFrame, LLMMessagesFrame, TextFrame
|
from pipecat.frames.frames import (
|
||||||
from pipecat.pipeline.parallel_pipeline import ParallelPipeline
|
Frame,
|
||||||
|
OutputAudioRawFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
|
URLImageRawFrame,
|
||||||
|
LLMMessagesFrame,
|
||||||
|
TextFrame)
|
||||||
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.sync_parallel_pipeline import SyncParallelPipeline
|
||||||
from pipecat.pipeline.task import PipelineTask
|
from pipecat.pipeline.task import PipelineTask
|
||||||
from pipecat.processors.aggregators.llm_response import LLMFullResponseAggregator
|
from pipecat.processors.aggregators.sentence import SentenceAggregator
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
|
||||||
from pipecat.services.fal import FalImageGenService
|
from pipecat.services.fal import FalImageGenService
|
||||||
from pipecat.transports.base_transport import TransportParams
|
from pipecat.transports.base_transport import TransportParams
|
||||||
from pipecat.transports.local.tk import TkLocalTransport
|
from pipecat.transports.local.tk import TkLocalTransport, TkOutputTransport
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -60,13 +66,14 @@ async def main():
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.audio = bytearray()
|
self.audio = bytearray()
|
||||||
|
self.frame = None
|
||||||
|
|
||||||
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)
|
||||||
|
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, TTSAudioRawFrame):
|
||||||
self.audio.extend(frame.audio)
|
self.audio.extend(frame.audio)
|
||||||
self.frame = AudioRawFrame(
|
self.frame = OutputAudioRawFrame(
|
||||||
bytes(self.audio), frame.sample_rate, frame.num_channels)
|
bytes(self.audio), frame.sample_rate, frame.num_channels)
|
||||||
|
|
||||||
class ImageGrabber(FrameProcessor):
|
class ImageGrabber(FrameProcessor):
|
||||||
@@ -84,9 +91,10 @@ async def main():
|
|||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
model="gpt-4o")
|
model="gpt-4o")
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id=os.getenv("ELEVENLABS_VOICE_ID"))
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
|
)
|
||||||
|
|
||||||
imagegen = FalImageGenService(
|
imagegen = FalImageGenService(
|
||||||
params=FalImageGenService.InputParams(
|
params=FalImageGenService.InputParams(
|
||||||
@@ -95,7 +103,7 @@ async def main():
|
|||||||
aiohttp_session=session,
|
aiohttp_session=session,
|
||||||
key=os.getenv("FAL_KEY"))
|
key=os.getenv("FAL_KEY"))
|
||||||
|
|
||||||
aggregator = LLMFullResponseAggregator()
|
sentence_aggregator = SentenceAggregator()
|
||||||
|
|
||||||
description = ImageDescription()
|
description = ImageDescription()
|
||||||
|
|
||||||
@@ -103,12 +111,22 @@ async def main():
|
|||||||
|
|
||||||
image_grabber = ImageGrabber()
|
image_grabber = ImageGrabber()
|
||||||
|
|
||||||
|
# With `SyncParallelPipeline` we synchronize audio and images by
|
||||||
|
# pushing them basically in order (e.g. I1 A1 A1 A1 I2 A2 A2 A2 A2
|
||||||
|
# I3 A3). To do that, each pipeline runs concurrently and
|
||||||
|
# `SyncParallelPipeline` will wait for the input frame to be
|
||||||
|
# processed.
|
||||||
|
#
|
||||||
|
# Note that `SyncParallelPipeline` requires all processors in it to
|
||||||
|
# be synchronous (which is the default for most processors).
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline([
|
||||||
llm,
|
llm, # LLM
|
||||||
aggregator,
|
sentence_aggregator, # Aggregates LLM output into full sentences
|
||||||
description,
|
description, # Store sentence
|
||||||
ParallelPipeline([tts, audio_grabber],
|
SyncParallelPipeline(
|
||||||
[imagegen, image_grabber])
|
[tts, audio_grabber], # Generate and store audio for the given sentence
|
||||||
|
[imagegen, image_grabber] # Generate and storeimage for the given sentence
|
||||||
|
)
|
||||||
])
|
])
|
||||||
|
|
||||||
task = PipelineTask(pipeline)
|
task = PipelineTask(pipeline)
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import os
|
|||||||
import sys
|
import sys
|
||||||
|
|
||||||
from pipecat.frames.frames import Frame, LLMMessagesFrame, MetricsFrame
|
from pipecat.frames.frames import Frame, LLMMessagesFrame, MetricsFrame
|
||||||
|
from pipecat.metrics.metrics import TTFBMetricsData, ProcessingMetricsData, LLMUsageMetricsData, TTSUsageMetricsData
|
||||||
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
|
||||||
@@ -37,8 +38,19 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
class MetricsLogger(FrameProcessor):
|
class MetricsLogger(FrameProcessor):
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
if isinstance(frame, MetricsFrame):
|
if isinstance(frame, MetricsFrame):
|
||||||
print(
|
for d in frame.data:
|
||||||
f"!!! MetricsFrame: {frame}, ttfb: {frame.ttfb}, processing: {frame.processing}, tokens: {frame.tokens}, characters: {frame.characters}")
|
if isinstance(d, TTFBMetricsData):
|
||||||
|
print(f"!!! MetricsFrame: {frame}, ttfb: {d.value}")
|
||||||
|
elif isinstance(d, ProcessingMetricsData):
|
||||||
|
print(f"!!! MetricsFrame: {frame}, processing: {d.value}")
|
||||||
|
elif isinstance(d, LLMUsageMetricsData):
|
||||||
|
tokens = d.value
|
||||||
|
print(
|
||||||
|
f"!!! MetricsFrame: {frame}, tokens: {
|
||||||
|
tokens.prompt_tokens}, characters: {
|
||||||
|
tokens.completion_tokens}")
|
||||||
|
elif isinstance(d, TTSUsageMetricsData):
|
||||||
|
print(f"!!! MetricsFrame: {frame}, characters: {d.value}")
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
@@ -90,11 +102,6 @@ async def main():
|
|||||||
])
|
])
|
||||||
|
|
||||||
task = PipelineTask(pipeline)
|
task = PipelineTask(pipeline)
|
||||||
task = PipelineTask(pipeline, PipelineParams(
|
|
||||||
allow_interruptions=True,
|
|
||||||
enable_metrics=True,
|
|
||||||
report_only_initial_ttfb=False,
|
|
||||||
))
|
|
||||||
|
|
||||||
@transport.event_handler("on_first_participant_joined")
|
@transport.event_handler("on_first_participant_joined")
|
||||||
async def on_first_participant_joined(transport, participant):
|
async def on_first_participant_joined(transport, participant):
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import sys
|
|||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from pipecat.frames.frames import ImageRawFrame, Frame, SystemFrame, TextFrame
|
from pipecat.frames.frames import Frame, OutputImageRawFrame, SystemFrame, TextFrame
|
||||||
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
|
||||||
@@ -20,8 +20,8 @@ from pipecat.processors.aggregators.llm_response import (
|
|||||||
LLMUserResponseAggregator,
|
LLMUserResponseAggregator,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
|
||||||
from pipecat.transports.services.daily import DailyTransport
|
from pipecat.transports.services.daily import DailyTransport
|
||||||
from pipecat.vad.silero import SileroVADAnalyzer
|
from pipecat.vad.silero import SileroVADAnalyzer
|
||||||
|
|
||||||
@@ -52,9 +52,16 @@ class ImageSyncAggregator(FrameProcessor):
|
|||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
if not isinstance(frame, SystemFrame) and direction == FrameDirection.DOWNSTREAM:
|
if not isinstance(frame, SystemFrame) and direction == FrameDirection.DOWNSTREAM:
|
||||||
await self.push_frame(ImageRawFrame(image=self._speaking_image_bytes, size=(1024, 1024), format=self._speaking_image_format))
|
await self.push_frame(OutputImageRawFrame(
|
||||||
|
image=self._speaking_image_bytes,
|
||||||
|
size=(1024, 1024),
|
||||||
|
format=self._speaking_image_format)
|
||||||
|
)
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
await self.push_frame(ImageRawFrame(image=self._waiting_image_bytes, size=(1024, 1024), format=self._waiting_image_format))
|
await self.push_frame(OutputImageRawFrame(
|
||||||
|
image=self._waiting_image_bytes,
|
||||||
|
size=(1024, 1024),
|
||||||
|
format=self._waiting_image_format))
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
@@ -78,9 +85,9 @@ async def main():
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
|
|||||||
@@ -5,26 +5,27 @@
|
|||||||
#
|
#
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import aiohttp
|
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from loguru import logger
|
||||||
|
from runner import configure
|
||||||
|
|
||||||
from pipecat.frames.frames import LLMMessagesFrame
|
from pipecat.frames.frames import LLMMessagesFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
from pipecat.processors.aggregators.llm_response import (
|
from pipecat.processors.aggregators.llm_response import (
|
||||||
LLMAssistantResponseAggregator, LLMUserResponseAggregator)
|
LLMAssistantResponseAggregator,
|
||||||
|
LLMUserResponseAggregator,
|
||||||
|
)
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
from pipecat.vad.silero import SileroVADAnalyzer
|
from pipecat.vad.silero import SileroVADAnalyzer
|
||||||
|
|
||||||
from runner import configure
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
load_dotenv(override=True)
|
load_dotenv(override=True)
|
||||||
|
|
||||||
logger.remove(0)
|
logger.remove(0)
|
||||||
@@ -43,8 +44,8 @@ async def main():
|
|||||||
audio_out_enabled=True,
|
audio_out_enabled=True,
|
||||||
transcription_enabled=True,
|
transcription_enabled=True,
|
||||||
vad_enabled=True,
|
vad_enabled=True,
|
||||||
vad_analyzer=SileroVADAnalyzer()
|
vad_analyzer=SileroVADAnalyzer(),
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = ElevenLabsTTSService(
|
||||||
@@ -52,9 +53,7 @@ async def main():
|
|||||||
voice_id=os.getenv("ELEVENLABS_VOICE_ID", ""),
|
voice_id=os.getenv("ELEVENLABS_VOICE_ID", ""),
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o")
|
||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
|
||||||
model="gpt-4o")
|
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
@@ -66,28 +65,32 @@ async def main():
|
|||||||
tma_in = LLMUserResponseAggregator(messages)
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
tma_out = LLMAssistantResponseAggregator(messages)
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
|
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline(
|
||||||
transport.input(), # Transport user input
|
[
|
||||||
tma_in, # User responses
|
transport.input(), # Transport user input
|
||||||
llm, # LLM
|
tma_in, # User responses
|
||||||
tts, # TTS
|
llm, # LLM
|
||||||
transport.output(), # Transport bot output
|
tts, # TTS
|
||||||
tma_out # Assistant spoken responses
|
transport.output(), # Transport bot output
|
||||||
])
|
tma_out, # Assistant spoken responses
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
task = PipelineTask(pipeline, PipelineParams(
|
task = PipelineTask(
|
||||||
allow_interruptions=True,
|
pipeline,
|
||||||
enable_metrics=True,
|
PipelineParams(
|
||||||
enable_usage_metrics=True,
|
allow_interruptions=True,
|
||||||
report_only_initial_ttfb=True,
|
enable_metrics=True,
|
||||||
))
|
enable_usage_metrics=True,
|
||||||
|
report_only_initial_ttfb=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
@transport.event_handler("on_first_participant_joined")
|
@transport.event_handler("on_first_participant_joined")
|
||||||
async def on_first_participant_joined(transport, participant):
|
async def on_first_participant_joined(transport, participant):
|
||||||
transport.capture_participant_transcription(participant["id"])
|
transport.capture_participant_transcription(participant["id"])
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
messages.append(
|
messages.append({"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
{"role": "system", "content": "Please introduce yourself to the user."})
|
|
||||||
await task.queue_frames([LLMMessagesFrame(messages)])
|
await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|||||||
102
examples/foundational/07l-interruptible-together.py
Normal file
102
examples/foundational/07l-interruptible-together.py
Normal file
@@ -0,0 +1,102 @@
|
|||||||
|
#
|
||||||
|
# Copyright (c) 2024, Daily
|
||||||
|
#
|
||||||
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
|
#
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import aiohttp
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
from pipecat.frames.frames import LLMMessagesFrame
|
||||||
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
|
from pipecat.processors.aggregators.llm_response import (
|
||||||
|
LLMAssistantResponseAggregator, LLMUserResponseAggregator)
|
||||||
|
from pipecat.services.cartesia import CartesiaTTSService
|
||||||
|
from pipecat.services.together import TogetherLLMService
|
||||||
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
from pipecat.vad.silero import SileroVADAnalyzer
|
||||||
|
|
||||||
|
from runner import configure
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
load_dotenv(override=True)
|
||||||
|
|
||||||
|
logger.remove(0)
|
||||||
|
logger.add(sys.stderr, level="DEBUG")
|
||||||
|
|
||||||
|
|
||||||
|
async def main():
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
(room_url, token) = await configure(session)
|
||||||
|
|
||||||
|
transport = DailyTransport(
|
||||||
|
room_url,
|
||||||
|
token,
|
||||||
|
"Respond bot",
|
||||||
|
DailyParams(
|
||||||
|
audio_out_enabled=True,
|
||||||
|
transcription_enabled=True,
|
||||||
|
vad_enabled=True,
|
||||||
|
vad_analyzer=SileroVADAnalyzer()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
tts = CartesiaTTSService(
|
||||||
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
|
)
|
||||||
|
|
||||||
|
llm = TogetherLLMService(
|
||||||
|
api_key=os.getenv("TOGETHER_API_KEY"),
|
||||||
|
model=os.getenv("TOGETHER_MODEL"),
|
||||||
|
params=TogetherLLMService.InputParams(
|
||||||
|
temperature=1.0,
|
||||||
|
top_p=0.9,
|
||||||
|
top_k=40,
|
||||||
|
extra={
|
||||||
|
"frequency_penalty": 2.0,
|
||||||
|
"presence_penalty": 0.0,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
messages = [
|
||||||
|
{
|
||||||
|
"role": "system",
|
||||||
|
"content": "You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
|
|
||||||
|
pipeline = Pipeline([
|
||||||
|
transport.input(), # Transport user input
|
||||||
|
tma_in, # User responses
|
||||||
|
llm, # LLM
|
||||||
|
tts, # TTS
|
||||||
|
transport.output(), # Transport bot output
|
||||||
|
tma_out # Assistant spoken responses
|
||||||
|
])
|
||||||
|
|
||||||
|
task = PipelineTask(pipeline, PipelineParams(allow_interruptions=True))
|
||||||
|
|
||||||
|
@transport.event_handler("on_first_participant_joined")
|
||||||
|
async def on_first_participant_joined(transport, participant):
|
||||||
|
transport.capture_participant_transcription(participant["id"])
|
||||||
|
# Kick off the conversation.
|
||||||
|
await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
|
runner = PipelineRunner()
|
||||||
|
|
||||||
|
await runner.run(task)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
@@ -3,14 +3,14 @@ import aiohttp
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from pipecat.pipeline.aggregators import SentenceAggregator
|
from pipecat.processors.aggregators import SentenceAggregator
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
|
||||||
from pipecat.transports.daily_transport import DailyTransport
|
from pipecat.transports.services.daily import DailyTransport
|
||||||
from pipecat.services.azure_ai_services import AzureLLMService, AzureTTSService
|
from pipecat.services.azure import AzureLLMService, AzureTTSService
|
||||||
from pipecat.services.elevenlabs_ai_services import ElevenLabsTTSService
|
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
||||||
from pipecat.services.fal_ai_services import FalImageGenService
|
from pipecat.services.fal import FalImageGenService
|
||||||
from pipecat.pipeline.frames import AudioFrame, EndFrame, ImageFrame, LLMMessagesFrame, TextFrame
|
from pipecat.frames.frames import AudioFrame, EndFrame, ImageFrame, LLMMessagesFrame, TextFrame
|
||||||
|
|
||||||
from runner import configure
|
from runner import configure
|
||||||
|
|
||||||
|
|||||||
@@ -8,9 +8,11 @@ import aiohttp
|
|||||||
import asyncio
|
import asyncio
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
from pipecat.frames.frames import Frame, InputAudioRawFrame, InputImageRawFrame, OutputAudioRawFrame, OutputImageRawFrame
|
||||||
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
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.transports.services.daily import DailyTransport, DailyParams
|
from pipecat.transports.services.daily import DailyTransport, DailyParams
|
||||||
|
|
||||||
from runner import configure
|
from runner import configure
|
||||||
@@ -24,6 +26,27 @@ logger.remove(0)
|
|||||||
logger.add(sys.stderr, level="DEBUG")
|
logger.add(sys.stderr, level="DEBUG")
|
||||||
|
|
||||||
|
|
||||||
|
class MirrorProcessor(FrameProcessor):
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if isinstance(frame, InputAudioRawFrame):
|
||||||
|
await self.push_frame(OutputAudioRawFrame(
|
||||||
|
audio=frame.audio,
|
||||||
|
sample_rate=frame.sample_rate,
|
||||||
|
num_channels=frame.num_channels)
|
||||||
|
)
|
||||||
|
elif isinstance(frame, InputImageRawFrame):
|
||||||
|
await self.push_frame(OutputImageRawFrame(
|
||||||
|
image=frame.image,
|
||||||
|
size=frame.size,
|
||||||
|
format=frame.format)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
async def main():
|
async def main():
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
(room_url, token) = await configure(session)
|
(room_url, token) = await configure(session)
|
||||||
@@ -44,7 +67,7 @@ async def main():
|
|||||||
async def on_first_participant_joined(transport, participant):
|
async def on_first_participant_joined(transport, participant):
|
||||||
transport.capture_participant_video(participant["id"])
|
transport.capture_participant_video(participant["id"])
|
||||||
|
|
||||||
pipeline = Pipeline([transport.input(), transport.output()])
|
pipeline = Pipeline([transport.input(), MirrorProcessor(), transport.output()])
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|
||||||
|
|||||||
@@ -10,9 +10,11 @@ import sys
|
|||||||
|
|
||||||
import tkinter as tk
|
import tkinter as tk
|
||||||
|
|
||||||
|
from pipecat.frames.frames import Frame, InputAudioRawFrame, InputImageRawFrame, OutputAudioRawFrame, OutputImageRawFrame
|
||||||
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
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.transports.base_transport import TransportParams
|
from pipecat.transports.base_transport import TransportParams
|
||||||
from pipecat.transports.local.tk import TkLocalTransport
|
from pipecat.transports.local.tk import TkLocalTransport
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
@@ -27,6 +29,25 @@ load_dotenv(override=True)
|
|||||||
logger.remove(0)
|
logger.remove(0)
|
||||||
logger.add(sys.stderr, level="DEBUG")
|
logger.add(sys.stderr, level="DEBUG")
|
||||||
|
|
||||||
|
class MirrorProcessor(FrameProcessor):
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if isinstance(frame, InputAudioRawFrame):
|
||||||
|
await self.push_frame(OutputAudioRawFrame(
|
||||||
|
audio=frame.audio,
|
||||||
|
sample_rate=frame.sample_rate,
|
||||||
|
num_channels=frame.num_channels)
|
||||||
|
)
|
||||||
|
elif isinstance(frame, InputImageRawFrame):
|
||||||
|
await self.push_frame(OutputImageRawFrame(
|
||||||
|
image=frame.image,
|
||||||
|
size=frame.size,
|
||||||
|
format=frame.format)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def main():
|
async def main():
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
@@ -52,7 +73,7 @@ async def main():
|
|||||||
async def on_first_participant_joined(transport, participant):
|
async def on_first_participant_joined(transport, participant):
|
||||||
transport.capture_participant_video(participant["id"])
|
transport.capture_participant_video(participant["id"])
|
||||||
|
|
||||||
pipeline = Pipeline([daily_transport.input(), tk_transport.output()])
|
pipeline = Pipeline([daily_transport.input(), MirrorProcessor(), tk_transport.output()])
|
||||||
|
|
||||||
task = PipelineTask(pipeline)
|
task = PipelineTask(pipeline)
|
||||||
|
|
||||||
|
|||||||
@@ -12,9 +12,9 @@ import wave
|
|||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
Frame,
|
Frame,
|
||||||
AudioRawFrame,
|
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
|
OutputAudioRawFrame,
|
||||||
)
|
)
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
@@ -25,7 +25,7 @@ from pipecat.processors.aggregators.llm_response import (
|
|||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.logger import FrameLogger
|
from pipecat.processors.logger import FrameLogger
|
||||||
from pipecat.services.elevenlabs import ElevenLabsTTSService
|
from pipecat.services.cartesia import CartesiaHttpTTSService
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
from pipecat.vad.silero import SileroVADAnalyzer
|
from pipecat.vad.silero import SileroVADAnalyzer
|
||||||
@@ -53,8 +53,8 @@ for file in sound_files:
|
|||||||
filename = os.path.splitext(os.path.basename(full_path))[0]
|
filename = os.path.splitext(os.path.basename(full_path))[0]
|
||||||
# Open the image and convert it to bytes
|
# Open the image and convert it to bytes
|
||||||
with wave.open(full_path) as audio_file:
|
with wave.open(full_path) as audio_file:
|
||||||
sounds[file] = AudioRawFrame(audio_file.readframes(-1),
|
sounds[file] = OutputAudioRawFrame(audio_file.readframes(-1),
|
||||||
audio_file.getframerate(), audio_file.getnchannels())
|
audio_file.getframerate(), audio_file.getnchannels())
|
||||||
|
|
||||||
|
|
||||||
class OutboundSoundEffectWrapper(FrameProcessor):
|
class OutboundSoundEffectWrapper(FrameProcessor):
|
||||||
@@ -103,9 +103,9 @@ async def main():
|
|||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
model="gpt-4o")
|
model="gpt-4o")
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = CartesiaHttpTTSService(
|
||||||
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id="ErXwobaYiN019PkySvjV",
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
|
|||||||
@@ -70,7 +70,7 @@ async def main():
|
|||||||
async def user_idle_callback(user_idle: UserIdleProcessor):
|
async def user_idle_callback(user_idle: UserIdleProcessor):
|
||||||
messages.append(
|
messages.append(
|
||||||
{"role": "system", "content": "Ask the user if they are still there and try to prompt for some input, but be short."})
|
{"role": "system", "content": "Ask the user if they are still there and try to prompt for some input, but be short."})
|
||||||
await user_idle.queue_frame(LLMMessagesFrame(messages))
|
await user_idle.push_frame(LLMMessagesFrame(messages))
|
||||||
|
|
||||||
user_idle = UserIdleProcessor(callback=user_idle_callback, timeout=5.0)
|
user_idle = UserIdleProcessor(callback=user_idle_callback, timeout=5.0)
|
||||||
|
|
||||||
|
|||||||
@@ -13,10 +13,11 @@ from PIL import Image
|
|||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
ImageRawFrame,
|
ImageRawFrame,
|
||||||
|
OutputImageRawFrame,
|
||||||
SpriteFrame,
|
SpriteFrame,
|
||||||
Frame,
|
Frame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
AudioRawFrame,
|
TTSAudioRawFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TextFrame,
|
TextFrame,
|
||||||
UserImageRawFrame,
|
UserImageRawFrame,
|
||||||
@@ -59,7 +60,11 @@ for i in range(1, 26):
|
|||||||
# Get the filename without the extension to use as the dictionary key
|
# Get the filename without the extension to use as the dictionary key
|
||||||
# Open the image and convert it to bytes
|
# Open the image and convert it to bytes
|
||||||
with Image.open(full_path) as img:
|
with Image.open(full_path) as img:
|
||||||
sprites.append(ImageRawFrame(image=img.tobytes(), size=img.size, format=img.format))
|
sprites.append(OutputImageRawFrame(
|
||||||
|
image=img.tobytes(),
|
||||||
|
size=img.size,
|
||||||
|
format=img.format)
|
||||||
|
)
|
||||||
|
|
||||||
flipped = sprites[::-1]
|
flipped = sprites[::-1]
|
||||||
sprites.extend(flipped)
|
sprites.extend(flipped)
|
||||||
@@ -82,7 +87,7 @@ class TalkingAnimation(FrameProcessor):
|
|||||||
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)
|
||||||
|
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, TTSAudioRawFrame):
|
||||||
if not self._is_talking:
|
if not self._is_talking:
|
||||||
await self.push_frame(talking_frame)
|
await self.push_frame(talking_frame)
|
||||||
self._is_talking = True
|
self._is_talking = True
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
python-dotenv
|
python-dotenv
|
||||||
fastapi[all]
|
fastapi[all]
|
||||||
uvicorn
|
uvicorn
|
||||||
pipecat-ai[daily,moondream,openai,silero]
|
pipecat-ai[daily,cartesia,moondream,openai,silero]
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ import os
|
|||||||
import sys
|
import sys
|
||||||
import wave
|
import wave
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame
|
from pipecat.frames.frames import OutputAudioRawFrame
|
||||||
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
|
||||||
@@ -49,8 +49,9 @@ for file in sound_files:
|
|||||||
filename = os.path.splitext(os.path.basename(full_path))[0]
|
filename = os.path.splitext(os.path.basename(full_path))[0]
|
||||||
# Open the sound and convert it to bytes
|
# Open the sound and convert it to bytes
|
||||||
with wave.open(full_path) as audio_file:
|
with wave.open(full_path) as audio_file:
|
||||||
sounds[file] = AudioRawFrame(audio_file.readframes(-1),
|
sounds[file] = OutputAudioRawFrame(audio_file.readframes(-1),
|
||||||
audio_file.getframerate(), audio_file.getnchannels())
|
audio_file.getframerate(),
|
||||||
|
audio_file.getnchannels())
|
||||||
|
|
||||||
|
|
||||||
class IntakeProcessor:
|
class IntakeProcessor:
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
python-dotenv
|
python-dotenv
|
||||||
fastapi[all]
|
fastapi[all]
|
||||||
uvicorn
|
uvicorn
|
||||||
pipecat-ai[daily,openai,silero]
|
pipecat-ai[daily,cartesia,openai,silero]
|
||||||
|
|||||||
@@ -16,11 +16,11 @@ from pipecat.pipeline.runner import PipelineRunner
|
|||||||
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
from pipecat.processors.aggregators.llm_response import LLMAssistantResponseAggregator, LLMUserResponseAggregator
|
from pipecat.processors.aggregators.llm_response import LLMAssistantResponseAggregator, LLMUserResponseAggregator
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
OutputImageRawFrame,
|
||||||
ImageRawFrame,
|
|
||||||
SpriteFrame,
|
SpriteFrame,
|
||||||
Frame,
|
Frame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStoppedFrame
|
TTSStoppedFrame
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
@@ -49,7 +49,11 @@ for i in range(1, 26):
|
|||||||
# Get the filename without the extension to use as the dictionary key
|
# Get the filename without the extension to use as the dictionary key
|
||||||
# Open the image and convert it to bytes
|
# Open the image and convert it to bytes
|
||||||
with Image.open(full_path) as img:
|
with Image.open(full_path) as img:
|
||||||
sprites.append(ImageRawFrame(image=img.tobytes(), size=img.size, format=img.format))
|
sprites.append(OutputImageRawFrame(
|
||||||
|
image=img.tobytes(),
|
||||||
|
size=img.size,
|
||||||
|
format=img.format)
|
||||||
|
)
|
||||||
|
|
||||||
flipped = sprites[::-1]
|
flipped = sprites[::-1]
|
||||||
sprites.extend(flipped)
|
sprites.extend(flipped)
|
||||||
@@ -72,7 +76,7 @@ class TalkingAnimation(FrameProcessor):
|
|||||||
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)
|
||||||
|
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, TTSAudioRawFrame):
|
||||||
if not self._is_talking:
|
if not self._is_talking:
|
||||||
await self.push_frame(talking_frame)
|
await self.push_frame(talking_frame)
|
||||||
self._is_talking = True
|
self._is_talking = True
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
python-dotenv
|
python-dotenv
|
||||||
fastapi[all]
|
fastapi[all]
|
||||||
uvicorn
|
uvicorn
|
||||||
pipecat-ai[daily,openai,silero]
|
pipecat-ai[daily,elevenlabs,openai,silero]
|
||||||
|
|||||||
@@ -2,4 +2,4 @@ async_timeout
|
|||||||
fastapi
|
fastapi
|
||||||
uvicorn
|
uvicorn
|
||||||
python-dotenv
|
python-dotenv
|
||||||
pipecat-ai[daily,openai,fal]
|
pipecat-ai[daily,elevenlabs,openai,fal]
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import os
|
|||||||
import wave
|
import wave
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, ImageRawFrame
|
from pipecat.frames.frames import OutputAudioRawFrame, OutputImageRawFrame
|
||||||
|
|
||||||
script_dir = os.path.dirname(__file__)
|
script_dir = os.path.dirname(__file__)
|
||||||
|
|
||||||
@@ -16,7 +16,8 @@ def load_images(image_files):
|
|||||||
filename = os.path.splitext(os.path.basename(full_path))[0]
|
filename = os.path.splitext(os.path.basename(full_path))[0]
|
||||||
# Open the image and convert it to bytes
|
# Open the image and convert it to bytes
|
||||||
with Image.open(full_path) as img:
|
with Image.open(full_path) as img:
|
||||||
images[filename] = ImageRawFrame(image=img.tobytes(), size=img.size, format=img.format)
|
images[filename] = OutputImageRawFrame(
|
||||||
|
image=img.tobytes(), size=img.size, format=img.format)
|
||||||
return images
|
return images
|
||||||
|
|
||||||
|
|
||||||
@@ -30,8 +31,8 @@ def load_sounds(sound_files):
|
|||||||
filename = os.path.splitext(os.path.basename(full_path))[0]
|
filename = os.path.splitext(os.path.basename(full_path))[0]
|
||||||
# Open the sound and convert it to bytes
|
# Open the sound and convert it to bytes
|
||||||
with wave.open(full_path) as audio_file:
|
with wave.open(full_path) as audio_file:
|
||||||
sounds[filename] = AudioRawFrame(audio=audio_file.readframes(-1),
|
sounds[filename] = OutputAudioRawFrame(audio=audio_file.readframes(-1),
|
||||||
sample_rate=audio_file.getframerate(),
|
sample_rate=audio_file.getframerate(),
|
||||||
num_channels=audio_file.getnchannels())
|
num_channels=audio_file.getnchannels())
|
||||||
|
|
||||||
return sounds
|
return sounds
|
||||||
|
|||||||
@@ -55,7 +55,7 @@ This project is a FastAPI-based chatbot that integrates with Twilio to handle We
|
|||||||
2. **Update the Twilio Webhook**:
|
2. **Update the Twilio Webhook**:
|
||||||
Copy the ngrok URL and update your Twilio phone number webhook URL to `http://<ngrok_url>/start_call`.
|
Copy the ngrok URL and update your Twilio phone number webhook URL to `http://<ngrok_url>/start_call`.
|
||||||
|
|
||||||
3. **Update the streams.xml**:
|
3. **Update streams.xml**:
|
||||||
Copy the ngrok URL and update templates/streams.xml with `wss://<ngrok_url>/ws`.
|
Copy the ngrok URL and update templates/streams.xml with `wss://<ngrok_url>/ws`.
|
||||||
|
|
||||||
## Running the Application
|
## Running the Application
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
import aiohttp
|
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
@@ -27,63 +26,62 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
|
|
||||||
async def run_bot(websocket_client, stream_sid):
|
async def run_bot(websocket_client, stream_sid):
|
||||||
async with aiohttp.ClientSession() as session:
|
transport = FastAPIWebsocketTransport(
|
||||||
transport = FastAPIWebsocketTransport(
|
websocket=websocket_client,
|
||||||
websocket=websocket_client,
|
params=FastAPIWebsocketParams(
|
||||||
params=FastAPIWebsocketParams(
|
audio_out_enabled=True,
|
||||||
audio_out_enabled=True,
|
add_wav_header=False,
|
||||||
add_wav_header=False,
|
vad_enabled=True,
|
||||||
vad_enabled=True,
|
vad_analyzer=SileroVADAnalyzer(),
|
||||||
vad_analyzer=SileroVADAnalyzer(),
|
vad_audio_passthrough=True,
|
||||||
vad_audio_passthrough=True,
|
serializer=TwilioFrameSerializer(stream_sid)
|
||||||
serializer=TwilioFrameSerializer(stream_sid)
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
model="gpt-4o")
|
model="gpt-4o")
|
||||||
|
|
||||||
stt = DeepgramSTTService(api_key=os.getenv('DEEPGRAM_API_KEY'))
|
stt = DeepgramSTTService(api_key=os.getenv('DEEPGRAM_API_KEY'))
|
||||||
|
|
||||||
tts = CartesiaTTSService(
|
tts = CartesiaTTSService(
|
||||||
api_key=os.getenv("CARTESIA_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
"role": "system",
|
"role": "system",
|
||||||
"content": "You are a helpful LLM in an audio call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
"content": "You are a helpful LLM in an audio call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
tma_in = LLMUserResponseAggregator(messages)
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
tma_out = LLMAssistantResponseAggregator(messages)
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
|
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline([
|
||||||
transport.input(), # Websocket input from client
|
transport.input(), # Websocket input from client
|
||||||
stt, # Speech-To-Text
|
stt, # Speech-To-Text
|
||||||
tma_in, # User responses
|
tma_in, # User responses
|
||||||
llm, # LLM
|
llm, # LLM
|
||||||
tts, # Text-To-Speech
|
tts, # Text-To-Speech
|
||||||
transport.output(), # Websocket output to client
|
transport.output(), # Websocket output to client
|
||||||
tma_out # LLM responses
|
tma_out # LLM responses
|
||||||
])
|
])
|
||||||
|
|
||||||
task = PipelineTask(pipeline, params=PipelineParams(allow_interruptions=True))
|
task = PipelineTask(pipeline, params=PipelineParams(allow_interruptions=True))
|
||||||
|
|
||||||
@transport.event_handler("on_client_connected")
|
@transport.event_handler("on_client_connected")
|
||||||
async def on_client_connected(transport, client):
|
async def on_client_connected(transport, client):
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
messages.append(
|
messages.append(
|
||||||
{"role": "system", "content": "Please introduce yourself to the user."})
|
{"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
await task.queue_frames([LLMMessagesFrame(messages)])
|
await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
@transport.event_handler("on_client_disconnected")
|
@transport.event_handler("on_client_disconnected")
|
||||||
async def on_client_disconnected(transport, client):
|
async def on_client_disconnected(transport, client):
|
||||||
await task.queue_frames([EndFrame()])
|
await task.queue_frames([EndFrame()])
|
||||||
|
|
||||||
runner = PipelineRunner(handle_sigint=False)
|
runner = PipelineRunner(handle_sigint=False)
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
pipecat-ai[daily,openai,silero,deepgram]
|
pipecat-ai[daily,cartesia,openai,silero,deepgram]
|
||||||
fastapi
|
fastapi
|
||||||
uvicorn
|
uvicorn
|
||||||
python-dotenv
|
python-dotenv
|
||||||
|
|||||||
@@ -4,7 +4,6 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import aiohttp
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
@@ -33,60 +32,59 @@ logger.add(sys.stderr, level="DEBUG")
|
|||||||
|
|
||||||
|
|
||||||
async def main():
|
async def main():
|
||||||
async with aiohttp.ClientSession() as session:
|
transport = WebsocketServerTransport(
|
||||||
transport = WebsocketServerTransport(
|
params=WebsocketServerParams(
|
||||||
params=WebsocketServerParams(
|
audio_out_enabled=True,
|
||||||
audio_out_enabled=True,
|
add_wav_header=True,
|
||||||
add_wav_header=True,
|
vad_enabled=True,
|
||||||
vad_enabled=True,
|
vad_analyzer=SileroVADAnalyzer(),
|
||||||
vad_analyzer=SileroVADAnalyzer(),
|
vad_audio_passthrough=True
|
||||||
vad_audio_passthrough=True
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
model="gpt-4o")
|
model="gpt-4o")
|
||||||
|
|
||||||
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
||||||
|
|
||||||
tts = CartesiaTTSService(
|
tts = CartesiaTTSService(
|
||||||
api_key=os.getenv("CARTESIA_API_KEY"),
|
api_key=os.getenv("CARTESIA_API_KEY"),
|
||||||
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
voice_id="79a125e8-cd45-4c13-8a67-188112f4dd22", # British Lady
|
||||||
)
|
)
|
||||||
|
|
||||||
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. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
"content": "You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be converted to audio so don't include special characters in your answers. Respond to what the user said in a creative and helpful way.",
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
|
|
||||||
tma_in = LLMUserResponseAggregator(messages)
|
tma_in = LLMUserResponseAggregator(messages)
|
||||||
tma_out = LLMAssistantResponseAggregator(messages)
|
tma_out = LLMAssistantResponseAggregator(messages)
|
||||||
|
|
||||||
pipeline = Pipeline([
|
pipeline = Pipeline([
|
||||||
transport.input(), # Websocket input from client
|
transport.input(), # Websocket input from client
|
||||||
stt, # Speech-To-Text
|
stt, # Speech-To-Text
|
||||||
tma_in, # User responses
|
tma_in, # User responses
|
||||||
llm, # LLM
|
llm, # LLM
|
||||||
tts, # Text-To-Speech
|
tts, # Text-To-Speech
|
||||||
transport.output(), # Websocket output to client
|
transport.output(), # Websocket output to client
|
||||||
tma_out # LLM responses
|
tma_out # LLM responses
|
||||||
])
|
])
|
||||||
|
|
||||||
task = PipelineTask(pipeline)
|
task = PipelineTask(pipeline)
|
||||||
|
|
||||||
@transport.event_handler("on_client_connected")
|
@transport.event_handler("on_client_connected")
|
||||||
async def on_client_connected(transport, client):
|
async def on_client_connected(transport, client):
|
||||||
# Kick off the conversation.
|
# Kick off the conversation.
|
||||||
messages.append(
|
messages.append(
|
||||||
{"role": "system", "content": "Please introduce yourself to the user."})
|
{"role": "system", "content": "Please introduce yourself to the user."})
|
||||||
await task.queue_frames([LLMMessagesFrame(messages)])
|
await task.queue_frames([LLMMessagesFrame(messages)])
|
||||||
|
|
||||||
runner = PipelineRunner()
|
runner = PipelineRunner()
|
||||||
|
|
||||||
await runner.run(task)
|
await runner.run(task)
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
asyncio.run(main())
|
asyncio.run(main())
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ message AudioRawFrame {
|
|||||||
bytes audio = 3;
|
bytes audio = 3;
|
||||||
uint32 sample_rate = 4;
|
uint32 sample_rate = 4;
|
||||||
uint32 num_channels = 5;
|
uint32 num_channels = 5;
|
||||||
|
optional uint64 pts = 6;
|
||||||
}
|
}
|
||||||
|
|
||||||
message TranscriptionFrame {
|
message TranscriptionFrame {
|
||||||
|
|||||||
@@ -1,2 +1,2 @@
|
|||||||
python-dotenv
|
python-dotenv
|
||||||
pipecat-ai[openai,silero,websocket,whisper]
|
pipecat-ai[cartesia,openai,silero,websocket,whisper]
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ Website = "https://pipecat.ai"
|
|||||||
[project.optional-dependencies]
|
[project.optional-dependencies]
|
||||||
anthropic = [ "anthropic~=0.34.0" ]
|
anthropic = [ "anthropic~=0.34.0" ]
|
||||||
azure = [ "azure-cognitiveservices-speech~=1.40.0" ]
|
azure = [ "azure-cognitiveservices-speech~=1.40.0" ]
|
||||||
cartesia = [ "websockets~=12.0" ]
|
cartesia = [ "cartesia~=1.0.13", "websockets~=12.0" ]
|
||||||
daily = [ "daily-python~=0.10.1" ]
|
daily = [ "daily-python~=0.10.1" ]
|
||||||
deepgram = [ "deepgram-sdk~=3.5.0" ]
|
deepgram = [ "deepgram-sdk~=3.5.0" ]
|
||||||
elevenlabs = [ "websockets~=12.0" ]
|
elevenlabs = [ "websockets~=12.0" ]
|
||||||
|
|||||||
@@ -24,6 +24,7 @@ message AudioRawFrame {
|
|||||||
bytes audio = 3;
|
bytes audio = 3;
|
||||||
uint32 sample_rate = 4;
|
uint32 sample_rate = 4;
|
||||||
uint32 num_channels = 5;
|
uint32 num_channels = 5;
|
||||||
|
optional uint64 pts = 6;
|
||||||
}
|
}
|
||||||
|
|
||||||
message TranscriptionFrame {
|
message TranscriptionFrame {
|
||||||
|
|||||||
@@ -4,11 +4,12 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
from typing import Any, List, Mapping, Optional, Tuple
|
from typing import Any, List, Optional, Tuple
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
|
|
||||||
from pipecat.clocks.base_clock import BaseClock
|
from pipecat.clocks.base_clock import BaseClock
|
||||||
|
from pipecat.metrics.metrics import MetricsData
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import nanoseconds_to_str
|
from pipecat.utils.time import nanoseconds_to_str
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
@@ -41,10 +42,7 @@ class DataFrame(Frame):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AudioRawFrame(DataFrame):
|
class AudioRawFrame(DataFrame):
|
||||||
"""A chunk of audio. Will be played by the transport if the transport's
|
"""A chunk of audio."""
|
||||||
microphone has been enabled.
|
|
||||||
|
|
||||||
"""
|
|
||||||
audio: bytes
|
audio: bytes
|
||||||
sample_rate: int
|
sample_rate: int
|
||||||
num_channels: int
|
num_channels: int
|
||||||
@@ -58,6 +56,31 @@ class AudioRawFrame(DataFrame):
|
|||||||
return f"{self.name}(pts: {pts}, size: {len(self.audio)}, frames: {self.num_frames}, sample_rate: {self.sample_rate}, channels: {self.num_channels})"
|
return f"{self.name}(pts: {pts}, size: {len(self.audio)}, frames: {self.num_frames}, sample_rate: {self.sample_rate}, channels: {self.num_channels})"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class InputAudioRawFrame(AudioRawFrame):
|
||||||
|
"""A chunk of audio usually coming from an input transport.
|
||||||
|
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OutputAudioRawFrame(AudioRawFrame):
|
||||||
|
"""A chunk of audio. Will be played by the output transport if the
|
||||||
|
transport's microphone has been enabled.
|
||||||
|
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TTSAudioRawFrame(OutputAudioRawFrame):
|
||||||
|
"""A chunk of output audio generated by a TTS service.
|
||||||
|
|
||||||
|
"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ImageRawFrame(DataFrame):
|
class ImageRawFrame(DataFrame):
|
||||||
"""An image. Will be shown by the transport if the transport's camera is
|
"""An image. Will be shown by the transport if the transport's camera is
|
||||||
@@ -74,20 +97,30 @@ class ImageRawFrame(DataFrame):
|
|||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class URLImageRawFrame(ImageRawFrame):
|
class InputImageRawFrame(ImageRawFrame):
|
||||||
"""An image with an associated URL. Will be shown by the transport if the
|
pass
|
||||||
transport's camera is enabled.
|
|
||||||
|
|
||||||
"""
|
|
||||||
url: str | None
|
|
||||||
|
|
||||||
def __str__(self):
|
|
||||||
pts = format_pts(self.pts)
|
|
||||||
return f"{self.name}(pts: {pts}, url: {self.url}, size: {self.size}, format: {self.format})"
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class VisionImageRawFrame(ImageRawFrame):
|
class OutputImageRawFrame(ImageRawFrame):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class UserImageRawFrame(InputImageRawFrame):
|
||||||
|
"""An image associated to a user. Will be shown by the transport if the
|
||||||
|
transport's camera is enabled.
|
||||||
|
|
||||||
|
"""
|
||||||
|
user_id: str
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
pts = format_pts(self.pts)
|
||||||
|
return f"{self.name}(pts: {pts}, user: {self.user_id}, size: {self.size}, format: {self.format})"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class VisionImageRawFrame(InputImageRawFrame):
|
||||||
"""An image with an associated text to ask for a description of it. Will be
|
"""An image with an associated text to ask for a description of it. Will be
|
||||||
shown by the transport if the transport's camera is enabled.
|
shown by the transport if the transport's camera is enabled.
|
||||||
|
|
||||||
@@ -100,16 +133,16 @@ class VisionImageRawFrame(ImageRawFrame):
|
|||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class UserImageRawFrame(ImageRawFrame):
|
class URLImageRawFrame(OutputImageRawFrame):
|
||||||
"""An image associated to a user. Will be shown by the transport if the
|
"""An image with an associated URL. Will be shown by the transport if the
|
||||||
transport's camera is enabled.
|
transport's camera is enabled.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
user_id: str
|
url: str | None
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
pts = format_pts(self.pts)
|
pts = format_pts(self.pts)
|
||||||
return f"{self.name}(pts: {pts}, user: {self.user_id}, size: {self.size}, format: {self.format})"
|
return f"{self.name}(pts: {pts}, url: {self.url}, size: {self.size}, format: {self.format})"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -248,6 +281,16 @@ class SystemFrame(Frame):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class StartFrame(SystemFrame):
|
||||||
|
"""This is the first frame that should be pushed down a pipeline."""
|
||||||
|
clock: BaseClock
|
||||||
|
allow_interruptions: bool = False
|
||||||
|
enable_metrics: bool = False
|
||||||
|
enable_usage_metrics: bool = False
|
||||||
|
report_only_initial_ttfb: bool = False
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class CancelFrame(SystemFrame):
|
class CancelFrame(SystemFrame):
|
||||||
"""Indicates that a pipeline needs to stop right away."""
|
"""Indicates that a pipeline needs to stop right away."""
|
||||||
@@ -323,10 +366,8 @@ class BotInterruptionFrame(SystemFrame):
|
|||||||
class MetricsFrame(SystemFrame):
|
class MetricsFrame(SystemFrame):
|
||||||
"""Emitted by processor that can compute metrics like latencies.
|
"""Emitted by processor that can compute metrics like latencies.
|
||||||
"""
|
"""
|
||||||
ttfb: List[Mapping[str, Any]] | None = None
|
data: List[MetricsData]
|
||||||
processing: List[Mapping[str, Any]] | None = None
|
|
||||||
tokens: List[Mapping[str, Any]] | None = None
|
|
||||||
characters: List[Mapping[str, Any]] | None = None
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Control frames
|
# Control frames
|
||||||
@@ -338,16 +379,6 @@ class ControlFrame(Frame):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class StartFrame(ControlFrame):
|
|
||||||
"""This is the first frame that should be pushed down a pipeline."""
|
|
||||||
clock: BaseClock
|
|
||||||
allow_interruptions: bool = False
|
|
||||||
enable_metrics: bool = False
|
|
||||||
enable_usage_metrics: bool = False
|
|
||||||
report_only_initial_ttfb: bool = False
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EndFrame(ControlFrame):
|
class EndFrame(ControlFrame):
|
||||||
"""Indicates that a pipeline has ended and frame processors and pipelines
|
"""Indicates that a pipeline has ended and frame processors and pipelines
|
||||||
@@ -420,10 +451,10 @@ class BotSpeakingFrame(ControlFrame):
|
|||||||
@dataclass
|
@dataclass
|
||||||
class TTSStartedFrame(ControlFrame):
|
class TTSStartedFrame(ControlFrame):
|
||||||
"""Used to indicate the beginning of a TTS response. Following
|
"""Used to indicate the beginning of a TTS response. Following
|
||||||
AudioRawFrames are part of the TTS response until an TTSEndFrame. These
|
TTSAudioRawFrames are part of the TTS response until an
|
||||||
frames can be used for aggregating audio frames in a transport to optimize
|
TTSStoppedFrame. These frames can be used for aggregating audio frames in a
|
||||||
the size of frames sent to the session, without needing to control this in
|
transport to optimize the size of frames sent to the session, without
|
||||||
the TTS service.
|
needing to control this in the TTS service.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
@@ -452,6 +483,66 @@ class LLMModelUpdateFrame(ControlFrame):
|
|||||||
model: str
|
model: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMTemperatureUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM temperature.
|
||||||
|
"""
|
||||||
|
temperature: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMTopKUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM top_k.
|
||||||
|
"""
|
||||||
|
top_k: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMTopPUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM top_p.
|
||||||
|
"""
|
||||||
|
top_p: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMFrequencyPenaltyUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM frequency
|
||||||
|
penalty.
|
||||||
|
|
||||||
|
"""
|
||||||
|
frequency_penalty: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMPresencePenaltyUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM presence
|
||||||
|
penalty.
|
||||||
|
|
||||||
|
"""
|
||||||
|
presence_penalty: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMMaxTokensUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM max tokens.
|
||||||
|
"""
|
||||||
|
max_tokens: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMSeedUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM seed.
|
||||||
|
"""
|
||||||
|
seed: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class LLMExtraUpdateFrame(ControlFrame):
|
||||||
|
"""A control frame containing a request to update to a new LLM extra params.
|
||||||
|
"""
|
||||||
|
extra: dict
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TTSModelUpdateFrame(ControlFrame):
|
class TTSModelUpdateFrame(ControlFrame):
|
||||||
"""A control frame containing a request to update the TTS model.
|
"""A control frame containing a request to update the TTS model.
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ _sym_db = _symbol_database.Default()
|
|||||||
|
|
||||||
|
|
||||||
|
|
||||||
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0c\x66rames.proto\x12\x07pipecat\"3\n\tTextFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04text\x18\x03 \x01(\t\"c\n\rAudioRawFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\r\n\x05\x61udio\x18\x03 \x01(\x0c\x12\x13\n\x0bsample_rate\x18\x04 \x01(\r\x12\x14\n\x0cnum_channels\x18\x05 \x01(\r\"`\n\x12TranscriptionFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04text\x18\x03 \x01(\t\x12\x0f\n\x07user_id\x18\x04 \x01(\t\x12\x11\n\ttimestamp\x18\x05 \x01(\t\"\x93\x01\n\x05\x46rame\x12\"\n\x04text\x18\x01 \x01(\x0b\x32\x12.pipecat.TextFrameH\x00\x12\'\n\x05\x61udio\x18\x02 \x01(\x0b\x32\x16.pipecat.AudioRawFrameH\x00\x12\x34\n\rtranscription\x18\x03 \x01(\x0b\x32\x1b.pipecat.TranscriptionFrameH\x00\x42\x07\n\x05\x66rameb\x06proto3')
|
DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0c\x66rames.proto\x12\x07pipecat\"3\n\tTextFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04text\x18\x03 \x01(\t\"}\n\rAudioRawFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\r\n\x05\x61udio\x18\x03 \x01(\x0c\x12\x13\n\x0bsample_rate\x18\x04 \x01(\r\x12\x14\n\x0cnum_channels\x18\x05 \x01(\r\x12\x10\n\x03pts\x18\x06 \x01(\x04H\x00\x88\x01\x01\x42\x06\n\x04_pts\"`\n\x12TranscriptionFrame\x12\n\n\x02id\x18\x01 \x01(\x04\x12\x0c\n\x04name\x18\x02 \x01(\t\x12\x0c\n\x04text\x18\x03 \x01(\t\x12\x0f\n\x07user_id\x18\x04 \x01(\t\x12\x11\n\ttimestamp\x18\x05 \x01(\t\"\x93\x01\n\x05\x46rame\x12\"\n\x04text\x18\x01 \x01(\x0b\x32\x12.pipecat.TextFrameH\x00\x12\'\n\x05\x61udio\x18\x02 \x01(\x0b\x32\x16.pipecat.AudioRawFrameH\x00\x12\x34\n\rtranscription\x18\x03 \x01(\x0b\x32\x1b.pipecat.TranscriptionFrameH\x00\x42\x07\n\x05\x66rameb\x06proto3')
|
||||||
|
|
||||||
_globals = globals()
|
_globals = globals()
|
||||||
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
|
_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals)
|
||||||
@@ -24,9 +24,9 @@ if _descriptor._USE_C_DESCRIPTORS == False:
|
|||||||
_globals['_TEXTFRAME']._serialized_start=25
|
_globals['_TEXTFRAME']._serialized_start=25
|
||||||
_globals['_TEXTFRAME']._serialized_end=76
|
_globals['_TEXTFRAME']._serialized_end=76
|
||||||
_globals['_AUDIORAWFRAME']._serialized_start=78
|
_globals['_AUDIORAWFRAME']._serialized_start=78
|
||||||
_globals['_AUDIORAWFRAME']._serialized_end=177
|
_globals['_AUDIORAWFRAME']._serialized_end=203
|
||||||
_globals['_TRANSCRIPTIONFRAME']._serialized_start=179
|
_globals['_TRANSCRIPTIONFRAME']._serialized_start=205
|
||||||
_globals['_TRANSCRIPTIONFRAME']._serialized_end=275
|
_globals['_TRANSCRIPTIONFRAME']._serialized_end=301
|
||||||
_globals['_FRAME']._serialized_start=278
|
_globals['_FRAME']._serialized_start=304
|
||||||
_globals['_FRAME']._serialized_end=425
|
_globals['_FRAME']._serialized_end=451
|
||||||
# @@protoc_insertion_point(module_scope)
|
# @@protoc_insertion_point(module_scope)
|
||||||
|
|||||||
0
src/pipecat/metrics/__init__.py
Normal file
0
src/pipecat/metrics/__init__.py
Normal file
31
src/pipecat/metrics/metrics.py
Normal file
31
src/pipecat/metrics/metrics.py
Normal file
@@ -0,0 +1,31 @@
|
|||||||
|
from typing import Optional
|
||||||
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
||||||
|
class MetricsData(BaseModel):
|
||||||
|
processor: str
|
||||||
|
model: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class TTFBMetricsData(MetricsData):
|
||||||
|
value: float
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessingMetricsData(MetricsData):
|
||||||
|
value: float
|
||||||
|
|
||||||
|
|
||||||
|
class LLMTokenUsage(BaseModel):
|
||||||
|
prompt_tokens: int
|
||||||
|
completion_tokens: int
|
||||||
|
total_tokens: int
|
||||||
|
cache_read_input_tokens: Optional[int] = None
|
||||||
|
cache_creation_input_tokens: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
|
class LLMUsageMetricsData(MetricsData):
|
||||||
|
value: LLMTokenUsage
|
||||||
|
|
||||||
|
|
||||||
|
class TTSUsageMetricsData(MetricsData):
|
||||||
|
value: int
|
||||||
@@ -49,12 +49,12 @@ class Sink(FrameProcessor):
|
|||||||
await self._down_queue.put(frame)
|
await self._down_queue.put(frame)
|
||||||
|
|
||||||
|
|
||||||
class ParallelTask(BasePipeline):
|
class SyncParallelPipeline(BasePipeline):
|
||||||
def __init__(self, *args):
|
def __init__(self, *args):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
if len(args) == 0:
|
if len(args) == 0:
|
||||||
raise Exception(f"ParallelTask needs at least one argument")
|
raise Exception(f"SyncParallelPipeline needs at least one argument")
|
||||||
|
|
||||||
self._sinks = []
|
self._sinks = []
|
||||||
self._sources = []
|
self._sources = []
|
||||||
@@ -66,7 +66,7 @@ class ParallelTask(BasePipeline):
|
|||||||
logger.debug(f"Creating {self} pipelines")
|
logger.debug(f"Creating {self} pipelines")
|
||||||
for processors in args:
|
for processors in args:
|
||||||
if not isinstance(processors, list):
|
if not isinstance(processors, list):
|
||||||
raise TypeError(f"ParallelTask argument {processors} is not a list")
|
raise TypeError(f"SyncParallelPipeline argument {processors} is not a list")
|
||||||
|
|
||||||
# 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.
|
||||||
source = Source(self._up_queue)
|
source = Source(self._up_queue)
|
||||||
@@ -20,6 +20,7 @@ from pipecat.frames.frames import (
|
|||||||
MetricsFrame,
|
MetricsFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
StopTaskFrame)
|
StopTaskFrame)
|
||||||
|
from pipecat.metrics.metrics import TTFBMetricsData, ProcessingMetricsData
|
||||||
from pipecat.pipeline.base_pipeline import BasePipeline
|
from pipecat.pipeline.base_pipeline import BasePipeline
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
@@ -118,9 +119,11 @@ class PipelineTask:
|
|||||||
|
|
||||||
def _initial_metrics_frame(self) -> MetricsFrame:
|
def _initial_metrics_frame(self) -> MetricsFrame:
|
||||||
processors = self._pipeline.processors_with_metrics()
|
processors = self._pipeline.processors_with_metrics()
|
||||||
ttfb = [{"processor": p.name, "value": 0.0} for p in processors]
|
data = []
|
||||||
processing = [{"processor": p.name, "value": 0.0} for p in processors]
|
for p in processors:
|
||||||
return MetricsFrame(ttfb=ttfb, processing=processing)
|
data.append(TTFBMetricsData(processor=p.name, value=0.0))
|
||||||
|
data.append(ProcessingMetricsData(processor=p.name, value=0.0))
|
||||||
|
return MetricsFrame(data=data)
|
||||||
|
|
||||||
async def _process_down_queue(self):
|
async def _process_down_queue(self):
|
||||||
self._clock.start()
|
self._clock.start()
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
from typing import List
|
from typing import List
|
||||||
from pipecat.pipeline.frames import EndFrame, EndPipeFrame
|
from pipecat.frames.frames import EndFrame, EndPipeFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
|
||||||
|
|
||||||
@@ -17,7 +17,8 @@ class GatedAggregator(FrameProcessor):
|
|||||||
Yields gate-opening frame before any accumulated frames, then ensuing frames
|
Yields gate-opening frame before any accumulated frames, then ensuing frames
|
||||||
until and not including the gate-closed frame.
|
until and not including the gate-closed frame.
|
||||||
|
|
||||||
>>> from pipecat.pipeline.frames import ImageFrame
|
Doctest: FIXME to work with asyncio
|
||||||
|
>>> from pipecat.frames.frames import ImageRawFrame
|
||||||
|
|
||||||
>>> async def print_frames(aggregator, frame):
|
>>> async def print_frames(aggregator, frame):
|
||||||
... async for frame in aggregator.process_frame(frame):
|
... async for frame in aggregator.process_frame(frame):
|
||||||
@@ -28,12 +29,12 @@ class GatedAggregator(FrameProcessor):
|
|||||||
|
|
||||||
>>> aggregator = GatedAggregator(
|
>>> aggregator = GatedAggregator(
|
||||||
... gate_close_fn=lambda x: isinstance(x, LLMResponseStartFrame),
|
... gate_close_fn=lambda x: isinstance(x, LLMResponseStartFrame),
|
||||||
... gate_open_fn=lambda x: isinstance(x, ImageFrame),
|
... gate_open_fn=lambda x: isinstance(x, ImageRawFrame),
|
||||||
... start_open=False)
|
... start_open=False)
|
||||||
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello")))
|
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello")))
|
||||||
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello again.")))
|
>>> asyncio.run(print_frames(aggregator, TextFrame("Hello again.")))
|
||||||
>>> asyncio.run(print_frames(aggregator, ImageFrame(image=bytes([]), size=(0, 0))))
|
>>> asyncio.run(print_frames(aggregator, ImageRawFrame(image=bytes([]), size=(0, 0))))
|
||||||
ImageFrame
|
ImageRawFrame
|
||||||
Hello
|
Hello
|
||||||
Hello again.
|
Hello again.
|
||||||
>>> asyncio.run(print_frames(aggregator, TextFrame("Goodbye.")))
|
>>> asyncio.run(print_frames(aggregator, TextFrame("Goodbye.")))
|
||||||
|
|||||||
@@ -4,8 +4,7 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import sys
|
from typing import List, Type
|
||||||
from typing import List
|
|
||||||
|
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame, OpenAILLMContext
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContextFrame, OpenAILLMContext
|
||||||
|
|
||||||
@@ -35,8 +34,8 @@ class LLMResponseAggregator(FrameProcessor):
|
|||||||
role: str,
|
role: str,
|
||||||
start_frame,
|
start_frame,
|
||||||
end_frame,
|
end_frame,
|
||||||
accumulator_frame: TextFrame,
|
accumulator_frame: Type[TextFrame],
|
||||||
interim_accumulator_frame: TextFrame | None = None,
|
interim_accumulator_frame: Type[TextFrame] | None = None,
|
||||||
handle_interruptions: bool = False
|
handle_interruptions: bool = False
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|||||||
@@ -13,7 +13,11 @@ from typing import Any, Awaitable, Callable, List
|
|||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from pipecat.frames.frames import Frame, VisionImageRawFrame, FunctionCallInProgressFrame, FunctionCallResultFrame
|
from pipecat.frames.frames import (
|
||||||
|
Frame,
|
||||||
|
VisionImageRawFrame,
|
||||||
|
FunctionCallInProgressFrame,
|
||||||
|
FunctionCallResultFrame)
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|||||||
@@ -16,7 +16,8 @@ class SentenceAggregator(FrameProcessor):
|
|||||||
TextFrame("Hello,") -> None
|
TextFrame("Hello,") -> None
|
||||||
TextFrame(" world.") -> TextFrame("Hello world.")
|
TextFrame(" world.") -> TextFrame("Hello world.")
|
||||||
|
|
||||||
Doctest:
|
Doctest: FIXME to work with asyncio
|
||||||
|
>>> import asyncio
|
||||||
>>> async def print_frames(aggregator, frame):
|
>>> async def print_frames(aggregator, frame):
|
||||||
... async for frame in aggregator.process_frame(frame):
|
... async for frame in aggregator.process_frame(frame):
|
||||||
... print(frame.text)
|
... print(frame.text)
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ class ResponseAggregator(FrameProcessor):
|
|||||||
TranscriptionFrame(" world.") -> None
|
TranscriptionFrame(" world.") -> None
|
||||||
UserStoppedSpeakingFrame() -> TextFrame("Hello world.")
|
UserStoppedSpeakingFrame() -> TextFrame("Hello world.")
|
||||||
|
|
||||||
Doctest:
|
Doctest: FIXME to work with asyncio
|
||||||
>>> async def print_frames(aggregator, frame):
|
>>> async def print_frames(aggregator, frame):
|
||||||
... async for frame in aggregator.process_frame(frame):
|
... async for frame in aggregator.process_frame(frame):
|
||||||
... if isinstance(frame, TextFrame):
|
... if isinstance(frame, TextFrame):
|
||||||
|
|||||||
@@ -4,15 +4,21 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
from pipecat.frames.frames import Frame, ImageRawFrame, TextFrame, VisionImageRawFrame
|
from pipecat.frames.frames import (
|
||||||
|
Frame,
|
||||||
|
InputImageRawFrame,
|
||||||
|
TextFrame,
|
||||||
|
VisionImageRawFrame
|
||||||
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
|
|
||||||
|
|
||||||
class VisionImageFrameAggregator(FrameProcessor):
|
class VisionImageFrameAggregator(FrameProcessor):
|
||||||
"""This aggregator waits for a consecutive TextFrame and an
|
"""This aggregator waits for a consecutive TextFrame and an
|
||||||
ImageFrame. After the ImageFrame arrives it will output a VisionImageFrame.
|
InputImageRawFrame. After the InputImageRawFrame arrives it will output a
|
||||||
|
VisionImageRawFrame.
|
||||||
|
|
||||||
>>> from pipecat.pipeline.frames import ImageFrame
|
>>> from pipecat.frames.frames import ImageFrame
|
||||||
|
|
||||||
>>> async def print_frames(aggregator, frame):
|
>>> async def print_frames(aggregator, frame):
|
||||||
... async for frame in aggregator.process_frame(frame):
|
... async for frame in aggregator.process_frame(frame):
|
||||||
@@ -34,7 +40,7 @@ class VisionImageFrameAggregator(FrameProcessor):
|
|||||||
|
|
||||||
if isinstance(frame, TextFrame):
|
if isinstance(frame, TextFrame):
|
||||||
self._describe_text = frame.text
|
self._describe_text = frame.text
|
||||||
elif isinstance(frame, ImageRawFrame):
|
elif isinstance(frame, InputImageRawFrame):
|
||||||
if self._describe_text:
|
if self._describe_text:
|
||||||
frame = VisionImageRawFrame(
|
frame = VisionImageRawFrame(
|
||||||
text=self._describe_text,
|
text=self._describe_text,
|
||||||
|
|||||||
@@ -1,64 +0,0 @@
|
|||||||
#
|
|
||||||
# Copyright (c) 2024, Daily
|
|
||||||
#
|
|
||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
|
||||||
#
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
|
|
||||||
from pipecat.frames.frames import EndFrame, Frame, StartInterruptionFrame
|
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
|
||||||
|
|
||||||
|
|
||||||
class AsyncFrameProcessor(FrameProcessor):
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
*,
|
|
||||||
name: str | None = None,
|
|
||||||
loop: asyncio.AbstractEventLoop | None = None,
|
|
||||||
**kwargs):
|
|
||||||
super().__init__(name=name, loop=loop, **kwargs)
|
|
||||||
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
||||||
await super().process_frame(frame, direction)
|
|
||||||
|
|
||||||
if isinstance(frame, StartInterruptionFrame):
|
|
||||||
await self._handle_interruptions(frame)
|
|
||||||
|
|
||||||
async def queue_frame(
|
|
||||||
self,
|
|
||||||
frame: Frame,
|
|
||||||
direction: FrameDirection = FrameDirection.DOWNSTREAM):
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def cleanup(self):
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
async def _handle_interruptions(self, frame: Frame):
|
|
||||||
# Cancel the task. This will stop pushing frames downstream.
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
# Push an out-of-band frame (i.e. not using the ordered push
|
|
||||||
# frame task).
|
|
||||||
await self.push_frame(frame)
|
|
||||||
# Create a new queue and task.
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
def _create_push_task(self):
|
|
||||||
self._push_queue = asyncio.Queue()
|
|
||||||
self._push_frame_task = self.get_event_loop().create_task(self._push_frame_task_handler())
|
|
||||||
|
|
||||||
async def _push_frame_task_handler(self):
|
|
||||||
running = True
|
|
||||||
while running:
|
|
||||||
try:
|
|
||||||
(frame, direction) = await self._push_queue.get()
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
running = not isinstance(frame, EndFrame)
|
|
||||||
self._push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
@@ -11,12 +11,21 @@ from enum import Enum
|
|||||||
|
|
||||||
from pipecat.clocks.base_clock import BaseClock
|
from pipecat.clocks.base_clock import BaseClock
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
|
EndFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
MetricsFrame,
|
MetricsFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
UserStoppedSpeakingFrame)
|
StopInterruptionFrame,
|
||||||
|
SystemFrame)
|
||||||
|
from pipecat.metrics.metrics import (
|
||||||
|
LLMTokenUsage,
|
||||||
|
LLMUsageMetricsData,
|
||||||
|
MetricsData,
|
||||||
|
ProcessingMetricsData,
|
||||||
|
TTFBMetricsData,
|
||||||
|
TTSUsageMetricsData)
|
||||||
from pipecat.utils.utils import obj_count, obj_id
|
from pipecat.utils.utils import obj_count, obj_id
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -29,11 +38,20 @@ class FrameDirection(Enum):
|
|||||||
|
|
||||||
class FrameProcessorMetrics:
|
class FrameProcessorMetrics:
|
||||||
def __init__(self, name: str):
|
def __init__(self, name: str):
|
||||||
self._name = name
|
self._core_metrics_data = MetricsData(processor=name)
|
||||||
self._start_ttfb_time = 0
|
self._start_ttfb_time = 0
|
||||||
self._start_processing_time = 0
|
self._start_processing_time = 0
|
||||||
self._should_report_ttfb = True
|
self._should_report_ttfb = True
|
||||||
|
|
||||||
|
def _processor_name(self):
|
||||||
|
return self._core_metrics_data.processor
|
||||||
|
|
||||||
|
def _model_name(self):
|
||||||
|
return self._core_metrics_data.model
|
||||||
|
|
||||||
|
def set_core_metrics_data(self, data: MetricsData):
|
||||||
|
self._core_metrics_data = data
|
||||||
|
|
||||||
async def start_ttfb_metrics(self, report_only_initial_ttfb):
|
async def start_ttfb_metrics(self, report_only_initial_ttfb):
|
||||||
if self._should_report_ttfb:
|
if self._should_report_ttfb:
|
||||||
self._start_ttfb_time = time.time()
|
self._start_ttfb_time = time.time()
|
||||||
@@ -44,13 +62,13 @@ class FrameProcessorMetrics:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
value = time.time() - self._start_ttfb_time
|
value = time.time() - self._start_ttfb_time
|
||||||
logger.debug(f"{self._name} TTFB: {value}")
|
logger.debug(f"{self._processor_name()} TTFB: {value}")
|
||||||
ttfb = {
|
ttfb = TTFBMetricsData(
|
||||||
"processor": self._name,
|
processor=self._processor_name(),
|
||||||
"value": value
|
value=value,
|
||||||
}
|
model=self._model_name())
|
||||||
self._start_ttfb_time = 0
|
self._start_ttfb_time = 0
|
||||||
return MetricsFrame(ttfb=[ttfb])
|
return MetricsFrame(data=[ttfb])
|
||||||
|
|
||||||
async def start_processing_metrics(self):
|
async def start_processing_metrics(self):
|
||||||
self._start_processing_time = time.time()
|
self._start_processing_time = time.time()
|
||||||
@@ -60,26 +78,28 @@ class FrameProcessorMetrics:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
value = time.time() - self._start_processing_time
|
value = time.time() - self._start_processing_time
|
||||||
logger.debug(f"{self._name} processing time: {value}")
|
logger.debug(f"{self._processor_name()} processing time: {value}")
|
||||||
processing = {
|
processing = ProcessingMetricsData(
|
||||||
"processor": self._name,
|
processor=self._processor_name(), value=value, model=self._model_name())
|
||||||
"value": value
|
|
||||||
}
|
|
||||||
self._start_processing_time = 0
|
self._start_processing_time = 0
|
||||||
return MetricsFrame(processing=[processing])
|
return MetricsFrame(data=[processing])
|
||||||
|
|
||||||
async def start_llm_usage_metrics(self, tokens: dict):
|
async def start_llm_usage_metrics(self, tokens: LLMTokenUsage):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"{self._name} prompt tokens: {tokens['prompt_tokens']}, completion tokens: {tokens['completion_tokens']}")
|
f"{self._processor_name()} prompt tokens: {tokens.prompt_tokens}, completion tokens: {tokens.completion_tokens}")
|
||||||
return MetricsFrame(tokens=[tokens])
|
value = LLMUsageMetricsData(
|
||||||
|
processor=self._processor_name(),
|
||||||
|
model=self._model_name(),
|
||||||
|
value=tokens)
|
||||||
|
return MetricsFrame(data=[value])
|
||||||
|
|
||||||
async def start_tts_usage_metrics(self, text: str):
|
async def start_tts_usage_metrics(self, text: str):
|
||||||
characters = {
|
characters = TTSUsageMetricsData(
|
||||||
"processor": self._name,
|
processor=self._processor_name(),
|
||||||
"value": len(text),
|
model=self._model_name(),
|
||||||
}
|
value=len(text))
|
||||||
logger.debug(f"{self._name} usage characters: {characters['value']}")
|
logger.debug(f"{self._processor_name()} usage characters: {characters.value}")
|
||||||
return MetricsFrame(characters=[characters])
|
return MetricsFrame(data=[characters])
|
||||||
|
|
||||||
|
|
||||||
class FrameProcessor:
|
class FrameProcessor:
|
||||||
@@ -88,6 +108,7 @@ class FrameProcessor:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
name: str | None = None,
|
name: str | None = None,
|
||||||
|
sync: bool = True,
|
||||||
loop: asyncio.AbstractEventLoop | None = None,
|
loop: asyncio.AbstractEventLoop | None = None,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
self.id: int = obj_id()
|
self.id: int = obj_id()
|
||||||
@@ -96,6 +117,7 @@ class FrameProcessor:
|
|||||||
self._prev: "FrameProcessor" | None = None
|
self._prev: "FrameProcessor" | None = None
|
||||||
self._next: "FrameProcessor" | None = None
|
self._next: "FrameProcessor" | None = None
|
||||||
self._loop: asyncio.AbstractEventLoop = loop or asyncio.get_running_loop()
|
self._loop: asyncio.AbstractEventLoop = loop or asyncio.get_running_loop()
|
||||||
|
self._sync = sync
|
||||||
|
|
||||||
# Clock
|
# Clock
|
||||||
self._clock: BaseClock | None = None
|
self._clock: BaseClock | None = None
|
||||||
@@ -109,6 +131,14 @@ class FrameProcessor:
|
|||||||
# Metrics
|
# Metrics
|
||||||
self._metrics = FrameProcessorMetrics(name=self.name)
|
self._metrics = FrameProcessorMetrics(name=self.name)
|
||||||
|
|
||||||
|
# Every processor in Pipecat should only output frames from a single
|
||||||
|
# task. This avoid problems like audio overlapping. System frames are
|
||||||
|
# the exception to this rule.
|
||||||
|
#
|
||||||
|
# This create this task.
|
||||||
|
if not self._sync:
|
||||||
|
self.__create_push_task()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def interruptions_allowed(self):
|
def interruptions_allowed(self):
|
||||||
return self._allow_interruptions
|
return self._allow_interruptions
|
||||||
@@ -128,6 +158,9 @@ class FrameProcessor:
|
|||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def set_core_metrics_data(self, data: MetricsData):
|
||||||
|
self._metrics.set_core_metrics_data(data)
|
||||||
|
|
||||||
async def start_ttfb_metrics(self):
|
async def start_ttfb_metrics(self):
|
||||||
if self.can_generate_metrics() and self.metrics_enabled:
|
if self.can_generate_metrics() and self.metrics_enabled:
|
||||||
await self._metrics.start_ttfb_metrics(self._report_only_initial_ttfb)
|
await self._metrics.start_ttfb_metrics(self._report_only_initial_ttfb)
|
||||||
@@ -148,7 +181,7 @@ class FrameProcessor:
|
|||||||
if frame:
|
if frame:
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def start_llm_usage_metrics(self, tokens: dict):
|
async def start_llm_usage_metrics(self, tokens: LLMTokenUsage):
|
||||||
if self.can_generate_metrics() and self.usage_metrics_enabled:
|
if self.can_generate_metrics() and self.usage_metrics_enabled:
|
||||||
frame = await self._metrics.start_llm_usage_metrics(tokens)
|
frame = await self._metrics.start_llm_usage_metrics(tokens)
|
||||||
if frame:
|
if frame:
|
||||||
@@ -192,14 +225,38 @@ class FrameProcessor:
|
|||||||
self._enable_usage_metrics = frame.enable_usage_metrics
|
self._enable_usage_metrics = frame.enable_usage_metrics
|
||||||
self._report_only_initial_ttfb = frame.report_only_initial_ttfb
|
self._report_only_initial_ttfb = frame.report_only_initial_ttfb
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
elif isinstance(frame, StartInterruptionFrame):
|
||||||
|
await self._start_interruption()
|
||||||
await self.stop_all_metrics()
|
await self.stop_all_metrics()
|
||||||
elif isinstance(frame, UserStoppedSpeakingFrame):
|
elif isinstance(frame, StopInterruptionFrame):
|
||||||
self._should_report_ttfb = True
|
self._should_report_ttfb = True
|
||||||
|
|
||||||
async def push_error(self, error: ErrorFrame):
|
async def push_error(self, error: ErrorFrame):
|
||||||
await self.push_frame(error, FrameDirection.UPSTREAM)
|
await self.push_frame(error, FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM):
|
async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM):
|
||||||
|
if self._sync or isinstance(frame, SystemFrame):
|
||||||
|
await self.__internal_push_frame(frame, direction)
|
||||||
|
else:
|
||||||
|
await self.__push_queue.put((frame, direction))
|
||||||
|
|
||||||
|
#
|
||||||
|
# Handle interruptions
|
||||||
|
#
|
||||||
|
|
||||||
|
async def _start_interruption(self):
|
||||||
|
if not self._sync:
|
||||||
|
# Cancel the task. This will stop pushing frames downstream.
|
||||||
|
self.__push_frame_task.cancel()
|
||||||
|
await self.__push_frame_task
|
||||||
|
|
||||||
|
# Create a new queue and task.
|
||||||
|
self.__create_push_task()
|
||||||
|
|
||||||
|
async def _stop_interruption(self):
|
||||||
|
# Nothing to do right now.
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def __internal_push_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
try:
|
try:
|
||||||
if direction == FrameDirection.DOWNSTREAM and self._next:
|
if direction == FrameDirection.DOWNSTREAM and self._next:
|
||||||
logger.trace(f"Pushing {frame} from {self} to {self._next}")
|
logger.trace(f"Pushing {frame} from {self} to {self._next}")
|
||||||
@@ -210,5 +267,20 @@ class FrameProcessor:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(f"Uncaught exception in {self}: {e}")
|
logger.exception(f"Uncaught exception in {self}: {e}")
|
||||||
|
|
||||||
|
def __create_push_task(self):
|
||||||
|
self.__push_queue = asyncio.Queue()
|
||||||
|
self.__push_frame_task = self.get_event_loop().create_task(self.__push_frame_task_handler())
|
||||||
|
|
||||||
|
async def __push_frame_task_handler(self):
|
||||||
|
running = True
|
||||||
|
while running:
|
||||||
|
try:
|
||||||
|
(frame, direction) = await self.__push_queue.get()
|
||||||
|
await self.__internal_push_frame(frame, direction)
|
||||||
|
running = not isinstance(frame, EndFrame)
|
||||||
|
self.__push_queue.task_done()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return self.name
|
return self.name
|
||||||
|
|||||||
@@ -272,8 +272,9 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
def __init__(self,
|
def __init__(self,
|
||||||
*,
|
*,
|
||||||
config: RTVIConfig = RTVIConfig(config=[]),
|
config: RTVIConfig = RTVIConfig(config=[]),
|
||||||
params: RTVIProcessorParams = RTVIProcessorParams()):
|
params: RTVIProcessorParams = RTVIProcessorParams(),
|
||||||
super().__init__()
|
**kwargs):
|
||||||
|
super().__init__(sync=False, **kwargs)
|
||||||
self._config = config
|
self._config = config
|
||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
@@ -286,9 +287,6 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
self._registered_actions: Dict[str, RTVIAction] = {}
|
self._registered_actions: Dict[str, RTVIAction] = {}
|
||||||
self._registered_services: Dict[str, RTVIService] = {}
|
self._registered_services: Dict[str, RTVIService] = {}
|
||||||
|
|
||||||
self._push_frame_task = self.get_event_loop().create_task(self._push_frame_task_handler())
|
|
||||||
self._push_queue = asyncio.Queue()
|
|
||||||
|
|
||||||
self._message_task = self.get_event_loop().create_task(self._message_task_handler())
|
self._message_task = self.get_event_loop().create_task(self._message_task_handler())
|
||||||
self._message_queue = asyncio.Queue()
|
self._message_queue = asyncio.Queue()
|
||||||
|
|
||||||
@@ -335,17 +333,16 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
message = RTVILLMFunctionCallStartMessage(data=fn)
|
message = RTVILLMFunctionCallStartMessage(data=fn)
|
||||||
await self._push_transport_message(message, exclude_none=False)
|
await self._push_transport_message(message, exclude_none=False)
|
||||||
|
|
||||||
async def push_frame(self, frame: Frame, direction: FrameDirection = FrameDirection.DOWNSTREAM):
|
|
||||||
if isinstance(frame, SystemFrame):
|
|
||||||
await super().push_frame(frame, direction)
|
|
||||||
else:
|
|
||||||
await self._internal_push_frame(frame, direction)
|
|
||||||
|
|
||||||
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)
|
||||||
|
|
||||||
# Specific system frames
|
# Specific system frames
|
||||||
if isinstance(frame, CancelFrame):
|
if isinstance(frame, StartFrame):
|
||||||
|
# Push StartFrame before start(), because we want StartFrame to be
|
||||||
|
# processed by every processor before any other frame is processed.
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
await self._start(frame)
|
||||||
|
elif isinstance(frame, CancelFrame):
|
||||||
await self._cancel(frame)
|
await self._cancel(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, ErrorFrame):
|
elif isinstance(frame, ErrorFrame):
|
||||||
@@ -355,11 +352,6 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
elif isinstance(frame, SystemFrame):
|
elif isinstance(frame, SystemFrame):
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
# Control frames
|
# Control frames
|
||||||
elif isinstance(frame, StartFrame):
|
|
||||||
# Push StartFrame before start(), because we want StartFrame to be
|
|
||||||
# processed by every processor before any other frame is processed.
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
await self._start(frame)
|
|
||||||
elif isinstance(frame, EndFrame):
|
elif isinstance(frame, EndFrame):
|
||||||
# Push EndFrame before stop(), because stop() waits on the task to
|
# Push EndFrame before stop(), because stop() waits on the task to
|
||||||
# finish and the task finishes when EndFrame is processed.
|
# finish and the task finishes when EndFrame is processed.
|
||||||
@@ -394,30 +386,10 @@ class RTVIProcessor(FrameProcessor):
|
|||||||
# processing EndFrames.
|
# processing EndFrames.
|
||||||
self._message_task.cancel()
|
self._message_task.cancel()
|
||||||
await self._message_task
|
await self._message_task
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
async def _cancel(self, frame: CancelFrame):
|
async def _cancel(self, frame: CancelFrame):
|
||||||
self._message_task.cancel()
|
self._message_task.cancel()
|
||||||
await self._message_task
|
await self._message_task
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
async def _internal_push_frame(
|
|
||||||
self,
|
|
||||||
frame: Frame | None,
|
|
||||||
direction: FrameDirection | None = FrameDirection.DOWNSTREAM):
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def _push_frame_task_handler(self):
|
|
||||||
running = True
|
|
||||||
while running:
|
|
||||||
try:
|
|
||||||
(frame, direction) = await self._push_queue.get()
|
|
||||||
await super().push_frame(frame, direction)
|
|
||||||
self._push_queue.task_done()
|
|
||||||
running = not isinstance(frame, EndFrame)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
|
async def _push_transport_message(self, model: BaseModel, exclude_none: bool = True):
|
||||||
frame = TransportMessageFrame(
|
frame = TransportMessageFrame(
|
||||||
|
|||||||
@@ -9,11 +9,11 @@ import asyncio
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
ImageRawFrame,
|
OutputAudioRawFrame,
|
||||||
|
OutputImageRawFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
SystemFrame)
|
SystemFrame)
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
@@ -41,7 +41,7 @@ class GStreamerPipelineSource(FrameProcessor):
|
|||||||
clock_sync: bool = True
|
clock_sync: bool = True
|
||||||
|
|
||||||
def __init__(self, *, pipeline: str, out_params: OutputParams = OutputParams(), **kwargs):
|
def __init__(self, *, pipeline: str, out_params: OutputParams = OutputParams(), **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._out_params = out_params
|
self._out_params = out_params
|
||||||
|
|
||||||
@@ -62,78 +62,42 @@ class GStreamerPipelineSource(FrameProcessor):
|
|||||||
bus.add_signal_watch()
|
bus.add_signal_watch()
|
||||||
bus.connect("message", self._on_gstreamer_message)
|
bus.connect("message", self._on_gstreamer_message)
|
||||||
|
|
||||||
# Create push frame task. This is the task that will push frames in
|
|
||||||
# order. We also guarantee that all frames are pushed in the same task.
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
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)
|
||||||
|
|
||||||
# Specific system frames
|
# Specific system frames
|
||||||
if isinstance(frame, CancelFrame):
|
if isinstance(frame, StartFrame):
|
||||||
|
# Push StartFrame before start(), because we want StartFrame to be
|
||||||
|
# processed by every processor before any other frame is processed.
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
await self._start(frame)
|
||||||
|
elif isinstance(frame, CancelFrame):
|
||||||
await self._cancel(frame)
|
await self._cancel(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
# All other system frames
|
# All other system frames
|
||||||
elif isinstance(frame, SystemFrame):
|
elif isinstance(frame, SystemFrame):
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
# Control frames
|
# Control frames
|
||||||
elif isinstance(frame, StartFrame):
|
|
||||||
# Push StartFrame before start(), because we want StartFrame to be
|
|
||||||
# processed by every processor before any other frame is processed.
|
|
||||||
await self._internal_push_frame(frame, direction)
|
|
||||||
await self._start(frame)
|
|
||||||
elif isinstance(frame, EndFrame):
|
elif isinstance(frame, EndFrame):
|
||||||
# Push EndFrame before stop(), because stop() waits on the task to
|
# Push EndFrame before stop(), because stop() waits on the task to
|
||||||
# finish and the task finishes when EndFrame is processed.
|
# finish and the task finishes when EndFrame is processed.
|
||||||
await self._internal_push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
await self._stop(frame)
|
await self._stop(frame)
|
||||||
# Other frames
|
# Other frames
|
||||||
else:
|
else:
|
||||||
await self._internal_push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
async def _start(self, frame: StartFrame):
|
async def _start(self, frame: StartFrame):
|
||||||
self._player.set_state(Gst.State.PLAYING)
|
self._player.set_state(Gst.State.PLAYING)
|
||||||
|
|
||||||
async def _stop(self, frame: EndFrame):
|
async def _stop(self, frame: EndFrame):
|
||||||
self._player.set_state(Gst.State.NULL)
|
self._player.set_state(Gst.State.NULL)
|
||||||
# Wait for the push frame task to finish. It will finish when the
|
|
||||||
# EndFrame is actually processed.
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
async def _cancel(self, frame: CancelFrame):
|
async def _cancel(self, frame: CancelFrame):
|
||||||
self._player.set_state(Gst.State.NULL)
|
self._player.set_state(Gst.State.NULL)
|
||||||
# Cancel all the tasks and wait for them to finish.
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Push frames task
|
# GStreamer
|
||||||
#
|
|
||||||
|
|
||||||
def _create_push_task(self):
|
|
||||||
loop = self.get_event_loop()
|
|
||||||
self._push_queue = asyncio.Queue()
|
|
||||||
self._push_frame_task = loop.create_task(self._push_frame_task_handler())
|
|
||||||
|
|
||||||
async def _internal_push_frame(
|
|
||||||
self,
|
|
||||||
frame: Frame | None,
|
|
||||||
direction: FrameDirection | None = FrameDirection.DOWNSTREAM):
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def _push_frame_task_handler(self):
|
|
||||||
running = True
|
|
||||||
while running:
|
|
||||||
try:
|
|
||||||
(frame, direction) = await self._push_queue.get()
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
running = not isinstance(frame, EndFrame)
|
|
||||||
self._push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
#
|
|
||||||
# GStreaner
|
|
||||||
#
|
#
|
||||||
|
|
||||||
def _on_gstreamer_message(self, bus: Gst.Bus, message: Gst.Message):
|
def _on_gstreamer_message(self, bus: Gst.Bus, message: Gst.Message):
|
||||||
@@ -218,20 +182,20 @@ class GStreamerPipelineSource(FrameProcessor):
|
|||||||
def _appsink_audio_new_sample(self, appsink: GstApp.AppSink):
|
def _appsink_audio_new_sample(self, appsink: GstApp.AppSink):
|
||||||
buffer = appsink.pull_sample().get_buffer()
|
buffer = appsink.pull_sample().get_buffer()
|
||||||
(_, info) = buffer.map(Gst.MapFlags.READ)
|
(_, info) = buffer.map(Gst.MapFlags.READ)
|
||||||
frame = AudioRawFrame(audio=info.data,
|
frame = OutputAudioRawFrame(audio=info.data,
|
||||||
sample_rate=self._out_params.audio_sample_rate,
|
sample_rate=self._out_params.audio_sample_rate,
|
||||||
num_channels=self._out_params.audio_channels)
|
num_channels=self._out_params.audio_channels)
|
||||||
asyncio.run_coroutine_threadsafe(self._internal_push_frame(frame), self.get_event_loop())
|
asyncio.run_coroutine_threadsafe(self.push_frame(frame), self.get_event_loop())
|
||||||
buffer.unmap(info)
|
buffer.unmap(info)
|
||||||
return Gst.FlowReturn.OK
|
return Gst.FlowReturn.OK
|
||||||
|
|
||||||
def _appsink_video_new_sample(self, appsink: GstApp.AppSink):
|
def _appsink_video_new_sample(self, appsink: GstApp.AppSink):
|
||||||
buffer = appsink.pull_sample().get_buffer()
|
buffer = appsink.pull_sample().get_buffer()
|
||||||
(_, info) = buffer.map(Gst.MapFlags.READ)
|
(_, info) = buffer.map(Gst.MapFlags.READ)
|
||||||
frame = ImageRawFrame(
|
frame = OutputImageRawFrame(
|
||||||
image=info.data,
|
image=info.data,
|
||||||
size=(self._out_params.video_width, self._out_params.video_height),
|
size=(self._out_params.video_width, self._out_params.video_height),
|
||||||
format="RGB")
|
format="RGB")
|
||||||
asyncio.run_coroutine_threadsafe(self._internal_push_frame(frame), self.get_event_loop())
|
asyncio.run_coroutine_threadsafe(self.push_frame(frame), self.get_event_loop())
|
||||||
buffer.unmap(info)
|
buffer.unmap(info)
|
||||||
return Gst.FlowReturn.OK
|
return Gst.FlowReturn.OK
|
||||||
|
|||||||
@@ -8,19 +8,14 @@ import asyncio
|
|||||||
|
|
||||||
from typing import Awaitable, Callable, List
|
from typing import Awaitable, Callable, List
|
||||||
|
|
||||||
from pipecat.frames.frames import Frame, SystemFrame
|
from pipecat.frames.frames import Frame
|
||||||
from pipecat.processors.async_frame_processor import AsyncFrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
|
||||||
|
|
||||||
|
|
||||||
class IdleFrameProcessor(AsyncFrameProcessor):
|
class IdleFrameProcessor(FrameProcessor):
|
||||||
"""This class waits to receive any frame or list of desired frames within a
|
"""This class waits to receive any frame or list of desired frames within a
|
||||||
given timeout. If the timeout is reached before receiving any of those
|
given timeout. If the timeout is reached before receiving any of those
|
||||||
frames the provided callback will be called.
|
frames the provided callback will be called.
|
||||||
|
|
||||||
The callback can then be used to push frames downstream by using
|
|
||||||
`queue_frame()` (or `push_frame()` for system frames).
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -30,7 +25,7 @@ class IdleFrameProcessor(AsyncFrameProcessor):
|
|||||||
timeout: float,
|
timeout: float,
|
||||||
types: List[type] = [],
|
types: List[type] = [],
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._callback = callback
|
self._callback = callback
|
||||||
self._timeout = timeout
|
self._timeout = timeout
|
||||||
@@ -41,10 +36,7 @@ class IdleFrameProcessor(AsyncFrameProcessor):
|
|||||||
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)
|
||||||
|
|
||||||
if isinstance(frame, SystemFrame):
|
await self.push_frame(frame, direction)
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
else:
|
|
||||||
await self.queue_frame(frame, direction)
|
|
||||||
|
|
||||||
# If we are not waiting for any specific frame set the event, otherwise
|
# If we are not waiting for any specific frame set the event, otherwise
|
||||||
# check if we have received one of the desired frames.
|
# check if we have received one of the desired frames.
|
||||||
@@ -55,7 +47,6 @@ class IdleFrameProcessor(AsyncFrameProcessor):
|
|||||||
if isinstance(frame, t):
|
if isinstance(frame, t):
|
||||||
self._idle_event.set()
|
self._idle_event.set()
|
||||||
|
|
||||||
# If we are not waiting for any specific frame set the event, otherwise
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
self._idle_task.cancel()
|
self._idle_task.cancel()
|
||||||
await self._idle_task
|
await self._idle_task
|
||||||
|
|||||||
@@ -11,21 +11,16 @@ from typing import Awaitable, Callable
|
|||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
BotSpeakingFrame,
|
BotSpeakingFrame,
|
||||||
Frame,
|
Frame,
|
||||||
SystemFrame,
|
|
||||||
UserStartedSpeakingFrame,
|
UserStartedSpeakingFrame,
|
||||||
UserStoppedSpeakingFrame)
|
UserStoppedSpeakingFrame)
|
||||||
from pipecat.processors.async_frame_processor import AsyncFrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
|
||||||
|
|
||||||
|
|
||||||
class UserIdleProcessor(AsyncFrameProcessor):
|
class UserIdleProcessor(FrameProcessor):
|
||||||
"""This class is useful to check if the user is interacting with the bot
|
"""This class is useful to check if the user is interacting with the bot
|
||||||
within a given timeout. If the timeout is reached before any interaction
|
within a given timeout. If the timeout is reached before any interaction
|
||||||
occurred the provided callback will be called.
|
occurred the provided callback will be called.
|
||||||
|
|
||||||
The callback can then be used to push frames downstream by using
|
|
||||||
`queue_frame()` (or `push_frame()` for system frames).
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -34,7 +29,7 @@ class UserIdleProcessor(AsyncFrameProcessor):
|
|||||||
callback: Callable[["UserIdleProcessor"], Awaitable[None]],
|
callback: Callable[["UserIdleProcessor"], Awaitable[None]],
|
||||||
timeout: float,
|
timeout: float,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._callback = callback
|
self._callback = callback
|
||||||
self._timeout = timeout
|
self._timeout = timeout
|
||||||
@@ -46,10 +41,7 @@ class UserIdleProcessor(AsyncFrameProcessor):
|
|||||||
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)
|
||||||
|
|
||||||
if isinstance(frame, SystemFrame):
|
await self.push_frame(frame, direction)
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
else:
|
|
||||||
await self.queue_frame(frame, direction)
|
|
||||||
|
|
||||||
# We shouldn't call the idle callback if the user or the bot are speaking.
|
# We shouldn't call the idle callback if the user or the bot are speaking.
|
||||||
if isinstance(frame, UserStartedSpeakingFrame):
|
if isinstance(frame, UserStartedSpeakingFrame):
|
||||||
|
|||||||
@@ -7,7 +7,10 @@
|
|||||||
import ctypes
|
import ctypes
|
||||||
import pickle
|
import pickle
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, Frame
|
from pipecat.frames.frames import (
|
||||||
|
Frame,
|
||||||
|
InputAudioRawFrame,
|
||||||
|
OutputAudioRawFrame)
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import FrameSerializer
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -22,12 +25,8 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
|
|
||||||
class LivekitFrameSerializer(FrameSerializer):
|
class LivekitFrameSerializer(FrameSerializer):
|
||||||
SERIALIZABLE_TYPES = {
|
|
||||||
AudioRawFrame: "audio",
|
|
||||||
}
|
|
||||||
|
|
||||||
def serialize(self, frame: Frame) -> str | bytes | None:
|
def serialize(self, frame: Frame) -> str | bytes | None:
|
||||||
if not isinstance(frame, AudioRawFrame):
|
if not isinstance(frame, OutputAudioRawFrame):
|
||||||
return None
|
return None
|
||||||
audio_frame = AudioFrame(
|
audio_frame = AudioFrame(
|
||||||
data=frame.audio,
|
data=frame.audio,
|
||||||
@@ -39,7 +38,7 @@ class LivekitFrameSerializer(FrameSerializer):
|
|||||||
|
|
||||||
def deserialize(self, data: str | bytes) -> Frame | None:
|
def deserialize(self, data: str | bytes) -> Frame | None:
|
||||||
audio_frame: AudioFrame = pickle.loads(data)['frame']
|
audio_frame: AudioFrame = pickle.loads(data)['frame']
|
||||||
return AudioRawFrame(
|
return InputAudioRawFrame(
|
||||||
audio=bytes(audio_frame.data),
|
audio=bytes(audio_frame.data),
|
||||||
sample_rate=audio_frame.sample_rate,
|
sample_rate=audio_frame.sample_rate,
|
||||||
num_channels=audio_frame.num_channels,
|
num_channels=audio_frame.num_channels,
|
||||||
|
|||||||
@@ -8,7 +8,11 @@ import dataclasses
|
|||||||
|
|
||||||
import pipecat.frames.protobufs.frames_pb2 as frame_protos
|
import pipecat.frames.protobufs.frames_pb2 as frame_protos
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, Frame, TextFrame, TranscriptionFrame
|
from pipecat.frames.frames import (
|
||||||
|
AudioRawFrame,
|
||||||
|
Frame,
|
||||||
|
TextFrame,
|
||||||
|
TranscriptionFrame)
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import FrameSerializer
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -29,14 +33,15 @@ class ProtobufFrameSerializer(FrameSerializer):
|
|||||||
def serialize(self, frame: Frame) -> str | bytes | None:
|
def serialize(self, frame: Frame) -> str | bytes | None:
|
||||||
proto_frame = frame_protos.Frame()
|
proto_frame = frame_protos.Frame()
|
||||||
if type(frame) not in self.SERIALIZABLE_TYPES:
|
if type(frame) not in self.SERIALIZABLE_TYPES:
|
||||||
raise ValueError(
|
logger.warning(f"Frame type {type(frame)} is not serializable")
|
||||||
f"Frame type {type(frame)} is not serializable. You may need to add it to ProtobufFrameSerializer.SERIALIZABLE_FIELDS.")
|
return None
|
||||||
|
|
||||||
# ignoring linter errors; we check that type(frame) is in this dict above
|
# ignoring linter errors; we check that type(frame) is in this dict above
|
||||||
proto_optional_name = self.SERIALIZABLE_TYPES[type(frame)] # type: ignore
|
proto_optional_name = self.SERIALIZABLE_TYPES[type(frame)] # type: ignore
|
||||||
for field in dataclasses.fields(frame): # type: ignore
|
for field in dataclasses.fields(frame): # type: ignore
|
||||||
setattr(getattr(proto_frame, proto_optional_name), field.name,
|
value = getattr(frame, field.name)
|
||||||
getattr(frame, field.name))
|
if value:
|
||||||
|
setattr(getattr(proto_frame, proto_optional_name), field.name, value)
|
||||||
|
|
||||||
result = proto_frame.SerializeToString()
|
result = proto_frame.SerializeToString()
|
||||||
return result
|
return result
|
||||||
@@ -48,8 +53,8 @@ class ProtobufFrameSerializer(FrameSerializer):
|
|||||||
|
|
||||||
>>> serializer = ProtobufFrameSerializer()
|
>>> serializer = ProtobufFrameSerializer()
|
||||||
>>> serializer.deserialize(
|
>>> serializer.deserialize(
|
||||||
... serializer.serialize(AudioFrame(data=b'1234567890')))
|
... serializer.serialize(OutputAudioFrame(data=b'1234567890')))
|
||||||
AudioFrame(data=b'1234567890')
|
InputAudioFrame(data=b'1234567890')
|
||||||
|
|
||||||
>>> serializer.deserialize(
|
>>> serializer.deserialize(
|
||||||
... serializer.serialize(TextFrame(text='hello world')))
|
... serializer.serialize(TextFrame(text='hello world')))
|
||||||
@@ -75,10 +80,13 @@ class ProtobufFrameSerializer(FrameSerializer):
|
|||||||
# Remove special fields if needed
|
# Remove special fields if needed
|
||||||
id = getattr(args, "id")
|
id = getattr(args, "id")
|
||||||
name = getattr(args, "name")
|
name = getattr(args, "name")
|
||||||
|
pts = getattr(args, "pts")
|
||||||
if not id:
|
if not id:
|
||||||
del args_dict["id"]
|
del args_dict["id"]
|
||||||
if not name:
|
if not name:
|
||||||
del args_dict["name"]
|
del args_dict["name"]
|
||||||
|
if not pts:
|
||||||
|
del args_dict["pts"]
|
||||||
|
|
||||||
# Create the instance
|
# Create the instance
|
||||||
instance = class_name(**args_dict)
|
instance = class_name(**args_dict)
|
||||||
@@ -88,5 +96,7 @@ class ProtobufFrameSerializer(FrameSerializer):
|
|||||||
setattr(instance, "id", getattr(args, "id"))
|
setattr(instance, "id", getattr(args, "id"))
|
||||||
if name:
|
if name:
|
||||||
setattr(instance, "name", getattr(args, "name"))
|
setattr(instance, "name", getattr(args, "name"))
|
||||||
|
if pts:
|
||||||
|
setattr(instance, "pts", getattr(args, "pts"))
|
||||||
|
|
||||||
return instance
|
return instance
|
||||||
|
|||||||
@@ -9,7 +9,10 @@ import json
|
|||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, Frame, StartInterruptionFrame
|
from pipecat.frames.frames import (
|
||||||
|
AudioRawFrame,
|
||||||
|
Frame,
|
||||||
|
StartInterruptionFrame)
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import FrameSerializer
|
||||||
from pipecat.utils.audio import ulaw_to_pcm, pcm_to_ulaw
|
from pipecat.utils.audio import ulaw_to_pcm, pcm_to_ulaw
|
||||||
|
|
||||||
@@ -19,10 +22,6 @@ class TwilioFrameSerializer(FrameSerializer):
|
|||||||
twilio_sample_rate: int = 8000
|
twilio_sample_rate: int = 8000
|
||||||
sample_rate: int = 16000
|
sample_rate: int = 16000
|
||||||
|
|
||||||
SERIALIZABLE_TYPES = {
|
|
||||||
AudioRawFrame: "audio",
|
|
||||||
}
|
|
||||||
|
|
||||||
def __init__(self, stream_sid: str, params: InputParams = InputParams()):
|
def __init__(self, stream_sid: str, params: InputParams = InputParams()):
|
||||||
self._stream_sid = stream_sid
|
self._stream_sid = stream_sid
|
||||||
self._params = params
|
self._params = params
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from pipecat.frames.frames import (
|
|||||||
STTModelUpdateFrame,
|
STTModelUpdateFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSLanguageUpdateFrame,
|
TTSLanguageUpdateFrame,
|
||||||
TTSModelUpdateFrame,
|
TTSModelUpdateFrame,
|
||||||
TTSSpeakFrame,
|
TTSSpeakFrame,
|
||||||
@@ -32,7 +33,7 @@ from pipecat.frames.frames import (
|
|||||||
UserImageRequestFrame,
|
UserImageRequestFrame,
|
||||||
VisionImageRawFrame
|
VisionImageRawFrame
|
||||||
)
|
)
|
||||||
from pipecat.processors.async_frame_processor import AsyncFrameProcessor
|
from pipecat.metrics.metrics import MetricsData
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.audio import calculate_audio_volume
|
from pipecat.utils.audio import calculate_audio_volume
|
||||||
@@ -47,6 +48,15 @@ from loguru import logger
|
|||||||
class AIService(FrameProcessor):
|
class AIService(FrameProcessor):
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
self._model_name: str = ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_name(self) -> str:
|
||||||
|
return self._model_name
|
||||||
|
|
||||||
|
def set_model_name(self, model: str):
|
||||||
|
self._model_name = model
|
||||||
|
self.set_core_metrics_data(MetricsData(processor=self.name, model=self._model_name))
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
pass
|
pass
|
||||||
@@ -67,7 +77,7 @@ class AIService(FrameProcessor):
|
|||||||
elif isinstance(frame, EndFrame):
|
elif isinstance(frame, EndFrame):
|
||||||
await self.stop(frame)
|
await self.stop(frame)
|
||||||
|
|
||||||
async def process_generator(self, generator: AsyncGenerator[Frame, None]):
|
async def process_generator(self, generator: AsyncGenerator[Frame | None, None]):
|
||||||
async for f in generator:
|
async for f in generator:
|
||||||
if f:
|
if f:
|
||||||
if isinstance(f, ErrorFrame):
|
if isinstance(f, ErrorFrame):
|
||||||
@@ -76,30 +86,6 @@ class AIService(FrameProcessor):
|
|||||||
await self.push_frame(f)
|
await self.push_frame(f)
|
||||||
|
|
||||||
|
|
||||||
class AsyncAIService(AsyncFrameProcessor):
|
|
||||||
def __init__(self, **kwargs):
|
|
||||||
super().__init__(**kwargs)
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
|
||||||
pass
|
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
||||||
await super().process_frame(frame, direction)
|
|
||||||
|
|
||||||
if isinstance(frame, StartFrame):
|
|
||||||
await self.start(frame)
|
|
||||||
elif isinstance(frame, CancelFrame):
|
|
||||||
await self.cancel(frame)
|
|
||||||
elif isinstance(frame, EndFrame):
|
|
||||||
await self.stop(frame)
|
|
||||||
|
|
||||||
|
|
||||||
class LLMService(AIService):
|
class LLMService(AIService):
|
||||||
"""This class is a no-op but serves as a base class for LLM services."""
|
"""This class is a no-op but serves as a base class for LLM services."""
|
||||||
|
|
||||||
@@ -165,25 +151,25 @@ class TTSService(AIService):
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
aggregate_sentences: bool = True,
|
aggregate_sentences: bool = True,
|
||||||
# if True, subclass is responsible for pushing TextFrames and LLMFullResponseEndFrames
|
# if True, TTSService will push TextFrames and LLMFullResponseEndFrames,
|
||||||
|
# otherwise subclass must do it
|
||||||
push_text_frames: bool = True,
|
push_text_frames: bool = True,
|
||||||
# if True, TTSService will push TTSStoppedFrames, otherwise subclass must do it
|
# TTS output sample rate
|
||||||
push_stop_frames: bool = False,
|
sample_rate: int = 16000,
|
||||||
# if push_stop_frames is True, wait for this idle period before pushing TTSStoppedFrame
|
|
||||||
stop_frame_timeout_s: float = 1.0,
|
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._aggregate_sentences: bool = aggregate_sentences
|
self._aggregate_sentences: bool = aggregate_sentences
|
||||||
self._push_text_frames: bool = push_text_frames
|
self._push_text_frames: bool = push_text_frames
|
||||||
self._push_stop_frames: bool = push_stop_frames
|
|
||||||
self._stop_frame_timeout_s: float = stop_frame_timeout_s
|
|
||||||
self._stop_frame_task: Optional[asyncio.Task] = None
|
|
||||||
self._stop_frame_queue: asyncio.Queue = asyncio.Queue()
|
|
||||||
self._current_sentence: str = ""
|
self._current_sentence: str = ""
|
||||||
|
self._sample_rate: int = sample_rate
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sample_rate(self) -> int:
|
||||||
|
return self._sample_rate
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
pass
|
self.set_model_name(model)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def set_voice(self, voice: str):
|
async def set_voice(self, voice: str):
|
||||||
@@ -218,7 +204,7 @@ class TTSService(AIService):
|
|||||||
if text:
|
if text:
|
||||||
await self._push_tts_frames(text)
|
await self._push_tts_frames(text)
|
||||||
|
|
||||||
async def _push_tts_frames(self, text: str, text_passthrough: bool = True):
|
async def _push_tts_frames(self, text: str):
|
||||||
text = text.strip()
|
text = text.strip()
|
||||||
if not text:
|
if not text:
|
||||||
return
|
return
|
||||||
@@ -248,7 +234,7 @@ class TTSService(AIService):
|
|||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, TTSSpeakFrame):
|
elif isinstance(frame, TTSSpeakFrame):
|
||||||
await self._push_tts_frames(frame.text, False)
|
await self._push_tts_frames(frame.text)
|
||||||
elif isinstance(frame, TTSModelUpdateFrame):
|
elif isinstance(frame, TTSModelUpdateFrame):
|
||||||
await self.set_model(frame.model)
|
await self.set_model(frame.model)
|
||||||
elif isinstance(frame, TTSVoiceUpdateFrame):
|
elif isinstance(frame, TTSVoiceUpdateFrame):
|
||||||
@@ -258,6 +244,25 @@ class TTSService(AIService):
|
|||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|
||||||
|
class AsyncTTSService(TTSService):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
# if True, TTSService will push TTSStoppedFrames, otherwise subclass must do it
|
||||||
|
push_stop_frames: bool = False,
|
||||||
|
# if push_stop_frames is True, wait for this idle period before pushing TTSStoppedFrame
|
||||||
|
stop_frame_timeout_s: float = 1.0,
|
||||||
|
**kwargs):
|
||||||
|
super().__init__(sync=False, **kwargs)
|
||||||
|
self._push_stop_frames: bool = push_stop_frames
|
||||||
|
self._stop_frame_timeout_s: float = stop_frame_timeout_s
|
||||||
|
self._stop_frame_task: Optional[asyncio.Task] = None
|
||||||
|
self._stop_frame_queue: asyncio.Queue = asyncio.Queue()
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def flush_audio(self):
|
||||||
|
pass
|
||||||
|
|
||||||
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:
|
||||||
@@ -283,7 +288,7 @@ class TTSService(AIService):
|
|||||||
if self._push_stop_frames and (
|
if self._push_stop_frames and (
|
||||||
isinstance(frame, StartInterruptionFrame) or
|
isinstance(frame, StartInterruptionFrame) or
|
||||||
isinstance(frame, TTSStartedFrame) or
|
isinstance(frame, TTSStartedFrame) or
|
||||||
isinstance(frame, AudioRawFrame) or
|
isinstance(frame, TTSAudioRawFrame) or
|
||||||
isinstance(frame, TTSStoppedFrame)):
|
isinstance(frame, TTSStoppedFrame)):
|
||||||
await self._stop_frame_queue.put(frame)
|
await self._stop_frame_queue.put(frame)
|
||||||
|
|
||||||
@@ -306,15 +311,6 @@ class TTSService(AIService):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
class AsyncTTSService(TTSService):
|
|
||||||
def __init__(self, **kwargs):
|
|
||||||
super().__init__(**kwargs)
|
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
async def flush_audio(self):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class AsyncWordTTSService(AsyncTTSService):
|
class AsyncWordTTSService(AsyncTTSService):
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
@@ -382,7 +378,7 @@ class STTService(AIService):
|
|||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
pass
|
self.set_model_name(model)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def set_language(self, language: Language):
|
async def set_language(self, language: Language):
|
||||||
|
|||||||
@@ -8,11 +8,12 @@ import base64
|
|||||||
import json
|
import json
|
||||||
import io
|
import io
|
||||||
import copy
|
import copy
|
||||||
from typing import List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from asyncio import CancelledError
|
from asyncio import CancelledError
|
||||||
import re
|
import re
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
Frame,
|
Frame,
|
||||||
@@ -29,6 +30,7 @@ from pipecat.frames.frames import (
|
|||||||
FunctionCallInProgressFrame,
|
FunctionCallInProgressFrame,
|
||||||
StartInterruptionFrame
|
StartInterruptionFrame
|
||||||
)
|
)
|
||||||
|
from pipecat.metrics.metrics import LLMTokenUsage
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.ai_services import LLMService
|
from pipecat.services.ai_services import LLMService
|
||||||
from pipecat.processors.aggregators.openai_llm_context import (
|
from pipecat.processors.aggregators.openai_llm_context import (
|
||||||
@@ -73,20 +75,30 @@ class AnthropicContextAggregatorPair:
|
|||||||
class AnthropicLLMService(LLMService):
|
class AnthropicLLMService(LLMService):
|
||||||
"""This class implements inference with Anthropic's AI models
|
"""This class implements inference with Anthropic's AI models
|
||||||
"""
|
"""
|
||||||
|
class InputParams(BaseModel):
|
||||||
|
enable_prompt_caching_beta: Optional[bool] = False
|
||||||
|
max_tokens: Optional[int] = Field(default_factory=lambda: 4096, ge=1)
|
||||||
|
temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
|
||||||
|
top_k: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=0)
|
||||||
|
top_p: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
model: str = "claude-3-5-sonnet-20240620",
|
model: str = "claude-3-5-sonnet-20240620",
|
||||||
max_tokens: int = 4096,
|
params: InputParams = InputParams(),
|
||||||
enable_prompt_caching_beta: bool = False,
|
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._client = AsyncAnthropic(api_key=api_key)
|
self._client = AsyncAnthropic(api_key=api_key)
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._max_tokens = max_tokens
|
self._max_tokens = params.max_tokens
|
||||||
self._enable_prompt_caching_beta = enable_prompt_caching_beta
|
self._enable_prompt_caching_beta: bool = params.enable_prompt_caching_beta or False
|
||||||
|
self._temperature = params.temperature
|
||||||
|
self._top_k = params.top_k
|
||||||
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -104,6 +116,30 @@ class AnthropicLLMService(LLMService):
|
|||||||
_assistant=assistant
|
_assistant=assistant
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def set_enable_prompt_caching_beta(self, enable_prompt_caching_beta: bool):
|
||||||
|
logger.debug(f"Switching LLM enable_prompt_caching_beta to: [{enable_prompt_caching_beta}]")
|
||||||
|
self._enable_prompt_caching_beta = enable_prompt_caching_beta
|
||||||
|
|
||||||
|
async def set_max_tokens(self, max_tokens: int):
|
||||||
|
logger.debug(f"Switching LLM max_tokens to: [{max_tokens}]")
|
||||||
|
self._max_tokens = max_tokens
|
||||||
|
|
||||||
|
async def set_temperature(self, temperature: float):
|
||||||
|
logger.debug(f"Switching LLM temperature to: [{temperature}]")
|
||||||
|
self._temperature = temperature
|
||||||
|
|
||||||
|
async def set_top_k(self, top_k: float):
|
||||||
|
logger.debug(f"Switching LLM top_k to: [{top_k}]")
|
||||||
|
self._top_k = top_k
|
||||||
|
|
||||||
|
async def set_top_p(self, top_p: float):
|
||||||
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def _process_context(self, context: OpenAILLMContext):
|
async def _process_context(self, context: OpenAILLMContext):
|
||||||
# Usage tracking. We track the usage reported by Anthropic in prompt_tokens and
|
# Usage tracking. We track the usage reported by Anthropic in prompt_tokens and
|
||||||
# completion_tokens. We also estimate the completion tokens from output text
|
# completion_tokens. We also estimate the completion tokens from output text
|
||||||
@@ -133,13 +169,21 @@ class AnthropicLLMService(LLMService):
|
|||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
response = await api_call(
|
params = {
|
||||||
tools=context.tools or [],
|
"tools": context.tools or [],
|
||||||
system=context.system,
|
"system": context.system,
|
||||||
messages=messages,
|
"messages": messages,
|
||||||
model=self._model,
|
"model": self.model_name,
|
||||||
max_tokens=self._max_tokens,
|
"max_tokens": self._max_tokens,
|
||||||
stream=True)
|
"stream": True,
|
||||||
|
"temperature": self._temperature,
|
||||||
|
"top_k": self._top_k,
|
||||||
|
"top_p": self._top_p
|
||||||
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
response = await api_call(**params)
|
||||||
|
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
|
|
||||||
@@ -231,7 +275,7 @@ class AnthropicLLMService(LLMService):
|
|||||||
context = AnthropicLLMContext.from_image_frame(frame)
|
context = AnthropicLLMContext.from_image_frame(frame)
|
||||||
elif isinstance(frame, LLMModelUpdateFrame):
|
elif isinstance(frame, LLMModelUpdateFrame):
|
||||||
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
||||||
self._model = frame.model
|
self.set_model_name(frame.model)
|
||||||
elif isinstance(frame, LLMEnablePromptCachingFrame):
|
elif isinstance(frame, LLMEnablePromptCachingFrame):
|
||||||
logger.debug(f"Setting enable prompt caching to: [{frame.enable}]")
|
logger.debug(f"Setting enable prompt caching to: [{frame.enable}]")
|
||||||
self._enable_prompt_caching_beta = frame.enable
|
self._enable_prompt_caching_beta = frame.enable
|
||||||
@@ -251,15 +295,13 @@ class AnthropicLLMService(LLMService):
|
|||||||
cache_creation_input_tokens: int,
|
cache_creation_input_tokens: int,
|
||||||
cache_read_input_tokens: int):
|
cache_read_input_tokens: int):
|
||||||
if prompt_tokens or completion_tokens or cache_creation_input_tokens or cache_read_input_tokens:
|
if prompt_tokens or completion_tokens or cache_creation_input_tokens or cache_read_input_tokens:
|
||||||
tokens = {
|
tokens = LLMTokenUsage(
|
||||||
"processor": self.name,
|
prompt_tokens=prompt_tokens,
|
||||||
"model": self._model,
|
completion_tokens=completion_tokens,
|
||||||
"prompt_tokens": prompt_tokens,
|
cache_creation_input_tokens=cache_creation_input_tokens,
|
||||||
"completion_tokens": completion_tokens,
|
cache_read_input_tokens=cache_read_input_tokens,
|
||||||
"cache_creation_input_tokens": cache_creation_input_tokens,
|
total_tokens=prompt_tokens + completion_tokens
|
||||||
"cache_read_input_tokens": cache_read_input_tokens,
|
)
|
||||||
"total_tokens": prompt_tokens + completion_tokens
|
|
||||||
}
|
|
||||||
await self.start_llm_usage_metrics(tokens)
|
await self.start_llm_usage_metrics(tokens)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,19 +12,19 @@ from PIL import Image
|
|||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
SystemFrame,
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
URLImageRawFrame)
|
URLImageRawFrame)
|
||||||
|
from pipecat.metrics.metrics import TTSUsageMetricsData
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.ai_services import AsyncAIService, TTSService, ImageGenService
|
from pipecat.services.ai_services import STTService, TTSService, ImageGenService
|
||||||
from pipecat.services.openai import BaseOpenAILLMService
|
from pipecat.services.openai import BaseOpenAILLMService
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
@@ -72,13 +72,21 @@ class AzureLLMService(BaseOpenAILLMService):
|
|||||||
|
|
||||||
|
|
||||||
class AzureTTSService(TTSService):
|
class AzureTTSService(TTSService):
|
||||||
def __init__(self, *, api_key: str, region: str, voice="en-US-SaraNeural", **kwargs):
|
def __init__(
|
||||||
super().__init__(**kwargs)
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str,
|
||||||
|
region: str,
|
||||||
|
voice="en-US-SaraNeural",
|
||||||
|
sample_rate: int = 16000,
|
||||||
|
**kwargs):
|
||||||
|
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
speech_config = SpeechConfig(subscription=api_key, region=region)
|
speech_config = SpeechConfig(subscription=api_key, region=region)
|
||||||
self._speech_synthesizer = SpeechSynthesizer(speech_config=speech_config, audio_config=None)
|
self._speech_synthesizer = SpeechSynthesizer(speech_config=speech_config, audio_config=None)
|
||||||
|
|
||||||
self._voice = voice
|
self._voice = voice
|
||||||
|
self._sample_rate = sample_rate
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -109,7 +117,7 @@ class AzureTTSService(TTSService):
|
|||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
await self.push_frame(TTSStartedFrame())
|
await self.push_frame(TTSStartedFrame())
|
||||||
# Azure always sends a 44-byte header. Strip it off.
|
# Azure always sends a 44-byte header. Strip it off.
|
||||||
yield AudioRawFrame(audio=result.audio_data[44:], sample_rate=16000, num_channels=1)
|
yield TTSAudioRawFrame(audio=result.audio_data[44:], sample_rate=self._sample_rate, num_channels=1)
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
elif result.reason == ResultReason.Canceled:
|
elif result.reason == ResultReason.Canceled:
|
||||||
cancellation_details = result.cancellation_details
|
cancellation_details = result.cancellation_details
|
||||||
@@ -118,7 +126,7 @@ class AzureTTSService(TTSService):
|
|||||||
logger.error(f"{self} error: {cancellation_details.error_details}")
|
logger.error(f"{self} error: {cancellation_details.error_details}")
|
||||||
|
|
||||||
|
|
||||||
class AzureSTTService(AsyncAIService):
|
class AzureSTTService(STTService):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -141,15 +149,11 @@ class AzureSTTService(AsyncAIService):
|
|||||||
speech_config=speech_config, audio_config=audio_config)
|
speech_config=speech_config, audio_config=audio_config)
|
||||||
self._speech_recognizer.recognized.connect(self._on_handle_recognized)
|
self._speech_recognizer.recognized.connect(self._on_handle_recognized)
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
await super().process_frame(frame, direction)
|
await self.start_processing_metrics()
|
||||||
|
self._audio_stream.write(audio)
|
||||||
if isinstance(frame, SystemFrame):
|
await self.stop_processing_metrics()
|
||||||
await self.push_frame(frame, direction)
|
yield None
|
||||||
elif isinstance(frame, AudioRawFrame):
|
|
||||||
self._audio_stream.write(frame.audio)
|
|
||||||
else:
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
@@ -168,7 +172,7 @@ class AzureSTTService(AsyncAIService):
|
|||||||
def _on_handle_recognized(self, event):
|
def _on_handle_recognized(self, event):
|
||||||
if event.result.reason == ResultReason.RecognizedSpeech and len(event.result.text) > 0:
|
if event.result.reason == ResultReason.RecognizedSpeech and len(event.result.text) > 0:
|
||||||
frame = TranscriptionFrame(event.result.text, "", time_now_iso8601())
|
frame = TranscriptionFrame(event.result.text, "", time_now_iso8601())
|
||||||
asyncio.run_coroutine_threadsafe(self.queue_frame(frame), self.get_event_loop())
|
asyncio.run_coroutine_threadsafe(self.push_frame(frame), self.get_event_loop())
|
||||||
|
|
||||||
|
|
||||||
class AzureImageGenServiceREST(ImageGenService):
|
class AzureImageGenServiceREST(ImageGenService):
|
||||||
@@ -188,7 +192,7 @@ class AzureImageGenServiceREST(ImageGenService):
|
|||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._azure_endpoint = endpoint
|
self._azure_endpoint = endpoint
|
||||||
self._api_version = api_version
|
self._api_version = api_version
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._image_size = image_size
|
self._image_size = image_size
|
||||||
self._aiohttp_session = aiohttp_session
|
self._aiohttp_session = aiohttp_session
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ import json
|
|||||||
import uuid
|
import uuid
|
||||||
import base64
|
import base64
|
||||||
import asyncio
|
import asyncio
|
||||||
import time
|
|
||||||
|
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
@@ -17,23 +16,23 @@ from pipecat.frames.frames import (
|
|||||||
CancelFrame,
|
CancelFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
AudioRawFrame,
|
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TextFrame,
|
|
||||||
LLMFullResponseEndFrame
|
LLMFullResponseEndFrame
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.services.ai_services import AsyncWordTTSService
|
from pipecat.services.ai_services import AsyncWordTTSService, TTSService
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
# See .env.example for Cartesia configuration needed
|
# See .env.example for Cartesia configuration needed
|
||||||
try:
|
try:
|
||||||
|
from cartesia import AsyncCartesia
|
||||||
import websockets
|
import websockets
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
@@ -90,13 +89,14 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
# if we're interrupted. Cartesia gives us word-by-word timestamps. We
|
# if we're interrupted. Cartesia gives us word-by-word timestamps. We
|
||||||
# can use those to generate text frames ourselves aligned with the
|
# can use those to generate text frames ourselves aligned with the
|
||||||
# playout timing of the audio!
|
# playout timing of the audio!
|
||||||
super().__init__(aggregate_sentences=True, push_text_frames=False, **kwargs)
|
super().__init__(aggregate_sentences=True,
|
||||||
|
push_text_frames=False, sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._cartesia_version = cartesia_version
|
self._cartesia_version = cartesia_version
|
||||||
self._url = url
|
self._url = url
|
||||||
self._voice_id = voice_id
|
self._voice_id = voice_id
|
||||||
self._model_id = model_id
|
self.set_model_name(model_id)
|
||||||
self._output_format = {
|
self._output_format = {
|
||||||
"container": params.container,
|
"container": params.container,
|
||||||
"encoding": params.encoding,
|
"encoding": params.encoding,
|
||||||
@@ -114,8 +114,8 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
|
await super().set_model(model)
|
||||||
logger.debug(f"Switching TTS model to: [{model}]")
|
logger.debug(f"Switching TTS model to: [{model}]")
|
||||||
self._model_id = model
|
|
||||||
|
|
||||||
async def set_voice(self, voice: str):
|
async def set_voice(self, voice: str):
|
||||||
logger.debug(f"Switching TTS voice to: [{voice}]")
|
logger.debug(f"Switching TTS voice to: [{voice}]")
|
||||||
@@ -173,6 +173,11 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error closing websocket: {e}")
|
logger.error(f"{self} error closing websocket: {e}")
|
||||||
|
|
||||||
|
def _get_websocket(self):
|
||||||
|
if self._websocket:
|
||||||
|
return self._websocket
|
||||||
|
raise Exception("Websocket not connected")
|
||||||
|
|
||||||
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
async def _handle_interruption(self, frame: StartInterruptionFrame, direction: FrameDirection):
|
||||||
await super()._handle_interruption(frame, direction)
|
await super()._handle_interruption(frame, direction)
|
||||||
await self.stop_all_metrics()
|
await self.stop_all_metrics()
|
||||||
@@ -182,12 +187,12 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
async def flush_audio(self):
|
async def flush_audio(self):
|
||||||
if not self._context_id or not self._websocket:
|
if not self._context_id or not self._websocket:
|
||||||
return
|
return
|
||||||
logger.debug("Flushing audio")
|
logger.trace("Flushing audio")
|
||||||
msg = {
|
msg = {
|
||||||
"transcript": "",
|
"transcript": "",
|
||||||
"continue": False,
|
"continue": False,
|
||||||
"context_id": self._context_id,
|
"context_id": self._context_id,
|
||||||
"model_id": self._model_id,
|
"model_id": self.model_name,
|
||||||
"voice": {
|
"voice": {
|
||||||
"mode": "id",
|
"mode": "id",
|
||||||
"id": self._voice_id
|
"id": self._voice_id
|
||||||
@@ -200,7 +205,7 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
|
|
||||||
async def _receive_task_handler(self):
|
async def _receive_task_handler(self):
|
||||||
try:
|
try:
|
||||||
async for message in self._websocket:
|
async for message in self._get_websocket():
|
||||||
msg = json.loads(message)
|
msg = json.loads(message)
|
||||||
if not msg or msg["context_id"] != self._context_id:
|
if not msg or msg["context_id"] != self._context_id:
|
||||||
continue
|
continue
|
||||||
@@ -220,7 +225,7 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
elif msg["type"] == "chunk":
|
elif msg["type"] == "chunk":
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
self.start_word_timestamps()
|
self.start_word_timestamps()
|
||||||
frame = AudioRawFrame(
|
frame = TTSAudioRawFrame(
|
||||||
audio=base64.b64decode(msg["data"]),
|
audio=base64.b64decode(msg["data"]),
|
||||||
sample_rate=self._output_format["sample_rate"],
|
sample_rate=self._output_format["sample_rate"],
|
||||||
num_channels=1
|
num_channels=1
|
||||||
@@ -274,7 +279,7 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
"add_timestamps": True,
|
"add_timestamps": True,
|
||||||
}
|
}
|
||||||
try:
|
try:
|
||||||
await self._websocket.send(json.dumps(msg))
|
await self._get_websocket().send(json.dumps(msg))
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} error sending message: {e}")
|
logger.error(f"{self} error sending message: {e}")
|
||||||
@@ -285,3 +290,84 @@ class CartesiaTTSService(AsyncWordTTSService):
|
|||||||
yield None
|
yield None
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} exception: {e}")
|
logger.error(f"{self} exception: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
class CartesiaHttpTTSService(TTSService):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str,
|
||||||
|
voice_id: str,
|
||||||
|
model_id: str = "sonic-english",
|
||||||
|
base_url: str = "https://api.cartesia.ai",
|
||||||
|
encoding: str = "pcm_s16le",
|
||||||
|
sample_rate: int = 16000,
|
||||||
|
language: str = "en",
|
||||||
|
**kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
self._api_key = api_key
|
||||||
|
self._voice_id = voice_id
|
||||||
|
self._model_id = model_id
|
||||||
|
self._output_format = {
|
||||||
|
"container": "raw",
|
||||||
|
"encoding": encoding,
|
||||||
|
"sample_rate": sample_rate,
|
||||||
|
}
|
||||||
|
self._language = language
|
||||||
|
|
||||||
|
self._client = AsyncCartesia(api_key=api_key, base_url=base_url)
|
||||||
|
|
||||||
|
def can_generate_metrics(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def set_model(self, model: str):
|
||||||
|
logger.debug(f"Switching TTS model to: [{model}]")
|
||||||
|
self._model_id = model
|
||||||
|
|
||||||
|
async def set_voice(self, voice: str):
|
||||||
|
logger.debug(f"Switching TTS voice to: [{voice}]")
|
||||||
|
self._voice_id = voice
|
||||||
|
|
||||||
|
async def set_language(self, language: Language):
|
||||||
|
logger.debug(f"Switching TTS language to: [{language}]")
|
||||||
|
self._language = language_to_cartesia_language(language)
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
await super().stop(frame)
|
||||||
|
await self._client.close()
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
await super().cancel(frame)
|
||||||
|
await self._client.close()
|
||||||
|
|
||||||
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
|
|
||||||
|
await self.push_frame(TTSStartedFrame())
|
||||||
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
|
try:
|
||||||
|
output = await self._client.tts.sse(
|
||||||
|
model_id=self._model_id,
|
||||||
|
transcript=text,
|
||||||
|
voice_id=self._voice_id,
|
||||||
|
output_format=self._output_format,
|
||||||
|
language=self._language,
|
||||||
|
stream=False
|
||||||
|
)
|
||||||
|
|
||||||
|
await self.stop_ttfb_metrics()
|
||||||
|
|
||||||
|
frame = TTSAudioRawFrame(
|
||||||
|
audio=output["audio"],
|
||||||
|
sample_rate=self._output_format["sample_rate"],
|
||||||
|
num_channels=1
|
||||||
|
)
|
||||||
|
yield frame
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} exception: {e}")
|
||||||
|
|
||||||
|
await self.start_tts_usage_metrics(text)
|
||||||
|
await self.push_frame(TTSStoppedFrame())
|
||||||
|
|||||||
@@ -9,13 +9,13 @@ import aiohttp
|
|||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InterimTranscriptionFrame,
|
InterimTranscriptionFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TranscriptionFrame)
|
TranscriptionFrame)
|
||||||
@@ -101,7 +101,8 @@ class DeepgramTTSService(TTSService):
|
|||||||
await self.push_frame(TTSStartedFrame())
|
await self.push_frame(TTSStartedFrame())
|
||||||
async for data in r.content:
|
async for data in r.content:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = AudioRawFrame(audio=data, sample_rate=self._sample_rate, num_channels=1)
|
frame = TTSAudioRawFrame(
|
||||||
|
audio=data, sample_rate=self._sample_rate, num_channels=1)
|
||||||
yield frame
|
yield frame
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -135,6 +136,7 @@ class DeepgramSTTService(STTService):
|
|||||||
self._connection.on(LiveTranscriptionEvents.Transcript, self._on_message)
|
self._connection.on(LiveTranscriptionEvents.Transcript, self._on_message)
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
|
await super().set_model(model)
|
||||||
logger.debug(f"Switching STT model to: [{model}]")
|
logger.debug(f"Switching STT model to: [{model}]")
|
||||||
self._live_options.model = model
|
self._live_options.model = model
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
@@ -161,8 +163,8 @@ class DeepgramSTTService(STTService):
|
|||||||
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
await self.start_processing_metrics()
|
await self.start_processing_metrics()
|
||||||
await self._connection.send(audio)
|
await self._connection.send(audio)
|
||||||
yield None
|
|
||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
|
yield None
|
||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
if await self._connection.start(self._live_options):
|
if await self._connection.start(self._live_options):
|
||||||
|
|||||||
@@ -7,31 +7,32 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
|
from typing import Any, AsyncGenerator, Dict, List, Literal, Mapping, Optional, Tuple
|
||||||
|
|
||||||
from typing import Any, AsyncGenerator, List, Literal, Mapping, Tuple
|
from loguru import logger
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, model_validator
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame)
|
TTSStoppedFrame,
|
||||||
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.ai_services import AsyncWordTTSService
|
from pipecat.services.ai_services import AsyncWordTTSService
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
# See .env.example for ElevenLabs configuration needed
|
# See .env.example for ElevenLabs configuration needed
|
||||||
try:
|
try:
|
||||||
import websockets
|
import websockets
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error(
|
logger.error(
|
||||||
"In order to use ElevenLabs, you need to `pip install pipecat-ai[elevenlabs]`. Also, set `ELEVENLABS_API_KEY` environment variable.")
|
"In order to use ElevenLabs, you need to `pip install pipecat-ai[elevenlabs]`. Also, set `ELEVENLABS_API_KEY` environment variable."
|
||||||
|
)
|
||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
|
||||||
@@ -49,7 +50,7 @@ def sample_rate_from_output_format(output_format: str) -> int:
|
|||||||
|
|
||||||
|
|
||||||
def calculate_word_times(
|
def calculate_word_times(
|
||||||
alignment_info: Mapping[str, Any], cumulative_time: float
|
alignment_info: Mapping[str, Any], cumulative_time: float
|
||||||
) -> List[Tuple[str, float]]:
|
) -> List[Tuple[str, float]]:
|
||||||
zipped_times = list(zip(alignment_info["chars"], alignment_info["charStartTimesMs"]))
|
zipped_times = list(zip(alignment_info["chars"], alignment_info["charStartTimesMs"]))
|
||||||
|
|
||||||
@@ -59,7 +60,7 @@ def calculate_word_times(
|
|||||||
# and using the previous word time, also taking into account there might not
|
# and using the previous word time, also taking into account there might not
|
||||||
# be a space at the end.
|
# be a space at the end.
|
||||||
times = []
|
times = []
|
||||||
for (i, (a, b)) in enumerate(zipped_times):
|
for i, (a, b) in enumerate(zipped_times):
|
||||||
if a == " " or i == len(zipped_times) - 1:
|
if a == " " or i == len(zipped_times) - 1:
|
||||||
t = cumulative_time + (zipped_times[i - 1][1] / 1000.0)
|
t = cumulative_time + (zipped_times[i - 1][1] / 1000.0)
|
||||||
times.append(t)
|
times.append(t)
|
||||||
@@ -72,16 +73,32 @@ def calculate_word_times(
|
|||||||
class ElevenLabsTTSService(AsyncWordTTSService):
|
class ElevenLabsTTSService(AsyncWordTTSService):
|
||||||
class InputParams(BaseModel):
|
class InputParams(BaseModel):
|
||||||
output_format: Literal["pcm_16000", "pcm_22050", "pcm_24000", "pcm_44100"] = "pcm_16000"
|
output_format: Literal["pcm_16000", "pcm_22050", "pcm_24000", "pcm_44100"] = "pcm_16000"
|
||||||
|
optimize_streaming_latency: Optional[str] = None
|
||||||
|
stability: Optional[float] = None
|
||||||
|
similarity_boost: Optional[float] = None
|
||||||
|
style: Optional[float] = None
|
||||||
|
use_speaker_boost: Optional[bool] = None
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def validate_voice_settings(self):
|
||||||
|
stability = self.stability
|
||||||
|
similarity_boost = self.similarity_boost
|
||||||
|
if (stability is None) != (similarity_boost is None):
|
||||||
|
raise ValueError(
|
||||||
|
"Both 'stability' and 'similarity_boost' must be provided when using voice settings"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
voice_id: str,
|
voice_id: str,
|
||||||
model: str = "eleven_turbo_v2_5",
|
model: str = "eleven_turbo_v2_5",
|
||||||
url: str = "wss://api.elevenlabs.io",
|
url: str = "wss://api.elevenlabs.io",
|
||||||
params: InputParams = InputParams(),
|
params: InputParams = InputParams(),
|
||||||
**kwargs):
|
**kwargs,
|
||||||
|
):
|
||||||
# Aggregating sentences still gives cleaner-sounding results and fewer
|
# Aggregating sentences still gives cleaner-sounding results and fewer
|
||||||
# artifacts than streaming one word at a time. On average, waiting for a
|
# artifacts than streaming one word at a time. On average, waiting for a
|
||||||
# full sentence should only "cost" us 15ms or so with GPT-4o or a Llama
|
# full sentence should only "cost" us 15ms or so with GPT-4o or a Llama
|
||||||
@@ -101,15 +118,16 @@ class ElevenLabsTTSService(AsyncWordTTSService):
|
|||||||
push_text_frames=False,
|
push_text_frames=False,
|
||||||
push_stop_frames=True,
|
push_stop_frames=True,
|
||||||
stop_frame_timeout_s=2.0,
|
stop_frame_timeout_s=2.0,
|
||||||
**kwargs
|
sample_rate=sample_rate_from_output_format(params.output_format),
|
||||||
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._voice_id = voice_id
|
self._voice_id = voice_id
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._url = url
|
self._url = url
|
||||||
self._params = params
|
self._params = params
|
||||||
self._sample_rate = sample_rate_from_output_format(params.output_format)
|
self._voice_settings = self._set_voice_settings()
|
||||||
|
|
||||||
# Websocket connection to ElevenLabs.
|
# Websocket connection to ElevenLabs.
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
@@ -121,9 +139,30 @@ class ElevenLabsTTSService(AsyncWordTTSService):
|
|||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def _set_voice_settings(self):
|
||||||
|
voice_settings = {}
|
||||||
|
if self._params.stability is not None and self._params.similarity_boost is not None:
|
||||||
|
voice_settings["stability"] = self._params.stability
|
||||||
|
voice_settings["similarity_boost"] = self._params.similarity_boost
|
||||||
|
if self._params.style is not None:
|
||||||
|
voice_settings["style"] = self._params.style
|
||||||
|
if self._params.use_speaker_boost is not None:
|
||||||
|
voice_settings["use_speaker_boost"] = self._params.use_speaker_boost
|
||||||
|
else:
|
||||||
|
if self._params.style is not None:
|
||||||
|
logger.warning(
|
||||||
|
"'style' is set but will not be applied because 'stability' and 'similarity_boost' are not both set."
|
||||||
|
)
|
||||||
|
if self._params.use_speaker_boost is not None:
|
||||||
|
logger.warning(
|
||||||
|
"'use_speaker_boost' is set but will not be applied because 'stability' and 'similarity_boost' are not both set."
|
||||||
|
)
|
||||||
|
|
||||||
|
return voice_settings or None
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
|
await super().set_model(model)
|
||||||
logger.debug(f"Switching TTS model to: [{model}]")
|
logger.debug(f"Switching TTS model to: [{model}]")
|
||||||
self._model = model
|
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
@@ -133,6 +172,28 @@ class ElevenLabsTTSService(AsyncWordTTSService):
|
|||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
|
async def set_voice_settings(
|
||||||
|
self,
|
||||||
|
stability: Optional[float] = None,
|
||||||
|
similarity_boost: Optional[float] = None,
|
||||||
|
style: Optional[float] = None,
|
||||||
|
use_speaker_boost: Optional[bool] = None,
|
||||||
|
):
|
||||||
|
self._params.stability = stability if stability is not None else self._params.stability
|
||||||
|
self._params.similarity_boost = (
|
||||||
|
similarity_boost if similarity_boost is not None else self._params.similarity_boost
|
||||||
|
)
|
||||||
|
self._params.style = style if style is not None else self._params.style
|
||||||
|
self._params.use_speaker_boost = (
|
||||||
|
use_speaker_boost if use_speaker_boost is not None else self._params.use_speaker_boost
|
||||||
|
)
|
||||||
|
|
||||||
|
self._set_voice_settings()
|
||||||
|
|
||||||
|
if self._websocket:
|
||||||
|
msg = {"voice_settings": self._voice_settings}
|
||||||
|
await self._websocket.send(json.dumps(msg))
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
await self._connect()
|
await self._connect()
|
||||||
@@ -160,18 +221,24 @@ class ElevenLabsTTSService(AsyncWordTTSService):
|
|||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
try:
|
try:
|
||||||
voice_id = self._voice_id
|
voice_id = self._voice_id
|
||||||
model = self._model
|
model = self.model_name
|
||||||
output_format = self._params.output_format
|
output_format = self._params.output_format
|
||||||
url = f"{self._url}/v1/text-to-speech/{voice_id}/stream-input?model_id={model}&output_format={output_format}"
|
url = f"{self._url}/v1/text-to-speech/{voice_id}/stream-input?model_id={model}&output_format={output_format}"
|
||||||
|
|
||||||
|
if self._params.optimize_streaming_latency:
|
||||||
|
url += f"&optimize_streaming_latency={self._params.optimize_streaming_latency}"
|
||||||
|
|
||||||
self._websocket = await websockets.connect(url)
|
self._websocket = await websockets.connect(url)
|
||||||
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
|
||||||
self._keepalive_task = self.get_event_loop().create_task(self._keepalive_task_handler())
|
self._keepalive_task = self.get_event_loop().create_task(self._keepalive_task_handler())
|
||||||
|
|
||||||
# According to ElevenLabs, we should always start with a single space.
|
# According to ElevenLabs, we should always start with a single space.
|
||||||
msg = {
|
msg: Dict[str, Any] = {
|
||||||
"text": " ",
|
"text": " ",
|
||||||
"xi_api_key": self._api_key,
|
"xi_api_key": self._api_key,
|
||||||
}
|
}
|
||||||
|
if self._voice_settings:
|
||||||
|
msg["voice_settings"] = self._voice_settings
|
||||||
await self._websocket.send(json.dumps(msg))
|
await self._websocket.send(json.dumps(msg))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self} initialization error: {e}")
|
logger.error(f"{self} initialization error: {e}")
|
||||||
@@ -209,7 +276,7 @@ class ElevenLabsTTSService(AsyncWordTTSService):
|
|||||||
self.start_word_timestamps()
|
self.start_word_timestamps()
|
||||||
|
|
||||||
audio = base64.b64decode(msg["audio"])
|
audio = base64.b64decode(msg["audio"])
|
||||||
frame = AudioRawFrame(audio, self._sample_rate, 1)
|
frame = TTSAudioRawFrame(audio, self.sample_rate, 1)
|
||||||
await self.push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
if msg.get("alignment"):
|
if msg.get("alignment"):
|
||||||
|
|||||||
@@ -43,9 +43,10 @@ class FalImageGenService(ImageGenService):
|
|||||||
aiohttp_session: aiohttp.ClientSession,
|
aiohttp_session: aiohttp.ClientSession,
|
||||||
model: str = "fal-ai/fast-sdxl",
|
model: str = "fal-ai/fast-sdxl",
|
||||||
key: str | None = None,
|
key: str | None = None,
|
||||||
|
**kwargs
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__(**kwargs)
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._params = params
|
self._params = params
|
||||||
self._aiohttp_session = aiohttp_session
|
self._aiohttp_session = aiohttp_session
|
||||||
if key:
|
if key:
|
||||||
@@ -55,7 +56,7 @@ class FalImageGenService(ImageGenService):
|
|||||||
logger.debug(f"Generating image from prompt: {prompt}")
|
logger.debug(f"Generating image from prompt: {prompt}")
|
||||||
|
|
||||||
response = await fal_client.run_async(
|
response = await fal_client.run_async(
|
||||||
self._model,
|
self.model_name,
|
||||||
arguments={"prompt": prompt, **self._params.model_dump(exclude_none=True)}
|
arguments={"prompt": prompt, **self._params.model_dump(exclude_none=True)}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -22,4 +22,4 @@ class FireworksLLMService(BaseOpenAILLMService):
|
|||||||
*,
|
*,
|
||||||
model: str = "accounts/fireworks/models/firefunction-v1",
|
model: str = "accounts/fireworks/models/firefunction-v1",
|
||||||
base_url: str = "https://api.fireworks.ai/inference/v1"):
|
base_url: str = "https://api.fireworks.ai/inference/v1"):
|
||||||
super().__init__(model, base_url)
|
super().__init__(model=model, base_url=base_url)
|
||||||
|
|||||||
@@ -7,20 +7,17 @@
|
|||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
|
|
||||||
from typing import Optional
|
from typing import AsyncGenerator, Optional
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
InterimTranscriptionFrame,
|
InterimTranscriptionFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
SystemFrame,
|
|
||||||
TranscriptionFrame)
|
TranscriptionFrame)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.services.ai_services import STTService
|
||||||
from pipecat.services.ai_services import AsyncAIService
|
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -35,7 +32,7 @@ except ModuleNotFoundError as e:
|
|||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
|
||||||
class GladiaSTTService(AsyncAIService):
|
class GladiaSTTService(STTService):
|
||||||
class InputParams(BaseModel):
|
class InputParams(BaseModel):
|
||||||
sample_rate: Optional[int] = 16000
|
sample_rate: Optional[int] = 16000
|
||||||
language: Optional[str] = "english"
|
language: Optional[str] = "english"
|
||||||
@@ -50,23 +47,13 @@ class GladiaSTTService(AsyncAIService):
|
|||||||
confidence: float = 0.5,
|
confidence: float = 0.5,
|
||||||
params: InputParams = InputParams(),
|
params: InputParams = InputParams(),
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._url = url
|
self._url = url
|
||||||
self._params = params
|
self._params = params
|
||||||
self._confidence = confidence
|
self._confidence = confidence
|
||||||
|
|
||||||
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
|
||||||
await super().process_frame(frame, direction)
|
|
||||||
|
|
||||||
if isinstance(frame, SystemFrame):
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
elif isinstance(frame, AudioRawFrame):
|
|
||||||
await self._send_audio(frame)
|
|
||||||
else:
|
|
||||||
await self.queue_frame(frame, direction)
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
self._websocket = await websockets.connect(self._url)
|
self._websocket = await websockets.connect(self._url)
|
||||||
@@ -81,6 +68,12 @@ class GladiaSTTService(AsyncAIService):
|
|||||||
await super().cancel(frame)
|
await super().cancel(frame)
|
||||||
await self._websocket.close()
|
await self._websocket.close()
|
||||||
|
|
||||||
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
|
await self.start_processing_metrics()
|
||||||
|
await self._send_audio(audio)
|
||||||
|
await self.stop_processing_metrics()
|
||||||
|
yield None
|
||||||
|
|
||||||
async def _setup_gladia(self):
|
async def _setup_gladia(self):
|
||||||
configuration = {
|
configuration = {
|
||||||
"x_gladia_key": self._api_key,
|
"x_gladia_key": self._api_key,
|
||||||
@@ -92,9 +85,9 @@ class GladiaSTTService(AsyncAIService):
|
|||||||
|
|
||||||
await self._websocket.send(json.dumps(configuration))
|
await self._websocket.send(json.dumps(configuration))
|
||||||
|
|
||||||
async def _send_audio(self, frame: AudioRawFrame):
|
async def _send_audio(self, audio: bytes):
|
||||||
message = {
|
message = {
|
||||||
'frames': base64.b64encode(frame.audio).decode("utf-8")
|
'frames': base64.b64encode(audio).decode("utf-8")
|
||||||
}
|
}
|
||||||
await self._websocket.send(json.dumps(message))
|
await self._websocket.send(json.dumps(message))
|
||||||
|
|
||||||
@@ -113,6 +106,6 @@ class GladiaSTTService(AsyncAIService):
|
|||||||
transcript = utterance["transcription"]
|
transcript = utterance["transcription"]
|
||||||
if confidence >= self._confidence:
|
if confidence >= self._confidence:
|
||||||
if type == "final":
|
if type == "final":
|
||||||
await self.queue_frame(TranscriptionFrame(transcript, "", time_now_iso8601()))
|
await self.push_frame(TranscriptionFrame(transcript, "", time_now_iso8601()))
|
||||||
else:
|
else:
|
||||||
await self.queue_frame(InterimTranscriptionFrame(transcript, "", time_now_iso8601()))
|
await self.push_frame(InterimTranscriptionFrame(transcript, "", time_now_iso8601()))
|
||||||
|
|||||||
@@ -50,6 +50,7 @@ class GoogleLLMService(LLMService):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def _create_client(self, model: str):
|
def _create_client(self, model: str):
|
||||||
|
self.set_model_name(model)
|
||||||
self._client = gai.GenerativeModel(model)
|
self._client = gai.GenerativeModel(model)
|
||||||
|
|
||||||
def _get_messages_from_openai_context(
|
def _get_messages_from_openai_context(
|
||||||
|
|||||||
@@ -10,13 +10,13 @@ from typing import AsyncGenerator
|
|||||||
|
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
)
|
)
|
||||||
@@ -46,7 +46,7 @@ class LmntTTSService(AsyncTTSService):
|
|||||||
**kwargs):
|
**kwargs):
|
||||||
# Let TTSService produce TTSStoppedFrames after a short delay of
|
# Let TTSService produce TTSStoppedFrames after a short delay of
|
||||||
# no activity.
|
# no activity.
|
||||||
super().__init__(push_stop_frames=True, **kwargs)
|
super().__init__(sync=False, push_stop_frames=True, sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._voice_id = voice_id
|
self._voice_id = voice_id
|
||||||
@@ -126,7 +126,7 @@ class LmntTTSService(AsyncTTSService):
|
|||||||
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
|
await self.push_error(ErrorFrame(f'{self} error: {msg["error"]}'))
|
||||||
elif "audio" in msg:
|
elif "audio" in msg:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = AudioRawFrame(
|
frame = TTSAudioRawFrame(
|
||||||
audio=msg["audio"],
|
audio=msg["audio"],
|
||||||
sample_rate=self._output_format["sample_rate"],
|
sample_rate=self._output_format["sample_rate"],
|
||||||
num_channels=1
|
num_channels=1
|
||||||
|
|||||||
@@ -46,12 +46,15 @@ def detect_device():
|
|||||||
class MoondreamService(VisionService):
|
class MoondreamService(VisionService):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
model="vikhyatk/moondream2",
|
model="vikhyatk/moondream2",
|
||||||
revision="2024-04-02",
|
revision="2024-08-26",
|
||||||
use_cpu=False
|
use_cpu=False,
|
||||||
|
**kwargs
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
|
self.set_model_name(model)
|
||||||
|
|
||||||
if not use_cpu:
|
if not use_cpu:
|
||||||
device, dtype = detect_device()
|
device, dtype = detect_device()
|
||||||
@@ -72,7 +75,7 @@ class MoondreamService(VisionService):
|
|||||||
|
|
||||||
async def run_vision(self, frame: VisionImageRawFrame) -> AsyncGenerator[Frame, None]:
|
async def run_vision(self, frame: VisionImageRawFrame) -> AsyncGenerator[Frame, None]:
|
||||||
if not self._model:
|
if not self._model:
|
||||||
logger.error(f"{self} error: Moondream model not available")
|
logger.error(f"{self} error: Moondream model not available ({self.model_name})")
|
||||||
yield ErrorFrame("Moondream model not available")
|
yield ErrorFrame("Moondream model not available")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -11,19 +11,20 @@ import json
|
|||||||
import httpx
|
import httpx
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
from typing import AsyncGenerator, List, Literal
|
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
LLMModelUpdateFrame,
|
LLMModelUpdateFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
TextFrame,
|
TextFrame,
|
||||||
@@ -33,6 +34,7 @@ from pipecat.frames.frames import (
|
|||||||
FunctionCallInProgressFrame,
|
FunctionCallInProgressFrame,
|
||||||
StartInterruptionFrame
|
StartInterruptionFrame
|
||||||
)
|
)
|
||||||
|
from pipecat.metrics.metrics import LLMTokenUsage
|
||||||
from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator, LLMAssistantContextAggregator
|
from pipecat.processors.aggregators.llm_response import LLMUserContextAggregator, LLMAssistantContextAggregator
|
||||||
|
|
||||||
from pipecat.processors.aggregators.openai_llm_context import (
|
from pipecat.processors.aggregators.openai_llm_context import (
|
||||||
@@ -47,7 +49,7 @@ from pipecat.services.ai_services import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from openai import AsyncOpenAI, AsyncStream, DefaultAsyncHttpxClient, BadRequestError
|
from openai import AsyncOpenAI, AsyncStream, DefaultAsyncHttpxClient, BadRequestError, NOT_GIVEN
|
||||||
from openai.types.chat import ChatCompletionChunk, ChatCompletionMessageParam
|
from openai.types.chat import ChatCompletionChunk, ChatCompletionMessageParam
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
@@ -55,6 +57,17 @@ except ModuleNotFoundError as e:
|
|||||||
"In order to use OpenAI, you need to `pip install pipecat-ai[openai]`. Also, set `OPENAI_API_KEY` environment variable.")
|
"In order to use OpenAI, you need to `pip install pipecat-ai[openai]`. Also, set `OPENAI_API_KEY` environment variable.")
|
||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
ValidVoice = Literal["alloy", "echo", "fable", "onyx", "nova", "shimmer"]
|
||||||
|
|
||||||
|
VALID_VOICES: Dict[str, ValidVoice] = {
|
||||||
|
"alloy": "alloy",
|
||||||
|
"echo": "echo",
|
||||||
|
"fable": "fable",
|
||||||
|
"onyx": "onyx",
|
||||||
|
"nova": "nova",
|
||||||
|
"shimmer": "shimmer",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
class OpenAIUnhandledFunctionException(Exception):
|
class OpenAIUnhandledFunctionException(Exception):
|
||||||
pass
|
pass
|
||||||
@@ -69,11 +82,33 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
as well as tool choices and the tool, which is used if requesting function
|
as well as tool choices and the tool, which is used if requesting function
|
||||||
calls from the LLM.
|
calls from the LLM.
|
||||||
"""
|
"""
|
||||||
|
class InputParams(BaseModel):
|
||||||
|
frequency_penalty: Optional[float] = Field(
|
||||||
|
default_factory=lambda: NOT_GIVEN, ge=-2.0, le=2.0)
|
||||||
|
presence_penalty: Optional[float] = Field(
|
||||||
|
default_factory=lambda: NOT_GIVEN, ge=-2.0, le=2.0)
|
||||||
|
seed: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=0)
|
||||||
|
temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=2.0)
|
||||||
|
top_p: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(self, *, model: str, api_key=None, base_url=None, **kwargs):
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
model: str,
|
||||||
|
api_key=None,
|
||||||
|
base_url=None,
|
||||||
|
params: InputParams = InputParams(),
|
||||||
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._model: str = model
|
self.set_model_name(model)
|
||||||
self._client = self.create_client(api_key=api_key, base_url=base_url, **kwargs)
|
self._client = self.create_client(api_key=api_key, base_url=base_url, **kwargs)
|
||||||
|
self._frequency_penalty = params.frequency_penalty
|
||||||
|
self._presence_penalty = params.presence_penalty
|
||||||
|
self._seed = params.seed
|
||||||
|
self._temperature = params.temperature
|
||||||
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def create_client(self, api_key=None, base_url=None, **kwargs):
|
def create_client(self, api_key=None, base_url=None, **kwargs):
|
||||||
return AsyncOpenAI(
|
return AsyncOpenAI(
|
||||||
@@ -88,18 +123,52 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def set_frequency_penalty(self, frequency_penalty: float):
|
||||||
|
logger.debug(f"Switching LLM frequency_penalty to: [{frequency_penalty}]")
|
||||||
|
self._frequency_penalty = frequency_penalty
|
||||||
|
|
||||||
|
async def set_presence_penalty(self, presence_penalty: float):
|
||||||
|
logger.debug(f"Switching LLM presence_penalty to: [{presence_penalty}]")
|
||||||
|
self._presence_penalty = presence_penalty
|
||||||
|
|
||||||
|
async def set_seed(self, seed: int):
|
||||||
|
logger.debug(f"Switching LLM seed to: [{seed}]")
|
||||||
|
self._seed = seed
|
||||||
|
|
||||||
|
async def set_temperature(self, temperature: float):
|
||||||
|
logger.debug(f"Switching LLM temperature to: [{temperature}]")
|
||||||
|
self._temperature = temperature
|
||||||
|
|
||||||
|
async def set_top_p(self, top_p: float):
|
||||||
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def get_chat_completions(
|
async def get_chat_completions(
|
||||||
self,
|
self,
|
||||||
context: OpenAILLMContext,
|
context: OpenAILLMContext,
|
||||||
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
||||||
chunks = await self._client.chat.completions.create(
|
|
||||||
model=self._model,
|
params = {
|
||||||
stream=True,
|
"model": self.model_name,
|
||||||
messages=messages,
|
"stream": True,
|
||||||
tools=context.tools,
|
"messages": messages,
|
||||||
tool_choice=context.tool_choice,
|
"tools": context.tools,
|
||||||
stream_options={"include_usage": True}
|
"tool_choice": context.tool_choice,
|
||||||
)
|
"stream_options": {"include_usage": True},
|
||||||
|
"frequency_penalty": self._frequency_penalty,
|
||||||
|
"presence_penalty": self._presence_penalty,
|
||||||
|
"seed": self._seed,
|
||||||
|
"temperature": self._temperature,
|
||||||
|
"top_p": self._top_p,
|
||||||
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
chunks = await self._client.chat.completions.create(**params)
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
async def _stream_chat_completions(
|
async def _stream_chat_completions(
|
||||||
@@ -137,13 +206,11 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
|
|
||||||
async for chunk in chunk_stream:
|
async for chunk in chunk_stream:
|
||||||
if chunk.usage:
|
if chunk.usage:
|
||||||
tokens = {
|
tokens = LLMTokenUsage(
|
||||||
"processor": self.name,
|
prompt_tokens=chunk.usage.prompt_tokens,
|
||||||
"model": self._model,
|
completion_tokens=chunk.usage.completion_tokens,
|
||||||
"prompt_tokens": chunk.usage.prompt_tokens,
|
total_tokens=chunk.usage.total_tokens
|
||||||
"completion_tokens": chunk.usage.completion_tokens,
|
)
|
||||||
"total_tokens": chunk.usage.total_tokens
|
|
||||||
}
|
|
||||||
await self.start_llm_usage_metrics(tokens)
|
await self.start_llm_usage_metrics(tokens)
|
||||||
|
|
||||||
if len(chunk.choices) == 0:
|
if len(chunk.choices) == 0:
|
||||||
@@ -212,7 +279,7 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
context = OpenAILLMContext.from_image_frame(frame)
|
context = OpenAILLMContext.from_image_frame(frame)
|
||||||
elif isinstance(frame, LLMModelUpdateFrame):
|
elif isinstance(frame, LLMModelUpdateFrame):
|
||||||
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
||||||
self._model = frame.model
|
self.set_model_name(frame.model)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
@@ -238,8 +305,13 @@ class OpenAIContextAggregatorPair:
|
|||||||
|
|
||||||
class OpenAILLMService(BaseOpenAILLMService):
|
class OpenAILLMService(BaseOpenAILLMService):
|
||||||
|
|
||||||
def __init__(self, *, model: str = "gpt-4o", **kwargs):
|
def __init__(
|
||||||
super().__init__(model=model, **kwargs)
|
self,
|
||||||
|
*,
|
||||||
|
model: str = "gpt-4o",
|
||||||
|
params: BaseOpenAILLMService.InputParams = BaseOpenAILLMService.InputParams(),
|
||||||
|
**kwargs):
|
||||||
|
super().__init__(model=model, params=params, **kwargs)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create_context_aggregator(context: OpenAILLMContext) -> OpenAIContextAggregatorPair:
|
def create_context_aggregator(context: OpenAILLMContext) -> OpenAIContextAggregatorPair:
|
||||||
@@ -262,7 +334,7 @@ class OpenAIImageGenService(ImageGenService):
|
|||||||
model: str = "dall-e-3",
|
model: str = "dall-e-3",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._image_size = image_size
|
self._image_size = image_size
|
||||||
self._client = AsyncOpenAI(api_key=api_key)
|
self._client = AsyncOpenAI(api_key=api_key)
|
||||||
self._aiohttp_session = aiohttp_session
|
self._aiohttp_session = aiohttp_session
|
||||||
@@ -272,7 +344,7 @@ class OpenAIImageGenService(ImageGenService):
|
|||||||
|
|
||||||
image = await self._client.images.generate(
|
image = await self._client.images.generate(
|
||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
model=self._model,
|
model=self.model_name,
|
||||||
n=1,
|
n=1,
|
||||||
size=self._image_size
|
size=self._image_size
|
||||||
)
|
)
|
||||||
@@ -307,13 +379,15 @@ class OpenAITTSService(TTSService):
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str | None = None,
|
api_key: str | None = None,
|
||||||
voice: Literal["alloy", "echo", "fable", "onyx", "nova", "shimmer"] = "alloy",
|
voice: str = "alloy",
|
||||||
model: Literal["tts-1", "tts-1-hd"] = "tts-1",
|
model: Literal["tts-1", "tts-1-hd"] = "tts-1",
|
||||||
|
sample_rate: int = 24000,
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
self._voice = voice
|
self._voice: ValidVoice = VALID_VOICES.get(voice, "alloy")
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
|
self._sample_rate = sample_rate
|
||||||
|
|
||||||
self._client = AsyncOpenAI(api_key=api_key)
|
self._client = AsyncOpenAI(api_key=api_key)
|
||||||
|
|
||||||
@@ -322,7 +396,11 @@ class OpenAITTSService(TTSService):
|
|||||||
|
|
||||||
async def set_voice(self, voice: str):
|
async def set_voice(self, voice: str):
|
||||||
logger.debug(f"Switching TTS voice to: [{voice}]")
|
logger.debug(f"Switching TTS voice to: [{voice}]")
|
||||||
self._voice = voice
|
self._voice = VALID_VOICES.get(voice, self._voice)
|
||||||
|
|
||||||
|
async def set_model(self, model: str):
|
||||||
|
logger.debug(f"Switching TTS model to: [{model}]")
|
||||||
|
self._model = model
|
||||||
|
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
logger.debug(f"Generating TTS: [{text}]")
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
@@ -331,7 +409,7 @@ class OpenAITTSService(TTSService):
|
|||||||
|
|
||||||
async with self._client.audio.speech.with_streaming_response.create(
|
async with self._client.audio.speech.with_streaming_response.create(
|
||||||
input=text,
|
input=text,
|
||||||
model=self._model,
|
model=self.model_name,
|
||||||
voice=self._voice,
|
voice=self._voice,
|
||||||
response_format="pcm",
|
response_format="pcm",
|
||||||
) as r:
|
) as r:
|
||||||
@@ -348,7 +426,7 @@ class OpenAITTSService(TTSService):
|
|||||||
async for chunk in r.iter_bytes(8192):
|
async for chunk in r.iter_bytes(8192):
|
||||||
if len(chunk) > 0:
|
if len(chunk) > 0:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = AudioRawFrame(chunk, 24_000, 1)
|
frame = TTSAudioRawFrame(chunk, self.sample_rate, 1)
|
||||||
yield frame
|
yield frame
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
except BadRequestError as e:
|
except BadRequestError as e:
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ class OpenPipeLLMService(BaseOpenAILLMService):
|
|||||||
context: OpenAILLMContext,
|
context: OpenAILLMContext,
|
||||||
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
||||||
chunks = await self._client.chat.completions.create(
|
chunks = await self._client.chat.completions.create(
|
||||||
model=self._model,
|
model=self.model_name,
|
||||||
stream=True,
|
stream=True,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
openpipe={
|
openpipe={
|
||||||
|
|||||||
@@ -9,7 +9,11 @@ import struct
|
|||||||
|
|
||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, Frame, TTSStartedFrame, TTSStoppedFrame
|
from pipecat.frames.frames import (
|
||||||
|
Frame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
|
TTSStartedFrame,
|
||||||
|
TTSStoppedFrame)
|
||||||
from pipecat.services.ai_services import TTSService
|
from pipecat.services.ai_services import TTSService
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -27,8 +31,15 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
class PlayHTTTSService(TTSService):
|
class PlayHTTTSService(TTSService):
|
||||||
|
|
||||||
def __init__(self, *, api_key: str, user_id: str, voice_url: str, **kwargs):
|
def __init__(
|
||||||
super().__init__(**kwargs)
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str,
|
||||||
|
user_id: str,
|
||||||
|
voice_url: str,
|
||||||
|
sample_rate: int = 16000,
|
||||||
|
**kwargs):
|
||||||
|
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||||
|
|
||||||
self._user_id = user_id
|
self._user_id = user_id
|
||||||
self._speech_key = api_key
|
self._speech_key = api_key
|
||||||
@@ -39,13 +50,17 @@ class PlayHTTTSService(TTSService):
|
|||||||
)
|
)
|
||||||
self._options = TTSOptions(
|
self._options = TTSOptions(
|
||||||
voice=voice_url,
|
voice=voice_url,
|
||||||
sample_rate=16000,
|
sample_rate=sample_rate,
|
||||||
quality="higher",
|
quality="higher",
|
||||||
format=Format.FORMAT_WAV)
|
format=Format.FORMAT_WAV)
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def set_voice(self, voice: str):
|
||||||
|
logger.debug(f"Switching TTS voice to: [{voice}]")
|
||||||
|
self._options.voice = voice
|
||||||
|
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
logger.debug(f"Generating TTS: [{text}]")
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
|
|
||||||
@@ -80,7 +95,7 @@ class PlayHTTTSService(TTSService):
|
|||||||
else:
|
else:
|
||||||
if len(chunk):
|
if len(chunk):
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = AudioRawFrame(chunk, 16000, 1)
|
frame = TTSAudioRawFrame(chunk, 16000, 1)
|
||||||
yield frame
|
yield frame
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -4,23 +4,20 @@
|
|||||||
# SPDX-License-Identifier: BSD 2-Clause License
|
# SPDX-License-Identifier: BSD 2-Clause License
|
||||||
#
|
#
|
||||||
|
|
||||||
import base64
|
|
||||||
import json
|
import json
|
||||||
import io
|
|
||||||
import copy
|
|
||||||
from typing import List, Optional
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from asyncio import CancelledError
|
|
||||||
import re
|
import re
|
||||||
import uuid
|
import uuid
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from asyncio import CancelledError
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
Frame,
|
Frame,
|
||||||
LLMModelUpdateFrame,
|
LLMModelUpdateFrame,
|
||||||
TextFrame,
|
TextFrame,
|
||||||
VisionImageRawFrame,
|
|
||||||
UserImageRequestFrame,
|
UserImageRequestFrame,
|
||||||
UserImageRawFrame,
|
|
||||||
LLMMessagesFrame,
|
LLMMessagesFrame,
|
||||||
LLMFullResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
LLMFullResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
@@ -28,6 +25,7 @@ from pipecat.frames.frames import (
|
|||||||
FunctionCallInProgressFrame,
|
FunctionCallInProgressFrame,
|
||||||
StartInterruptionFrame
|
StartInterruptionFrame
|
||||||
)
|
)
|
||||||
|
from pipecat.metrics.metrics import LLMTokenUsage
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.ai_services import LLMService
|
from pipecat.services.ai_services import LLMService
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext, OpenAILLMContextFrame
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext, OpenAILLMContextFrame
|
||||||
@@ -59,18 +57,32 @@ class TogetherContextAggregatorPair:
|
|||||||
class TogetherLLMService(LLMService):
|
class TogetherLLMService(LLMService):
|
||||||
"""This class implements inference with Together's Llama 3.1 models
|
"""This class implements inference with Together's Llama 3.1 models
|
||||||
"""
|
"""
|
||||||
|
class InputParams(BaseModel):
|
||||||
|
frequency_penalty: Optional[float] = Field(default=None, ge=-2.0, le=2.0)
|
||||||
|
max_tokens: Optional[int] = Field(default=4096, ge=1)
|
||||||
|
presence_penalty: Optional[float] = Field(default=None, ge=-2.0, le=2.0)
|
||||||
|
temperature: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
||||||
|
top_k: Optional[int] = Field(default=None, ge=0)
|
||||||
|
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
model: str = "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo",
|
model: str = "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo",
|
||||||
max_tokens: int = 4096,
|
params: InputParams = InputParams(),
|
||||||
**kwargs):
|
**kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._client = AsyncTogether(api_key=api_key)
|
self._client = AsyncTogether(api_key=api_key)
|
||||||
self._model = model
|
self.set_model_name(model)
|
||||||
self._max_tokens = max_tokens
|
self._max_tokens = params.max_tokens
|
||||||
|
self._frequency_penalty = params.frequency_penalty
|
||||||
|
self._presence_penalty = params.presence_penalty
|
||||||
|
self._temperature = params.temperature
|
||||||
|
self._top_k = params.top_k
|
||||||
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -84,6 +96,34 @@ class TogetherLLMService(LLMService):
|
|||||||
_assistant=assistant
|
_assistant=assistant
|
||||||
)
|
)
|
||||||
|
|
||||||
|
async def set_frequency_penalty(self, frequency_penalty: float):
|
||||||
|
logger.debug(f"Switching LLM frequency_penalty to: [{frequency_penalty}]")
|
||||||
|
self._frequency_penalty = frequency_penalty
|
||||||
|
|
||||||
|
async def set_max_tokens(self, max_tokens: int):
|
||||||
|
logger.debug(f"Switching LLM max_tokens to: [{max_tokens}]")
|
||||||
|
self._max_tokens = max_tokens
|
||||||
|
|
||||||
|
async def set_presence_penalty(self, presence_penalty: float):
|
||||||
|
logger.debug(f"Switching LLM presence_penalty to: [{presence_penalty}]")
|
||||||
|
self._presence_penalty = presence_penalty
|
||||||
|
|
||||||
|
async def set_temperature(self, temperature: float):
|
||||||
|
logger.debug(f"Switching LLM temperature to: [{temperature}]")
|
||||||
|
self._temperature = temperature
|
||||||
|
|
||||||
|
async def set_top_k(self, top_k: float):
|
||||||
|
logger.debug(f"Switching LLM top_k to: [{top_k}]")
|
||||||
|
self._top_k = top_k
|
||||||
|
|
||||||
|
async def set_top_p(self, top_p: float):
|
||||||
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def _process_context(self, context: OpenAILLMContext):
|
async def _process_context(self, context: OpenAILLMContext):
|
||||||
try:
|
try:
|
||||||
await self.push_frame(LLMFullResponseStartFrame())
|
await self.push_frame(LLMFullResponseStartFrame())
|
||||||
@@ -93,12 +133,21 @@ class TogetherLLMService(LLMService):
|
|||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
stream = await self._client.chat.completions.create(
|
params = {
|
||||||
messages=context.messages,
|
"messages": context.messages,
|
||||||
model=self._model,
|
"model": self.model_name,
|
||||||
max_tokens=self._max_tokens,
|
"max_tokens": self._max_tokens,
|
||||||
stream=True,
|
"stream": True,
|
||||||
)
|
"frequency_penalty": self._frequency_penalty,
|
||||||
|
"presence_penalty": self._presence_penalty,
|
||||||
|
"temperature": self._temperature,
|
||||||
|
"top_k": self._top_k,
|
||||||
|
"top_p": self._top_p
|
||||||
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
stream = await self._client.chat.completions.create(**params)
|
||||||
|
|
||||||
# Function calling
|
# Function calling
|
||||||
got_first_chunk = False
|
got_first_chunk = False
|
||||||
@@ -108,13 +157,11 @@ class TogetherLLMService(LLMService):
|
|||||||
async for chunk in stream:
|
async for chunk in stream:
|
||||||
# logger.debug(f"Together LLM event: {chunk}")
|
# logger.debug(f"Together LLM event: {chunk}")
|
||||||
if chunk.usage:
|
if chunk.usage:
|
||||||
tokens = {
|
tokens = LLMTokenUsage(
|
||||||
"processor": self.name,
|
prompt_tokens=chunk.usage.prompt_tokens,
|
||||||
"model": self._model,
|
completion_tokens=chunk.usage.completion_tokens,
|
||||||
"prompt_tokens": chunk.usage.prompt_tokens,
|
total_tokens=chunk.usage.total_tokens
|
||||||
"completion_tokens": chunk.usage.completion_tokens,
|
)
|
||||||
"total_tokens": chunk.usage.total_tokens
|
|
||||||
}
|
|
||||||
await self.start_llm_usage_metrics(tokens)
|
await self.start_llm_usage_metrics(tokens)
|
||||||
|
|
||||||
if len(chunk.choices) == 0:
|
if len(chunk.choices) == 0:
|
||||||
@@ -156,7 +203,7 @@ class TogetherLLMService(LLMService):
|
|||||||
context = TogetherLLMContext.from_messages(frame.messages)
|
context = TogetherLLMContext.from_messages(frame.messages)
|
||||||
elif isinstance(frame, LLMModelUpdateFrame):
|
elif isinstance(frame, LLMModelUpdateFrame):
|
||||||
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
logger.debug(f"Switching LLM model to: [{frame.model}]")
|
||||||
self._model = frame.model
|
self.set_model_name(frame.model)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._device: str = device
|
self._device: str = device
|
||||||
self._compute_type = compute_type
|
self._compute_type = compute_type
|
||||||
self._model_name: str | Model = model
|
self.set_model_name(model if isinstance(model, str) else model.value)
|
||||||
self._no_speech_prob = no_speech_prob
|
self._no_speech_prob = no_speech_prob
|
||||||
self._model: WhisperModel | None = None
|
self._model: WhisperModel | None = None
|
||||||
self._load()
|
self._load()
|
||||||
@@ -65,7 +65,7 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
this model is being run, it will take time to download."""
|
this model is being run, it will take time to download."""
|
||||||
logger.debug("Loading Whisper model...")
|
logger.debug("Loading Whisper model...")
|
||||||
self._model = WhisperModel(
|
self._model = WhisperModel(
|
||||||
self._model_name.value if isinstance(self._model_name, Enum) else self._model_name,
|
self.model_name,
|
||||||
device=self._device,
|
device=self._device,
|
||||||
compute_type=self._compute_type)
|
compute_type=self._compute_type)
|
||||||
logger.debug("Loaded Whisper model")
|
logger.debug("Loaded Whisper model")
|
||||||
|
|||||||
@@ -9,10 +9,10 @@ import aiohttp
|
|||||||
from typing import Any, AsyncGenerator, Dict
|
from typing import Any, AsyncGenerator, Dict
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
ErrorFrame,
|
ErrorFrame,
|
||||||
Frame,
|
Frame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
|
TTSAudioRawFrame,
|
||||||
TTSStartedFrame,
|
TTSStartedFrame,
|
||||||
TTSStoppedFrame)
|
TTSStoppedFrame)
|
||||||
from pipecat.services.ai_services import TTSService
|
from pipecat.services.ai_services import TTSService
|
||||||
@@ -128,7 +128,7 @@ class XTTSService(TTSService):
|
|||||||
# Convert the numpy array back to bytes
|
# Convert the numpy array back to bytes
|
||||||
resampled_audio_bytes = resampled_audio.astype(np.int16).tobytes()
|
resampled_audio_bytes = resampled_audio.astype(np.int16).tobytes()
|
||||||
# Create the frame with the resampled audio
|
# Create the frame with the resampled audio
|
||||||
frame = AudioRawFrame(resampled_audio_bytes, 16000, 1)
|
frame = TTSAudioRawFrame(resampled_audio_bytes, 16000, 1)
|
||||||
yield frame
|
yield frame
|
||||||
|
|
||||||
# Process any remaining data in the buffer
|
# Process any remaining data in the buffer
|
||||||
@@ -136,7 +136,7 @@ class XTTSService(TTSService):
|
|||||||
audio_np = np.frombuffer(buffer, dtype=np.int16)
|
audio_np = np.frombuffer(buffer, dtype=np.int16)
|
||||||
resampled_audio = resampy.resample(audio_np, 24000, 16000)
|
resampled_audio = resampy.resample(audio_np, 24000, 16000)
|
||||||
resampled_audio_bytes = resampled_audio.astype(np.int16).tobytes()
|
resampled_audio_bytes = resampled_audio.astype(np.int16).tobytes()
|
||||||
frame = AudioRawFrame(resampled_audio_bytes, 16000, 1)
|
frame = TTSAudioRawFrame(resampled_audio_bytes, 16000, 1)
|
||||||
yield frame
|
yield frame
|
||||||
|
|
||||||
await self.push_frame(TTSStoppedFrame())
|
await self.push_frame(TTSStoppedFrame())
|
||||||
|
|||||||
@@ -10,9 +10,9 @@ from concurrent.futures import ThreadPoolExecutor
|
|||||||
|
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
BotInterruptionFrame,
|
BotInterruptionFrame,
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
|
InputAudioRawFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
@@ -31,16 +31,12 @@ from loguru import logger
|
|||||||
class BaseInputTransport(FrameProcessor):
|
class BaseInputTransport(FrameProcessor):
|
||||||
|
|
||||||
def __init__(self, params: TransportParams, **kwargs):
|
def __init__(self, params: TransportParams, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
self._executor = ThreadPoolExecutor(max_workers=5)
|
self._executor = ThreadPoolExecutor(max_workers=5)
|
||||||
|
|
||||||
# Create push frame task. This is the task that will push frames in
|
|
||||||
# order. We also guarantee that all frames are pushed in the same task.
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
# 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:
|
||||||
@@ -53,10 +49,6 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
self._audio_task.cancel()
|
self._audio_task.cancel()
|
||||||
await self._audio_task
|
await self._audio_task
|
||||||
|
|
||||||
# Wait for the push frame task to finish. It will finish when the
|
|
||||||
# EndFrame is actually processed.
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
# Cancel all the tasks and wait for them to finish.
|
# Cancel all the tasks and wait for them to finish.
|
||||||
|
|
||||||
@@ -64,13 +56,10 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
self._audio_task.cancel()
|
self._audio_task.cancel()
|
||||||
await self._audio_task
|
await self._audio_task
|
||||||
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
|
|
||||||
def vad_analyzer(self) -> VADAnalyzer | None:
|
def vad_analyzer(self) -> VADAnalyzer | None:
|
||||||
return self._params.vad_analyzer
|
return self._params.vad_analyzer
|
||||||
|
|
||||||
async def push_audio_frame(self, frame: AudioRawFrame):
|
async def push_audio_frame(self, frame: InputAudioRawFrame):
|
||||||
if self._params.audio_in_enabled or self._params.vad_enabled:
|
if self._params.audio_in_enabled or self._params.vad_enabled:
|
||||||
await self._audio_in_queue.put(frame)
|
await self._audio_in_queue.put(frame)
|
||||||
|
|
||||||
@@ -82,28 +71,25 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
await super().process_frame(frame, direction)
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
# Specific system frames
|
# Specific system frames
|
||||||
if isinstance(frame, CancelFrame):
|
if isinstance(frame, StartFrame):
|
||||||
|
# Push StartFrame before start(), because we want StartFrame to be
|
||||||
|
# processed by every processor before any other frame is processed.
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
await self.start(frame)
|
||||||
|
elif isinstance(frame, CancelFrame):
|
||||||
await self.cancel(frame)
|
await self.cancel(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, BotInterruptionFrame):
|
elif isinstance(frame, BotInterruptionFrame):
|
||||||
await self._handle_interruptions(frame, False)
|
logger.debug("Bot interruption")
|
||||||
elif isinstance(frame, StartInterruptionFrame):
|
|
||||||
await self._start_interruption()
|
await self._start_interruption()
|
||||||
elif isinstance(frame, StopInterruptionFrame):
|
|
||||||
await self._stop_interruption()
|
|
||||||
# All other system frames
|
# All other system frames
|
||||||
elif isinstance(frame, SystemFrame):
|
elif isinstance(frame, SystemFrame):
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
# Control frames
|
# Control frames
|
||||||
elif isinstance(frame, StartFrame):
|
|
||||||
# Push StartFrame before start(), because we want StartFrame to be
|
|
||||||
# processed by every processor before any other frame is processed.
|
|
||||||
await self._internal_push_frame(frame, direction)
|
|
||||||
await self.start(frame)
|
|
||||||
elif isinstance(frame, EndFrame):
|
elif isinstance(frame, EndFrame):
|
||||||
# Push EndFrame before stop(), because stop() waits on the task to
|
# Push EndFrame before stop(), because stop() waits on the task to
|
||||||
# finish and the task finishes when EndFrame is processed.
|
# finish and the task finishes when EndFrame is processed.
|
||||||
await self._internal_push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
await self.stop(frame)
|
await self.stop(frame)
|
||||||
elif isinstance(frame, VADParamsUpdateFrame):
|
elif isinstance(frame, VADParamsUpdateFrame):
|
||||||
vad_analyzer = self.vad_analyzer()
|
vad_analyzer = self.vad_analyzer()
|
||||||
@@ -111,73 +97,28 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
vad_analyzer.set_params(frame.params)
|
vad_analyzer.set_params(frame.params)
|
||||||
# Other frames
|
# Other frames
|
||||||
else:
|
else:
|
||||||
await self._internal_push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
|
|
||||||
#
|
|
||||||
# Push frames task
|
|
||||||
#
|
|
||||||
|
|
||||||
def _create_push_task(self):
|
|
||||||
loop = self.get_event_loop()
|
|
||||||
self._push_queue = asyncio.Queue()
|
|
||||||
self._push_frame_task = loop.create_task(self._push_frame_task_handler())
|
|
||||||
|
|
||||||
async def _internal_push_frame(
|
|
||||||
self,
|
|
||||||
frame: Frame | None,
|
|
||||||
direction: FrameDirection | None = FrameDirection.DOWNSTREAM):
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def _push_frame_task_handler(self):
|
|
||||||
running = True
|
|
||||||
while running:
|
|
||||||
try:
|
|
||||||
(frame, direction) = await self._push_queue.get()
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
running = not isinstance(frame, EndFrame)
|
|
||||||
self._push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Handle interruptions
|
# Handle interruptions
|
||||||
#
|
#
|
||||||
|
|
||||||
async def _start_interruption(self):
|
async def _handle_interruptions(self, frame: Frame):
|
||||||
if not self.interruptions_allowed:
|
|
||||||
return
|
|
||||||
|
|
||||||
# Cancel the task. This will stop pushing frames downstream.
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
# Push an out-of-band frame (i.e. not using the ordered push
|
|
||||||
# frame task) to stop everything, specially at the output
|
|
||||||
# transport.
|
|
||||||
await self.push_frame(StartInterruptionFrame())
|
|
||||||
# Create a new queue and task.
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
async def _stop_interruption(self):
|
|
||||||
if not self.interruptions_allowed:
|
|
||||||
return
|
|
||||||
|
|
||||||
await self.push_frame(StopInterruptionFrame())
|
|
||||||
|
|
||||||
async def _handle_interruptions(self, frame: Frame, push_frame: bool):
|
|
||||||
if self.interruptions_allowed:
|
if self.interruptions_allowed:
|
||||||
# Make sure we notify about interruptions quickly out-of-band
|
# Make sure we notify about interruptions quickly out-of-band.
|
||||||
if isinstance(frame, BotInterruptionFrame):
|
if isinstance(frame, UserStartedSpeakingFrame):
|
||||||
logger.debug("Bot interruption")
|
|
||||||
await self._start_interruption()
|
|
||||||
elif isinstance(frame, UserStartedSpeakingFrame):
|
|
||||||
logger.debug("User started speaking")
|
logger.debug("User started speaking")
|
||||||
await self._start_interruption()
|
await self._start_interruption()
|
||||||
|
# Push an out-of-band frame (i.e. not using the ordered push
|
||||||
|
# frame task) to stop everything, specially at the output
|
||||||
|
# transport.
|
||||||
|
await self.push_frame(StartInterruptionFrame())
|
||||||
elif isinstance(frame, UserStoppedSpeakingFrame):
|
elif isinstance(frame, UserStoppedSpeakingFrame):
|
||||||
logger.debug("User stopped speaking")
|
logger.debug("User stopped speaking")
|
||||||
await self._stop_interruption()
|
await self._stop_interruption()
|
||||||
|
await self.push_frame(StopInterruptionFrame())
|
||||||
|
|
||||||
if push_frame:
|
await self.push_frame(frame)
|
||||||
await self._internal_push_frame(frame)
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Audio input
|
# Audio input
|
||||||
@@ -201,7 +142,7 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
frame = UserStoppedSpeakingFrame()
|
frame = UserStoppedSpeakingFrame()
|
||||||
|
|
||||||
if frame:
|
if frame:
|
||||||
await self._handle_interruptions(frame, True)
|
await self._handle_interruptions(frame)
|
||||||
|
|
||||||
vad_state = new_vad_state
|
vad_state = new_vad_state
|
||||||
return vad_state
|
return vad_state
|
||||||
@@ -210,7 +151,7 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
vad_state: VADState = VADState.QUIET
|
vad_state: VADState = VADState.QUIET
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
frame: AudioRawFrame = await self._audio_in_queue.get()
|
frame: InputAudioRawFrame = await self._audio_in_queue.get()
|
||||||
|
|
||||||
audio_passthrough = True
|
audio_passthrough = True
|
||||||
|
|
||||||
@@ -222,7 +163,7 @@ class BaseInputTransport(FrameProcessor):
|
|||||||
|
|
||||||
# Push audio downstream if passthrough.
|
# Push audio downstream if passthrough.
|
||||||
if audio_passthrough:
|
if audio_passthrough:
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
self._audio_in_queue.task_done()
|
self._audio_in_queue.task_done()
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
|
|||||||
@@ -15,17 +15,17 @@ from typing import List
|
|||||||
|
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
BotSpeakingFrame,
|
BotSpeakingFrame,
|
||||||
BotStartedSpeakingFrame,
|
BotStartedSpeakingFrame,
|
||||||
BotStoppedSpeakingFrame,
|
BotStoppedSpeakingFrame,
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
MetricsFrame,
|
MetricsFrame,
|
||||||
|
OutputAudioRawFrame,
|
||||||
|
OutputImageRawFrame,
|
||||||
SpriteFrame,
|
SpriteFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
ImageRawFrame,
|
|
||||||
StartInterruptionFrame,
|
StartInterruptionFrame,
|
||||||
StopInterruptionFrame,
|
StopInterruptionFrame,
|
||||||
SystemFrame,
|
SystemFrame,
|
||||||
@@ -43,7 +43,7 @@ from pipecat.utils.time import nanoseconds_to_seconds
|
|||||||
class BaseOutputTransport(FrameProcessor):
|
class BaseOutputTransport(FrameProcessor):
|
||||||
|
|
||||||
def __init__(self, params: TransportParams, **kwargs):
|
def __init__(self, params: TransportParams, **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(sync=False, **kwargs)
|
||||||
|
|
||||||
self._params = params
|
self._params = params
|
||||||
|
|
||||||
@@ -70,10 +70,6 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
# generating frames upstream while, for example, the audio is playing.
|
# generating frames upstream while, for example, the audio is playing.
|
||||||
self._create_sink_tasks()
|
self._create_sink_tasks()
|
||||||
|
|
||||||
# Create push frame task. This is the task that will push frames in
|
|
||||||
# order. We also guarantee that all frames are pushed in the same task.
|
|
||||||
self._create_push_task()
|
|
||||||
|
|
||||||
async def start(self, frame: StartFrame):
|
async def start(self, frame: StartFrame):
|
||||||
# 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:
|
||||||
@@ -85,6 +81,13 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
self._audio_out_task = self.get_event_loop().create_task(self._audio_out_task_handler())
|
self._audio_out_task = self.get_event_loop().create_task(self._audio_out_task_handler())
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
|
# At this point we have enqueued an EndFrame and we need to wait for
|
||||||
|
# that EndFrame to be processed by the sink tasks. We also need to wait
|
||||||
|
# for these tasks before cancelling the camera and audio tasks below
|
||||||
|
# because they might be still rendering.
|
||||||
|
await self._sink_task
|
||||||
|
await self._sink_clock_task
|
||||||
|
|
||||||
# Cancel and wait for the camera output task to finish.
|
# Cancel and wait for the camera output task to finish.
|
||||||
if self._params.camera_out_enabled:
|
if self._params.camera_out_enabled:
|
||||||
self._camera_out_task.cancel()
|
self._camera_out_task.cancel()
|
||||||
@@ -95,23 +98,23 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
self._audio_out_task.cancel()
|
self._audio_out_task.cancel()
|
||||||
await self._audio_out_task
|
await self._audio_out_task
|
||||||
|
|
||||||
# Wait for the push frame and sink tasks to finish. They will finish when
|
|
||||||
# the EndFrame is actually processed.
|
|
||||||
await self._push_frame_task
|
|
||||||
await self._sink_task
|
|
||||||
|
|
||||||
async def cancel(self, frame: CancelFrame):
|
async def cancel(self, frame: CancelFrame):
|
||||||
# Cancel all the tasks and wait for them to finish.
|
# Since we are cancelling everything it doesn't matter if we cancel sink
|
||||||
|
# tasks first or not.
|
||||||
|
self._sink_task.cancel()
|
||||||
|
self._sink_clock_task.cancel()
|
||||||
|
await self._sink_task
|
||||||
|
await self._sink_clock_task
|
||||||
|
|
||||||
|
# Cancel and wait for the camera output task to finish.
|
||||||
if self._params.camera_out_enabled:
|
if self._params.camera_out_enabled:
|
||||||
self._camera_out_task.cancel()
|
self._camera_out_task.cancel()
|
||||||
await self._camera_out_task
|
await self._camera_out_task
|
||||||
|
|
||||||
self._push_frame_task.cancel()
|
# Cancel and wait for the audio output task to finish.
|
||||||
await self._push_frame_task
|
if self._params.audio_out_enabled and self._params.audio_out_is_live:
|
||||||
|
self._audio_out_task.cancel()
|
||||||
self._sink_task.cancel()
|
await self._audio_out_task
|
||||||
await self._sink_task
|
|
||||||
|
|
||||||
async def send_message(self, frame: TransportMessageFrame):
|
async def send_message(self, frame: TransportMessageFrame):
|
||||||
pass
|
pass
|
||||||
@@ -119,7 +122,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
async def send_metrics(self, frame: MetricsFrame):
|
async def send_metrics(self, frame: MetricsFrame):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
async def write_frame_to_camera(self, frame: ImageRawFrame):
|
async def write_frame_to_camera(self, frame: OutputImageRawFrame):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
async def write_raw_audio_frames(self, frames: bytes):
|
async def write_raw_audio_frames(self, frames: bytes):
|
||||||
@@ -137,7 +140,12 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
# immediately. Other frames require order so they are put in the sink
|
# immediately. Other frames require order so they are put in the sink
|
||||||
# queue.
|
# queue.
|
||||||
#
|
#
|
||||||
if isinstance(frame, CancelFrame):
|
if isinstance(frame, StartFrame):
|
||||||
|
# Push StartFrame before start(), because we want StartFrame to be
|
||||||
|
# processed by every processor before any other frame is processed.
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
await self.start(frame)
|
||||||
|
elif isinstance(frame, CancelFrame):
|
||||||
await self.cancel(frame)
|
await self.cancel(frame)
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
elif isinstance(frame, StartInterruptionFrame) or isinstance(frame, StopInterruptionFrame):
|
elif isinstance(frame, StartInterruptionFrame) or isinstance(frame, StopInterruptionFrame):
|
||||||
@@ -149,17 +157,14 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
elif isinstance(frame, SystemFrame):
|
elif isinstance(frame, SystemFrame):
|
||||||
await self.push_frame(frame, direction)
|
await self.push_frame(frame, direction)
|
||||||
# Control frames.
|
# Control frames.
|
||||||
elif isinstance(frame, StartFrame):
|
|
||||||
await self._sink_queue.put(frame)
|
|
||||||
await self.start(frame)
|
|
||||||
elif isinstance(frame, EndFrame):
|
elif isinstance(frame, EndFrame):
|
||||||
await self._sink_clock_queue.put((sys.maxsize, frame.id, frame))
|
await self._sink_clock_queue.put((sys.maxsize, frame.id, frame))
|
||||||
await self._sink_queue.put(frame)
|
await self._sink_queue.put(frame)
|
||||||
await self.stop(frame)
|
await self.stop(frame)
|
||||||
# Other frames.
|
# Other frames.
|
||||||
elif isinstance(frame, AudioRawFrame):
|
elif isinstance(frame, OutputAudioRawFrame):
|
||||||
await self._handle_audio(frame)
|
await self._handle_audio(frame)
|
||||||
elif isinstance(frame, ImageRawFrame) or isinstance(frame, SpriteFrame):
|
elif isinstance(frame, OutputImageRawFrame) or isinstance(frame, SpriteFrame):
|
||||||
await self._handle_image(frame)
|
await self._handle_image(frame)
|
||||||
elif isinstance(frame, TransportMessageFrame) and frame.urgent:
|
elif isinstance(frame, TransportMessageFrame) and frame.urgent:
|
||||||
await self.send_message(frame)
|
await self.send_message(frame)
|
||||||
@@ -182,15 +187,11 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
await self._sink_clock_task
|
await self._sink_clock_task
|
||||||
# Create sink tasks.
|
# Create sink tasks.
|
||||||
self._create_sink_tasks()
|
self._create_sink_tasks()
|
||||||
# Stop push task.
|
|
||||||
self._push_frame_task.cancel()
|
|
||||||
await self._push_frame_task
|
|
||||||
self._create_push_task()
|
|
||||||
# Let's send a bot stopped speaking if we have to.
|
# Let's send a bot stopped speaking if we have to.
|
||||||
if self._bot_speaking:
|
if self._bot_speaking:
|
||||||
await self._bot_stopped_speaking()
|
await self._bot_stopped_speaking()
|
||||||
|
|
||||||
async def _handle_audio(self, frame: AudioRawFrame):
|
async def _handle_audio(self, frame: OutputAudioRawFrame):
|
||||||
if not self._params.audio_out_enabled:
|
if not self._params.audio_out_enabled:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -199,12 +200,14 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
else:
|
else:
|
||||||
self._audio_buffer.extend(frame.audio)
|
self._audio_buffer.extend(frame.audio)
|
||||||
while len(self._audio_buffer) >= self._audio_chunk_size:
|
while len(self._audio_buffer) >= self._audio_chunk_size:
|
||||||
chunk = AudioRawFrame(bytes(self._audio_buffer[:self._audio_chunk_size]),
|
chunk = OutputAudioRawFrame(
|
||||||
sample_rate=frame.sample_rate, num_channels=frame.num_channels)
|
bytes(self._audio_buffer[:self._audio_chunk_size]),
|
||||||
|
sample_rate=frame.sample_rate, num_channels=frame.num_channels
|
||||||
|
)
|
||||||
await self._sink_queue.put(chunk)
|
await self._sink_queue.put(chunk)
|
||||||
self._audio_buffer = self._audio_buffer[self._audio_chunk_size:]
|
self._audio_buffer = self._audio_buffer[self._audio_chunk_size:]
|
||||||
|
|
||||||
async def _handle_image(self, frame: ImageRawFrame | SpriteFrame):
|
async def _handle_image(self, frame: OutputImageRawFrame | SpriteFrame):
|
||||||
if not self._params.camera_out_enabled:
|
if not self._params.camera_out_enabled:
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -225,11 +228,11 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
self._sink_clock_task = loop.create_task(self._sink_clock_task_handler())
|
self._sink_clock_task = loop.create_task(self._sink_clock_task_handler())
|
||||||
|
|
||||||
async def _sink_frame_handler(self, frame: Frame):
|
async def _sink_frame_handler(self, frame: Frame):
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, OutputAudioRawFrame):
|
||||||
await self.write_raw_audio_frames(frame.audio)
|
await self.write_raw_audio_frames(frame.audio)
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
||||||
elif isinstance(frame, ImageRawFrame):
|
elif isinstance(frame, OutputImageRawFrame):
|
||||||
await self._set_camera_image(frame)
|
await self._set_camera_image(frame)
|
||||||
elif isinstance(frame, SpriteFrame):
|
elif isinstance(frame, SpriteFrame):
|
||||||
await self._set_camera_images(frame.images)
|
await self._set_camera_images(frame.images)
|
||||||
@@ -237,12 +240,12 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
await self.send_message(frame)
|
await self.send_message(frame)
|
||||||
elif isinstance(frame, TTSStartedFrame):
|
elif isinstance(frame, TTSStartedFrame):
|
||||||
await self._bot_started_speaking()
|
await self._bot_started_speaking()
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
elif isinstance(frame, TTSStoppedFrame):
|
elif isinstance(frame, TTSStoppedFrame):
|
||||||
await self._bot_stopped_speaking()
|
await self._bot_stopped_speaking()
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
else:
|
else:
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def _sink_task_handler(self):
|
async def _sink_task_handler(self):
|
||||||
running = True
|
running = True
|
||||||
@@ -261,7 +264,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
# TODO(aleix): For now we just process TextFrame. But we should process
|
# TODO(aleix): For now we just process TextFrame. But we should process
|
||||||
# audio and video as well.
|
# audio and video as well.
|
||||||
if isinstance(frame, TextFrame):
|
if isinstance(frame, TextFrame):
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def _sink_clock_task_handler(self):
|
async def _sink_clock_task_handler(self):
|
||||||
running = True
|
running = True
|
||||||
@@ -269,7 +272,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
try:
|
try:
|
||||||
timestamp, _, frame = await self._sink_clock_queue.get()
|
timestamp, _, frame = await self._sink_clock_queue.get()
|
||||||
|
|
||||||
# If we hit an EndFrame, we cna finish right away.
|
# If we hit an EndFrame, we can finish right away.
|
||||||
running = not isinstance(frame, EndFrame)
|
running = not isinstance(frame, EndFrame)
|
||||||
|
|
||||||
# If we have a frame we check it's presentation timestamp. If it
|
# If we have a frame we check it's presentation timestamp. If it
|
||||||
@@ -293,47 +296,21 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
async def _bot_started_speaking(self):
|
async def _bot_started_speaking(self):
|
||||||
logger.debug("Bot started speaking")
|
logger.debug("Bot started speaking")
|
||||||
self._bot_speaking = True
|
self._bot_speaking = True
|
||||||
await self._internal_push_frame(BotStartedSpeakingFrame(), FrameDirection.UPSTREAM)
|
await self.push_frame(BotStartedSpeakingFrame(), FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
async def _bot_stopped_speaking(self):
|
async def _bot_stopped_speaking(self):
|
||||||
logger.debug("Bot stopped speaking")
|
logger.debug("Bot stopped speaking")
|
||||||
self._bot_speaking = False
|
self._bot_speaking = False
|
||||||
await self._internal_push_frame(BotStoppedSpeakingFrame(), FrameDirection.UPSTREAM)
|
await self.push_frame(BotStoppedSpeakingFrame(), FrameDirection.UPSTREAM)
|
||||||
|
|
||||||
#
|
|
||||||
# Push frames task
|
|
||||||
#
|
|
||||||
|
|
||||||
def _create_push_task(self):
|
|
||||||
loop = self.get_event_loop()
|
|
||||||
self._push_queue = asyncio.Queue()
|
|
||||||
self._push_frame_task = loop.create_task(self._push_frame_task_handler())
|
|
||||||
|
|
||||||
async def _internal_push_frame(
|
|
||||||
self,
|
|
||||||
frame: Frame | None,
|
|
||||||
direction: FrameDirection | None = FrameDirection.DOWNSTREAM):
|
|
||||||
await self._push_queue.put((frame, direction))
|
|
||||||
|
|
||||||
async def _push_frame_task_handler(self):
|
|
||||||
running = True
|
|
||||||
while running:
|
|
||||||
try:
|
|
||||||
(frame, direction) = await self._push_queue.get()
|
|
||||||
await self.push_frame(frame, direction)
|
|
||||||
running = not isinstance(frame, EndFrame)
|
|
||||||
self._push_queue.task_done()
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
break
|
|
||||||
|
|
||||||
#
|
#
|
||||||
# Camera out
|
# Camera out
|
||||||
#
|
#
|
||||||
|
|
||||||
async def send_image(self, frame: ImageRawFrame | SpriteFrame):
|
async def send_image(self, frame: OutputImageRawFrame | SpriteFrame):
|
||||||
await self.process_frame(frame, FrameDirection.DOWNSTREAM)
|
await self.process_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
async def _draw_image(self, frame: ImageRawFrame):
|
async def _draw_image(self, frame: OutputImageRawFrame):
|
||||||
desired_size = (self._params.camera_out_width, self._params.camera_out_height)
|
desired_size = (self._params.camera_out_width, self._params.camera_out_height)
|
||||||
|
|
||||||
if frame.size != desired_size:
|
if frame.size != desired_size:
|
||||||
@@ -341,14 +318,17 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
resized_image = image.resize(desired_size)
|
resized_image = image.resize(desired_size)
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"{frame} does not have the expected size {desired_size}, resizing")
|
f"{frame} does not have the expected size {desired_size}, resizing")
|
||||||
frame = ImageRawFrame(resized_image.tobytes(), resized_image.size, resized_image.format)
|
frame = OutputImageRawFrame(
|
||||||
|
resized_image.tobytes(),
|
||||||
|
resized_image.size,
|
||||||
|
resized_image.format)
|
||||||
|
|
||||||
await self.write_frame_to_camera(frame)
|
await self.write_frame_to_camera(frame)
|
||||||
|
|
||||||
async def _set_camera_image(self, image: ImageRawFrame):
|
async def _set_camera_image(self, image: OutputImageRawFrame):
|
||||||
self._camera_images = itertools.cycle([image])
|
self._camera_images = itertools.cycle([image])
|
||||||
|
|
||||||
async def _set_camera_images(self, images: List[ImageRawFrame]):
|
async def _set_camera_images(self, images: List[OutputImageRawFrame]):
|
||||||
self._camera_images = itertools.cycle(images)
|
self._camera_images = itertools.cycle(images)
|
||||||
|
|
||||||
async def _camera_out_task_handler(self):
|
async def _camera_out_task_handler(self):
|
||||||
@@ -363,9 +343,9 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
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(1.0 / self._params.camera_out_framerate)
|
await asyncio.sleep(self._camera_out_frame_duration)
|
||||||
else:
|
else:
|
||||||
await asyncio.sleep(1.0 / self._params.camera_out_framerate)
|
await asyncio.sleep(self._camera_out_frame_duration)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
break
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -400,7 +380,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
# Audio out
|
# Audio out
|
||||||
#
|
#
|
||||||
|
|
||||||
async def send_audio(self, frame: AudioRawFrame):
|
async def send_audio(self, frame: OutputAudioRawFrame):
|
||||||
await self.process_frame(frame, FrameDirection.DOWNSTREAM)
|
await self.process_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
async def _audio_out_task_handler(self):
|
async def _audio_out_task_handler(self):
|
||||||
@@ -408,7 +388,7 @@ class BaseOutputTransport(FrameProcessor):
|
|||||||
try:
|
try:
|
||||||
frame = await self._audio_out_queue.get()
|
frame = await self._audio_out_queue.get()
|
||||||
await self.write_raw_audio_frames(frame.audio)
|
await self.write_raw_audio_frames(frame.audio)
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
await self.push_frame(BotSpeakingFrame(), FrameDirection.UPSTREAM)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
break
|
break
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import asyncio
|
|||||||
|
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, StartFrame
|
from pipecat.frames.frames import InputAudioRawFrame, StartFrame
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
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
|
||||||
@@ -54,9 +54,9 @@ class LocalAudioInputTransport(BaseInputTransport):
|
|||||||
self._in_stream.close()
|
self._in_stream.close()
|
||||||
|
|
||||||
def _audio_in_callback(self, in_data, frame_count, time_info, status):
|
def _audio_in_callback(self, in_data, frame_count, time_info, status):
|
||||||
frame = AudioRawFrame(audio=in_data,
|
frame = InputAudioRawFrame(audio=in_data,
|
||||||
sample_rate=self._params.audio_in_sample_rate,
|
sample_rate=self._params.audio_in_sample_rate,
|
||||||
num_channels=self._params.audio_in_channels)
|
num_channels=self._params.audio_in_channels)
|
||||||
|
|
||||||
asyncio.run_coroutine_threadsafe(self.push_audio_frame(frame), self.get_event_loop())
|
asyncio.run_coroutine_threadsafe(self.push_audio_frame(frame), self.get_event_loop())
|
||||||
|
|
||||||
|
|||||||
@@ -11,8 +11,7 @@ from concurrent.futures import ThreadPoolExecutor
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import tkinter as tk
|
import tkinter as tk
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, ImageRawFrame, StartFrame
|
from pipecat.frames.frames import InputAudioRawFrame, OutputImageRawFrame, StartFrame
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
|
||||||
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
|
||||||
@@ -64,9 +63,9 @@ class TkInputTransport(BaseInputTransport):
|
|||||||
self._in_stream.close()
|
self._in_stream.close()
|
||||||
|
|
||||||
def _audio_in_callback(self, in_data, frame_count, time_info, status):
|
def _audio_in_callback(self, in_data, frame_count, time_info, status):
|
||||||
frame = AudioRawFrame(audio=in_data,
|
frame = InputAudioRawFrame(audio=in_data,
|
||||||
sample_rate=self._params.audio_in_sample_rate,
|
sample_rate=self._params.audio_in_sample_rate,
|
||||||
num_channels=self._params.audio_in_channels)
|
num_channels=self._params.audio_in_channels)
|
||||||
|
|
||||||
asyncio.run_coroutine_threadsafe(self.push_audio_frame(frame), self.get_event_loop())
|
asyncio.run_coroutine_threadsafe(self.push_audio_frame(frame), self.get_event_loop())
|
||||||
|
|
||||||
@@ -108,10 +107,10 @@ class TkOutputTransport(BaseOutputTransport):
|
|||||||
async def write_raw_audio_frames(self, frames: bytes):
|
async def write_raw_audio_frames(self, frames: bytes):
|
||||||
await self.get_event_loop().run_in_executor(self._executor, self._out_stream.write, frames)
|
await self.get_event_loop().run_in_executor(self._executor, self._out_stream.write, frames)
|
||||||
|
|
||||||
async def write_frame_to_camera(self, frame: ImageRawFrame):
|
async def write_frame_to_camera(self, frame: OutputImageRawFrame):
|
||||||
self.get_event_loop().call_soon(self._write_frame_to_tk, frame)
|
self.get_event_loop().call_soon(self._write_frame_to_tk, frame)
|
||||||
|
|
||||||
def _write_frame_to_tk(self, frame: ImageRawFrame):
|
def _write_frame_to_tk(self, frame: OutputImageRawFrame):
|
||||||
width = frame.size[0]
|
width = frame.size[0]
|
||||||
height = frame.size[1]
|
height = frame.size[1]
|
||||||
data = f"P6 {width} {height} 255 ".encode() + frame.image
|
data = f"P6 {width} {height} 255 ".encode() + frame.image
|
||||||
@@ -141,12 +140,12 @@ class TkLocalTransport(BaseTransport):
|
|||||||
# BaseTransport
|
# BaseTransport
|
||||||
#
|
#
|
||||||
|
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> TkInputTransport:
|
||||||
if not self._input:
|
if not self._input:
|
||||||
self._input = TkInputTransport(self._pyaudio, self._params)
|
self._input = TkInputTransport(self._pyaudio, self._params)
|
||||||
return self._input
|
return self._input
|
||||||
|
|
||||||
def output(self) -> FrameProcessor:
|
def output(self) -> TkOutputTransport:
|
||||||
if not self._output:
|
if not self._output:
|
||||||
self._output = TkOutputTransport(self._tk_root, self._pyaudio, self._params)
|
self._output = TkOutputTransport(self._tk_root, self._pyaudio, self._params)
|
||||||
return self._output
|
return self._output
|
||||||
|
|||||||
@@ -12,8 +12,16 @@ import wave
|
|||||||
from typing import Awaitable, Callable
|
from typing import Awaitable, Callable
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, CancelFrame, EndFrame, Frame, StartFrame, StartInterruptionFrame
|
from pipecat.frames.frames import (
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
AudioRawFrame,
|
||||||
|
CancelFrame,
|
||||||
|
EndFrame,
|
||||||
|
Frame,
|
||||||
|
InputAudioRawFrame,
|
||||||
|
StartFrame,
|
||||||
|
StartInterruptionFrame
|
||||||
|
)
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import 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
|
||||||
@@ -79,7 +87,11 @@ class FastAPIWebsocketInputTransport(BaseInputTransport):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, AudioRawFrame):
|
||||||
await self.push_audio_frame(frame)
|
await self.push_audio_frame(InputAudioRawFrame(
|
||||||
|
audio=frame.audio,
|
||||||
|
sample_rate=frame.sample_rate,
|
||||||
|
num_channels=frame.num_channels)
|
||||||
|
)
|
||||||
|
|
||||||
await self._callbacks.on_client_disconnected(self._websocket)
|
await self._callbacks.on_client_disconnected(self._websocket)
|
||||||
|
|
||||||
@@ -164,10 +176,10 @@ class FastAPIWebsocketTransport(BaseTransport):
|
|||||||
self._register_event_handler("on_client_connected")
|
self._register_event_handler("on_client_connected")
|
||||||
self._register_event_handler("on_client_disconnected")
|
self._register_event_handler("on_client_disconnected")
|
||||||
|
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> FastAPIWebsocketInputTransport:
|
||||||
return self._input
|
return self._input
|
||||||
|
|
||||||
def output(self) -> FrameProcessor:
|
def output(self) -> FastAPIWebsocketOutputTransport:
|
||||||
return self._output
|
return self._output
|
||||||
|
|
||||||
async def _on_client_connected(self, websocket):
|
async def _on_client_connected(self, websocket):
|
||||||
|
|||||||
@@ -11,8 +11,7 @@ import wave
|
|||||||
from typing import Awaitable, Callable
|
from typing import Awaitable, Callable
|
||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import AudioRawFrame, CancelFrame, EndFrame, StartFrame
|
from pipecat.frames.frames import AudioRawFrame, CancelFrame, EndFrame, InputAudioRawFrame, StartFrame
|
||||||
from pipecat.processors.frame_processor import FrameProcessor
|
|
||||||
from pipecat.serializers.base_serializer import FrameSerializer
|
from pipecat.serializers.base_serializer import FrameSerializer
|
||||||
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
@@ -98,9 +97,13 @@ class WebsocketServerInputTransport(BaseInputTransport):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if isinstance(frame, AudioRawFrame):
|
if isinstance(frame, AudioRawFrame):
|
||||||
await self.push_audio_frame(frame)
|
await self.push_audio_frame(InputAudioRawFrame(
|
||||||
|
audio=frame.audio,
|
||||||
|
sample_rate=frame.sample_rate,
|
||||||
|
num_channels=frame.num_channels)
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
# Notify disconnection
|
# Notify disconnection
|
||||||
await self._callbacks.on_client_disconnected(websocket)
|
await self._callbacks.on_client_disconnected(websocket)
|
||||||
@@ -190,13 +193,13 @@ class WebsocketServerTransport(BaseTransport):
|
|||||||
self._register_event_handler("on_client_connected")
|
self._register_event_handler("on_client_connected")
|
||||||
self._register_event_handler("on_client_disconnected")
|
self._register_event_handler("on_client_disconnected")
|
||||||
|
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> WebsocketServerInputTransport:
|
||||||
if not self._input:
|
if not self._input:
|
||||||
self._input = WebsocketServerInputTransport(
|
self._input = WebsocketServerInputTransport(
|
||||||
self._host, self._port, self._params, self._callbacks, name=self._input_name)
|
self._host, self._port, self._params, self._callbacks, name=self._input_name)
|
||||||
return self._input
|
return self._input
|
||||||
|
|
||||||
def output(self) -> FrameProcessor:
|
def output(self) -> WebsocketServerOutputTransport:
|
||||||
if not self._output:
|
if not self._output:
|
||||||
self._output = WebsocketServerOutputTransport(self._params, name=self._output_name)
|
self._output = WebsocketServerOutputTransport(self._params, name=self._output_name)
|
||||||
return self._output
|
return self._output
|
||||||
|
|||||||
@@ -22,19 +22,21 @@ from daily import (
|
|||||||
from pydantic.main import BaseModel
|
from pydantic.main import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
|
||||||
CancelFrame,
|
CancelFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
Frame,
|
Frame,
|
||||||
ImageRawFrame,
|
InputAudioRawFrame,
|
||||||
InterimTranscriptionFrame,
|
InterimTranscriptionFrame,
|
||||||
MetricsFrame,
|
MetricsFrame,
|
||||||
|
OutputAudioRawFrame,
|
||||||
|
OutputImageRawFrame,
|
||||||
SpriteFrame,
|
SpriteFrame,
|
||||||
StartFrame,
|
StartFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
TransportMessageFrame,
|
TransportMessageFrame,
|
||||||
UserImageRawFrame,
|
UserImageRawFrame,
|
||||||
UserImageRequestFrame)
|
UserImageRequestFrame)
|
||||||
|
from pipecat.metrics.metrics import LLMUsageMetricsData, ProcessingMetricsData, TTFBMetricsData, TTSUsageMetricsData
|
||||||
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.transports.base_input import BaseInputTransport
|
from pipecat.transports.base_input import BaseInputTransport
|
||||||
@@ -239,7 +241,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
completion=completion_callback(future))
|
completion=completion_callback(future))
|
||||||
await future
|
await future
|
||||||
|
|
||||||
async def read_next_audio_frame(self) -> AudioRawFrame | None:
|
async def read_next_audio_frame(self) -> InputAudioRawFrame | None:
|
||||||
if not self._speaker:
|
if not self._speaker:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -252,7 +254,10 @@ class DailyTransportClient(EventHandler):
|
|||||||
audio = await future
|
audio = await future
|
||||||
|
|
||||||
if len(audio) > 0:
|
if len(audio) > 0:
|
||||||
return AudioRawFrame(audio=audio, sample_rate=sample_rate, num_channels=num_channels)
|
return InputAudioRawFrame(
|
||||||
|
audio=audio,
|
||||||
|
sample_rate=sample_rate,
|
||||||
|
num_channels=num_channels)
|
||||||
else:
|
else:
|
||||||
# If we don't read any audio it could be there's no participant
|
# If we don't read any audio it could be there's no participant
|
||||||
# connected. daily-python will return immediately if that's the
|
# connected. daily-python will return immediately if that's the
|
||||||
@@ -268,7 +273,7 @@ class DailyTransportClient(EventHandler):
|
|||||||
self._mic.write_frames(frames, completion=completion_callback(future))
|
self._mic.write_frames(frames, completion=completion_callback(future))
|
||||||
await future
|
await future
|
||||||
|
|
||||||
async def write_frame_to_camera(self, frame: ImageRawFrame):
|
async def write_frame_to_camera(self, frame: OutputImageRawFrame):
|
||||||
if not self._camera:
|
if not self._camera:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -625,11 +630,11 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
#
|
#
|
||||||
|
|
||||||
async def push_transcription_frame(self, frame: TranscriptionFrame | InterimTranscriptionFrame):
|
async def push_transcription_frame(self, frame: TranscriptionFrame | InterimTranscriptionFrame):
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
async def push_app_message(self, message: Any, sender: str):
|
async def push_app_message(self, message: Any, sender: str):
|
||||||
frame = DailyTransportMessageFrame(message=message, participant_id=sender)
|
frame = DailyTransportMessageFrame(message=message, participant_id=sender)
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
#
|
#
|
||||||
# Audio in
|
# Audio in
|
||||||
@@ -692,7 +697,7 @@ class DailyInputTransport(BaseInputTransport):
|
|||||||
image=buffer,
|
image=buffer,
|
||||||
size=size,
|
size=size,
|
||||||
format=format)
|
format=format)
|
||||||
await self._internal_push_frame(frame)
|
await self.push_frame(frame)
|
||||||
|
|
||||||
self._video_renderers[participant_id]["timestamp"] = curr_time
|
self._video_renderers[participant_id]["timestamp"] = curr_time
|
||||||
|
|
||||||
@@ -731,14 +736,23 @@ class DailyOutputTransport(BaseOutputTransport):
|
|||||||
|
|
||||||
async def send_metrics(self, frame: MetricsFrame):
|
async def send_metrics(self, frame: MetricsFrame):
|
||||||
metrics = {}
|
metrics = {}
|
||||||
if frame.ttfb:
|
for d in frame.data:
|
||||||
metrics["ttfb"] = frame.ttfb
|
if isinstance(d, TTFBMetricsData):
|
||||||
if frame.processing:
|
if "ttfb" not in metrics:
|
||||||
metrics["processing"] = frame.processing
|
metrics["ttfb"] = []
|
||||||
if frame.tokens:
|
metrics["ttfb"].append(d.model_dump(exclude_none=True))
|
||||||
metrics["tokens"] = frame.tokens
|
elif isinstance(d, ProcessingMetricsData):
|
||||||
if frame.characters:
|
if "processing" not in metrics:
|
||||||
metrics["characters"] = frame.characters
|
metrics["processing"] = []
|
||||||
|
metrics["processing"].append(d.model_dump(exclude_none=True))
|
||||||
|
elif isinstance(d, LLMUsageMetricsData):
|
||||||
|
if "tokens" not in metrics:
|
||||||
|
metrics["tokens"] = []
|
||||||
|
metrics["tokens"].append(d.value.model_dump(exclude_none=True))
|
||||||
|
elif isinstance(d, TTSUsageMetricsData):
|
||||||
|
if "characters" not in metrics:
|
||||||
|
metrics["characters"] = []
|
||||||
|
metrics["characters"].append(d.model_dump(exclude_none=True))
|
||||||
|
|
||||||
message = DailyTransportMessageFrame(message={
|
message = DailyTransportMessageFrame(message={
|
||||||
"type": "pipecat-metrics",
|
"type": "pipecat-metrics",
|
||||||
@@ -749,7 +763,7 @@ class DailyOutputTransport(BaseOutputTransport):
|
|||||||
async def write_raw_audio_frames(self, frames: bytes):
|
async def write_raw_audio_frames(self, frames: bytes):
|
||||||
await self._client.write_raw_audio_frames(frames)
|
await self._client.write_raw_audio_frames(frames)
|
||||||
|
|
||||||
async def write_frame_to_camera(self, frame: ImageRawFrame):
|
async def write_frame_to_camera(self, frame: OutputImageRawFrame):
|
||||||
await self._client.write_frame_to_camera(frame)
|
await self._client.write_frame_to_camera(frame)
|
||||||
|
|
||||||
|
|
||||||
@@ -811,12 +825,12 @@ class DailyTransport(BaseTransport):
|
|||||||
# BaseTransport
|
# BaseTransport
|
||||||
#
|
#
|
||||||
|
|
||||||
def input(self) -> FrameProcessor:
|
def input(self) -> DailyInputTransport:
|
||||||
if not self._input:
|
if not self._input:
|
||||||
self._input = DailyInputTransport(self._client, self._params, name=self._input_name)
|
self._input = DailyInputTransport(self._client, self._params, name=self._input_name)
|
||||||
return self._input
|
return self._input
|
||||||
|
|
||||||
def output(self) -> FrameProcessor:
|
def output(self) -> DailyOutputTransport:
|
||||||
if not self._output:
|
if not self._output:
|
||||||
self._output = DailyOutputTransport(self._client, self._params, name=self._output_name)
|
self._output = DailyOutputTransport(self._client, self._params, name=self._output_name)
|
||||||
return self._output
|
return self._output
|
||||||
@@ -829,11 +843,11 @@ class DailyTransport(BaseTransport):
|
|||||||
def participant_id(self) -> str:
|
def participant_id(self) -> str:
|
||||||
return self._client.participant_id
|
return self._client.participant_id
|
||||||
|
|
||||||
async def send_image(self, frame: ImageRawFrame | SpriteFrame):
|
async def send_image(self, frame: OutputImageRawFrame | SpriteFrame):
|
||||||
if self._output:
|
if self._output:
|
||||||
await self._output.process_frame(frame, FrameDirection.DOWNSTREAM)
|
await self._output.process_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
async def send_audio(self, frame: AudioRawFrame):
|
async def send_audio(self, frame: OutputAudioRawFrame):
|
||||||
if self._output:
|
if self._output:
|
||||||
await self._output.process_frame(frame, FrameDirection.DOWNSTREAM)
|
await self._output.process_frame(frame, FrameDirection.DOWNSTREAM)
|
||||||
|
|
||||||
|
|||||||
35
test-requirements.txt
Normal file
35
test-requirements.txt
Normal file
@@ -0,0 +1,35 @@
|
|||||||
|
aiohttp~=3.10.3
|
||||||
|
anthropic
|
||||||
|
autopep8~=2.3.1
|
||||||
|
azure-cognitiveservices-speech~=1.40.0
|
||||||
|
build~=1.2.1
|
||||||
|
daily-python~=0.10.1
|
||||||
|
deepgram-sdk~=3.5.0
|
||||||
|
fal-client~=0.4.1
|
||||||
|
fastapi~=0.112.1
|
||||||
|
faster-whisper~=1.0.3
|
||||||
|
google-generativeai~=0.7.2
|
||||||
|
grpcio-tools~=1.62.2
|
||||||
|
langchain~=0.2.14
|
||||||
|
livekit~=0.13.1
|
||||||
|
lmnt~=1.1.4
|
||||||
|
loguru~=0.7.2
|
||||||
|
numpy~=1.26.4
|
||||||
|
openai~=1.37.2
|
||||||
|
openpipe~=4.24.0
|
||||||
|
Pillow~=10.4.0
|
||||||
|
pip-tools~=7.4.1
|
||||||
|
pyaudio~=0.2.14
|
||||||
|
pydantic~=2.8.2
|
||||||
|
pyloudnorm~=0.1.1
|
||||||
|
pyht~=0.0.28
|
||||||
|
pyright~=1.1.376
|
||||||
|
pytest~=8.3.2
|
||||||
|
python-dotenv~=1.0.1
|
||||||
|
resampy~=0.4.3
|
||||||
|
setuptools~=72.2.0
|
||||||
|
setuptools_scm~=8.1.0
|
||||||
|
silero-vad~=5.1
|
||||||
|
together~=1.2.7
|
||||||
|
transformers~=4.44.0
|
||||||
|
websockets~=12.0
|
||||||
@@ -1,14 +1,19 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
from pipecat.pipeline.openai_frames import OpenAILLMContextFrame
|
from pipecat.processors.aggregators.openai_llm_context import (
|
||||||
from pipecat.services.azure_ai_services import AzureLLMService
|
OpenAILLMContext,
|
||||||
from pipecat.services.openai_llm_context import OpenAILLMContext
|
OpenAILLMContextFrame
|
||||||
|
)
|
||||||
|
from pipecat.services.azure import AzureLLMService
|
||||||
|
|
||||||
from openai.types.chat import (
|
from openai.types.chat import (
|
||||||
ChatCompletionSystemMessageParam,
|
ChatCompletionSystemMessageParam,
|
||||||
)
|
)
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
@unittest.skip("Skip azure integration test")
|
||||||
async def test_chat():
|
async def test_chat():
|
||||||
llm = AzureLLMService(
|
llm = AzureLLMService(
|
||||||
api_key=os.getenv("AZURE_CHATGPT_API_KEY"),
|
api_key=os.getenv("AZURE_CHATGPT_API_KEY"),
|
||||||
|
|||||||
@@ -1,13 +1,18 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from pipecat.pipeline.openai_frames import OpenAILLMContextFrame
|
from pipecat.processors.aggregators.openai_llm_context import (
|
||||||
from pipecat.services.openai_llm_context import OpenAILLMContext
|
OpenAILLMContext,
|
||||||
|
OpenAILLMContextFrame
|
||||||
|
)
|
||||||
|
|
||||||
from openai.types.chat import (
|
from openai.types.chat import (
|
||||||
ChatCompletionSystemMessageParam,
|
ChatCompletionSystemMessageParam,
|
||||||
)
|
)
|
||||||
from pipecat.services.ollama_ai_services import OLLamaLLMService
|
from pipecat.services.ollama import OLLamaLLMService
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
@unittest.skip("Skip azure integration test")
|
||||||
async def test_chat():
|
async def test_chat():
|
||||||
llm = OLLamaLLMService()
|
llm = OLLamaLLMService()
|
||||||
context = OpenAILLMContext()
|
context = OpenAILLMContext()
|
||||||
|
|||||||
@@ -3,18 +3,18 @@ import doctest
|
|||||||
import functools
|
import functools
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from pipecat.pipeline.aggregators import (
|
from pipecat.processors.aggregators.gated import GatedAggregator
|
||||||
GatedAggregator,
|
from pipecat.processors.aggregators.sentence import SentenceAggregator
|
||||||
ParallelPipeline,
|
from pipecat.processors.text_transformer import StatelessTextTransformer
|
||||||
SentenceAggregator,
|
|
||||||
StatelessTextTransformer,
|
from pipecat.pipeline.parallel_pipeline import ParallelPipeline
|
||||||
)
|
|
||||||
from pipecat.pipeline.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioFrame,
|
AudioRawFrame,
|
||||||
EndFrame,
|
EndFrame,
|
||||||
ImageFrame,
|
ImageRawFrame,
|
||||||
LLMResponseEndFrame,
|
LLMFullResponseEndFrame,
|
||||||
LLMResponseStartFrame,
|
LLMFullResponseStartFrame,
|
||||||
Frame,
|
Frame,
|
||||||
TextFrame,
|
TextFrame,
|
||||||
)
|
)
|
||||||
@@ -23,6 +23,7 @@ from pipecat.pipeline.pipeline import Pipeline
|
|||||||
|
|
||||||
|
|
||||||
class TestDailyFrameAggregators(unittest.IsolatedAsyncioTestCase):
|
class TestDailyFrameAggregators(unittest.IsolatedAsyncioTestCase):
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_sentence_aggregator(self):
|
async def test_sentence_aggregator(self):
|
||||||
sentence = "Hello, world. How are you? I am fine"
|
sentence = "Hello, world. How are you? I am fine"
|
||||||
expected_sentences = ["Hello, world.", " How are you?", " I am fine "]
|
expected_sentences = ["Hello, world.", " How are you?", " I am fine "]
|
||||||
@@ -43,36 +44,38 @@ class TestDailyFrameAggregators(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
self.assertEqual(expected_sentences, [])
|
self.assertEqual(expected_sentences, [])
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_gated_accumulator(self):
|
async def test_gated_accumulator(self):
|
||||||
gated_aggregator = GatedAggregator(
|
gated_aggregator = GatedAggregator(
|
||||||
gate_open_fn=lambda frame: isinstance(
|
gate_open_fn=lambda frame: isinstance(
|
||||||
frame, ImageFrame), gate_close_fn=lambda frame: isinstance(
|
frame, ImageRawFrame), gate_close_fn=lambda frame: isinstance(
|
||||||
frame, LLMResponseStartFrame), start_open=False, )
|
frame, LLMFullResponseStartFrame), start_open=False, )
|
||||||
|
|
||||||
frames = [
|
frames = [
|
||||||
LLMResponseStartFrame(),
|
LLMFullResponseStartFrame(),
|
||||||
TextFrame("Hello, "),
|
TextFrame("Hello, "),
|
||||||
TextFrame("world."),
|
TextFrame("world."),
|
||||||
AudioFrame(b"hello"),
|
AudioRawFrame(b"hello"),
|
||||||
ImageFrame(b"image", (0, 0)),
|
ImageRawFrame(b"image", (0, 0)),
|
||||||
AudioFrame(b"world"),
|
AudioRawFrame(b"world"),
|
||||||
LLMResponseEndFrame(),
|
LLMFullResponseEndFrame(),
|
||||||
]
|
]
|
||||||
|
|
||||||
expected_output_frames = [
|
expected_output_frames = [
|
||||||
ImageFrame(b"image", (0, 0)),
|
ImageRawFrame(b"image", (0, 0)),
|
||||||
LLMResponseStartFrame(),
|
LLMFullResponseStartFrame(),
|
||||||
TextFrame("Hello, "),
|
TextFrame("Hello, "),
|
||||||
TextFrame("world."),
|
TextFrame("world."),
|
||||||
AudioFrame(b"hello"),
|
AudioRawFrame(b"hello"),
|
||||||
AudioFrame(b"world"),
|
AudioRawFrame(b"world"),
|
||||||
LLMResponseEndFrame(),
|
LLMFullResponseEndFrame(),
|
||||||
]
|
]
|
||||||
for frame in frames:
|
for frame in frames:
|
||||||
async for out_frame in gated_aggregator.process_frame(frame):
|
async for out_frame in gated_aggregator.process_frame(frame):
|
||||||
self.assertEqual(out_frame, expected_output_frames.pop(0))
|
self.assertEqual(out_frame, expected_output_frames.pop(0))
|
||||||
self.assertEqual(expected_output_frames, [])
|
self.assertEqual(expected_output_frames, [])
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_parallel_pipeline(self):
|
async def test_parallel_pipeline(self):
|
||||||
|
|
||||||
async def slow_add(sleep_time: float, name: str, x: str):
|
async def slow_add(sleep_time: float, name: str, x: str):
|
||||||
@@ -124,6 +127,6 @@ class TestDailyFrameAggregators(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
def load_tests(loader, tests, ignore):
|
def load_tests(loader, tests, ignore):
|
||||||
""" Run doctests on the aggregators module. """
|
""" Run doctests on the aggregators module. """
|
||||||
from pipecat.pipeline import aggregators
|
from pipecat.processors import aggregators
|
||||||
tests.addTests(doctest.DocTestSuite(aggregators))
|
tests.addTests(doctest.DocTestSuite(aggregators))
|
||||||
return tests
|
return tests
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ import unittest
|
|||||||
|
|
||||||
class TestDailyTransport(unittest.IsolatedAsyncioTestCase):
|
class TestDailyTransport(unittest.IsolatedAsyncioTestCase):
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_event_handler(self):
|
async def test_event_handler(self):
|
||||||
from pipecat.transports.daily_transport import DailyTransport
|
from pipecat.transports.daily_transport import DailyTransport
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ load_dotenv()
|
|||||||
|
|
||||||
|
|
||||||
class TestWhisperOpenAIService(unittest.IsolatedAsyncioTestCase):
|
class TestWhisperOpenAIService(unittest.IsolatedAsyncioTestCase):
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_whisper_tts(self):
|
async def test_whisper_tts(self):
|
||||||
pa = pyaudio.PyAudio()
|
pa = pyaudio.PyAudio()
|
||||||
stream = pa.open(format=pyaudio.paInt16,
|
stream = pa.open(format=pyaudio.paInt16,
|
||||||
|
|||||||
@@ -2,15 +2,17 @@ import asyncio
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import Mock
|
from unittest.mock import Mock
|
||||||
|
|
||||||
from pipecat.pipeline.aggregators import SentenceAggregator, StatelessTextTransformer
|
from pipecat.processors.aggregators.sentence import SentenceAggregator
|
||||||
from pipecat.pipeline.frame_processor import FrameProcessor
|
from pipecat.processors.text_transformer import StatelessTextTransformer
|
||||||
from pipecat.pipeline.frames import EndFrame, TextFrame
|
from pipecat.processors.frame_processor import FrameProcessor
|
||||||
|
from pipecat.frames.frames import EndFrame, TextFrame
|
||||||
|
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
from pipecat.pipeline.pipeline import Pipeline
|
||||||
|
|
||||||
|
|
||||||
class TestDailyPipeline(unittest.IsolatedAsyncioTestCase):
|
class TestDailyPipeline(unittest.IsolatedAsyncioTestCase):
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_pipeline_simple(self):
|
async def test_pipeline_simple(self):
|
||||||
aggregator = SentenceAggregator()
|
aggregator = SentenceAggregator()
|
||||||
|
|
||||||
@@ -27,6 +29,7 @@ class TestDailyPipeline(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.assertEqual(await outgoing_queue.get(), TextFrame("Hello, world."))
|
self.assertEqual(await outgoing_queue.get(), TextFrame("Hello, world."))
|
||||||
self.assertIsInstance(await outgoing_queue.get(), EndFrame)
|
self.assertIsInstance(await outgoing_queue.get(), EndFrame)
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_pipeline_multiple_stages(self):
|
async def test_pipeline_multiple_stages(self):
|
||||||
sentence_aggregator = SentenceAggregator()
|
sentence_aggregator = SentenceAggregator()
|
||||||
to_upper = StatelessTextTransformer(lambda x: x.upper())
|
to_upper = StatelessTextTransformer(lambda x: x.upper())
|
||||||
@@ -78,18 +81,21 @@ class TestLogFrame(unittest.TestCase):
|
|||||||
self.pipeline._name = 'MyClass'
|
self.pipeline._name = 'MyClass'
|
||||||
self.pipeline._logger = Mock()
|
self.pipeline._logger = Mock()
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
def test_log_frame_from_source(self):
|
def test_log_frame_from_source(self):
|
||||||
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
||||||
self.pipeline._log_frame(frame, depth=1)
|
self.pipeline._log_frame(frame, depth=1)
|
||||||
self.pipeline._logger.debug.assert_called_once_with(
|
self.pipeline._logger.debug.assert_called_once_with(
|
||||||
'MyClass source -> MyFrame -> processor1')
|
'MyClass source -> MyFrame -> processor1')
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
def test_log_frame_to_sink(self):
|
def test_log_frame_to_sink(self):
|
||||||
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
||||||
self.pipeline._log_frame(frame, depth=3)
|
self.pipeline._log_frame(frame, depth=3)
|
||||||
self.pipeline._logger.debug.assert_called_once_with(
|
self.pipeline._logger.debug.assert_called_once_with(
|
||||||
'MyClass processor2 -> MyFrame -> sink')
|
'MyClass processor2 -> MyFrame -> sink')
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
def test_log_frame_repeated_log(self):
|
def test_log_frame_repeated_log(self):
|
||||||
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
frame = Mock(__class__=Mock(__name__='MyFrame'))
|
||||||
self.pipeline._log_frame(frame, depth=2)
|
self.pipeline._log_frame(frame, depth=2)
|
||||||
@@ -98,6 +104,7 @@ class TestLogFrame(unittest.TestCase):
|
|||||||
self.pipeline._log_frame(frame, depth=2)
|
self.pipeline._log_frame(frame, depth=2)
|
||||||
self.pipeline._logger.debug.assert_called_with('MyClass ... repeated')
|
self.pipeline._logger.debug.assert_called_with('MyClass ... repeated')
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
def test_log_frame_reset_repeated_log(self):
|
def test_log_frame_reset_repeated_log(self):
|
||||||
frame1 = Mock(__class__=Mock(__name__='MyFrame1'))
|
frame1 = Mock(__class__=Mock(__name__='MyFrame1'))
|
||||||
frame2 = Mock(__class__=Mock(__name__='MyFrame2'))
|
frame2 = Mock(__class__=Mock(__name__='MyFrame2'))
|
||||||
|
|||||||
@@ -1,13 +1,14 @@
|
|||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from pipecat.pipeline.frames import AudioFrame, TextFrame, TranscriptionFrame
|
from pipecat.frames.frames import AudioRawFrame, TextFrame, TranscriptionFrame
|
||||||
from pipecat.serializers.protobuf_serializer import ProtobufFrameSerializer
|
from pipecat.serializers.protobuf import ProtobufFrameSerializer
|
||||||
|
|
||||||
|
|
||||||
class TestProtobufFrameSerializer(unittest.IsolatedAsyncioTestCase):
|
class TestProtobufFrameSerializer(unittest.IsolatedAsyncioTestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.serializer = ProtobufFrameSerializer()
|
self.serializer = ProtobufFrameSerializer()
|
||||||
|
|
||||||
|
@unittest.skip("FIXME: This test is failing")
|
||||||
async def test_roundtrip(self):
|
async def test_roundtrip(self):
|
||||||
text_frame = TextFrame(text='hello world')
|
text_frame = TextFrame(text='hello world')
|
||||||
frame = self.serializer.deserialize(
|
frame = self.serializer.deserialize(
|
||||||
@@ -20,7 +21,7 @@ class TestProtobufFrameSerializer(unittest.IsolatedAsyncioTestCase):
|
|||||||
self.serializer.serialize(transcription_frame))
|
self.serializer.serialize(transcription_frame))
|
||||||
self.assertEqual(frame, transcription_frame)
|
self.assertEqual(frame, transcription_frame)
|
||||||
|
|
||||||
audio_frame = AudioFrame(data=b'1234567890')
|
audio_frame = AudioRawFrame(data=b'1234567890')
|
||||||
frame = self.serializer.deserialize(
|
frame = self.serializer.deserialize(
|
||||||
self.serializer.serialize(audio_frame))
|
self.serializer.serialize(audio_frame))
|
||||||
self.assertEqual(frame, audio_frame)
|
self.assertEqual(frame, audio_frame)
|
||||||
|
|||||||
@@ -1,113 +1,113 @@
|
|||||||
import asyncio
|
# import asyncio
|
||||||
import unittest
|
# import unittest
|
||||||
from unittest.mock import AsyncMock, patch, Mock
|
# from unittest.mock import AsyncMock, patch, Mock
|
||||||
|
|
||||||
from pipecat.pipeline.frames import AudioFrame, EndFrame, TextFrame, TTSEndFrame, TTSStartFrame
|
# from pipecat.pipeline.frames import AudioFrame, EndFrame, TextFrame, TTSEndFrame, TTSStartFrame
|
||||||
from pipecat.pipeline.pipeline import Pipeline
|
# from pipecat.pipeline.pipeline import Pipeline
|
||||||
from pipecat.transports.websocket_transport import WebSocketFrameProcessor, WebsocketTransport
|
# from pipecat.transports.websocket_transport import WebSocketFrameProcessor, WebsocketTransport
|
||||||
|
|
||||||
|
|
||||||
class TestWebSocketTransportService(unittest.IsolatedAsyncioTestCase):
|
# class TestWebSocketTransportService(unittest.IsolatedAsyncioTestCase):
|
||||||
def setUp(self):
|
# def setUp(self):
|
||||||
self.transport = WebsocketTransport(host="localhost", port=8765)
|
# self.transport = WebsocketTransport(host="localhost", port=8765)
|
||||||
self.pipeline = Pipeline([])
|
# self.pipeline = Pipeline([])
|
||||||
self.sample_frame = TextFrame("Hello there!")
|
# self.sample_frame = TextFrame("Hello there!")
|
||||||
self.serialized_sample_frame = self.transport._serializer.serialize(
|
# self.serialized_sample_frame = self.transport._serializer.serialize(
|
||||||
self.sample_frame)
|
# self.sample_frame)
|
||||||
|
|
||||||
async def queue_frame(self):
|
# async def queue_frame(self):
|
||||||
await asyncio.sleep(0.1)
|
# await asyncio.sleep(0.1)
|
||||||
await self.pipeline.queue_frames([self.sample_frame, EndFrame()])
|
# await self.pipeline.queue_frames([self.sample_frame, EndFrame()])
|
||||||
|
|
||||||
async def test_websocket_handler(self):
|
# async def test_websocket_handler(self):
|
||||||
mock_websocket = AsyncMock()
|
# mock_websocket = AsyncMock()
|
||||||
|
|
||||||
with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
# with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
||||||
mock_serve.return_value.__anext__.return_value = (
|
# mock_serve.return_value.__anext__.return_value = (
|
||||||
mock_websocket, "/")
|
# mock_websocket, "/")
|
||||||
|
|
||||||
await self.transport._websocket_handler(mock_websocket, "/")
|
# await self.transport._websocket_handler(mock_websocket, "/")
|
||||||
|
|
||||||
await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
# await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
||||||
self.assertEqual(mock_websocket.send.call_count, 1)
|
# self.assertEqual(mock_websocket.send.call_count, 1)
|
||||||
|
|
||||||
self.assertEqual(
|
# self.assertEqual(
|
||||||
mock_websocket.send.call_args[0][0], self.serialized_sample_frame)
|
# mock_websocket.send.call_args[0][0], self.serialized_sample_frame)
|
||||||
|
|
||||||
async def test_on_connection_decorator(self):
|
# async def test_on_connection_decorator(self):
|
||||||
mock_websocket = AsyncMock()
|
# mock_websocket = AsyncMock()
|
||||||
|
|
||||||
connection_handler_called = asyncio.Event()
|
# connection_handler_called = asyncio.Event()
|
||||||
|
|
||||||
@self.transport.on_connection
|
# @self.transport.on_connection
|
||||||
async def connection_handler():
|
# async def connection_handler():
|
||||||
connection_handler_called.set()
|
# connection_handler_called.set()
|
||||||
|
|
||||||
with patch("websockets.serve", return_value=AsyncMock()):
|
# with patch("websockets.serve", return_value=AsyncMock()):
|
||||||
await self.transport._websocket_handler(mock_websocket, "/")
|
# await self.transport._websocket_handler(mock_websocket, "/")
|
||||||
|
|
||||||
self.assertTrue(connection_handler_called.is_set())
|
# self.assertTrue(connection_handler_called.is_set())
|
||||||
|
|
||||||
async def test_frame_processor(self):
|
# async def test_frame_processor(self):
|
||||||
processor = WebSocketFrameProcessor(audio_frame_size=4)
|
# processor = WebSocketFrameProcessor(audio_frame_size=4)
|
||||||
|
|
||||||
source_frames = [
|
# source_frames = [
|
||||||
TTSStartFrame(),
|
# TTSStartFrame(),
|
||||||
AudioFrame(b"1234"),
|
# AudioFrame(b"1234"),
|
||||||
AudioFrame(b"5678"),
|
# AudioFrame(b"5678"),
|
||||||
TTSEndFrame(),
|
# TTSEndFrame(),
|
||||||
TextFrame("hello world")
|
# TextFrame("hello world")
|
||||||
]
|
# ]
|
||||||
|
|
||||||
frames = []
|
# frames = []
|
||||||
for frame in source_frames:
|
# for frame in source_frames:
|
||||||
async for output_frame in processor.process_frame(frame):
|
# async for output_frame in processor.process_frame(frame):
|
||||||
frames.append(output_frame)
|
# frames.append(output_frame)
|
||||||
|
|
||||||
self.assertEqual(len(frames), 3)
|
# self.assertEqual(len(frames), 3)
|
||||||
self.assertIsInstance(frames[0], AudioFrame)
|
# self.assertIsInstance(frames[0], AudioFrame)
|
||||||
self.assertEqual(frames[0].data, b"1234")
|
# self.assertEqual(frames[0].data, b"1234")
|
||||||
self.assertIsInstance(frames[1], AudioFrame)
|
# self.assertIsInstance(frames[1], AudioFrame)
|
||||||
self.assertEqual(frames[1].data, b"5678")
|
# self.assertEqual(frames[1].data, b"5678")
|
||||||
self.assertIsInstance(frames[2], TextFrame)
|
# self.assertIsInstance(frames[2], TextFrame)
|
||||||
self.assertEqual(frames[2].text, "hello world")
|
# self.assertEqual(frames[2].text, "hello world")
|
||||||
|
|
||||||
async def test_serializer_parameter(self):
|
# async def test_serializer_parameter(self):
|
||||||
mock_websocket = AsyncMock()
|
# mock_websocket = AsyncMock()
|
||||||
|
|
||||||
# Test with ProtobufFrameSerializer (default)
|
# # Test with ProtobufFrameSerializer (default)
|
||||||
with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
# with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
||||||
mock_serve.return_value.__anext__.return_value = (
|
# mock_serve.return_value.__anext__.return_value = (
|
||||||
mock_websocket, "/")
|
# mock_websocket, "/")
|
||||||
|
|
||||||
await self.transport._websocket_handler(mock_websocket, "/")
|
# await self.transport._websocket_handler(mock_websocket, "/")
|
||||||
|
|
||||||
await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
# await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
||||||
self.assertEqual(mock_websocket.send.call_count, 1)
|
# self.assertEqual(mock_websocket.send.call_count, 1)
|
||||||
self.assertEqual(
|
# self.assertEqual(
|
||||||
mock_websocket.send.call_args[0][0],
|
# mock_websocket.send.call_args[0][0],
|
||||||
self.serialized_sample_frame,
|
# self.serialized_sample_frame,
|
||||||
)
|
# )
|
||||||
|
|
||||||
# Test with a mock serializer
|
# # Test with a mock serializer
|
||||||
mock_serializer = Mock()
|
# mock_serializer = Mock()
|
||||||
mock_serializer.serialize.return_value = b"mock_serialized_data"
|
# mock_serializer.serialize.return_value = b"mock_serialized_data"
|
||||||
self.transport = WebsocketTransport(
|
# self.transport = WebsocketTransport(
|
||||||
host="localhost", port=8765, serializer=mock_serializer
|
# host="localhost", port=8765, serializer=mock_serializer
|
||||||
)
|
# )
|
||||||
mock_websocket.reset_mock()
|
# mock_websocket.reset_mock()
|
||||||
with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
# with patch("websockets.serve", return_value=AsyncMock()) as mock_serve:
|
||||||
mock_serve.return_value.__anext__.return_value = (
|
# mock_serve.return_value.__anext__.return_value = (
|
||||||
mock_websocket, "/")
|
# mock_websocket, "/")
|
||||||
|
|
||||||
await self.transport._websocket_handler(mock_websocket, "/")
|
# await self.transport._websocket_handler(mock_websocket, "/")
|
||||||
await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
# await asyncio.gather(self.transport.run(self.pipeline), self.queue_frame())
|
||||||
self.assertEqual(mock_websocket.send.call_count, 1)
|
# self.assertEqual(mock_websocket.send.call_count, 1)
|
||||||
self.assertEqual(
|
# self.assertEqual(
|
||||||
mock_websocket.send.call_args[0][0], b"mock_serialized_data")
|
# mock_websocket.send.call_args[0][0], b"mock_serialized_data")
|
||||||
mock_serializer.serialize.assert_called_once_with(
|
# mock_serializer.serialize.assert_called_once_with(
|
||||||
TextFrame("Hello there!"))
|
# TextFrame("Hello there!"))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
# if __name__ == "__main__":
|
||||||
unittest.main()
|
# unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user