Merge pull request #8 from daily-co/queueframe-refactor

Refactor QueueFrame
This commit is contained in:
Moishe Lettvin
2024-01-23 13:15:11 -05:00
committed by GitHub
10 changed files with 113 additions and 107 deletions

View File

@@ -1,6 +1,6 @@
import asyncio import asyncio
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import LLMMessagesQueueFrame, QueueFrame, TextQueueFrame
from dailyai.services.ai_services import AIService from dailyai.services.ai_services import AIService
from typing import AsyncGenerator, List from typing import AsyncGenerator, List
@@ -34,26 +34,14 @@ class LLMContextAggregator(AIService):
async def process_frame(self, frame:QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame:QueueFrame) -> AsyncGenerator[QueueFrame, None]:
content: str = "" content: str = ""
if frame.frame_type == FrameType.TRANSCRIPTION: # TODO: split up transcription by participant
message = frame.frame_data if isinstance(frame, TextQueueFrame):
if not isinstance(message, dict): content = frame.text
return
if message["session_id"] == self.bot_participant_id:
return
content = message["text"]
elif frame.frame_type == FrameType.TEXT:
if not isinstance(frame.frame_data, str):
return
content = frame.frame_data
# todo: we should differentiate between transcriptions from different participants
self.sentence += content self.sentence += content
if self.sentence.endswith((".", "?", "!")): if self.sentence.endswith((".", "?", "!")):
self.messages.append({"role": self.role, "content": self.sentence}) self.messages.append({"role": self.role, "content": self.sentence})
self.sentence = "" self.sentence = ""
yield QueueFrame(FrameType.LLM_MESSAGE, self.messages) yield LLMMessagesQueueFrame(self.messages)
yield frame yield frame

View File

@@ -1,18 +1,38 @@
from enum import Enum from enum import Enum
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any
class FrameType(Enum):
NOOP = -1
START_STREAM = 0
END_STREAM = 1
AUDIO = 2
IMAGE = 3
TEXT = 4
TRANSCRIPTION = 5
LLM_MESSAGE = 6
APP_MESSAGE = 7
@dataclass(frozen=True)
class QueueFrame: class QueueFrame:
frame_type: FrameType pass
frame_data: str | dict | bytes | list | None
class StartStreamQueueFrame(QueueFrame):
pass
class EndStreamQueueFrame(QueueFrame):
pass
@dataclass()
class AudioQueueFrame(QueueFrame):
data: bytes
@dataclass()
class ImageQueueFrame(QueueFrame):
url: str | None
image: bytes
@dataclass()
class TextQueueFrame(QueueFrame):
text: str
@dataclass()
class TranscriptionQueueFrame(TextQueueFrame):
participantId: str
timestamp: str
@dataclass()
class LLMMessagesQueueFrame(QueueFrame):
messages: list[dict[str,str]] # TODO: define this more concretely!
class AppMessageQueueFrame(QueueFrame):
message: Any
participantId: str

View File

