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/gemini.rs | 73 ++++++++++++++++++++++++++++++++++++++++++++-------- 1 file changed, 62 insertions(+), 11 deletions(-) (limited to 'src/client/gemini.rs') diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 5cc45c5..03eef7a 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -1,11 +1,10 @@ -use super::{ - vertexai::*, ChatCompletionsData, Client, ExtraConfig, GeminiClient, Model, ModelData, - ModelPatches, PromptAction, PromptKind, -}; +use super::vertexai::*; +use super::*; -use anyhow::Result; +use anyhow::{Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; +use serde_json::{json, Value}; const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/"; @@ -38,13 +37,41 @@ impl GeminiClient { }; let mut body = gemini_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let model = &self.model.name(); + let url = format!("{API_BASE}{}:{}?key={}", &self.model.name(), func, api_key); - let url = format!("{API_BASE}{}:{}?key={}", model, func, api_key); + debug!("Gemini Chat Completions Request: {url} {body}"); - debug!("Gemini Request: {url} {body}"); + let builder = client.post(url).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let body = json!({ + "content": { + "parts": [ + { + "text": data.texts[0], + } + ] + } + }); + + let url = format!( + "{API_BASE}{}:embedContent?key={}", + &self.model.name(), + api_key + ); + + debug!("Gemini Embeddings Request: {url} {body}"); let builder = client.post(url).json(&body); @@ -54,6 +81,30 @@ impl GeminiClient { impl_client_trait!( GeminiClient, - crate::client::vertexai::gemini_chat_completions, - crate::client::vertexai::gemini_chat_completions_streaming + gemini_chat_completions, + gemini_chat_completions_streaming, + gemini_embeddings ); + +async fn gemini_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 = vec![res_body.embedding.values]; + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embedding: EmbeddingsResBodyEmbedding, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbedding { + values: Vec, +} -- cgit v1.2.3