From 669f2c602c4631db1c91fd7a27098b7685027f9a Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 17 Aug 2024 16:01:39 +0800 Subject: feat: enable custom `api_base` for most clients (#793) --- config.example.yaml | 30 +++++++++++++++++++----------- src/client/claude.rs | 10 ++++++++-- src/client/cloudflare.rs | 8 +++++++- src/client/cohere.rs | 27 ++++++++++++++++++++------- src/client/gemini.rs | 21 ++++++++++++++++++--- src/client/openai.rs | 2 +- src/client/openai_compatible.rs | 20 ++++++++++++-------- src/client/qianwen.rs | 33 ++++++++++++++++++++++++--------- 8 files changed, 109 insertions(+), 42 deletions(-) diff --git a/config.example.yaml b/config.example.yaml index 32e7f83..d536998 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -124,6 +124,7 @@ clients: # See https://ai.google.dev/docs - type: gemini api_key: xxx + api_base: https://generativelanguage.googleapis.com/v1beta # Optional patch: chat_completions: '.*': @@ -141,6 +142,7 @@ clients: # See https://docs.anthropic.com/claude/reference/getting-started-with-the-api - type: claude api_key: sk-ant-xxx + api_base: https://api.anthropic.com/v1 # Optional # See https://docs.mistral.ai/ - type: openai-compatible @@ -151,18 +153,19 @@ clients: # See https://docs.cohere.com/docs/the-cohere-platform - type: cohere api_key: xxx + api_base: https://api.cohere.ai/v1 # Optional # See https://docs.perplexity.ai/docs/getting-started - type: openai-compatible name: perplexity - api_base: https://api.perplexity.ai api_key: pplx-xxx + api_base: https://api.perplexity.ai # See https://console.groq.com/docs/quickstart - type: openai-compatible name: groq - api_base: https://api.groq.com/openai/v1 api_key: gsk_xxx + api_base: https://api.groq.com/openai/v1 # See https://github.com/jmorganca/ollama - type: ollama @@ -179,8 +182,8 @@ clients: # See https://learn.microsoft.com/en-us/azure/ai-services/openai/chatgpt-quickstart - type: azure-openai - api_base: https://{RESOURCE}.openai.azure.com api_key: xxx + api_base: https://{RESOURCE}.openai.azure.com models: - name: gpt-4o # Model deployment name max_input_tokens: 128000 @@ -219,6 +222,7 @@ clients: - type: cloudflare account_id: xxx api_key: xxx + api_base: https://api.cloudflare.com/client/v4 # Optional # See https://replicate.com/docs - type: replicate @@ -232,66 +236,70 @@ clients: # See https://help.aliyun.com/zh/dashscope/ - type: qianwen api_key: sk-xxx + api_base: https://dashscope.aliyuncs.com/api/v1 # Optional # See https://platform.moonshot.cn/docs/intro - type: openai-compatible name: moonshot - api_base: https://api.moonshot.cn/v1 api_key: sk-xxx + api_base: https://api.moonshot.cn/v1 # See https://platform.deepseek.com/api-docs/ - type: openai-compatible name: deepseek api_key: sk-xxx + api_base: https://api.deepseek.com # See https://open.bigmodel.cn/dev/howuse/introduction - type: openai-compatible name: zhipuai api_key: xxx + api_base: https://open.bigmodel.cn/api/paas/v4 # See https://platform.lingyiwanwu.com/docs - type: openai-compatible name: lingyiwanwu api_key: xxx + api_base: https://api.lingyiwanwu.com/v1 # See https://deepinfra.com/docs - type: openai-compatible name: deepinfra - api_base: https://api.deepinfra.com/v1/openai api_key: xxx + api_base: https://api.deepinfra.com/v1/openai # See https://readme.fireworks.ai/docs/quickstart - type: openai-compatible name: fireworks - api_base: https://api.fireworks.ai/inference/v1 api_key: xxx + api_base: https://api.fireworks.ai/inference/v1 # See https://openrouter.ai/docs#quick-start - type: openai-compatible name: openrouter - api_base: https://openrouter.ai/api/v1 api_key: xxx + api_base: https://openrouter.ai/api/v1 # See https://octo.ai/docs/getting-started/quickstart - type: openai-compatible name: octoai - api_base: https://text.octoai.run/v1 api_key: xxx + api_base: https://text.octoai.run/v1 # See https://docs.together.ai/docs/quickstart - type: openai-compatible name: together - api_base: https://api.together.xyz/v1 api_key: xxx + api_base: https://api.together.xyz/v1 # See https://jina.ai - type: openai-compatible name: jina - api_base: https://api.jina.ai/v1 api_key: xxx + api_base: https://api.jina.ai/v1 # See https://docs.voyageai.com/docs/introduction - type: openai-compatible name: voyageai + api_key: xxx api_base: https://api.voyageai.ai/v1 - api_key: xxx \ No newline at end of file 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, pub api_key: Option, + pub api_base: Option, #[serde(default)] pub models: Vec, pub patch: Option, @@ -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 { 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, pub account_id: Option, + pub api_base: Option, pub api_key: Option, #[serde(default)] pub models: Vec, @@ -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 { 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, pub api_key: Option, + pub api_base: Option, #[serde(default)] pub models: Vec, pub patch: Option, @@ -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 { 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 { 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 Result Result { 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, pub api_key: Option, + pub api_base: Option, #[serde(default)] pub models: Vec, pub patch: Option, @@ -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 { 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 { 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 { } } }; - Ok(api_base) + Ok(api_base.trim_end_matches('/').to_string()) } pub async fn generic_rerank(builder: RequestBuilder, _model: &Model) -> Result { @@ -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, pub api_key: Option, + pub api_base: Option, #[serde(default)] pub models: Vec, pub patch: Option, @@ -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 { 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 { 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