Storytelling chatbot updates (#1066)
* initial changes for gemini storybot * storybot updates for gemini * more storybot updates * interim interruptible commit * cleanup * cleanup * cleanup * cleanup
This commit is contained in:
@@ -37,8 +37,7 @@ class StoryPromptFrame(TextFrame):
|
||||
|
||||
|
||||
class StoryImageProcessor(FrameProcessor):
|
||||
"""
|
||||
Processor for image prompt frames that will be sent to the FAL service.
|
||||
"""Processor for image prompt frames that will be sent to the FAL service.
|
||||
|
||||
This processor is responsible for consuming frames of type `StoryImageFrame`.
|
||||
It processes them by passing it to the FAL service.
|
||||
@@ -68,8 +67,7 @@ class StoryImageProcessor(FrameProcessor):
|
||||
|
||||
|
||||
class StoryProcessor(FrameProcessor):
|
||||
"""
|
||||
Primary frame processor. It takes the frames generated by the LLM
|
||||
"""Primary frame processor. It takes the frames generated by the LLM
|
||||
and processes them into image prompts and story pages (sentences).
|
||||
For a clearer picture of how this works, reference prompts.py
|
||||
|
||||
@@ -97,44 +95,10 @@ class StoryProcessor(FrameProcessor):
|
||||
await self.push_frame(sounds["talking"])
|
||||
|
||||
elif isinstance(frame, TextFrame):
|
||||
# We want to look for sentence breaks in the text
|
||||
# but since TextFrames are streamed from the LLM
|
||||
# we need to keep a buffer of the text we've seen so far
|
||||
# Add new text to the buffer
|
||||
self._text += frame.text
|
||||
|
||||
# IMAGE PROMPT
|
||||
# Looking for: < [image prompt] > in the LLM response
|
||||
# We prompted our LLM to add an image prompt in the response
|
||||
# so we use regex matching to find it and yield a StoryImageFrame
|
||||
if re.search(r"<.*?>", self._text):
|
||||
if not re.search(r"<.*?>.*?>", self._text):
|
||||
# Pass any frames until we have a closing bracket
|
||||
# otherwise the image prompt will be passed to TTS
|
||||
pass
|
||||
# Extract the image prompt from the text using regex
|
||||
image_prompt = re.search(r"<(.*?)>", self._text).group(1)
|
||||
# Remove the image prompt from the text
|
||||
self._text = re.sub(r"<.*?>", "", self._text, count=1)
|
||||
# Process the image prompt frame
|
||||
await self.push_frame(StoryImageFrame(image_prompt))
|
||||
|
||||
# STORY PAGE
|
||||
# Looking for: [break] in the LLM response
|
||||
# We prompted our LLM to add a [break] after each sentence
|
||||
# so we use regex matching to find it in the LLM response
|
||||
if re.search(r".*\[[bB]reak\].*", self._text):
|
||||
# Remove the [break] token from the text
|
||||
# so it isn't spoken out loud by the TTS
|
||||
self._text = re.sub(r"\[[bB]reak\]", "", self._text, flags=re.IGNORECASE)
|
||||
self._text = self._text.replace("\n", " ")
|
||||
if len(self._text) > 2:
|
||||
# Append the sentence to the story
|
||||
self._story.append(self._text)
|
||||
await self.push_frame(StoryPageFrame(self._text))
|
||||
# Assert that it's the LLMs turn, until we're finished
|
||||
await self.push_frame(DailyTransportMessageFrame(CUE_ASSISTANT_TURN))
|
||||
# Clear the buffer
|
||||
self._text = ""
|
||||
# Process any complete patterns in the order they appear
|
||||
await self.process_text_content()
|
||||
|
||||
# End of a full LLM response
|
||||
# Driven by the prompt, the LLM should have asked the user for input
|
||||
@@ -150,3 +114,38 @@ class StoryProcessor(FrameProcessor):
|
||||
# Anything that is not a TextFrame pass through
|
||||
else:
|
||||
await self.push_frame(frame)
|
||||
|
||||
async def process_text_content(self):
|
||||
"""Process text content in order of appearance, handling both image prompts and story breaks."""
|
||||
while True:
|
||||
# Find the first occurrence of each pattern
|
||||
image_match = re.search(r"<(.*?)>", self._text)
|
||||
break_match = re.search(r"\[[bB]reak\]", self._text)
|
||||
|
||||
# If neither pattern is found, we're done processing
|
||||
if not image_match and not break_match:
|
||||
break
|
||||
|
||||
# Find which pattern comes first in the text
|
||||
image_pos = image_match.start() if image_match else float("inf")
|
||||
break_pos = break_match.start() if break_match else float("inf")
|
||||
|
||||
if image_pos < break_pos:
|
||||
# Process image prompt first
|
||||
image_prompt = image_match.group(1)
|
||||
# Remove the image prompt from the text
|
||||
self._text = self._text[: image_match.start()] + self._text[image_match.end() :]
|
||||
await self.push_frame(StoryImageFrame(image_prompt))
|
||||
else:
|
||||
# Process story break first
|
||||
parts = re.split(r"\[[bB]reak\]", self._text, flags=re.IGNORECASE, maxsplit=1)
|
||||
before_break = parts[0].replace("\n", " ").strip()
|
||||
|
||||
if len(before_break) > 2:
|
||||
self._story.append(before_break)
|
||||
await self.push_frame(StoryPageFrame(before_break))
|
||||
# await self.push_frame(sounds["ding"])
|
||||
await self.push_frame(DailyTransportMessageFrame(CUE_ASSISTANT_TURN))
|
||||
|
||||
# Keep the remainder (if any) in the buffer
|
||||
self._text = parts[1].strip() if len(parts) > 1 else ""
|
||||
|
||||
Reference in New Issue
Block a user