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/openai_compatible.rs | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) (limited to 'src/client/openai_compatible.rs') 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 +} -- cgit v1.2.3