diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 10:39:28 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 10:39:28 +0800 |
| commit | a62e461e38482ade15c6826e656393d5f867488a (patch) | |
| tree | bb8c0d76e6863e761f273059744bee722c0ab260 | |
| parent | a7f2da156c0692121e35dbe20bdacc76baadd5c4 (diff) | |
| download | aichat-a62e461e38482ade15c6826e656393d5f867488a.tar.gz | |
feat: support conversation (#48)
| -rw-r--r-- | src/client.rs | 15 | ||||
| -rw-r--r-- | src/config/conversation.rs | 83 | ||||
| -rw-r--r-- | src/config/mod.rs (renamed from src/config.rs) | 108 | ||||
| -rw-r--r-- | src/repl/handler.rs | 28 | ||||
| -rw-r--r-- | src/repl/init.rs | 8 | ||||
| -rw-r--r-- | src/repl/mod.rs | 14 |
6 files changed, 200 insertions, 56 deletions
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<String> { 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<RequestBuilder> { - 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/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<Message>, +} + +#[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<Value> = 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.rs b/src/config/mod.rs index dae3462..6555b14 100644 --- a/src/config.rs +++ b/src/config/mod.rs @@ -1,9 +1,14 @@ +mod conversation; + +use self::conversation::Session; + use crate::utils::{emphasis, now}; -use anyhow::{anyhow, Context, Result}; +use anyhow::{anyhow, bail, Context, Result}; use inquire::{Confirm, Text}; use parking_lot::Mutex; -use serde::{Deserialize, Serialize}; +use serde::Deserialize; +use serde_json::{json, Value}; use std::{ env, fs::{create_dir_all, read_to_string, File, OpenOptions}, @@ -17,7 +22,7 @@ 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 TEMP_ROLE_NAME: &str = "%PROMPT%"; const SET_COMPLETIONS: [&str; 9] = [ ".set api_key", ".set temperature", @@ -53,6 +58,9 @@ pub struct Config { /// Current selected role #[serde(default, skip)] pub role: Option<Role>, + /// Current conversation + #[serde(default, skip)] + pub conversation: Option<Session>, } pub type SharedConfig = Arc<Mutex<Config>>; @@ -149,7 +157,11 @@ impl Config { Self::local_file(MESSAGE_FILE_NAME) } - pub fn change_role(&mut self, name: &str) -> String { + 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) => { let temperature = match role.temperature { @@ -166,28 +178,20 @@ impl Config { temperature ); self.role = Some(role); - output + Ok(output) } - None => "Error: Unknown role".into(), + None => bail!("Error: Unknown role"), } } - pub fn create_temp_role(&mut self, prompt: &str) { + 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, }); - } - - pub fn get_prompt(&self) -> Option<String> { - self.role.as_ref().and_then(|v| { - if v.prompt.is_empty() { - None - } else { - Some(v.prompt.to_string()) - } - }) + Ok(()) } pub fn get_temperature(&self) -> Option<f64> { @@ -197,10 +201,25 @@ impl Config { .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 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]) } } @@ -253,10 +272,10 @@ impl Config { completion } - pub fn update(&mut self, data: &str) -> Result<String> { + pub fn update(&mut self, data: &str) -> Result<()> { let parts: Vec<&str> = data.split_whitespace().collect(); if parts.len() != 2 { - return Ok("Usage: .set <key> <value>. If value is null, unset key.".into()); + bail!("Usage: .set <key> <value>. If value is null, unset key."); } let key = parts[0]; let value = parts[1]; @@ -264,7 +283,7 @@ impl Config { match key { "api_key" => { if unset { - return Ok("Error: Not allowed".into()); + bail!("Error: Not allowed"); } else { self.api_key = value.to_string(); } @@ -296,9 +315,37 @@ impl Config { let value = value.parse().with_context(|| "Invalid value")?; self.dry_run = value; } - _ => return Ok(format!("Error: Unknown key `{key}`")), + _ => bail!("Error: Unknown key `{key}`"), } - Ok("".into()) + 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<File> { @@ -310,6 +357,13 @@ 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() { @@ -324,7 +378,7 @@ impl Config { } } -#[derive(Debug, Clone, Deserialize, Serialize)] +#[derive(Debug, Clone, Deserialize)] pub struct Role { /// Role name pub name: String, 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<str> { - 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<str> { 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 <name>\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::<Vec<String>>() .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, ); } |
