From ba884c9fc65d820f8603859acafd055b98897b56 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 17 Jun 2024 13:19:00 +0800 Subject: feat: add `.continue` repl command (#608) --- src/config/input.rs | 42 ++++++++++++++++++++++++++++++++---------- 1 file changed, 32 insertions(+), 10 deletions(-) (limited to 'src/config/input.rs') diff --git a/src/config/input.rs b/src/config/input.rs index 48b0359..350f74c 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -30,27 +30,31 @@ pub struct Input { config: GlobalConfig, text: String, patched_text: Option, + continue_output: Option, medias: Vec, data_urls: HashMap, tool_call: Option, rag_name: Option, role: Role, with_session: bool, + with_bot: bool, } impl Input { pub fn from_str(config: &GlobalConfig, text: &str, role: Option) -> Self { - let (role, with_session) = resolve_role(&config.read(), role); + let (role, with_session, with_bot) = resolve_role(&config.read(), role); Self { config: config.clone(), text: text.to_string(), patched_text: None, + continue_output: None, medias: Default::default(), data_urls: Default::default(), tool_call: None, rag_name: None, role, with_session, + with_bot, } } @@ -96,17 +100,19 @@ impl Input { } } - let (role, session) = resolve_role(&config.read(), role); + let (role, with_session, with_bot) = resolve_role(&config.read(), role); Ok(Self { config: config.clone(), text: texts.join("\n"), patched_text: None, + continue_output: None, medias, data_urls, tool_call: Default::default(), rag_name: None, role, - with_session: session, + with_session, + with_bot, }) } @@ -129,6 +135,18 @@ impl Input { self.text = text; } + pub fn continue_output(&self) -> Option<&str> { + self.continue_output.as_deref() + } + + pub fn set_continue_output(&mut self, output: &str) { + let output = match &self.continue_output { + Some(v) => format!("{v}{output}"), + None => output.to_string(), + }; + self.continue_output = Some(output); + } + pub async fn use_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> { if self.text.is_empty() { return Ok(()); @@ -150,10 +168,6 @@ impl Input { self.rag_name.as_deref() } - pub fn clear_patch_text(&mut self) { - self.patched_text.take(); - } - pub fn merge_tool_call(mut self, output: String, tool_call_results: Vec) -> Self { match self.tool_call.as_mut() { Some(exist_tool_call_results) => { @@ -234,6 +248,10 @@ impl Input { } } + pub fn with_bot(&self) -> bool { + self.with_bot + } + pub fn summary(&self) -> String { let text: String = self .text @@ -296,10 +314,14 @@ impl Input { } } -fn resolve_role(config: &Config, role: Option) -> (Role, bool) { +fn resolve_role(config: &Config, role: Option) -> (Role, bool, bool) { match role { - Some(v) => (v, false), - None => (config.extract_role(), config.session.is_some()), + Some(v) => (v, false, false), + None => ( + config.extract_role(), + config.session.is_some(), + config.bot.is_some(), + ), } } -- cgit v1.2.3