From 6b6413b02df588faf6557053751b7e9c591539f7 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 1 Dec 2024 11:35:12 +0800 Subject: feat: add cli option `--agent-variable ` (#1027) --- src/cli.rs | 5 ++++- src/config/agent.rs | 16 ++++++---------- src/config/mod.rs | 36 +++++++++++++++++++++++++++--------- src/main.rs | 11 ++++++++++- 4 files changed, 47 insertions(+), 21 deletions(-) (limited to 'src') diff --git a/src/cli.rs b/src/cli.rs index 21734b8..c3af5cc 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -24,8 +24,11 @@ pub struct Cli { /// Start a agent #[clap(short = 'a', long)] pub agent: Option, + /// Set agent variables + #[clap(long, value_names = ["NAME", "VALUE"], num_args = 2)] + pub agent_variable: Vec, /// Start a RAG - #[clap(short = 'R', long)] + #[clap(long)] pub rag: Option, /// Serve the LLM API and WebAPP #[clap(long, value_name = "ADDRESS")] diff --git a/src/config/agent.rs b/src/config/agent.rs index 87ac719..5e54b26 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -15,24 +15,17 @@ const DEFAULT_AGENT_NAME: &str = "rag"; pub type AgentVariables = IndexMap; -#[derive(Debug, Clone, Serialize)] +#[derive(Debug, Clone)] pub struct Agent { name: String, config: AgentConfig, definition: AgentDefinition, - #[serde(skip)] shared_variables: AgentVariables, - #[serde(skip)] session_variables: Option, - #[serde(skip)] shared_dynamic_instructions: Option, - #[serde(skip)] session_dynamic_instructions: Option, - #[serde(skip)] functions: Functions, - #[serde(skip)] rag: Option>, - #[serde(skip)] model: Model, } @@ -75,7 +68,7 @@ impl Agent { let rag = if rag_path.exists() { Some(Arc::new(Rag::load(config, DEFAULT_AGENT_NAME, &rag_path)?)) - } else if !definition.documents.is_empty() && !config.read().print_info_only { + } else if !definition.documents.is_empty() && !config.read().cli_info_flag { let mut ans = false; if *IS_STDOUT_TERMINAL { ans = Confirm::new("The agent has the documents, init RAG?") @@ -182,11 +175,14 @@ impl Agent { pub fn export(&self) -> Result { let mut agent = self.clone(); agent.definition.instructions = self.interpolated_instructions(); - let mut value = serde_json::json!(agent); + let mut value = json!({}); + value["name"] = json!(self.name()); let variables = self.variables(); if !variables.is_empty() { value["variables"] = serde_json::to_value(variables)?; } + value["config"] = json!(self.config); + value["definition"] = json!(self.definition); value["functions_dir"] = Config::agent_functions_dir(&self.name) .display() .to_string() diff --git a/src/config/mod.rs b/src/config/mod.rs index e0e1efa..c98c81d 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -156,9 +156,12 @@ pub struct Config { #[serde(skip)] pub working_mode: WorkingMode, #[serde(skip)] - pub print_info_only: bool, - #[serde(skip)] pub last_message: Option<(Input, String)>, + + #[serde(skip)] + pub cli_info_flag: bool, + #[serde(skip)] + pub cli_agent_variables: Option, } impl Default for Config { @@ -218,8 +221,10 @@ impl Default for Config { model: Default::default(), functions: Default::default(), working_mode: WorkingMode::Cmd, - print_info_only: false, last_message: None, + + cli_info_flag: false, + cli_agent_variables: None, } } } @@ -1508,6 +1513,7 @@ impl Config { if self.agent.take().is_some() { self.rag.take(); self.last_message = None; + self.cli_agent_variables = None; } Ok(()) } @@ -1997,14 +2003,20 @@ impl Config { None => return Ok(()), }; if !agent.defined_variables().is_empty() && agent.shared_variables().is_empty() { + let mut config_variables = agent.config_variables().clone(); + if let Some(v) = &self.cli_agent_variables { + config_variables.extend(v.clone()); + } let new_variables = Agent::init_agent_variables( agent.defined_variables(), - agent.config_variables(), - self.print_info_only, + &config_variables, + self.cli_info_flag, )?; agent.set_shared_variables(new_variables); } - agent.update_shared_dynamic_instructions(false)?; + if !self.cli_info_flag { + agent.update_shared_dynamic_instructions(false)?; + } Ok(()) } @@ -2017,10 +2029,14 @@ impl Config { let shared_variables = agent.shared_variables().clone(); let session_variables = if !agent.defined_variables().is_empty() && shared_variables.is_empty() { + let mut config_variables = agent.config_variables().clone(); + if let Some(v) = &self.cli_agent_variables { + config_variables.extend(v.clone()); + } let new_variables = Agent::init_agent_variables( agent.defined_variables(), - agent.config_variables(), - self.print_info_only, + &config_variables, + self.cli_info_flag, )?; agent.set_shared_variables(new_variables.clone()); new_variables @@ -2028,7 +2044,9 @@ impl Config { shared_variables }; agent.set_session_variables(session_variables); - agent.update_session_dynamic_instructions(None)?; + if !self.cli_info_flag { + agent.update_session_dynamic_instructions(None)?; + } session.sync_agent(agent); } else { let variables = session.agent_variables(); diff --git a/src/main.rs b/src/main.rs index 2cefaf3..77b72c2 100644 --- a/src/main.rs +++ b/src/main.rs @@ -65,7 +65,7 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option) -> Result<()> return serve::run(config, addr).await; } if cli.info { - config.write().print_info_only = true; + config.write().cli_info_flag = true; } if cli.list_models { @@ -98,6 +98,15 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option) -> Result<()> Some(v) => v.as_str(), None => TEMP_SESSION_NAME, }); + if !cli.agent_variable.is_empty() { + config.write().cli_agent_variables = Some( + cli.agent_variable + .chunks(2) + .map(|v| (v[0].to_string(), v[1].to_string())) + .collect(), + ); + } + Config::use_agent(&config, agent, session, abort_signal.clone()).await? } else { if let Some(prompt) = &cli.prompt { -- cgit v1.2.3