diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 23:30:12 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 23:30:12 +0800 |
| commit | 0044399509dc26df4f0d1d776ca97edc73b041e5 (patch) | |
| tree | b01602d523c3eb0f284a7a98821cb217ea0501c6 /src | |
| parent | ad282ae74c29bb96a213f2d02412f080bc2ce938 (diff) | |
| download | aichat-0044399509dc26df4f0d1d776ca97edc73b041e5.tar.gz | |
feat: add config.conversation_first (#55)
If set true, start a conversation immediately upon repl.
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/conversation.rs | 18 | ||||
| -rw-r--r-- | src/config/mod.rs | 29 |
2 files changed, 32 insertions, 15 deletions
diff --git a/src/config/conversation.rs b/src/config/conversation.rs index 4472781..7fe843d 100644 --- a/src/config/conversation.rs +++ b/src/config/conversation.rs @@ -2,7 +2,7 @@ use super::message::{num_tokens_from_messages, Message, MessageRole}; use super::role::Role; use super::MAX_TOKENS; -use anyhow::Result; +use anyhow::{bail, Result}; use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Deserialize, Serialize)] @@ -19,10 +19,24 @@ impl Conversation { role, messages: vec![], }; - value.tokens = num_tokens_from_messages(&value.build_emssages("")); + value.update_tokens(); value } + pub fn update_role(&mut self, role: &Role) -> Result<()> { + if self.messages.is_empty() { + self.role = Some(role.clone()); + self.update_tokens(); + } else { + bail!("Error: Cannot perform this action in the middle of conversation") + } + Ok(()) + } + + pub fn update_tokens(&mut self) { + self.tokens = num_tokens_from_messages(&self.build_emssages("")); + } + pub fn add_message(&mut self, input: &str, output: &str) -> Result<()> { let mut need_add_msg = true; if self.messages.is_empty() { diff --git a/src/config/mod.rs b/src/config/mod.rs index 2c96161..644d162 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -55,6 +55,9 @@ pub struct Config { /// Used only for debugging #[serde(default)] pub dry_run: bool, + /// If set ture, start a conversation immediately upon repl + #[serde(default)] + pub conversation_first: bool, /// Predefined roles #[serde(skip)] pub roles: Vec<Role>, @@ -79,6 +82,10 @@ impl Config { let mut config: Config = serde_yaml::from_str(&content) .with_context(|| format!("Invalid config at {}", config_path.display()))?; config.load_roles()?; + if config.conversation_first { + config.start_conversation()?; + } + Ok(config) } @@ -161,12 +168,11 @@ impl Config { } pub fn change_role(&mut self, name: &str) -> Result<String> { - self.ensure_no_conversation()?; - if self.conversation.is_some() { - bail!("") - } match self.find_role(name) { Some(role) => { + if let Some(conversation) = self.conversation.as_mut() { + conversation.update_role(&role)?; + } let output = serde_yaml::to_string(&role).unwrap_or("Unable to echo role details".into()); self.role = Some(role); @@ -177,8 +183,11 @@ impl Config { } pub fn create_temp_role(&mut self, prompt: &str) -> Result<()> { - self.ensure_no_conversation()?; - self.role = Some(Role::new(prompt, self.temperature)); + let role = Role::new(prompt, self.temperature); + if let Some(conversation) = self.conversation.as_mut() { + conversation.update_role(&role)?; + } + self.role = Some(role); Ok(()) } @@ -239,6 +248,7 @@ impl Config { ("save", self.save.to_string()), ("highlight", self.highlight.to_string()), ("proxy", proxy), + ("conversation_first", self.conversation_first.to_string()), ("dry_run", self.dry_run.to_string()), ]; let mut output = String::new(); @@ -342,13 +352,6 @@ impl Config { .with_context(|| format!("Failed to create/append {}", path.display())) } - fn ensure_no_conversation(&self) -> Result<()> { - if self.conversation.is_some() { - bail!("Error: Cannot perform this action in a conversation"); - } - Ok(()) - } - fn load_roles(&mut self) -> Result<()> { let path = Self::roles_file()?; if !path.exists() { |
