summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-10-05 18:39:42 +0800
committerGitHub <noreply@github.com>2024-10-05 18:39:42 +0800
commitea4c2131e40be1124110fe4203bd018a8018760c (patch)
tree6f52dcdfb06b34401651f8dbd815bed2338d17ec /src/config
parent11a706a68d1a1a22817f2264bdbbbe7036897378 (diff)
downloadaichat-ea4c2131e40be1124110fe4203bd018a8018760c.tar.gz
feat: add `.compress session` REPL command (#907)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/mod.rs24
-rw-r--r--src/config/session.rs9
2 files changed, 25 insertions, 8 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 3243c26..97cd302 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1133,11 +1133,28 @@ impl Config {
false
}
- pub fn compress_session(&mut self, summary: &str) {
- if let Some(session) = self.session.as_mut() {
- let summary_prompt = self.summary_prompt.as_deref().unwrap_or(SUMMARY_PROMPT);
+ pub async fn compress_session(config: &GlobalConfig) -> Result<()> {
+ match config.read().session.as_ref() {
+ Some(session) => {
+ if !session.has_user_messages() {
+ bail!("No need to compress since there are no messages in the session")
+ }
+ }
+ None => bail!("No session"),
+ }
+ let input = Input::from_str(config, config.read().summarize_prompt(), None);
+ let client = input.create_client()?;
+ let summary = client.chat_completions(input).await?.text;
+ let summary_prompt = config
+ .read()
+ .summary_prompt
+ .clone()
+ .unwrap_or_else(|| SUMMARY_PROMPT.into());
+ if let Some(session) = config.write().session.as_mut() {
session.compress(format!("{}{}", summary_prompt, summary));
}
+ config.write().last_message = None;
+ Ok(())
}
pub fn summarize_prompt(&self) -> &str {
@@ -1155,7 +1172,6 @@ impl Config {
if let Some(session) = self.session.as_mut() {
session.set_compressing(false);
}
- self.last_message = None;
}
pub async fn use_rag(
diff --git a/src/config/session.rs b/src/config/session.rs
index 0bb55e8..4bda347 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -110,6 +110,10 @@ impl Session {
self.model().total_tokens(&self.messages)
}
+ pub fn has_user_messages(&self) -> bool {
+ self.messages.iter().any(|v| v.role.is_user())
+ }
+
pub fn user_messages_len(&self) -> usize {
self.messages.iter().filter(|v| v.role.is_user()).count()
}
@@ -372,12 +376,9 @@ impl Session {
}
}
} else {
- let mut need_add_msg = true;
if self.messages.is_empty() {
self.messages.extend(input.role().build_messages(input));
- need_add_msg = false;
- }
- if need_add_msg {
+ } else {
self.messages
.push(Message::new(MessageRole::User, input.message_content()));
}