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) --- src/client/qianwen.rs | 33 ++++++++++++++++++++++++--------- 1 file changed, 24 insertions(+), 9 deletions(-) (limited to 'src/client/qianwen.rs') 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