summaryrefslogtreecommitdiffstats
path: root/src/config/session.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-19 18:10:07 +0800
committerGitHub <noreply@github.com>2024-11-19 18:10:07 +0800
commit8c7e58ad1d76a8e74f73d0a95fa12ffcda400e9f (patch)
treeb3f84b1c87063b646a3641e3f40cdbf8139abd2c /src/config/session.rs
parentbb1c34d0b84e342232c4e0b7f3f0c1fb3a9334c0 (diff)
downloadaichat-8c7e58ad1d76a8e74f73d0a95fa12ffcda400e9f.tar.gz
feat: support autonaming session (#1001)
Diffstat (limited to 'src/config/session.rs')
-rw-r--r--src/config/session.rs133
1 files changed, 110 insertions, 23 deletions
diff --git a/src/config/session.rs b/src/config/session.rs
index 633f621..23658fd 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -5,6 +5,7 @@ use crate::client::{Message, MessageContent, MessageRole};
use crate::render::MarkdownRender;
use anyhow::{bail, Context, Result};
+use fancy_regex::Regex;
use inquire::{validator::Validation, Confirm, Text};
use serde::{Deserialize, Serialize};
use serde_json::json;
@@ -12,6 +13,10 @@ use std::collections::HashMap;
use std::fs::{read_to_string, write};
use std::path::Path;
+lazy_static::lazy_static! {
+ static ref RE_AUTONAME_PREFIX: Regex = Regex::new(r"\d{8}T\d{6}-").unwrap();
+}
+
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Session {
#[serde(rename(serialize = "model", deserialize = "model"))]
@@ -32,12 +37,11 @@ pub struct Session {
#[serde(default, skip_serializing_if = "IndexMap::is_empty")]
agent_variables: IndexMap<String, String>,
- #[serde(default, skip_serializing_if = "HashMap::is_empty")]
- data_urls: HashMap<String, String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
compressed_messages: Vec<Message>,
-
messages: Vec<Message>,
+ #[serde(default, skip_serializing_if = "HashMap::is_empty")]
+ data_urls: HashMap<String, String>,
#[serde(skip)]
model: Model,
@@ -50,9 +54,11 @@ pub struct Session {
#[serde(skip)]
dirty: bool,
#[serde(skip)]
- append_conversation: bool,
+ save_session_this_time: bool,
#[serde(skip)]
compressing: bool,
+ #[serde(skip)]
+ autoname: Option<AutoName>,
}
impl Session {
@@ -75,8 +81,17 @@ impl Session {
serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?;
session.model = Model::retrieve_chat(config, &session.model_id)?;
- session.name = name.to_string();
- session.path = Some(path.display().to_string());
+
+ if let Some(autoname) = name.strip_prefix("_/") {
+ session.name = TEMP_SESSION_NAME.to_string();
+ session.path = None;
+ if let Ok(true) = RE_AUTONAME_PREFIX.is_match(autoname) {
+ session.autoname = Some(AutoName::new(autoname[16..].to_string()));
+ }
+ } else {
+ session.name = name.to_string();
+ session.path = Some(path.display().to_string());
+ }
if let Some(role_name) = &session.role_name {
if let Ok(role) = config.retrieve_role(role_name) {
@@ -103,19 +118,10 @@ impl Session {
self.dirty
}
- pub fn compressing(&self) -> bool {
- self.compressing
- }
-
pub fn save_session(&self) -> Option<bool> {
self.save_session
}
- pub fn need_compress(&self, global_compress_threshold: usize) -> bool {
- let threshold = self.compress_threshold.unwrap_or(global_compress_threshold);
- threshold > 0 && self.tokens() > threshold
- }
-
pub fn tokens(&self) -> usize {
self.model().total_tokens(&self.messages)
}
@@ -171,6 +177,10 @@ impl Session {
items.push(("path", path.to_string()));
}
+ if let Some(autoname) = self.autoname() {
+ items.push(("autoname", autoname.to_string()));
+ }
+
items.push(("model", self.model().id()));
if let Some(temperature) = self.temperature() {
@@ -279,17 +289,14 @@ impl Session {
}
pub fn set_save_session(&mut self, value: Option<bool>) {
- if self.name == TEMP_SESSION_NAME {
- return;
- }
if self.save_session != value {
self.save_session = value;
self.dirty = true;
}
}
- pub fn set_append_conversation(&mut self) {
- self.append_conversation = true;
+ pub fn set_save_session_this_time(&mut self) {
+ self.save_session_this_time = true;
}
pub fn set_compress_threshold(&mut self, value: Option<usize>) {
@@ -299,6 +306,21 @@ impl Session {
}
}
+ pub fn need_compress(&self, global_compress_threshold: usize) -> bool {
+ if self.compressing {
+ return false;
+ }
+ let threshold = self.compress_threshold.unwrap_or(global_compress_threshold);
+ if threshold < 1 {
+ return false;
+ }
+ self.tokens() > threshold
+ }
+
+ pub fn compressing(&self) -> bool {
+ self.compressing
+ }
+
pub fn set_compressing(&mut self, compressing: bool) {
self.compressing = compressing;
}
@@ -323,12 +345,39 @@ impl Session {
self.dirty = true;
}
+ pub fn need_autoname(&self) -> bool {
+ self.autoname.as_ref().map(|v| v.need()).unwrap_or_default()
+ }
+
+ pub fn set_autonaming(&mut self, naming: bool) {
+ if let Some(v) = self.autoname.as_mut() {
+ v.naming = naming;
+ }
+ }
+
+ pub fn chat_history_for_autonaming(&self) -> Option<String> {
+ self.autoname.as_ref().and_then(|v| v.chat_history.clone())
+ }
+
+ pub fn autoname(&self) -> Option<&str> {
+ self.autoname.as_ref().and_then(|v| v.name.as_deref())
+ }
+
+ pub fn set_autoname(&mut self, value: &str) {
+ let name = value
+ .chars()
+ .map(|v| if v.is_alphanumeric() { v } else { '-' })
+ .collect();
+ self.autoname = Some(AutoName::new(name));
+ }
+
pub fn exit(&mut self, session_dir: &Path, is_repl: bool) -> Result<()> {
let mut save_session = self.save_session();
- if self.append_conversation {
+ if self.save_session_this_time {
save_session = Some(true);
}
if self.dirty && save_session != Some(false) {
+ let mut session_dir = session_dir.to_path_buf();
let mut session_name = self.name().to_string();
if save_session.is_none() {
if !is_repl {
@@ -353,9 +402,16 @@ impl Session {
.prompt()?;
}
} else if save_session == Some(true) && session_name == TEMP_SESSION_NAME {
+ session_dir = session_dir.join("_");
+ ensure_parent_exists(&session_dir).with_context(|| {
+ format!("Failed to create directory '{}'", session_dir.display())
+ })?;
+
let now = chrono::Local::now();
- let formatted_time = now.format("%Y%m%dT%H:%M:%S").to_string();
- session_name = format!("{TEMP_SESSION_NAME}-{formatted_time}");
+ session_name = now.format("%Y%m%dT%H%M%S").to_string();
+ if let Some(autoname) = self.autoname() {
+ session_name = format!("{session_name}-{autoname}")
+ }
}
let session_path = session_dir.join(format!("{session_name}.yaml"));
self.save(&session_name, &session_path, is_repl)?;
@@ -413,6 +469,11 @@ impl Session {
}
} else {
if self.messages.is_empty() {
+ if self.name == TEMP_SESSION_NAME && self.save_session == Some(true) {
+ let raw_input = input.raw();
+ let chat_history = format!("USER: {raw_input}\nASSISTANT: {output}\n");
+ self.autoname = Some(AutoName::new_from_chat_history(chat_history));
+ }
self.messages.extend(input.role().build_messages(input));
} else {
self.messages
@@ -438,6 +499,7 @@ impl Session {
self.messages.clear();
self.compressed_messages.clear();
self.data_urls.clear();
+ self.autoname = None;
self.dirty = true;
}
@@ -532,3 +594,28 @@ impl RoleLike for Session {
}
}
}
+
+#[derive(Debug, Clone, Default)]
+struct AutoName {
+ naming: bool,
+ chat_history: Option<String>,
+ name: Option<String>,
+}
+
+impl AutoName {
+ pub fn new(name: String) -> Self {
+ Self {
+ name: Some(name),
+ ..Default::default()
+ }
+ }
+ pub fn new_from_chat_history(chat_history: String) -> Self {
+ Self {
+ chat_history: Some(chat_history),
+ ..Default::default()
+ }
+ }
+ pub fn need(&self) -> bool {
+ !self.naming && self.chat_history.is_some() && self.name.is_none()
+ }
+}