Skip to content

Commit 6f1d90d

Browse files
committed
Refactor AI class to remove reply methods and integrate AIAgent for handling replies in MessageHandler
1 parent 21f0368 commit 6f1d90d

3 files changed

Lines changed: 36 additions & 17 deletions

File tree

bot/rp_bot/ai_agent/agent_tools/agent.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
from typing import AsyncIterator
12
from omnimodkit import ModelsToolkit
23
from motor.motor_asyncio import AsyncIOMotorDatabase
34
from .agent_toolkit import AIAgentToolkit
@@ -18,3 +19,19 @@ def __init__(
1819
self.toolkit = AIAgentToolkit(
1920
person, context, message, db, models_toolkit, prompt_manager
2021
)
22+
self.models_toolkit = models_toolkit
23+
24+
async def get_reply(self, user_input: str, system_prompt: str) -> str:
25+
response = await self.models_toolkit.text_model.arun_default(
26+
user_input, system_prompt
27+
)
28+
return response.text
29+
30+
async def get_streaming_reply(
31+
self, user_input: str, system_prompt: str
32+
) -> AsyncIterator[str]:
33+
async for response in self.models_toolkit.text_model.astream_default(
34+
user_input, system_prompt
35+
):
36+
if response is not None:
37+
yield response.text_chunk

bot/rp_bot/ai_agent/ai.py

Lines changed: 0 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -74,21 +74,6 @@ async def generate_image(
7474
response = await self.models_toolkit.image_generation_model.arun_default(prompt)
7575
return response.image_url
7676

77-
async def get_reply(self, user_input: str, system_prompt: str) -> str:
78-
response = await self.models_toolkit.text_model.arun_default(
79-
user_input, system_prompt
80-
)
81-
return response.text
82-
83-
async def get_streaming_reply(
84-
self, user_input: str, system_prompt: str
85-
) -> AsyncIterator[str]:
86-
async for response in self.models_toolkit.text_model.astream_default(
87-
user_input, system_prompt
88-
):
89-
if response is not None:
90-
yield response.text_chunk
91-
9277
async def get_price(
9378
self,
9479
message: Message,

bot/rp_bot/messages/message_handler.py

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from ...models.handlers_input import Person, Context, Message, TranscribedMessage
66
from ..auth import AllowedUser, BotAdmin, NotBanned
77
from ..rp_bot_handlers import RPBotMessageHandler
8+
from ..ai_agent.agent_tools.agent import AIAgent
89

910

1011
class MessageHandler(RPBotMessageHandler):
@@ -144,7 +145,15 @@ async def get_response(
144145
person, context, transcribed_message
145146
)
146147
system_prompt = self.prompt_manager.get_reply_system_prompt(context)
147-
response_message = await self.ai.get_reply(prompt, system_prompt)
148+
ai_agent = AIAgent(
149+
person,
150+
context,
151+
message,
152+
self.db,
153+
self.ai.models_toolkit,
154+
self.prompt_manager,
155+
)
156+
response_message = await ai_agent.get_reply(prompt, system_prompt)
148157
result = CommandResponse(
149158
text="message_response", kwargs={"response_text": response_message}
150159
)
@@ -167,7 +176,15 @@ async def stream_get_response(
167176
)
168177
response_message = ""
169178
system_prompt = self.prompt_manager.get_reply_system_prompt(context)
170-
async for response_message_chunk in self.ai.get_streaming_reply(
179+
ai_agent = AIAgent(
180+
person,
181+
context,
182+
message,
183+
self.db,
184+
self.ai.models_toolkit,
185+
self.prompt_manager,
186+
)
187+
async for response_message_chunk in ai_agent.get_streaming_reply(
171188
prompt, system_prompt
172189
):
173190
if not response_message_chunk:

0 commit comments

Comments
 (0)