diff options
Diffstat (limited to 'src/config/session.rs')
| -rw-r--r-- | src/config/session.rs | 57 |
1 files changed, 44 insertions, 13 deletions
diff --git a/src/config/session.rs b/src/config/session.rs index 076fa27..e824c43 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -22,6 +22,9 @@ pub struct Session { messages: Vec<Message>, #[serde(default)] data_urls: HashMap<String, String>, + #[serde(default)] + compressed_messages: Vec<Message>, + compress_threshold: Option<usize>, #[serde(skip)] pub name: String, #[serde(skip)] @@ -29,6 +32,8 @@ pub struct Session { #[serde(skip)] pub dirty: bool, #[serde(skip)] + pub compressing: bool, + #[serde(skip)] pub role: Option<Role>, #[serde(skip)] pub model: Model, @@ -41,10 +46,13 @@ impl Session { model_id: model.id(), temperature, messages: vec![], + compressed_messages: vec![], + compress_threshold: None, data_urls: Default::default(), name: name.to_string(), path: None, dirty: false, + compressing: false, role, model, } @@ -74,6 +82,13 @@ impl Session { self.temperature } + pub fn need_compress(&self, current_compress_threshold: usize) -> bool { + let threshold = self + .compress_threshold + .unwrap_or(current_compress_threshold); + threshold >= 1000 && self.tokens() > threshold + } + pub fn tokens(&self) -> usize { self.model.total_tokens(&self.messages) } @@ -106,7 +121,7 @@ impl Session { Ok(output) } - pub fn render(&self, render: &mut MarkdownRender) -> Result<String> { + pub fn info(&self, render: &mut MarkdownRender) -> Result<String> { let mut items = vec![]; if let Some(path) = &self.path { @@ -119,6 +134,10 @@ impl Session { items.push(("temperature", temperature.to_string())); } + if let Some(compress_threshold) = self.compress_threshold { + items.push(("compress_threshold", compress_threshold.to_string())); + } + if let Some(max_tokens) = self.model.max_tokens { items.push(("max_tokens", max_tokens.to_string())); } @@ -135,7 +154,7 @@ impl Session { for message in &self.messages { match message.role { MessageRole::System => { - continue; + lines.push(render.render(&message.content.render_input(resolve_url_fn))); } MessageRole::Assistant => { if let MessageContent::Text(text) = &message.content { @@ -181,14 +200,28 @@ impl Session { self.temperature = value; } + pub fn set_compress_threshold(&mut self, value: usize) { + self.compress_threshold = Some(value); + } + pub fn set_model(&mut self, model: Model) -> Result<()> { self.model_id = model.id(); self.model = model; Ok(()) } + pub fn compress(&mut self, prompt: String) { + self.compressed_messages.append(&mut self.messages); + self.messages.push(Message { + role: MessageRole::System, + content: MessageContent::Text(prompt), + }); + self.role = None; + self.dirty = true; + } + pub fn save(&mut self, session_path: &Path) -> Result<()> { - if !self.should_save() { + if !self.dirty { return Ok(()); } self.path = Some(session_path.display().to_string()); @@ -208,10 +241,6 @@ impl Session { Ok(()) } - pub fn should_save(&self) -> bool { - !self.is_empty() && self.dirty - } - pub fn guard_save(&self) -> Result<()> { if self.path.is_none() { bail!("Not found session '{}'", self.name) @@ -258,11 +287,9 @@ impl Session { Ok(()) } - pub fn clear_messgaes(&mut self) { - if self.messages.is_empty() { - return; - } + pub fn clear_messages(&mut self) { self.messages.clear(); + self.compressed_messages.clear(); self.data_urls.clear(); self.dirty = true; } @@ -275,12 +302,16 @@ impl Session { pub fn build_emssages(&self, input: &Input) -> Vec<Message> { let mut messages = self.messages.clone(); let mut need_add_msg = true; - if messages.is_empty() { + let len = messages.len(); + if len == 0 { if let Some(role) = self.role.as_ref() { messages = role.build_messages(input); need_add_msg = false; } - }; + } else if len == 1 && self.compressed_messages.len() >= 2 { + messages + .extend(self.compressed_messages[self.compressed_messages.len() - 2..].to_vec()); + } if need_add_msg { messages.push(Message { role: MessageRole::User, |
