summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-10 18:35:34 +0800
committerGitHub <noreply@github.com>2024-09-10 18:35:34 +0800
commite181ae9b0d4814a7c005c94837426126dde63c6c (patch)
treeed727036d586f0673198a4b30217bbe2e34a094b /src
parent84e9515509c559ed01e4b0a67539f10cd2c065e6 (diff)
downloadaichat-e181ae9b0d4814a7c005c94837426126dde63c6c.tar.gz
refactor: extract built-in roles to embedded files (#853)
Diffstat (limited to 'src')
-rw-r--r--src/config/agent.rs21
-rw-r--r--src/config/mod.rs19
-rw-r--r--src/config/role.rs106
-rw-r--r--src/config/session.rs9
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/utils/command.rs15
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/variables.rs33
8 files changed, 79 insertions, 128 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs
index bf4ef54..36af15e 100644
--- a/src/config/agent.rs
+++ b/src/config/agent.rs
@@ -292,9 +292,7 @@ impl AgentDefinition {
for variable in &self.variables {
output = output.replace(&format!("{{{{{}}}}}", variable.name), &variable.value)
}
- for (key, value) in builtin_variables() {
- output = output.replace(&format!("{{{{{}}}}}", key), &value);
- }
+ interpolate_variables(&mut output);
output
}
@@ -412,20 +410,3 @@ fn save_variables(variables_path: &Path, variables: &[AgentVariable]) -> Result<
.with_context(|| format!("Failed to save variables to '{}'", variables_path.display()))?;
Ok(())
}
-
-fn builtin_variables() -> Vec<(&'static str, String)> {
- vec![
- ("__os__", env::consts::OS.to_string()),
- ("__os_family__", env::consts::FAMILY.to_string()),
- ("__arch__", env::consts::ARCH.to_string()),
- ("__shell__", SHELL.name.clone()),
- ("__locale__", sys_locale::get_locale().unwrap_or_default()),
- ("__now__", now()),
- (
- "__cwd__",
- env::current_dir()
- .map(|v| v.display().to_string())
- .unwrap_or_default(),
- ),
- ]
-}
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 5cd266f..dae2018 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -5,7 +5,7 @@ mod session;
pub use self::agent::{list_agents, Agent};
pub use self::input::Input;
-pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE};
+pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE};
use self::session::Session;
use crate::client::{
@@ -774,11 +774,7 @@ impl Config {
let content = read_to_string(&path)?;
Role::new(name, &content)
} else {
- BUILTIN_ROLES
- .iter()
- .find(|v| v.name() == name)
- .cloned()
- .ok_or_else(|| anyhow!("Unknown role `{name}`"))?
+ Role::builtin(name)?
};
match role.model_id() {
Some(model_id) => {
@@ -839,8 +835,11 @@ impl Config {
if role_name == TEMP_ROLE_NAME {
role_name = Text::new("Role name:")
.with_validator(|input: &str| {
- if input.trim().is_empty() {
- Ok(Validation::Invalid("This field is required".into()))
+ let input = input.trim();
+ if input.is_empty() {
+ Ok(Validation::Invalid("This name is required".into()))
+ } else if input == TEMP_ROLE_NAME {
+ Ok(Validation::Invalid("This name is reserved".into()))
} else {
Ok(Validation::Valid)
}
@@ -862,7 +861,7 @@ impl Config {
}
pub fn all_roles() -> Vec<Role> {
- let mut roles: HashMap<String, Role> = BUILTIN_ROLES
+ let mut roles: HashMap<String, Role> = Role::list_builtin_roles()
.iter()
.map(|v| (v.name().to_string(), v.clone()))
.collect();
@@ -894,7 +893,7 @@ impl Config {
}
}
if with_builtin {
- names.extend(BUILTIN_ROLES.iter().map(|v| v.name().to_string()));
+ names.extend(Role::list_builtin_role_names());
}
let mut names: Vec<_> = names.into_iter().collect();
names.sort_unstable();
diff --git a/src/config/role.rs b/src/config/role.rs
index 74d9953..68366d0 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -4,6 +4,7 @@ use crate::client::{Message, MessageContent, MessageRole, Model};
use anyhow::Result;
use fancy_regex::Regex;
+use rust_embed::Embed;
use serde::{Deserialize, Serialize};
use serde_json::Value;
@@ -13,68 +14,11 @@ pub const CODE_ROLE: &str = "%code%";
pub const INPUT_PLACEHOLDER: &str = "__INPUT__";
-lazy_static::lazy_static! {
- pub static ref BUILTIN_ROLES: Vec<Role> = {
- [
- (SHELL_ROLE, shell_prompt()),
- (
- EXPLAIN_SHELL_ROLE,
- r#"Provide a terse, single sentence description of the given shell command.
-Describe each argument and option of the command.
-Provide short responses in about 80 words.
-APPLY MARKDOWN formatting when possible."#
- .into(),
- ),
- (
- CODE_ROLE,
- r#"Provide only code without comments or explanations.
-### INPUT:
-async sleep in js
-### OUTPUT:
-```javascript
-async function timeout(ms) {
- return new Promise(resolve => setTimeout(resolve, ms));
-}
-```
-"#
- .into(),
- ),
- (
- "%create-prompt%",
- r#"As a professional Prompt Engineer, your role is to create effective and innovative prompts for interacting with AI models.
-
-Your core skills include:
-1. **CO-STAR Framework Application**: Utilize the CO-STAR framework to build efficient prompts, ensuring effective communication with large language models.
-2. **Contextual Awareness**: Construct prompts that adapt to complex conversation contexts, ensuring relevant and coherent responses.
-3. **Chain-of-Thought Prompting**: Create prompts that elicit AI models to demonstrate their reasoning process, enhancing the transparency and accuracy of answers.
-4. **Zero-shot Learning**: Design prompts that enable AI models to perform specific tasks without requiring examples, reducing dependence on training data.
-5. **Few-shot Learning**: Guide AI models to quickly learn and execute new tasks through a few examples.
-
-Your output format should include:
-- **Context**: Provide comprehensive background information for the task to ensure the AI understands the specific scenario and offers relevant feedback.
-- **Objective**: Clearly define the task objective, guiding the AI to focus on achieving specific goals.
-- **Style**: Specify writing styles according to requirements, such as imitating a particular person or industry expert.
-- **Tone**: Set an appropriate emotional tone to ensure the AI's response aligns with the expected emotional context.
-- **Audience**: Tailor AI responses for a specific audience, ensuring content appropriateness and ease of understanding.
-- **Response**: Specify output formats for easy execution of downstream tasks, such as lists, JSON, or professional reports.
-- **Workflow**: Instruct the AI on how to step-by-step complete tasks, clarifying inputs, outputs, and specific actions for each step.
-- **Examples**: Show a case of input and output that fits the scenario.
-
-Your workflow should be:
-1. **Analyze User Input**: Extract key information from user requests to determine design objectives.
-2. **Conceive New Prompts**: Based on user needs, create prompts that meet requirements, with each part being professional and detailed.
-3. **Generate Output**: Must only output the newly generated and optimized prompts, without explanation, and without wrapping it in markdown code block."#.into(),
- ),
- ("%functions%", r#"---
-use_tools: all
----
- "#.into()),
- ]
- .into_iter()
- .map(|(name, content)| Role::new(name, &content))
- .collect()
- };
+#[derive(Embed)]
+#[folder = "assets/roles/"]
+struct RolesAsset;
+lazy_static::lazy_static! {
static ref RE_METADATA: Regex = Regex::new(r"(?s)-{3,}\s*(.*?)\s*-{3,}\s*(.*)").unwrap();
}
@@ -122,9 +66,11 @@ impl Role {
prompt = prompt_value.as_str().trim();
}
}
+ let mut prompt = complete_prompt_args(prompt, name);
+ interpolate_variables(&mut prompt);
let mut role = Self {
name: name.to_string(),
- prompt: complete_prompt_args(prompt, name),
+ prompt,
..Default::default()
};
if !metadata.is_empty() {
@@ -145,6 +91,25 @@ impl Role {
role
}
+ pub fn builtin(name: &str) -> Result<Self> {
+ let content = RolesAsset::get(&format!("{name}.md"))
+ .ok_or_else(|| anyhow!("Unknown role `{name}`"))?;
+ let content = unsafe { std::str::from_utf8_unchecked(&content.data) };
+ Ok(Role::new(name, content))
+ }
+
+ pub fn list_builtin_role_names() -> Vec<String> {
+ RolesAsset::iter()
+ .filter_map(|v| v.strip_suffix(".md").map(|v| v.to_string()))
+ .collect()
+ }
+
+ pub fn list_builtin_roles() -> Vec<Self> {
+ RolesAsset::iter()
+ .filter_map(|v| Role::builtin(&v).ok())
+ .collect()
+ }
+
pub fn match_name(names: &[String], name: &str) -> Option<String> {
if names.contains(&name.to_string()) {
Some(name.to_string())
@@ -412,23 +377,6 @@ fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) {
(prompt, vec![])
}
-fn shell_prompt() -> String {
- let os = OS.as_str();
- let shell = SHELL.name.as_str();
- let combinator = if shell == "powershell" {
- "If multiple steps required try to combine them together using ';'.\nIf it already combined with '&&' try to replace it with ';'.".to_string()
- } else {
- "If multiple steps required try to combine them together using '&&'.".to_string()
- };
- format!(
- r#"Provide only {shell} commands for {os} without any description.
-Ensure the output is a valid {shell} command.
-{combinator}
-If there is a lack of details, provide most logical solution.
-Output plain text only, without any markdown formatting."#
- )
-}
-
#[cfg(test)]
mod tests {
use super::*;
diff --git a/src/config/session.rs b/src/config/session.rs
index 476ad51..0bb55e8 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -298,11 +298,14 @@ impl Session {
if !ans {
return Ok(());
}
- while session_name == TEMP_SESSION_NAME {
+ if session_name == TEMP_SESSION_NAME {
session_name = Text::new("Session name:")
.with_validator(|input: &str| {
- if input.trim().is_empty() {
- Ok(Validation::Invalid("This field is required".into()))
+ let input = input.trim();
+ if input.is_empty() {
+ Ok(Validation::Invalid("This name is required".into()))
+ } else if input == TEMP_SESSION_NAME {
+ Ok(Validation::Invalid("This name is reserved".into()))
} else {
Ok(Validation::Valid)
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index b6c71a9..8c68f78 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -326,7 +326,7 @@ impl Repl {
self.config.write().save_session(name)?;
}
_ => {
- println!(r#"Usage: .save session [name]"#)
+ println!(r#"Usage: .save <role|session> [name]"#)
}
}
}
diff --git a/src/utils/command.rs b/src/utils/command.rs
index caf70c7..8afef14 100644
--- a/src/utils/command.rs
+++ b/src/utils/command.rs
@@ -5,24 +5,9 @@ use std::{collections::HashMap, env, ffi::OsStr, path::Path, process::Command};
use anyhow::{anyhow, bail, Context, Result};
lazy_static::lazy_static! {
- pub static ref OS: String = detect_os();
pub static ref SHELL: Shell = detect_shell();
}
-pub fn detect_os() -> String {
- let os = env::consts::OS;
- if os == "linux" {
- if let Ok(contents) = std::fs::read_to_string("/etc/os-release") {
- for line in contents.lines() {
- if let Some(id) = line.strip_prefix("ID=") {
- return format!("{os} ({id})");
- }
- }
- }
- }
- os.to_string()
-}
-
pub struct Shell {
pub name: String,
pub cmd: String,
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 1090c4e..e937e6b 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -8,6 +8,7 @@ mod prompt_input;
mod render_prompt;
mod request;
mod spinner;
+mod variables;
pub use self::abort_signal::*;
pub use self::clipboard::set_text;
@@ -19,6 +20,7 @@ pub use self::prompt_input::*;
pub use self::render_prompt::render_prompt;
pub use self::request::*;
pub use self::spinner::*;
+pub use self::variables::*;
use anyhow::{Context, Result};
use fancy_regex::Regex;
diff --git a/src/utils/variables.rs b/src/utils/variables.rs
new file mode 100644
index 0000000..91c9912
--- /dev/null
+++ b/src/utils/variables.rs
@@ -0,0 +1,33 @@
+use super::*;
+use fancy_regex::{Captures, Regex};
+
+lazy_static::lazy_static! {
+ pub static ref RE_VARIABLE: Regex = Regex::new(r"\{\{(\w+)\}\}").unwrap();
+}
+pub fn interpolate_variables(text: &mut String) {
+ *text = RE_VARIABLE
+ .replace_all(text, |caps: &Captures<'_>| {
+ let key = &caps[1];
+ match key {
+ "__os__" => env::consts::OS.to_string(),
+ "__os_distro__" => {
+ let info = os_info::get();
+ if env::consts::OS == "linux" {
+ format!("{info} (linux)")
+ } else {
+ info.to_string()
+ }
+ }
+ "__os_family__" => env::consts::FAMILY.to_string(),
+ "__arch__" => env::consts::ARCH.to_string(),
+ "__shell__" => SHELL.name.clone(),
+ "__locale__" => sys_locale::get_locale().unwrap_or_default(),
+ "__now__" => now(),
+ "__cwd__" => env::current_dir()
+ .map(|v| v.display().to_string())
+ .unwrap_or_default(),
+ _ => format!("{{{{{}}}}}", key),
+ }
+ })
+ .to_string();
+}