From 9a8b302432a3f9bfa1e467dde027fc92dacce3e2 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 7 Nov 2023 10:56:28 +0800 Subject: refactor: remove Model.client_index, match client by name (#218) --- src/client/azure_openai.rs | 4 ++-- src/client/common.rs | 29 +++++++++++++---------------- src/client/ernie.rs | 6 +++--- src/client/localai.rs | 4 ++-- src/client/model.rs | 6 ++---- src/client/openai.rs | 4 ++-- src/client/palm.rs | 16 ++++++---------- src/client/qianwen.rs | 4 ++-- 8 files changed, 32 insertions(+), 41 deletions(-) (limited to 'src') 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 { + pub fn list_models(local_config: &AzureOpenAIConfig) -> Vec { 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> { 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 { + pub fn list_models(local_config: &ErnieConfig) -> Vec { 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 { + pub fn list_models(local_config: &LocalAIConfig) -> Vec { 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, @@ -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 { + pub fn list_models(local_config: &OpenAIConfig) -> Vec { 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 { + pub fn list_models(local_config: &PaLMConfig) -> Vec { 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 { + pub fn list_models(local_config: &QianwenConfig) -> Vec { 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() } -- cgit v1.2.3