summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-07 10:56:28 +0800
committerGitHub <noreply@github.com>2023-11-07 10:56:28 +0800
commit9a8b302432a3f9bfa1e467dde027fc92dacce3e2 (patch)
treeee18de8b077980619ff9a1dee24464a1d13b801d /src
parent87aec71e080dc9c15f3dd0059dbc91686f3de4e6 (diff)
downloadaichat-9a8b302432a3f9bfa1e467dde027fc92dacce3e2.tar.gz
refactor: remove Model.client_index, match client by name (#218)
Diffstat (limited to 'src')
-rw-r--r--src/client/azure_openai.rs4
-rw-r--r--src/client/common.rs29
-rw-r--r--src/client/ernie.rs6
-rw-r--r--src/client/localai.rs4
-rw-r--r--src/client/model.rs6
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/palm.rs16
-rw-r--r--src/client/qianwen.rs4
8 files changed, 32 insertions, 41 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 8fa3e10..f4b7916 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -41,14 +41,14 @@ impl AzureOpenAIClient {
),
];
- pub fn list_models(local_config: &AzureOpenAIConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &AzureOpenAIConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
local_config
.models
.iter()
.map(|v| {
- Model::new(client_index, client_name, &v.name)
+ Model::new(client_name, &v.name)
.set_max_tokens(v.max_tokens)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/common.rs b/src/client/common.rs
index 336450f..27dd81d 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -54,13 +54,15 @@ macro_rules! register_client {
pub fn init(global_config: &$crate::config::GlobalConfig) -> Option<Box<dyn Client>> {
let model = global_config.read().model.clone();
- let config = {
- if let ClientConfig::$config(c) = &global_config.read().clients[model.client_index] {
- c.clone()
- } else {
- return None;
+ let config = global_config.read().clients.iter().find_map(|client_config| {
+ if let ClientConfig::$config(c) = client_config {
+ if Self::name(c) == &model.client_name {
+ return Some(c.clone())
+ }
}
- };
+ None
+ })?;
+
Some(Box::new(Self {
global_config: global_config.clone(),
config,
@@ -68,8 +70,8 @@ macro_rules! register_client {
}))
}
- pub fn name(local_config: &$config) -> &str {
- local_config.name.as_deref().unwrap_or(Self::NAME)
+ pub fn name(config: &$config) -> &str {
+ config.name.as_deref().unwrap_or(Self::NAME)
}
}
@@ -80,11 +82,7 @@ macro_rules! register_client {
$(.or_else(|| $client::init(config)))+
.ok_or_else(|| {
let model = config.read().model.clone();
- anyhow::anyhow!(
- "Unknown client '{}' at config.clients[{}]",
- &model.client_name,
- &model.client_index
- )
+ anyhow::anyhow!("Unknown client '{}'", &model.client_name)
})
}
@@ -105,9 +103,8 @@ macro_rules! register_client {
config
.clients
.iter()
- .enumerate()
- .flat_map(|(i, v)| match v {
- $(ClientConfig::$config(c) => $client::list_models(c, i),)+
+ .flat_map(|v| match v {
+ $(ClientConfig::$config(c) => $client::list_models(c),)+
ClientConfig::Unknown => vec![],
})
.collect()
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index c871450..5e32a81 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -64,11 +64,11 @@ impl ErnieClient {
("secret_key", "Secret Key:", true, PromptKind::String),
];
- pub fn list_models(local_config: &ErnieConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &ErnieConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, _)| Model::new(client_index, client_name, name))
+ .map(|(name, _)| Model::new(client_name, name))
.collect()
}
@@ -79,7 +79,7 @@ impl ErnieClient {
let (_, chat_endpoint) = MODELS
.iter()
.find(|(v, _)| v == &model)
- .ok_or_else(|| anyhow!("Miss Model '{}' in {}", model, self.model.client_name))?;
+ .ok_or_else(|| anyhow!("Miss Model '{}'", self.model.id()))?;
let url = format!("{API_BASE}{chat_endpoint}?access_token={}", unsafe {
&ACCESS_TOKEN
diff --git a/src/client/localai.rs b/src/client/localai.rs
index ecb0e63..93853dc 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -41,14 +41,14 @@ impl LocalAIClient {
),
];
- pub fn list_models(local_config: &LocalAIConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &LocalAIConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
local_config
.models
.iter()
.map(|v| {
- Model::new(client_index, client_name, &v.name)
+ Model::new(client_name, &v.name)
.set_max_tokens(v.max_tokens)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/model.rs b/src/client/model.rs
index dc30dcf..82d47c9 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -8,7 +8,6 @@ pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
#[derive(Debug, Clone)]
pub struct Model {
- pub client_index: usize,
pub client_name: String,
pub name: String,
pub max_tokens: Option<usize>,
@@ -17,14 +16,13 @@ pub struct Model {
impl Default for Model {
fn default() -> Self {
- Model::new(0, "", "")
+ Model::new("", "")
}
}
impl Model {
- pub fn new(client_index: usize, client_name: &str, name: &str) -> Self {
+ pub fn new(client_name: &str, name: &str) -> Self {
Self {
- client_index,
client_name: client_name.into(),
name: name.into(),
max_tokens: None,
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 7c10a4d..6b5edab 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -44,12 +44,12 @@ impl OpenAIClient {
pub const PROMPTS: [PromptType<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
- pub fn list_models(local_config: &OpenAIConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &OpenAIConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
MODELS
.into_iter()
.map(|(name, max_tokens)| {
- Model::new(client_index, client_name, name)
+ Model::new(client_name, name)
.set_max_tokens(Some(max_tokens))
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/palm.rs b/src/client/palm.rs
index 3720796..45496fa 100644
--- a/src/client/palm.rs
+++ b/src/client/palm.rs
@@ -8,9 +8,9 @@ use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
-const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta2";
+const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta2/models/";
-const MODELS: [(&str, usize, &str); 1] = [("chat-bison-001", 4096, "/models/chat-bison-001")];
+const MODELS: [(&str, usize); 1] = [("chat-bison-001", 4096)];
const TOKENS_COUNT_FACTORS: TokensCountFactors = (3, 8);
@@ -49,12 +49,12 @@ impl PaLMClient {
pub const PROMPTS: [PromptType<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
- pub fn list_models(local_config: &PaLMConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &PaLMConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, _)| {
- Model::new(client_index, client_name, name)
+ .map(|(name, max_tokens)| {
+ Model::new(client_name, name)
.set_max_tokens(Some(max_tokens))
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})
@@ -67,12 +67,8 @@ impl PaLMClient {
let body = build_body(data, self.model.name.clone());
let model = self.model.name.clone();
- let (_, _, endpoint) = MODELS
- .iter()
- .find(|(v, _, _)| v == &model)
- .ok_or_else(|| anyhow!("Miss Model '{}' in {}", model, self.model.client_name))?;
- let url = format!("{API_BASE}{endpoint}:generateMessage?key={}", api_key);
+ let url = format!("{API_BASE}{}:generateMessage?key={}", model, api_key);
let builder = client.post(url).json(&body);
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 162fcc2..b8e0d6f 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -54,11 +54,11 @@ impl QianwenClient {
pub const PROMPTS: [PromptType<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
- pub fn list_models(local_config: &QianwenConfig, client_index: usize) -> Vec<Model> {
+ pub fn list_models(local_config: &QianwenConfig) -> Vec<Model> {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens)| Model::new(client_index, client_name, name).set_max_tokens(Some(max_tokens)))
+ .map(|(name, max_tokens)| Model::new(client_name, name).set_max_tokens(Some(max_tokens)))
.collect()
}