summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/mod.rs6
-rw-r--r--src/client/vertexai.rs69
-rw-r--r--src/client/vertexai_claude.rs86
3 files changed, 59 insertions, 102 deletions
diff --git a/src/client/mod.rs b/src/client/mod.rs
index b4c0091..dd1d0b1 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -38,12 +38,6 @@ register_client!(
AzureOpenAIClient
),
(vertexai, "vertexai", VertexAIConfig, VertexAIClient),
- (
- vertexai_claude,
- "vertexai-claude",
- VertexAIClaudeConfig,
- VertexAIClaudeClient
- ),
(bedrock, "bedrock", BedrockConfig, BedrockClient),
(cloudflare, "cloudflare", CloudflareConfig, CloudflareClient),
(replicate, "replicate", ReplicateConfig, ReplicateClient),
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 828c300..c9ae648 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,4 +1,5 @@
use super::access_token::*;
+use super::claude::*;
use super::*;
use anyhow::{anyhow, bail, Context, Result};
@@ -7,7 +8,7 @@ use chrono::{Duration, Utc};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
-use std::path::PathBuf;
+use std::{path::PathBuf, str::FromStr};
#[derive(Debug, Clone, Deserialize, Default)]
pub struct VertexAIConfig {
@@ -34,6 +35,7 @@ impl VertexAIClient {
&self,
client: &ReqwestClient,
data: ChatCompletionsData,
+ model_category: &ModelCategory,
) -> Result<RequestBuilder> {
let project_id = self.get_project_id()?;
let location = self.get_location()?;
@@ -41,13 +43,32 @@ impl VertexAIClient {
let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
- let func = match data.stream {
- true => "streamGenerateContent",
- false => "generateContent",
+ let model_name = self.model.name();
+
+ let url = match model_category {
+ ModelCategory::Gemini => {
+ let func = match data.stream {
+ true => "streamGenerateContent",
+ false => "generateContent",
+ };
+ format!("{base_url}/google/models/{model_name}:{func}")
+ }
+ ModelCategory::Claude => {
+ format!("{base_url}/anthropic/models/{model_name}:streamRawPredict")
+ }
};
- let url = format!("{base_url}/google/models/{}:{func}", self.model.name());
- let mut body = gemini_build_chat_completions_body(data, &self.model)?;
+ let mut 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)?;
+ if let Some(body_obj) = body.as_object_mut() {
+ body_obj.remove("model");
+ }
+ body["anthropic_version"] = "vertex-2023-10-16".into();
+ body
+ }
+ };
self.patch_chat_completions_body(&mut body);
debug!("VertexAI Chat Completions Request: {url} {body}");
@@ -96,8 +117,12 @@ impl Client for VertexAIClient {
data: ChatCompletionsData,
) -> Result<ChatCompletionsOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.chat_completions_builder(client, data)?;
- gemini_chat_completions(builder).await
+ let model_category = ModelCategory::from_str(self.model.name())?;
+ let builder = self.chat_completions_builder(client, data, &model_category)?;
+ match model_category {
+ ModelCategory::Gemini => gemini_chat_completions(builder).await,
+ ModelCategory::Claude => claude_chat_completions(builder).await,
+ }
}
async fn chat_completions_streaming_inner(
@@ -107,8 +132,12 @@ impl Client for VertexAIClient {
data: ChatCompletionsData,
) -> Result<()> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.chat_completions_builder(client, data)?;
- gemini_chat_completions_streaming(builder, handler).await
+ let model_category = ModelCategory::from_str(self.model.name())?;
+ let builder = self.chat_completions_builder(client, data, &model_category)?;
+ match model_category {
+ ModelCategory::Gemini => gemini_chat_completions_streaming(builder, handler).await,
+ ModelCategory::Claude => claude_chat_completions_streaming(builder, handler).await,
+ }
}
async fn embeddings_inner(
@@ -366,6 +395,26 @@ pub fn gemini_build_chat_completions_body(
Ok(body)
}
+#[derive(Debug, Clone, Copy, PartialEq, Eq)]
+enum ModelCategory {
+ Gemini,
+ Claude,
+}
+
+impl FromStr for ModelCategory {
+ type Err = anyhow::Error;
+
+ fn from_str(s: &str) -> std::result::Result<Self, Self::Err> {
+ if s.starts_with("gemini-") {
+ Ok(ModelCategory::Gemini)
+ } else if s.starts_with("claude-") {
+ Ok(ModelCategory::Claude)
+ } else {
+ unsupported_model!(s)
+ }
+ }
+}
+
pub async fn prepare_gcloud_access_token(
client: &reqwest::Client,
client_name: &str,
diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs
deleted file mode 100644
index 946fb6f..0000000
--- a/src/client/vertexai_claude.rs
+++ /dev/null
@@ -1,86 +0,0 @@
-use super::access_token::*;
-use super::claude::*;
-use super::vertexai::*;
-use super::*;
-
-use anyhow::Result;
-use async_trait::async_trait;
-use reqwest::{Client as ReqwestClient, RequestBuilder};
-use serde::Deserialize;
-
-#[derive(Debug, Clone, Deserialize, Default)]
-pub struct VertexAIClaudeConfig {
- pub name: Option<String>,
- pub project_id: Option<String>,
- pub location: Option<String>,
- pub adc_file: Option<String>,
- #[serde(default)]
- pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
- pub extra: Option<ExtraConfig>,
-}
-
-impl VertexAIClaudeClient {
- config_get_fn!(project_id, get_project_id);
- config_get_fn!(location, get_location);
-
- pub const PROMPTS: [PromptAction<'static>; 2] = [
- ("project_id", "Project ID", true, PromptKind::String),
- ("location", "Location", true, PromptKind::String),
- ];
-
- fn chat_completions_builder(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<RequestBuilder> {
- let project_id = self.get_project_id()?;
- let location = self.get_location()?;
- let access_token = get_access_token(self.name())?;
-
- let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
- let url = format!(
- "{base_url}/anthropic/models/{}:streamRawPredict",
- self.model.name()
- );
-
- let mut body = claude_build_chat_completions_body(data, &self.model)?;
- self.patch_chat_completions_body(&mut body);
- if let Some(body_obj) = body.as_object_mut() {
- body_obj.remove("model");
- }
- body["anthropic_version"] = "vertex-2023-10-16".into();
-
- debug!("VertexAIClaude Request: {url} {body}");
-
- let builder = client.post(url).bearer_auth(access_token).json(&body);
-
- Ok(builder)
- }
-}
-
-#[async_trait]
-impl Client for VertexAIClaudeClient {
- client_common_fns!();
-
- async fn chat_completions_inner(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<ChatCompletionsOutput> {
- prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.chat_completions_builder(client, data)?;
- claude_chat_completions(builder).await
- }
-
- async fn chat_completions_streaming_inner(
- &self,
- client: &ReqwestClient,
- handler: &mut SseHandler,
- data: ChatCompletionsData,
- ) -> Result<()> {
- prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.chat_completions_builder(client, data)?;
- claude_chat_completions_streaming(builder, handler).await
- }
-}