probably will undo most of this
This commit is contained in:
@@ -18,11 +18,19 @@ from dailyai.queue_frame import (
|
|||||||
from abc import abstractmethod
|
from abc import abstractmethod
|
||||||
from typing import AsyncGenerator, AsyncIterable, BinaryIO, Iterable
|
from typing import AsyncGenerator, AsyncIterable, BinaryIO, Iterable
|
||||||
|
|
||||||
class AbstractPipeService():
|
class AbstractPipeService:
|
||||||
def __init__(self, source:asyncio.Queue[QueueFrame], sink:asyncio.Queue[QueueFrame]):
|
|
||||||
self.source: asyncio.Queue[QueueFrame] = source
|
def __init__(
|
||||||
self.sink: asyncio.Queue[QueueFrame] = sink
|
self,
|
||||||
pass
|
):
|
||||||
|
self.source_queue: asyncio.Queue[QueueFrame] = asyncio.Queue()
|
||||||
|
self.sink_queue: asyncio.Queue[QueueFrame] = asyncio.Queue()
|
||||||
|
|
||||||
|
async def get(self) -> QueueFrame:
|
||||||
|
return await self.sink_queue.get()
|
||||||
|
|
||||||
|
async def put(self, frame: QueueFrame) -> None:
|
||||||
|
await self.source_queue.put(frame)
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def process_queue(self):
|
async def process_queue(self):
|
||||||
@@ -32,34 +40,20 @@ class PipeService(AbstractPipeService):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
source: asyncio.Queue[QueueFrame] | AbstractPipeService | None=None,
|
source: AbstractPipeService | None = None,
|
||||||
sink: asyncio.Queue[QueueFrame] | AbstractPipeService | None=None,
|
sink: AbstractPipeService | None = None,
|
||||||
):
|
):
|
||||||
|
super().__init__()
|
||||||
self.logger: logging.Logger = logging.getLogger("dailyai")
|
self.logger: logging.Logger = logging.getLogger("dailyai")
|
||||||
|
|
||||||
if not source:
|
|
||||||
source = asyncio.Queue()
|
|
||||||
elif isinstance(source, AbstractPipeService):
|
|
||||||
source = source.sink
|
|
||||||
|
|
||||||
if not sink:
|
|
||||||
sink = asyncio.Queue()
|
|
||||||
elif isinstance(sink, AbstractPipeService):
|
|
||||||
sink = sink.source
|
|
||||||
|
|
||||||
self.source = source
|
self.source = source
|
||||||
self.sink = sink
|
self.sink = sink
|
||||||
|
|
||||||
def chain(self, service: AbstractPipeService) -> AbstractPipeService:
|
|
||||||
self.source = service.sink
|
|
||||||
return self
|
|
||||||
|
|
||||||
async def process_queue(self):
|
async def process_queue(self):
|
||||||
if not self.source:
|
if not self.source:
|
||||||
return
|
return
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
frame = await self.source.get()
|
frame: QueueFrame = await self.source.get()
|
||||||
async for output_frame in self.process_frame(frame):
|
async for output_frame in self.process_frame(frame):
|
||||||
await self.sink.put(output_frame)
|
await self.sink.put(output_frame)
|
||||||
if isinstance(frame, EndStreamQueueFrame):
|
if isinstance(frame, EndStreamQueueFrame):
|
||||||
|
|||||||
10
src/dailyai/tests/test_pipe_service.py
Normal file
10
src/dailyai/tests/test_pipe_service.py
Normal file
@@ -0,0 +1,10 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from dailyai.queue_frame import AudioQueueFrame, ImageQueueFrame
|
||||||
|
from dailyai.services.ai_services import PipeService
|
||||||
|
|
||||||
|
class TestDailyTransport(unittest.IsolatedAsyncioTestCase):
|
||||||
|
def test_pipe_chain(self):
|
||||||
|
pipe1 = PipeService()
|
||||||
@@ -1,8 +1,10 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
from typing import Any, AsyncGenerator, Callable, Tuple
|
||||||
import aiohttp
|
import aiohttp
|
||||||
import os
|
import os
|
||||||
|
|
||||||
from dailyai.queue_frame import AudioQueueFrame, ImageQueueFrame
|
from dailyai.queue_frame import AudioQueueFrame, EndStreamQueueFrame, ImageQueueFrame, LLMResponseEndQueueFrame, QueueFrame, TextQueueFrame
|
||||||
|
from dailyai.services.ai_services import PipeService
|
||||||
from dailyai.services.azure_ai_services import AzureLLMService, AzureImageGenServiceREST, AzureTTSService
|
from dailyai.services.azure_ai_services import AzureLLMService, AzureImageGenServiceREST, AzureTTSService
|
||||||
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
|
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
|
||||||
from dailyai.services.daily_transport_service import DailyTransportService
|
from dailyai.services.daily_transport_service import DailyTransportService
|
||||||
@@ -27,23 +29,109 @@ async def main(room_url):
|
|||||||
camera_height=1024
|
camera_height=1024
|
||||||
)
|
)
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
/ TTS \
|
||||||
|
Month prompt -> LLM -> Fork -> -> Gate -> Transport
|
||||||
|
\ ImageGen /
|
||||||
|
"""
|
||||||
|
|
||||||
|
class QueueFork(PipeService):
|
||||||
|
def __init__(self, source:PipeService, sinks: list[asyncio.Queue[QueueFrame]]):
|
||||||
|
self.source = source
|
||||||
|
self.sinks = sinks
|
||||||
|
|
||||||
|
async def process_queue(self):
|
||||||
|
while True:
|
||||||
|
frame = await self.source.get()
|
||||||
|
for sink in self.sinks:
|
||||||
|
await sink.put(frame)
|
||||||
|
if isinstance(frame, EndStreamQueueFrame):
|
||||||
|
break
|
||||||
|
|
||||||
|
|
||||||
|
class QueueGateOnFrame(PipeService):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
source: asyncio.Queue[QueueFrame],
|
||||||
|
sink: asyncio.Queue[QueueFrame],
|
||||||
|
aggregator: Callable[[Any, QueueFrame], Tuple[Any, QueueFrame | None]]
|
||||||
|
):
|
||||||
|
self.source = source
|
||||||
|
self.sink = sink
|
||||||
|
self.aggregator = aggregator
|
||||||
|
self.accumulation = None
|
||||||
|
|
||||||
|
async def process_frame(
|
||||||
|
self, frame: QueueFrame
|
||||||
|
) -> AsyncGenerator[QueueFrame, None]:
|
||||||
|
output_frame: QueueFrame | None = None
|
||||||
|
(self.aggregation, output_frame) = self.aggregator(
|
||||||
|
self.aggregation, frame
|
||||||
|
)
|
||||||
|
if output_frame:
|
||||||
|
yield output_frame
|
||||||
|
|
||||||
|
class QueueMergeGateOnFirst(PipeService):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, queue_fork_service: QueueFork, sink: asyncio.Queue[QueueFrame]
|
||||||
|
):
|
||||||
|
self.queue_fork_service = queue_fork_service
|
||||||
|
self.sink = sink
|
||||||
|
|
||||||
|
async def process_queue(self):
|
||||||
|
(frames): list[QueueFrame] = await asyncio.gather(
|
||||||
|
*[source.get() for source in self.queue_fork_service.get_end_sinks()]
|
||||||
|
)
|
||||||
|
for idx, frame in enumerate(frames):
|
||||||
|
await self.sink.put(frame)
|
||||||
|
|
||||||
|
# if the first frame we got from a source is an EndStreamQueueFrame, remove that source
|
||||||
|
if isinstance(frame, EndStreamQueueFrame):
|
||||||
|
self.sources.pop(idx)
|
||||||
|
|
||||||
|
async def pass_through(sink, source):
|
||||||
|
while True:
|
||||||
|
frame = await source.get()
|
||||||
|
await sink.put(frame)
|
||||||
|
if isinstance(frame, EndStreamQueueFrame):
|
||||||
|
break
|
||||||
|
|
||||||
|
await asyncio.gather(*[pass_through(self.sink, source) for source in self.sources])
|
||||||
|
|
||||||
llm = AzureLLMService(
|
llm = AzureLLMService(
|
||||||
api_key=os.getenv("AZURE_CHATGPT_API_KEY"),
|
api_key=os.getenv("AZURE_CHATGPT_API_KEY"),
|
||||||
endpoint=os.getenv("AZURE_CHATGPT_ENDPOINT"),
|
endpoint=os.getenv("AZURE_CHATGPT_ENDPOINT"),
|
||||||
model=os.getenv("AZURE_CHATGPT_MODEL"))
|
model=os.getenv("AZURE_CHATGPT_MODEL"))
|
||||||
|
|
||||||
tts = ElevenLabsTTSService(
|
tts = ElevenLabsTTSService(
|
||||||
aiohttp_session=session,
|
aiohttp_session=session,
|
||||||
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
api_key=os.getenv("ELEVENLABS_API_KEY"),
|
||||||
voice_id="ErXwobaYiN019PkySvjV")
|
voice_id="ErXwobaYiN019PkySvjV")
|
||||||
# tts = AzureTTSService(api_key=os.getenv("AZURE_SPEECH_API_KEY"), region=os.getenv("AZURE_SPEECH_REGION"))
|
|
||||||
|
|
||||||
dalle = FalImageGenService(
|
dalle = FalImageGenService(
|
||||||
image_size="1024x1024",
|
image_size="1024x1024",
|
||||||
aiohttp_session=session,
|
aiohttp_session=session,
|
||||||
key_id=os.getenv("FAL_KEY_ID"),
|
key_id=os.getenv("FAL_KEY_ID"),
|
||||||
key_secret=os.getenv("FAL_KEY_SECRET"))
|
key_secret=os.getenv("FAL_KEY_SECRET"),
|
||||||
# dalle = OpenAIImageGenService(aiohttp_session=session, api_key=os.getenv("OPENAI_DALLE_API_KEY"), image_size="1024x1024")
|
)
|
||||||
# dalle = AzureImageGenServiceREST(image_size="1024x1024", aiohttp_session=session, api_key=os.getenv("AZURE_DALLE_API_KEY"), endpoint=os.getenv("AZURE_DALLE_ENDPOINT"), model=os.getenv("AZURE_DALLE_MODEL"))
|
|
||||||
|
def aggregator(
|
||||||
|
accumulation, frame: QueueFrame
|
||||||
|
) -> tuple[Any, QueueFrame | None]:
|
||||||
|
if isinstance(frame, TextQueueFrame):
|
||||||
|
accumulation += frame.text
|
||||||
|
return (accumulation, None)
|
||||||
|
elif isinstance(frame, LLMResponseEndQueueFrame):
|
||||||
|
return ("", TextQueueFrame(accumulation))
|
||||||
|
else:
|
||||||
|
return (accumulation, frame)
|
||||||
|
|
||||||
|
llm_image_gate = QueueGateOnFrame(llm.sink, dalle.source, aggregator)
|
||||||
|
fork_audio_image = QueueFork(llm.sink, [tts.source, llm_image_gate.source])
|
||||||
|
audio_image_gate = QueueMergeGateOnFirst([tts.sink, dalle.sink], transport.send_queue)
|
||||||
|
|
||||||
# Get a complete audio chunk from the given text. Splitting this into its own
|
# Get a complete audio chunk from the given text. Splitting this into its own
|
||||||
# coroutine lets us ensure proper ordering of the audio chunks on the send queue.
|
# coroutine lets us ensure proper ordering of the audio chunks on the send queue.
|
||||||
|
|||||||
Reference in New Issue
Block a user