diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-10 08:01:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-10 08:01:21 +0800 |
| commit | d4087dbaf7208da5ae842ec5ba2e9c442bfa4f94 (patch) | |
| tree | b6cd9a8d5a1e2c3b1ad6b641bf2c6c509fefb282 /src/client/common.rs | |
| parent | 216c332c9a6c0143497703e4213b25741e4ce0c7 (diff) | |
| download | aichat-d4087dbaf7208da5ae842ec5ba2e9c442bfa4f94.tar.gz | |
feat: supports fetching models during configuration initialization (#1161)
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 123 |
1 files changed, 92 insertions, 31 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 5090bab..cfbb9d6 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -10,7 +10,9 @@ use crate::{ use anyhow::{bail, Context, Result}; use fancy_regex::Regex; use indexmap::IndexMap; -use inquire::{required, Select, Text}; +use inquire::{ + list_option::ListOption, required, validator::Validation, MultiSelect, Select, Text, +}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -23,6 +25,7 @@ lazy_static::lazy_static! { pub static ref ALL_PROVIDER_MODELS: Vec<ProviderModels> = { Config::loal_models_override().ok().unwrap_or_else(|| serde_yaml::from_str(MODELS_YAML).unwrap()) }; + static ref EMBEDDING_MODEL_RE: Regex = Regex::new(r"(^(bge-|e5-|uae-|gte-|text-)|embed|multilingual|minilm)").unwrap(); static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap(); } @@ -331,16 +334,29 @@ pub struct RerankResult { pub type PromptAction<'a> = (&'a str, &'a str, Option<&'a str>); -pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { +pub async fn create_config( + prompts: &[PromptAction<'static>], + client: &str, +) -> Result<(String, Value)> { let mut config = json!({ "type": client, }); - let model = set_client_config(prompts, &mut config, client)?; + for (key, desc, help_message) in prompts { + let env_name = format!("{client}_{key}").to_ascii_uppercase(); + let required = std::env::var(&env_name).is_err(); + let value = prompt_input_string(desc, required, *help_message)?; + if !value.is_empty() { + config[key] = value.into(); + } + } + let model = set_client_models_config(&mut config, client).await?; let clients = json!(vec![config]); Ok((model, clients)) } -pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> { +pub async fn create_openai_compatible_client_config( + client: &str, +) -> Result<Option<(String, Value)>> { let api_base = super::OPENAI_COMPATIBLE_PROVIDERS .into_iter() .find(|(name, _)| client == *name) @@ -371,7 +387,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St config["api_key"] = api_key.into(); } - let model = set_client_models_config(&mut config, &name)?; + let model = set_client_models_config(&mut config, &name).await?; let clients = json!(vec![config]); Ok(Some((model, clients))) } @@ -512,23 +528,7 @@ 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<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(); - let value = prompt_input_string(desc, required, *help_message)?; - if !value.is_empty() { - client_config[key] = value.into(); - } - } - set_client_models_config(client_config, client) -} - -fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<String> { +async 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 @@ -539,13 +539,46 @@ fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<S let model_name = select_model(models)?; return Ok(format!("{client}:{model_name}")); } - - let model_names = prompt_input_string( - "LLM models", - true, - Some("Separated by commas, e.g. llama3.3,qwen2.5"), - )?; - let model_names = model_names + let mut model_names = vec![]; + if let (Some(true), Some(api_base), api_key) = ( + client_config["type"] + .as_str() + .map(|v| v == OpenAICompatibleClient::NAME), + client_config["api_base"].as_str(), + client_config["api_key"] + .as_str() + .map(|v| v.to_string()) + .or_else(|| { + let env_name = format!("{client}_api_key").to_ascii_uppercase(); + std::env::var(&env_name).ok() + }), + ) { + if let Ok(fetched_models) = abortable_run_with_spinner( + fetch_models(api_base, api_key.as_deref()), + "Fetching models", + create_abort_signal(), + ) + .await + { + model_names = MultiSelect::new("LLM models (required):", fetched_models) + .with_validator(|list: &[ListOption<&String>]| { + if list.is_empty() { + Ok(Validation::Invalid( + "At least one item must be selected".into(), + )) + } else { + Ok(Validation::Valid) + } + }) + .prompt()?; + } + } + if model_names.is_empty() { + model_names = prompt_input_string( + "LLM models", + true, + Some("Separated by commas, e.g. llama3.3,qwen2.5"), + )? .split(',') .filter_map(|v| { let v = v.trim(); @@ -556,10 +589,38 @@ fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<S } }) .collect::<Vec<_>>(); + } if model_names.is_empty() { bail!("No models"); } - let models: Vec<Value> = model_names.iter().map(|v| json!({"name": v})).collect(); + let models: Vec<Value> = model_names + .iter() + .map(|v| { + let l = v.to_lowercase(); + if l.contains("rank") { + json!({ + "name": v, + "type": "reranker", + }) + } else if let Ok(true) = EMBEDDING_MODEL_RE.is_match(&l) { + json!({ + "name": v, + "type": "embedding", + "default_chunk_size": 1000, + "max_batch_size": 16 + }) + } else if v.contains("vision") { + json!({ + "name": v, + "supports_vision": true + }) + } else { + json!({ + "name": v, + }) + } + }) + .collect(); client_config["models"] = models.into(); let model_name = select_model(model_names)?; Ok(format!("{client}:{model_name}")) @@ -572,7 +633,7 @@ fn select_model(model_names: Vec<String>) -> Result<String> { let model = if model_names.len() == 1 { model_names[0].clone() } else { - Select::new("Select model:", model_names).prompt()? + Select::new("Default Model (required):", model_names).prompt()? }; Ok(model) } |
