From a62e461e38482ade15c6826e656393d5f867488a Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 9 Mar 2023 10:39:28 +0800 Subject: feat: support conversation (#48) --- src/client.rs | 15 +- src/config.rs | 375 --------------------------------------- src/config/conversation.rs | 83 +++++++++ src/config/mod.rs | 429 +++++++++++++++++++++++++++++++++++++++++++++ src/repl/handler.rs | 28 +-- src/repl/init.rs | 8 +- src/repl/mod.rs | 14 +- 7 files changed, 548 insertions(+), 404 deletions(-) delete mode 100644 src/config.rs create mode 100644 src/config/conversation.rs create mode 100644 src/config/mod.rs (limited to 'src') diff --git a/src/client.rs b/src/client.rs index 960af0c..18ee97c 100644 --- a/src/client.rs +++ b/src/client.rs @@ -70,7 +70,7 @@ impl ChatGptClient { async fn send_message_inner(&self, content: &str) -> Result { if self.config.lock().dry_run { - return Ok(self.config.lock().merge_prompt(content)); + return Ok(self.config.lock().echo_messages(content)); } let builder = self.request_builder(content, false)?; @@ -89,7 +89,7 @@ impl ChatGptClient { handler: &mut ReplyStreamHandler, ) -> Result<()> { if self.config.lock().dry_run { - handler.text(&self.config.lock().merge_prompt(content))?; + handler.text(&self.config.lock().echo_messages(content))?; return Ok(()); } let builder = self.request_builder(content, true)?; @@ -133,16 +133,7 @@ impl ChatGptClient { } fn request_builder(&self, content: &str, stream: bool) -> Result { - let user_message = json!({ "role": "user", "content": content }); - let messages = match self.config.lock().get_prompt() { - Some(prompt) => { - let system_message = json!({ "role": "system", "content": prompt.trim() }); - json!([system_message, user_message]) - } - None => { - json!([user_message]) - } - }; + let messages = self.config.lock().build_messages(content); let mut body = json!({ "model": MODEL, "messages": messages, diff --git a/src/config.rs b/src/config.rs deleted file mode 100644 index dae3462..0000000 --- a/src/config.rs +++ /dev/null @@ -1,375 +0,0 @@ -use crate::utils::{emphasis, now}; - -use anyhow::{anyhow, Context, Result}; -use inquire::{Confirm, Text}; -use parking_lot::Mutex; -use serde::{Deserialize, Serialize}; -use std::{ - env, - fs::{create_dir_all, read_to_string, File, OpenOptions}, - io::Write, - path::{Path, PathBuf}, - process::exit, - sync::Arc, -}; - -const CONFIG_FILE_NAME: &str = "config.yaml"; -const ROLES_FILE_NAME: &str = "roles.yaml"; -const HISTORY_FILE_NAME: &str = "history.txt"; -const MESSAGE_FILE_NAME: &str = "messages.md"; -const TEMP_ROLE_NAME: &str = "%TEMP%"; -const SET_COMPLETIONS: [&str; 9] = [ - ".set api_key", - ".set temperature", - ".set save true", - ".set save false", - ".set highlight true", - ".set highlight false", - ".set proxy", - ".set dry_run true", - ".set dry_run false", -]; - -#[derive(Debug, Clone, Deserialize)] -pub struct Config { - /// Openai api key - pub api_key: String, - /// What sampling temperature to use, between 0 and 2 - pub temperature: Option, - /// Whether to persistently save chat messages - #[serde(default)] - pub save: bool, - /// Whether to disable highlight - #[serde(default = "highlight_value")] - pub highlight: bool, - /// Set proxy - pub proxy: Option, - /// Used only for debugging - #[serde(default)] - pub dry_run: bool, - /// Predefined roles - #[serde(default, skip)] - pub roles: Vec, - /// Current selected role - #[serde(default, skip)] - pub role: Option, -} - -pub type SharedConfig = Arc>; - -impl Config { - pub fn init(is_interactive: bool) -> Result { - let config_path = Config::config_file()?; - if is_interactive && !config_path.exists() { - create_config_file(&config_path)?; - } - let content = read_to_string(&config_path) - .with_context(|| format!("Failed to load config at {}", config_path.display()))?; - let mut config: Config = serde_yaml::from_str(&content) - .with_context(|| format!("Invalid config at {}", config_path.display()))?; - config.load_roles()?; - Ok(config) - } - - pub fn find_role(&self, name: &str) -> Option { - self.roles.iter().find(|v| v.name == name).cloned() - } - - pub fn config_dir() -> Result { - let env_name = format!( - "{}_CONFIG_DIR", - env!("CARGO_CRATE_NAME").to_ascii_uppercase() - ); - let path = match env::var(env_name) { - Ok(v) => PathBuf::from(v), - Err(_) => { - let mut dir = dirs::config_dir().ok_or_else(|| anyhow!("Not found config dir"))?; - dir.push(env!("CARGO_CRATE_NAME")); - dir - } - }; - if !path.exists() { - create_dir_all(&path).map_err(|err| { - anyhow!("Failed to create config dir at {}, {err}", path.display()) - })?; - } - Ok(path) - } - - pub fn local_file(name: &str) -> Result { - let mut path = Self::config_dir()?; - path.push(name); - Ok(path) - } - - pub fn save_message(&self, input: &str, output: &str) -> Result<()> { - if !self.save { - return Ok(()); - } - let mut file = self.open_message_file()?; - if output.is_empty() || !self.save { - return Ok(()); - } - let timestamp = now(); - let output = match self.role.as_ref() { - None => { - format!("# CHAT:[{timestamp}]\n{input}\n--------\n{output}\n--------\n\n",) - } - Some(v) => { - if v.name == TEMP_ROLE_NAME { - format!( - "# CHAT:[{timestamp}]\n{}\n{input}\n--------\n{output}\n--------\n\n", - v.prompt - ) - } else { - format!( - "# CHAT:[{timestamp}] ({})\n{input}\n--------\n{output}\n--------\n\n", - v.name, - ) - } - } - }; - file.write_all(output.as_bytes()) - .with_context(|| "Failed to save message") - } - - pub fn config_file() -> Result { - Self::local_file(CONFIG_FILE_NAME) - } - - pub fn roles_file() -> Result { - Self::local_file(ROLES_FILE_NAME) - } - - pub fn history_file() -> Result { - Self::local_file(HISTORY_FILE_NAME) - } - - pub fn messages_file() -> Result { - Self::local_file(MESSAGE_FILE_NAME) - } - - pub fn change_role(&mut self, name: &str) -> String { - match self.find_role(name) { - Some(role) => { - let temperature = match role.temperature { - Some(v) => format!("{v}"), - None => "null".into(), - }; - let output = format!( - "{}: {}\n{}: {}\n{}: {}", - emphasis("name"), - role.name, - emphasis("prompt"), - role.prompt.trim(), - emphasis("temperature"), - temperature - ); - self.role = Some(role); - output - } - None => "Error: Unknown role".into(), - } - } - - pub fn create_temp_role(&mut self, prompt: &str) { - self.role = Some(Role { - name: TEMP_ROLE_NAME.into(), - prompt: prompt.into(), - temperature: self.temperature, - }); - } - - pub fn get_prompt(&self) -> Option { - self.role.as_ref().and_then(|v| { - if v.prompt.is_empty() { - None - } else { - Some(v.prompt.to_string()) - } - }) - } - - pub fn get_temperature(&self) -> Option { - self.role - .as_ref() - .and_then(|v| v.temperature) - .or(self.temperature) - } - - pub fn merge_prompt(&self, content: &str) -> String { - match self.get_prompt() { - Some(prompt) => format!("{}\n{content}", prompt.trim()), - None => content.to_string(), - } - } - - pub fn info(&self) -> Result { - let file_info = |path: &Path| { - let state = if path.exists() { "" } else { " ⚠️" }; - format!("{}{state}", path.display()) - }; - let proxy = self - .proxy - .as_ref() - .map(|v| v.to_string()) - .unwrap_or("-".into()); - let temperature = self - .temperature - .map(|v| v.to_string()) - .unwrap_or("-".into()); - let role_name = self - .role - .as_ref() - .map(|v| v.name.to_string()) - .unwrap_or("-".into()); - let items = vec![ - ("config_file", file_info(&Config::config_file()?)), - ("roles_file", file_info(&Config::roles_file()?)), - ("messages_file", file_info(&Config::messages_file()?)), - ("role", role_name), - ("api_key", self.api_key.clone()), - ("temperature", temperature), - ("save", self.save.to_string()), - ("highlight", self.highlight.to_string()), - ("proxy", proxy), - ("dry_run", self.dry_run.to_string()), - ]; - let mut output = String::new(); - for (name, value) in items { - output.push_str(&format!("{name:<20}{value}\n")); - } - Ok(output) - } - - pub fn repl_completions(&self) -> Vec { - let mut completion: Vec = self - .roles - .iter() - .map(|v| format!(".role {}", v.name)) - .collect(); - - completion.extend(SET_COMPLETIONS.map(|v| v.to_string())); - completion - } - - pub fn update(&mut self, data: &str) -> Result { - let parts: Vec<&str> = data.split_whitespace().collect(); - if parts.len() != 2 { - return Ok("Usage: .set . If value is null, unset key.".into()); - } - let key = parts[0]; - let value = parts[1]; - let unset = value == "null"; - match key { - "api_key" => { - if unset { - return Ok("Error: Not allowed".into()); - } else { - self.api_key = value.to_string(); - } - } - "temperature" => { - if unset { - self.temperature = None; - } else { - let value = value.parse().with_context(|| "Invalid value")?; - self.temperature = Some(value); - } - } - "save" => { - let value = value.parse().with_context(|| "Invalid value")?; - self.save = value; - } - "highlight" => { - let value = value.parse().with_context(|| "Invalid value")?; - self.highlight = value; - } - "proxy" => { - if unset { - self.proxy = None; - } else { - self.proxy = Some(value.to_string()); - } - } - "dry_run" => { - let value = value.parse().with_context(|| "Invalid value")?; - self.dry_run = value; - } - _ => return Ok(format!("Error: Unknown key `{key}`")), - } - Ok("".into()) - } - - fn open_message_file(&self) -> Result { - let path = Config::messages_file()?; - OpenOptions::new() - .create(true) - .append(true) - .open(&path) - .with_context(|| format!("Failed to create/append {}", path.display())) - } - - fn load_roles(&mut self) -> Result<()> { - let path = Self::roles_file()?; - if !path.exists() { - return Ok(()); - } - let content = read_to_string(&path) - .with_context(|| format!("Failed to load roles at {}", path.display()))?; - let roles: Vec = - serde_yaml::from_str(&content).with_context(|| "Invalid roles config")?; - self.roles = roles; - Ok(()) - } -} - -#[derive(Debug, Clone, Deserialize, Serialize)] -pub struct Role { - /// Role name - pub name: String, - /// Prompt text send to ai for setting up a role - pub prompt: String, - /// What sampling temperature to use, between 0 and 2 - pub temperature: Option, -} - -fn create_config_file(config_path: &Path) -> Result<()> { - let confirm_map_err = |_| anyhow!("Error with questionnaire, try again later"); - let text_map_err = |_| anyhow!("An error happened when asking for your key, try again later."); - let ans = Confirm::new("No config file, create a new one?") - .with_default(true) - .prompt() - .map_err(confirm_map_err)?; - if !ans { - exit(0); - } - let api_key = Text::new("Openai API Key:") - .prompt() - .map_err(text_map_err)?; - let mut raw_config = format!("api_key: {api_key}\n"); - - let ans = Confirm::new("Use proxy?") - .with_default(false) - .prompt() - .map_err(confirm_map_err)?; - if ans { - let proxy = Text::new("Set proxy:").prompt().map_err(text_map_err)?; - raw_config.push_str(&format!("proxy: {proxy}\n")); - } - - let ans = Confirm::new("Save chat messages") - .with_default(false) - .prompt() - .map_err(confirm_map_err)?; - if ans { - raw_config.push_str("save: true\n"); - } - - std::fs::write(config_path, raw_config).with_context(|| "Failed to write to config file")?; - Ok(()) -} - -fn highlight_value() -> bool { - true -} diff --git a/src/config/conversation.rs b/src/config/conversation.rs new file mode 100644 index 0000000..ca50233 --- /dev/null +++ b/src/config/conversation.rs @@ -0,0 +1,83 @@ +use anyhow::Result; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct Session { + pub tokens: usize, + pub messages: Vec, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct Message { + pub role: MessageRole, + pub content: String, +} + +impl Session { + pub fn new() -> Self { + Self { + tokens: 0, + messages: vec![], + } + } + + pub fn add_conversatoin(&mut self, input: &str, output: &str) -> Result<()> { + self.messages.push(Message { + role: MessageRole::User, + content: input.to_string(), + }); + self.messages.push(Message { + role: MessageRole::Assistant, + content: output.to_string(), + }); + Ok(()) + } + + /// Readline prompt + pub fn add_prompt(&mut self, prompt: &str) { + self.messages.push(Message { + role: MessageRole::System, + content: prompt.into(), + }); + } + + pub fn echo_messages(&self, content: &str) -> String { + let mut messages = self.messages.to_vec(); + messages.push(Message { + role: MessageRole::User, + content: content.into(), + }); + serde_yaml::to_string(&messages).unwrap_or("Unable to echo message".into()) + } + + pub fn build_emssages(&self, content: &str) -> Value { + let mut messages: Vec = self.messages.iter().map(msg_to_value).collect(); + messages.push(msg_to_value(&Message { + role: MessageRole::User, + content: content.into(), + })); + json!(messages) + } +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub enum MessageRole { + System, + Assistant, + User, +} + +impl MessageRole { + pub fn name(&self) -> &'static str { + match self { + MessageRole::System => "system", + MessageRole::Assistant => "assistant", + MessageRole::User => "user", + } + } +} + +fn msg_to_value(msg: &Message) -> Value { + json!({ "role": msg.role.name(), "content": msg.content }) +} diff --git a/src/config/mod.rs b/src/config/mod.rs new file mode 100644 index 0000000..6555b14 --- /dev/null +++ b/src/config/mod.rs @@ -0,0 +1,429 @@ +mod conversation; + +use self::conversation::Session; + +use crate::utils::{emphasis, now}; + +use anyhow::{anyhow, bail, Context, Result}; +use inquire::{Confirm, Text}; +use parking_lot::Mutex; +use serde::Deserialize; +use serde_json::{json, Value}; +use std::{ + env, + fs::{create_dir_all, read_to_string, File, OpenOptions}, + io::Write, + path::{Path, PathBuf}, + process::exit, + sync::Arc, +}; + +const CONFIG_FILE_NAME: &str = "config.yaml"; +const ROLES_FILE_NAME: &str = "roles.yaml"; +const HISTORY_FILE_NAME: &str = "history.txt"; +const MESSAGE_FILE_NAME: &str = "messages.md"; +const TEMP_ROLE_NAME: &str = "%PROMPT%"; +const SET_COMPLETIONS: [&str; 9] = [ + ".set api_key", + ".set temperature", + ".set save true", + ".set save false", + ".set highlight true", + ".set highlight false", + ".set proxy", + ".set dry_run true", + ".set dry_run false", +]; + +#[derive(Debug, Clone, Deserialize)] +pub struct Config { + /// Openai api key + pub api_key: String, + /// What sampling temperature to use, between 0 and 2 + pub temperature: Option, + /// Whether to persistently save chat messages + #[serde(default)] + pub save: bool, + /// Whether to disable highlight + #[serde(default = "highlight_value")] + pub highlight: bool, + /// Set proxy + pub proxy: Option, + /// Used only for debugging + #[serde(default)] + pub dry_run: bool, + /// Predefined roles + #[serde(default, skip)] + pub roles: Vec, + /// Current selected role + #[serde(default, skip)] + pub role: Option, + /// Current conversation + #[serde(default, skip)] + pub conversation: Option, +} + +pub type SharedConfig = Arc>; + +impl Config { + pub fn init(is_interactive: bool) -> Result { + let config_path = Config::config_file()?; + if is_interactive && !config_path.exists() { + create_config_file(&config_path)?; + } + let content = read_to_string(&config_path) + .with_context(|| format!("Failed to load config at {}", config_path.display()))?; + let mut config: Config = serde_yaml::from_str(&content) + .with_context(|| format!("Invalid config at {}", config_path.display()))?; + config.load_roles()?; + Ok(config) + } + + pub fn find_role(&self, name: &str) -> Option { + self.roles.iter().find(|v| v.name == name).cloned() + } + + pub fn config_dir() -> Result { + let env_name = format!( + "{}_CONFIG_DIR", + env!("CARGO_CRATE_NAME").to_ascii_uppercase() + ); + let path = match env::var(env_name) { + Ok(v) => PathBuf::from(v), + Err(_) => { + let mut dir = dirs::config_dir().ok_or_else(|| anyhow!("Not found config dir"))?; + dir.push(env!("CARGO_CRATE_NAME")); + dir + } + }; + if !path.exists() { + create_dir_all(&path).map_err(|err| { + anyhow!("Failed to create config dir at {}, {err}", path.display()) + })?; + } + Ok(path) + } + + pub fn local_file(name: &str) -> Result { + let mut path = Self::config_dir()?; + path.push(name); + Ok(path) + } + + pub fn save_message(&self, input: &str, output: &str) -> Result<()> { + if !self.save { + return Ok(()); + } + let mut file = self.open_message_file()?; + if output.is_empty() || !self.save { + return Ok(()); + } + let timestamp = now(); + let output = match self.role.as_ref() { + None => { + format!("# CHAT:[{timestamp}]\n{input}\n--------\n{output}\n--------\n\n",) + } + Some(v) => { + if v.name == TEMP_ROLE_NAME { + format!( + "# CHAT:[{timestamp}]\n{}\n{input}\n--------\n{output}\n--------\n\n", + v.prompt + ) + } else { + format!( + "# CHAT:[{timestamp}] ({})\n{input}\n--------\n{output}\n--------\n\n", + v.name, + ) + } + } + }; + file.write_all(output.as_bytes()) + .with_context(|| "Failed to save message") + } + + pub fn config_file() -> Result { + Self::local_file(CONFIG_FILE_NAME) + } + + pub fn roles_file() -> Result { + Self::local_file(ROLES_FILE_NAME) + } + + pub fn history_file() -> Result { + Self::local_file(HISTORY_FILE_NAME) + } + + pub fn messages_file() -> Result { + Self::local_file(MESSAGE_FILE_NAME) + } + + pub fn change_role(&mut self, name: &str) -> Result { + self.ensure_no_conversation()?; + if self.conversation.is_some() { + bail!("") + } + match self.find_role(name) { + Some(role) => { + let temperature = match role.temperature { + Some(v) => format!("{v}"), + None => "null".into(), + }; + let output = format!( + "{}: {}\n{}: {}\n{}: {}", + emphasis("name"), + role.name, + emphasis("prompt"), + role.prompt.trim(), + emphasis("temperature"), + temperature + ); + self.role = Some(role); + Ok(output) + } + None => bail!("Error: Unknown role"), + } + } + + pub fn create_temp_role(&mut self, prompt: &str) -> Result<()> { + self.ensure_no_conversation()?; + self.role = Some(Role { + name: TEMP_ROLE_NAME.into(), + prompt: prompt.into(), + temperature: self.temperature, + }); + Ok(()) + } + + pub fn get_temperature(&self) -> Option { + self.role + .as_ref() + .and_then(|v| v.temperature) + .or(self.temperature) + } + + pub fn echo_messages(&self, content: &str) -> String { + if let Some(conversation) = self.conversation.as_ref() { + conversation.echo_messages(content) + } else if let Some(role) = self.role.as_ref() { + format!("{}\n{content}", role.prompt.trim()) + } else { + content.to_string() + } + } + + pub fn build_messages(&self, content: &str) -> Value { + let user_message = json!({ "role": "user", "content": content }); + if let Some(conversation) = self.conversation.as_ref() { + conversation.build_emssages(content) + } else if let Some(role) = self.role.as_ref() { + let system_message = json!({ "role": "system", "content": role.prompt.trim() }); + json!([system_message, user_message]) + } else { + json!([user_message]) + } + } + + pub fn info(&self) -> Result { + let file_info = |path: &Path| { + let state = if path.exists() { "" } else { " ⚠️" }; + format!("{}{state}", path.display()) + }; + let proxy = self + .proxy + .as_ref() + .map(|v| v.to_string()) + .unwrap_or("-".into()); + let temperature = self + .temperature + .map(|v| v.to_string()) + .unwrap_or("-".into()); + let role_name = self + .role + .as_ref() + .map(|v| v.name.to_string()) + .unwrap_or("-".into()); + let items = vec![ + ("config_file", file_info(&Config::config_file()?)), + ("roles_file", file_info(&Config::roles_file()?)), + ("messages_file", file_info(&Config::messages_file()?)), + ("role", role_name), + ("api_key", self.api_key.clone()), + ("temperature", temperature), + ("save", self.save.to_string()), + ("highlight", self.highlight.to_string()), + ("proxy", proxy), + ("dry_run", self.dry_run.to_string()), + ]; + let mut output = String::new(); + for (name, value) in items { + output.push_str(&format!("{name:<20}{value}\n")); + } + Ok(output) + } + + pub fn repl_completions(&self) -> Vec { + let mut completion: Vec = self + .roles + .iter() + .map(|v| format!(".role {}", v.name)) + .collect(); + + completion.extend(SET_COMPLETIONS.map(|v| v.to_string())); + completion + } + + pub fn update(&mut self, data: &str) -> Result<()> { + let parts: Vec<&str> = data.split_whitespace().collect(); + if parts.len() != 2 { + bail!("Usage: .set . If value is null, unset key."); + } + let key = parts[0]; + let value = parts[1]; + let unset = value == "null"; + match key { + "api_key" => { + if unset { + bail!("Error: Not allowed"); + } else { + self.api_key = value.to_string(); + } + } + "temperature" => { + if unset { + self.temperature = None; + } else { + let value = value.parse().with_context(|| "Invalid value")?; + self.temperature = Some(value); + } + } + "save" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.save = value; + } + "highlight" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.highlight = value; + } + "proxy" => { + if unset { + self.proxy = None; + } else { + self.proxy = Some(value.to_string()); + } + } + "dry_run" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.dry_run = value; + } + _ => bail!("Error: Unknown key `{key}`"), + } + Ok(()) + } + + pub fn start_conversation(&mut self) -> Result<()> { + if self.conversation.is_some() { + let ans = Confirm::new("Already in a conversation, start a new one?") + .with_default(true) + .prompt()?; + if !ans { + return Ok(()); + } + } + let mut conversation = Session::new(); + if let Some(role) = self.role.as_ref() { + conversation.add_prompt(&role.prompt); + } + self.conversation = Some(conversation); + Ok(()) + } + + pub fn end_conversation(&mut self) { + self.conversation = None; + } + + pub fn record_conversation(&mut self, input: &str, output: &str) -> Result<()> { + if let Some(conversation) = self.conversation.as_mut() { + conversation.add_conversatoin(input, output)?; + } + Ok(()) + } + + fn open_message_file(&self) -> Result { + let path = Config::messages_file()?; + OpenOptions::new() + .create(true) + .append(true) + .open(&path) + .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() { + return Ok(()); + } + let content = read_to_string(&path) + .with_context(|| format!("Failed to load roles at {}", path.display()))?; + let roles: Vec = + serde_yaml::from_str(&content).with_context(|| "Invalid roles config")?; + self.roles = roles; + Ok(()) + } +} + +#[derive(Debug, Clone, Deserialize)] +pub struct Role { + /// Role name + pub name: String, + /// Prompt text send to ai for setting up a role + pub prompt: String, + /// What sampling temperature to use, between 0 and 2 + pub temperature: Option, +} + +fn create_config_file(config_path: &Path) -> Result<()> { + let confirm_map_err = |_| anyhow!("Error with questionnaire, try again later"); + let text_map_err = |_| anyhow!("An error happened when asking for your key, try again later."); + let ans = Confirm::new("No config file, create a new one?") + .with_default(true) + .prompt() + .map_err(confirm_map_err)?; + if !ans { + exit(0); + } + let api_key = Text::new("Openai API Key:") + .prompt() + .map_err(text_map_err)?; + let mut raw_config = format!("api_key: {api_key}\n"); + + let ans = Confirm::new("Use proxy?") + .with_default(false) + .prompt() + .map_err(confirm_map_err)?; + if ans { + let proxy = Text::new("Set proxy:").prompt().map_err(text_map_err)?; + raw_config.push_str(&format!("proxy: {proxy}\n")); + } + + let ans = Confirm::new("Save chat messages") + .with_default(false) + .prompt() + .map_err(confirm_map_err)?; + if ans { + raw_config.push_str("save: true\n"); + } + + std::fs::write(config_path, raw_config).with_context(|| "Failed to write to config file")?; + Ok(()) +} + +fn highlight_value() -> bool { + true +} diff --git a/src/repl/handler.rs b/src/repl/handler.rs index 1b2eb20..979fc5e 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -16,7 +16,9 @@ pub enum ReplCmd { UpdateConfig(String), Prompt(String), ClearRole, - Info, + ViewInfo, + StartConversation, + EndConversatoin, } pub struct ReplCmdHandler { @@ -61,10 +63,11 @@ impl ReplCmdHandler { wg.wait(); let buffer = ret?; self.config.lock().save_message(&input, &buffer)?; + self.config.lock().record_conversation(&input, &buffer)?; *self.reply.borrow_mut() = buffer; } ReplCmd::SetRole(name) => { - let output = self.config.lock().change_role(&name); + let output = self.config.lock().change_role(&name)?; print_now!("{}\n\n", output.trim_end()); } ReplCmd::ClearRole => { @@ -72,21 +75,24 @@ impl ReplCmdHandler { print_now!("\n"); } ReplCmd::Prompt(prompt) => { - self.config.lock().create_temp_role(&prompt); + self.config.lock().create_temp_role(&prompt)?; print_now!("\n"); } - ReplCmd::Info => { + ReplCmd::ViewInfo => { let output = self.config.lock().info()?; print_now!("{}\n\n", output.trim_end()); } ReplCmd::UpdateConfig(input) => { - let output = self.config.lock().update(&input)?; - let output = output.trim(); - if output.is_empty() { - print_now!("\n"); - } else { - print_now!("{}\n\n", output); - } + self.config.lock().update(&input)?; + print_now!("\n"); + } + ReplCmd::StartConversation => { + self.config.lock().start_conversation()?; + print_now!("\n"); + } + ReplCmd::EndConversatoin => { + self.config.lock().end_conversation(); + print_now!("\n"); } } Ok(()) diff --git a/src/repl/init.rs b/src/repl/init.rs index 998b9d0..a14265c 100644 --- a/src/repl/init.rs +++ b/src/repl/init.rs @@ -11,7 +11,6 @@ use reedline::{ use std::borrow::Cow; const MENU_NAME: &str = "completion_menu"; -const DEFAULT_PROMPT_INDICATOR: &str = "〉"; const DEFAULT_MULTILINE_INDICATOR: &str = "::: "; pub struct Repl { @@ -140,7 +139,12 @@ impl Prompt for ReplPrompt { } fn render_prompt_indicator(&self, _prompt_mode: reedline::PromptEditMode) -> Cow { - Cow::Borrowed(DEFAULT_PROMPT_INDICATOR) + let config = self.0.lock(); + if config.conversation.is_some() { + Cow::Borrowed("$") + } else { + Cow::Borrowed("〉") + } } fn render_prompt_multiline_indicator(&self) -> Cow { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index e9a0b51..9407bcf 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -15,12 +15,14 @@ use anyhow::{Context, Result}; use reedline::Signal; use std::sync::Arc; -pub const REPL_COMMANDS: [(&str, &str, bool); 10] = [ +pub const REPL_COMMANDS: [(&str, &str, bool); 12] = [ (".info", "Print the information", false), (".set", "Modify the configuration temporarily", false), (".prompt", "Add a GPT prompt", true), (".role", "Select a role", false), (".clear role", "Clear the currently selected role", false), + (".conversation", "Start a conversation.", false), + (".clear conversation", "End the conversation.", false), (".history", "Print the history", false), (".clear history", "Clear the history", false), (".editor", "Enter editor mode for multiline input", true), @@ -102,6 +104,7 @@ impl Repl { print_now!("\n"); } Some("role") => handler.handle(ReplCmd::ClearRole)?, + Some("conversation") => handler.handle(ReplCmd::EndConversatoin)?, _ => dump_unknown_command(), }, ".history" => { @@ -113,7 +116,7 @@ impl Repl { None => print_now!("Usage: .role \n\n"), }, ".info" => { - handler.handle(ReplCmd::Info)?; + handler.handle(ReplCmd::ViewInfo)?; } ".editor" => { let mut text = args.unwrap_or_default().to_string(); @@ -140,6 +143,9 @@ impl Repl { handler.handle(ReplCmd::Prompt(text))?; } } + ".conversation" => { + handler.handle(ReplCmd::StartConversation)?; + } _ => dump_unknown_command(), } } else { @@ -157,11 +163,11 @@ fn dump_unknown_command() { fn dump_repl_help() { let head = REPL_COMMANDS .iter() - .map(|(name, desc, _)| format!("{name:<15} {desc}")) + .map(|(name, desc, _)| format!("{name:<24} {desc}")) .collect::>() .join("\n"); print_now!( - "{}\n\nPress Ctrl+C to abort session, Ctrl+D to exit the REPL\n\n", + "{}\n\nPress Ctrl+C to abort conversation, Ctrl+D to exit the REPL\n\n", head, ); } -- cgit v1.2.3