@@ -1,10 +1,14 @@
import asyncio import asyncio
import logging import logging
import re
from httpx import request from dailyai.queue_frame import (
AudioQueueFrame,
from dailyai.queue_frame import QueueFrame, FrameType EndStreamQueueFrame,
ImageQueueFrame,
LLMMessagesQueueFrame,
QueueFrame,
TextQueueFrame,
)
from abc import abstractmethod from abc import abstractmethod
from typing import AsyncGenerator, AsyncIterable, Iterable from typing import AsyncGenerator, AsyncIterable, Iterable
@@ -24,7 +28,7 @@ class AIService:
await queue.put(frame) await queue.put(frame)
if add_end_of_stream: if add_end_of_stream:
await queue.put(QueueFrame(FrameType.END_STREAM, None)) await queue.put(EndStreamQueueFrame())
async def run( async def run(
self, self,
@@ -46,7 +50,7 @@ class AIService:
frame = await frames.get() frame = await frames.get()
async for output_frame in self.process_frame(frame): async for output_frame in self.process_frame(frame):
yield output_frame yield output_frame
if frame.frame_type == FrameType.END_STREAM: if isinstance(frame, EndStreamQueueFrame):
break break
else: else:
raise Exception("Frames must be an iterable or async iterable") raise Exception("Frames must be an iterable or async iterable")
@@ -61,21 +65,15 @@ class AIService:
async def process_frame(self, frame:QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame:QueueFrame) -> AsyncGenerator[QueueFrame, None]:
# This is a trick for the interpreter (and linter) to know that this is a generator. # This is a trick for the interpreter (and linter) to know that this is a generator.
if False: if False:
yield QueueFrame(FrameType.NOOP, None) yield QueueFrame()
@abstractmethod @abstractmethod
async def finalize(self) -> AsyncGenerator[QueueFrame, None]: async def finalize(self) -> AsyncGenerator[QueueFrame, None]:
# This is a trick for the interpreter (and linter) to know that this is a generator. # This is a trick for the interpreter (and linter) to know that this is a generator.
if False: if False:
yield QueueFrame(FrameType.NOOP, None) yield QueueFrame()
class LLMService(AIService): class LLMService(AIService):
def allowed_input_frame_types(self) -> set[FrameType]:
return set([FrameType.LLM_MESSAGE])
def allowed_output_frame_types(self) -> set[FrameType]:
return set([FrameType.TEXT])
@abstractmethod @abstractmethod
async def run_llm_async(self, messages) -> AsyncGenerator[str, None]: async def run_llm_async(self, messages) -> AsyncGenerator[str, None]:
yield "" yield ""
@@ -85,13 +83,9 @@ class LLMService(AIService):
pass pass
async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]:
if frame.frame_type == FrameType.LLM_MESSAGE: if isinstance(frame, LLMMessagesQueueFrame):
if type(frame.frame_data) != list: async for text_chunk in self.run_llm_async(frame.messages):
raise Exception("LLM service requires a dict for the data field") yield TextQueueFrame(text_chunk)
messages: list[dict[str, str]] = frame.frame_data
async for text_chunk in self.run_llm_async(messages):
yield QueueFrame(FrameType.TEXT, text_chunk)
class TTSService(AIService): class TTSService(AIService):
@@ -112,34 +106,31 @@ class TTSService(AIService):
yield bytes() yield bytes()
async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]:
if frame.frame_type != FrameType.TEXT: if not isinstance(frame, TextQueueFrame):
yield frame yield frame
return return
if not isinstance(frame.frame_data, str):
raise(Exception(f"Invalid data type in frame type: {frame.frame_type}, type: {type(frame.frame_data)}"))
text: str | None = None text: str | None = None
if not self.aggregate_sentences: if not self.aggregate_sentences:
text = frame.frame_data text = frame.text
else: else:
self.current_sentence += frame.frame_data self.current_sentence += frame.text
if self.current_sentence.endswith((".", "?", "!")): if self.current_sentence.endswith((".", "?", "!")):
text = self.current_sentence text = self.current_sentence
self.current_sentence = "" self.current_sentence = ""
if text: if text:
async for audio_chunk in self.run_tts(text): async for audio_chunk in self.run_tts(text):
yield QueueFrame(FrameType.AUDIO, audio_chunk) yield AudioQueueFrame(audio_chunk)
async def finalize(self): async def finalize(self):
if self.current_sentence: if self.current_sentence:
async for audio_chunk in self.run_tts(self.current_sentence): async for audio_chunk in self.run_tts(self.current_sentence):
yield QueueFrame(FrameType.AUDIO, audio_chunk) yield AudioQueueFrame(audio_chunk)
# Convenience function to send the audio for a sentence to the given queue # Convenience function to send the audio for a sentence to the given queue
async def say(self, sentence, queue: asyncio.Queue): async def say(self, sentence, queue: asyncio.Queue):
await self.run_to_queue(queue, [QueueFrame(FrameType.TEXT, sentence)]) await self.run_to_queue(queue, [TextQueueFrame(sentence)])
class ImageGenService(AIService): class ImageGenService(AIService):
@@ -149,15 +140,15 @@ class ImageGenService(AIService):
# Renders the image. Returns an Image object. # Renders the image. Returns an Image object.
@abstractmethod @abstractmethod
async def run_image_gen(self, sentence) -> tuple[str, bytes]: async def run_image_gen(self, sentence:str) -> tuple[str, bytes]:
pass pass
async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]:
if type(frame.frame_data) != str: if not isinstance(frame, TextQueueFrame):
raise Exception("Image service requires a string for the data field") return
(_, image_data) = await self.run_image_gen(frame.frame_data) (url, image_data) = await self.run_image_gen(frame.text)
yield QueueFrame(FrameType.IMAGE, image_data) yield ImageQueueFrame(url, image_data)
@dataclass @dataclass

