summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
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 /src/client/common.rs
parente522289b61bbf0d253f03433120ab3222da5e4d3 (diff)
downloadaichat-b666fc6bd24a039d2d23a312fa453d0fe3ed79ab.tar.gz
refactor: optimize configuration initialization (#1110)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs169
1 files changed, 73 insertions, 96 deletions
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)
}