From 5e4210980d6ed92c3850042c7b57ca7eef028be0 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 15 Feb 2024 08:10:30 +0800 Subject: feat: support vertexai (#308) --- config.example.yaml | 9 +++-- src/client/gemini.rs | 6 ++-- src/client/mod.rs | 1 + src/client/vertexai.rs | 92 ++++++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 103 insertions(+), 5 deletions(-) create mode 100644 src/client/vertexai.rs diff --git a/config.example.yaml b/config.example.yaml index 08fb529..77ee50b 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -56,7 +56,7 @@ clients: # See https://learn.microsoft.com/en-us/azure/ai-services/openai/chatgpt-quickstart - type: azure-openai - api_base: https://RESOURCE.openai.azure.com + api_base: https://{RESOURCE}.openai.azure.com api_key: xxx models: - name: MyGPT4 # Model deployment name @@ -69,4 +69,9 @@ clients: # See https://help.aliyun.com/zh/dashscope/ - type: qianwen - api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx \ No newline at end of file + api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + + # See https://cloud.google.com/vertex-ai + - type: vertexai + api_base: https://{REGION}-aiplatform.googleapis.com/v1/projects/{PROJECT_ID}/locations/{REGION}/publishers/google/models + api_key: xxx \ No newline at end of file diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 6f3141c..e846ab1 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -89,7 +89,7 @@ impl GeminiClient { } } -async fn send_message(builder: RequestBuilder) -> Result { +pub(crate) async fn send_message(builder: RequestBuilder) -> Result { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -102,7 +102,7 @@ async fn send_message(builder: RequestBuilder) -> Result { Ok(output.to_string()) } -async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { +pub(crate) async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { let res = builder.send().await?; if res.status() != 200 { let data: Value = res.json().await?; @@ -178,7 +178,7 @@ fn check_error(data: &Value) -> Result<()> { } } -fn build_body(data: SendData, _model: String) -> Result { +pub(crate) fn build_body(data: SendData, _model: String) -> Result { let SendData { mut messages, temperature, diff --git a/src/client/mod.rs b/src/client/mod.rs index 7c157ef..320ddc4 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -20,4 +20,5 @@ register_client!( ), (ernie, "ernie", ErnieConfig, ErnieClient), (qianwen, "qianwen", QianwenConfig, QianwenClient), + (vertexai, "vertexai", VertexAIConfig, VertexAIClient), ); diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs new file mode 100644 index 0000000..890aedd --- /dev/null +++ b/src/client/vertexai.rs @@ -0,0 +1,92 @@ +use super::{ + Client, ExtraConfig, VertexAIClient, Model, PromptType, + SendData, TokensCountFactors, +}; +use super::gemini::{build_body, send_message, send_message_streaming}; + +use crate::{render::ReplyHandler, utils::PromptKind}; + +use anyhow::Result; +use async_trait::async_trait; +use reqwest::{Client as ReqwestClient, RequestBuilder}; +use serde::Deserialize; + +const MODELS: [(&str, usize, &str); 2] = [ + ("gemini-pro", 32760, "text"), + ("gemini-pro-vision", 16384, "text,vision"), +]; + +const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); + +#[derive(Debug, Clone, Deserialize, Default)] +pub struct VertexAIConfig { + pub name: Option, + pub api_base: Option, + pub api_key: Option, + pub extra: Option, +} + +#[async_trait] +impl Client for VertexAIClient { + client_common_fns!(); + + async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result { + let builder = self.request_builder(client, data)?; + send_message(builder).await + } + + async fn send_message_streaming_inner( + &self, + client: &ReqwestClient, + handler: &mut ReplyHandler, + data: SendData, + ) -> Result<()> { + let builder = self.request_builder(client, data)?; + send_message_streaming(builder, handler).await + } +} + +impl VertexAIClient { + config_get_fn!(api_base, get_api_base); + config_get_fn!(api_key, get_api_key); + + pub const PROMPTS: [PromptType<'static>; 2] = [ + ("api_base", "API Base:", true, PromptKind::String), + ("api_key", "API Key:", true, PromptKind::String), + ]; + + pub fn list_models(local_config: &VertexAIConfig) -> Vec { + let client_name = Self::name(local_config); + MODELS + .into_iter() + .map(|(name, max_tokens, capabilities)| { + Model::new(client_name, name) + .set_capabilities(capabilities.into()) + .set_max_tokens(Some(max_tokens)) + .set_tokens_count_factors(TOKENS_COUNT_FACTORS) + }) + .collect() + } + + fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { + let api_base = self.get_api_base()?; + let api_key = self.get_api_key()?; + + let func = match data.stream { + true => "streamGenerateContent", + false => "generateContent", + }; + + let body = build_body(data, self.model.name.clone())?; + + let model = self.model.name.clone(); + + let url = format!("{api_base}/{}:{}", model, func); + + debug!("VertexAI Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } +} -- cgit v1.2.3