summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs36
-rw-r--r--src/client/bedrock.rs84
-rw-r--r--src/client/claude.rs24
-rw-r--r--src/client/cloudflare.rs41
-rw-r--r--src/client/cohere.rs45
-rw-r--r--src/client/common.rs164
-rw-r--r--src/client/ernie.rs65
-rw-r--r--src/client/gemini.rs43
-rw-r--r--src/client/macros.rs29
-rw-r--r--src/client/ollama.rs39
-rw-r--r--src/client/openai.rs41
-rw-r--r--src/client/openai_compatible.rs51
-rw-r--r--src/client/qianwen.rs46
-rw-r--r--src/client/rag_dedicated.rs40
-rw-r--r--src/client/replicate.rs24
-rw-r--r--src/client/vertexai.rs38
16 files changed, 415 insertions, 395 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 2c4df05..8b583e0 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -2,7 +2,6 @@ use super::openai::*;
use super::*;
use anyhow::Result;
-use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
#[derive(Debug, Clone, Deserialize)]
@@ -12,7 +11,7 @@ pub struct AzureOpenAIConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -32,51 +31,42 @@ impl AzureOpenAIClient {
),
];
- fn chat_completions_builder(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<RequestBuilder> {
+ fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
let api_base = self.get_api_base()?;
let api_key = self.get_api_key()?;
- let mut body = openai_build_chat_completions_body(data, &self.model);
- self.patch_chat_completions_body(&mut body);
-
let url = format!(
"{}/openai/deployments/{}/chat/completions?api-version=2024-02-01",
&api_base,
self.model.name()
);
- debug!("AzureOpenAI Chat Completions Request: {url} {body}");
+ let body = openai_build_chat_completions_body(data, &self.model);
+
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).header("api-key", api_key).json(&body);
+ request_data.header("api-key", 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_base = self.get_api_base()?;
let api_key = self.get_api_key()?;
- let body = openai_build_embeddings_body(data, &self.model);
-
let url = format!(
"{}/openai/deployments/{}/embeddings?api-version=2024-02-01",
&api_base,
self.model.name()
);
- debug!("AzureOpenAI Embeddings Request: {url} {body}");
+ let body = openai_build_embeddings_body(data, &self.model);
+
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).header("api-key", api_key).json(&body);
+ request_data.header("api-key", api_key);
- Ok(builder)
+ Ok(request_data)
}
}
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 50ce6f1..2b1ca6e 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -3,19 +3,16 @@ use super::*;
use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256};
use anyhow::{bail, Context, Result};
+use async_trait::async_trait;
use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder};
use aws_smithy_eventstream::smithy::parse_response_headers;
use bytes::BytesMut;
use chrono::{DateTime, Utc};
use futures_util::StreamExt;
use indexmap::IndexMap;
-use reqwest::{
- header::{HeaderMap, HeaderName, HeaderValue},
- Client as ReqwestClient, Method, RequestBuilder,
-};
+use reqwest::{Client as ReqwestClient, Method, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
-use std::str::FromStr;
#[derive(Debug, Clone, Deserialize)]
pub struct BedrockConfig {
@@ -25,7 +22,7 @@ pub struct BedrockConfig {
pub region: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -61,16 +58,22 @@ impl BedrockClient {
let host = format!("bedrock-runtime.{region}.amazonaws.com");
let model_name = &self.model.name();
+
let uri = if data.stream {
format!("/model/{model_name}/converse-stream")
} else {
format!("/model/{model_name}/converse")
};
- let headers = IndexMap::new();
+ let body = build_chat_completions_body(data, &self.model)?;
- let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
+ let mut request_data = RequestData::new("", body);
+ self.patch_request_data(&mut request_data, ApiType::ChatCompletions);
+ let RequestData {
+ url: _,
+ headers,
+ body,
+ } = request_data;
let builder = aws_fetch(
client,
@@ -105,8 +108,6 @@ impl BedrockClient {
let uri = format!("/model/{}/invoke", self.model.name());
- let headers = IndexMap::new();
-
let input_type = match data.query {
true => "search_query",
false => "search_document",
@@ -117,6 +118,14 @@ impl BedrockClient {
"input_type": input_type,
});
+ let mut request_data = RequestData::new("", body);
+ self.patch_request_data(&mut request_data, ApiType::Embeddings);
+ let RequestData {
+ url: _,
+ headers,
+ body,
+ } = request_data;
+
let builder = aws_fetch(
client,
&AwsCredentials {
@@ -139,12 +148,38 @@ impl BedrockClient {
}
}
-impl_client_trait!(
- BedrockClient,
- chat_completions,
- chat_completions_streaming,
- embeddings
-);
+#[async_trait]
+impl Client for BedrockClient {
+ client_common_fns!();
+
+ async fn chat_completions_inner(
+ &self,
+ client: &ReqwestClient,
+ data: ChatCompletionsData,
+ ) -> Result<ChatCompletionsOutput> {
+ let builder = self.chat_completions_builder(client, data)?;
+ chat_completions(builder).await
+ }
+
+ async fn chat_completions_streaming_inner(
+ &self,
+ client: &ReqwestClient,
+ handler: &mut SseHandler,
+ data: ChatCompletionsData,
+ ) -> Result<()> {
+ let builder = self.chat_completions_builder(client, data)?;
+ chat_completions_streaming(builder, handler).await
+ }
+
+ async fn embeddings_inner(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<EmbeddingsOutput> {
+ let builder = self.embeddings_builder(client, data)?;
+ embeddings(builder).await
+ }
+}
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -550,17 +585,14 @@ fn aws_fetch(
headers.insert("authorization".into(), authorization_header);
- let mut req_headers = HeaderMap::new();
- for (k, v) in &headers {
- req_headers.insert(HeaderName::from_str(k)?, HeaderValue::from_str(v)?);
- }
+ debug!("Request {endpoint} {body}");
- debug!("Bedrock Request: {endpoint} {body}");
+ let mut request_builder = client.request(method, endpoint).body(body);
+
+ for (key, value) in &headers {
+ request_builder = request_builder.header(key, value);
+ }
- let request_builder = client
- .request(method, endpoint)
- .headers(req_headers)
- .body(body);
Ok(request_builder)
}
diff --git a/src/client/claude.rs b/src/client/claude.rs
index df0c034..8a7b14c 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.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};
@@ -13,7 +13,7 @@ pub struct ClaudeConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -23,27 +23,19 @@ impl ClaudeClient {
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().ok();
- let mut body = claude_build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
+ let body = claude_build_chat_completions_body(data, &self.model)?;
- let url = API_BASE;
+ let mut request_data = RequestData::new(API_BASE, body);
- debug!("Claude Request: {url} {body}");
-
- let mut builder = client.post(url).json(&body);
- builder = builder.header("anthropic-version", "2023-06-01");
+ request_data.header("anthropic-version", "2023-06-01");
if let Some(api_key) = api_key {
- builder = builder.header("x-api-key", api_key)
+ request_data.header("x-api-key", api_key)
}
- Ok(builder)
+ Ok(request_data)
}
}
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 439abac..3ee0a91 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,7 +1,7 @@
use super::*;
use anyhow::{anyhow, Context, Result};
-use reqwest::{Client as ReqwestClient, RequestBuilder};
+use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
@@ -14,7 +14,7 @@ pub struct CloudflareConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -27,51 +27,42 @@ impl CloudflareClient {
("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 account_id = self.get_account_id()?;
let api_key = self.get_api_key()?;
- let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
-
let url = format!(
"{API_BASE}/accounts/{account_id}/ai/run/{}",
self.model.name()
);
- debug!("Cloudflare Chat Completions Request: {url} {body}");
+ let body = build_chat_completions_body(data, &self.model)?;
+
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ 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 account_id = self.get_account_id()?;
let api_key = self.get_api_key()?;
- let body = json!({
- "text": data.texts,
- });
-
let url = format!(
"{API_BASE}/accounts/{account_id}/ai/run/{}",
self.model.name()
);
- debug!("Cloudflare Embeddings Request: {url} {body}");
+ let body = json!({
+ "text": data.texts,
+ });
+
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ request_data.bearer_auth(api_key);
- Ok(builder)
+ Ok(request_data)
}
}
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 4ab6f38..b6e9755 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -2,7 +2,7 @@ use super::rag_dedicated::*;
use super::*;
use anyhow::{bail, Context, Result};
-use reqwest::{Client as ReqwestClient, RequestBuilder};
+use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
@@ -16,7 +16,7 @@ pub struct CohereConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -26,30 +26,19 @@ impl CohereClient {
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 mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
+ let body = build_chat_completions_body(data, &self.model)?;
- let url = CHAT_COMPLETIONS_API_URL;
+ let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body);
- debug!("Cohere Chat Completions Request: {url} {body}");
+ request_data.bearer_auth(api_key);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
-
- 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 input_type = match data.query {
@@ -63,27 +52,23 @@ impl CohereClient {
"input_type": input_type,
});
- let url = EMBEDDINGS_API_URL;
-
- debug!("Cohere Embeddings Request: {url} {body}");
+ let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ 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()?;
let body = rag_dedicated_build_rerank_body(data, &self.model);
- let url = RERANK_API_URL;
-
- debug!("Cohere Rerank Request: {url} {body}");
+ let mut request_data = RequestData::new(RERANK_API_URL, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ request_data.bearer_auth(api_key);
- Ok(builder)
+ Ok(request_data)
}
}
diff --git a/src/client/common.rs b/src/client/common.rs
index ea7aee5..3321511 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -31,7 +31,7 @@ pub trait Client: Sync + Send {
fn extra_config(&self) -> Option<&ExtraConfig>;
- fn patch_config(&self) -> Option<&ModelPatch>;
+ fn patch_config(&self) -> Option<&RequestPatch>;
fn name(&self) -> &str;
@@ -110,13 +110,40 @@ pub trait Client: Sync + Send {
.context("Failed to call rerank api")
}
- fn patch_chat_completions_body(&self, body: &mut Value) {
- if let Some(patch) = extract_chat_completions_body_patch(
- self.patch_config().map(|v| v.chat_completions_body.clone()),
- self.model(),
- ) {
- if body.is_object() && patch.is_object() {
- json_patch::merge(body, &patch)
+ fn request_builder(
+ &self,
+ client: &reqwest::Client,
+ mut request_data: RequestData,
+ api_type: ApiType,
+ ) -> RequestBuilder {
+ self.patch_request_data(&mut request_data, api_type);
+ request_data.into_builder(client)
+ }
+
+ fn patch_request_data(&self, request_data: &mut RequestData, api_type: ApiType) {
+ let map = std::env::var(get_env_name(&format!(
+ "patch_{}_{}",
+ self.model().client_name(),
+ api_type.name(),
+ )))
+ .ok()
+ .and_then(|v| serde_json::from_str(&v).ok())
+ .or_else(|| {
+ self.patch_config()
+ .and_then(|v| api_type.extract_patch(v))
+ .cloned()
+ });
+ let map = match map {
+ Some(v) => v,
+ _ => return,
+ };
+ for (key, patch) in map {
+ let key = ESCAPE_SLASH_RE.replace_all(&key, r"\/");
+ if let Ok(regex) = Regex::new(&format!("^({key})$")) {
+ if let Ok(true) = regex.is_match(self.model().name()) {
+ request_data.apply_patch(patch);
+ return;
+ }
}
}
}
@@ -164,32 +191,99 @@ pub struct ExtraConfig {
}
#[derive(Debug, Clone, Deserialize, Default)]
-pub struct ModelPatch {
- pub chat_completions_body: ChatCompletionsBodyPatch,
+pub struct RequestPatch {
+ pub chat_completions: Option<ApiPatch>,
+ pub embeddings: Option<ApiPatch>,
+ pub rerank: Option<ApiPatch>,
}
-pub type ChatCompletionsBodyPatch = IndexMap<String, Value>;
-
-pub fn extract_chat_completions_body_patch(
- patch: Option<ChatCompletionsBodyPatch>,
- model: &Model,
-) -> Option<Value> {
- let patch = std::env::var(get_env_name(&format!(
- "{}_chat_completions_body_patch",
- model.client_name()
- )))
- .ok()
- .and_then(|v| serde_json::from_str(&v).ok())
- .or(patch)?;
- for (key, patch_data) in patch {
- let key = ESCAPE_SLASH_RE.replace_all(&key, r"\/");
- if let Ok(regex) = Regex::new(&format!("^({key})$")) {
- if let Ok(true) = regex.is_match(model.name()) {
- return Some(patch_data);
+pub type ApiPatch = IndexMap<String, Value>;
+
+#[derive(Debug, Clone, Copy, PartialEq, Eq)]
+pub enum ApiType {
+ ChatCompletions,
+ Embeddings,
+ Rerank,
+}
+
+impl ApiType {
+ pub fn name(&self) -> &str {
+ match self {
+ ApiType::ChatCompletions => "chat_completions",
+ ApiType::Embeddings => "embeddings",
+ ApiType::Rerank => "rerank",
+ }
+ }
+ pub fn extract_patch<'a>(&self, patch: &'a RequestPatch) -> Option<&'a ApiPatch> {
+ match self {
+ ApiType::ChatCompletions => patch.chat_completions.as_ref(),
+ ApiType::Embeddings => patch.embeddings.as_ref(),
+ ApiType::Rerank => patch.rerank.as_ref(),
+ }
+ }
+}
+
+pub struct RequestData {
+ pub url: String,
+ pub headers: IndexMap<String, String>,
+ pub body: Value,
+}
+
+impl RequestData {
+ pub fn new<T>(url: T, body: Value) -> Self
+ where
+ T: std::fmt::Display,
+ {
+ Self {
+ url: url.to_string(),
+ headers: Default::default(),
+ body,
+ }
+ }
+
+ pub fn bearer_auth<T>(&mut self, auth: T)
+ where
+ T: std::fmt::Display,
+ {
+ self.headers
+ .insert("authorization".into(), format!("Bearer {auth}"));
+ }
+
+ pub fn header<K, V>(&mut self, key: K, value: V)
+ where
+ K: std::fmt::Display,
+ V: std::fmt::Display,
+ {
+ self.headers.insert(key.to_string(), value.to_string());
+ }
+
+ pub fn into_builder(self, client: &ReqwestClient) -> RequestBuilder {
+ let RequestData { url, headers, body } = self;
+ debug!("Request {url} {body}");
+
+ let mut builder = client.post(url);
+ for (key, value) in headers {
+ builder = builder.header(key, value);
+ }
+ builder = builder.json(&body);
+ builder
+ }
+
+ pub fn apply_patch(&mut self, patch: Value) {
+ if let Some(patch_url) = patch["url"].as_str() {
+ self.url = patch_url.into();
+ }
+ if let Some(patch_body) = patch.get("body") {
+ json_patch::merge(&mut self.body, patch_body)
+ }
+ if let Some(patch_headers) = patch["headers"].as_object() {
+ for (key, value) in patch_headers {
+ if let Some(value) = value.as_str() {
+ self.header(key, value)
+ }
}
}
}
- None
}
#[derive(Debug)]
@@ -445,24 +539,24 @@ fn set_client_config(
fn set_client_config_value(client_config: &mut Value, path: &str, kind: &PromptKind, value: &str) {
let segs: Vec<&str> = path.split('.').collect();
match segs.as_slice() {
- [name] => client_config[name] = to_json(kind, value),
+ [name] => client_config[name] = prompt_value_to_json(kind, value),
[scope, name] => match scope.split_once('[') {
None => {
if client_config.get(scope).is_none() {
let mut obj = json!({});
- obj[name] = to_json(kind, value);
+ obj[name] = prompt_value_to_json(kind, value);
client_config[scope] = obj;
} else {
- client_config[scope][name] = to_json(kind, value);
+ client_config[scope][name] = prompt_value_to_json(kind, value);
}
}
Some((scope, _)) => {
if client_config.get(scope).is_none() {
let mut obj = json!({});
- obj[name] = to_json(kind, value);
+ obj[name] = prompt_value_to_json(kind, value);
client_config[scope] = json!([obj]);
} else {
- client_config[scope][0][name] = to_json(kind, value);
+ client_config[scope][0][name] = prompt_value_to_json(kind, value);
}
}
},
@@ -470,7 +564,7 @@ fn set_client_config_value(client_config: &mut Value, path: &str, kind: &PromptK
}
}
-fn to_json(kind: &PromptKind, value: &str) -> Value {
+fn prompt_value_to_json(kind: &PromptKind, value: &str) -> Value {
if value.is_empty() {
return Value::Null;
}
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index f1d432d..e508978 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -19,7 +19,7 @@ pub struct ErnieConfig {
pub secret_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -29,54 +29,46 @@ impl ErnieClient {
("secret_key", "Secret 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 access_token = get_access_token(self.name())?;
- let mut body = build_chat_completions_body(data, &self.model);
- self.patch_chat_completions_body(&mut body);
-
let url = format!(
"{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}",
&self.model.name(),
);
- debug!("Ernie Chat Completions Request: {url} {body}");
+ let body = build_chat_completions_body(data, &self.model);
- let builder = client.post(url).json(&body);
+ let request_data = RequestData::new(url, body);
- Ok(builder)
+ Ok(request_data)
}
- fn embeddings_builder(
- &self,
- client: &ReqwestClient,
- data: EmbeddingsData,
- ) -> Result<RequestBuilder> {
+ fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
let access_token = get_access_token(self.name())?;
- let body = json!({
- "input": data.texts,
- });
-
let url = format!(
"{API_BASE}/wenxinworkshop/embeddings/{}?access_token={access_token}",
&self.model.name(),
);
- debug!("Ernie Embeddings Request: {url} {body}");
+ let body = json!({
+ "input": data.texts,
+ });
- let builder = client.post(url).json(&body);
+ let request_data = RequestData::new(url, body);
- Ok(builder)
+ Ok(request_data)
}
- fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> {
+ fn prepare_rerank(&self, data: RerankData) -> Result<RequestData> {
let access_token = get_access_token(self.name())?;
+ let url = format!(
+ "{API_BASE}/wenxinworkshop/reranker/{}?access_token={access_token}",
+ &self.model.name(),
+ );
+
let RerankData {
query,
documents,
@@ -89,16 +81,9 @@ impl ErnieClient {
"top_n": top_n
});
- let url = format!(
- "{API_BASE}/wenxinworkshop/reranker/{}?access_token={access_token}",
- &self.model.name(),
- );
-
- debug!("Ernie Rerank Request: {url} {body}");
-
- let builder = client.post(url).json(&body);
+ let request_data = RequestData::new(url, body);
- Ok(builder)
+ Ok(request_data)
}
async fn prepare_access_token(&self) -> Result<()> {
@@ -135,7 +120,8 @@ impl Client for ErnieClient {
data: ChatCompletionsData,
) -> Result<ChatCompletionsOutput> {
self.prepare_access_token().await?;
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions(builder).await
}
@@ -146,7 +132,8 @@ impl Client for ErnieClient {
data: ChatCompletionsData,
) -> Result<()> {
self.prepare_access_token().await?;
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions_streaming(builder, handler).await
}
@@ -156,12 +143,14 @@ impl Client for ErnieClient {
data: EmbeddingsData,
) -> Result<EmbeddingsOutput> {
self.prepare_access_token().await?;
- let builder = self.embeddings_builder(client, data)?;
+ let request_data = self.prepare_embeddings(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Embeddings);
embeddings(builder).await
}
async fn rerank_inner(&self, client: &ReqwestClient, data: RerankData) -> Result<RerankOutput> {
- let builder = self.rerank_builder(client, data)?;
+ let request_data = self.prepare_rerank(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Rerank);
rerank(builder).await
}
}
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 37f7c12..aa1a5b1 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -2,7 +2,7 @@ use super::vertexai::*;
use super::*;
use anyhow::{Context, Result};
-use reqwest::{Client as ReqwestClient, RequestBuilder};
+use reqwest::RequestBuilder;
use serde::Deserialize;
use serde_json::{json, Value};
@@ -14,7 +14,7 @@ pub struct GeminiConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -24,11 +24,7 @@ impl GeminiClient {
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 func = match data.stream {
@@ -36,25 +32,24 @@ impl GeminiClient {
false => "generateContent",
};
- let mut body = gemini_build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
-
let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key);
- debug!("Gemini Chat Completions Request: {url} {body}");
+ let body = gemini_build_chat_completions_body(data, &self.model)?;
- let builder = client.post(url).json(&body);
+ let request_data = RequestData::new(url, body);
- 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 url = format!(
+ "{API_BASE}{}:embedContent?key={}",
+ &self.model.name(),
+ api_key
+ );
+
let body = json!({
"content": {
"parts": [
@@ -65,17 +60,9 @@ impl GeminiClient {
}
});
- let url = format!(
- "{API_BASE}{}:embedContent?key={}",
- &self.model.name(),
- api_key
- );
-
- debug!("Gemini Embeddings Request: {url} {body}");
-
- let builder = client.post(url).json(&body);
+ let request_data = RequestData::new(url, body);
- Ok(builder)
+ Ok(request_data)
}
}
diff --git a/src/client/macros.rs b/src/client/macros.rs
index 778daf4..19e7734 100644
--- a/src/client/macros.rs
+++ b/src/client/macros.rs
@@ -141,7 +141,7 @@ macro_rules! client_common_fns {
self.config.extra.as_ref()
}
- fn patch_config(&self) -> Option<&$crate::client::ModelPatch> {
+ fn patch_config(&self) -> Option<&$crate::client::RequestPatch> {
self.config.patch.as_ref()
}
@@ -171,7 +171,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::ChatCompletionsData,
) -> anyhow::Result<$crate::client::ChatCompletionsOutput> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions(builder).await
}
@@ -181,7 +182,8 @@ macro_rules! impl_client_trait {
handler: &mut $crate::client::SseHandler,
data: $crate::client::ChatCompletionsData,
) -> Result<()> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions_streaming(builder, handler).await
}
}
@@ -196,7 +198,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::ChatCompletionsData,
) -> anyhow::Result<$crate::client::ChatCompletionsOutput> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions(builder).await
}
@@ -206,7 +209,8 @@ macro_rules! impl_client_trait {
handler: &mut $crate::client::SseHandler,
data: $crate::client::ChatCompletionsData,
) -> Result<()> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions_streaming(builder, handler).await
}
@@ -215,7 +219,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::EmbeddingsData,
) -> Result<$crate::client::EmbeddingsOutput> {
- let builder = self.embeddings_builder(client, data)?;
+ let request_data = self.prepare_embeddings(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Embeddings);
$embeddings(builder).await
}
}
@@ -230,7 +235,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::ChatCompletionsData,
) -> anyhow::Result<$crate::client::ChatCompletionsOutput> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions(builder).await
}
@@ -240,7 +246,8 @@ macro_rules! impl_client_trait {
handler: &mut $crate::client::SseHandler,
data: $crate::client::ChatCompletionsData,
) -> Result<()> {
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
$chat_completions_streaming(builder, handler).await
}
@@ -249,7 +256,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::EmbeddingsData,
) -> Result<$crate::client::EmbeddingsOutput> {
- let builder = self.embeddings_builder(client, data)?;
+ let request_data = self.prepare_embeddings(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Embeddings);
$embeddings(builder).await
}
@@ -258,7 +266,8 @@ macro_rules! impl_client_trait {
client: &reqwest::Client,
data: $crate::client::RerankData,
) -> Result<$crate::client::RerankOutput> {
- let builder = self.rerank_builder(client, data)?;
+ let request_data = self.prepare_rerank(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Rerank);
$rerank(builder).await
}
}
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 299cae7..5c26b99 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.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};
@@ -12,7 +12,7 @@ pub struct OllamaConfig {
pub api_auth: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -32,52 +32,41 @@ impl OllamaClient {
),
];
- fn chat_completions_builder(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<RequestBuilder> {
+ fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
let api_base = self.get_api_base()?;
let api_auth = self.get_api_auth().ok();
- let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
-
let url = format!("{api_base}/api/chat");
- debug!("Ollama Chat Completions Request: {url} {body}");
+ let body = 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_auth) = api_auth {
- builder = builder.header("Authorization", api_auth)
+ request_data.header("Authorization", api_auth)
}
- 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_base = self.get_api_base()?;
let api_auth = self.get_api_auth().ok();
+ let url = format!("{api_base}/api/embed");
+
let body = json!({
"model": self.model.name(),
"input": data.texts,
});
- let url = format!("{api_base}/api/embed");
-
- debug!("Ollama Embeddings Request: {url} {body}");
+ let mut request_data = RequestData::new(url, body);
- let mut builder = client.post(url).json(&body);
if let Some(api_auth) = api_auth {
- builder = builder.header("Authorization", api_auth)
+ request_data.header("Authorization", api_auth)
}
- Ok(builder)
+ Ok(request_data)
}
}
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)
}
}
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> {
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index e3aea31..4aa67ae 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -27,7 +27,7 @@ pub struct QianwenConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -37,11 +37,7 @@ impl QianwenClient {
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 stream = data.stream;
@@ -50,27 +46,24 @@ impl QianwenClient {
true => CHAT_COMPLETIONS_API_URL_VL,
false => CHAT_COMPLETIONS_API_URL,
};
- let (mut body, has_upload) = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
- debug!("Qianwen Chat Completions Request: {url} {body}");
+ let (body, has_upload) = build_chat_completions_body(data, &self.model)?;
+
+ let mut request_data = RequestData::new(url, body);
+
+ request_data.bearer_auth(api_key);
- let mut builder = client.post(url).bearer_auth(api_key).json(&body);
if stream {
- builder = builder.header("X-DashScope-SSE", "enable");
+ request_data.header("X-DashScope-SSE", "enable");
}
if has_upload {
- builder = builder.header("X-DashScope-OssResourceResolve", "enable");
+ request_data.header("X-DashScope-OssResourceResolve", "enable");
}
- 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 text_type = match data.query {
@@ -88,13 +81,11 @@ impl QianwenClient {
}
});
- let url = EMBEDDINGS_API_URL;
-
- debug!("Qianwen Embeddings Request: {url} {body}");
+ let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ request_data.bearer_auth(api_key);
- Ok(builder)
+ Ok(request_data)
}
}
@@ -109,7 +100,8 @@ impl Client for QianwenClient {
) -> Result<ChatCompletionsOutput> {
let api_key = self.get_api_key()?;
patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions(builder, &self.model).await
}
@@ -121,7 +113,8 @@ impl Client for QianwenClient {
) -> Result<()> {
let api_key = self.get_api_key()?;
patch_messages(self.model.name(), &api_key, &mut data.messages).await?;
- let builder = self.chat_completions_builder(client, data)?;
+ let request_data = self.prepare_chat_completions(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions_streaming(builder, handler, &self.model).await
}
@@ -130,7 +123,8 @@ impl Client for QianwenClient {
client: &ReqwestClient,
data: EmbeddingsData,
) -> Result<Vec<Vec<f32>>> {
- let builder = self.embeddings_builder(client, data)?;
+ let request_data = self.prepare_embeddings(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Embeddings);
embeddings(builder).await
}
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -26,52 +26,42 @@ impl RagDedicatedClient {
pub const PROMPTS: [PromptAction<'static>; 0] = [];
- fn chat_completions_builder(
- &self,
- _client: &ReqwestClient,
- _data: ChatCompletionsData,
- ) -> Result<RequestBuilder> {
+ fn prepare_chat_completions(&self, _data: ChatCompletionsData) -> Result<RequestData> {
bail!("The client doesn't support chat-completions api");
}
- 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!("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<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!("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<String> {
diff --git a/src/client/replicate.rs b/src/client/replicate.rs
index 53097e6..d6ca401 100644
--- a/src/client/replicate.rs
+++ b/src/client/replicate.rs
@@ -16,7 +16,7 @@ pub struct ReplicateConfig {
pub api_key: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -26,22 +26,20 @@ impl ReplicateClient {
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
- fn chat_completions_builder(
+ fn prepare_chat_completions(
&self,
- client: &ReqwestClient,
data: ChatCompletionsData,
api_key: &str,
- ) -> Result<RequestBuilder> {
- let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
-
+ ) -> Result<RequestData> {
let url = format!("{API_BASE}/models/{}/predictions", self.model.name());
- debug!("Replicate Request: {url} {body}");
+ let body = build_chat_completions_body(data, &self.model)?;
+
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ request_data.bearer_auth(api_key);
- Ok(builder)
+ Ok(request_data)
}
}
@@ -55,7 +53,8 @@ impl Client for ReplicateClient {
data: ChatCompletionsData,
) -> Result<ChatCompletionsOutput> {
let api_key = self.get_api_key()?;
- let builder = self.chat_completions_builder(client, data, &api_key)?;
+ let request_data = self.prepare_chat_completions(data, &api_key)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions(client, builder, &api_key).await
}
@@ -66,7 +65,8 @@ impl Client for ReplicateClient {
data: ChatCompletionsData,
) -> Result<()> {
let api_key = self.get_api_key()?;
- let builder = self.chat_completions_builder(client, data, &api_key)?;
+ let request_data = self.prepare_chat_completions(data, &api_key)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
chat_completions_streaming(client, builder, handler).await
}
}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 4fc2ee1..4ad07b8 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -19,7 +19,7 @@ pub struct VertexAIConfig {
pub adc_file: Option<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patch: Option<ModelPatch>,
+ pub patch: Option<RequestPatch>,
pub extra: Option<ExtraConfig>,
}
@@ -32,12 +32,11 @@ impl VertexAIClient {
("location", "Location", true, PromptKind::String),
];
- fn chat_completions_builder(
+ fn prepare_chat_completions(
&self,
- client: &ReqwestClient,
data: ChatCompletionsData,
model_category: &ModelCategory,
- ) -> Result<RequestBuilder> {
+ ) -> Result<RequestData> {
let project_id = self.get_project_id()?;
let location = self.get_location()?;
let access_token = get_access_token(self.name())?;
@@ -66,7 +65,7 @@ impl VertexAIClient {
}
};
- let mut body = match model_category {
+ let body = match model_category {
ModelCategory::Gemini => gemini_build_chat_completions_body(data, &self.model)?,
ModelCategory::Claude => {
let mut body = claude_build_chat_completions_body(data, &self.model)?;
@@ -84,20 +83,15 @@ impl VertexAIClient {
body
}
};
- self.patch_chat_completions_body(&mut body);
- debug!("VertexAI Chat Completions Request: {url} {body}");
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).bearer_auth(access_token).json(&body);
+ request_data.bearer_auth(access_token);
- Ok(builder)
+ Ok(request_data)
}
- fn embeddings_builder(
- &self,
- client: &ReqwestClient,
- data: EmbeddingsData,
- ) -> Result<RequestBuilder> {
+ fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
let project_id = self.get_project_id()?;
let location = self.get_location()?;
let access_token = get_access_token(self.name())?;
@@ -110,15 +104,16 @@ impl VertexAIClient {
.into_iter()
.map(|v| json!({"content": v}))
.collect();
+
let body = json!({
"instances": instances,
});
- debug!("VertexAI Embeddings Request: {url} {body}");
+ let mut request_data = RequestData::new(url, body);
- let builder = client.post(url).bearer_auth(access_token).json(&body);
+ request_data.bearer_auth(access_token);
- Ok(builder)
+ Ok(request_data)
}
}
@@ -133,7 +128,8 @@ impl Client for VertexAIClient {
) -> Result<ChatCompletionsOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let model_category = ModelCategory::from_str(self.model.name())?;
- let builder = self.chat_completions_builder(client, data, &model_category)?;
+ let request_data = self.prepare_chat_completions(data, &model_category)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
match model_category {
ModelCategory::Gemini => gemini_chat_completions(builder).await,
ModelCategory::Claude => claude_chat_completions(builder).await,
@@ -149,7 +145,8 @@ impl Client for VertexAIClient {
) -> Result<()> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let model_category = ModelCategory::from_str(self.model.name())?;
- let builder = self.chat_completions_builder(client, data, &model_category)?;
+ let request_data = self.prepare_chat_completions(data, &model_category)?;
+ let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
match model_category {
ModelCategory::Gemini => gemini_chat_completions_streaming(builder, handler).await,
ModelCategory::Claude => claude_chat_completions_streaming(builder, handler).await,
@@ -163,7 +160,8 @@ impl Client for VertexAIClient {
data: EmbeddingsData,
) -> Result<Vec<Vec<f32>>> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.embeddings_builder(client, data)?;
+ let request_data = self.prepare_embeddings(data)?;
+ let builder = self.request_builder(client, request_data, ApiType::Embeddings);
embeddings(builder).await
}
}