Skip to content

Commit 9d42455

Browse files
committed
Refactor RPBot to replace AI class with ModelsToolkit for handling model interactions and streamline message processing in MessageHandler
1 parent 6f1d90d commit 9d42455

4 files changed

Lines changed: 47 additions & 126 deletions

File tree

bot/rp_bot/ai_agent/ai.py

Lines changed: 0 additions & 106 deletions
This file was deleted.

bot/rp_bot/bot.py

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from typing import List, Union, Type
22
from logging import Logger
33
from omnimodkit.ai_config import AIConfig
4+
from omnimodkit.models_toolkit import ModelsToolkit
45
from ..models.base_bot import BaseBot
56
from .rp_bot_handlers import (
67
RPBotCommandHandler,
@@ -47,11 +48,9 @@ def get_rp_bot(
4748
default_language=bot_config.default_language,
4849
)
4950
prompt_manager = PromptManager(db=db)
50-
ai = AI(
51+
models_toolkit = ModelsToolkit(
5152
openai_api_key=openai_api_key,
5253
ai_config=ai_config,
53-
prompt_manager=prompt_manager,
54-
db=db,
5554
)
5655
auth = Auth(
5756
allowed_handles=allowed_handles,
@@ -60,7 +59,7 @@ def get_rp_bot(
6059
)
6160
return RPBot(
6261
db=db,
63-
ai=ai,
62+
models_toolkit=models_toolkit,
6463
localizer=localizer,
6564
prompt_manager=prompt_manager,
6665
auth=auth,
@@ -73,15 +72,15 @@ class RPBot(BaseBot):
7372
def __init__(
7473
self,
7574
db: DB,
76-
ai: AI,
75+
models_toolkit: ModelsToolkit,
7776
localizer: Localizer,
7877
prompt_manager: PromptManager,
7978
auth: Auth,
8079
bot_config: BotConfig,
8180
logger: Logger,
8281
):
8382
self.db = db
84-
self.ai = ai
83+
self.models_toolkit = models_toolkit
8584
self.localizer = localizer
8685
self.prompt_manager = prompt_manager
8786
self.auth = auth
@@ -101,7 +100,7 @@ def _init_handlers(
101100
return [
102101
handler(
103102
db=self.db,
104-
ai=self.ai,
103+
models_toolkit=self.models_toolkit,
105104
localizer=self.localizer,
106105
prompt_manager=self.prompt_manager,
107106
auth=self.auth,

bot/rp_bot/messages/message_handler.py

Lines changed: 34 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -11,23 +11,44 @@
1111
class MessageHandler(RPBotMessageHandler):
1212
permission_classes = (AllowedUser, BotAdmin, NotBanned)
1313

14+
async def estimate_price(self, message: Message) -> float:
15+
"""
16+
Estimates the price of the message
17+
"""
18+
return self.models_toolkit.estimate_price(
19+
input_text=message.message_text,
20+
input_image=message.in_file_image,
21+
input_audio=message.in_file_audio,
22+
)
23+
1424
async def _get_user_usage(
1525
self, input_message: Message, generated_message: str
1626
) -> int:
17-
return self.ai.get_price(
18-
message=input_message, generated_message=generated_message
27+
return self.models_toolkit.get_price(
28+
input_text=input_message.message_text,
29+
output_text=generated_message,
30+
input_image=input_message.in_file_image,
31+
input_audio=input_message.in_file_audio,
1932
)
2033

2134
async def _get_transcribed_message(self, message: Message) -> TranscribedMessage:
2235
# Note that here the responsibility to pass NULL images and Audio is on the
2336
# outer level bot processing (TG bot or other bot)
2437
image_description = (
25-
await self.ai.describe_image(message.in_file_image)
38+
str(
39+
await self.models_toolkit.vision_model.arun_default(
40+
message.in_file_image
41+
)
42+
)
2643
if message.in_file_image
2744
else None
2845
)
2946
voice_description = (
30-
await self.ai.transcribe_audio(message.in_file_audio)
47+
str(
48+
await self.models_toolkit.audio_recognition_model.arun_default(
49+
message.in_file_audio
50+
)
51+
)
3152
if message.in_file_audio
3253
else None
3354
)
@@ -66,7 +87,12 @@ async def prepare_get_reply(
6687
autoengage_state = await self.db.chats.get_autoengage_state(context)
6788
engage_is_needed = False
6889
if autoengage_state:
69-
engage_is_needed = await self.ai.engage_is_needed(message)
90+
prompt = await self.prompt_manager.compose_engage_needed_prompt(
91+
message.message_text
92+
)
93+
engage_is_needed = (
94+
await self.models_toolkit.text_model.async_ask_yes_no_question(prompt)
95+
)
7096
if not context.is_bot_mentioned and not engage_is_needed:
7197
self.logger.info(
7298
f"Saving the message from {person.user_handle} in chat {context.chat_id} "
@@ -83,7 +109,7 @@ async def prepare_get_reply(
83109
async def is_usage_under_limit(
84110
self, person: Person, context: Context, transcribed_message: TranscribedMessage
85111
) -> bool:
86-
estimated_usage = await self.ai.estimate_price(transcribed_message)
112+
estimated_usage = await self.estimate_price(transcribed_message)
87113
user_usage = await self.db.user_usage.get_user_usage(person)
88114
user_limit = await self.db.user_usage.get_user_usage_limit(person)
89115
return user_usage + estimated_usage < user_limit
@@ -150,7 +176,7 @@ async def get_response(
150176
context,
151177
message,
152178
self.db,
153-
self.ai.models_toolkit,
179+
self.models_toolkit,
154180
self.prompt_manager,
155181
)
156182
response_message = await ai_agent.get_reply(prompt, system_prompt)
@@ -181,7 +207,7 @@ async def stream_get_response(
181207
context,
182208
message,
183209
self.db,
184-
self.ai.models_toolkit,
210+
self.models_toolkit,
185211
self.prompt_manager,
186212
)
187213
async for response_message_chunk in ai_agent.get_streaming_reply(

bot/rp_bot/rp_bot_handlers.py

Lines changed: 7 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,7 @@
11
from abc import ABC, abstractmethod
22
from typing import List, Optional
3-
3+
from omnimodkit.models_toolkit import ModelsToolkit
44
from .db import DB
5-
from ai_agent.ai import AI
65
from .prompt_manager import PromptManager
76
from .auth import Auth
87
from .localizer import Localizer
@@ -20,7 +19,7 @@ class RPBotHandlerMixin(ABC):
2019
def __init__(
2120
self,
2221
db: DB,
23-
ai: AI,
22+
models_toolkit: ModelsToolkit,
2423
localizer: Localizer,
2524
prompt_manager: PromptManager,
2625
auth: Auth,
@@ -29,7 +28,7 @@ def __init__(
2928
**kwargs
3029
):
3130
self.db = db
32-
self.ai = ai
31+
self.models_toolkit = models_toolkit
3332
self.localizer = localizer
3433
self.prompt_manager = prompt_manager
3534
self.auth = auth
@@ -48,7 +47,10 @@ def get_initialized_permissions(self) -> List[BasePermission]:
4847
return result
4948

5049
async def get_localized_text(
51-
self, text: str, kwargs: Optional[dict] = None, context: Optional[Context] = None
50+
self,
51+
text: str,
52+
kwargs: Optional[dict] = None,
53+
context: Optional[Context] = None,
5254
) -> Optional[str]:
5355
return await self.localizer.get_command_response(
5456
text=text, kwargs=kwargs, context=context

0 commit comments

Comments
 (0)