summaryrefslogtreecommitdiffstats
path: root/src/client/openai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-27 21:33:04 +0800
committerGitHub <noreply@github.com>2024-07-27 21:33:04 +0800
commitf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (patch)
tree6f6a6ce299d7383d290a7015b52ce777566d1e8e /src/client/openai.rs
parentadf6716c8436b53854eb55588d258f4326580a4d (diff)
downloadaichat-f5e8e872b08d2aa9b3c063b0a98c59f842f94ecc.tar.gz
feat: support patching request url, headers and body (#756)
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs41
1 files changed, 17 insertions, 24 deletions
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -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<RequestBuilder> {
+ fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
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<RequestBuilder> {
+ fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
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)
}
}