diff options
| author | sigoden <sigoden@gmail.com> | 2024-11-06 11:38:11 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-11-06 11:38:11 +0800 |
| commit | 0fac7fea5090642eb5155b2b24f041476b89c813 (patch) | |
| tree | b293bb3a4b672e2126564254dc910362b2dd10d2 | |
| parent | b9d2e7e0cfce139f83163d14a387576ac8d331dd (diff) | |
| download | aichat-0fac7fea5090642eb5155b2b24f041476b89c813.tar.gz | |
refactor: several improvements (#973)
| -rw-r--r-- | src/config/agent.rs | 44 | ||||
| -rw-r--r-- | src/config/mod.rs | 21 | ||||
| -rw-r--r-- | src/config/role.rs | 2 | ||||
| -rw-r--r-- | src/config/session.rs | 2 | ||||
| -rw-r--r-- | src/rag/mod.rs | 6 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 |
6 files changed, 41 insertions, 36 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index 5d7278b..5589f0e 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -92,11 +92,13 @@ impl Agent { None }; + let shared_variables = agent_config.variables.clone(); + Ok(Self { name: name.to_string(), config: agent_config, definition, - shared_variables: Default::default(), + shared_variables, session_variables: None, functions, rag, @@ -113,6 +115,7 @@ impl Agent { return Ok(output); } let mut printed = false; + let mut unset_variables = vec![]; for agent_variable in agent_variables { let key = agent_variable.name.clone(); match variables.get(&key) { @@ -126,25 +129,38 @@ impl Agent { } if *IS_STDOUT_TERMINAL { if !printed { - println!("🚀 Init agent variables..."); + println!("⚙ Init agent variables..."); printed = true; } - let value = Text::new(&agent_variable.description) - .with_validator(|input: &str| { - if input.trim().is_empty() { - Ok(Validation::Invalid("This field is required".into())) - } else { - Ok(Validation::Valid) - } - }) - .prompt()?; + let value = Text::new(&format!( + "{} ({}):", + agent_variable.name, agent_variable.description + )) + .with_validator(|input: &str| { + if input.trim().is_empty() { + Ok(Validation::Invalid("This field is required".into())) + } else { + Ok(Validation::Valid) + } + }) + .prompt()?; output.insert(key, value); } else { - bail!("Failed to init agent variables in non-interactive mode"); + unset_variables.push(agent_variable) } } } } + if !unset_variables.is_empty() { + bail!( + "The following agent variables are required:\n{}", + unset_variables + .iter() + .map(|v| format!(" - {}: {}", v.name, v.description)) + .collect::<Vec<_>>() + .join("\n") + ) + } Ok(output) } @@ -211,10 +227,6 @@ impl Agent { } } - pub fn config_variables(&self) -> &IndexMap<String, String> { - &self.config.variables - } - pub fn shared_variables(&self) -> &IndexMap<String, String> { &self.shared_variables } diff --git a/src/config/mod.rs b/src/config/mod.rs index 124eaa5..b9d9d8e 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -711,7 +711,7 @@ impl Config { } } } - println!("✨ Successfully deleted {kind}."); + println!("✓ Successfully deleted {kind}."); Ok(()) } @@ -1062,8 +1062,8 @@ 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); + if self.agent.is_some() { + self.init_agent_shared_variables()?; } Ok(()) } @@ -1863,8 +1863,9 @@ impl Config { None => return Ok(()), }; let new_variables = - Agent::init_agent_variables(agent.defined_variables(), agent.config_variables())?; + Agent::init_agent_variables(agent.defined_variables(), agent.shared_variables())?; agent.set_shared_variables(new_variables); + agent.set_session_variables(None); Ok(()) } @@ -1873,18 +1874,10 @@ impl Config { (Some(agent), Some(session)) => (agent, session), _ => return Ok(()), }; - let config_variables = agent.config_variables(); let shared_variables = agent.shared_variables(); - let mut all_variables = if shared_variables.is_empty() { - config_variables.clone() - } else { - shared_variables.clone() - }; + let mut all_variables = shared_variables.clone(); all_variables.extend(session.agent_variables().clone()); let new_variables = Agent::init_agent_variables(agent.defined_variables(), &all_variables)?; - if shared_variables.is_empty() { - agent.set_shared_variables(new_variables.clone()); - } agent.set_session_variables(Some(new_variables)); session.sync_agent(agent, false); Ok(()) @@ -2210,7 +2203,7 @@ fn create_config_file(config_path: &Path) -> Result<()> { std::fs::set_permissions(config_path, perms)?; } - println!("✨ Saved config file to '{}'.\n", config_path.display()); + println!("✓ Saved config file to '{}'.\n", config_path.display()); Ok(()) } diff --git a/src/config/role.rs b/src/config/role.rs index 8c56f0c..ecd7eb0 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -167,7 +167,7 @@ impl Role { })?; if is_repl { - println!("✨ Saved role to '{}'.", role_path.display()); + println!("✓ Saved role to '{}'.", role_path.display()); } if role_name != self.name { diff --git a/src/config/session.rs b/src/config/session.rs index 8cdd793..31a6e02 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -360,7 +360,7 @@ impl Session { })?; if is_repl { - println!("✨ Saved session to '{}'.", session_path.display()); + println!("✓ Saved session to '{}'.", session_path.display()); } if self.name() != session_name { diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 0add998..1fd0376 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -68,7 +68,7 @@ impl Rag { if !*IS_STDOUT_TERMINAL { bail!("Failed to init rag in non-interactive mode"); } - println!("🚀 Initializing RAG..."); + println!("⚙ Initializing RAG..."); let (embedding_model, chunk_size, chunk_overlap) = Self::create_config(config)?; let (reranker_model, top_k) = { let config = config.read(); @@ -100,7 +100,7 @@ impl Rag { }, }; if rag.save()? { - println!("✨ Saved RAG to '{}'.", save_path.display()); + println!("✓ Saved RAG to '{}'.", save_path.display()); } Ok(rag) } @@ -155,7 +155,7 @@ impl Rag { }, }; if self.save()? { - println!("✨ Saved rag to '{}'.", self.path); + println!("✓ Saved rag to '{}'.", self.path); } Ok(()) } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 2b4fc8f..7f971b1 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -366,7 +366,7 @@ impl Repl { let ret = Config::compress_session(&self.config).await; spinner.stop(); ret?; - println!("✨ Successfully compressed the session."); + println!("✓ Successfully compressed the session."); } _ => { println!(r#"Usage: .compress session"#) |
