diff options
| author | sigoden <sigoden@gmail.com> | 2024-08-17 16:01:39 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-08-17 16:01:39 +0800 |
| commit | 669f2c602c4631db1c91fd7a27098b7685027f9a (patch) | |
| tree | 3f743d4c2fa6e4eb847c6f9f2618b140f75b95b4 /src/client | |
| parent | 580ed6bea370345f76ca69ecb4c1cc30afa689c5 (diff) | |
| download | aichat-669f2c602c4631db1c91fd7a27098b7685027f9a.tar.gz | |
feat: enable custom `api_base` for most clients (#793)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/claude.rs | 10 | ||||
| -rw-r--r-- | src/client/cloudflare.rs | 8 | ||||
| -rw-r--r-- | src/client/cohere.rs | 27 | ||||
| -rw-r--r-- | src/client/gemini.rs | 21 | ||||
| -rw-r--r-- | src/client/openai.rs | 2 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 20 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 33 |
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); |
