Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
238 changes: 158 additions & 80 deletions src/chat_gpt_handler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use redis::aio::ConnectionManager;
use regex::Regex;
use serde::{Deserialize, Serialize};
use teloxide::prelude::*;
use teloxide::types::ReplyParameters;
use teloxide::types::{MessageId, ReplyParameters, ThreadId};
use teloxide::RequestError;

const FEDOR_CHAT_GPT_SYSTEM_CONTEXT: &str = "Предоставь грубый ответ. \
Expand Down Expand Up @@ -76,55 +76,37 @@ pub async fn handle_chat_gpt_question(
return Ok(());
};
info!("gpt invocation: chat_id: {chat_id}, message: {message}");
let bot_profiles = &*BOT_PROFILES;
let bot_configuration = bot_profiles
.iter()
.find(|&x| x.is_correct_config(message))
.unwrap_or(&bot_profiles[0]);
let bot_context_key = &format!("{:#?}:chat:{:#?}", bot_configuration.profile, chat_id.0);

let bot_configuration = bot_configuration_for_message(message);
let bot_context_key = format!("{:#?}:chat:{:#?}", bot_configuration.profile, chat_id.0);
let user_message = ChatMessage {
role: User,
content: message.to_string(),
};
let summary_request_regex = &*CHAT_SUMMARY_REQUEST_REGEX;
let mut redis_cm = gpt_parameters.redis_connection_manager.clone();
let context = match summary_request_regex.is_match(message) {
true => {
fetch_chat_summary_context(
&mut redis_cm,
chat_id.0,
&user_message,
bot_configuration.gpt_system_context,
)
.await
}
false => {
fetch_bot_context(
&mut redis_cm,
bot_context_key,
&user_message,
bot_configuration.gpt_system_context,
)
.await
}
};
let context = build_question_context(
&mut redis_cm,
chat_id,
&bot_context_key,
&user_message,
bot_configuration,
)
.await;

let gpt_response_message = gpt_service::chat_gpt_call(gpt_parameters, chat_id, context).await;
let bot_reply_msg_response = if let Some(thread_id) = msg.thread_id {
bot.send_message(chat_id, &gpt_response_message.content)
.reply_parameters(ReplyParameters::new(msg.id))
.message_thread_id(thread_id)
.await
} else {
bot.send_message(chat_id, &gpt_response_message.content)
.reply_parameters(ReplyParameters::new(msg.id))
.await
};
let bot_reply_msg_response = send_gpt_reply(
&bot,
chat_id,
msg.id,
msg.thread_id,
&gpt_response_message.content,
)
.await;

