summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-08-17 16:01:39 +0800
committerGitHub <noreply@github.com>2024-08-17 16:01:39 +0800
commit669f2c602c4631db1c91fd7a27098b7685027f9a (patch)
tree3f743d4c2fa6e4eb847c6f9f2618b140f75b95b4 /src
parent580ed6bea370345f76ca69ecb4c1cc30afa689c5 (diff)
downloadaichat-669f2c602c4631db1c91fd7a27098b7685027f9a.tar.gz
feat: enable custom `api_base` for most clients (#793)
Diffstat (limited to 'src')
-rw-r--r--src/client/claude.rs10
-rw-r--r--src/client/cloudflare.rs8
-rw-r--r--src/client/cohere.rs27
-rw-r--r--src/client/gemini.rs21
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/openai_compatible.rs20
-rw-r--r--src/client/qianwen.rs33
7 files changed, 90 insertions, 31 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 8811476..7b472a7 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -5,12 +5,13 @@ use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
-const API_BASE: &str = "https://api.anthropic.com/v1/messages";
+const API_BASE: &str = "https://api.anthropic.com/v1";
#[derive(Debug, Clone, Deserialize)]
pub struct ClaudeConfig {
pub name: Option<String>,
pub api_key: Option<String>,
+ pub api_base: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patch: Option<RequestPatch>,
@@ -19,6 +20,7 @@ pub struct ClaudeConfig {
impl ClaudeClient {
config_get_fn!(api_key, get_api_key);
+ config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
@@ -40,10 +42,14 @@ fn prepare_chat_completions(
data: ChatCompletionsData,
) -> Result<RequestData> {
let api_key = self_.get_api_key().ok();
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
+ let url = format!("{}/messages", api_base.trim_end_matches('/'));
let body = claude_build_chat_completions_body(data, &self_.model)?;
- let mut request_data = RequestData::new(API_BASE, body);
+ let mut request_data = RequestData::new(url, body);
request_data.header("anthropic-version", "2023-06-01");
if let Some(api_key) = api_key {
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index a24a1c6..3626c73 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -11,6 +11,7 @@ const API_BASE: &str = "https://api.cloudflare.com/client/v4";
pub struct CloudflareConfig {
pub name: Option<String>,
pub account_id: Option<String>,
+ pub api_base: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
@@ -21,6 +22,7 @@ pub struct CloudflareConfig {
impl CloudflareClient {
config_get_fn!(account_id, get_account_id);
config_get_fn!(api_key, get_api_key);
+ config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptAction<'static>; 2] = [
("account_id", "Account ID:", true, PromptKind::String),
@@ -45,9 +47,13 @@ fn prepare_chat_completions(
) -> Result<RequestData> {
let account_id = self_.get_account_id()?;
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
let url = format!(
- "{API_BASE}/accounts/{account_id}/ai/run/{}",
+ "{}/accounts/{account_id}/ai/run/{}",
+ api_base.trim_end_matches('/'),
self_.model.name()
);
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index aff919e..64263b7 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,19 +1,18 @@
-use super::*;
use super::openai_compatible::*;
+use super::*;
use anyhow::{bail, Context, Result};
use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
-const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat";
-const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed";
-const RERANK_API_URL: &str = "https://api.cohere.ai/v1/rerank";
+const API_BASE: &str = "https://api.cohere.ai/v1";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct CohereConfig {
pub name: Option<String>,
pub api_key: Option<String>,
+ pub api_base: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patch: Option<RequestPatch>,
@@ -22,6 +21,7 @@ pub struct CohereConfig {
impl CohereClient {
config_get_fn!(api_key, get_api_key);
+ config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
@@ -43,10 +43,14 @@ fn prepare_chat_completions(
data: ChatCompletionsData,
) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
+ let url = format!("{}/chat", api_base.trim_end_matches('/'));
let body = build_chat_completions_body(data, &self_.model)?;
- let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body);
+ let mut request_data = RequestData::new(url, body);
request_data.bearer_auth(api_key);
@@ -55,6 +59,11 @@ fn prepare_chat_completions(
fn prepare_embeddings(self_: &CohereClient, data: EmbeddingsData) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
+
+ let url = format!("{}/embed", api_base.trim_end_matches('/'));
let input_type = match data.query {
true => "search_query",
@@ -67,7 +76,7 @@ fn prepare_embeddings(self_: &CohereClient, data: EmbeddingsData) -> Result<Requ
"input_type": input_type,
});
- let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
+ let mut request_data = RequestData::new(url, body);
request_data.bearer_auth(api_key);
@@ -76,10 +85,14 @@ fn prepare_embeddings(self_: &CohereClient, data: EmbeddingsData) -> Result<Requ
fn prepare_rerank(self_: &CohereClient, data: RerankData) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
+ let url = format!("{}/rerank", api_base.trim_end_matches('/'));
let body = generic_build_rerank_body(data, &self_.model);
- let mut request_data = RequestData::new(RERANK_API_URL, body);
+ let mut request_data = RequestData::new(url, body);
request_data.bearer_auth(api_key);
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 2616218..572e082 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -6,12 +6,13 @@ use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
-const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/";
+const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct GeminiConfig {
pub name: Option<String>,
pub api_key: Option<String>,
+ pub api_base: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patch: Option<RequestPatch>,
@@ -20,6 +21,7 @@ pub struct GeminiConfig {
impl GeminiClient {
config_get_fn!(api_key, get_api_key);
+ config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
@@ -41,13 +43,22 @@ fn prepare_chat_completions(
data: ChatCompletionsData,
) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
let func = match data.stream {
true => "streamGenerateContent",
false => "generateContent",
};
- let url = format!("{API_BASE}{}:{}?key={}", self_.model.name(), func, api_key);
+ let url = format!(
+ "{}/models/{}:{}?key={}",
+ api_base.trim_end_matches('/'),
+ self_.model.name(),
+ func,
+ api_key
+ );
let body = gemini_build_chat_completions_body(data, &self_.model)?;
@@ -58,9 +69,13 @@ fn prepare_chat_completions(
fn prepare_embeddings(self_: &GeminiClient, data: EmbeddingsData) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
let url = format!(
- "{API_BASE}{}:embedContent?key={}",
+ "{}/models/{}:embedContent?key={}",
+ api_base.trim_end_matches('/'),
self_.model.name(),
api_key
);
diff --git a/src/client/openai.rs b/src/client/openai.rs
index ec9cb9d..3e1707f 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -47,7 +47,7 @@ fn prepare_chat_completions(
.get_api_base()
.unwrap_or_else(|_| API_BASE.to_string());
- let url = format!("{api_base}/chat/completions");
+ let url = format!("{}/chat/completions", api_base.trim_end_matches('/'));
let body = openai_build_chat_completions_body(data, &self_.model);
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 8acac58..2bde884 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -36,7 +36,6 @@ impl OpenAICompatibleClient {
];
}
-
impl_client_trait!(
OpenAICompatibleClient,
(
@@ -55,11 +54,16 @@ fn prepare_chat_completions(
let api_key = self_.get_api_key().ok();
let api_base = get_api_base_ext(self_)?;
- let chat_endpoint = self_
- .config
- .chat_endpoint
- .as_deref()
- .unwrap_or("/chat/completions");
+ let chat_endpoint = match self_.config.chat_endpoint.clone() {
+ Some(v) => {
+ if v.starts_with('/') {
+ v
+ } else {
+ format!("/{}", v)
+ }
+ }
+ None => "/chat/completions".into(),
+ };
let url = format!("{api_base}{chat_endpoint}");
@@ -126,7 +130,7 @@ fn get_api_base_ext(self_: &OpenAICompatibleClient) -> Result<String> {
}
}
};
- Ok(api_base)
+ Ok(api_base.trim_end_matches('/').to_string())
}
pub async fn generic_rerank(builder: RequestBuilder, _model: &Model) -> Result<RerankOutput> {
@@ -171,4 +175,4 @@ pub fn generic_build_rerank_body(data: RerankData, model: &Model) -> Value {
body["top_n"] = top_n.into()
}
body
-} \ No newline at end of file
+}
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 3c246f9..38534d8 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -11,19 +11,19 @@ use serde::Deserialize;
use serde_json::{json, Value};
use std::borrow::BorrowMut;
-const CHAT_COMPLETIONS_API_URL: &str =
- "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation";
+const API_BASE: &str = "https://dashscope.aliyuncs.com/api/v1";
-const CHAT_COMPLETIONS_API_URL_VL: &str =
- "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation";
+const CHAT_COMPLETIONS_ENDPOINT: &str = "/services/aigc/text-generation/generation";
-const EMBEDDINGS_API_URL: &str =
- "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding";
+const CHAT_COMPLETIONS_VL_ENDPOINT: &str = "/services/aigc/multimodal-generation/generation";
+
+const EMBEDDINGS_ENDPOINT: &str = "/services/embeddings/text-embedding/text-embedding";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct QianwenConfig {
pub name: Option<String>,
pub api_key: Option<String>,
+ pub api_base: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patch: Option<RequestPatch>,
@@ -32,6 +32,7 @@ pub struct QianwenConfig {
impl QianwenClient {
config_get_fn!(api_key, get_api_key);
+ config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
@@ -82,12 +83,21 @@ fn prepare_chat_completions(
data: ChatCompletionsData,
) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
let stream = data.stream;
let url = match self_.model().supports_vision() {
- true => CHAT_COMPLETIONS_API_URL_VL,
- false => CHAT_COMPLETIONS_API_URL,
+ true => format!(
+ "{}{CHAT_COMPLETIONS_VL_ENDPOINT}",
+ api_base.trim_end_matches('/'),
+ ),
+ false => format!(
+ "{}{CHAT_COMPLETIONS_ENDPOINT}",
+ api_base.trim_end_matches('/'),
+ ),
};
let (body, has_upload) = build_chat_completions_body(data, &self_.model)?;
@@ -108,6 +118,11 @@ fn prepare_chat_completions(
fn prepare_embeddings(self_: &QianwenClient, data: EmbeddingsData) -> Result<RequestData> {
let api_key = self_.get_api_key()?;
+ let api_base = self_
+ .get_api_base()
+ .unwrap_or_else(|_| API_BASE.to_string());
+
+ let url = format!("{}{EMBEDDINGS_ENDPOINT}", api_base.trim_end_matches('/'),);
let text_type = match data.query {
true => "query",
@@ -124,7 +139,7 @@ fn prepare_embeddings(self_: &QianwenClient, data: EmbeddingsData) -> Result<Req
}
});
- let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
+ let mut request_data = RequestData::new(url, body);
request_data.bearer_auth(api_key);