View File

@@ -8,9 +8,16 @@ import types
from functools import partial from functools import partial
from queue import Queue, Empty from queue import Queue, Empty
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import (
AudioQueueFrame,
EndStreamQueueFrame,
ImageQueueFrame,
QueueFrame,
StartStreamQueueFrame,
TranscriptionQueueFrame,
)
from threading import Thread, Event, Timer from threading import Thread, Event
from daily import ( from daily import (
EventHandler, EventHandler,
@@ -201,7 +208,7 @@ class DailyTransportService(EventHandler):
while True: while True:
frame = await self.receive_queue.get() frame = await self.receive_queue.get()
yield frame yield frame
if frame.frame_type == FrameType.END_STREAM: if isinstance(frame, EndStreamQueueFrame):
break break
def get_async_send_queue(self): def get_async_send_queue(self):
@@ -212,7 +219,7 @@ class DailyTransportService(EventHandler):
frame: QueueFrame | list = await self.send_queue.get() frame: QueueFrame | list = await self.send_queue.get()
self.threadsafe_send_queue.put(frame) self.threadsafe_send_queue.put(frame)
self.send_queue.task_done() self.send_queue.task_done()
if type(frame) == QueueFrame and frame.frame_type == FrameType.END_STREAM: if isinstance(frame, EndStreamQueueFrame):
break break
async def wait_for_send_queue_to_empty(self): async def wait_for_send_queue_to_empty(self):
@@ -242,8 +249,8 @@ class DailyTransportService(EventHandler):
self.stop_threads.set() self.stop_threads.set()
await self.receive_queue.put(QueueFrame(FrameType.END_STREAM, None)) await self.receive_queue.put(EndStreamQueueFrame())
await self.send_queue.put(QueueFrame(FrameType.END_STREAM, None)) await self.send_queue.put(EndStreamQueueFrame())
await async_output_queue_marshal_task await async_output_queue_marshal_task
if self.camera_thread and self.camera_thread.is_alive(): if self.camera_thread and self.camera_thread.is_alive():
@@ -281,7 +288,12 @@ class DailyTransportService(EventHandler):
def on_transcription_message(self, message:dict): def on_transcription_message(self, message:dict):
if self.loop: if self.loop:
frame = QueueFrame(FrameType.TRANSCRIPTION, message) participantId = ""
if "participantId" in message:
participantId = message["participantId"]
elif "session_id" in message:
participantId = message["session_id"]
frame = TranscriptionQueueFrame(message["text"], participantId, message["timestamp"])
asyncio.run_coroutine_threadsafe(self.receive_queue.put(frame), self.loop) asyncio.run_coroutine_threadsafe(self.receive_queue.put(frame), self.loop)
def on_transcription_stopped(self, stopped_by, stopped_by_error): def on_transcription_stopped(self, stopped_by, stopped_by_error):
@@ -314,15 +326,15 @@ class DailyTransportService(EventHandler):
while True: while True:
try: try:
frames_or_frame: QueueFrame | list[QueueFrame] = self.threadsafe_send_queue.get() frames_or_frame: QueueFrame | list[QueueFrame] = self.threadsafe_send_queue.get()
if type(frames_or_frame) == QueueFrame: if isinstance(frames_or_frame, QueueFrame):
frames: list[QueueFrame] = [frames_or_frame] frames: list[QueueFrame] = [frames_or_frame]
elif type(frames_or_frame) == list: elif isinstance(frames_or_frame, list):
frames: list[QueueFrame] = frames_or_frame frames: list[QueueFrame] = frames_or_frame
else: else:
raise Exception("Unknown type in output queue") raise Exception("Unknown type in output queue")
for frame in frames: for frame in frames:
if frame.frame_type == FrameType.END_STREAM: if isinstance(frame, EndStreamQueueFrame):
self.logger.info("Stopping frame consumer thread") self.logger.info("Stopping frame consumer thread")
self.threadsafe_send_queue.task_done() self.threadsafe_send_queue.task_done()
return return
@@ -330,8 +342,8 @@ class DailyTransportService(EventHandler):
# if interrupted, we just pull frames off the queue and discard them # if interrupted, we just pull frames off the queue and discard them
if not self.is_interrupted.is_set(): if not self.is_interrupted.is_set():
if frame: if frame:
if frame.frame_type == FrameType.AUDIO: if isinstance(frame, AudioQueueFrame):
chunk = frame.frame_data chunk = frame.data
all_audio_frames.extend(chunk) all_audio_frames.extend(chunk)
@@ -340,8 +352,8 @@ class DailyTransportService(EventHandler):
if l: if l:
self.mic.write_frames(bytes(b[:l])) self.mic.write_frames(bytes(b[:l]))
b = b[l:] b = b[l:]
elif frame.frame_type == FrameType.IMAGE: elif isinstance(frame, ImageQueueFrame):
self.set_image(frame.frame_data) self.set_image(frame.image)
elif len(b): elif len(b):
self.mic.write_frames(bytes(b)) self.mic.write_frames(bytes(b))
b = bytearray() b = bytearray()
@@ -352,7 +364,7 @@ class DailyTransportService(EventHandler):
) )
self.interrupt_time = None self.interrupt_time = None
if frame.frame_type == FrameType.START_STREAM: if isinstance(frame, StartStreamQueueFrame):
self.is_interrupted.clear() self.is_interrupted.clear()
self.threadsafe_send_queue.task_done() self.threadsafe_send_queue.task_done()

