summaryrefslogtreecommitdiffstats
path: root/src/client/gemini.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/client/gemini.rs
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
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)
Diffstat (limited to 'src/client/gemini.rs')
-rw-r--r--src/client/gemini.rs73
1 files changed, 62 insertions, 11 deletions
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<RequestBuilder> {
+ 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<EmbeddingsOutput> {
+ 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<f32>,
+}