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/azure_openai.rs | 31 ++- src/client/bedrock.rs | 10 +- src/client/claude.rs | 9 +- src/client/cloudflare.rs | 7 +- src/client/cohere.rs | 69 ++++- src/client/common.rs | 140 ++++++---- src/client/ernie.rs | 9 +- src/client/gemini.rs | 73 +++++- src/client/model.rs | 44 +++- src/client/ollama.rs | 67 ++++- src/client/openai.rs | 67 ++++- src/client/openai_compatible.rs | 73 ++++-- src/client/qianwen.rs | 86 +++++- src/client/replicate.rs | 9 +- src/client/vertexai.rs | 79 +++++- src/client/vertexai_claude.rs | 11 +- src/config/input.rs | 59 +++-- src/config/mod.rs | 185 ++++++++++--- src/config/session.rs | 2 +- src/main.rs | 32 ++- src/rag/loader.rs | 146 +++++++++++ src/rag/mod.rs | 425 ++++++++++++++++++++++++++++++ src/rag/splitter.rs | 564 ++++++++++++++++++++++++++++++++++++++++ src/render/stream.rs | 15 +- src/repl/mod.rs | 61 +++-- src/serve.rs | 11 +- src/utils/abort_signal.rs | 9 + src/utils/mod.rs | 2 +- src/utils/spinner.rs | 39 ++- 29 files changed, 2048 insertions(+), 286 deletions(-) create mode 100644 src/rag/loader.rs create mode 100644 src/rag/mod.rs create mode 100644 src/rag/splitter.rs (limited to 'src') diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 52c8a34..19d234a 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,7 +1,5 @@ -use super::{ - openai::*, AzureOpenAIClient, ChatCompletionsData, Client, ExtraConfig, Model, ModelData, - ModelPatches, PromptAction, PromptKind, -}; +use super::*; +use super::openai::*; use anyhow::Result; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -12,6 +10,7 @@ pub struct AzureOpenAIConfig { pub name: Option, pub api_base: Option, pub api_key: Option, + #[serde(default)] pub models: Vec, pub patches: Option, pub extra: Option, @@ -42,7 +41,7 @@ impl AzureOpenAIClient { let api_key = self.get_api_key()?; let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!( "{}/openai/deployments/{}/chat/completions?api-version=2024-02-01", @@ -56,10 +55,28 @@ impl AzureOpenAIClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_base = self.get_api_base()?; + let api_key = self.get_api_key()?; + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } impl_client_trait!( AzureOpenAIClient, - crate::client::openai::openai_chat_completions, - crate::client::openai::openai_chat_completions_streaming + openai_chat_completions, + openai_chat_completions_streaming, + openai_embeddings ); diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 3dfa977..981d1cb 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,8 +1,6 @@ -use super::{ - catch_error, claude::*, prompt_format::*, BedrockClient, ChatCompletionsData, - ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, - PromptKind, SseHandler, -}; +use super::*; +use super::claude::*; +use super::prompt_format::*; use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; @@ -102,7 +100,7 @@ impl BedrockClient { let headers = IndexMap::new(); let mut body = build_chat_completions_body(data, &self.model, model_category)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let builder = aws_fetch( client, diff --git a/src/client/claude.rs b/src/client/claude.rs index 6533a9a..a16722b 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,9 +1,4 @@ -use super::{ - catch_error, extract_system_message, message::*, sse_stream, ChatCompletionsData, - ChatCompletionsOutput, ClaudeClient, Client, ExtraConfig, ImageUrl, MessageContent, - MessageContentPart, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, - SseMmessage, ToolCall, -}; +use super::*; use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -36,7 +31,7 @@ impl ClaudeClient { let api_key = self.get_api_key().ok(); let mut body = claude_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = API_BASE; diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 966cee4..965f20a 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -1,7 +1,4 @@ -use super::{ - catch_error, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, CloudflareClient, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, SseMmessage, -}; +use super::*; use anyhow::{anyhow, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -39,7 +36,7 @@ impl CloudflareClient { let api_key = self.get_api_key()?; let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!( "{API_BASE}/accounts/{account_id}/ai/run/{}", diff --git a/src/client/cohere.rs b/src/client/cohere.rs index e0a5eec..69c343b 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,15 +1,12 @@ -use super::{ - catch_error, extract_system_message, json_stream, message::*, ChatCompletionsData, - ChatCompletionsOutput, Client, CohereClient, ExtraConfig, Model, ModelData, ModelPatches, - PromptAction, PromptKind, SseHandler, ToolCall, -}; +use super::*; -use anyhow::{bail, Result}; +use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; -const API_URL: &str = "https://api.cohere.ai/v1/chat"; +const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat"; +const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed"; #[derive(Debug, Clone, Deserialize, Default)] pub struct CohereConfig { @@ -35,11 +32,38 @@ impl CohereClient { let api_key = self.get_api_key()?; let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let url = API_URL; + let url = CHAT_COMPLETIONS_API_URL; - debug!("Cohere Request: {url} {body}"); + debug!("Cohere Chat Completions Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let input_type = match data.query { + true => "search_query", + false => "search_document", + }; + + let body = json!({ + "model": self.model.name(), + "texts": data.texts, + "input_type": input_type, + }); + + let url = EMBEDDINGS_API_URL; + + debug!("Cohere Embeddings Request: {url} {body}"); let builder = client.post(url).bearer_auth(api_key).json(&body); @@ -47,7 +71,12 @@ impl CohereClient { } } -impl_client_trait!(CohereClient, chat_completions, chat_completions_streaming); +impl_client_trait!( + CohereClient, + chat_completions, + chat_completions_streaming, + embeddings +); async fn chat_completions(builder: RequestBuilder) -> Result { let res = builder.send().await?; @@ -100,6 +129,24 @@ async fn 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")?; + Ok(res_body.embeddings) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embeddings: Vec>, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result { let ChatCompletionsData { mut messages, diff --git a/src/client/common.rs b/src/client/common.rs index 96ec90b..1055b84 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,10 +1,13 @@ -use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, SseHandler}; +use super::*; use crate::{ config::{GlobalConfig, Input}, function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolCallResult}, render::{render_error, render_stream}, - utils::{prompt_input_integer, prompt_input_string, tokenize, AbortSignal, PromptKind}, + utils::{ + prompt_input_integer, prompt_input_string, tokenize, watch_abort_signal, AbortSignal, + PromptKind, + }, }; use anyhow::{bail, Context, Result}; @@ -16,13 +19,12 @@ use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; use std::{env, future::Future, time::Duration}; -use tokio::{sync::mpsc::unbounded_channel, time::sleep}; +use tokio::sync::mpsc::unbounded_channel; const MODELS_YAML: &str = include_str!("../../models.yaml"); lazy_static! { - pub static ref ALL_CLIENT_MODELS: Vec = - serde_yaml::from_str(MODELS_YAML).unwrap(); + pub static ref ALL_MODELS: Vec = serde_yaml::from_str(MODELS_YAML).unwrap(); static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(? Vec { let client_name = Self::name(local_config); if local_config.models.is_empty() { - if let Some(client_models) = $crate::client::ALL_CLIENT_MODELS.iter().find(|v| { + if let Some(models) = $crate::client::ALL_MODELS.iter().find(|v| { v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform)) }) { - return Model::from_config(client_name, &client_models.models); + return Model::from_config(client_name, &models.models); } vec![] } else { @@ -137,10 +139,10 @@ macro_rules! register_client { anyhow::bail!("Unknown client '{}'", client) } - static mut ALL_CLIENTS: Option> = None; + static mut ALL_CLIENT_MODELS: Option> = None; - pub fn list_models(config: &$crate::config::Config) -> Vec<&$crate::client::Model> { - if unsafe { ALL_CLIENTS.is_none() } { + pub fn list_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + if unsafe { ALL_CLIENT_MODELS.is_none() } { let models: Vec<_> = config .clients .iter() @@ -149,9 +151,17 @@ macro_rules! register_client { ClientConfig::Unknown => vec![], }) .collect(); - unsafe { ALL_CLIENTS = Some(models) }; + unsafe { ALL_CLIENT_MODELS = Some(models) }; } - unsafe { ALL_CLIENTS.as_ref().unwrap().iter().collect() } + unsafe { ALL_CLIENT_MODELS.as_ref().unwrap().iter().collect() } + } + + pub fn list_chat_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + list_models(config).into_iter().filter(|v| v.mode() == "chat").collect() + } + + pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> { + list_models(config).into_iter().filter(|v| v.mode() == "embedding").collect() } }; } @@ -171,10 +181,6 @@ macro_rules! client_common_fns { self.config.patches.as_ref() } - fn list_models(&self) -> Vec { - Self::list_models(&self.config) - } - fn name(&self) -> &str { Self::name(&self.config) } @@ -186,10 +192,6 @@ macro_rules! client_common_fns { fn model_mut(&mut self) -> &mut Model { &mut self.model } - - fn set_model(&mut self, model: Model) { - self.model = model; - } }; } @@ -220,6 +222,40 @@ macro_rules! impl_client_trait { } } }; + ($client:ident, $chat_completions:path, $chat_completions_streaming:path, $embeddings:path) => { + #[async_trait::async_trait] + impl $crate::client::Client for $crate::client::$client { + client_common_fns!(); + + async fn chat_completions_inner( + &self, + client: &reqwest::Client, + data: $crate::client::ChatCompletionsData, + ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions(builder).await + } + + async fn chat_completions_streaming_inner( + &self, + client: &reqwest::Client, + handler: &mut $crate::client::SseHandler, + data: $crate::client::ChatCompletionsData, + ) -> Result<()> { + let builder = self.chat_completions_builder(client, data)?; + $chat_completions_streaming(builder, handler).await + } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + let builder = self.embeddings_builder(client, data)?; + $embeddings(builder).await + } + } + }; } #[macro_export] @@ -256,19 +292,12 @@ pub trait Client: Sync + Send { fn patches_config(&self) -> Option<&ModelPatches>; - #[allow(unused)] fn name(&self) -> &str; - #[allow(unused)] - fn list_models(&self) -> Vec; - fn model(&self) -> &Model; fn model_mut(&mut self) -> &mut Model; - #[allow(unused)] - fn set_model(&mut self, model: Model); - fn build_client(&self) -> Result { let mut builder = ReqwestClient::builder(); let extra = self.extra_config(); @@ -288,11 +317,10 @@ pub trait Client: Sync + Send { return Ok(ChatCompletionsOutput::new(&content)); } let client = self.build_client()?; - let data = input.prepare_completion_data(self.model(), false)?; self.chat_completions_inner(&client, data) .await - .with_context(|| "Failed to get answer") + .with_context(|| "Failed to get chat completions") } async fn chat_completions_streaming( @@ -300,15 +328,7 @@ pub trait Client: Sync + Send { input: &Input, handler: &mut SseHandler, ) -> Result<()> { - async fn watch_abort(abort: AbortSignal) { - loop { - if abort.aborted() { - break; - } - sleep(Duration::from_millis(100)).await; - } - } - let abort = handler.get_abort(); + let abort_signal = handler.get_abort(); let input = input.clone(); tokio::select! { ret = async { @@ -326,20 +346,28 @@ pub trait Client: Sync + Send { self.chat_completions_streaming_inner(&client, handler, data).await } => { handler.done()?; - ret.with_context(|| "Failed to get answer") + ret.with_context(|| "Failed to get chat completions") } - _ = watch_abort(abort.clone()) => { + _ = watch_abort_signal(abort_signal) => { handler.done()?; Ok(()) }, } } - fn patch_request_body(&self, body: &mut Value) { + async fn embeddings(&self, data: EmbeddingsData) -> Result>> { + let client = self.build_client()?; + self.model().guard_max_concurrent_chunks(&data)?; + self.embeddings_inner(&client, data) + .await + .with_context(|| "Failed to get embeddings") + } + + fn patch_chat_completions_body(&self, body: &mut Value) { let model_name = self.model().name(); if let Some(patch_data) = select_model_patch(self.patches_config(), model_name) { - if body.is_object() && patch_data.request_body.is_object() { - json_patch::merge(body, &patch_data.request_body) + if body.is_object() && patch_data.chat_completions_body.is_object() { + json_patch::merge(body, &patch_data.chat_completions_body) } } } @@ -356,6 +384,14 @@ pub trait Client: Sync + Send { handler: &mut SseHandler, data: ChatCompletionsData, ) -> Result<()>; + + async fn embeddings_inner( + &self, + _client: &ReqwestClient, + _data: EmbeddingsData, + ) -> Result>> { + bail!("No embeddings api") + } } impl Default for ClientConfig { @@ -375,7 +411,7 @@ pub type ModelPatches = IndexMap; #[derive(Debug, Clone, Deserialize)] pub struct ModelPatch { #[serde(default)] - pub request_body: Value, + pub chat_completions_body: Value, } pub fn select_model_patch<'a>( @@ -421,6 +457,20 @@ impl ChatCompletionsOutput { } } +#[derive(Debug)] +pub struct EmbeddingsData { + pub texts: Vec, + pub query: bool, +} + +impl EmbeddingsData { + pub fn new(texts: Vec, query: bool) -> Self { + Self { texts, query } + } +} + +pub type EmbeddingsOutput = Vec>; + pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind); pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { @@ -445,7 +495,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result Result { let mut body = build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let access_token = get_access_token(self.name())?; 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, +} diff --git a/src/client/model.rs b/src/client/model.rs index 65e4143..e16cb4e 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,4 +1,7 @@ -use super::message::{Message, MessageContent}; +use super::{ + message::{Message, MessageContent}, + EmbeddingsData, +}; use crate::utils::{estimate_token_length, format_option_value}; @@ -81,6 +84,10 @@ impl Model { &self.data.name } + pub fn mode(&self) -> &str { + &self.data.mode + } + pub fn data(&self) -> &ModelData { &self.data } @@ -137,6 +144,14 @@ impl Model { self.data.supports_function_calling } + pub fn default_chunk_size(&self) -> usize { + self.data.default_chunk_size.unwrap_or(1000) + } + + pub fn max_concurrent_chunks(&self) -> usize { + self.data.max_concurrent_chunks.unwrap_or(1) + } + pub fn max_tokens_param(&self) -> Option { if self.data.pass_max_tokens { self.data.max_output_tokens @@ -182,30 +197,45 @@ impl Model { } } - pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> { + pub fn guard_max_input_tokens(&self, messages: &[Message]) -> Result<()> { let total_tokens = self.total_tokens(messages) + BASIS_TOKENS; if let Some(max_input_tokens) = self.data.max_input_tokens { if total_tokens >= max_input_tokens { - bail!("Exceed max input tokens limit") + bail!("Exceed max_input_tokens limit") } } Ok(()) } + + pub fn guard_max_concurrent_chunks(&self, data: &EmbeddingsData) -> Result<()> { + if data.texts.len() > self.max_concurrent_chunks() { + bail!("Exceed max_concurrent_chunks limit"); + } + Ok(()) + } } #[derive(Debug, Clone, Default, Deserialize)] pub struct ModelData { pub name: String, + #[serde(default = "default_model_mode")] + pub mode: String, pub max_input_tokens: Option, + pub input_price: Option, + pub output_price: Option, + + // chat-only properties pub max_output_tokens: Option, #[serde(default)] pub pass_max_tokens: bool, - pub input_price: Option, - pub output_price: Option, #[serde(default)] pub supports_vision: bool, #[serde(default)] pub supports_function_calling: bool, + + // embedding-only properties + pub default_chunk_size: Option, + pub max_concurrent_chunks: Option, } impl ModelData { @@ -222,3 +252,7 @@ pub struct BuiltinModels { pub platform: String, pub models: Vec, } + +fn default_model_mode() -> String { + "chat".into() +} diff --git a/src/client/ollama.rs b/src/client/ollama.rs index beba8a1..f9bf8d5 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,10 +1,6 @@ -use super::{ - catch_error, json_stream, message::*, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, OllamaClient, PromptAction, PromptKind, - SseHandler, -}; +use super::*; -use anyhow::{anyhow, bail, Result}; +use anyhow::{anyhow, bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -14,7 +10,7 @@ pub struct OllamaConfig { pub name: Option, pub api_base: Option, pub api_auth: Option, - pub chat_endpoint: Option, + #[serde(default)] pub models: Vec, pub patches: Option, pub extra: Option, @@ -45,13 +41,36 @@ impl OllamaClient { let api_auth = self.get_api_auth().ok(); let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - let chat_endpoint = self.config.chat_endpoint.as_deref().unwrap_or("/api/chat"); + let url = format!("{api_base}/api/chat"); - let url = format!("{api_base}{chat_endpoint}"); + debug!("Ollama Chat Completions Request: {url} {body}"); - debug!("Ollama Request: {url} {body}"); + let mut builder = client.post(url).json(&body); + if let Some(api_auth) = api_auth { + builder = builder.header("Authorization", api_auth) + } + + Ok(builder) + } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_base = self.get_api_base()?; + let api_auth = self.get_api_auth().ok(); + + let body = json!({ + "model": self.model.name(), + "prompt": data.texts[0], + }); + + let url = format!("{api_base}/api/embeddings"); + + debug!("Ollama Embeddings Request: {url} {body}"); let mut builder = client.post(url).json(&body); if let Some(api_auth) = api_auth { @@ -62,7 +81,12 @@ impl OllamaClient { } } -impl_client_trait!(OllamaClient, chat_completions, chat_completions_streaming); +impl_client_trait!( + OllamaClient, + chat_completions, + chat_completions_streaming, + embeddings +); async fn chat_completions(builder: RequestBuilder) -> Result { let res = builder.send().await?; @@ -109,6 +133,25 @@ async fn chat_completions_streaming( Ok(()) } +async fn embeddings( + builder: RequestBuilder, +) -> Result { + let res = builder.send().await?; + let status = res.status(); + let data = 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]; + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + embedding: Vec, +} + fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result { let ChatCompletionsData { messages, diff --git a/src/client/openai.rs b/src/client/openai.rs index 3cdea24..0da8166 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,10 +1,6 @@ -use super::{ - catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, OpenAIClient, PromptAction, PromptKind, - SseHandler, SseMmessage, ToolCall, -}; +use super::*; -use anyhow::{bail, Result}; +use anyhow::{bail, Context, Result}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -39,11 +35,11 @@ impl OpenAIClient { let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string()); let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!("{api_base}/chat/completions"); - debug!("OpenAI Request: {url} {body}"); + debug!("OpenAI Chat Completions Request: {url} {body}"); let mut builder = client.post(url).bearer_auth(api_key).json(&body); @@ -53,6 +49,25 @@ impl OpenAIClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + let api_base = self.get_api_base().unwrap_or_else(|_| API_BASE.to_string()); + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + debug!("OpenAI Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } pub async fn openai_chat_completions(builder: RequestBuilder) -> Result { @@ -125,6 +140,30 @@ pub async fn openai_chat_completions_streaming( sse_stream(builder, handle).await } +pub async fn openai_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.data.into_iter().map(|v| v.embedding).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + data: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyEmbedding { + embedding: Vec, +} + pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { let ChatCompletionsData { messages, @@ -201,6 +240,15 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod body } + +pub fn openai_build_embeddings_body(data: EmbeddingsData, model: &Model) -> Value { + json!({ + "input": data.texts, + "model": model.name(), + "encoding_format": "float", + }) +} + pub fn openai_extract_chat_completions(data: &Value) -> Result { let text = data["choices"][0]["message"]["content"] .as_str() @@ -244,5 +292,6 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result Result { - let api_base = match self.get_api_base() { - Ok(v) => v, - Err(err) => { - match OPENAI_COMPATIBLE_PLATFORMS - .into_iter() - .find_map(|(name, api_base)| { - if name == self.model.client_name() { - Some(api_base.to_string()) - } else { - None - } - }) { - Some(v) => v, - None => return Err(err), - } - } - }; let api_key = self.get_api_key().ok(); + let api_base = self.get_api_base_ext()?; let mut body = openai_build_chat_completions_body(data, &self.model); - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let chat_endpoint = self .config @@ -71,7 +53,7 @@ impl OpenAICompatibleClient { let url = format!("{api_base}{chat_endpoint}"); - debug!("OpenAICompatible Request: {url} {body}"); + debug!("OpenAICompatible Chat Completions Request: {url} {body}"); let mut builder = client.post(url).json(&body); if let Some(api_key) = api_key { @@ -80,10 +62,51 @@ impl OpenAICompatibleClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + let api_base = self.get_api_base_ext()?; + + let body = openai_build_embeddings_body(data, &self.model); + + let url = format!("{api_base}/embeddings"); + + debug!("OpenAICompatible Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } + + fn get_api_base_ext(&self) -> Result { + let api_base = match self.get_api_base() { + Ok(v) => v, + Err(err) => { + match OPENAI_COMPATIBLE_PLATFORMS + .into_iter() + .find_map(|(name, api_base)| { + if name == self.model.client_name() { + Some(api_base.to_string()) + } else { + None + } + }) { + Some(v) => v, + None => return Err(err), + } + } + }; + Ok(api_base) + } } impl_client_trait!( OpenAICompatibleClient, - crate::client::openai::openai_chat_completions, - crate::client::openai::openai_chat_completions_streaming + openai_chat_completions, + openai_chat_completions_streaming, + openai_embeddings ); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 0230e21..c34e409 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -1,8 +1,4 @@ -use super::{ - maybe_catch_error, message::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, QianwenClient, - SseHandler, SseMmessage, -}; +use super::*; use crate::utils::{base64_decode, sha256}; @@ -16,12 +12,15 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::borrow::BorrowMut; -const API_URL: &str = +const CHAT_COMPLETIONS_API_URL: &str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation"; -const API_URL_VL: &str = +const CHAT_COMPLETIONS_API_URL_VL: &str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"; +const EMBEDDINGS_API_URL: &str = + "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"; + #[derive(Debug, Clone, Deserialize, Default)] pub struct QianwenConfig { pub name: Option, @@ -48,13 +47,13 @@ impl QianwenClient { let stream = data.stream; let url = match self.model.supports_vision() { - true => API_URL_VL, - false => API_URL, + true => CHAT_COMPLETIONS_API_URL_VL, + false => CHAT_COMPLETIONS_API_URL, }; let (mut body, has_upload) = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); - debug!("Qianwen Request: {url} {body}"); + debug!("Qianwen Chat Completions Request: {url} {body}"); let mut builder = client.post(url).bearer_auth(api_key).json(&body); if stream { @@ -66,6 +65,37 @@ impl QianwenClient { Ok(builder) } + + fn embeddings_builder( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result { + let api_key = self.get_api_key()?; + + let text_type = match data.query { + true => "query", + false => "document", + }; + + let body = json!({ + "model": self.model.name(), + "input": { + "texts": data.texts, + }, + "parameters": { + "text_type": text_type, + } + }); + + let url = EMBEDDINGS_API_URL; + + debug!("Qianwen Embeddings Request: {url} {body}"); + + let builder = client.post(url).bearer_auth(api_key).json(&body); + + Ok(builder) + } } #[async_trait] @@ -94,6 +124,15 @@ impl Client for QianwenClient { let builder = self.chat_completions_builder(client, data)?; chat_completions_streaming(builder, handler, &self.model).await } + + async fn embeddings_inner( + &self, + client: &ReqwestClient, + data: EmbeddingsData, + ) -> Result>> { + let builder = self.embeddings_builder(client, data)?; + embeddings(builder).await + } } async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result { @@ -210,6 +249,31 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu Ok((body, has_upload)) } +async fn embeddings( + builder: RequestBuilder, +) -> Result { + let data: Value = builder.send().await?.json().await?; + maybe_catch_error(&data)?; + let res_body: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?; + let output = res_body.output.embeddings.into_iter().map(|v| v.embedding).collect(); + Ok(output) +} + +#[derive(Deserialize)] +struct EmbeddingsResBody { + output: EmbeddingsResBodyOutput, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyOutput { + embeddings: Vec, +} + +#[derive(Deserialize)] +struct EmbeddingsResBodyOutputEmbedding { + embedding: Vec, +} + fn extract_chat_completions_text(data: &Value, model: &Model) -> Result { let err = || anyhow!("Invalid response data: {data}"); let text = if model.name() == "qwen-long" { diff --git a/src/client/replicate.rs b/src/client/replicate.rs index 92c7e18..e96ed64 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -1,8 +1,5 @@ -use super::{ - catch_error, prompt_format::*, sse_stream, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, ReplicateClient, - SseHandler, SseMmessage, -}; +use super::*; +use super::prompt_format::*; use anyhow::{anyhow, Result}; use async_trait::async_trait; @@ -36,7 +33,7 @@ impl ReplicateClient { api_key: &str, ) -> Result { let mut body = build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); let url = format!("{API_BASE}/models/{}/predictions", self.model.name()); 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 { diff --git a/src/client/vertexai_claude.rs b/src/client/vertexai_claude.rs index bdce7d8..3993078 100644 --- a/src/client/vertexai_claude.rs +++ b/src/client/vertexai_claude.rs @@ -1,8 +1,7 @@ -use super::{ - access_token::*, claude::*, vertexai::*, ChatCompletionsData, ChatCompletionsOutput, Client, - ExtraConfig, Model, ModelData, ModelPatches, PromptAction, PromptKind, SseHandler, - VertexAIClaudeClient, -}; +use super::*; +use super::access_token::*; +use super::claude::*; +use super::vertexai::*; use anyhow::Result; use async_trait::async_trait; @@ -46,7 +45,7 @@ impl VertexAIClaudeClient { ); let mut body = claude_build_chat_completions_body(data, &self.model)?; - self.patch_request_body(&mut body); + self.patch_chat_completions_body(&mut body); if let Some(body_obj) = body.as_object_mut() { body_obj.remove("model"); } diff --git a/src/config/input.rs b/src/config/input.rs index ae94799..56ae5ed 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,11 +1,11 @@ use super::{role::Role, session::Session, GlobalConfig}; use crate::client::{ - init_client, list_models, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, + init_client, list_chat_models, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, MessageContentPart, MessageRole, Model, }; use crate::function::{ToolCallResult, ToolResults}; -use crate::utils::{base64_encode, sha256}; +use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; @@ -29,9 +29,11 @@ lazy_static! { pub struct Input { config: GlobalConfig, text: String, + patch_text: Option, medias: Vec, data_urls: HashMap, tool_call: Option, + rag: Option, context: InputContext, } @@ -40,9 +42,11 @@ impl Input { Self { config: config.clone(), text: text.to_string(), + patch_text: None, medias: Default::default(), data_urls: Default::default(), tool_call: None, + rag: None, context: context.unwrap_or_else(|| InputContext::from_config(config)), } } @@ -92,9 +96,11 @@ impl Input { Ok(Self { config: config.clone(), text: texts.join("\n"), + patch_text: None, medias, data_urls, tool_call: Default::default(), + rag: None, context: context.unwrap_or_else(|| InputContext::from_config(config)), }) } @@ -108,13 +114,41 @@ impl Input { } pub fn text(&self) -> String { - self.text.clone() + match self.patch_text.clone() { + Some(text) => text, + None => self.text.clone(), + } } pub fn set_text(&mut self, text: String) { self.text = text; } + pub async fn maybe_embeddings(&mut self, abort_signal: AbortSignal) -> Result<()> { + if self.text.is_empty() { + return Ok(()); + } + if !self.text.is_empty() { + let rag = self.config.read().rag.clone(); + if let Some(rag) = rag { + let top_k = self.config.read().rag_top_k; + let embeddings = rag.search(&self.text, top_k, abort_signal).await?; + let text = self.config.read().rag_template(&embeddings, &self.text); + self.patch_text = Some(text); + self.rag = Some(rag.name().to_string()); + } + } + Ok(()) + } + + pub fn rag(&self) -> Option<&str> { + self.rag.as_deref() + } + + pub fn clear_patch_text(&mut self) { + self.patch_text.take(); + } + pub fn merge_tool_call( mut self, output: String, @@ -134,7 +168,7 @@ impl Input { let model = self.config.read().model.clone(); if let Some(model_id) = self.role().and_then(|v| v.model_id.clone()) { if model.id() != model_id { - if let Some(model) = list_models(&self.config.read()) + if let Some(model) = list_chat_models(&self.config.read()) .into_iter() .find(|v| v.id() == model_id) { @@ -158,7 +192,7 @@ impl Input { bail!("The current model does not support vision."); } let messages = self.build_messages()?; - self.config.read().model.max_input_tokens_limit(&messages)?; + self.config.read().model.guard_max_input_tokens(&messages)?; let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session) { (session.temperature(), session.top_p()) @@ -262,12 +296,12 @@ impl Input { pub fn render(&self) -> String { if self.medias.is_empty() { - return self.text.clone(); + return self.text(); } let text = if self.text.is_empty() { - self.text.to_string() + String::new() } else { - format!(" -- {}", self.text) + format!(" -- {}", self.text()) }; let files: Vec = self .medias @@ -280,7 +314,7 @@ impl Input { pub fn message_content(&self) -> MessageContent { if self.medias.is_empty() { - MessageContent::Text(self.text.clone()) + MessageContent::Text(self.text()) } else { let mut list: Vec = self .medias @@ -291,12 +325,7 @@ impl Input { }) .collect(); if !self.text.is_empty() { - list.insert( - 0, - MessageContentPart::Text { - text: self.text.clone(), - }, - ); + list.insert(0, MessageContentPart::Text { text: self.text() }); } MessageContent::Array(list) } diff --git a/src/config/mod.rs b/src/config/mod.rs index 5ae6bec..fc2df17 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -7,14 +7,15 @@ pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - create_client_config, list_client_types, list_models, ClientConfig, Model, + create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{Function, ToolCallResult}; +use crate::rag::{Rag, TEMP_RAG_NAME}; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, - set_text, + set_text, AbortSignal, }; use anyhow::{anyhow, bail, Context, Result}; @@ -42,6 +43,7 @@ const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; +const RAGS_DIR_NAME: &str = "rags"; const FUNCTIONS_DIR_NAME: &str = "functions"; const CLIENTS_FIELD: &str = "clients"; @@ -49,7 +51,16 @@ const CLIENTS_FIELD: &str = "clients"; const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; -const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{color.cyan}{?session )}{!session >}{color.reset} "; + +const RAG_TEMPLATE: &str = r#"Answer the following question based only on the provided context: + +__CONTEXT__ + + +Question: __INPUT__ +"#; + +const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{?rag #{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; const RIGHT_PROMPT: &str = "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"; #[derive(Debug, Clone, Deserialize)] @@ -71,6 +82,9 @@ pub struct Config { pub keybindings: Keybindings, pub prelude: Option, pub buffer_editor: Option, + pub embedding_model: Option, + pub rag_top_k: usize, + pub rag_template: Option, pub function_calling: bool, pub compress_threshold: usize, pub summarize_prompt: Option, @@ -85,6 +99,8 @@ pub struct Config { #[serde(skip)] pub session: Option, #[serde(skip)] + pub rag: Option>, + #[serde(skip)] pub model: Model, #[serde(skip)] pub function: Function, @@ -111,6 +127,9 @@ impl Default for Config { keybindings: Default::default(), prelude: None, buffer_editor: None, + embedding_model: None, + rag_top_k: 4, + rag_template: None, function_calling: false, compress_threshold: 4000, summarize_prompt: None, @@ -121,6 +140,7 @@ impl Default for Config { roles: vec![], role: None, session: None, + rag: None, model: Default::default(), function: Default::default(), working_mode: WorkingMode::Command, @@ -170,12 +190,12 @@ impl Config { match prelude.split_once(':') { Some(("role", name)) => { if self.role.is_none() && self.session.is_none() { - self.set_role(name).with_context(err_msg)?; + self.use_role(name).with_context(err_msg)?; } } Some(("session", name)) => { if self.session.is_none() { - self.start_session(Some(name)).with_context(err_msg)?; + self.use_session(Some(name)).with_context(err_msg)?; } } _ => { @@ -223,10 +243,11 @@ impl Config { pub fn save_message( &mut self, - input: &Input, + input: &mut Input, output: &str, tool_call_results: &[ToolCallResult], ) -> Result<()> { + input.clear_patch_text(); self.last_message = Some((input.clone(), output.to_string())); if self.dry_run || output.is_empty() || !tool_call_results.is_empty() { @@ -248,17 +269,13 @@ impl Config { let timestamp = now(); let summary = input.summary(); let input_markdown = input.render(); - let output = match input.role() { - None => { - format!("# CHAT: {summary} [{timestamp}]\n{input_markdown}\n--------\n{output}\n--------\n\n",) - } - Some(v) => { - format!( - "# CHAT: {summary} [{timestamp}] ({})\n{input_markdown}\n--------\n{output}\n--------\n\n", - v.name, - ) - } + let scope = match (input.role().map(|v| v.name.as_str()), input.rag()) { + (Some(role), Some(rag)) => format!(" ({role}#{rag})"), + (Some(role), _) => format!(" ({role})"), + (None, Some(rag)) => format!(" (#{rag})"), + _ => String::new(), }; + let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); file.write_all(output.as_bytes()) .with_context(|| "Failed to save message") } @@ -289,6 +306,10 @@ impl Config { Self::local_path(SESSIONS_DIR_NAME) } + pub fn rags_dir() -> Result { + Self::local_path(RAGS_DIR_NAME) + } + pub fn functions_dir() -> Result { Self::local_path(FUNCTIONS_DIR_NAME) } @@ -299,17 +320,23 @@ impl Config { Ok(path) } - pub fn set_prompt(&mut self, prompt: &str) -> Result<()> { + pub fn rag_file(name: &str) -> Result { + let mut path = Self::rags_dir()?; + path.push(&format!("{name}.bin")); + Ok(path) + } + + pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { let role = Role::temp(prompt); - self.set_role_obj(role) + self.use_role_obj(role) } - pub fn set_role(&mut self, name: &str) -> Result<()> { + pub fn use_role(&mut self, name: &str) -> Result<()> { let role = self.retrieve_role(name)?; - self.set_role_obj(role) + self.use_role_obj(role) } - pub fn set_role_obj(&mut self, role: Role) -> Result<()> { + pub fn use_role_obj(&mut self, role: Role) -> Result<()> { if let Some(session) = self.session.as_mut() { session.guard_empty()?; session.set_role_properties(&role); @@ -321,7 +348,7 @@ impl Config { Ok(()) } - pub fn clear_role(&mut self) -> Result<()> { + pub fn exit_role(&mut self) -> Result<()> { self.role = None; self.restore_model()?; Ok(()) @@ -337,7 +364,10 @@ impl Config { } } if self.role.is_some() { - flags |= StateFlags::ROLE + flags |= StateFlags::ROLE; + } + if self.rag.is_some() { + flags |= StateFlags::RAG; } flags } @@ -393,7 +423,7 @@ impl Config { } pub fn set_model(&mut self, value: &str) -> Result<()> { - let models = list_models(self); + let models = list_chat_models(self); let model = Model::find(&models, value); match model { None => bail!("No model '{}'", value), @@ -442,6 +472,7 @@ impl Config { ), ("temperature", format_option_value(&temperature)), ("top_p", format_option_value(&top_p)), + ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), ("dry_run", self.dry_run.to_string()), @@ -458,6 +489,7 @@ impl Config { ("roles_file", display_path(&Self::roles_file()?)), ("messages_file", display_path(&Self::messages_file()?)), ("sessions_dir", display_path(&Self::sessions_dir()?)), + ("rags_dir", display_path(&Self::rags_dir()?)), ("functions_dir", display_path(&Self::functions_dir()?)), ]; let output = items @@ -486,11 +518,21 @@ impl Config { } } + pub fn rag_info(&self) -> Result { + if let Some(rag) = &self.rag { + rag.export() + } else { + bail!("No rag") + } + } + pub fn info(&self) -> Result { if let Some(session) = &self.session { session.export() } else if let Some(role) = &self.role { role.export() + } else if let Some(rag) = &self.rag { + rag.export() } else { self.system_info() } @@ -511,7 +553,7 @@ impl Config { .iter() .map(|v| (v.name.clone(), String::new())) .collect(), - ".model" => list_models(self) + ".model" => list_chat_models(self) .into_iter() .map(|v| (v.id(), v.description())) .collect(), @@ -520,10 +562,16 @@ impl Config { .into_iter() .map(|v| (v.clone(), String::new())) .collect(), + ".rag" => self + .list_rags() + .into_iter() + .map(|v| (v.clone(), String::new())) + .collect(), ".set" => vec![ "max_output_tokens", "temperature", "top_p", + "rag_top_k", "function_calling", "compress_threshold", "save", @@ -592,6 +640,11 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "rag_top_k" => { + if let Some(value) = parse_value(value)? { + self.rag_top_k = value; + } + } "function_calling" => { let value = value.parse().with_context(|| "Invalid value")?; self.function_calling = value; @@ -625,7 +678,7 @@ impl Config { Ok(()) } - pub fn start_session(&mut self, session: Option<&str>) -> Result<()> { + pub fn use_session(&mut self, session: Option<&str>) -> Result<()> { if self.session.is_some() { bail!( "Already in a session, please run '.exit session' first to exit the current session." @@ -671,7 +724,7 @@ impl Config { Ok(()) } - pub fn end_session(&mut self) -> Result<()> { + pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { self.last_message = None; let save_session = session.save_session(); @@ -767,6 +820,74 @@ impl Config { } } + pub async fn use_rag( + config: &GlobalConfig, + rag: Option<&str>, + abort_signal: AbortSignal, + ) -> Result<()> { + if config.read().rag.is_some() { + bail!("Already in a rag, please run '.exit rag' first to exit the current rag."); + } + let rag = match rag { + None => { + let rag_path = Self::rag_file(TEMP_RAG_NAME)?; + if rag_path.exists() { + remove_file(&rag_path).with_context(|| { + format!("Failed to cleanup previous '{TEMP_RAG_NAME}' rag") + })?; + } + Rag::init(config, TEMP_RAG_NAME, &rag_path, abort_signal).await? + } + Some(name) => { + let rag_path = Self::rag_file(name)?; + if !rag_path.exists() { + Rag::init(config, name, &rag_path, abort_signal).await? + } else { + Rag::load(config, name, &rag_path)? + } + } + }; + config.write().rag = Some(Arc::new(rag)); + Ok(()) + } + + pub fn exit_rag(&mut self) -> Result<()> { + self.rag.take(); + Ok(()) + } + + pub fn list_rags(&self) -> Vec { + let rags_dir = match Self::rags_dir() { + Ok(dir) => dir, + Err(_) => return vec![], + }; + match read_dir(rags_dir) { + Ok(rd) => { + let mut names = vec![]; + for entry in rd.flatten() { + let name = entry.file_name(); + if let Some(name) = name.to_string_lossy().strip_suffix(".bin") { + names.push(name.to_string()); + } + } + names.sort_unstable(); + names + } + Err(_) => vec![], + } + } + + pub fn rag_template(&self, embeddings: &str, text: &str) -> String { + if embeddings.is_empty() { + return text.to_string(); + } + self.rag_template + .as_deref() + .unwrap_or(RAG_TEMPLATE) + .replace("__CONTEXT__", embeddings) + .replace("__INPUT__", text) + } + pub fn get_render_options(&self) -> Result { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -858,6 +979,9 @@ impl Config { output.insert("consume_percent", percent.to_string()); output.insert("user_messages_len", session.user_messages_len().to_string()); } + if let Some(rag) = &self.rag { + output.insert("rag", rag.name().to_string()); + } if self.highlight { output.insert("color.reset", "\u{1b}[0m".to_string()); @@ -974,7 +1098,7 @@ impl Config { fn setup_model(&mut self) -> Result<()> { let model_id = if self.model_id.is_empty() { - let models = list_models(self); + let models = list_chat_models(self); if models.is_empty() { bail!("No available model"); } @@ -1049,6 +1173,7 @@ bitflags::bitflags! { const ROLE = 1 << 0; const SESSION_EMPTY = 1 << 1; const SESSION = 1 << 2; + const RAG = 1 << 3; } } @@ -1090,12 +1215,12 @@ fn create_config_file(config_path: &Path) -> Result<()> { std::fs::set_permissions(config_path, perms)?; } - println!("✨ Saved config file to {}\n", config_path.display()); + println!("✨ Saved config file to '{}'\n", config_path.display()); Ok(()) } -fn ensure_parent_exists(path: &Path) -> Result<()> { +pub(crate) fn ensure_parent_exists(path: &Path) -> Result<()> { if path.exists() { return Ok(()); } diff --git a/src/config/session.rs b/src/config/session.rs index 8ac5ad9..14e0731 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -160,7 +160,7 @@ impl Session { data["messages"] = json!(self.messages); let output = serde_yaml::to_string(&data) - .with_context(|| format!("Unable to show info about session {}", &self.name))?; + .with_context(|| format!("Unable to show info about session '{}'", &self.name))?; Ok(output) } diff --git a/src/main.rs b/src/main.rs index 0a5e404..3222285 100644 --- a/src/main.rs +++ b/src/main.rs @@ -3,6 +3,7 @@ mod client; mod config; mod function; mod logger; +mod rag; mod render; mod repl; mod serve; @@ -13,12 +14,12 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{list_models, send_stream, ChatCompletionsOutput}; +use crate::client::{list_chat_models, send_stream, ChatCompletionsOutput}; use crate::config::{ Config, GlobalConfig, Input, InputContext, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, }; -use crate::function::eval_tool_calls; +use crate::function::{eval_tool_calls, need_send_call_results}; use crate::render::{render_error, MarkdownRender}; use crate::repl::Repl; use crate::utils::{ @@ -29,14 +30,12 @@ use crate::utils::{ use anyhow::{bail, Result}; use async_recursion::async_recursion; use clap::Parser; -use function::need_send_call_results; use inquire::{Select, Text}; use is_terminal::IsTerminal; use parking_lot::RwLock; use std::io::{stderr, stdin, stdout, Read}; use std::process; use std::sync::Arc; -use tokio::sync::oneshot; #[tokio::main] async fn main() -> Result<()> { @@ -67,7 +66,7 @@ async fn main() -> Result<()> { return Ok(()); } if cli.list_models { - for model in list_models(&config.read()) { + for model in list_chat_models(&config.read()) { println!("{}", model.id()); } return Ok(()); @@ -87,18 +86,18 @@ async fn main() -> Result<()> { config.write().dry_run = true; } if let Some(prompt) = &cli.prompt { - config.write().set_prompt(prompt)?; + config.write().use_prompt(prompt)?; } else if let Some(name) = &cli.role { - config.write().set_role(name)?; + config.write().use_role(name)?; } else if cli.execute { - config.write().set_role(SHELL_ROLE)?; + config.write().use_role(SHELL_ROLE)?; } else if cli.code { - config.write().set_role(CODE_ROLE)?; + config.write().use_role(CODE_ROLE)?; } if let Some(session) = &cli.session { config .write() - .start_session(session.as_ref().map(|v| v.as_str()))?; + .use_session(session.as_ref().map(|v| v.as_str()))?; } if let Some(model) = &cli.model { config.write().set_model(model)?; @@ -142,7 +141,7 @@ async fn main() -> Result<()> { #[async_recursion] async fn start_directive( config: &GlobalConfig, - input: Input, + mut input: Input, no_stream: bool, code_mode: bool, ) -> Result<()> { @@ -176,8 +175,8 @@ async fn start_directive( }; config .write() - .save_message(&input, &output, &tool_call_results)?; - config.write().end_session()?; + .save_message(&mut input, &output, &tool_call_results)?; + config.write().exit_session()?; if need_send_call_results(&tool_call_results) { start_directive( config, @@ -201,10 +200,9 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let client = input.create_client()?; let is_terminal_stdout = stdout().is_terminal(); let ret = if is_terminal_stdout { - let (spinner_tx, spinner_rx) = oneshot::channel(); - tokio::spawn(run_spinner(" Generating", spinner_rx)); + let (stop_spinner_tx, _) = run_spinner("Generating").await; let ret = client.chat_completions(input.clone()).await; - let _ = spinner_tx.send(()); + let _ = stop_spinner_tx.send(()); ret } else { client.chat_completions(input.clone()).await @@ -213,7 +211,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } - config.write().save_message(&input, &eval_str, &[])?; + config.write().save_message(&mut input, &eval_str, &[])?; config.read().maybe_copy(&eval_str); let render_options = config.read().get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; diff --git a/src/rag/loader.rs b/src/rag/loader.rs new file mode 100644 index 0000000..106802a --- /dev/null +++ b/src/rag/loader.rs @@ -0,0 +1,146 @@ +use super::RagDocument; + +use anyhow::{bail, Context, Result}; +use async_recursion::async_recursion; +use std::{path::Path, process::Command}; +use tokio::fs; + +pub async fn load(path: &str, extension: &str) -> Result> { + match extension { + "docx" | "epub" | "ipynb" => load_pandoc(path) + .await + .context("Failed to load with pandoc"), + "pdf" => load_pdf(path).await, + _ => load_plain(path).await, + } +} + +async fn load_plain(path: &str) -> Result> { + let contents = fs::read_to_string(path).await?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pdf(path: &str) -> Result> { + let contents = pdf_extract::extract_text(path)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +async fn load_pandoc(path: &str) -> Result> { + let output = Command::new("pandoc") + .arg("--to") + .arg("plain") + .arg(path) + .output()?; + + if !output.status.success() { + let stderr = String::from_utf8_lossy(&output.stderr); + bail!( + "Pandoc conversion failed with exit code {:?}: {}", + output.status.code(), + stderr + ); + } + + let contents = std::str::from_utf8(&output.stdout)?; + let document = RagDocument::new(contents); + Ok(vec![document]) +} + +pub fn parse_glob(path_str: &str) -> Result<(String, Vec)> { + if let Some(start) = path_str.find("/**/*.").or_else(|| path_str.find(r"\**\*.")) { + let base_path = path_str[..start].to_string(); + if let Some(curly_brace_end) = path_str[start..].find('}') { + let end = start + curly_brace_end; + let extensions_str = &path_str[start + 6..end + 1]; + let extensions = if extensions_str.starts_with('{') && extensions_str.ends_with('}') { + extensions_str[1..extensions_str.len() - 1] + .split(',') + .map(|s| s.to_string()) + .collect::>() + } else { + bail!("Invalid path '{path_str}'"); + }; + Ok((base_path, extensions)) + } else { + let extensions_str = &path_str[start + 6..]; + let extensions = vec![extensions_str.to_string()]; + Ok((base_path, extensions)) + } + } else { + Ok((path_str.to_string(), vec![])) + } +} + +#[async_recursion] +pub async fn list_files( + files: &mut Vec, + entry_path: &Path, + suffixes: Option<&Vec>, +) -> Result<()> { + if !entry_path.exists() { + bail!("Not found: {:?}", entry_path); + } + if entry_path.is_file() { + add_file(files, suffixes, entry_path); + return Ok(()); + } + if !entry_path.is_dir() { + bail!("Not a directory: {:?}", entry_path); + } + let mut reader = fs::read_dir(entry_path).await?; + while let Some(entry) = reader.next_entry().await? { + let path = entry.path(); + if path.is_file() { + add_file(files, suffixes, &path); + } else if path.is_dir() { + list_files(files, &path, suffixes).await?; + } + } + Ok(()) +} + +fn add_file(files: &mut Vec, suffixes: Option<&Vec>, path: &Path) { + if is_valid_extension(suffixes, path) { + files.push(path.display().to_string()); + } +} + +fn is_valid_extension(suffixes: Option<&Vec>, path: &Path) -> bool { + if let Some(suffixes) = suffixes { + if !suffixes.is_empty() { + if let Some(extension) = path.extension().map(|v| v.to_string_lossy().to_string()) { + return suffixes.contains(&extension); + } + return false; + } + } + true +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_glob() { + assert_eq!(parse_glob("dir").unwrap(), ("dir".into(), vec![])); + assert_eq!( + parse_glob("dir/file.md").unwrap(), + ("dir/file.md".into(), vec![]) + ); + assert_eq!( + parse_glob("dir/**/*.md").unwrap(), + ("dir".into(), vec!["md".into()]) + ); + assert_eq!( + parse_glob("dir/**/*.{md,txt}").unwrap(), + ("dir".into(), vec!["md".into(), "txt".into()]) + ); + assert_eq!( + parse_glob("C:\\dir\\**\\*.{md,txt}").unwrap(), + ("C:\\dir".into(), vec!["md".into(), "txt".into()]) + ); + } +} diff --git a/src/rag/mod.rs b/src/rag/mod.rs new file mode 100644 index 0000000..387d3d9 --- /dev/null +++ b/src/rag/mod.rs @@ -0,0 +1,425 @@ +use self::loader::*; +use self::splitter::*; + +use crate::client::*; +use crate::config::*; +use crate::utils::*; + +mod loader; +mod splitter; + +use anyhow::bail; +use anyhow::{anyhow, Context, Result}; +use hnsw_rs::prelude::*; +use indexmap::IndexMap; +use inquire::{required, validator::Validation, Select, Text}; +use path_absolutize::Absolutize; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use std::fmt::Debug; +use std::{io::BufReader, path::Path}; +use tokio::sync::mpsc; + +pub const TEMP_RAG_NAME: &str = "temp"; +pub const CHUNK_OVERLAP: usize = 20; +pub const SIMILARITY_THRESHOLD: f32 = 0.25; + +pub struct Rag { + client: Box, + name: String, + path: String, + model: Model, + hnsw: Hnsw<'static, f32, DistCosine>, + data: RagData, +} + +impl Debug for Rag { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Rag") + .field("name", &self.name) + .field("path", &self.path) + .field("model", &self.model) + .field("data", &self.data) + .finish() + } +} + +impl Rag { + pub async fn init( + config: &GlobalConfig, + name: &str, + path: &Path, + abort_signal: AbortSignal, + ) -> Result { + debug!("init rag: {name}"); + let model = select_embedding_model(config)?; + let chunk_size = model.default_chunk_size(); + let chunk_size = set_chunk_size(chunk_size)?; + let data = RagData::new(&model.id(), chunk_size); + let mut rag = Self::create(config, name, path, data)?; + let paths = add_document_paths()?; + debug!("document paths: {paths:?}"); + let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; + tokio::select! { + ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { + let _ = stop_spinner_tx.send(()); + ret?; + } + _ = watch_abort_signal(abort_signal) => { + let _ = stop_spinner_tx.send(()); + bail!("Aborted!") + }, + }; + if !rag.is_temp() { + rag.save(path)?; + println!("✨ Saved rag to '{}'", path.display()); + } + Ok(rag) + } + + pub fn load(config: &GlobalConfig, name: &str, path: &Path) -> Result { + let err = || format!("Failed to load rag '{name}'"); + let file = std::fs::File::open(path).with_context(err)?; + let reader = BufReader::new(file); + let data: RagData = bincode::deserialize_from(reader).with_context(err)?; + Self::create(config, name, path, data) + } + + pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result { + let hnsw = data.build_hnsw(); + let model = retrieve_embedding_model(&config.read(), &data.model)?; + let client = init_client(config, Some(model.clone()))?; + let rag = Rag { + client, + name: name.to_string(), + path: path.display().to_string(), + data, + model, + hnsw, + }; + Ok(rag) + } + + pub fn save(&self, path: &Path) -> Result<()> { + ensure_parent_exists(path)?; + let mut file = std::fs::File::create(path)?; + bincode::serialize_into(&mut file, &self.data) + .with_context(|| format!("Failed to save rag '{}'", self.name))?; + Ok(()) + } + + pub fn export(&self) -> Result { + let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect(); + let data = json!({ + "path": self.path, + "model": self.model.id(), + "chunk_size": self.data.chunk_size, + "files": files, + }); + let output = serde_yaml::to_string(&data) + .with_context(|| format!("Unable to show info about rag '{}'", self.name))?; + Ok(output) + } + + pub fn name(&self) -> &str { + &self.name + } + + pub fn is_temp(&self) -> bool { + self.name == TEMP_RAG_NAME + } + + pub async fn search( + &self, + text: &str, + top_k: usize, + abort_signal: AbortSignal, + ) -> Result { + let (stop_spinner_tx, _) = run_spinner("Embedding").await; + let ret = tokio::select! { + ret = self.search_impl(text, top_k) => { + ret + } + _ = watch_abort_signal(abort_signal) => { + bail!("Aborted!") + }, + }; + let _ = stop_spinner_tx.send(()); + let output = ret?.join("\n\n"); + Ok(output) + } + + pub async fn add_paths>( + &mut self, + paths: &[T], + progress_tx: Option>, + ) -> Result<()> { + // List files + let mut file_paths = vec![]; + progress(&progress_tx, "Listing paths".into()); + for path in paths { + let path = path + .as_ref() + .absolutize() + .with_context(|| anyhow!("Invalid path '{}'", path.as_ref().display()))?; + let path_str = path.display().to_string(); + if self.data.files.iter().any(|v| v.path == path_str) { + continue; + } + let (path_str, suffixes) = parse_glob(&path_str)?; + let suffixes = if suffixes.is_empty() { + None + } else { + Some(&suffixes) + }; + list_files(&mut file_paths, Path::new(&path_str), suffixes).await?; + } + + // Load files + let mut rag_files = vec![]; + let file_paths_len = file_paths.len(); + progress(&progress_tx, format!("Loading files [1/{file_paths_len}]")); + for path in file_paths { + let extension = Path::new(&path) + .extension() + .map(|v| v.to_string_lossy().to_lowercase()) + .unwrap_or_default(); + let separator = autodetect_separator(&extension); + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, separator); + let documents = load(&path, &extension) + .await + .with_context(|| format!("Failed to load text at '{path}'"))?; + let documents = + splitter.split_documents(&documents, &SplitterChunkHeaderOptions::default()); + rag_files.push(RagFile { path, documents }); + progress( + &progress_tx, + format!("Loading files [{}/{file_paths_len}]", rag_files.len()), + ); + } + + if rag_files.is_empty() { + return Ok(()); + } + + // Convert vectors + let mut vector_ids = vec![]; + let mut texts = vec![]; + for (file_index, file) in rag_files.iter().enumerate() { + for (document_index, doc) in file.documents.iter().enumerate() { + vector_ids.push(combine_vector_id(file_index, document_index)); + texts.push(doc.page_content.clone()) + } + } + + let embeddings_data = EmbeddingsData::new(texts, false); + let embeddings = self + .create_embeddings(embeddings_data, progress_tx.clone()) + .await?; + + self.data.add(rag_files, vector_ids, embeddings); + progress(&progress_tx, "Building vector store".into()); + self.hnsw = self.data.build_hnsw(); + + Ok(()) + } + + async fn search_impl(&self, text: &str, top_k: usize) -> Result> { + let splitter = Splitter::new(self.data.chunk_size, CHUNK_OVERLAP, &DEFAULT_SEPARATES); + let texts = splitter.split_text(text); + let embeddings_data = EmbeddingsData::new(texts, true); + let embeddings = self.create_embeddings(embeddings_data, None).await?; + let output = self + .hnsw + .parallel_search(&embeddings, top_k, 30) + .into_iter() + .flat_map(|list| { + list.into_iter() + .filter_map(|v| { + if v.distance < SIMILARITY_THRESHOLD { + return None; + } + let (file_index, document_index) = split_vector_id(v.d_id); + let text = self.data.files[file_index].documents[document_index] + .page_content + .clone(); + Some(text) + }) + .collect::>() + }) + .collect(); + Ok(output) + } + + async fn create_embeddings( + &self, + data: EmbeddingsData, + progress_tx: Option>, + ) -> Result { + let EmbeddingsData { texts, query } = data; + let mut output = vec![]; + let chunks = texts.chunks(self.model.max_concurrent_chunks()); + let chunks_len = chunks.len(); + progress( + &progress_tx, + format!("Creating embeddings [1/{chunks_len}]"), + ); + for (index, texts) in chunks.enumerate() { + let chunk_data = EmbeddingsData { + texts: texts.to_vec(), + query, + }; + let chunk_output = self + .client + .embeddings(chunk_data) + .await + .context("Failed to create embedding")?; + output.extend(chunk_output); + progress( + &progress_tx, + format!("Creating embeddings [{}/{chunks_len}]", index + 1), + ); + } + Ok(output) + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagData { + pub model: String, + pub chunk_size: usize, + pub files: Vec, + pub vectors: IndexMap>, +} + +impl RagData { + pub fn new(model: &str, chunk_size: usize) -> Self { + Self { + model: model.to_string(), + chunk_size, + files: Default::default(), + vectors: Default::default(), + } + } + + pub fn add( + &mut self, + files: Vec, + vector_ids: Vec, + embeddings: EmbeddingsOutput, + ) { + self.files.extend(files); + self.vectors.extend(vector_ids.into_iter().zip(embeddings)); + } + + pub fn build_hnsw(&self) -> Hnsw<'static, f32, DistCosine> { + let hnsw = Hnsw::new(32, self.vectors.len(), 16, 200, DistCosine {}); + let list: Vec<_> = self.vectors.iter().map(|(k, v)| (v, *k)).collect(); + hnsw.parallel_insert(&list); + hnsw + } +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagFile { + path: String, + documents: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RagDocument { + pub page_content: String, + pub metadata: RagMetadata, +} + +impl RagDocument { + pub fn new>(page_content: S) -> Self { + RagDocument { + page_content: page_content.into(), + metadata: IndexMap::new(), + } + } + + #[allow(unused)] + pub fn with_metadata(mut self, metadata: RagMetadata) -> Self { + self.metadata = metadata; + self + } +} + +impl Default for RagDocument { + fn default() -> Self { + RagDocument { + page_content: "".to_string(), + metadata: IndexMap::new(), + } + } +} + +pub type RagMetadata = IndexMap; + +pub type VectorID = usize; + +pub fn combine_vector_id(file_index: usize, document_index: usize) -> VectorID { + file_index << (usize::BITS / 2) | document_index +} + +pub fn split_vector_id(value: VectorID) -> (usize, usize) { + let low_mask = (1 << (usize::BITS / 2)) - 1; + let low = value & low_mask; + let high = value >> (usize::BITS / 2); + (high, low) +} + +fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result { + let models = list_embedding_models(config); + let model = + Model::find(&models, model_id).ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; + Ok(model) +} + +fn select_embedding_model(config: &GlobalConfig) -> Result { + let config = config.read(); + let model = match config.embedding_model.clone() { + Some(model_id) => retrieve_embedding_model(&config, &model_id)?, + None => { + let models = list_embedding_models(&config); + if models.is_empty() { + bail!("No embedding model"); + } + let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); + let model_id = Select::new("Select embedding model:", model_ids).prompt()?; + retrieve_embedding_model(&config, &model_id)? + } + }; + Ok(model) +} + +fn set_chunk_size(chunk_size: usize) -> Result { + let value = Text::new("Set chunk size:") + .with_default(&chunk_size.to_string()) + .with_validator(move |text: &str| { + let out = match text.parse::() { + Ok(_) => Validation::Valid, + Err(_) => Validation::Invalid("Must be a integer".into()), + }; + Ok(out) + }) + .prompt()?; + value.parse().map_err(|_| anyhow!("Invalid chunk_size")) +} + +fn add_document_paths() -> Result> { + let text = Text::new("Add document paths:") + .with_validator(required!("This field is required")) + .with_help_message("e.g. file1;dir2/;dir3/**/*.md") + .prompt()?; + let paths = text.split(';').map(|v| v.to_string()).collect(); + Ok(paths) +} + +fn progress(spinner_message_tx: &Option>, message: String) { + if let Some(tx) = spinner_message_tx { + let _ = tx.send(message); + } +} diff --git a/src/rag/splitter.rs b/src/rag/splitter.rs new file mode 100644 index 0000000..5fdacee --- /dev/null +++ b/src/rag/splitter.rs @@ -0,0 +1,564 @@ +use super::{RagDocument, RagMetadata}; + +use std::cmp::Ordering; + +pub const DEFAULT_SEPARATES: [&str; 4] = ["\n\n", "\n", " ", ""]; +pub const HTML_SEPARATES: [&str; 28] = [ + // First, try to split along HTML tags + "", "
", "

", "
", "

  • ", "

    ", "

    ", "

    ", "

    ", "

    ", "
    ", + "", "", "", "
    ", "", "
      ", "
        ", "
        ", "