summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-17 19:54:24 +0800
committerGitHub <noreply@github.com>2024-06-17 19:54:24 +0800
commitb965c63be05c726c996ece298d6f5f2514a19316 (patch)
treedecc1e021caa638de171e0fb94e6e4f5b3e39f63 /src/config
parent62b297e8bb09e257a61154afddc4017e48bd12f3 (diff)
downloadaichat-b965c63be05c726c996ece298d6f5f2514a19316.tar.gz
refactor: minor improvement (#611)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/mod.rs16
-rw-r--r--src/config/session.rs8
2 files changed, 16 insertions, 8 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index cb510ef..41e449a 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -583,7 +583,8 @@ impl Config {
}
pub fn use_prompt(&mut self, prompt: &str) -> Result<()> {
- let role = Role::new(TEMP_ROLE_NAME, prompt);
+ let mut role = Role::new(TEMP_ROLE_NAME, prompt);
+ role.set_model(&self.model);
self.use_role_obj(role)
}
@@ -705,9 +706,8 @@ impl Config {
pub fn exit_session(&mut self) -> Result<()> {
if let Some(mut session) = self.session.take() {
- let is_repl = self.working_mode == WorkingMode::Repl;
let sessions_dir = self.sessions_dir()?;
- session.exit(&sessions_dir, is_repl)?;
+ session.exit(&sessions_dir, self.working_mode.is_repl())?;
self.last_message = None;
}
Ok(())
@@ -723,7 +723,7 @@ impl Config {
};
let session_path = self.session_file(&name)?;
if let Some(session) = self.session.as_mut() {
- session.save(&session_path)?;
+ session.save(&session_path, self.working_mode.is_repl())?;
}
Ok(())
}
@@ -933,8 +933,8 @@ impl Config {
}
pub fn exit_bot(&mut self) -> Result<()> {
+ self.exit_session()?;
if self.bot.take().is_some() {
- self.exit_session()?;
self.rag.take();
self.last_message = None;
}
@@ -1420,6 +1420,12 @@ pub enum WorkingMode {
Serve,
}
+impl WorkingMode {
+ pub fn is_repl(&self) -> bool {
+ *self == WorkingMode::Repl
+ }
+}
+
bitflags::bitflags! {
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct StateFlags: u32 {
diff --git a/src/config/session.rs b/src/config/session.rs
index 2149ba5..726249c 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -312,12 +312,12 @@ impl Session {
}
}
let session_path = session_dir.join(format!("{}.yaml", self.name()));
- self.save(&session_path)?;
+ self.save(&session_path, is_repl)?;
}
Ok(())
}
- pub fn save(&mut self, session_path: &Path) -> Result<()> {
+ pub fn save(&mut self, session_path: &Path, is_repl: bool) -> Result<()> {
if let Some(sessions_dir) = session_path.parent() {
if !sessions_dir.exists() {
create_dir_all(sessions_dir).with_context(|| {
@@ -338,7 +338,9 @@ impl Session {
)
})?;
- println!("✨ Saved session to '{}'", session_path.display());
+ if is_repl {
+ println!("✨ Saved session to '{}'", session_path.display());
+ }
self.dirty = false;