From 1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 5 Jun 2024 09:02:23 +0800 Subject: feat: support RAG (#560) * feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb) --- src/client/vertexai.rs | 79 +++++++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 71 insertions(+), 8 deletions(-) (limited to 'src/client/vertexai.rs') diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index b40247d..a9f84f8 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,8 +1,5 @@ -use super::{ - access_token::*, catch_error, json_stream, message::*, patch_system_message, - ChatCompletionsData, ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData, - ModelPatches, PromptAction, PromptKind, SseHandler, ToolCall, VertexAIClient, -}; +use super::*; +use super::access_token::*; use anyhow::{anyhow, bail, Context, Result}; use async_trait::async_trait; @@ -51,9 +48,37 @@ impl VertexAIClient { let url = format!("{base_url}/google/models/{}:{func}", self.model.name()); let mut body = gemini_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - debug!("VertexAI Request: {url} {body}"); + debug!("VertexAI Chat Completions Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(access_token).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> 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}/google/models/{}:predict", self.model.name()); + + let task_type = match data.query { + true => "RETRIEVAL_DOCUMENT", + false => "QUESTION_ANSWERING", + }; + let instances: Vec<_> = data.texts.into_iter().map(|v| json!({"task_type": task_type, "content": v})).collect(); + let body = json!({ + "instances": instances, + }); + + debug!("VertexAI Embeddings Request: {url} {body}"); let builder = client.post(url).bearer_auth(access_token).json(&body); @@ -85,6 +110,16 @@ impl Client for VertexAIClient { let builder = self.chat_completions_builder(client, data)?; gemini_chat_completions_streaming(builder, handler).await } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?; + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } } pub async fn gemini_chat_completions(builder: RequestBuilder) -> Result { @@ -138,6 +173,34 @@ pub async fn gemini_chat_completions_streaming( Ok(()) } +async fn embeddings(builder: RequestBuilder) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data: Value = res.json().await?; + if !status.is_success() { + catch_error(&data, status.as_u16())?; + } + let res_body: EmbeddingsResBody = + serde_json::from_value(data).context("Invalid request data")?; + let output = res_body.predictions.into_iter().map(|v| v.embeddings.values).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + predictions: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPrediction { + embeddings: EmbeddingsResBodyPredictionEmbeddings, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyPredictionEmbeddings { + values: Vec +} + fn gemini_extract_chat_completions_text(data: &Value) -> Result { let text = data["candidates"][0]["content"]["parts"][0]["text"] .as_str() @@ -179,7 +242,7 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result Result { -- cgit v1.2.3