diff options
| author | sigoden <sigoden@gmail.com> | 2024-12-01 07:28:12 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-12-01 07:28:12 +0800 |
| commit | 431d16363ede42f299dbd32c4e38390713620809 (patch) | |
| tree | d49f1a25ea152ab9e6be9117b67f0acbe9115164 /src | |
| parent | 50bbfc4541bf2cb2ecb28f2fed3b9b3ab7d5005b (diff) | |
| download | aichat-431d16363ede42f299dbd32c4e38390713620809.tar.gz | |
feat: agent supports dynamic instructions (#1023)
* feat: agent supports dynamic instructions
* change tool calls' null output to 'TODO'
* REPL don't print banner if use agent/rag
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/agent.rs | 100 | ||||
| -rw-r--r-- | src/config/mod.rs | 82 | ||||
| -rw-r--r-- | src/config/session.rs | 12 | ||||
| -rw-r--r-- | src/function.rs | 33 | ||||
| -rw-r--r-- | src/main.rs | 3 | ||||
| -rw-r--r-- | src/repl/mod.rs | 27 |
6 files changed, 179 insertions, 78 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index 29121ee..87ac719 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -1,6 +1,9 @@ use super::*; -use crate::{client::Model, function::Functions}; +use crate::{ + client::Model, + function::{run_llm_function, Functions}, +}; use anyhow::{Context, Result}; use inquire::{validator::Validation, Text}; @@ -22,6 +25,10 @@ pub struct Agent { #[serde(skip)] session_variables: Option<AgentVariables>, #[serde(skip)] + shared_dynamic_instructions: Option<String>, + #[serde(skip)] + session_dynamic_instructions: Option<String>, + #[serde(skip)] functions: Functions, #[serde(skip)] rag: Option<Arc<Rag>>, @@ -102,6 +109,8 @@ impl Agent { definition, shared_variables: Default::default(), session_variables: None, + shared_dynamic_instructions: None, + session_dynamic_instructions: None, functions, rag, model, @@ -215,7 +224,16 @@ impl Agent { } pub fn interpolated_instructions(&self) -> String { - self.definition.interpolated_instructions(self.variables()) + let mut output = self + .session_dynamic_instructions + .clone() + .or_else(|| self.shared_dynamic_instructions.clone()) + .unwrap_or_else(|| self.definition.instructions.clone()); + for (k, v) in self.variables() { + output = output.replace(&format!("{{{{{k}}}}}"), v) + } + interpolate_variables(&mut output); + output } pub fn agent_prelude(&self) -> Option<&str> { @@ -257,11 +275,8 @@ impl Agent { self.shared_variables = shared_variables; } - pub fn set_session_variables(&mut self, session_variables: Option<AgentVariables>) { - if self.shared_variables.is_empty() { - self.shared_variables = session_variables.clone().unwrap_or_default(); - } - self.session_variables = session_variables; + pub fn set_session_variables(&mut self, session_variables: AgentVariables) { + self.session_variables = Some(session_variables); } pub fn set_variable(&mut self, key: &str, value: &str) -> Result<()> { @@ -269,16 +284,67 @@ impl Agent { Some(v) => v, None => &mut self.shared_variables, }; - if !variables.contains_key(key) { + let Some(old_value) = variables.get(key) else { bail!("Unknown variable: '{key}'") + }; + if old_value == value { + return Ok(()); } variables.insert(key.to_string(), value.to_string()); + if self.session_variables.is_some() { + self.update_session_dynamic_instructions(None)?; + } else { + self.update_shared_dynamic_instructions(true)?; + } + Ok(()) } pub fn defined_variables(&self) -> &[AgentVariable] { &self.definition.variables } + + pub fn exit_session(&mut self) { + self.session_variables = None; + self.session_dynamic_instructions = None; + } + + pub fn is_dynamic_instructions(&self) -> bool { + self.definition.dynamic_instructions + } + + pub fn update_shared_dynamic_instructions(&mut self, force: bool) -> Result<()> { + if self.is_dynamic_instructions() && (force || self.shared_dynamic_instructions.is_none()) { + self.shared_dynamic_instructions = Some(self.run_instructions_fn()?); + } + Ok(()) + } + + pub fn update_session_dynamic_instructions(&mut self, value: Option<String>) -> Result<()> { + if self.is_dynamic_instructions() { + let value = match value { + Some(v) => v, + None => self.run_instructions_fn()?, + }; + self.session_dynamic_instructions = Some(value); + } + Ok(()) + } + + fn run_instructions_fn(&self) -> Result<String> { + let value = run_llm_function( + self.name().to_string(), + vec!["_instructions".into(), "{}".into()], + self.variable_envs(), + )?; + match value { + Some(v) => { + println!(); + Ok(v) + } + _ => bail!("No return value from '_instructions' function"), + } + } } impl RoleLike for Agent { @@ -331,11 +397,15 @@ impl RoleLike for Agent { pub struct AgentConfig { #[serde(rename(serialize = "model", deserialize = "model"))] pub model_id: Option<String>, + #[serde(skip_serializing_if = "Option::is_none")] pub temperature: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] pub top_p: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] pub use_tools: Option<String>, + #[serde(skip_serializing_if = "Option::is_none")] pub agent_prelude: Option<String>, - #[serde(default)] + #[serde(default, skip_serializing_if = "IndexMap::is_empty")] pub variables: AgentVariables, } @@ -389,8 +459,11 @@ pub struct AgentDefinition { pub description: String, #[serde(default)] pub version: String, + #[serde(default)] pub instructions: String, #[serde(default)] + pub dynamic_instructions: bool, + #[serde(default)] pub variables: Vec<AgentVariable>, #[serde(default)] pub conversation_starters: Vec<String>, @@ -436,15 +509,6 @@ impl AgentDefinition { ) } - fn interpolated_instructions(&self, variables: &AgentVariables) -> String { - let mut output = self.instructions.clone(); - for (k, v) in variables { - output = output.replace(&format!("{{{{{k}}}}}"), v) - } - interpolate_variables(&mut output); - output - } - fn replace_tools_placeholder(&mut self, functions: &Functions) { let tools_placeholder: &str = "{{__tools__}}"; if self.instructions.contains(tools_placeholder) { diff --git a/src/config/mod.rs b/src/config/mod.rs index 31337cd..39c63a3 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1076,9 +1076,6 @@ impl Config { session.exit(&sessions_dir, self.working_mode.is_repl())?; self.last_message = None; } - if let Some(agent) = self.agent.as_mut() { - agent.set_session_variables(None); - } Ok(()) } @@ -1122,7 +1119,7 @@ impl Config { pub fn empty_session(&mut self) -> Result<()> { if let Some(session) = self.session.as_mut() { if let Some(agent) = self.agent.as_ref() { - session.sync_agent(agent, false); + session.sync_agent(agent); } session.clear_messages(); } else { @@ -1496,7 +1493,7 @@ impl Config { let value = parts[1]; agent.set_variable(key, value)?; if let Some(session) = self.session.as_mut() { - session.sync_agent(agent, true); + session.sync_agent(agent); } } None => bail!("No agent"), @@ -1513,6 +1510,17 @@ impl Config { Ok(()) } + pub fn exit_agent_session(&mut self) -> Result<()> { + self.exit_session()?; + if let Some(agent) = self.agent.as_mut() { + agent.exit_session(); + if self.working_mode.is_repl() { + self.init_agent_shared_variables()?; + } + } + Ok(()) + } + pub fn apply_prelude(&mut self) -> Result<()> { if !self.state().is_empty() { return Ok(()); @@ -1986,12 +1994,15 @@ impl Config { Some(v) => v, None => return Ok(()), }; - let new_variables = Agent::init_agent_variables( - agent.defined_variables(), - agent.config_variables(), - self.print_info_only, - )?; - agent.set_shared_variables(new_variables); + if !agent.defined_variables().is_empty() && agent.shared_variables().is_empty() { + let new_variables = Agent::init_agent_variables( + agent.defined_variables(), + agent.config_variables(), + self.print_info_only, + )?; + agent.set_shared_variables(new_variables); + } + agent.update_shared_dynamic_instructions(false)?; Ok(()) } @@ -2000,20 +2011,30 @@ impl Config { (Some(agent), Some(session)) => (agent, session), _ => return Ok(()), }; - let shared_variables = agent.shared_variables(); - let mut all_variables = if shared_variables.is_empty() { - agent.config_variables().clone() + if session.is_empty() { + let shared_variables = agent.shared_variables().clone(); + let session_variables = + if !agent.defined_variables().is_empty() && shared_variables.is_empty() { + let new_variables = Agent::init_agent_variables( + agent.defined_variables(), + agent.config_variables(), + self.print_info_only, + )?; + agent.set_shared_variables(new_variables.clone()); + new_variables + } else { + shared_variables + }; + agent.set_session_variables(session_variables); + agent.update_session_dynamic_instructions(None)?; + session.sync_agent(agent); } else { - shared_variables.clone() - }; - all_variables.extend(session.agent_variables().clone()); - let new_variables = Agent::init_agent_variables( - agent.defined_variables(), - &all_variables, - self.print_info_only, - )?; - agent.set_session_variables(Some(new_variables)); - session.sync_agent(agent, false); + let variables = session.agent_variables(); + agent.set_session_variables(variables.clone()); + agent.update_session_dynamic_instructions(Some( + session.agent_instructions().to_string(), + ))?; + } Ok(()) } @@ -2306,9 +2327,22 @@ impl AssertState { pub fn pass() -> Self { AssertState::False(StateFlags::empty()) } + pub fn bare() -> Self { AssertState::Equal(StateFlags::empty()) } + + pub fn assert(self, flags: StateFlags) -> bool { + match self { + AssertState::True(true_flags) => true_flags & flags != StateFlags::empty(), + AssertState::False(false_flags) => false_flags & flags == StateFlags::empty(), + AssertState::TrueFalse(true_flags, false_flags) => { + (true_flags & flags != StateFlags::empty()) + && (false_flags & flags == StateFlags::empty()) + } + AssertState::Equal(check_flags) => check_flags == flags, + } + } } fn create_config_file(config_path: &Path) -> Result<()> { diff --git a/src/config/session.rs b/src/config/session.rs index 4b3da76..820fa12 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -36,6 +36,8 @@ pub struct Session { role_name: Option<String>, #[serde(default, skip_serializing_if = "IndexMap::is_empty")] agent_variables: AgentVariables, + #[serde(default, skip_serializing_if = "String::is_empty")] + agent_instructions: String, #[serde(default, skip_serializing_if = "Vec::is_empty")] compressed_messages: Vec<Message>, @@ -275,19 +277,21 @@ impl Session { self.role_prompt.clear(); } - pub fn sync_agent(&mut self, agent: &Agent, set_dirty: bool) { + pub fn sync_agent(&mut self, agent: &Agent) { self.role_name = None; self.role_prompt = agent.interpolated_instructions(); self.agent_variables = agent.variables().clone(); - if set_dirty { - self.dirty = true; - } + self.agent_instructions = self.role_prompt.clone(); } pub fn agent_variables(&self) -> &AgentVariables { &self.agent_variables } + pub fn agent_instructions(&self) -> &str { + &self.agent_instructions + } + pub fn set_save_session(&mut self, value: Option<bool>) { if self.save_session != value { self.save_session = value; diff --git a/src/function.rs b/src/function.rs index 698c137..6f407d6 100644 --- a/src/function.rs +++ b/src/function.rs @@ -27,17 +27,22 @@ pub fn eval_tool_calls(config: &GlobalConfig, mut calls: Vec<ToolCall>) -> Resul if calls.is_empty() { bail!("The request was aborted because an infinite loop of function calls was detected.") } + let mut is_all_null = true; for call in calls { - let result = call.eval(config)?; + let mut result = call.eval(config)?; + if result.is_null() { + result = json!("DONE"); + } else { + is_all_null = false; + } output.push(ToolResult::new(call, result)); } + if is_all_null { + output = vec![]; + } Ok(output) } -pub fn need_send_tool_results(arr: &[ToolResult]) -> bool { - arr.iter().any(|v| !v.output.is_null()) -} - #[derive(Debug, Clone, Deserialize, Serialize)] pub struct ToolResult { pub call: ToolCall, @@ -159,17 +164,16 @@ impl ToolCall { pub fn eval(&self, config: &GlobalConfig) -> Result<Value> { let function_name = self.name.clone(); - let (call_name, cmd_name, mut cmd_args, envs, agent_name) = match &config.read().agent { + let (call_name, cmd_name, mut cmd_args, envs) = match &config.read().agent { Some(agent) => match agent.functions().find(&function_name) { Some(function) => { let agent_name = agent.name().to_string(); if function.agent { ( format!("{agent_name}-{function_name}"), - agent_name.clone(), + agent_name, vec![function_name], agent.variable_envs(), - Some(agent_name), ) } else { ( @@ -177,7 +181,6 @@ impl ToolCall { function_name, vec![], Default::default(), - Some(agent_name), ) } } @@ -189,7 +192,6 @@ impl ToolCall { function_name, vec![], Default::default(), - None, ), false => bail!("Unexpected call: {function_name} {}", self.arguments), }, @@ -210,10 +212,10 @@ impl ToolCall { cmd_args.push(json_data.to_string()); - let output = match run_llm_function(cmd_name, cmd_args, envs, agent_name)? { + let output = match run_llm_function(cmd_name, cmd_args, envs)? { Some(contents) => serde_json::from_str(&contents) .ok() - .unwrap_or_else(|| json!({"result": contents})), + .unwrap_or_else(|| json!({"output": contents})), None => Value::Null, }; @@ -221,17 +223,16 @@ impl ToolCall { } } -fn run_llm_function( +pub fn run_llm_function( cmd_name: String, cmd_args: Vec<String>, mut envs: HashMap<String, String>, - agent_name: Option<String>, ) -> Result<Option<String>> { let prompt = format!("Call {cmd_name} {}", cmd_args.join(" ")); let mut bin_dirs: Vec<PathBuf> = vec![]; - if let Some(agent_name) = agent_name { - let dir = Config::agent_functions_dir(&agent_name).join("bin"); + if cmd_args.len() > 1 { + let dir = Config::agent_functions_dir(&cmd_name).join("bin"); if dir.exists() { bin_dirs.push(dir); } diff --git a/src/main.rs b/src/main.rs index c41b1e7..d8f66f8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -18,7 +18,6 @@ use crate::config::{ ensure_parent_exists, list_agents, load_env_file, Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, TEMP_SESSION_NAME, }; -use crate::function::need_send_tool_results; use crate::render::render_error; use crate::repl::Repl; use crate::utils::*; @@ -185,7 +184,7 @@ async fn start_directive( .write() .after_chat_completion(&input, &output, &tool_results)?; - if need_send_tool_results(&tool_results) { + if !tool_results.is_empty() { start_directive( config, input.merge_tool_results(output, tool_results), diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 4519e23..f3eab56 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -8,7 +8,6 @@ use self::prompt::ReplPrompt; use crate::client::{call_chat_completions, call_chat_completions_streaming}; use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; -use crate::function::need_send_tool_results; use crate::render::render_error; use crate::utils::{ abortable_run_with_spinner, create_abort_signal, set_text, temp_file, AbortSignal, @@ -195,7 +194,11 @@ impl Repl { } pub async fn run(&mut self) -> Result<()> { - self.banner(); + if AssertState::False(StateFlags::AGENT | StateFlags::RAG) + .assert(self.config.read().state()) + { + self.banner(); + } loop { if self.abort_signal.aborted_ctrld() { @@ -228,7 +231,7 @@ impl Repl { _ => {} } } - self.handle(".exit session").await?; + self.config.write().exit_session()?; Ok(()) } @@ -459,7 +462,11 @@ impl Repl { self.config.write().exit_role()?; } Some("session") => { - self.config.write().exit_session()?; + if self.config.read().agent.is_some() { + self.config.write().exit_agent_session()?; + } else { + self.config.write().exit_session()?; + } } Some("rag") => { self.config.write().exit_rag()?; @@ -603,15 +610,7 @@ impl ReplCommand { } fn is_valid(&self, flags: StateFlags) -> bool { - match self.state { - AssertState::True(true_flags) => true_flags & flags != StateFlags::empty(), - AssertState::False(false_flags) => false_flags & flags == StateFlags::empty(), - AssertState::TrueFalse(true_flags, false_flags) => { - (true_flags & flags != StateFlags::empty()) - && (false_flags & flags == StateFlags::empty()) - } - AssertState::Equal(check_flags) => check_flags == flags, - } + self.state.assert(flags) } } @@ -656,7 +655,7 @@ async fn ask( config .write() .after_chat_completion(&input, &output, &tool_results)?; - if need_send_tool_results(&tool_results) { + if !tool_results.is_empty() { ask( config, abort_signal, |
