diff options
| -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 | ||||
| -rw-r--r-- | src/config/mod.rs | 2 | ||||
| -rw-r--r-- | src/rag/mod.rs | 18 | ||||
| -rw-r--r-- | src/utils/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/prompt_input.rs | 74 |
16 files changed, 115 insertions, 222 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), ]; } diff --git a/src/config/mod.rs b/src/config/mod.rs index 5dfb2d9..8813930 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -2568,7 +2568,7 @@ fn create_config_file(config_path: &Path) -> Result<()> { process::exit(0); } - let client = Select::new("Platform:", list_client_types()).prompt()?; + let client = Select::new("API Provider (required):", list_client_types()).prompt()?; let mut config = serde_json::json!({}); let (model, clients_config) = create_client_config(client)?; diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 8855f4f..6165969 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -842,6 +842,24 @@ fn select_embedding_model(models: &[&Model]) -> Result<String> { Ok(result.value) } +#[derive(Debug)] +struct SelectOption { + pub value: String, + pub description: String, +} + +impl SelectOption { + pub fn new(value: String, description: String) -> Self { + Self { value, description } + } +} + +impl std::fmt::Display for SelectOption { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{} ({})", self.value, self.description) + } +} + fn set_chunk_size(model: &Model) -> Result<usize> { let default_value = model.default_chunk_size().to_string(); let help_message = model diff --git a/src/utils/mod.rs b/src/utils/mod.rs index ecc81aa..8d66470 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -5,7 +5,6 @@ mod crypto; mod html_to_md; mod loader; mod path; -mod prompt_input; mod render_prompt; mod request; mod spinner; @@ -18,7 +17,6 @@ pub use self::crypto::*; pub use self::html_to_md::*; pub use self::loader::*; pub use self::path::*; -pub use self::prompt_input::*; pub use self::render_prompt::render_prompt; pub use self::request::*; pub use self::spinner::*; diff --git a/src/utils/prompt_input.rs b/src/utils/prompt_input.rs deleted file mode 100644 index 26343a2..0000000 --- a/src/utils/prompt_input.rs +++ /dev/null @@ -1,74 +0,0 @@ -use inquire::{required, validator::Validation, Text}; - -const MSG_REQUIRED: &str = "This field is required"; -const MSG_OPTIONAL: &str = "Optional field - Press ↵ to skip"; - -pub fn prompt_input_string(desc: &str, required: bool) -> anyhow::Result<String> { - let mut text = Text::new(desc); - if required { - text = text.with_validator(required!(MSG_REQUIRED)) - } else { - text = text.with_help_message(MSG_OPTIONAL) - } - let text = text.prompt()?; - Ok(text) -} - -pub fn prompt_input_integer(desc: &str, required: bool) -> anyhow::Result<String> { - let mut text = Text::new(desc); - if required { - text = text.with_validator(|text: &str| { - let out = if text.is_empty() { - Validation::Invalid(MSG_REQUIRED.into()) - } else { - validate_integer(text) - }; - Ok(out) - }) - } else { - text = text - .with_validator(|text: &str| { - let out = if text.is_empty() { - Validation::Valid - } else { - validate_integer(text) - }; - Ok(out) - }) - .with_help_message(MSG_OPTIONAL) - } - let text = text.prompt()?; - Ok(text) -} - -#[derive(Debug, Clone, Copy)] -pub enum PromptKind { - String, - Integer, -} - -fn validate_integer(text: &str) -> Validation { - if text.parse::<i32>().is_err() { - Validation::Invalid("Must be a integer".into()) - } else { - Validation::Valid - } -} - -#[derive(Debug)] -pub struct SelectOption { - pub value: String, - pub description: String, -} - -impl SelectOption { - pub fn new(value: String, description: String) -> Self { - Self { value, description } - } -} - -impl std::fmt::Display for SelectOption { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{} ({})", self.value, self.description) - } -} |
