From f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 27 Jul 2024 21:33:04 +0800 Subject: feat: support patching request url, headers and body (#756) --- src/client/rag_dedicated.rs | 40 +++++++++++++++------------------------- 1 file changed, 15 insertions(+), 25 deletions(-) (limited to 'src/client/rag_dedicated.rs') diff --git a/src/client/rag_dedicated.rs b/src/client/rag_dedicated.rs index 19f1626..7d2b846 100644 --- a/src/client/rag_dedicated.rs +++ b/src/client/rag_dedicated.rs @@ -4,7 +4,7 @@ use super::*; use anyhow::bail; use anyhow::Context; use anyhow::Result; -use reqwest::{Client as ReqwestClient, RequestBuilder}; +use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::json; use serde_json::Value; @@ -16,7 +16,7 @@ pub struct RagDedicatedConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -26,52 +26,42 @@ impl RagDedicatedClient { pub const PROMPTS: [PromptAction<'static>; 0] = []; - fn chat_completions_builder( - &self, - _client: &ReqwestClient, - _data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, _data: ChatCompletionsData) -> Result { bail!("The client doesn't support chat-completions api"); } - fn embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { let api_key = self.get_api_key().ok(); let api_base = self.get_api_base_ext()?; - let body = openai_build_embeddings_body(data, &self.model); - let url = format!("{api_base}/embeddings"); - debug!("RagDedicated Embeddings Request: {url} {body}"); + let body = openai_build_embeddings_body(data, &self.model); + + let mut request_data = RequestData::new(url, body); - let mut builder = client.post(url).json(&body); if let Some(api_key) = api_key { - builder = builder.bearer_auth(api_key); + request_data.bearer_auth(api_key); } - Ok(builder) + Ok(request_data) } - fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result { + fn prepare_rerank(&self, data: RerankData) -> Result { let api_key = self.get_api_key().ok(); let api_base = self.get_api_base_ext()?; - let body = rag_dedicated_build_rerank_body(data, &self.model); - let url = format!("{api_base}/rerank"); - debug!("RagDedicated Rerank Request: {url} {body}"); + let body = rag_dedicated_build_rerank_body(data, &self.model); + + let mut request_data = RequestData::new(url, body); - let mut builder = client.post(url).json(&body); if let Some(api_key) = api_key { - builder = builder.bearer_auth(api_key); + request_data.bearer_auth(api_key); } - Ok(builder) + Ok(request_data) } fn get_api_base_ext(&self) -> Result { -- cgit v1.2.3