diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-08 18:31:22 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-08 18:31:22 +0800 |
| commit | 9b3ca7a4b0de876ac58adbb9e2a6038fd66c2303 (patch) | |
| tree | eba7e05b02a169fec1bfbf182d096c80ab30cfa9 | |
| parent | d6e5214153ecd9c99331c6f4c27d0553451a18a3 (diff) | |
| download | aichat-9b3ca7a4b0de876ac58adbb9e2a6038fd66c2303.tar.gz | |
feat: supports selecting LLM during configuration initialization (#1158)
| -rw-r--r-- | src/client/common.rs | 61 | ||||
| -rw-r--r-- | src/client/macros.rs | 4 |
2 files changed, 50 insertions, 15 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 5ad27ae..5090bab 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -10,7 +10,7 @@ use crate::{ use anyhow::{bail, Context, Result}; use fancy_regex::Regex; use indexmap::IndexMap; -use inquire::{required, Text}; +use inquire::{required, Select, Text}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -335,9 +335,9 @@ pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, let mut config = json!({ "type": client, }); - set_client_config(prompts, &mut config, client)?; + let model = set_client_config(prompts, &mut config, client)?; let clients = json!(vec![config]); - Ok((client.to_string(), clients)) + Ok((model, clients)) } pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> { @@ -371,9 +371,9 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St config["api_key"] = api_key.into(); } - set_client_models_config(&mut config, &name)?; + let model = set_client_models_config(&mut config, &name)?; let clients = json!(vec![config]); - Ok(Some((name, clients))) + Ok(Some((model, clients))) } pub async fn call_chat_completions( @@ -512,7 +512,11 @@ pub fn json_str_from_map<'a>( map.get(field_name).and_then(|v| v.as_str()) } -fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: &str) -> Result<()> { +fn set_client_config( + list: &[PromptAction], + client_config: &mut Value, + client: &str, +) -> Result<String> { for (key, desc, help_message) in list { let env_name = format!("{client}_{key}").to_ascii_uppercase(); let required = std::env::var(&env_name).is_err(); @@ -524,9 +528,16 @@ fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: & set_client_models_config(client_config, client) } -fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<()> { - if ALL_PROVIDER_MODELS.iter().any(|v| v.provider == client) { - return Ok(()); +fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<String> { + if let Some(provider) = ALL_PROVIDER_MODELS.iter().find(|v| v.provider == client) { + let models: Vec<String> = provider + .models + .iter() + .filter(|v| v.model_type == "chat") + .map(|v| v.name.clone()) + .collect(); + let model_name = select_model(models)?; + return Ok(format!("{client}:{model_name}")); } let model_names = prompt_input_string( @@ -534,12 +545,36 @@ fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<( true, Some("Separated by commas, e.g. llama3.3,qwen2.5"), )?; - let models: Vec<Value> = model_names + let model_names = model_names .split(',') - .map(|v| json!({"name": v.trim()})) - .collect(); + .filter_map(|v| { + let v = v.trim(); + if v.is_empty() { + None + } else { + Some(v.to_string()) + } + }) + .collect::<Vec<_>>(); + if model_names.is_empty() { + bail!("No models"); + } + let models: Vec<Value> = model_names.iter().map(|v| json!({"name": v})).collect(); client_config["models"] = models.into(); - Ok(()) + let model_name = select_model(model_names)?; + Ok(format!("{client}:{model_name}")) +} + +fn select_model(model_names: Vec<String>) -> Result<String> { + if model_names.is_empty() { + bail!("No models"); + } + let model = if model_names.len() == 1 { + model_names[0].clone() + } else { + Select::new("Select model:", model_names).prompt()? + }; + Ok(model) } fn prompt_input_string( diff --git a/src/client/macros.rs b/src/client/macros.rs index 97171db..6a06236 100644 --- a/src/client/macros.rs +++ b/src/client/macros.rs @@ -52,12 +52,12 @@ macro_rules! register_client { pub fn list_models(local_config: &$config) -> Vec<Model> { let client_name = Self::name(local_config); if local_config.models.is_empty() { - if let Some(models) = $crate::client::ALL_PROVIDER_MODELS.iter().find(|v| { + if let Some(v) = $crate::client::ALL_PROVIDER_MODELS.iter().find(|v| { v.provider == $name || ($name == OpenAICompatibleClient::NAME && local_config.name.as_ref().map(|name| name.starts_with(&v.provider)).unwrap_or_default()) }) { - return Model::from_config(client_name, &models.models); + return Model::from_config(client_name, &v.models); } vec![] } else { |
