From 9b283024b47f57ad7bbf032fa015eab42f8162a9 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 6 May 2024 08:19:42 +0800 Subject: feat: extract vertexai-claude client (#485) --- src/client/vertexai_claude.rs | 100 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 100 insertions(+) create mode 100644 src/client/vertexai_claude.rs (limited to 'src/client/vertexai_claude.rs') diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs new file mode 100644 index 0000000..78adb72 --- /dev/null +++ b/src/client/vertexai_claude.rs @@ -0,0 +1,100 @@ +use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming}; +use super::vertexai::fetch_gcloud_access_token; +use super::{ + Client, CompletionDetails, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, + SseHandler, VertexAIClaudeClient, +}; + +use anyhow::{anyhow, Context, Result}; +use async_trait::async_trait; +use chrono::{Duration, Utc}; +use reqwest::{Client as ReqwestClient, RequestBuilder}; +use serde::Deserialize; + +static mut ACCESS_TOKEN: (String, i64) = (String::new(), 0); // safe under linear operation + +#[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 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 request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { + let project_id = self.get_project_id()?; + let location = self.get_location()?; + + 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_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(); + + debug!("VertexAIClaude Request: {url} {body}"); + + let builder = client + .post(url) + .bearer_auth(unsafe { &ACCESS_TOKEN.0 }) + .json(&body); + + Ok(builder) + } +} + +#[async_trait] +impl Client for VertexAIClaudeClient { + client_common_fns!(); + + async fn send_message_inner( + &self, + client: &ReqwestClient, + data: SendData, + ) -> Result<(String, CompletionDetails)> { + prepare_access_token(client, &self.config.adc_file).await?; + let builder = self.request_builder(client, data)?; + claude_send_message(builder).await + } + + async fn send_message_streaming_inner( + &self, + client: &ReqwestClient, + handler: &mut SseHandler, + data: SendData, + ) -> Result<()> { + prepare_access_token(client, &self.config.adc_file).await?; + let builder = self.request_builder(client, data)?; + claude_send_message_streaming(builder, handler).await + } +} + +async fn prepare_access_token(client: &reqwest::Client, adc_file: &Option) -> Result<()> { + if unsafe { ACCESS_TOKEN.0.is_empty() || Utc::now().timestamp() > ACCESS_TOKEN.1 } { + let (token, expires_in) = fetch_gcloud_access_token(client, adc_file) + .await + .with_context(|| "Failed to fetch access token")?; + let expires_at = Utc::now() + + Duration::try_seconds(expires_in) + .ok_or_else(|| anyhow!("Failed to parse expires_in of access_token"))?; + unsafe { ACCESS_TOKEN = (token, expires_at.timestamp()) }; + } + Ok(()) +} -- cgit v1.2.3