diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-25 10:52:05 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-25 10:52:05 +0800 |
| commit | eec041c111c0ee170dab65942184e66c41479fcd (patch) | |
| tree | 0e048d834fd2ae66ecd9f9c47aa7e9d21250e63e /src/client/openai_compatible.rs | |
| parent | bbd0c287261f8876d8c66a4e8015f5a0d8ce6548 (diff) | |
| download | aichat-eec041c111c0ee170dab65942184e66c41479fcd.tar.gz | |
feat: rename client localai to openai-compatible (#373)
BREAKING CHANGE: rename client localai to openai-compatible
Diffstat (limited to 'src/client/openai_compatible.rs')
| -rw-r--r-- | src/client/openai_compatible.rs | 77 |
1 files changed, 77 insertions, 0 deletions
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<String>, + pub api_base: String, + pub api_key: Option<String>, + pub chat_endpoint: Option<String>, + pub models: Vec<ModelConfig>, + pub extra: Option<ExtraConfig>, +} + +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<Model> { + 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<RequestBuilder> { + 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) + } +} |
