summaryrefslogtreecommitdiffstats
path: root/src/client/openai_compatible.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-25 10:52:05 +0800
committerGitHub <noreply@github.com>2024-03-25 10:52:05 +0800
commiteec041c111c0ee170dab65942184e66c41479fcd (patch)
tree0e048d834fd2ae66ecd9f9c47aa7e9d21250e63e /src/client/openai_compatible.rs
parentbbd0c287261f8876d8c66a4e8015f5a0d8ce6548 (diff)
downloadaichat-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.rs77
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)
+ }
+}