From 74c86babed710149bea45529362ce6f2815c72d8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 27 Jul 2024 15:50:08 +0800 Subject: feat: abandon AICHAT_PLATFORM and several improvements (#752) --- src/client/model.rs | 6 +++--- src/config/mod.rs | 34 +++++++++++++++++++--------------- src/main.rs | 48 +++++++++++++++++++++++++++++++----------------- 3 files changed, 53 insertions(+), 35 deletions(-) (limited to 'src') diff --git a/src/client/model.rs b/src/client/model.rs index acf04a3..17646c9 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -45,21 +45,21 @@ impl Model { pub fn retrieve_chat(config: &Config, model_id: &str) -> Result { match Self::find(&list_chat_models(config), model_id) { Some(v) => Ok(v), - None => bail!("Invalid chat model '{model_id}'"), + None => bail!("Unknown chat model '{model_id}'"), } } pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result { match Self::find(&list_embedding_models(config), model_id) { Some(v) => Ok(v), - None => bail!("Invalid embedding model '{model_id}'"), + None => bail!("Unknown embedding model '{model_id}'"), } } pub fn retrieve_reranker(config: &Config, model_id: &str) -> Result { match Self::find(&list_reranker_models(config), model_id) { Some(v) => Ok(v), - None => bail!("Invalid reranker model '{model_id}'"), + None => bail!("Unknown reranker model '{model_id}'"), } } diff --git a/src/config/mod.rs b/src/config/mod.rs index 07d7b33..6d196d6 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -210,17 +210,20 @@ impl Default for Config { pub type GlobalConfig = Arc>; impl Config { - pub fn init(working_mode: WorkingMode) -> Result { + pub fn init(working_mode: WorkingMode, model_id: Option<&str>) -> Result { let config_path = Self::config_file()?; - - let platform = env::var(get_env_name("platform")).ok(); - if *IS_STDOUT_TERMINAL && platform.is_none() && !config_path.exists() { - create_config_file(&config_path)?; - } - let mut config = if platform.is_some() { - Self::load_dynamic_config(&platform.unwrap())? + let mut config = if !config_path.exists() { + match model_id { + Some(v) => Self::load_dynamic(v)?, + None => { + if *IS_STDOUT_TERMINAL { + create_config_file(&config_path)?; + } + Self::load_from_file(&config_path)? + } + } } else { - Self::load_config_file(&config_path)? + Self::load_from_file(&config_path)? }; config.working_mode = working_mode; @@ -1498,7 +1501,7 @@ impl Config { .with_context(|| format!("Failed to create/append {}", path.display())) } - fn load_config_file(config_path: &Path) -> Result { + fn load_from_file(config_path: &Path) -> Result { let content = read_to_string(config_path) .with_context(|| format!("Failed to load config at {}", config_path.display()))?; let config: Self = serde_yaml::from_str(&content).map_err(|err| { @@ -1520,7 +1523,11 @@ impl Config { Ok(config) } - fn load_dynamic_config(platform: &str) -> Result { + fn load_dynamic(model_id: &str) -> Result { + let platform = match model_id.split_once(':') { + Some((v, _)) => v, + _ => model_id, + }; let is_openai_compatible = OPENAI_COMPATIBLE_PLATFORMS .into_iter() .any(|(name, _)| platform == name); @@ -1530,7 +1537,7 @@ impl Config { json!({ "type": platform }) }; let config = json!({ - "model": platform.to_string(), + "model": model_id.to_string(), "save": false, "clients": vec![client], }); @@ -1540,9 +1547,6 @@ impl Config { } fn load_envs(&mut self) { - if let Ok(v) = env::var(get_env_name("model")) { - self.model_id = v; - } if let Some(v) = read_env_value::("temperature") { self.temperature = v; } diff --git a/src/main.rs b/src/main.rs index 50394fb..609b51b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -46,18 +46,39 @@ async fn main() -> Result<()> { let cli = Cli::parse(); let text = cli.text(); let text = aggregate_text(text)?; - let file = &cli.file; - let no_input = text.is_none() && file.is_empty(); let working_mode = if cli.serve.is_some() { WorkingMode::Serve - } else if no_input { + } else if text.is_none() && cli.file.is_empty() { WorkingMode::Repl } else { WorkingMode::Command }; + let mut model_id = cli.model.clone(); + if model_id.is_none() { + if let Ok(v) = env::var(get_env_name("model")) { + model_id = Some(v); + } + } setup_logger(working_mode.is_serve())?; - let config = Arc::new(RwLock::new(Config::init(working_mode)?)); + let config = Arc::new(RwLock::new(Config::init( + working_mode, + model_id.as_deref(), + )?)); + let highlight = config.read().highlight; + if let Err(err) = run(config, cli, text, model_id).await { + let highlight = stderr().is_terminal() && highlight; + render_error(err, highlight); + std::process::exit(1); + } + Ok(()) +} +async fn run( + config: GlobalConfig, + cli: Cli, + text: Option, + model_id: Option, +) -> Result<()> { let abort_signal = create_abort_signal(); if let Some(addr) = cli.serve { @@ -127,7 +148,7 @@ async fn main() -> Result<()> { println!("{sessions}"); return Ok(()); } - if let Some(model_id) = &cli.model { + if let Some(model_id) = &model_id { config.write().set_model(model_id)?; } if cli.save_session { @@ -141,29 +162,22 @@ async fn main() -> Result<()> { println!("{}", info); return Ok(()); } - if cli.execute { - if no_input { - bail!("No input"); - } - let input = create_input(&config, text, file).await?; + let is_repl = config.read().working_mode.is_repl(); + if cli.execute && !is_repl { + let input = create_input(&config, text, &cli.file).await?; let shell = detect_shell(); shell_execute(&config, &shell, input).await?; return Ok(()); } config.write().apply_prelude()?; - if let Err(err) = match no_input { + match is_repl { false => { - let mut input = create_input(&config, text, file).await?; + let mut input = create_input(&config, text, &cli.file).await?; input.use_embeddings(abort_signal.clone()).await?; start_directive(&config, input, cli.no_stream, cli.code, abort_signal).await } true => start_interactive(&config).await, - } { - let highlight = stderr().is_terminal() && config.read().highlight; - render_error(err, highlight); - std::process::exit(1); } - Ok(()) } #[async_recursion] -- cgit v1.2.3