From 3205491166e190512608bf01754815cadae47a92 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Andreas=20K=C3=B6pf?= Date: Thu, 22 Dec 2022 14:51:12 +0100 Subject: [PATCH] add channel handler async msg routing --- .pre-commit-config.yaml | 2 +- bot/bot.py | 258 ++++++++----------------- bot/bot_base.py | 53 +++++ bot/channel_handlers.py | 83 ++++++++ bot/message_templates.py | 18 ++ bot/task_handlers.py | 153 +++++++++++++++ bot/templates/boot.msg | 2 +- bot/templates/task_initial_prompt.msg | 4 +- bot/templates/task_summarize_story.msg | 2 + pyproject.toml | 1 - 10 files changed, 389 insertions(+), 187 deletions(-) create mode 100644 bot/bot_base.py create mode 100644 bot/channel_handlers.py create mode 100644 bot/message_templates.py create mode 100644 bot/task_handlers.py create mode 100644 bot/templates/task_summarize_story.msg diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index f2775e86..cccb2167 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,4 +1,4 @@ -exclude: "build|stubs|bot/templates/.*msg" +exclude: "build|stubs|^bot/templates/" default_language_version: python: python3 diff --git a/bot/bot.py b/bot/bot.py index a420b6c1..b3e2e309 100644 --- a/bot/bot.py +++ b/bot/bot.py @@ -1,62 +1,26 @@ # -*- coding: utf-8 -*- +from __future__ import annotations + import asyncio from datetime import timedelta from pathlib import Path -from typing import Any, Optional, Union +from typing import Optional, Union import discord -import discord.ui as ui -import jinja2 +import task_handlers from api_client import ApiClient, TaskType +from bot_base import BotBase from discord import app_commands from loguru import logger +from message_templates import MessageTemplates from oasst_shared.schemas import protocol as protocol_schema from utils import get_git_head_hash, utcnow -__version__ = "0.0.1" +__version__ = "0.0.2" BOT_NAME = "Open-Assistant Junior" -class RatingButton(discord.ui.Button): - def __init__(self, label, value, response_handler): - super().__init__(label=label, style=discord.ButtonStyle.green) - self.value = value - self.response_handler = response_handler - - async def callback(self, interaction): - await self.response_handler(self.value, interaction) - - -def generate_rating_view(lo: int, hi: int, response_handler) -> discord.ui.View: - view = discord.ui.View() - for i in range(lo, hi + 1): - view.add_item(RatingButton(str(i), i, response_handler)) - return view - - -class Questionnaire(ui.Modal, title="Questionnaire Response"): - name = ui.TextInput(label="Name") - answer = ui.TextInput(label="Answer", style=discord.TextStyle.paragraph) - - async def on_submit(self, interaction: discord.Interaction): - await interaction.response.send_message(f"Thanks for your response, {self.name}!", ephemeral=True) - - -class MessageTemplates: - def __init__(self, template_dir="./templates"): - self.env = jinja2.Environment( - loader=jinja2.FileSystemLoader(template_dir), - autoescape=jinja2.select_autoescape(disabled_extensions=("msg",), default=False, default_for_string=False), - ) - - def render(self, template_name, **kwargs): - template = self.env.get_template(template_name) - txt = template.render(kwargs) - logger.info(txt) - return txt - - -class OpenAssistantBot: +class OpenAssistantBot(BotBase): def __init__( self, bot_token: str, @@ -67,6 +31,8 @@ class OpenAssistantBot: template_dir: str = "./templates", debug: bool = False, ): + super().__init__() + self.template_dir = Path(template_dir) self.bot_channel_name = bot_channel_name self.templates = MessageTemplates(template_dir) @@ -82,10 +48,11 @@ class OpenAssistantBot: self.bot_token = bot_token client = discord.Client(intents=intents) self.client = client + self.loop = client.loop self.bot_channel: discord.TextChannel = None self.backend = ApiClient(backend_url, api_key) - self.reply_handlers = {} # handlers by msg_id + self.tree = app_commands.CommandTree(self.client, fallback_to_global=True) @client.event @@ -109,6 +76,9 @@ class OpenAssistantBot: @self.tree.command() async def tutorial(interaction: discord.Interaction): """Start the Open-Assistant tutorial via DMs.""" + + dm = await self.client.create_dm(discord.Object(interaction.user.id)) + await dm.send("Tutorial coming soon... :-)") await interaction.response.send_message(f"tutorial command by {interaction.user.name}") @self.tree.command() @@ -119,27 +89,12 @@ class OpenAssistantBot: @self.tree.command() async def work(interaction: discord.Interaction): """Request a new personalized task""" + # task = self.backend.fetch_task(protocol_schema.TaskRequestType.rate_summary, user=None) # task = self.backend.fetch_random_task(user=None) - q = Questionnaire() + q = task_handlers.Questionnaire() await interaction.response.send_modal(q) - def ensure_bot_channel(self) -> None: - if self.bot_channel is None: - raise RuntimeError(f"bot channel '{self.bot_channel_name}' not found") - - async def post(self, content: str, view: discord.ui.View = None) -> discord.Message: - self.ensure_bot_channel() - return await self.bot_channel.send(content=content) - - async def post_template_view(self, name: str, *, view: discord.ui.View, **kwargs: Any) -> discord.Message: - logger.info(f"rendering {name}") - text = self.templates.render(name, **kwargs) - return await self.post(text, view) - - async def post_template(self, name: str, **kwargs: Any) -> discord.Message: - return await self.post_template_view(name=name, view=None, **kwargs) - async def post_boot_message(self) -> discord.Message: return await self.post_template( "boot.msg", bot_name=BOT_NAME, version=__version__, git_hash=get_git_head_hash(), debug=self.debug @@ -163,140 +118,64 @@ class OpenAssistantBot: msg: discord.Message = await self.bot_channel.send(f"\n:point_right: {title} :point_left:\n") return msg - async def generate_summarize_story(self, task: protocol_schema.SummarizeStoryTask): - text = f"Summarize to the following story:\n{task.story}" - msg: discord.Message = await self.bot_channel.send(text) - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="Summaries") - - async def on_reply(message: discord.Message): - logger.info("on_summarize_story_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - async def generate_rate_summary(self, task: protocol_schema.RateSummaryTask): - async def rating_response_handler(score, interaction: discord.Interaction): - logger.info("rating_response_handler", score) - await interaction.response.send_message(f"got your feedback: {score}") - - view = generate_rating_view(task.scale.min, task.scale.max, rating_response_handler) - msg = await self.post_template_view("task_rate_summary.msg", view=view, task=task) - - async def on_reply(message: discord.Message): - logger.info("on_summary_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - async def generate_initial_prompt(self, task: protocol_schema.InitialPromptTask): - msg = await self.post_template("task_initial_prompt.msg", task=task) - - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="Prompts") - - async def on_reply(message: discord.Message): - logger.info("on_initial_prompt_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - def _render_message(self, message: protocol_schema.ConversationMessage) -> str: - """Render a message to the user.""" - if message.is_assistant: - return f":robot: Assistant:\n{message.text}" - else: - return f":person_red_hair: User:\n**{message.text}**" - - async def generate_user_reply(self, task: protocol_schema.UserReplyTask): - msg = await self.post_template("task_user_reply.msg", task=task) - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="User responses") - - async def on_reply(message: discord.Message): - logger.info("on_user_reply_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - async def generate_assistant_reply(self, task: protocol_schema.AssistantReplyTask): - msg = await self.post_template("task_assistant_reply.msg", task=task) - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="Agent responses") - - async def on_reply(message: discord.Message): - logger.info("on_assistant_reply_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - async def generate_rank_initial_prompts(self, task: protocol_schema.RankInitialPromptsTask): - msg = await self.post_template("task_rank_initial_prompts.msg", task=task) - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="User responses") - - async def on_reply(message: discord.Message): - logger.info("on_rank_initial_prompts_reply", message) - await message.add_reaction("✅") - - self.reply_handlers[msg.id] = on_reply - - return msg - - async def generate_rank_conversation(self, task: protocol_schema.RankConversationRepliesTask): - msg = await self.post_template("task_rank_conversation_replies.msg", task=task) - await self.bot_channel.create_thread(message=discord.Object(msg.id), name="User responses") - - async def on_reply(message: discord.Message): - logger.info("on_rank_conversation_reply", message) - await message.add_reaction("✅") - message - - self.reply_handlers[msg.id] = on_reply - - return msg - async def next_task(self): - task_type = protocol_schema.TaskRequestType.random + task_type = protocol_schema.TaskRequestType.rate_summary task = self.backend.fetch_task(task_type, user=None) await self.print_separtor("New Task") - msg: discord.Message = None + handler: task_handlers.ChannelTaskBase = None match task.type: case TaskType.summarize_story: - msg = await self.generate_summarize_story(task) + handler = task_handlers.SummarizeStoryHandler() case TaskType.rate_summary: - msg = await self.generate_rate_summary(task) + handler = task_handlers.RateSummaryHandler() case TaskType.initial_prompt: - msg = await self.generate_initial_prompt(task) + handler = task_handlers.InitialPromptHandler() case TaskType.user_reply: - msg = await self.generate_user_reply(task) + handler = task_handlers.UserReplyHandler() case TaskType.assistant_reply: - msg = await self.generate_assistant_reply(task) + handler = task_handlers.AssistantReplyHandler() case TaskType.rank_initial_prompts: - msg = await self.generate_rank_initial_prompts(task) + handler = task_handlers.RankInitialPromptsHandler() case TaskType.rank_user_replies | TaskType.rank_assistant_replies: - msg = await self.generate_rank_conversation(task) + handler = task_handlers.RankConversationsHandler() + case _: + logger.warning(f"Unsupported task type received: {task.type}") + self.backend.nack_task(task.id, "not supported") - if msg is not None: - self.backend.ack_task(task.id, msg.id) - else: - self.backend.nack_task(task.id, "not supported") + if handler: + try: + logger.info(f"strarting task {task.id}") + msg = await handler.start(self, task) + self.backend.ack_task(task.id, msg.id) + except Exception: + logger.exception("Starting task failed.") + self.backend.nack_task(task.id, "faled") async def background_timer(self): + next_remove_completed = utcnow() + timedelta(seconds=10) + next_fetch_task = utcnow() + timedelta(seconds=1) while True: + now = utcnow() + if self.bot_channel: - try: - await self.next_task() - except Exception: - logger.exception("fetching next task failed") - await asyncio.sleep(5) + if now > next_fetch_task: + next_fetch_task = utcnow() + timedelta(seconds=600) + + try: + await self.next_task() + except Exception: + logger.exception("fetching next task failed") + + for x in self.reply_handlers.values(): + x.handler.tick(now) + + if now > next_remove_completed: + next_remove_completed = utcnow() + timedelta(seconds=10) + await self.remove_completed_handlers() + + await asyncio.sleep(1) async def _sync(self, command: str, message: discord.Message): @@ -341,18 +220,33 @@ class OpenAssistantBot: if isinstance(message.channel, discord.Thread): handler = self.reply_handlers.get(message.channel.id) - if handler: - await handler(message) + if handler and not handler.handler.completed: + handler.handler.on_reply(message) if message.reference: handler = self.reply_handlers.get(message.reference.message_id) - if handler: - await handler(message) + if handler and not handler.handler.completed: + handler.handler.on_reply(message) logger.debug( f"{message.type} {message.channel.type} from ({user_display_name}) {user_id}: {message.content} ({type(message.content)})" ) + async def remove_completed_handlers(self): + completed = [k for k, v in self.reply_handlers.items() if v.handler is None or v.handler.completed] + if len(completed) == 0: + return + + for c in completed: + handler = self.reply_handlers[c] + del self.reply_handlers[c] + try: + await handler.handler.finalize() + except Exception: + logger.exception("handler finalize failed") + + logger.info(f"removed {len(completed)} completed handlers (remaining: {len(self.reply_handlers)})") + def get_text_channel_by_name(self, channel_name) -> discord.TextChannel: for channel in self.client.get_all_channels(): if channel.type == discord.ChannelType.text and channel.name == channel_name: diff --git a/bot/bot_base.py b/bot/bot_base.py new file mode 100644 index 00000000..52b98f2c --- /dev/null +++ b/bot/bot_base.py @@ -0,0 +1,53 @@ +# -*- coding: utf-8 -*- +from __future__ import annotations + +import asyncio +from abc import ABC +from dataclasses import dataclass +from typing import Any + +import discord +from channel_handlers import ChannelHandlerBase +from loguru import logger +from message_templates import MessageTemplates + + +@dataclass +class ReplyHandlerInfo: + msg_id: int + handler_task: asyncio.Task + handler: ChannelHandlerBase + + +class BotBase(ABC): + bot_channel_name: str + debug: bool + client: discord.Client + loop: asyncio.BaseEventLoop + owner_id: int + bot_channel: discord.TextChannel + templates: MessageTemplates + reply_handlers: dict[int, ReplyHandlerInfo] + + def __init__(self): + self.reply_handlers = {} # handlers by msg_id + + def ensure_bot_channel(self) -> None: + if self.bot_channel is None: + raise RuntimeError(f"bot channel '{self.bot_channel_name}' not found") + + async def post(self, content: str, view: discord.ui.View = None) -> discord.Message: + self.ensure_bot_channel() + return await self.bot_channel.send(content=content, view=view) + + async def post_template(self, name: str, *, view: discord.ui.View = None, **kwargs: Any) -> discord.Message: + logger.debug(f"rendering {name}") + text = self.templates.render(name, **kwargs) + return await self.post(text, view) + + def register_reply_handler(self, msg_id: int, handler: ChannelHandlerBase): + if msg_id in self.reply_handlers: + raise RuntimeError(f"Handler already registered for msg_id: {msg_id}") + task = asyncio.create_task(coro=handler.handler_loop(), name=f"reply_handler(msg_id={msg_id})") + task.add_done_callback(lambda t: handler.on_completed()) + self.reply_handlers[msg_id] = ReplyHandlerInfo(msg_id=msg_id, handler_task=task, handler=handler) diff --git a/bot/channel_handlers.py b/bot/channel_handlers.py new file mode 100644 index 00000000..74d88414 --- /dev/null +++ b/bot/channel_handlers.py @@ -0,0 +1,83 @@ +# -*- coding: utf-8 -*- +import asyncio +from abc import ABC, abstractmethod +from datetime import datetime + +import discord +from loguru import logger + + +class ChannelExpiredException(Exception): + pass + + +class ChannelHandlerBase(ABC): + queue: asyncio.Queue + completed: bool + expiry_date: datetime + expired: bool + + def __init__(self, *, expiry_date: datetime = None): + self.expiry_date = expiry_date + self.expired = False + self.queue = asyncio.Queue() + self.completed = False + + async def read(self) -> discord.Message: + """Call this method to read the next message from the user in the handler method.""" + msg = await self.queue.get() + if msg is None and self.expired: + raise ChannelExpiredException() + return msg + + def on_reply(self, message: discord.Message) -> None: + self.queue.put_nowait(message) + + def on_expire(self) -> None: + logger.info("ChannelHandler: on_expire") + self.expired = True + self.queue.put_nowait(None) + + def on_completed(self) -> None: + logger.info("ChannelHandler: on_completed") + self.completed = True + + def tick(self, now: datetime): + if now > self.expiry_date and not self.expired: + self.on_expire() + + @abstractmethod + async def handler_loop(self): + ... + + async def finalize(self): + pass + + +class AutoDestructThreadHandler(ChannelHandlerBase): + first_message: discord.Message + thread: discord.Thread + + def __init__(self, *, expiry_date: datetime = None): + super().__init__(expiry_date=expiry_date) + + async def read(self) -> discord.Message: + try: + return await super().read() + except ChannelExpiredException: + await self.cleanup() + raise + + async def cleanup(self): + if self.thread: + logger.debug(f"[expired] deleting thread: {self.thread.name}") + await self.thread.delete() + self.thread = None + if self.first_message: + logger.debug(f"[expired] deleting first_message: {self.first_message.content}") + await self.first_message.delete() + self.first_message = None + + async def finalize(self): + await self.cleanup() + return await super().finalize() diff --git a/bot/message_templates.py b/bot/message_templates.py new file mode 100644 index 00000000..df3ef1ac --- /dev/null +++ b/bot/message_templates.py @@ -0,0 +1,18 @@ +# -*- coding: utf-8 -*- +import jinja2 +from loguru import logger + + +class MessageTemplates: + def __init__(self, template_dir="./templates"): + self.env = jinja2.Environment( + loader=jinja2.FileSystemLoader(template_dir), + autoescape=jinja2.select_autoescape(disabled_extensions=("msg",), default=False, default_for_string=False), + ) + + def render(self, template_name, **kwargs): + template = self.env.get_template(template_name) + txt = template.render(kwargs) + logger.debug(txt) + + return txt diff --git a/bot/task_handlers.py b/bot/task_handlers.py new file mode 100644 index 00000000..a18f666d --- /dev/null +++ b/bot/task_handlers.py @@ -0,0 +1,153 @@ +# -*- coding: utf-8 -*- +from __future__ import annotations + +from abc import abstractmethod +from datetime import timedelta + +import discord +from bot_base import BotBase +from channel_handlers import AutoDestructThreadHandler +from loguru import logger +from oasst_shared.schemas import protocol as protocol_schema +from utils import utcnow + + +class Questionnaire(discord.ui.Modal, title="Questionnaire Response"): + name = discord.ui.TextInput(label="Name") + answer = discord.ui.TextInput(label="Answer", style=discord.TextStyle.paragraph) + + async def on_submit(self, interaction: discord.Interaction): + await interaction.response.send_message(f"Thanks for your response, {self.name}!", ephemeral=True) + + +class ChannelTaskBase(AutoDestructThreadHandler): + thread_name: str = "Replies" + expires_after: timedelta = timedelta(minutes=5) + + async def start(self, bot: BotBase, task: protocol_schema.Task) -> discord.Message: + self.bot = bot + self.task = task + msg = await self.send_first_message() + self.first_message = msg + self.thread = await bot.bot_channel.create_thread(message=discord.Object(msg.id), name=self.thread_name) + self.expiry_date = utcnow() + self.expires_after if self.expires_after else None + bot.register_reply_handler(msg_id=msg.id, handler=self) + return msg + + @abstractmethod + async def send_first_message(self) -> discord.message: + ... + + +class SummarizeStoryHandler(ChannelTaskBase): + task: protocol_schema.SummarizeStoryTask + thread_name: str = "Summaries" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_summarize_story.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + print("received: ", msg, type(msg)) + logger.info("on_summarize_story_reply") + await msg.add_reaction("✅") + + +class InitialPromptHandler(ChannelTaskBase): + task: protocol_schema.InitialPromptTask + thread_name: str = "Prompts" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_initial_prompt.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_initial_prompt_reply") + await msg.add_reaction("✅") + + +class UserReplyHandler(ChannelTaskBase): + task: protocol_schema.UserReplyTask + thread_name: str = "User replies" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_user_reply.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_user_reply_reply") + await msg.add_reaction("✅") + + +class AssistantReplyHandler(ChannelTaskBase): + task: protocol_schema.AssistantReplyTask + thread_name: str = "Assistant replies" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_assistant_reply.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_assistant_reply_reply") + await msg.add_reaction("✅") + + +class RankInitialPromptsHandler(ChannelTaskBase): + task: protocol_schema.RankInitialPromptsTask + thread_name: str = "User Responses" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_rank_initial_prompts.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_rank_initial_prompts_reply") + await msg.add_reaction("✅") + + +class RankConversationsHandler(ChannelTaskBase): + task: protocol_schema.RankConversationRepliesTask + thread_name: str = "Rankings" + + async def send_first_message(self) -> discord.message: + return await self.bot.post_template("task_rank_conversation_replies.msg", task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_rank_conversation_reply") + await msg.add_reaction("✅") + + +class RatingButton(discord.ui.Button): + def __init__(self, label, value, response_handler): + super().__init__(label=label, style=discord.ButtonStyle.green) + self.value = value + self.response_handler = response_handler + + async def callback(self, interaction): + await self.response_handler(self.value, interaction) + + +def generate_rating_view(lo: int, hi: int, response_handler) -> discord.ui.View: + view = discord.ui.View() + for i in range(lo, hi + 1): + view.add_item(RatingButton(str(i), i, response_handler)) + return view + + +class RateSummaryHandler(ChannelTaskBase): + task: protocol_schema.RateSummaryTask + thread_name: str = "Rate" + + async def send_first_message(self) -> discord.message: + async def rating_response_handler(score, interaction: discord.Interaction): + logger.info("rating_response_handler", score) + await interaction.response.send_message(f"got your feedback: {score}") + + view = generate_rating_view(self.task.scale.min, self.task.scale.max, rating_response_handler) + return await self.bot.post_template("task_rate_summary.msg", view=view, task=self.task) + + async def handler_loop(self): + msg = await self.read() + logger.info("on_rate_summary_reply") + await msg.add_reaction("✅") diff --git a/bot/templates/boot.msg b/bot/templates/boot.msg index 0561c8be..a3629715 100644 --- a/bot/templates/boot.msg +++ b/bot/templates/boot.msg @@ -10,4 +10,4 @@ ________ __ git hash: {{git_hash}} debug_mode: {{debug}} ``` -https://github.com/LAION-AI/Open-Assistant \ No newline at end of file +https://github.com/LAION-AI/Open-Assistant diff --git a/bot/templates/task_initial_prompt.msg b/bot/templates/task_initial_prompt.msg index dc3b10d3..47cf0f45 100644 --- a/bot/templates/task_initial_prompt.msg +++ b/bot/templates/task_initial_prompt.msg @@ -1,4 +1,4 @@ Please provide an initial prompt to the assistant. -{% if task.hint %} -Hint: {task.hint}" +{% if task.hint is not none %} +Hint: {{task.hint}} {% endif %} \ No newline at end of file diff --git a/bot/templates/task_summarize_story.msg b/bot/templates/task_summarize_story.msg new file mode 100644 index 00000000..24753841 --- /dev/null +++ b/bot/templates/task_summarize_story.msg @@ -0,0 +1,2 @@ +Summarize to the following story: +{{task.story}} diff --git a/pyproject.toml b/pyproject.toml index 30541eec..83b614a2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,4 +11,3 @@ line_length = 120 [tool.black] line-length = 120 target-version = ['py310'] -exclude = ["bot/templates"]