diff options
| -rw-r--r-- | models.yaml | 36 | ||||
| -rw-r--r-- | src/client/mod.rs | 6 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 69 | ||||
| -rw-r--r-- | src/client/vertexai_claude.rs | 86 |
4 files changed, 74 insertions, 123 deletions
diff --git a/models.yaml b/models.yaml index 0df9944..cdced59 100644 --- a/models.yaml +++ b/models.yaml @@ -277,6 +277,7 @@ - platform: vertexai # docs: # - https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models + # - https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/use-claude # - https://cloud.google.com/vertex-ai/generative-ai/pricing # - https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/gemini # notes: @@ -302,27 +303,6 @@ input_price: 0.125 output_price: 0.375 supports_function_calling: true - - name: text-embedding-004 - type: embedding - max_input_tokens: 3072 - input_price: 0.025 - output_vector_size: 768 - default_chunk_size: 1500 - max_batch_size: 5 - - name: text-multilingual-embedding-002 - type: embedding - max_input_tokens: 3072 - input_price: 0.2 - output_vector_size: 768 - default_chunk_size: 1500 - max_batch_size: 5 - -- platform: vertexai-claude - # docs: - # - https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/use-claude - # notes: - # - get max_output_tokens info from models doc - models: - name: claude-3-5-sonnet@20240620 max_input_tokens: 200000 max_output_tokens: 4096 @@ -355,6 +335,20 @@ output_price: 1.25 supports_vision: true supports_function_calling: true + - name: text-embedding-004 + type: embedding + max_input_tokens: 3072 + input_price: 0.025 + output_vector_size: 768 + default_chunk_size: 1500 + max_batch_size: 5 + - name: text-multilingual-embedding-002 + type: embedding + max_input_tokens: 3072 + input_price: 0.2 + output_vector_size: 768 + default_chunk_size: 1500 + max_batch_size: 5 - platform: bedrock # docs: 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 - } -} |
