use crate::chat_client::openai_api::message::{
AssistantMessage, Message, SystemMessage, UserMessage,
};
use iter_accumulate::IterAccumulate;
#[derive(Default, Clone)]
pub struct Context {
system_message: Option<String>,
conversation: Vec<(String, String)>,
tokenizer: Option<tiktoken_rs::CoreBPE>,
min_history_tokens: Option<usize>,
max_history_tokens: Option<usize>,
}
impl Context {
pub fn new(system_message: Option<String>) -> Self {
Self {
system_message,
conversation: Vec::new(),
tokenizer: None,
min_history_tokens: None,
max_history_tokens: None,
}
}
pub fn new_with_rolling_window(
system_message: Option<String>,
tokenizer: tiktoken_rs::CoreBPE,
min_history_tokens: Option<usize>,
max_history_tokens: Option<usize>,
) -> Self {
debug_assert!(min_history_tokens.is_some() || max_history_tokens.is_some());
Self {
system_message,
conversation: Vec::new(),
tokenizer: Some(tokenizer),
min_history_tokens,
max_history_tokens,
}
}
pub fn with_request(&self, request: String) -> impl Iterator<Item = Message> + '_ {
self.system_message
.iter()
.map(|system_message| SystemMessage::new(system_message.clone()).into())
.chain(self.conversation.iter().flat_map(|(request, response)| {
[
UserMessage::new(request.clone()).into(),
AssistantMessage::new(response.clone()).into(),
]
.into_iter()
}))
.chain(std::iter::once(UserMessage::new(request).into()))
}
pub fn push(&mut self, request: String, response: String) {
self.conversation.push((request, response));
self.keep_recent();
}
fn keep_recent(&mut self) {
let Some(ref tokenizer) = self.tokenizer else {
return;
};
debug_assert!(self.min_history_tokens.is_some() || self.max_history_tokens.is_some());
let min_tokens = self.min_history_tokens.unwrap_or(usize::MAX);
let max_tokens = self.max_history_tokens.unwrap_or(usize::MAX);
let num_tokens = |m| tokenizer.encode_with_special_tokens(m).len();
let system_tokens = self
.system_message
.as_ref()
.map(|m| num_tokens(m))
.unwrap_or_default();
let keep = self
.conversation
.iter()
.rev()
.map(|transaction| num_tokens(&transaction.0) + num_tokens(&transaction.1))
.accumulate((0, system_tokens), |(_, acc), x| (acc, acc + x))
.map_while(|(prev, current)| (prev < min_tokens).then_some(current))
.take_while(|current| *current <= max_tokens)
.count();
let discard = self.conversation.len() - keep;
self.conversation.drain(0..discard);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty() {
let context = Context::default();
assert_eq!(
context
.with_request(String::from("req"))
.collect::<Vec<_>>(),
vec![UserMessage::new(String::from("req")).into()],
);
}
#[test]
fn non_empty() {
let mut context = Context::default();
context.push(String::from("req1"), String::from("resp1"));
assert_eq!(
context
.with_request(String::from("req2"))
.collect::<Vec<_>>(),
vec![
UserMessage::new(String::from("req1")).into(),
AssistantMessage::new(String::from("resp1")).into(),
UserMessage::new(String::from("req2")).into(),
],
);
}
#[test]
fn empty_with_system_message() {
let context = Context::new(Some(String::from("system")));
assert_eq!(
context
.with_request(String::from("req"))
.collect::<Vec<_>>(),
vec![
SystemMessage::new(String::from("system")).into(),
UserMessage::new(String::from("req")).into(),
]
);
}
#[test]
fn non_empty_with_system_message() {
let mut context = Context::new(Some(String::from("system")));
context.push(String::from("req1"), String::from("resp1"));
assert_eq!(
context
.with_request(String::from("req2"))
.collect::<Vec<_>>(),
vec![
SystemMessage::new(String::from("system")).into(),
UserMessage::new(String::from("req1")).into(),
AssistantMessage::new(String::from("resp1")).into(),
UserMessage::new(String::from("req2")).into(),
]
);
}
#[test]
fn min_history_tokens() {
let tokenizer = tiktoken_rs::o200k_base().unwrap();
let num_tokens = |m| tokenizer.encode_with_special_tokens(m).len();
let system = "to to to to to".to_string();
let request = "do do do do do".to_string();
let response = "be be be be be".to_string();
assert_eq!(num_tokens(&system), 5);
assert_eq!(num_tokens(&request), 5);
assert_eq!(num_tokens(&response), 5);
let mut context = Context::new_with_rolling_window(
Some(system.to_string()),
tokenizer.clone(),
Some(20),
None,
);
assert!(context.conversation.is_empty());
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 1);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
}
#[test]
fn min_history_tokens_exact() {
let tokenizer = tiktoken_rs::o200k_base().unwrap();
let num_tokens = |m| tokenizer.encode_with_special_tokens(m).len();
let request = "do do do do do".to_string();
let response = "be be be be be".to_string();
assert_eq!(num_tokens(&request), 5);
assert_eq!(num_tokens(&response), 5);
let mut context = Context::new_with_rolling_window(None, tokenizer.clone(), Some(20), None);
assert!(context.conversation.is_empty());
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 1);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
}
#[test]
fn max_history_tokens() {
let tokenizer = tiktoken_rs::o200k_base().unwrap();
let num_tokens = |m| tokenizer.encode_with_special_tokens(m).len();
let system = "to to to to to".to_string();
let request = "do do do do do".to_string();
let response = "be be be be be".to_string();
assert_eq!(num_tokens(&system), 5);
assert_eq!(num_tokens(&request), 5);
assert_eq!(num_tokens(&response), 5);
let mut context = Context::new_with_rolling_window(
Some(system.to_string()),
tokenizer.clone(),
None,
Some(30),
);
assert!(context.conversation.is_empty());
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 1);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
}
#[test]
fn max_history_tokens_exact() {
let tokenizer = tiktoken_rs::o200k_base().unwrap();
let num_tokens = |m| tokenizer.encode_with_special_tokens(m).len();
let request = "do do do do do".to_string();
let response = "be be be be be".to_string();
assert_eq!(num_tokens(&request), 5);
assert_eq!(num_tokens(&response), 5);
let mut context = Context::new_with_rolling_window(None, tokenizer.clone(), None, Some(30));
assert!(context.conversation.is_empty());
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 1);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 2);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 3);
context.push(request.clone(), response.clone());
assert_eq!(context.conversation.len(), 3);
}
}