summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-08 18:31:22 +0800
committerGitHub <noreply@github.com>2025-02-08 18:31:22 +0800
commit9b3ca7a4b0de876ac58adbb9e2a6038fd66c2303 (patch)
treeeba7e05b02a169fec1bfbf182d096c80ab30cfa9
parentd6e5214153ecd9c99331c6f4c27d0553451a18a3 (diff)
downloadaichat-9b3ca7a4b0de876ac58adbb9e2a6038fd66c2303.tar.gz
feat: supports selecting LLM during configuration initialization (#1158)
-rw-r--r--src/client/common.rs61
-rw-r--r--src/client/macros.rs4
2 files changed, 50 insertions, 15 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 5ad27ae..5090bab 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -10,7 +10,7 @@ use crate::{
use anyhow::{bail, Context, Result};
use fancy_regex::Regex;
use indexmap::IndexMap;
-use inquire::{required, Text};
+use inquire::{required, Select, Text};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
@@ -335,9 +335,9 @@ pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String,
let mut config = json!({
"type": client,
});
- set_client_config(prompts, &mut config, client)?;
+ let model = set_client_config(prompts, &mut config, client)?;
let clients = json!(vec![config]);
- Ok((client.to_string(), clients))
+ Ok((model, clients))
}
pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> {
@@ -371,9 +371,9 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
config["api_key"] = api_key.into();
}
- set_client_models_config(&mut config, &name)?;
+ let model = set_client_models_config(&mut config, &name)?;
let clients = json!(vec![config]);
- Ok(Some((name, clients)))
+ Ok(Some((model, clients)))
}
pub async fn call_chat_completions(
@@ -512,7 +512,11 @@ 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<()> {
+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();
@@ -524,9 +528,16 @@ fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: &
set_client_models_config(client_config, client)
}
-fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<()> {
- if ALL_PROVIDER_MODELS.iter().any(|v| v.provider == client) {
- return Ok(());
+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
+ .iter()
+ .filter(|v| v.model_type == "chat")
+ .map(|v| v.name.clone())
+ .collect();
+ let model_name = select_model(models)?;
+ return Ok(format!("{client}:{model_name}"));
}
let model_names = prompt_input_string(
@@ -534,12 +545,36 @@ fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<(
true,
Some("Separated by commas, e.g. llama3.3,qwen2.5"),
)?;
- let models: Vec<Value> = model_names
+ let model_names = model_names
.split(',')
- .map(|v| json!({"name": v.trim()}))
- .collect();
+ .filter_map(|v| {
+ let v = v.trim();
+ if v.is_empty() {
+ None
+ } else {
+ Some(v.to_string())
+ }
+ })
+ .collect::<Vec<_>>();
+ if model_names.is_empty() {
+ bail!("No models");
+ }
+ let models: Vec<Value> = model_names.iter().map(|v| json!({"name": v})).collect();
client_config["models"] = models.into();
- Ok(())
+ let model_name = select_model(model_names)?;
+ Ok(format!("{client}:{model_name}"))
+}
+
+fn select_model(model_names: Vec<String>) -> Result<String> {
+ if model_names.is_empty() {
+ bail!("No models");
+ }
+ let model = if model_names.len() == 1 {
+ model_names[0].clone()
+ } else {
+ Select::new("Select model:", model_names).prompt()?
+ };
+ Ok(model)
}
fn prompt_input_string(
diff --git a/src/client/macros.rs b/src/client/macros.rs
index 97171db..6a06236 100644
--- a/src/client/macros.rs
+++ b/src/client/macros.rs
@@ -52,12 +52,12 @@ macro_rules! register_client {
pub fn list_models(local_config: &$config) -> Vec<Model> {
let client_name = Self::name(local_config);
if local_config.models.is_empty() {
- if let Some(models) = $crate::client::ALL_PROVIDER_MODELS.iter().find(|v| {
+ if let Some(v) = $crate::client::ALL_PROVIDER_MODELS.iter().find(|v| {
v.provider == $name ||
($name == OpenAICompatibleClient::NAME
&& local_config.name.as_ref().map(|name| name.starts_with(&v.provider)).unwrap_or_default())
}) {
- return Model::from_config(client_name, &models.models);
+ return Model::from_config(client_name, &v.models);
}
vec![]
} else {