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/azure_openai.rs | 36 ++++----- src/client/bedrock.rs | 84 +++++++++++++------- src/client/claude.rs | 24 ++---- src/client/cloudflare.rs | 41 ++++------ src/client/cohere.rs | 45 ++++------- src/client/common.rs | 164 +++++++++++++++++++++++++++++++--------- src/client/ernie.rs | 65 +++++++--------- src/client/gemini.rs | 43 ++++------- src/client/macros.rs | 29 ++++--- src/client/ollama.rs | 39 ++++------ src/client/openai.rs | 41 +++++----- src/client/openai_compatible.rs | 51 +++++-------- src/client/qianwen.rs | 46 +++++------ src/client/rag_dedicated.rs | 40 ++++------ src/client/replicate.rs | 24 +++--- src/client/vertexai.rs | 38 +++++----- 16 files changed, 415 insertions(+), 395 deletions(-) (limited to 'src') 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -32,51 +31,42 @@ impl AzureOpenAIClient { ), ]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + 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 { + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } +} async fn chat_completions(builder: RequestBuilder) -> Result { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -27,51 +27,42 @@ impl CloudflareClient { ("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 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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 { + fn prepare_rerank(&self, data: RerankData) -> Result { 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, + pub embeddings: Option, + pub rerank: Option, } -pub type ChatCompletionsBodyPatch = IndexMap; - -pub fn extract_chat_completions_body_patch( - patch: Option, - model: &Model, -) -> Option { - 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; + +#[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, + pub body: Value, +} + +impl RequestData { + pub fn new(url: T, body: Value) -> Self + where + T: std::fmt::Display, + { + Self { + url: url.to_string(), + headers: Default::default(), + body, + } + } + + pub fn bearer_auth(&mut self, auth: T) + where + T: std::fmt::Display, + { + self.headers + .insert("authorization".into(), format!("Bearer {auth}")); + } + + pub fn header(&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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -29,54 +29,46 @@ impl ErnieClient { ("secret_key", "Secret Key:", true, PromptKind::String), ]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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 { + fn prepare_rerank(&self, data: RerankData) -> Result { 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 { 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 { 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 { - 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -32,52 +32,41 @@ impl OllamaClient { ), ]; - fn chat_completions_builder( - &self, - client: &ReqwestClient, - data: ChatCompletionsData, - ) -> Result { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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, #[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) } } 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -35,17 +34,10 @@ impl OpenAICompatibleClient { ), ]; - 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().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 { + 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!("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 { + 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!("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 { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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 { 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>> { - 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, #[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 { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { - let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_chat_completions_body(&mut body); - + ) -> Result { 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 { 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, #[serde(default)] pub models: Vec, - pub patch: Option, + pub patch: Option, pub extra: Option, } @@ -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 { + ) -> Result { 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 { + fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { 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 { 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>> { 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 } } -- cgit v1.2.3