From 96ad64352dccfb9328374ef754c2a90a31c0a339 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 25 Jul 2024 15:56:56 -0700 Subject: feat: merge vertexai-cluade with vertexai (#745) --- src/client/mod.rs | 6 --- src/client/vertexai.rs | 69 +++++++++++++++++++++++++++++----- src/client/vertexai_claude.rs | 86 ------------------------------------------- 3 files changed, 59 insertions(+), 102 deletions(-) delete mode 100644 src/client/vertexai_claude.rs (limited to 'src') 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 { 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 { 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 { + 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, - pub project_id: Option, - pub location: Option, - pub adc_file: Option, - #[serde(default)] - pub models: Vec, - pub patches: Option, - pub extra: Option, -} - -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 { - 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 { - 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 - } -} -- cgit v1.2.3