mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-08-05 12:40:38 +08:00
add channel handler async msg routing
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
exclude: "build|stubs|bot/templates/.*msg"
|
||||
exclude: "build|stubs|^bot/templates/"
|
||||
|
||||
default_language_version:
|
||||
python: python3
|
||||
|
||||
+76
-182
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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("✅")
|
||||
@@ -10,4 +10,4 @@ ________ __
|
||||
git hash: {{git_hash}}
|
||||
debug_mode: {{debug}}
|
||||
```
|
||||
https://github.com/LAION-AI/Open-Assistant
|
||||
https://github.com/LAION-AI/Open-Assistant
|
||||
|
||||
@@ -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 %}
|
||||
@@ -0,0 +1,2 @@
|
||||
Summarize to the following story:
|
||||
{{task.story}}
|
||||
@@ -11,4 +11,3 @@ line_length = 120
|
||||
[tool.black]
|
||||
line-length = 120
|
||||
target-version = ['py310']
|
||||
exclude = ["bot/templates"]
|
||||
|
||||
Reference in New Issue
Block a user