View File

@@ -4,7 +4,7 @@ import unittest
from typing import AsyncGenerator, Generator from typing import AsyncGenerator, Generator
from dailyai.services.ai_services import AIService from dailyai.services.ai_services import AIService
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import EndStreamQueueFrame, QueueFrame, TextQueueFrame
class SimpleAIService(AIService): class SimpleAIService(AIService):
async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]: async def process_frame(self, frame: QueueFrame) -> AsyncGenerator[QueueFrame, None]:
@@ -15,8 +15,8 @@ class TestBaseAIService(unittest.IsolatedAsyncioTestCase):
service = SimpleAIService() service = SimpleAIService()
input_frames = [ input_frames = [
QueueFrame(FrameType.TEXT, "hello"), TextQueueFrame("hello"),
QueueFrame(FrameType.END_STREAM, None), EndStreamQueueFrame()
] ]
async def iterate_frames() -> AsyncGenerator[QueueFrame, None]: async def iterate_frames() -> AsyncGenerator[QueueFrame, None]:
for frame in input_frames: for frame in input_frames:
@@ -31,10 +31,7 @@ class TestBaseAIService(unittest.IsolatedAsyncioTestCase):
async def test_nonasync_input(self): async def test_nonasync_input(self):
service = SimpleAIService() service = SimpleAIService()
input_frames = [ input_frames = [TextQueueFrame("hello"), EndStreamQueueFrame()]
QueueFrame(FrameType.TEXT, "hello"),
QueueFrame(FrameType.END_STREAM, None),
]
def iterate_frames() -> Generator[QueueFrame, None, None]: def iterate_frames() -> Generator[QueueFrame, None, None]:
for frame in input_frames: for frame in input_frames:

View File

@@ -1,10 +1,7 @@
import argparse import argparse
import asyncio import asyncio
from typing import AsyncGenerator
from dailyai.queue_frame import QueueFrame, FrameType
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
from dailyai.services.azure_ai_services import AzureTTSService
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
async def main(room_url): async def main(room_url):

View File

@@ -1,8 +1,7 @@
import argparse import argparse
import asyncio import asyncio
from typing import AsyncGenerator
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import LLMMessagesQueueFrame
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
from dailyai.services.azure_ai_services import AzureLLMService from dailyai.services.azure_ai_services import AzureLLMService
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
@@ -27,7 +26,7 @@ async def main(room_url):
tts_task = asyncio.create_task( tts_task = asyncio.create_task(
tts.run_to_queue( tts.run_to_queue(
transport.send_queue, transport.send_queue,
llm.run([QueueFrame(FrameType.LLM_MESSAGE, messages)]) llm.run([LLMMessagesQueueFrame(messages)]),
) )
) )

View File

