summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-12-24 16:04:18 +0800
committerGitHub <noreply@github.com>2023-12-24 16:04:18 +0800
commit1c9ca1b002badb511ff21d58b3f20986741ab31f (patch)
tree171edf24341f96ac4a93e00b315a3ce8b90455fc /src/config
parent89fefb4d1a3239496e90daf74be8b08decc603d2 (diff)
downloadaichat-1c9ca1b002badb511ff21d58b3f20986741ab31f.tar.gz
feat: custom REPL prompt (#283)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/mod.rs92
-rw-r--r--src/config/session.rs4
2 files changed, 84 insertions, 12 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 52203ba..71e47da 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -11,13 +11,14 @@ use crate::client::{
Model, OpenAIClient, SendData,
};
use crate::render::{MarkdownRender, RenderOptions};
-use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, prompt_op_err};
+use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, prompt_op_err, render_prompt};
use anyhow::{anyhow, bail, Context, Result};
use inquire::{Confirm, Select, Text};
use is_terminal::IsTerminal;
use parking_lot::RwLock;
use serde::Deserialize;
+use std::collections::HashMap;
use std::{
env,
fs::{create_dir_all, read_dir, read_to_string, remove_file, File, OpenOptions},
@@ -66,6 +67,10 @@ pub struct Config {
pub keybindings: Keybindings,
/// Set a default role or session (role:<name>, session:<name>)
pub prelude: String,
+ /// REPL left prompt
+ pub left_prompt: String,
+ /// REPL right prompt
+ pub right_prompt: String,
/// Setup clients
pub clients: Vec<ClientConfig>,
/// Predefined roles
@@ -99,6 +104,9 @@ impl Default for Config {
auto_copy: false,
keybindings: Default::default(),
prelude: String::new(),
+ left_prompt: "{color.green}{?session {session}{?role /}}{role}{color.cyan}{?session )}{!session >}{color.reset} ".to_string(),
+ right_prompt: "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"
+ .to_string(),
clients: vec![ClientConfig::default()],
roles: vec![],
role: None,
@@ -648,18 +656,14 @@ impl Config {
Ok(RenderOptions::new(theme, wrap, self.wrap_code))
}
+ pub fn render_prompt_left(&self) -> String {
+ let variables = self.generate_prompt_context();
+ render_prompt(&self.left_prompt, &variables)
+ }
+
pub fn render_prompt_right(&self) -> String {
- if let Some(session) = &self.session {
- let (tokens, percent) = session.tokens_and_percent();
- let percent = if percent == 0.0 {
- String::new()
- } else {
- format!("({percent}%)")
- };
- format!("{tokens}{percent}")
- } else {
- String::new()
- }
+ let variables = self.generate_prompt_context();
+ render_prompt(&self.right_prompt, &variables)
}
pub fn prepare_send_data(&self, input: &Input, stream: bool) -> Result<SendData> {
@@ -681,6 +685,70 @@ impl Config {
}
}
+ fn generate_prompt_context(&self) -> HashMap<&str, String> {
+ let mut output = HashMap::new();
+ output.insert("model", self.model.id());
+ output.insert("client_name", self.model.client_name.clone());
+ output.insert("model_name", self.model.name.clone());
+ output.insert(
+ "max_tokens",
+ self.model.max_tokens.unwrap_or_default().to_string(),
+ );
+ if let Some(temperature) = self.temperature {
+ if temperature != 0.0 {
+ output.insert("temperature", temperature.to_string());
+ }
+ }
+ if self.dry_run {
+ output.insert("dry_run", "true".to_string());
+ }
+ if self.save {
+ output.insert("save", "true".to_string());
+ }
+ if let Some(wrap) = &self.wrap {
+ if wrap != "no" {
+ output.insert("wrap", wrap.clone());
+ }
+ }
+ if self.auto_copy {
+ output.insert("auto_copy", "true".to_string());
+ }
+ if let Some(role) = &self.role {
+ output.insert("role", role.name.clone());
+ }
+ if let Some(session) = &self.session {
+ output.insert("session", session.name().to_string());
+ let (tokens, percent) = session.tokens_and_percent();
+ output.insert("consume_tokens", tokens.to_string());
+ output.insert("consume_percent", percent.to_string());
+ output.insert("user_messages_len", session.user_messages_len().to_string());
+ }
+
+ if self.highlight {
+ output.insert("color.reset", "\u{1b}[0m".to_string());
+ output.insert("color.black", "\u{1b}[30m".to_string());
+ output.insert("color.dark_gray", "\u{1b}[90m".to_string());
+ output.insert("color.red", "\u{1b}[31m".to_string());
+ output.insert("color.light_red", "\u{1b}[91m".to_string());
+ output.insert("color.green", "\u{1b}[32m".to_string());
+ output.insert("color.light_green", "\u{1b}[92m".to_string());
+ output.insert("color.yellow", "\u{1b}[33m".to_string());
+ output.insert("color.light_yellow", "\u{1b}[93m".to_string());
+ output.insert("color.blue", "\u{1b}[34m".to_string());
+ output.insert("color.light_blue", "\u{1b}[94m".to_string());
+ output.insert("color.purple", "\u{1b}[35m".to_string());
+ output.insert("color.light_purple", "\u{1b}[95m".to_string());
+ output.insert("color.magenta", "\u{1b}[35m".to_string());
+ output.insert("color.light_magenta", "\u{1b}[95m".to_string());
+ output.insert("color.cyan", "\u{1b}[36m".to_string());
+ output.insert("color.light_cyan", "\u{1b}[96m".to_string());
+ output.insert("color.white", "\u{1b}[37m".to_string());
+ output.insert("color.light_gray", "\u{1b}[97m".to_string());
+ }
+
+ output
+ }
+
fn open_message_file(&self) -> Result<File> {
let path = Self::messages_file()?;
ensure_parent_exists(&path)?;
diff --git a/src/config/session.rs b/src/config/session.rs
index e135c58..cbfe1c2 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -78,6 +78,10 @@ impl Session {
self.model.total_tokens(&self.messages)
}
+ pub fn user_messages_len(&self) -> usize {
+ self.messages.iter().filter(|v| v.role.is_user()).count()
+ }
+
pub fn export(&self) -> Result<String> {
self.guard_save()?;
let (tokens, percent) = self.tokens_and_percent();