From 3a3388375be05758d5f5574cce1d0f8eb7bdd604 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 4 Dec 2024 21:03:59 +0800 Subject: refactor: improve retrieve model (#1036) - check the model type while retrieve model - select chat/reranker model even if it is missed in client models - find predefined-models for openai-compatible client with startsWith - remove client::ApiType --- src/client/macros.rs | 45 +++++++++++++++++++++++++++------------------ 1 file changed, 27 insertions(+), 18 deletions(-) (limited to 'src/client/macros.rs') diff --git a/src/client/macros.rs b/src/client/macros.rs index 0543cd6..4f52044 100644 --- a/src/client/macros.rs +++ b/src/client/macros.rs @@ -52,9 +52,10 @@ macro_rules! register_client { pub fn list_models(local_config: &$config) -> Vec { let client_name = Self::name(local_config); if local_config.models.is_empty() { - if let Some(models) = $crate::client::ALL_MODELS.iter().find(|v| { + if let Some(models) = $crate::client::ALL_PREDEFINED_MODELS.iter().find(|v| { v.platform == $name || - ($name == OpenAICompatibleClient::NAME && local_config.name.as_deref() == Some(&v.platform)) + ($name == OpenAICompatibleClient::NAME + && local_config.name.as_ref().map(|name| name.starts_with(&v.platform)).unwrap_or_default()) }) { return Model::from_config(client_name, &models.models); } @@ -98,32 +99,40 @@ macro_rules! register_client { anyhow::bail!("Unknown client '{}'", client) } - static ALL_CLIENT_MODELS: std::sync::OnceLock> = std::sync::OnceLock::new(); + static ALL_CLIENT_NAMES: std::sync::OnceLock> = std::sync::OnceLock::new(); - pub fn list_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { - let models = ALL_CLIENT_MODELS.get_or_init(|| { + pub fn list_client_names(config: &$crate::config::Config) -> Vec<&'static String> { + let names = ALL_CLIENT_NAMES.get_or_init(|| { config .clients .iter() .flat_map(|v| match v { - $(ClientConfig::$config(c) => $client::list_models(c),)+ + $(ClientConfig::$config(c) => vec![$client::name(c).to_string()],)+ ClientConfig::Unknown => vec![], }) .collect() }); - models.iter().collect() + names.iter().collect() } - pub fn list_chat_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { - list_models(config).into_iter().filter(|v| v.model_type() == "chat").collect() - } + static ALL_MODELS: std::sync::OnceLock> = std::sync::OnceLock::new(); - pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { - list_models(config).into_iter().filter(|v| v.model_type() == "embedding").collect() + pub fn list_all_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + let models = ALL_MODELS.get_or_init(|| { + config + .clients + .iter() + .flat_map(|v| match v { + $(ClientConfig::$config(c) => $client::list_models(c),)+ + ClientConfig::Unknown => vec![], + }) + .collect() + }); + models.iter().collect() } - pub fn list_reranker_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { - list_models(config).into_iter().filter(|v| v.model_type() == "reranker").collect() + pub fn list_models(config: &$crate::config::Config, model_type: $crate::client::ModelType) -> Vec<&'static $crate::client::Model> { + list_all_models(config).into_iter().filter(|v| v.model_type() == model_type).collect() } }; } @@ -175,7 +184,7 @@ macro_rules! impl_client_trait { data: $crate::client::ChatCompletionsData, ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> { let request_data = $prepare_chat_completions(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); + let builder = self.request_builder(client, request_data); $chat_completions(builder, self.model()).await } @@ -186,7 +195,7 @@ macro_rules! impl_client_trait { data: $crate::client::ChatCompletionsData, ) -> Result<()> { let request_data = $prepare_chat_completions(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::ChatCompletions); + let builder = self.request_builder(client, request_data); $chat_completions_streaming(builder, handler, self.model()).await } @@ -196,7 +205,7 @@ macro_rules! impl_client_trait { data: &$crate::client::EmbeddingsData, ) -> Result<$crate::client::EmbeddingsOutput> { let request_data = $prepare_embeddings(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::Embeddings); + let builder = self.request_builder(client, request_data); $embeddings(builder, self.model()).await } @@ -206,7 +215,7 @@ macro_rules! impl_client_trait { data: &$crate::client::RerankData, ) -> Result<$crate::client::RerankOutput> { let request_data = $prepare_rerank(self, data)?; - let builder = self.request_builder(client, request_data, ApiType::Rerank); + let builder = self.request_builder(client, request_data); $rerank(builder, self.model()).await } } -- cgit v1.2.3