summaryrefslogtreecommitdiffstats
path: root/src/client
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/client
parenteec041c111c0ee170dab65942184e66c41479fcd (diff)
downloadaichat-0ebc7955da67b488877cab7929019d40688fb3ed.tar.gz
refactor: improve creating config for openai-compatible client (#374)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/common.rs10
-rw-r--r--src/client/openai_compatible.rs3
2 files changed, 9 insertions, 4 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),