diff options
| author | sigoden <sigoden@gmail.com> | 2025-01-21 21:45:20 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-01-21 21:45:20 +0800 |
| commit | b666fc6bd24a039d2d23a312fa453d0fe3ed79ab (patch) | |
| tree | c19b5c07ca3fba94ed99e2f34b0e96b04d43ba6d /src/client | |
| parent | e522289b61bbf0d253f03433120ab3222da5e4d3 (diff) | |
| download | aichat-b666fc6bd24a039d2d23a312fa453d0fe3ed79ab.tar.gz | |
refactor: optimize configuration initialization (#1110)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/azure_openai.rs | 13 | ||||
| -rw-r--r-- | src/client/bedrock.rs | 16 | ||||
| -rw-r--r-- | src/client/claude.rs | 3 | ||||
| -rw-r--r-- | src/client/cohere.rs | 3 | ||||
| -rw-r--r-- | src/client/common.rs | 169 | ||||
| -rw-r--r-- | src/client/ernie.rs | 4 | ||||
| -rw-r--r-- | src/client/gemini.rs | 3 | ||||
| -rw-r--r-- | src/client/macros.rs | 2 | ||||
| -rw-r--r-- | src/client/mod.rs | 8 | ||||
| -rw-r--r-- | src/client/openai.rs | 3 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 13 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 4 |
12 files changed, 96 insertions, 145 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index cd1d9e4..fc856b2 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -19,16 +19,13 @@ impl AzureOpenAIClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptAction<'static>; 4] = [ - ("api_base", "API Base:", true, PromptKind::String), - ("api_key", "API Key:", true, PromptKind::String), - ("models[].name", "Model Name:", true, PromptKind::String), + pub const PROMPTS: [PromptAction<'static>; 2] = [ ( - "models[].max_input_tokens", - "Max Input Tokens:", - false, - PromptKind::Integer, + "api_base", + "API Base", + Some("e.g. https://{RESOURCE}.openai.azure.com"), ), + ("api_key", "API Key", None), ]; } diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 7cd289c..435aa57 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -31,19 +31,9 @@ impl BedrockClient { config_get_fn!(region, get_region); pub const PROMPTS: [PromptAction<'static>; 3] = [ - ( - "access_key_id", - "AWS Access Key ID", - true, - PromptKind::String, - ), - ( - "secret_access_key", - "AWS Secret Access Key", - true, - PromptKind::String, - ), - ("region", "AWS Region", true, PromptKind::String), + ("access_key_id", "AWS Access Key ID", None), + ("secret_access_key", "AWS Secret Access Key", None), + ("region", "AWS Region", None), ]; fn chat_completions_builder( diff --git a/src/client/claude.rs b/src/client/claude.rs index f982e14..9a0e1f6 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -22,8 +22,7 @@ impl ClaudeClient { config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); - pub const PROMPTS: [PromptAction<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key", None)]; } impl_client_trait!( diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 5f61454..ae96977 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -24,8 +24,7 @@ impl CohereClient { config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); - pub const PROMPTS: [PromptAction<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key", None)]; } impl_client_trait!( diff --git a/src/client/common.rs b/src/client/common.rs index a4e171a..b4d01ce 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -10,6 +10,7 @@ use crate::{ use anyhow::{bail, Context, Result}; use fancy_regex::Regex; use indexmap::IndexMap; +use inquire::{required, Text}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -325,53 +326,50 @@ pub struct RerankResult { pub relevance_score: f64, } -pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind); +pub type PromptAction<'a> = (&'a str, &'a str, Option<&'a str>); pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { let mut config = json!({ "type": client, }); - let mut model = client.to_string(); - set_client_config(prompts, &mut model, &mut config)?; + set_client_config(prompts, &mut config, client)?; let clients = json!(vec![config]); - Ok((model, clients)) + Ok((client.to_string(), clients)) } pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> { - match super::OPENAI_COMPATIBLE_PLATFORMS + let api_base = super::OPENAI_COMPATIBLE_PLATFORMS .into_iter() .find(|(name, _)| client == *name) - { - None => Ok(None), - Some((name, api_base)) => { - let mut config = json!({ - "type": OpenAICompatibleClient::NAME, - "name": name, - }); - let mut prompts = vec![]; - if api_base.is_empty() { - prompts.push(("api_base", "API Base:", true, PromptKind::String)); - } else { - config["api_base"] = api_base.into(); - } - prompts.push(("api_key", "API Key:", false, PromptKind::String)); - if !ALL_PREDEFINED_MODELS.iter().any(|v| v.platform == name) { - prompts.extend([ - ("models[].name", "Model Name:", true, PromptKind::String), - ( - "models[].max_input_tokens", - "Max Input Tokens:", - false, - PromptKind::Integer, - ), - ]); - }; - let mut model = client.to_string(); - set_client_config(&prompts, &mut model, &mut config)?; - let clients = json!(vec![config]); - Ok(Some((model, clients))) - } + .map(|(_, api_base)| api_base) + .unwrap_or("http(s)://{API_ADDR}/v1"); + + let name = if client == OpenAICompatibleClient::NAME { + prompt_input_string("Provider Name", true, None)? + } else { + client.to_string() + }; + + let mut config = json!({ + "type": OpenAICompatibleClient::NAME, + "name": &name, + }); + + let api_base = if api_base.contains('{') { + prompt_input_string("API Base", true, Some(&format!("e.g. {api_base}")))? + } else { + api_base.to_string() + }; + config["api_base"] = api_base.into(); + + let api_key = prompt_input_string("API Key", false, None)?; + if !api_key.is_empty() { + config["api_key"] = api_key.into(); } + + set_client_models_config(&mut config, &name)?; + let clients = json!(vec![config]); + Ok(Some((name, clients))) } pub async fn call_chat_completions( @@ -537,74 +535,53 @@ pub fn maybe_catch_error(data: &Value) -> Result<()> { Ok(()) } -fn set_client_config( - list: &[PromptAction], - model: &mut String, - client_config: &mut Value, -) -> Result<()> { - let env_prefix = model.clone(); - for (path, desc, required, kind) in list { - let mut required = *required; - if required { - let env_name = format!("{env_prefix}_{path}").to_ascii_uppercase(); - if std::env::var(&env_name).is_ok() { - required = false; - } - } - match kind { - PromptKind::String => { - let value = prompt_input_string(desc, required)?; - set_client_config_value(client_config, path, kind, &value); - if *path == "name" { - *model = value; - } - } - PromptKind::Integer => { - let value = prompt_input_integer(desc, required)?; - set_client_config_value(client_config, path, kind, &value); - } +fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: &str) -> Result<()> { + 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(); } } - Ok(()) + set_client_models_config(client_config, client) } -fn set_client_config_value(client_config: &mut Value, path: &str, kind: &PromptKind, value: &str) { - let segs: Vec<&str> = path.split('.').collect(); - match segs.as_slice() { - [name] => client_config[name] = prompt_value_to_json(kind, value), - [scope, name] => match scope.split_once('[') { - None => { - if client_config.get(scope).is_none() { - let mut obj = json!({}); - obj[name] = prompt_value_to_json(kind, value); - client_config[scope] = obj; - } else { - client_config[scope][name] = prompt_value_to_json(kind, value); - } - } - Some((scope, _)) => { - if client_config.get(scope).is_none() { - let mut obj = json!({}); - obj[name] = prompt_value_to_json(kind, value); - client_config[scope] = json!([obj]); - } else { - client_config[scope][0][name] = prompt_value_to_json(kind, value); - } - } - }, - _ => {} +fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<()> { + if ALL_PREDEFINED_MODELS.iter().any(|v| v.platform == client) { + return Ok(()); } + + let model_names = prompt_input_string( + "LLM models", + true, + Some("Separated by commas, e.g. llama3.3,qwen2.5"), + )?; + let models: Vec<Value> = model_names + .split(',') + .map(|v| json!({"name": v.trim()})) + .collect(); + client_config["models"] = models.into(); + Ok(()) } -fn prompt_value_to_json(kind: &PromptKind, value: &str) -> Value { - if value.is_empty() { - return Value::Null; +fn prompt_input_string( + desc: &str, + required: bool, + help_message: Option<&str>, +) -> anyhow::Result<String> { + let desc = if required { + format!("{desc} (required):") + } else { + format!("{desc} (optional):") + }; + let mut text = Text::new(&desc); + if required { + text = text.with_validator(required!("This field is required")) } - match kind { - PromptKind::String => value.into(), - PromptKind::Integer => match value.parse::<i32>() { - Ok(value) => value.into(), - Err(_) => value.into(), - }, + if let Some(help_message) = help_message { + text = text.with_help_message(help_message); } + let text = text.prompt()?; + Ok(text) } diff --git a/src/client/ernie.rs b/src/client/ernie.rs index d0fe7b5..9a6b962 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -25,8 +25,8 @@ impl ErnieClient { config_get_fn!(api_key, get_api_key); config_get_fn!(secret_key, get_secret_key); pub const PROMPTS: [PromptAction<'static>; 2] = [ - ("api_key", "API Key:", true, PromptKind::String), - ("secret_key", "Secret Key:", true, PromptKind::String), + ("api_key", "API Key", None), + ("secret_key", "Secret Key", None), ]; } diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 5d0ee0d..85917c3 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -23,8 +23,7 @@ impl GeminiClient { config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); - pub const PROMPTS: [PromptAction<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key", None)]; } impl_client_trait!( diff --git a/src/client/macros.rs b/src/client/macros.rs index 4f52044..a76e62b 100644 --- a/src/client/macros.rs +++ b/src/client/macros.rs @@ -89,7 +89,7 @@ macro_rules! register_client { pub fn create_client_config(client: &str) -> anyhow::Result<(String, serde_json::Value)> { $( - if client == $client::NAME { + if client == $client::NAME && client != $crate::client::OpenAICompatibleClient::NAME { return create_config(&$client::PROMPTS, $client::NAME) } )+ diff --git a/src/client/mod.rs b/src/client/mod.rs index 9ccfad8..3d8d4da 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -7,7 +7,6 @@ mod model; mod stream; pub use crate::function::ToolCall; -pub use crate::utils::PromptKind; pub use common::*; pub use message::*; pub use model::*; @@ -37,7 +36,10 @@ register_client!( pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 22] = [ ("ai21", "https://api.ai21.com/studio/v1"), - ("cloudflare", ""), + ( + "cloudflare", + "https://api.cloudflare.com/client/v4/accounts/{ACCOUNT_ID}/ai/v1", + ), ("deepinfra", "https://api.deepinfra.com/v1/openai"), ("deepseek", "https://api.deepseek.com"), ("fireworks", "https://api.fireworks.ai/inference/v1"), @@ -49,7 +51,7 @@ pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 22] = [ ("mistral", "https://api.mistral.ai/v1"), ("moonshot", "https://api.moonshot.cn/v1"), ("openrouter", "https://openrouter.ai/api/v1"), - ("ollama", "http://127.0.0.1:11434/v1"), + ("ollama", "http://{OLLAMA_HOST}:11434/v1"), ("perplexity", "https://api.perplexity.ai"), ( "qianwen", diff --git a/src/client/openai.rs b/src/client/openai.rs index 253bf21..ce00de7 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -23,8 +23,7 @@ impl OpenAIClient { config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); - pub const PROMPTS: [PromptAction<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key", None)]; } impl_client_trait!( diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index f2b7ae3..18acafb 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -21,18 +21,7 @@ impl OpenAICompatibleClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptAction<'static>; 5] = [ - ("name", "Platform Name:", true, PromptKind::String), - ("api_base", "API Base:", true, PromptKind::String), - ("api_key", "API Key:", false, PromptKind::String), - ("models[].name", "Model Name:", true, PromptKind::String), - ( - "models[].max_input_tokens", - "Max Input Tokens:", - false, - PromptKind::Integer, - ), - ]; + pub const PROMPTS: [PromptAction<'static>; 0] = []; } impl_client_trait!( diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 1612b34..19c8436 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -27,8 +27,8 @@ impl VertexAIClient { config_get_fn!(location, get_location); pub const PROMPTS: [PromptAction<'static>; 2] = [ - ("project_id", "Project ID", true, PromptKind::String), - ("location", "Location", true, PromptKind::String), + ("project_id", "Project ID", None), + ("location", "Location", None), ]; } |