update_bot_context_and_identifiers(
&mut redis_cm,
bot_configuration.profile,
bot_context_key,
&bot_context_key,
&user_message,
&gpt_response_message,
bot_reply_msg_response,
Expand All @@ -133,6 +115,72 @@ pub async fn handle_chat_gpt_question(
Ok(())
}

/// Pick the bot profile whose mention regex matches the message, falling back
/// to the first profile when none match.
fn bot_configuration_for_message(message: &str) -> &'static BotConfiguration<'static> {
let bot_profiles = &*BOT_PROFILES;
bot_profiles
.iter()
.find(|&config| config.is_correct_config(message))
.unwrap_or(&bot_profiles[0])
}

/// Look up the bot profile a previous bot message was sent under, falling back
/// to the first profile when it is unknown.
fn bot_configuration_for_profile(profile: BotProfile) -> &'static BotConfiguration<'static> {
let bot_profiles = &*BOT_PROFILES;
bot_profiles
.iter()
.find(|&config| config.profile == profile)
.unwrap_or(&bot_profiles[0])
}

/// Build the GPT context for a fresh question: a chat-history summary when the
/// message asks "what's going on", otherwise the profile's rolling context.
async fn build_question_context(
redis_cm: &mut ConnectionManager,
chat_id: ChatId,
bot_context_key: &String,
user_message: &ChatMessage,
bot_configuration: &BotConfiguration<'_>,
) -> Vec<ChatMessage> {
if CHAT_SUMMARY_REQUEST_REGEX.is_match(&user_message.content) {
fetch_chat_summary_context(
redis_cm,
chat_id.0,
user_message,
bot_configuration.gpt_system_context,
)
.await
} else {
fetch_bot_context(
redis_cm,
bot_context_key,
user_message,
bot_configuration.gpt_system_context,
)
.await
}
}

/// Send a bot reply, routing it into the originating message thread when there
/// is one.
async fn send_gpt_reply(
bot: &Bot,
chat_id: ChatId,
reply_to: MessageId,
thread_id: Option<ThreadId>,
content: &str,
) -> Result<Message, RequestError> {
let request = bot
.send_message(chat_id, content.to_owned())
.reply_parameters(ReplyParameters::new(reply_to));
match thread_id {
Some(thread_id) => request.message_thread_id(thread_id).await,
None => request.await,
}
}

async fn update_bot_context_and_identifiers(
redis_connection_manager: &mut ConnectionManager,
bot_profile: BotProfile,
Expand Down Expand Up @@ -182,51 +230,48 @@ pub async fn handle_reply(
info!("chat_key: {chat_key:?}");
let reply_msg_id = reply_msg.id.0;
let mut redis_cm = gpt_parameters.redis_connection_manager.clone();
if let Ok(reply_msg_bot_profile) =
let Ok(reply_msg_bot_profile) =
chat_repository::get_bot_msg_profile(&mut redis_cm, chat_key, reply_msg_id).await
{
info!(
"handle msg of bot msg reply_msg_id:'{reply_msg_id:#?}' under bot profile:'{reply_msg_bot_profile:#?}'"
);
info!(
"truing to reply chat_id:{chat_id:#?}, msg_id: {:?}, thread_id: {:#?}",
msg.id, msg.thread_id
);
let bot_profiles = &*BOT_PROFILES;
let bot_configuration = bot_profiles
.iter()
.find(|&x| x.profile == reply_msg_bot_profile)
.unwrap_or(&bot_profiles[0]);
let bot_context_key = &format!("{:#?}:chat:{:#?}", bot_configuration.profile, chat_id.0);
let user_message = ChatMessage {
role: User,
content: message.to_string(),
};
let context = fetch_bot_context(
&mut redis_cm,
bot_context_key,
&user_message,
bot_configuration.gpt_system_context,
)
.await;
else {
return Ok(());
};
info!(
"handle msg of bot msg reply_msg_id:'{reply_msg_id:#?}' under bot profile:'{reply_msg_bot_profile:#?}'"
);
info!(
"truing to reply chat_id:{chat_id:#?}, msg_id: {:?}, thread_id: {:#?}",
msg.id, msg.thread_id
);

let gpt_response_message =
gpt_service::chat_gpt_call(gpt_parameters, chat_id, context).await;
let bot_reply_msg_response = bot
.send_message(chat_id, &gpt_response_message.content)
.reply_parameters(ReplyParameters::new(msg.id))
.await;
let bot_configuration = bot_configuration_for_profile(reply_msg_bot_profile);
let bot_context_key = format!("{:#?}:chat:{:#?}", bot_configuration.profile, chat_id.0);
let user_message = ChatMessage {
role: User,
content: message.to_string(),
};
let context = fetch_bot_context(
&mut redis_cm,
&bot_context_key,
&user_message,
bot_configuration.gpt_system_context,
)
.await;

update_bot_context_and_identifiers(
&mut redis_cm,
bot_configuration.profile,
bot_context_key,
&user_message,
&gpt_response_message,
bot_reply_msg_response,
)
let gpt_response_message = gpt_service::chat_gpt_call(gpt_parameters, chat_id, context).await;
let bot_reply_msg_response = bot
.send_message(chat_id, &gpt_response_message.content)
.reply_parameters(ReplyParameters::new(msg.id))
.await;
}

update_bot_context_and_identifiers(
&mut redis_cm,
bot_configuration.profile,
&bot_context_key,
&user_message,
&gpt_response_message,
bot_reply_msg_response,
)
.await;
Ok(())
}

Expand Down Expand Up @@ -300,6 +345,7 @@ pub enum BotProfile {

#[cfg(test)]
mod tests {
use super::{bot_configuration_for_message, bot_configuration_for_profile, BotProfile};
use crate::chat_gpt_handler::SUMMARY_REQUEST_REGEX;
use regex::Regex;

Expand All @@ -311,4 +357,36 @@ mod tests {
assert!(summary_regex.is_match("Фёдор, что тут происходит?"));
assert!(!summary_regex.is_match("Fedor, kak dela?"));
}

#[test]
fn bot_configuration_for_message_selects_by_mention() {
assert_eq!(
bot_configuration_for_message("привет fedor").profile,
BotProfile::Fedor
);
assert_eq!(
bot_configuration_for_message("а вот и феликс").profile,
BotProfile::Felix
);
assert_eq!(
bot_configuration_for_message("ferris the crab").profile,
BotProfile::Ferris
);
}

#[test]
fn bot_configuration_for_message_defaults_to_first_profile() {
// No mention keyword -> first configured profile (Fedor).
assert_eq!(
bot_configuration_for_message("nothing relevant here").profile,
BotProfile::Fedor
);
}

#[test]
fn bot_configuration_for_profile_round_trips() {
for profile in [BotProfile::Fedor, BotProfile::Felix, BotProfile::Ferris] {
assert_eq!(bot_configuration_for_profile(profile).profile, profile);
}
}
}
Loading