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/openai.rs | 41 +++++++++++++++++------------------------ 1 file changed, 17 insertions(+), 24 deletions(-) (limited to 'src/client/openai.rs') diff --git a/src/client/openai.rs b/src/client/openai.rs index a494056..2b83b7d 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,7 +1,7 @@ use super::*; use anyhow::{bail, Context, Result}; -use reqwest::{Client as ReqwestClient, RequestBuilder}; +use reqwest::RequestBuilder; use serde::Deserialize; use serde_json::{json, Value}; @@ -15,7 +15,7 @@ pub struct OpenAIConfig { pub organization_id: Option, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -26,47 +26,40 @@ impl OpenAIClient { pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, 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 mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_chat_completions_body(&mut body); - let url = format!("{api_base}/chat/completions"); - debug!("OpenAI Chat Completions Request: {url} {body}"); + let body = openai_build_chat_completions_body(data, &self.model); - let mut builder = client.post(url).bearer_auth(api_key).json(&body); + let mut request_data = RequestData::new(url, body); + request_data.bearer_auth(api_key); if let Some(organization_id) = &self.config.organization_id { - builder = builder.header("OpenAI-Organization", organization_id); + request_data.header("OpenAI-Organization", organization_id); } - Ok(builder) + Ok(request_data) } - fn embeddings_builder( - &self, - client: &ReqwestClient, - data: EmbeddingsData, - ) -> Result { + fn prepare_embeddings(&self, 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 body = openai_build_embeddings_body(data, &self.model); - let url = format!("{api_base}/embeddings"); - debug!("OpenAI Embeddings Request: {url} {body}"); + let body = openai_build_embeddings_body(data, &self.model); - let builder = client.post(url).bearer_auth(api_key).json(&body); + let mut request_data = RequestData::new(url, body); + + request_data.bearer_auth(api_key); + if let Some(organization_id) = &self.config.organization_id { + request_data.header("OpenAI-Organization", organization_id); + } - Ok(builder) + Ok(request_data) } } -- cgit v1.2.3