diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-27 21:33:04 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-27 21:33:04 +0800 |
| commit | f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (patch) | |
| tree | 6f6a6ce299d7383d290a7015b52ce777566d1e8e /src/client/openai_compatible.rs | |
| parent | adf6716c8436b53854eb55588d258f4326580a4d (diff) | |
| download | aichat-f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc.tar.gz | |
feat: support patching request url, headers and body (#756)
Diffstat (limited to 'src/client/openai_compatible.rs')
| -rw-r--r-- | src/client/openai_compatible.rs | 51 |
1 files changed, 19 insertions, 32 deletions
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index c59eff6..a2302de 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -3,7 +3,6 @@ use super::rag_dedicated::*; use super::*; use anyhow::Result; -use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; #[derive(Debug, Clone, Deserialize)] @@ -14,7 +13,7 @@ pub struct OpenAICompatibleConfig { pub chat_endpoint: Option<String>, #[serde(default)] pub models: Vec<ModelData>, - pub patch: Option<ModelPatch>, + pub patch: Option<RequestPatch>, pub extra: Option<ExtraConfig>, } @@ -35,17 +34,10 @@ impl OpenAICompatibleClient { ), ]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result<RequestBuilder> { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> { let api_key = self.get_api_key().ok(); let api_base = self.get_api_base_ext()?; - let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_chat_completions_body(&mut body); - let chat_endpoint = self .config .chat_endpoint @@ -54,54 +46,49 @@ impl OpenAICompatibleClient { let url = format!("{api_base}{chat_endpoint}"); - debug!("OpenAICompatible Chat Completions Request: {url} {body}"); + let body = openai_build_chat_completions_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 embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result<RequestBuilder> { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> { 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!("OpenAICompatible 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<RequestBuilder> { + fn prepare_rerank(&self, data: RerankData) -> Result<RequestData> { 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!("OpenAICompatible 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<String> { |
