summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-09 10:39:28 +0800
committerGitHub <noreply@github.com>2023-03-09 10:39:28 +0800
commita62e461e38482ade15c6826e656393d5f867488a (patch)
treebb8c0d76e6863e761f273059744bee722c0ab260 /src
parenta7f2da156c0692121e35dbe20bdacc76baadd5c4 (diff)
downloadaichat-a62e461e38482ade15c6826e656393d5f867488a.tar.gz
feat: support conversation (#48)
Diffstat (limited to 'src')
-rw-r--r--src/client.rs15
-rw-r--r--src/config/conversation.rs83
-rw-r--r--src/config/mod.rs (renamed from src/config.rs)108
-rw-r--r--src/repl/handler.rs28
-rw-r--r--src/repl/init.rs8
-rw-r--r--src/repl/mod.rs14
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,
);
}