summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs196
-rw-r--r--src/config/role.rs197
-rw-r--r--src/config/session.rs30
-rw-r--r--src/main.rs7
-rw-r--r--src/repl/mod.rs55
-rw-r--r--src/serve.rs2
6 files changed, 315 insertions, 172 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 508b768..72f6d5a 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, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE};
+pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE};
use self::session::Session;
use crate::client::{
@@ -19,7 +19,7 @@ use crate::utils::*;
use anyhow::{anyhow, bail, Context, Result};
use indexmap::IndexMap;
-use inquire::{Confirm, Select};
+use inquire::{validator::Validation, Confirm, Select, Text};
use parking_lot::RwLock;
use serde::Deserialize;
use serde_json::json;
@@ -40,7 +40,7 @@ const DARK_THEME: &[u8] = include_bytes!("../../assets/monokai-extended.theme.bi
const LIGHT_THEME: &[u8] = include_bytes!("../../assets/monokai-extended-light.theme.bin");
const CONFIG_FILE_NAME: &str = "config.yaml";
-const ROLES_FILE_NAME: &str = "roles.yaml";
+const ROLES_DIR_NAME: &str = "roles";
const ENV_FILE_NAME: &str = ".env";
const MESSAGES_FILE_NAME: &str = "messages.md";
const SESSIONS_DIR_NAME: &str = "sessions";
@@ -129,8 +129,6 @@ pub struct Config {
pub clients: Vec<ClientConfig>,
#[serde(skip)]
- pub roles: Vec<Role>,
- #[serde(skip)]
pub role: Option<Role>,
#[serde(skip)]
pub session: Option<Session>,
@@ -195,7 +193,6 @@ impl Default for Config {
clients: vec![],
- roles: vec![],
role: None,
session: None,
rag: None,
@@ -236,7 +233,6 @@ impl Config {
}
config.load_functions()?;
- config.load_roles()?;
config.setup_model()?;
config.setup_document_loaders();
@@ -269,13 +265,17 @@ impl Config {
}
}
- pub fn roles_file() -> Result<PathBuf> {
- match env::var(get_env_name("roles_file")) {
+ pub fn roles_dir() -> Result<PathBuf> {
+ match env::var(get_env_name("roles_dir")) {
Ok(value) => Ok(PathBuf::from(value)),
- Err(_) => Self::local_path(ROLES_FILE_NAME),
+ Err(_) => Self::local_path(ROLES_DIR_NAME),
}
}
+ pub fn role_file(name: &str) -> Result<PathBuf> {
+ Ok(Self::roles_dir()?.join(format!("{name}.md")))
+ }
+
pub fn env_file() -> Result<PathBuf> {
match env::var(get_env_name("env_file")) {
Ok(value) => Ok(PathBuf::from(value)),
@@ -487,7 +487,7 @@ impl Config {
} else if let Some(session) = &self.session {
session.export()
} else if let Some(role) = &self.role {
- role.export()
+ Ok(role.export())
} else if let Some(rag) = &self.rag {
rag.export()
} else {
@@ -531,7 +531,7 @@ impl Config {
("highlight", self.highlight.to_string()),
("light_theme", self.light_theme.to_string()),
("config_file", display_path(&Self::config_file()?)),
- ("roles_file", display_path(&Self::roles_file()?)),
+ ("roles_dir", display_path(&Self::roles_dir()?)),
("env_file", display_path(&Self::env_file()?)),
("functions_dir", display_path(&Self::functions_dir()?)),
("rags_dir", display_path(&Self::rags_dir()?)),
@@ -543,9 +543,9 @@ impl Config {
}
let output = items
.iter()
- .map(|(name, value)| format!("{name:<24}{value}"))
+ .map(|(name, value)| format!("{name:<24}{value}\n"))
.collect::<Vec<String>>()
- .join("\n");
+ .join("");
Ok(output)
}
@@ -716,7 +716,7 @@ impl Config {
pub fn role_info(&self) -> Result<String> {
if let Some(role) = &self.role {
- role.export()
+ Ok(role.export())
} else {
bail!("No role")
}
@@ -733,17 +733,17 @@ impl Config {
}
pub fn retrieve_role(&self, name: &str) -> Result<Role> {
- let mut role = self
- .roles
- .iter()
- .find(|v| v.match_name(name))
- .map(|v| {
- let mut role = v.clone();
- role.complete_prompt_args(name);
- role
- })
- .ok_or_else(|| anyhow!("Unknown role `{name}`"))?;
-
+ let mut role = if Self::list_roles(false).contains(&name.to_string()) {
+ let path = Self::role_file(name)?;
+ 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}`"))?
+ };
match role.model_id() {
Some(model_id) => {
if self.model.id() != model_id {
@@ -758,6 +758,110 @@ impl Config {
Ok(role)
}
+ pub fn new_role(&mut self, name: &str) -> Result<()> {
+ let ans = Confirm::new("Create a new role?")
+ .with_default(true)
+ .prompt()?;
+ if ans {
+ self.upsert_role(name)?;
+ }
+ Ok(())
+ }
+
+ pub fn edit_role(&mut self) -> Result<()> {
+ if let Some(name) = self.role.as_ref().map(|v| v.name().to_string()) {
+ self.upsert_role(&name)
+ } else {
+ bail!("No role")
+ }
+ }
+
+ pub fn upsert_role(&mut self, name: &str) -> Result<()> {
+ let role_path = Self::role_file(name)?;
+ ensure_parent_exists(&role_path)?;
+ let editor = self.editor()?;
+ edit_file(&editor, &role_path)?;
+ self.use_role(name)?;
+ Ok(())
+ }
+
+ pub fn save_role(&mut self, name: Option<&str>) -> Result<()> {
+ let mut role_name = match &self.role {
+ Some(role) => match name {
+ Some(v) => v.to_string(),
+ None => role.name().to_string(),
+ },
+ None => bail!("No role"),
+ };
+ 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()))
+ } else {
+ Ok(Validation::Valid)
+ }
+ })
+ .prompt()?;
+ }
+ let role_path = Self::role_file(&role_name)?;
+ if let Some(role) = self.role.as_mut() {
+ let old_name = role.name().to_string();
+ role.save(&role_name, &role_path, self.working_mode.is_repl())?;
+ if old_name != role_name {
+ if let Ok(path) = Self::role_file(&old_name) {
+ let _ = remove_file(&path);
+ }
+ }
+ }
+
+ Ok(())
+ }
+
+ pub fn all_roles() -> Vec<Role> {
+ let mut roles: HashMap<String, Role> = BUILTIN_ROLES
+ .iter()
+ .map(|v| (v.name().to_string(), v.clone()))
+ .collect();
+ let names = Self::list_roles(false);
+ for name in names {
+ if let Ok(path) = Self::role_file(&name) {
+ if let Ok(content) = read_to_string(&path) {
+ let role = Role::new(&name, &content);
+ roles.insert(name, role);
+ }
+ }
+ }
+ let mut roles: Vec<_> = roles.into_values().collect();
+ roles.sort_unstable_by(|a, b| a.name().cmp(b.name()));
+ roles
+ }
+
+ pub fn list_roles(with_builtin: bool) -> Vec<String> {
+ let mut names = HashSet::new();
+ if let Some(rd) = Self::roles_dir().ok().and_then(|dir| read_dir(dir).ok()) {
+ for entry in rd.flatten() {
+ if let Some(name) = entry
+ .file_name()
+ .to_str()
+ .and_then(|v| v.strip_suffix(".md"))
+ {
+ names.insert(name.to_string());
+ }
+ }
+ }
+ if with_builtin {
+ names.extend(BUILTIN_ROLES.iter().map(|v| v.name().to_string()));
+ }
+ let mut names: Vec<_> = names.into_iter().collect();
+ names.sort_unstable();
+ names
+ }
+
+ pub fn has_role(name: &str) -> bool {
+ Self::list_roles(true).iter().any(|v| v == name)
+ }
+
pub fn use_session(&mut self, session_name: Option<&str>) -> Result<()> {
if self.session.is_some() {
bail!(
@@ -824,16 +928,22 @@ impl Config {
}
pub fn save_session(&mut self, name: Option<&str>) -> Result<()> {
- let name = match &self.session {
+ let session_name = match &self.session {
Some(session) => match name {
Some(v) => v.to_string(),
None => session.name().to_string(),
},
None => bail!("No session"),
};
- let session_path = self.session_file(&name)?;
+ let session_path = self.session_file(&session_name)?;
if let Some(session) = self.session.as_mut() {
- session.save(&session_path, self.working_mode.is_repl())?;
+ let old_name = session.name().to_string();
+ session.save(&session_name, &session_path, self.working_mode.is_repl())?;
+ if old_name != session_name {
+ if let Ok(path) = self.session_file(&old_name) {
+ let _ = remove_file(&path);
+ }
+ }
}
Ok(())
}
@@ -843,9 +953,9 @@ impl Config {
Some(session) => session.name().to_string(),
None => bail!("No session"),
};
- let editor = self.editor()?;
let session_path = self.session_file(&name)?;
self.save_session(Some(&name))?;
+ let editor = self.editor()?;
edit_file(&editor, &session_path).with_context(|| {
format!(
"Failed to edit '{}' with '{editor}'",
@@ -1201,10 +1311,9 @@ impl Config {
let mut filter = "";
if args.len() == 1 {
values = match cmd {
- ".role" => self
- .roles
- .iter()
- .map(|v| (v.name().to_string(), None))
+ ".role" => Self::list_roles(true)
+ .into_iter()
+ .map(|v| (v, None))
.collect(),
".model" => list_chat_models(self)
.into_iter()
@@ -1702,25 +1811,6 @@ impl Config {
Ok(())
}
- fn load_roles(&mut self) -> Result<()> {
- let path = Self::roles_file()?;
- self.roles = if !path.exists() {
- vec![]
- } else {
- let content = read_to_string(&path)
- .with_context(|| format!("Failed to load roles at {}", path.display()))?;
- serde_yaml::from_str(&content).with_context(|| "Invalid roles config")?
- };
- let exist_roles: HashSet<_> = self.roles.iter().map(|v| v.name().to_string()).collect();
- let builtin_roles = Role::builtin();
- for role in builtin_roles {
- if !exist_roles.contains(role.name()) {
- self.roles.push(role);
- }
- }
- Ok(())
- }
-
fn setup_model(&mut self) -> Result<()> {
let mut model_id = self.model_id.clone();
if model_id.is_empty() {
diff --git a/src/config/role.rs b/src/config/role.rs
index 31e894b..c61aa53 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -2,15 +2,57 @@ use super::*;
use crate::client::{Message, MessageContent, MessageRole, Model};
-use anyhow::{Context, Result};
+use anyhow::Result;
+use fancy_regex::Regex;
use serde::{Deserialize, Serialize};
+use serde_json::Value;
pub const SHELL_ROLE: &str = "%shell%";
pub const EXPLAIN_SHELL_ROLE: &str = "%explain-shell%";
pub const CODE_ROLE: &str = "%code%";
+pub const FUNCTIONS_ROLE: &str = "%functions%";
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(),
+ ),
+ (FUNCTIONS_ROLE, r#"---
+use_tools: all
+---
+ "#.into()),
+ ]
+ .into_iter()
+ .map(|(name, content)| Role::new(name, &content))
+ .collect()
+ };
+
+ static ref RE_METADATA: Regex = Regex::new(r"(?s)-{3,}\s*(.*?)\s*-{3,}\s*(.*)").unwrap();
+}
+
pub trait RoleLike {
fn to_role(&self) -> Role;
fn model(&self) -> &Model;
@@ -46,57 +88,82 @@ pub struct Role {
}
impl Role {
- pub fn new(name: &str, prompt: &str) -> Self {
- Self {
- name: name.into(),
- prompt: prompt.into(),
+ pub fn new(name: &str, content: &str) -> Self {
+ let mut metadata = "";
+ let mut prompt = content.trim();
+ if let Ok(Some(caps)) = RE_METADATA.captures(content) {
+ if let (Some(metadata_value), Some(prompt_value)) = (caps.get(1), caps.get(2)) {
+ metadata = metadata_value.as_str().trim();
+ prompt = prompt_value.as_str().trim();
+ }
+ }
+ let mut role = Self {
+ name: name.to_string(),
+ prompt: prompt.to_string(),
..Default::default()
+ };
+ if !metadata.is_empty() {
+ if let Ok(value) = serde_yaml::from_str::<Value>(metadata) {
+ if let Some(value) = value.as_object() {
+ for (key, value) in value {
+ match key.as_str() {
+ "model" => role.model_id = value.as_str().map(|v| v.to_string()),
+ "temperature" => role.temperature = value.as_f64(),
+ "top_p" => role.top_p = value.as_f64(),
+ "use_tools" => role.use_tools = value.as_str().map(|v| v.to_string()),
+ _ => (),
+ }
+ }
+ }
+ }
}
+ role
}
- pub fn builtin() -> Vec<Role> {
- [
- (SHELL_ROLE, shell_prompt(), None),
- (
- 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(),
- None,
- ),
- (
- 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(),
- None,
- ),
- ("%functions%", String::new(), Some("all".into())),
- ]
- .into_iter()
- .map(|(name, prompt, use_tools)| Self {
- name: name.into(),
- prompt,
- use_tools,
- ..Default::default()
- })
- .collect()
+ pub fn export(&self) -> String {
+ let mut metadata = vec![];
+ if let Some(model) = self.model_id() {
+ metadata.push(format!("model: {}", model));
+ }
+ if let Some(temperature) = self.temperature() {
+ metadata.push(format!("temperature: {}", temperature));
+ }
+ if let Some(top_p) = self.top_p() {
+ metadata.push(format!("top_p: {}", top_p));
+ }
+ if let Some(use_tools) = self.use_tools() {
+ metadata.push(format!("use_tools: {}", use_tools));
+ }
+ if metadata.is_empty() {
+ format!("{}\n", self.prompt)
+ } else if self.prompt.is_empty() {
+ format!("---\n{}\n---\n", metadata.join("\n"))
+ } else {
+ format!("---\n{}\n---\n\n{}\n", metadata.join("\n"), self.prompt)
+ }
}
- pub fn export(&self) -> Result<String> {
- let output = serde_yaml::to_string(&self)
- .with_context(|| format!("Unable to show info about role {}", &self.name))?;
- Ok(output.trim_end().to_string())
+ pub fn save(&mut self, role_name: &str, role_path: &Path, is_repl: bool) -> Result<()> {
+ ensure_parent_exists(role_path)?;
+
+ let content = self.export();
+ std::fs::write(role_path, content).with_context(|| {
+ format!(
+ "Failed to write role {} to {}",
+ self.name,
+ role_path.display()
+ )
+ })?;
+
+ if is_repl {
+ println!("✨ Saved role to '{}'", role_path.display());
+ }
+
+ if role_name != self.name {
+ self.name = role_name.to_string();
+ }
+
+ Ok(())
}
pub fn sync<T: RoleLike>(&mut self, role_like: &T) {
@@ -150,21 +217,6 @@ async function timeout(ms) {
self.prompt.contains(INPUT_PLACEHOLDER)
}
- pub fn complete_prompt_args(&mut self, name: &str) {
- self.name = name.to_string();
- self.prompt = complete_prompt_args(&self.prompt, &self.name);
- }
-
- pub fn match_name(&self, name: &str) -> bool {
- if self.name.contains(':') {
- let role_name_parts: Vec<&str> = self.name.split(':').collect();
- let name_parts: Vec<&str> = name.split(':').collect();
- role_name_parts[0] == name_parts[0] && role_name_parts.len() == name_parts.len()
- } else {
- self.name == name
- }
- }
-
pub fn echo_messages(&self, input: &Input) -> String {
let input_markdown = input.render();
if self.is_empty_prompt() {
@@ -239,6 +291,9 @@ impl RoleLike for Role {
}
fn set_model(&mut self, model: &Model) {
+ if !self.model().id().is_empty() {
+ self.model_id = Some(model.id().to_string());
+ }
self.model = model.clone();
}
@@ -255,14 +310,6 @@ impl RoleLike for Role {
}
}
-fn complete_prompt_args(prompt: &str, name: &str) -> String {
- let mut prompt = prompt.trim().to_string();
- for (i, arg) in name.split(':').skip(1).enumerate() {
- prompt = prompt.replace(&format!("__ARG{}__", i + 1), arg);
- }
- prompt
-}
-
fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) {
let mut text = prompt;
let mut search_input = true;
@@ -333,18 +380,6 @@ mod tests {
use super::*;
#[test]
- fn test_merge_prompt_name() {
- assert_eq!(
- complete_prompt_args("convert __ARG1__", "convert:foo"),
- "convert foo"
- );
- assert_eq!(
- complete_prompt_args("convert __ARG1__ to __ARG2__", "convert:foo:bar"),
- "convert foo to bar"
- );
- }
-
- #[test]
fn test_parse_structure_prompt1() {
let prompt = r#"
System message
diff --git a/src/config/session.rs b/src/config/session.rs
index eccffa9..c36c9db 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -9,7 +9,7 @@ use inquire::{validator::Validation, Confirm, Text};
use serde::{Deserialize, Serialize};
use serde_json::json;
use std::collections::HashMap;
-use std::fs::{self, read_to_string};
+use std::fs::{read_to_string, write};
use std::path::Path;
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
@@ -86,10 +86,6 @@ impl Session {
Ok(session)
}
- pub fn is_temp(&self) -> bool {
- self.name == TEMP_SESSION_NAME
- }
-
pub fn is_empty(&self) -> bool {
self.messages.is_empty() && self.compressed_messages.is_empty()
}
@@ -218,12 +214,7 @@ impl Session {
}
}
- if lines.last() == Some(&String::new()) {
- lines.pop();
- }
-
- let output = lines.join("\n");
- Ok(output)
+ Ok(lines.join("\n"))
}
pub fn tokens_usage(&self) -> (usize, f32) {
@@ -302,6 +293,7 @@ impl Session {
pub fn exit(&mut self, session_dir: &Path, is_repl: bool) -> Result<()> {
let save_session = self.save_session();
if self.dirty && save_session != Some(false) {
+ let mut session_name = self.name().to_string();
if save_session.is_none() {
if !is_repl {
return Ok(());
@@ -310,8 +302,8 @@ impl Session {
if !ans {
return Ok(());
}
- if self.is_temp() {
- self.name = Text::new("Session name:")
+ while 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()))
@@ -322,20 +314,20 @@ impl Session {
.prompt()?;
}
}
- let session_path = session_dir.join(format!("{}.yaml", self.name()));
- self.save(&session_path, is_repl)?;
+ let session_path = session_dir.join(format!("{session_name}.yaml"));
+ self.save(&session_name, &session_path, is_repl)?;
}
Ok(())
}
- pub fn save(&mut self, session_path: &Path, is_repl: bool) -> Result<()> {
+ pub fn save(&mut self, session_name: &str, session_path: &Path, is_repl: bool) -> Result<()> {
ensure_parent_exists(session_path)?;
self.path = Some(session_path.display().to_string());
let content = serde_yaml::to_string(&self)
.with_context(|| format!("Failed to serde session {}", self.name))?;
- fs::write(session_path, content).with_context(|| {
+ write(session_path, content).with_context(|| {
format!(
"Failed to write session {} to {}",
self.name,
@@ -347,6 +339,10 @@ impl Session {
println!("✨ Saved session to '{}'", session_path.display());
}
+ if self.name() != session_name {
+ self.name = session_name.to_string()
+ }
+
self.dirty = false;
Ok(())
diff --git a/src/main.rs b/src/main.rs
index 5c22f7d..83d2536 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -78,11 +78,8 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()>
return Ok(());
}
if cli.list_roles {
- config
- .read()
- .roles
- .iter()
- .for_each(|v| println!("{}", v.name()));
+ let roles = Config::list_roles(true).join("\n");
+ println!("{roles}");
return Ok(());
}
if cli.list_agents {
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 490cc89..0d9353b 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -31,7 +31,7 @@ lazy_static::lazy_static! {
const MENU_NAME: &str = "completion_menu";
lazy_static::lazy_static! {
- static ref REPL_COMMANDS: [ReplCommand; 28] = [
+ static ref REPL_COMMANDS: [ReplCommand; 30] = [
ReplCommand::new(".help", "Show this help message", AssertState::pass()),
ReplCommand::new(".info", "View system info", AssertState::pass()),
ReplCommand::new(".model", "Change the current LLM", AssertState::pass()),
@@ -42,7 +42,7 @@ lazy_static::lazy_static! {
),
ReplCommand::new(
".role",
- "Switch to a specific role",
+ "Create or switch to a specific role",
AssertState::False(StateFlags::SESSION | StateFlags::AGENT)
),
ReplCommand::new(
@@ -51,6 +51,16 @@ lazy_static::lazy_static! {
AssertState::True(StateFlags::ROLE),
),
ReplCommand::new(
+ ".edit role",
+ "Edit the current role",
+ AssertState::TrueFalse(StateFlags::ROLE, StateFlags::SESSION_EMPTY | StateFlags::SESSION),
+ ),
+ ReplCommand::new(
+ ".save role",
+ "Save the current role to file",
+ AssertState::True(StateFlags::ROLE)
+ ),
+ ReplCommand::new(
".exit role",
"Leave the role",
AssertState::True(StateFlags::ROLE),
@@ -66,13 +76,8 @@ lazy_static::lazy_static! {
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION),
),
ReplCommand::new(
- ".save session",
- "Save the current session to file",
- AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
- ),
- ReplCommand::new(
".edit session",
- "Edit the current session with an editor",
+ "Edit the current session",
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
),
ReplCommand::new(
@@ -81,6 +86,11 @@ lazy_static::lazy_static! {
AssertState::True(StateFlags::SESSION)
),
ReplCommand::new(
+ ".save session",
+ "Save the current session to file",
+ AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
+ ),
+ ReplCommand::new(
".exit session",
"End the session",
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
@@ -219,24 +229,24 @@ impl Repl {
".info" => match args {
Some("role") => {
let info = self.config.read().role_info()?;
- println!("{}", info);
+ print!("{}", info);
}
Some("session") => {
let info = self.config.read().session_info()?;
- println!("{}", info);
+ print!("{}", info);
}
Some("rag") => {
let info = self.config.read().rag_info()?;
- println!("{}", info);
+ print!("{}", info);
}
Some("agent") => {
let info = self.config.read().agent_info()?;
- println!("{}", info);
+ print!("{}", info);
}
Some(_) => unknown_command()?,
None => {
let output = self.config.read().sysinfo()?;
- println!("{}", output);
+ print!("{}", output);
}
},
".model" => match args {
@@ -259,10 +269,19 @@ impl Repl {
ask(&self.config, self.abort_signal.clone(), input, false).await?;
}
None => {
- self.config.write().use_role(args)?;
+ let name = args;
+ if Config::has_role(name) {
+ self.config.write().use_role(name)?;
+ } else {
+ self.config.write().new_role(name)?;
+ }
}
},
- None => println!(r#"Usage: .role <name> [text]..."#),
+ None => println!(
+ r#"Usage:
+ .role <name> # If the role exists, switch to it; otherwise, create a new role
+ .role <name> [text]... # Temporarily switch to the role, send the text, and switch back"#
+ ),
},
".session" => {
self.config.write().use_session(args)?;
@@ -300,6 +319,9 @@ impl Repl {
Some((subcmd, args)) => (subcmd, Some(args.trim())),
None => (v, None),
}) {
+ Some(("role", name)) => {
+ self.config.write().save_role(name)?;
+ }
Some(("session", name)) => {
self.config.write().save_session(name)?;
}
@@ -313,6 +335,9 @@ impl Repl {
Some((subcmd, args)) => (subcmd, Some(args.trim())),
None => (v, None),
}) {
+ Some(("role", _)) => {
+ self.config.write().edit_role()?;
+ }
Some(("session", _)) => {
self.config.write().edit_session()?;
}
diff --git a/src/serve.rs b/src/serve.rs
index 0b0156d..65ac4cf 100644
--- a/src/serve.rs
+++ b/src/serve.rs
@@ -76,7 +76,7 @@ impl Server {
let config = config.read();
let clients = config.clients.clone();
let model = config.model.clone();
- let roles = config.roles.clone();
+ let roles = Config::all_roles();
let mut models = list_models(&config);
let mut default_model = model.clone();
default_model.data_mut().name = DEFAULT_MODEL_NAME.into();