mirror of
https://github.com/katanemo/plano.git
synced 2026-07-23 16:51:04 +02:00
add warning if first/last message is not from user
This commit is contained in:
parent
a39ef5f215
commit
c80a04eb39
1 changed files with 20 additions and 28 deletions
|
|
@ -3,7 +3,7 @@ use common::{
|
||||||
consts::{SYSTEM_ROLE, USER_ROLE},
|
consts::{SYSTEM_ROLE, USER_ROLE},
|
||||||
};
|
};
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use tracing::debug;
|
use tracing::{debug, warn};
|
||||||
|
|
||||||
use super::router_model::{RouterModel, RoutingModelError};
|
use super::router_model::{RouterModel, RoutingModelError};
|
||||||
|
|
||||||
|
|
@ -61,17 +61,12 @@ impl RouterModel for RouterModelV1 {
|
||||||
let messages_vec = messages
|
let messages_vec = messages
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|m| m.role != SYSTEM_ROLE)
|
.filter(|m| m.role != SYSTEM_ROLE)
|
||||||
// .map(|m| {
|
|
||||||
// let content_json_str = serde_json::to_string(&m.content).unwrap_or_default();
|
|
||||||
// format!("{}: {}", m.role, content_json_str)
|
|
||||||
// })
|
|
||||||
// .collect::<Vec<String>>();
|
|
||||||
.collect::<Vec<&Message>>();
|
.collect::<Vec<&Message>>();
|
||||||
|
|
||||||
// Following code is to ensure that the conversation does not exceed max token length
|
// Following code is to ensure that the conversation does not exceed max token length
|
||||||
// Note: we use a simple heuristic to estimate token count based on character length to optimize for performance
|
// Note: we use a simple heuristic to estimate token count based on character length to optimize for performance
|
||||||
let mut token_count = ARCH_ROUTER_V1_SYSTEM_PROMPT.len() / TOKEN_LENGTH_DIVISOR;
|
let mut token_count = ARCH_ROUTER_V1_SYSTEM_PROMPT.len() / TOKEN_LENGTH_DIVISOR;
|
||||||
let mut selected_messages_list: Vec<&Message> = vec![];
|
let mut selected_messages_list_reversed: Vec<&Message> = vec![];
|
||||||
for (selected_messsage_count, message) in messages_vec.iter().rev().enumerate() {
|
for (selected_messsage_count, message) in messages_vec.iter().rev().enumerate() {
|
||||||
let message_token_count = message
|
let message_token_count = message
|
||||||
.content
|
.content
|
||||||
|
|
@ -91,41 +86,38 @@ impl RouterModel for RouterModelV1 {
|
||||||
);
|
);
|
||||||
if message.role == USER_ROLE {
|
if message.role == USER_ROLE {
|
||||||
// If message that exceeds max token length is from user, we need to keep it
|
// If message that exceeds max token length is from user, we need to keep it
|
||||||
selected_messages_list.push(message);
|
selected_messages_list_reversed.push(message);
|
||||||
}
|
}
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
// If we are here, it means that the message is within the max token length
|
// If we are here, it means that the message is within the max token length
|
||||||
selected_messages_list.push(message);
|
selected_messages_list_reversed.push(message);
|
||||||
}
|
}
|
||||||
|
|
||||||
if selected_messages_list.is_empty() {
|
if selected_messages_list_reversed.is_empty() {
|
||||||
debug!(
|
debug!(
|
||||||
"RouterModelV1: no messages selected, using the last message in the conversation"
|
"RouterModelV1: no messages selected, using the last message in the conversation"
|
||||||
);
|
);
|
||||||
if let Some(last_message) = messages_vec.last() {
|
if let Some(last_message) = messages_vec.last() {
|
||||||
selected_messages_list.push(last_message);
|
selected_messages_list_reversed.push(last_message);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// if selected_messsage_count == 0 {
|
// ensure that first and last selected message is from user
|
||||||
// debug!("RouterModelV1: most recent message in conversation history exceeds max token length {}, keeping only the last message (even if it exceeds max token length)",
|
if let Some(first_message) = selected_messages_list_reversed.first() {
|
||||||
// self.max_token_length);
|
if first_message.role != USER_ROLE {
|
||||||
// messages_vec = messages_vec
|
warn!("RouterModelV1: last message in the conversation is not from user, this may lead to incorrect routing");
|
||||||
// .last()
|
}
|
||||||
// .map_or_else(Vec::new, |last_message| vec![last_message.to_string()]);
|
}
|
||||||
// } else {
|
if let Some(last_message) = selected_messages_list_reversed.last() {
|
||||||
// let skip_messages_count = messages_vec.len() - selected_messsage_count;
|
if last_message.role != USER_ROLE {
|
||||||
// if skip_messages_count > 0 {
|
warn!("RouterModelV1: first message in the conversation is not from user, this may lead to incorrect routing");
|
||||||
// debug!(
|
}
|
||||||
// "RouterModelV1: skipping first {} messages from the beginning of the conversation",
|
}
|
||||||
// skip_messages_count
|
|
||||||
// );
|
|
||||||
// messages_vec = messages_vec.into_iter().skip(skip_messages_count).collect();
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
|
|
||||||
let selected_conversation_list_str = selected_messages_list
|
// Reverse the selected messages to maintain the conversation order
|
||||||
|
|
||||||
|
let selected_conversation_list_str = selected_messages_list_reversed
|
||||||
.iter()
|
.iter()
|
||||||
.rev()
|
.rev()
|
||||||
.map(|m| {
|
.map(|m| {
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue