summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-17 17:03:28 +0800
committerGitHub <noreply@github.com>2024-06-17 17:03:28 +0800
commit62b297e8bb09e257a61154afddc4017e48bd12f3 (patch)
treeb402281f625b304ad99ba392597f65d9fca0b21e /src/config
parentff284779d9757c128acaf690ddf64a87a986e413 (diff)
downloadaichat-62b297e8bb09e257a61154afddc4017e48bd12f3.tar.gz
feat: add `.regenerate` repl command (#610)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs15
-rw-r--r--src/config/mod.rs7
-rw-r--r--src/config/session.rs46
3 files changed, 46 insertions, 22 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 350f74c..4ddc06f 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -31,6 +31,7 @@ pub struct Input {
text: String,
patched_text: Option<String>,
continue_output: Option<String>,
+ regenerate: bool,
medias: Vec<String>,
data_urls: HashMap<String, String>,
tool_call: Option<ToolResults>,
@@ -48,6 +49,7 @@ impl Input {
text: text.to_string(),
patched_text: None,
continue_output: None,
+ regenerate: false,
medias: Default::default(),
data_urls: Default::default(),
tool_call: None,
@@ -106,6 +108,7 @@ impl Input {
text: texts.join("\n"),
patched_text: None,
continue_output: None,
+ regenerate: false,
medias,
data_urls,
tool_call: Default::default(),
@@ -147,6 +150,18 @@ impl Input {
self.continue_output = Some(output);
}
+ pub fn regenerate(&self) -> bool {
+ self.regenerate
+ }
+
+ pub fn set_regenerate(&mut self) {
+ let role = self.config.read().extract_role();
+ if role.name() == self.role().name() {
+ self.role = role;
+ }
+ self.regenerate = true;
+ }
+
pub async fn use_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> {
if self.text.is_empty() {
return Ok(());
diff --git a/src/config/mod.rs b/src/config/mod.rs
index a27245d..cb510ef 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -933,8 +933,11 @@ impl Config {
}
pub fn exit_bot(&mut self) -> Result<()> {
- self.rag.take();
- self.bot.take();
+ if self.bot.take().is_some() {
+ self.exit_session()?;
+ self.rag.take();
+ self.last_message = None;
+ }
Ok(())
}
diff --git a/src/config/session.rs b/src/config/session.rs
index 0b48cbd..2149ba5 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -353,30 +353,33 @@ impl Session {
}
pub fn add_message(&mut self, input: &Input, output: &str) -> Result<()> {
- match input.continue_output() {
- Some(_) => {
- if let Some(message) = self.messages.last_mut() {
- if let MessageContent::Text(text) = &mut message.content {
- *text = format!("{text}{output}");
- }
+ if input.continue_output().is_some() {
+ if let Some(message) = self.messages.last_mut() {
+ if let MessageContent::Text(text) = &mut message.content {
+ *text = format!("{text}{output}");
}
}
- None => {
- let mut need_add_msg = true;
- if self.messages.is_empty() {
- self.messages.extend(input.role().build_messages(input));
- need_add_msg = false;
+ } else if input.regenerate() {
+ if let Some(message) = self.messages.last_mut() {
+ if let MessageContent::Text(text) = &mut message.content {
+ *text = output.to_string();
}
- if need_add_msg {
- self.messages
- .push(Message::new(MessageRole::User, input.message_content()));
- }
- self.data_urls.extend(input.data_urls());
- self.messages.push(Message::new(
- MessageRole::Assistant,
- MessageContent::Text(output.to_string()),
- ));
}
+ } 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 {
+ self.messages
+ .push(Message::new(MessageRole::User, input.message_content()));
+ }
+ self.data_urls.extend(input.data_urls());
+ self.messages.push(Message::new(
+ MessageRole::Assistant,
+ MessageContent::Text(output.to_string()),
+ ));
}
self.dirty = true;
Ok(())
@@ -398,6 +401,9 @@ impl Session {
let mut messages = self.messages.clone();
if input.continue_output().is_some() {
return messages;
+ } else if input.regenerate() {
+ messages.pop();
+ return messages;
}
let mut need_add_msg = true;
let len = messages.len();