From eec041c111c0ee170dab65942184e66c41479fcd Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 25 Mar 2024 10:52:05 +0800 Subject: feat: rename client localai to openai-compatible (#373) BREAKING CHANGE: rename client localai to openai-compatible --- src/client/localai.rs | 77 ----------------------------------------- src/client/mod.rs | 7 +++- src/client/openai_compatible.rs | 77 +++++++++++++++++++++++++++++++++++++++++ 3 files changed, 83 insertions(+), 78 deletions(-) delete mode 100644 src/client/localai.rs create mode 100644 src/client/openai_compatible.rs (limited to 'src/client') diff --git a/src/client/localai.rs b/src/client/localai.rs deleted file mode 100644 index 0e9db0e..0000000 --- a/src/client/localai.rs +++ /dev/null @@ -1,77 +0,0 @@ -use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{ExtraConfig, LocalAIClient, Model, ModelConfig, PromptType, SendData}; - -use crate::utils::PromptKind; - -use anyhow::Result; -use async_trait::async_trait; -use reqwest::{Client as ReqwestClient, RequestBuilder}; -use serde::Deserialize; - -#[derive(Debug, Clone, Deserialize)] -pub struct LocalAIConfig { - pub name: Option, - pub api_base: String, - pub api_key: Option, - pub chat_endpoint: Option, - pub models: Vec, - pub extra: Option, -} - -openai_compatible_client!(LocalAIClient); - -impl LocalAIClient { - config_get_fn!(api_key, get_api_key); - - pub const PROMPTS: [PromptType<'static>; 4] = [ - ("api_base", "API Base:", true, PromptKind::String), - ("api_key", "API Key:", false, PromptKind::String), - ("models[].name", "Model Name:", true, PromptKind::String), - ( - "models[].max_input_tokens", - "Max Input Tokens:", - false, - PromptKind::Integer, - ), - ]; - - pub fn list_models(local_config: &LocalAIConfig) -> Vec { - let client_name = Self::name(local_config); - - local_config - .models - .iter() - .map(|v| { - Model::new(client_name, &v.name) - .set_capabilities(v.capabilities) - .set_max_input_tokens(v.max_input_tokens) - .set_extra_fields(v.extra_fields.clone()) - .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) - }) - .collect() - } - - fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { - let api_key = self.get_api_key().ok(); - - let mut body = openai_build_body(data, self.model.name.clone()); - self.model.merge_extra_fields(&mut body); - - let chat_endpoint = self - .config - .chat_endpoint - .as_deref() - .unwrap_or("/chat/completions"); - - let url = format!("{}{chat_endpoint}", self.config.api_base); - - debug!("LocalAI Request: {url} {body}"); - - let mut builder = client.post(url).json(&body); - if let Some(api_key) = api_key { - builder = builder.bearer_auth(api_key); - } - - Ok(builder) - } -} diff --git a/src/client/mod.rs b/src/client/mod.rs index 7c7cfba..37775f3 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -12,7 +12,12 @@ register_client!( (gemini, "gemini", GeminiConfig, GeminiClient), (claude, "claude", ClaudeConfig, ClaudeClient), (mistral, "mistral", MistralConfig, MistralClient), - (localai, "localai", LocalAIConfig, LocalAIClient), + ( + openai_compatible, + "openai-compatible", + OpenAICompatibleConfig, + OpenAICompatibleClient + ), (ollama, "ollama", OllamaConfig, OllamaClient), ( azure_openai, diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs new file mode 100644 index 0000000..ec3333c --- /dev/null +++ b/src/client/openai_compatible.rs @@ -0,0 +1,77 @@ +use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; +use super::{ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptType, SendData}; + +use crate::utils::PromptKind; + +use anyhow::Result; +use async_trait::async_trait; +use reqwest::{Client as ReqwestClient, RequestBuilder}; +use serde::Deserialize; + +#[derive(Debug, Clone, Deserialize)] +pub struct OpenAICompatibleConfig { + pub name: Option, + pub api_base: String, + pub api_key: Option, + pub chat_endpoint: Option, + pub models: Vec, + pub extra: Option, +} + +openai_compatible_client!(OpenAICompatibleClient); + +impl OpenAICompatibleClient { + config_get_fn!(api_key, get_api_key); + + pub const PROMPTS: [PromptType<'static>; 4] = [ + ("api_base", "API Base:", true, PromptKind::String), + ("api_key", "API Key:", false, PromptKind::String), + ("models[].name", "Model Name:", true, PromptKind::String), + ( + "models[].max_input_tokens", + "Max Input Tokens:", + false, + PromptKind::Integer, + ), + ]; + + pub fn list_models(local_config: &OpenAICompatibleConfig) -> Vec { + let client_name = Self::name(local_config); + + local_config + .models + .iter() + .map(|v| { + Model::new(client_name, &v.name) + .set_capabilities(v.capabilities) + .set_max_input_tokens(v.max_input_tokens) + .set_extra_fields(v.extra_fields.clone()) + .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) + }) + .collect() + } + + fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { + let api_key = self.get_api_key().ok(); + + let mut body = openai_build_body(data, self.model.name.clone()); + self.model.merge_extra_fields(&mut body); + + let chat_endpoint = self + .config + .chat_endpoint + .as_deref() + .unwrap_or("/chat/completions"); + + let url = format!("{}{chat_endpoint}", self.config.api_base); + + debug!("OpenAICompatible Request: {url} {body}"); + + let mut builder = client.post(url).json(&body); + if let Some(api_key) = api_key { + builder = builder.bearer_auth(api_key); + } + + Ok(builder) + } +} -- cgit v1.2.3