@@ -1,7 +1,7 @@
import argparse import argparse
import asyncio import asyncio
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import TextQueueFrame
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
from dailyai.services.open_ai_services import OpenAIImageGenService from dailyai.services.open_ai_services import OpenAIImageGenService
@@ -23,7 +23,7 @@ async def main(room_url):
imagegen = OpenAIImageGenService(image_size="1024x1024") imagegen = OpenAIImageGenService(image_size="1024x1024")
image_task = asyncio.create_task( image_task = asyncio.create_task(
imagegen.run_to_queue(transport.send_queue, [QueueFrame(FrameType.TEXT, "a cat in the style of picasso")]) imagegen.run_to_queue(transport.send_queue, [TextQueueFrame("a cat in the style of picasso")])
) )
@transport.event_handler("on_participant_joined") @transport.event_handler("on_participant_joined")

View File

@@ -4,7 +4,7 @@ import re
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
from dailyai.services.azure_ai_services import AzureLLMService, AzureTTSService from dailyai.services.azure_ai_services import AzureLLMService, AzureTTSService
from dailyai.queue_frame import QueueFrame, FrameType from dailyai.queue_frame import EndStreamQueueFrame, LLMMessagesQueueFrame
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
async def main(room_url:str): async def main(room_url:str):
@@ -35,7 +35,7 @@ async def main(room_url:str):
llm_response_task = asyncio.create_task( llm_response_task = asyncio.create_task(
elevenlabs_tts.run_to_queue( elevenlabs_tts.run_to_queue(
buffer_queue, buffer_queue,
llm.run([QueueFrame(FrameType.LLM_MESSAGE, messages)]), llm.run([LLMMessagesQueueFrame(messages)]),
True, True,
) )
) )
@@ -52,7 +52,7 @@ async def main(room_url:str):
frame = await buffer_queue.get() frame = await buffer_queue.get()
await transport.send_queue.put(frame) await transport.send_queue.put(frame)
buffer_queue.task_done() buffer_queue.task_done()
if frame.frame_type == FrameType.END_STREAM: if isinstance(frame, EndStreamQueueFrame):
break break
await asyncio.gather(llm_response_task, buffer_to_send_queue()) await asyncio.gather(llm_response_task, buffer_to_send_queue())

View File

@@ -1,13 +1,9 @@
import argparse import argparse
import asyncio import asyncio
from asyncio.queues import Queue from dailyai.queue_frame import AudioQueueFrame, ImageQueueFrame
import re
from dailyai.queue_frame import QueueFrame, FrameType
from dailyai.services.azure_ai_services import AzureLLMService from dailyai.services.azure_ai_services import AzureLLMService
from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService from dailyai.services.elevenlabs_ai_service import ElevenLabsTTSService
from dailyai.services.open_ai_services import OpenAIImageGenService
from dailyai.services.daily_transport_service import DailyTransportService from dailyai.services.daily_transport_service import DailyTransportService
from dailyai.services.fal_ai_services import FalImageGenService from dailyai.services.fal_ai_services import FalImageGenService
@@ -48,14 +44,20 @@ async def main(room_url):
] ]
image_description = await llm.run_llm(messages) image_description = await llm.run_llm(messages)
if not image_description:
return
to_speak = f"{month}: {image_description}" to_speak = f"{month}: {image_description}"
audio_task = asyncio.create_task(get_all_audio(to_speak))
image_task = asyncio.create_task(dalle.run_image_gen(image_description))
(audio, image_data) = await asyncio.gather( (audio, image_data) = await asyncio.gather(
get_all_audio(to_speak), dalle.run_image_gen(image_description) audio_task, image_task
) )
return { return {
"month": month, "month": month,
"text": image_description, "text": image_description,
"image_url": image_data[0],
"image": image_data[1], "image": image_data[1],
"audio": audio, "audio": audio,
} }
@@ -84,8 +86,8 @@ async def main(room_url):
data = await month_data_task data = await month_data_task
await transport.send_queue.put( await transport.send_queue.put(
[ [
QueueFrame(FrameType.IMAGE, data["image"]), ImageQueueFrame(data["image_url"], data["image"]),
QueueFrame(FrameType.AUDIO, data["audio"]), AudioQueueFrame(data["audio"]),
] ]
) )