summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-01-21 21:45:20 +0800
committerGitHub <noreply@github.com>2025-01-21 21:45:20 +0800
commitb666fc6bd24a039d2d23a312fa453d0fe3ed79ab (patch)
treec19b5c07ca3fba94ed99e2f34b0e96b04d43ba6d
parente522289b61bbf0d253f03433120ab3222da5e4d3 (diff)
downloadaichat-b666fc6bd24a039d2d23a312fa453d0fe3ed79ab.tar.gz
refactor: optimize configuration initialization (#1110)
-rw-r--r--src/client/azure_openai.rs13
-rw-r--r--src/client/bedrock.rs16
-rw-r--r--src/client/claude.rs3
-rw-r--r--src/client/cohere.rs3
-rw-r--r--src/client/common.rs169
-rw-r--r--src/client/ernie.rs4
-rw-r--r--src/client/gemini.rs3
-rw-r--r--src/client/macros.rs2
-rw-r--r--src/client/mod.rs8
-rw-r--r--src/client/openai.rs3
-rw-r--r--src/client/openai_compatible.rs13
-rw-r--r--src/client/vertexai.rs4
-rw-r--r--src/config/mod.rs2
-rw-r--r--src/rag/mod.rs18
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/prompt_input.rs74
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)
- }
-}