summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-25 11:13:54 +0800
committerGitHub <noreply@github.com>2024-03-25 11:13:54 +0800
commit0ebc7955da67b488877cab7929019d40688fb3ed (patch)
treee08d1c71a3f12a3470e9518af8828e3d2da40b82 /src
parenteec041c111c0ee170dab65942184e66c41479fcd (diff)
downloadaichat-0ebc7955da67b488877cab7929019d40688fb3ed.tar.gz
refactor: improve creating config for openai-compatible client (#374)
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs10
-rw-r--r--src/client/openai_compatible.rs3
-rw-r--r--src/config/mod.rs5
3 files changed, 12 insertions, 6 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index c35ba1b..4acc707 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -104,7 +104,7 @@ macro_rules! register_client {
vec![$($client::NAME,)+]
}
- pub fn create_client_config(client: &str) -> anyhow::Result<serde_json::Value> {
+ pub fn create_client_config(client: &str) -> anyhow::Result<(String, serde_json::Value)> {
$(
if client == $client::NAME {
return create_config(&$client::PROMPTS, $client::NAME)
@@ -310,15 +310,19 @@ pub struct SendData {
pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind);
-pub fn create_config(list: &[PromptType], client: &str) -> Result<Value> {
+pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> {
let mut config = json!({
"type": client,
});
+ let mut model = client.to_string();
for (path, desc, required, kind) in list {
match kind {
PromptKind::String => {
let value = prompt_input_string(desc, *required)?;
set_config_value(&mut config, path, kind, &value);
+ if *path == "name" {
+ model = value;
+ }
}
PromptKind::Integer => {
let value = prompt_input_integer(desc, *required)?;
@@ -328,7 +332,7 @@ pub fn create_config(list: &[PromptType], client: &str) -> Result<Value> {
}
let clients = json!(vec![config]);
- Ok(clients)
+ Ok((model, clients))
}
#[allow(unused)]
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index ec3333c..ec36623 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -23,7 +23,8 @@ openai_compatible_client!(OpenAICompatibleClient);
impl OpenAICompatibleClient {
config_get_fn!(api_key, get_api_key);
- pub const PROMPTS: [PromptType<'static>; 4] = [
+ pub const PROMPTS: [PromptType<'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),
diff --git a/src/config/mod.rs b/src/config/mod.rs
index b6bc730..bcebdd0 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1046,8 +1046,9 @@ fn create_config_file(config_path: &Path) -> Result<()> {
let client = Select::new("Platform:", list_client_types()).prompt()?;
let mut config = serde_json::json!({});
- config["model"] = client.into();
- config[CLIENTS_FIELD] = create_client_config(client)?;
+ let (model, clients_config) = create_client_config(client)?;
+ config["model"] = model.into();
+ config[CLIENTS_FIELD] = clients_config;
let config_data = serde_yaml::to_string(&config).with_context(|| "Failed to create config")?;