From 927b73665ffc8025a7d18a162165e38e3af08747 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 9 Jan 2025 08:44:05 +0800 Subject: feat: `.file` supports including last reply with `%%` (#1079) --- src/config/input.rs | 34 ++++++++++++++++++++++------ src/config/mod.rs | 64 ++++++++++++++++++++++++++++++++++++----------------- src/repl/mod.rs | 49 +++++++++++++++++++++++++++++++--------- 3 files changed, 109 insertions(+), 38 deletions(-) diff --git a/src/config/input.rs b/src/config/input.rs index 2af4c2b..70a01cf 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -21,6 +21,7 @@ pub struct Input { text: String, raw: (String, Vec), patched_text: Option, + last_reply: Option, continue_output: Option, regenerate: bool, medias: Vec, @@ -40,6 +41,7 @@ impl Input { text: text.to_string(), raw: (text.to_string(), vec![]), patched_text: None, + last_reply: None, continue_output: None, regenerate: false, medias: Default::default(), @@ -62,10 +64,15 @@ impl Input { let mut external_cmds = vec![]; let mut local_paths = vec![]; let mut remote_urls = vec![]; + let mut last_reply = None; + let mut with_last_reply = false; for path in paths { match resolve_local_path(&path) { Some(v) => { - if v.len() > 2 && v.starts_with('`') && v.ends_with('`') { + if v == "%%" { + with_last_reply = true; + raw_paths.push(v); + } else if v.len() > 2 && v.starts_with('`') && v.ends_with('`') { external_cmds.push(v[1..v.len() - 1].to_string()); raw_paths.push(v); } else { @@ -89,12 +96,24 @@ impl Input { if !raw_text.is_empty() { texts.push(raw_text.to_string()); }; - if !files.is_empty() { - texts.push(String::new()); + if with_last_reply { + if let Some(LastMessage { input, output, .. }) = config.read().last_message.as_ref() { + if !output.is_empty() { + last_reply = Some(output.clone()) + } else if let Some(v) = input.last_reply.as_ref() { + last_reply = Some(v.clone()); + } + if let Some(v) = last_reply.clone() { + texts.push(format!("\n{v}\n")); + } + } + if last_reply.is_none() && files.is_empty() && medias.is_empty() { + bail!("No last reply found"); + } } for (kind, path, contents) in files { texts.push(format!( - "============ {kind}: {path} ============\n{contents}\n" + "\n============ {kind}: {path} ============\n{contents}" )); } let (role, with_session, with_agent) = resolve_role(&config.read(), role); @@ -103,6 +122,7 @@ impl Input { text: texts.join("\n"), raw: (raw_text.to_string(), raw_paths), patched_text: None, + last_reply, continue_output: None, regenerate: false, medias, @@ -407,10 +427,10 @@ async fn load_documents( let loaders = config.read().document_loaders.clone(); for file_path in local_files { if is_image(&file_path) { - let data_url = read_media_to_data_url(&file_path) + let contents = read_media_to_data_url(&file_path) .with_context(|| format!("Unable to read media file '{file_path}'"))?; - data_urls.insert(sha256(&data_url), file_path); - medias.push(data_url) + data_urls.insert(sha256(&contents), file_path); + medias.push(contents) } else { let document = load_file(&loaders, &file_path) .await diff --git a/src/config/mod.rs b/src/config/mod.rs index 48ca390..ca0990a 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -156,7 +156,7 @@ pub struct Config { #[serde(skip)] pub working_mode: WorkingMode, #[serde(skip)] - pub last_message: Option<(Input, String)>, + pub last_message: Option, #[serde(skip)] pub cli_info_flag: bool, @@ -1051,8 +1051,15 @@ impl Config { if let Some(session) = session.as_mut() { if session.is_empty() { new_session = true; - if let Some((input, output)) = &self.last_message { - if self.agent.is_some() == input.with_agent() { + if let Some(LastMessage { + input, + output, + continuous, + }) = &self.last_message + { + if (*continuous && !output.is_empty()) + && self.agent.is_some() == input.with_agent() + { let ans = Confirm::new( "Start a session that incorporates the last question and answer?", ) @@ -1093,7 +1100,7 @@ impl Config { if let Some(mut session) = self.session.take() { let sessions_dir = self.sessions_dir(); session.exit(&sessions_dir, self.working_mode.is_repl())?; - self.last_message = None; + self.discontinuous_last_message(); } Ok(()) } @@ -1131,7 +1138,7 @@ impl Config { ) })?; self.session = Some(Session::load(self, &name, &session_path)?); - self.last_message = None; + self.discontinuous_last_message(); Ok(()) } @@ -1144,7 +1151,7 @@ impl Config { } else { bail!("No session") } - self.last_message = None; + self.discontinuous_last_message(); Ok(()) } @@ -1225,7 +1232,7 @@ impl Config { if let Some(session) = config.write().session.as_mut() { session.compress(format!("{}{}", summary_prompt, summary)); } - config.write().last_message = None; + config.write().discontinuous_last_message(); Ok(()) } @@ -1524,7 +1531,7 @@ impl Config { self.exit_session()?; if self.agent.take().is_some() { self.rag.take(); - self.last_message = None; + self.discontinuous_last_message(); self.cli_agent_variables = None; } Ok(()) @@ -1793,13 +1800,6 @@ impl Config { .collect() } - pub fn last_reply(&self) -> &str { - self.last_message - .as_ref() - .map(|(_, reply)| reply.as_str()) - .unwrap_or_default() - } - pub fn render_options(&self) -> Result { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -1939,7 +1939,7 @@ impl Config { } pub fn before_chat_completion(&mut self, input: &Input) -> Result<()> { - self.last_message = Some((input.clone(), String::new())); + self.last_message = Some(LastMessage::new(input.clone(), String::new())); Ok(()) } @@ -1949,15 +1949,22 @@ impl Config { output: &str, tool_results: &[ToolResult], ) -> Result<()> { - if self.dry_run || output.is_empty() || !tool_results.is_empty() { - self.last_message = None; + if output.is_empty() || !tool_results.is_empty() { return Ok(()); } - self.last_message = Some((input.clone(), output.to_string())); - self.save_message(input, output)?; + self.last_message = Some(LastMessage::new(input.clone(), output.to_string())); + if !self.dry_run { + self.save_message(input, output)?; + } Ok(()) } + fn discontinuous_last_message(&mut self) { + if let Some(last_message) = self.last_message.as_mut() { + last_message.continuous = false; + } + } + fn save_message(&mut self, input: &Input, output: &str) -> Result<()> { let mut input = input.clone(); input.clear_patch(); @@ -2342,6 +2349,23 @@ impl WorkingMode { } } +#[derive(Debug, Clone)] +pub struct LastMessage { + pub input: Input, + pub output: String, + pub continuous: bool, +} + +impl LastMessage { + pub fn new(input: Input, output: String) -> Self { + Self { + input, + output, + continuous: true, + } + } +} + bitflags::bitflags! { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub struct StateFlags: u32 { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 1c9ba25..136d113 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -7,7 +7,7 @@ use self::highlighter::ReplHighlighter; use self::prompt::ReplPrompt; use crate::client::{call_chat_completions, call_chat_completions_streaming}; -use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; +use crate::config::{AssertState, Config, GlobalConfig, Input, LastMessage, StateFlags}; use crate::render::render_error; use crate::utils::{ abortable_run_with_spinner, create_abort_signal, set_text, temp_file, AbortSignal, @@ -154,10 +154,10 @@ lazy_static::lazy_static! { ReplCommand::new(".continue", "Continue the response", AssertState::pass()), ReplCommand::new( ".regenerate", - "Regenerate the last response", + "Regenerate the response", AssertState::pass() ), - ReplCommand::new(".copy", "Copy the last response", AssertState::pass()), + ReplCommand::new(".copy", "Copy the last chat response", AssertState::pass()), ReplCommand::new(".set", "Adjust runtime configuration", AssertState::pass()), ReplCommand::new(".delete", "Delete roles/sessions/RAGs/agents", AssertState::pass()), ReplCommand::new(".exit", "Exit the REPL", AssertState::pass()), @@ -413,27 +413,44 @@ impl Repl { ask(&self.config, self.abort_signal.clone(), input, true).await?; } None => println!( - r#"Usage: .file ... [-- ...] + r#"Usage: .file ... [-- ...] .file /tmp/file.txt .file src/ Cargo.toml -- analyze .file https://example.com/file.txt -- summarize .file https://example.com/image.png -- recognize text +.file %% -- translate last reply to english .file `git diff` -- Generate git commit message"# ), }, ".continue" => { - let (mut input, output) = match self.config.read().last_message.clone() { + let LastMessage { + mut input, output, .. + } = match self + .config + .read() + .last_message + .as_ref() + .filter(|v| v.continuous && !v.output.is_empty()) + .cloned() + { Some(v) => v, - None => bail!("Unable to continue response"), + None => bail!("Unable to continue the response"), }; input.set_continue_output(&output); ask(&self.config, self.abort_signal.clone(), input, true).await?; } ".regenerate" => { - let (mut input, _) = match self.config.read().last_message.clone() { + let LastMessage { mut input, .. } = match self + .config + .read() + .last_message + .as_ref() + .filter(|v| v.continuous) + .cloned() + { Some(v) => v, - None => bail!("Unable to regenerate the last response"), + None => bail!("Unable to regenerate the response"), }; input.set_regenerate(); ask(&self.config, self.abort_signal.clone(), input, true).await?; @@ -455,9 +472,19 @@ impl Repl { } }, ".copy" => { - let config = self.config.read(); - self.copy(config.last_reply()) - .with_context(|| "Failed to copy the last response")?; + let output = match self + .config + .read() + .last_message + .as_ref() + .filter(|v| v.continuous && !v.output.is_empty()) + .map(|v| v.output.clone()) + { + Some(v) => v, + None => bail!("No chat response to copy"), + }; + self.copy(&output) + .with_context(|| "Failed to copy the last chat response")?; } ".exit" => match args { Some("role") => { -- cgit v1.2.3