summaryrefslogtreecommitdiffstats
path: root/src
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
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')
-rw-r--r--src/client/azure_openai.rs31
-rw-r--r--src/client/bedrock.rs10
-rw-r--r--src/client/claude.rs9
-rw-r--r--src/client/cloudflare.rs7
-rw-r--r--src/client/cohere.rs69
-rw-r--r--src/client/common.rs140
-rw-r--r--src/client/ernie.rs9
-rw-r--r--src/client/gemini.rs73
-rw-r--r--src/client/model.rs44
-rw-r--r--src/client/ollama.rs67
-rw-r--r--src/client/openai.rs67
-rw-r--r--src/client/openai_compatible.rs73
-rw-r--r--src/client/qianwen.rs86
-rw-r--r--src/client/replicate.rs9
-rw-r--r--src/client/vertexai.rs79
-rw-r--r--src/client/vertexai_claude.rs11
-rw-r--r--src/config/input.rs59
-rw-r--r--src/config/mod.rs185
-rw-r--r--src/config/session.rs2
-rw-r--r--src/main.rs32
-rw-r--r--src/rag/loader.rs146
-rw-r--r--src/rag/mod.rs425
-rw-r--r--src/rag/splitter.rs564
-rw-r--r--src/render/stream.rs15
-rw-r--r--src/repl/mod.rs61
-rw-r--r--src/serve.rs11
-rw-r--r--src/utils/abort_signal.rs9
-rw-r--r--src/utils/mod.rs2
-rw-r--r--src/utils/spinner.rs39
29 files changed, 2048 insertions, 286 deletions
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<String>,
pub api_base: Option<String>,
pub api_key: Option<String>,
+ #[serde(default)]
pub models: Vec<ModelData>,
pub patches: Option<ModelPatches>,
pub extra: Option<ExtraConfig>,
@@ -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<RequestBuilder> {
+ 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<RequestBuilder> {
+ 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<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -100,6 +129,24 @@ async fn chat_completions_streaming(
Ok(())
}
+async fn 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")?;
+ Ok(res_body.embeddings)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ embeddings: Vec<Vec<f32>>,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
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<BuiltinModels> =
- serde_yaml::from_str(MODELS_YAML).unwrap();
+ pub static ref ALL_MODELS: Vec<BuiltinModels> = serde_yaml::from_str(MODELS_YAML).unwrap();
static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap();
}
@@ -92,10 +94,10 @@ macro_rules! register_client {
pub fn list_models(local_config: &$config) -> Vec<Model> {
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<Vec<$crate::client::Model>> = None;
+ static mut ALL_CLIENT_MODELS: Option<Vec<$crate::client::Model>> = 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<Model> {
- 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<Vec<Vec<f32>>> {
+ 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<Model>;
-
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<ReqwestClient> {
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<Vec<Vec<f32>>> {
+ 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<Vec<Vec<f32>>> {
+ bail!("No embeddings api")
+ }
}
impl Default for ClientConfig {
@@ -375,7 +411,7 @@ pub type ModelPatches = IndexMap<String, ModelPatch>;
#[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<String>,
+ pub query: bool,
+}
+
+impl EmbeddingsData {
+ pub fn new(texts: Vec<String>, query: bool) -> Self {
+ Self { texts, query }
+ }
+}
+
+pub type EmbeddingsOutput = Vec<Vec<f32>>;
+
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<Option<(St
"name": name,
"api_base": api_base,
});
- let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) {
+ let prompts = if ALL_MODELS.iter().any(|v| &v.platform == name) {
vec![("api_key", "API Key:", false, PromptKind::String)]
} else {
vec![
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 097ee68..77f4741 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,8 +1,5 @@
-use super::{
- access_token::*, maybe_catch_error, patch_system_message, sse_stream, ChatCompletionsData,
- ChatCompletionsOutput, Client, ErnieClient, ExtraConfig, Model, ModelData, ModelPatches,
- PromptAction, PromptKind, SseHandler, SseMmessage,
-};
+use super::*;
+use super::access_token::*;
use anyhow::{anyhow, Context, Result};
use async_trait::async_trait;
@@ -37,7 +34,7 @@ impl ErnieClient {
data: ChatCompletionsData,
) -> Result<RequestBuilder> {
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<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>,
+}
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<isize> {
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<usize>,
+ pub input_price: Option<f64>,
+ pub output_price: Option<f64>,
+
+ // chat-only properties
pub max_output_tokens: Option<isize>,
#[serde(default)]
pub pass_max_tokens: bool,
- pub input_price: Option<f64>,
- pub output_price: Option<f64>,
#[serde(default)]
pub supports_vision: bool,
#[serde(default)]
pub supports_function_calling: bool,
+
+ // embedding-only properties
+ pub default_chunk_size: Option<usize>,
+ pub max_concurrent_chunks: Option<usize>,
}
impl ModelData {
@@ -222,3 +252,7 @@ pub struct BuiltinModels {
pub platform: String,
pub models: Vec<ModelData>,
}
+
+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<String>,
pub api_base: Option<String>,
pub api_auth: Option<String>,
- pub chat_endpoint: Option<String>,
+ #[serde(default)]
pub models: Vec<ModelData>,
pub patches: Option<ModelPatches>,
pub extra: Option<ExtraConfig>,
@@ -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<RequestBuilder> {
+ 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<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -109,6 +133,25 @@ async fn chat_completions_streaming(
Ok(())
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ 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<f32>,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
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<RequestBuilder> {
+ 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<ChatCompletionsOutput> {
@@ -125,6 +140,30 @@ pub async fn openai_chat_completions_streaming(
sse_stream(builder, handle).await
}
+pub async fn openai_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 = res_body.data.into_iter().map(|v| v.embedding).collect();
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ data: Vec<EmbeddingsResBodyEmbedding>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyEmbedding {
+ embedding: Vec<f32>,
+}
+
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<ChatCompletionsOutput> {
let text = data["choices"][0]["message"]["content"]
.as_str()
@@ -244,5 +292,6 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
impl_client_trait!(
OpenAIClient,
openai_chat_completions,
- openai_chat_completions_streaming
+ openai_chat_completions_streaming,
+ openai_embeddings
);
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 74cd954..af7cd0e 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -1,7 +1,5 @@
-use super::{
- openai::*, ChatCompletionsData, Client, ExtraConfig, Model, ModelData, ModelPatches,
- OpenAICompatibleClient, PromptAction, PromptKind, OPENAI_COMPATIBLE_PLATFORMS,
-};
+use super::*;
+use super::openai::*;
use anyhow::Result;
use reqwest::{Client as ReqwestClient, RequestBuilder};
@@ -41,27 +39,11 @@ impl OpenAICompatibleClient {
client: &ReqwestClient,
data: ChatCompletionsData,
) -> Result<RequestBuilder> {
- 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<RequestBuilder> {
+ 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<String> {
+ 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<String>,
@@ -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<RequestBuilder> {
+ 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<Vec<Vec<f32>>> {
+ let builder = self.embeddings_builder(client, data)?;
+ embeddings(builder).await
+ }
}
async fn chat_completions(builder: RequestBuilder, model: &Model) -> Result<ChatCompletionsOutput> {
@@ -210,6 +249,31 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
Ok((body, has_upload))
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ 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<EmbeddingsResBodyOutputEmbedding>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyOutputEmbedding {
+ embedding: Vec<f32>,
+}
+
fn extract_chat_completions_text(data: &Value, model: &Model) -> Result<ChatCompletionsOutput> {
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<RequestBuilder> {
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<RequestBuilder> {
+ 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<Vec<Vec<f32>>> {
+ 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<ChatCompletionsOutput> {
@@ -138,6 +173,34 @@ pub async fn gemini_chat_completions_streaming(
Ok(())
}
+async fn 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 = res_body.predictions.into_iter().map(|v| v.embeddings.values).collect();
+ Ok(output)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ predictions: Vec<EmbeddingsResBodyPrediction>,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyPrediction {
+ embeddings: EmbeddingsResBodyPredictionEmbeddings,
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBodyPredictionEmbeddings {
+ values: Vec<f32>
+}
+
fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["candidates"][0]["content"]["parts"][0]["text"]
.as_str()
@@ -179,7 +242,7 @@ fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsO
Ok(output)
}
-pub(crate) fn gemini_build_chat_completions_body(
+pub fn gemini_build_chat_completions_body(
data: ChatCompletionsData,
model: &Model,
) -> Result<Value> {
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<String>,
medias: Vec<String>,
data_urls: HashMap<String, String>,
tool_call: Option<ToolResults>,
+ rag: Option<String>,
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<String> = 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<MessageContentPart> = 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>
+__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<String>,
pub buffer_editor: Option<String>,
+ pub embedding_model: Option<String>,
+ pub rag_top_k: usize,
+ pub rag_template: Option<String>,
pub function_calling: bool,
pub compress_threshold: usize,
pub summarize_prompt: Option<String>,
@@ -85,6 +99,8 @@ pub struct Config {
#[serde(skip)]
pub session: Option<Session>,
#[serde(skip)]
+ pub rag: Option<Arc<Rag>>,
+ #[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<PathBuf> {
+ Self::local_path(RAGS_DIR_NAME)
+ }
+
pub fn functions_dir() -> Result<PathBuf> {
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<PathBuf> {
+ 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<String> {
+ if let Some(rag) = &self.rag {
+ rag.export()
+ } else {
+ bail!("No rag")
+ }
+ }
+
pub fn info(&self) -> Result<String> {
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<String> {
+ 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<RenderOptions> {
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<Vec<RagDocument>> {
+ 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<Vec<RagDocument>> {
+ let contents = fs::read_to_string(path).await?;
+ let document = RagDocument::new(contents);
+ Ok(vec![document])
+}
+
+async fn load_pdf(path: &str) -> Result<Vec<RagDocument>> {
+ let contents = pdf_extract::extract_text(path)?;
+ let document = RagDocument::new(contents);
+ Ok(vec![document])
+}
+
+async fn load_pandoc(path: &str) -> Result<Vec<RagDocument>> {
+ 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<String>)> {
+ 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::<Vec<String>>()
+ } 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<String>,
+ entry_path: &Path,
+ suffixes: Option<&Vec<String>>,
+) -> 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<String>, suffixes: Option<&Vec<String>>, path: &Path) {
+ if is_valid_extension(suffixes, path) {
+ files.push(path.display().to_string());
+ }
+}
+
+fn is_valid_extension(suffixes: Option<&Vec<String>>, 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<dyn Client>,
+ 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<Self> {
+ 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<Self> {
+ 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<Self> {
+ 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<String> {
+ 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<String> {
+ 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<T: AsRef<Path>>(
+ &mut self,
+ paths: &[T],
+ progress_tx: Option<mpsc::UnboundedSender<String>>,
+ ) -> 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<Vec<String>> {
+ 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::<Vec<_>>()
+ })
+ .collect();
+ Ok(output)
+ }
+
+ async fn create_embeddings(
+ &self,
+ data: EmbeddingsData,
+ progress_tx: Option<mpsc::UnboundedSender<String>>,
+ ) -> Result<EmbeddingsOutput> {
+ 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<RagFile>,
+ pub vectors: IndexMap<VectorID, Vec<f32>>,
+}
+
+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<RagFile>,
+ vector_ids: Vec<VectorID>,
+ 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<RagDocument>,
+}
+
+#[derive(Debug, Clone, Serialize, Deserialize)]
+pub struct RagDocument {
+ pub page_content: String,
+ pub metadata: RagMetadata,
+}
+
+impl RagDocument {
+ pub fn new<S: Into<String>>(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<String, String>;
+
+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<Model> {
+ 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<Model> {
+ 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<usize> {
+ let value = Text::new("Set chunk size:")
+ .with_default(&chunk_size.to_string())
+ .with_validator(move |text: &str| {
+ let out = match text.parse::<usize>() {
+ 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<Vec<String>> {
+ 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<mpsc::UnboundedSender<String>>, 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
+ "<body>", "<div>", "<p>", "<br>", "<li>", "<h1>", "<h2>", "<h3>", "<h4>", "<h5>", "<h6>",
+ "<span>", "<table>", "<tr>", "<td>", "<th>", "<ul>", "<ol>", "<header>", "<footer>", "<nav>",
+ // Head
+ "<head>", "<style>", "<script>", "<meta>", "<title>", // Normal type of lines
+ " ", "",
+];
+pub const MARKDOWN_SEPARATES: [&str; 13] = [
+ // First, try to split along Markdown headings (starting with level 2)
+ "\n## ",
+ "\n### ",
+ "\n#### ",
+ "\n##### ",
+ "\n###### ",
+ // Note the alternative syntax for headings (below) is not handled here
+ // Heading level 2
+ // ---------------
+ // End of code block
+ "```\n\n",
+ // Horizontal lines
+ "\n\n***\n\n",
+ "\n\n---\n\n",
+ "\n\n___\n\n",
+ // Note that this splitter doesn't handle horizontal lines defined
+ // by *three or more* of ***, ---, or ___, but this is not handled
+ "\n\n",
+ "\n",
+ " ",
+ "",
+];
+pub const LATEX_SEPARATES: [&str; 19] = [
+ // First, try to split along Latex sections
+ "\n\\chapter{",
+ "\n\\section{",
+ "\n\\subsection{",
+ "\n\\subsubsection{",
+ // Now split by environments
+ "\n\\begin{enumerate}",
+ "\n\\begin{itemize}",
+ "\n\\begin{description}",
+ "\n\\begin{list}",
+ "\n\\begin{quote}",
+ "\n\\begin{quotation}",
+ "\n\\begin{verse}",
+ "\n\\begin{verbatim}",
+ // Now split by math environments
+ "\n\\begin{align}",
+ "$$",
+ "$",
+ // Now split by the normal type of lines
+ "\n\n",
+ "\n",
+ " ",
+ "",
+];
+
+pub fn autodetect_separator(extension: &str) -> &[&'static str] {
+ match extension {
+ "md" | "mkd" => &MARKDOWN_SEPARATES,
+ "htm" | "html" => &HTML_SEPARATES,
+ "tex" => &LATEX_SEPARATES,
+ _ => &DEFAULT_SEPARATES,
+ }
+}
+
+pub struct Splitter {
+ pub chunk_size: usize,
+ pub chunk_overlap: usize,
+ pub separators: Vec<String>,
+ pub length_function: Box<dyn Fn(&str) -> usize + Send + Sync>,
+}
+
+impl Default for Splitter {
+ fn default() -> Self {
+ Self {
+ chunk_size: 1000,
+ chunk_overlap: 20,
+ separators: DEFAULT_SEPARATES.iter().map(|v| v.to_string()).collect(),
+ length_function: Box::new(|text| text.len()),
+ }
+ }
+}
+
+// Builder pattern for Options struct
+impl Splitter {
+ pub fn new(chunk_size: usize, chunk_overlap: usize, separators: &[&str]) -> Self {
+ Self::default()
+ .with_chunk_size(chunk_size)
+ .with_chunk_overlap(chunk_overlap)
+ .with_separators(separators)
+ }
+
+ pub fn with_chunk_size(mut self, chunk_size: usize) -> Self {
+ self.chunk_size = chunk_size;
+ self
+ }
+
+ pub fn with_chunk_overlap(mut self, chunk_overlap: usize) -> Self {
+ self.chunk_overlap = chunk_overlap;
+ self
+ }
+
+ pub fn with_separators(mut self, separators: &[&str]) -> Self {
+ self.separators = separators.iter().map(|v| v.to_string()).collect();
+ self
+ }
+
+ #[allow(unused)]
+ pub fn with_length_function<F>(mut self, length_function: F) -> Self
+ where
+ F: Fn(&str) -> usize + Send + Sync + 'static,
+ {
+ self.length_function = Box::new(length_function);
+ self
+ }
+
+ pub fn split_documents(
+ &self,
+ documents: &[RagDocument],
+ chunk_header_options: &SplitterChunkHeaderOptions,
+ ) -> Vec<RagDocument> {
+ let mut texts: Vec<String> = Vec::new();
+ let mut metadatas: Vec<RagMetadata> = Vec::new();
+ documents.iter().for_each(|d| {
+ if !d.page_content.is_empty() {
+ texts.push(d.page_content.clone());
+ metadatas.push(d.metadata.clone());
+ }
+ });
+
+ self.create_documents(&texts, &metadatas, chunk_header_options)
+ }
+
+ pub fn create_documents(
+ &self,
+ texts: &[String],
+ metadatas: &[RagMetadata],
+ chunk_header_options: &SplitterChunkHeaderOptions,
+ ) -> Vec<RagDocument> {
+ let SplitterChunkHeaderOptions {
+ chunk_header,
+ chunk_overlap_header,
+ append_chunk_overlap_header,
+ } = chunk_header_options;
+
+ let mut documents = Vec::new();
+ for (i, text) in texts.iter().enumerate() {
+ let mut line_counter_index = 1;
+ let mut prev_chunk = None;
+ let mut index_prev_chunk = None;
+
+ for chunk in self.split_text(text) {
+ let mut page_content = chunk_header.clone();
+
+ let index_chunk = {
+ let idx = match index_prev_chunk {
+ Some(v) => v + 1,
+ None => 0,
+ };
+ text[idx..].find(&chunk).map(|i| i + idx).unwrap_or(0)
+ };
+ if prev_chunk.is_none() {
+ line_counter_index += self.number_of_newlines(text, 0, index_chunk);
+ } else {
+ let index_end_prev_chunk: usize = index_prev_chunk.unwrap_or_default()
+ + (self.length_function)(prev_chunk.as_deref().unwrap_or_default());
+
+ match index_end_prev_chunk.cmp(&index_chunk) {
+ Ordering::Less => {
+ line_counter_index +=
+ self.number_of_newlines(text, index_end_prev_chunk, index_chunk);
+ }
+ Ordering::Greater => {
+ let number =
+ self.number_of_newlines(text, index_chunk, index_end_prev_chunk);
+ line_counter_index = line_counter_index.saturating_sub(number);
+ }
+ Ordering::Equal => {}
+ }
+
+ if *append_chunk_overlap_header {
+ page_content += chunk_overlap_header;
+ }
+ }
+
+ let newlines_count = self.number_of_newlines(&chunk, 0, chunk.len());
+
+ let mut metadata = metadatas[i].clone();
+ metadata.insert(
+ "loc".into(),
+ format!(
+ "{}:{}",
+ line_counter_index,
+ line_counter_index + newlines_count
+ ),
+ );
+ page_content += &chunk;
+ documents.push(RagDocument {
+ page_content,
+ metadata,
+ });
+
+ line_counter_index += newlines_count;
+ prev_chunk = Some(chunk);
+ index_prev_chunk = Some(index_chunk);
+ }
+ }
+
+ documents
+ }
+
+ fn number_of_newlines(&self, text: &str, start: usize, end: usize) -> usize {
+ text[start..end].matches('\n').count()
+ }
+
+ pub fn split_text(&self, text: &str) -> Vec<String> {
+ let keep_separator = self
+ .separators
+ .iter()
+ .any(|v| v.chars().any(|v| !v.is_whitespace()));
+ self.split_text_impl(text, &self.separators, keep_separator)
+ }
+
+ fn split_text_impl(
+ &self,
+ text: &str,
+ separators: &[String],
+ keep_separator: bool,
+ ) -> Vec<String> {
+ let mut final_chunks = Vec::new();
+
+ let mut separator: String = separators.last().cloned().unwrap_or_default();
+ let mut new_separators: Vec<String> = vec![];
+ for (i, s) in separators.iter().enumerate() {
+ if s.is_empty() {
+ separator.clone_from(s);
+ break;
+ }
+ if text.contains(s) {
+ separator.clone_from(s);
+ new_separators = separators[i + 1..].to_vec();
+ break;
+ }
+ }
+
+ // Now that we have the separator, split the text
+ let splits = split_on_separator(text, &separator, keep_separator);
+
+ // Now go merging things, recursively splitting longer texts.
+ let mut good_splits = Vec::new();
+ let _separator = if keep_separator { "" } else { &separator };
+ for s in splits {
+ if (self.length_function)(s) < self.chunk_size {
+ good_splits.push(s.to_string());
+ } else {
+ if !good_splits.is_empty() {
+ let merged_text = self.merge_splits(&good_splits, _separator);
+ final_chunks.extend(merged_text);
+ good_splits.clear();
+ }
+ if new_separators.is_empty() {
+ final_chunks.push(s.to_string());
+ } else {
+ let other_info = self.split_text_impl(s, &new_separators, keep_separator);
+ final_chunks.extend(other_info);
+ }
+ }
+ }
+ if !good_splits.is_empty() {
+ let merged_text = self.merge_splits(&good_splits, _separator);
+ final_chunks.extend(merged_text);
+ }
+ final_chunks
+ }
+
+ fn merge_splits(&self, splits: &[String], separator: &str) -> Vec<String> {
+ let mut docs = Vec::new();
+ let mut current_doc = Vec::new();
+ let mut total = 0;
+ for d in splits {
+ let _len = (self.length_function)(d);
+ if total + _len + current_doc.len() * separator.len() > self.chunk_size {
+ if total > self.chunk_size {
+ // warn!("Warning: Created a chunk of size {}, which is longer than the specified {}", total, self.chunk_size);
+ }
+ if !current_doc.is_empty() {
+ let doc = self.join_docs(&current_doc, separator);
+ if let Some(doc) = doc {
+ docs.push(doc);
+ }
+ // Keep on popping if:
+ // - we have a larger chunk than in the chunk overlap
+ // - or if we still have any chunks and the length is long
+ while total > self.chunk_overlap
+ || (total + _len + current_doc.len() * separator.len() > self.chunk_size
+ && total > 0)
+ {
+ total -= (self.length_function)(&current_doc[0]);
+ current_doc.remove(0);
+ }
+ }
+ }
+ current_doc.push(d.to_string());
+ total += _len;
+ }
+ let doc = self.join_docs(&current_doc, separator);
+ if let Some(doc) = doc {
+ docs.push(doc);
+ }
+ docs
+ }
+
+ fn join_docs(&self, docs: &[String], separator: &str) -> Option<String> {
+ let text = docs.join(separator).trim().to_string();
+ if text.is_empty() {
+ None
+ } else {
+ Some(text)
+ }
+ }
+}
+
+pub struct SplitterChunkHeaderOptions {
+ pub chunk_header: String,
+ pub chunk_overlap_header: String,
+ pub append_chunk_overlap_header: bool,
+}
+
+impl Default for SplitterChunkHeaderOptions {
+ fn default() -> Self {
+ Self {
+ chunk_header: "".into(),
+ chunk_overlap_header: "(cont'd) ".into(),
+ append_chunk_overlap_header: false,
+ }
+ }
+}
+
+impl SplitterChunkHeaderOptions {
+ // Set the value of chunk_header
+ #[allow(unused)]
+ pub fn with_chunk_header(mut self, header: &str) -> Self {
+ self.chunk_header = header.to_string();
+ self
+ }
+
+ // Set the value of chunk_overlap_header
+ #[allow(unused)]
+ pub fn with_chunk_overlap_header(mut self, overlap_header: &str) -> Self {
+ self.chunk_overlap_header = overlap_header.to_string();
+ self
+ }
+
+ // Set the value of append_chunk_overlap_header
+ #[allow(unused)]
+ pub fn with_append_chunk_overlap_header(mut self, value: bool) -> Self {
+ self.append_chunk_overlap_header = value;
+ self
+ }
+}
+
+fn split_on_separator<'a>(text: &'a str, separator: &str, keep_separator: bool) -> Vec<&'a str> {
+ let splits: Vec<&str> = if !separator.is_empty() {
+ if keep_separator {
+ let mut splits = Vec::new();
+ let mut prev_idx = 0;
+ let sep_len = separator.len();
+
+ while let Some(idx) = text[prev_idx..].find(separator) {
+ splits.push(&text[prev_idx.saturating_sub(sep_len)..prev_idx + idx]);
+ prev_idx += idx + sep_len;
+ }
+
+ if prev_idx < text.len() {
+ splits.push(&text[prev_idx.saturating_sub(sep_len)..]);
+ }
+
+ splits
+ } else {
+ text.split(separator).collect()
+ }
+ } else {
+ text.split("").collect()
+ };
+ splits.into_iter().filter(|s| !s.is_empty()).collect()
+}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use indexmap::IndexMap;
+ use pretty_assertions::assert_eq;
+ use serde_json::{json, Value};
+
+ fn build_metadata(source: &str, loc_from_line: usize, loc_to_line: usize) -> Value {
+ json!({
+ "source": source,
+ "loc": format!("{loc_from_line}:{loc_to_line}"),
+ })
+ }
+ #[test]
+ fn test_split_text() {
+ let splitter = Splitter {
+ chunk_size: 7,
+ chunk_overlap: 3,
+ separators: vec![" ".into()],
+ ..Default::default()
+ };
+ let output = splitter.split_text("foo bar baz 123");
+ assert_eq!(output, vec!["foo bar", "bar baz", "baz 123"]);
+ }
+
+ #[test]
+ fn test_create_document() {
+ let splitter = Splitter::new(3, 0, &[" "]);
+ let chunk_header_options = SplitterChunkHeaderOptions::default();
+ let mut metadata1 = IndexMap::new();
+ metadata1.insert("source".into(), "1".into());
+ let mut metadata2 = IndexMap::new();
+ metadata2.insert("source".into(), "2".into());
+ let output = splitter.create_documents(
+ &["foo bar".into(), "baz".into()],
+ &[metadata1, metadata2],
+ &chunk_header_options,
+ );
+ let output = json!(output);
+ assert_eq!(
+ output,
+ json!([
+ {
+ "page_content": "foo",
+ "metadata": build_metadata("1", 1, 1),
+ },
+ {
+ "page_content": "bar",
+ "metadata": build_metadata("1", 1, 1),
+ },
+ {
+ "page_content": "baz",
+ "metadata": build_metadata("2", 1, 1),
+ },
+ ])
+ );
+ }
+
+ #[test]
+ fn test_chunk_header() {
+ let splitter = Splitter::new(3, 0, &[" "]);
+ let chunk_header_options = SplitterChunkHeaderOptions::default()
+ .with_chunk_header("SOURCE NAME: testing\n-----\n")
+ .with_append_chunk_overlap_header(true);
+ let mut metadata1 = IndexMap::new();
+ metadata1.insert("source".into(), "1".into());
+ let mut metadata2 = IndexMap::new();
+ metadata2.insert("source".into(), "2".into());
+ let output = splitter.create_documents(
+ &["foo bar".into(), "baz".into()],
+ &[metadata1, metadata2],
+ &chunk_header_options,
+ );
+ let output = json!(output);
+ assert_eq!(
+ output,
+ json!([
+ {
+ "page_content": "SOURCE NAME: testing\n-----\nfoo",
+ "metadata": build_metadata("1", 1, 1),
+ },
+ {
+ "page_content": "SOURCE NAME: testing\n-----\n(cont'd) bar",
+ "metadata": build_metadata("1", 1, 1),
+ },
+ {
+ "page_content": "SOURCE NAME: testing\n-----\nbaz",
+ "metadata": build_metadata("2", 1, 1),
+ },
+ ])
+ );
+ }
+
+ #[test]
+ fn test_markdown_splitter() {
+ let text = r#"# šŸ¦œļøšŸ”— LangChain
+
+⚔ Building applications with LLMs through composability ⚔
+
+## Quick Install
+
+```bash
+# Hopefully this code block isn't split
+pip install langchain
+```
+
+As an open source project in a rapidly developing field, we are extremely open to contributions."#;
+ let splitter = Splitter::new(100, 0, &MARKDOWN_SEPARATES);
+ let output = splitter.split_text(text);
+ let expected_output = vec![
+ "# šŸ¦œļøšŸ”— LangChain\n\n⚔ Building applications with LLMs through composability ⚔",
+ "## Quick Install\n\n```bash\n# Hopefully this code block isn't split\npip install langchain",
+ "```",
+ "As an open source project in a rapidly developing field, we are extremely open to contributions.",
+ ];
+ assert_eq!(output, expected_output);
+ }
+
+ #[test]
+ fn test_html_splitter() {
+ let text = r#"<!DOCTYPE html>
+<html>
+ <head>
+ <title>šŸ¦œļøšŸ”— LangChain</title>
+ <style>
+ body {
+ font-family: Arial, sans-serif;
+ }
+ h1 {
+ color: darkblue;
+ }
+ </style>
+ </head>
+ <body>
+ <div>
+ <h1>šŸ¦œļøšŸ”— LangChain</h1>
+ <p>⚔ Building applications with LLMs through composability ⚔</p>
+ </div>
+ <div>
+ As an open source project in a rapidly developing field, we are extremely open to contributions.
+ </div>
+ </body>
+</html>"#;
+ let splitter = Splitter::new(175, 20, &HTML_SEPARATES);
+ let output = splitter.split_text(text);
+ let expected_output = vec![
+ "<!DOCTYPE html>\n<html>",
+ "<head>\n <title>šŸ¦œļøšŸ”— LangChain</title>",
+ r#"<style>
+ body {
+ font-family: Arial, sans-serif;
+ }
+ h1 {
+ color: darkblue;
+ }
+ </style>
+ </head>"#,
+ r#"<body>
+ <div>
+ <h1>šŸ¦œļøšŸ”— LangChain</h1>
+ <p>⚔ Building applications with LLMs through composability ⚔</p>
+ </div>"#,
+ r#"<div>
+ As an open source project in a rapidly developing field, we are extremely open to contributions.
+ </div>
+ </body>
+</html>"#,
+ ];
+ assert_eq!(output, expected_output);
+ }
+}
diff --git a/src/render/stream.rs b/src/render/stream.rs
index f35831c..0c70bce 100644
--- a/src/render/stream.rs
+++ b/src/render/stream.rs
@@ -14,7 +14,7 @@ use std::{
time::Duration,
};
use textwrap::core::display_width;
-use tokio::sync::{mpsc::UnboundedReceiver, oneshot};
+use tokio::sync::mpsc::UnboundedReceiver;
pub async fn markdown_stream(
rx: UnboundedReceiver<SseEvent>,
@@ -62,17 +62,16 @@ async fn markdown_stream_inner(
let columns = terminal::size()?.0;
- let (spinner_tx, spinner_rx) = oneshot::channel();
- let mut spinner_tx = Some(spinner_tx);
- tokio::spawn(run_spinner(" Generating", spinner_rx));
+ let (stop_spinner_tx, _) = run_spinner("Generating").await;
+ let mut stop_spinner_tx = Some(stop_spinner_tx);
'outer: loop {
if abort.aborted() {
return Ok(());
}
for reply_event in gather_events(&mut rx).await {
- if let Some(spinner_tx) = spinner_tx.take() {
- let _ = spinner_tx.send(());
+ if let Some(stop_spinner_tx) = stop_spinner_tx.take() {
+ let _ = stop_spinner_tx.send(());
}
match reply_event {
@@ -150,8 +149,8 @@ async fn markdown_stream_inner(
}
}
- if let Some(spinner_tx) = spinner_tx.take() {
- let _ = spinner_tx.send(());
+ if let Some(stop_spinner_tx) = stop_spinner_tx.take() {
+ let _ = stop_spinner_tx.send(());
}
Ok(())
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index ba506dd..291bcd2 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -7,7 +7,7 @@ use self::highlighter::ReplHighlighter;
use self::prompt::ReplPrompt;
use crate::client::send_stream;
-use crate::config::{AssertState, GlobalConfig, Input, InputContext, StateFlags};
+use crate::config::{AssertState, Config, GlobalConfig, Input, InputContext, StateFlags};
use crate::function::need_send_call_results;
use crate::render::render_error;
use crate::utils::{create_abort_signal, set_text, AbortSignal};
@@ -33,7 +33,7 @@ lazy_static! {
const MENU_NAME: &str = "completion_menu";
lazy_static! {
- static ref REPL_COMMANDS: [ReplCommand; 16] = [
+ static ref REPL_COMMANDS: [ReplCommand; 19] = [
ReplCommand::new(".help", "Show this help message", AssertState::any()),
ReplCommand::new(".info", "View system info", AssertState::any()),
ReplCommand::new(".model", "Change the current LLM", AssertState::any()),
@@ -82,6 +82,17 @@ lazy_static! {
"End the current session",
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
),
+ ReplCommand::new(".rag", "Init or use a rag", AssertState::any()),
+ ReplCommand::new(
+ ".info rag",
+ "View rag info",
+ AssertState::True(StateFlags::RAG),
+ ),
+ ReplCommand::new(
+ ".exit rag",
+ "Leave the rag",
+ AssertState::True(StateFlags::RAG)
+ ),
ReplCommand::new(
".file",
"Include files with the message",
@@ -99,7 +110,7 @@ pub struct Repl {
config: GlobalConfig,
editor: Reedline,
prompt: ReplPrompt,
- abort: AbortSignal,
+ abort_signal: AbortSignal,
}
impl Repl {
@@ -114,7 +125,7 @@ impl Repl {
config: config.clone(),
editor,
prompt,
- abort,
+ abort_signal: abort,
})
}
@@ -122,13 +133,13 @@ impl Repl {
self.banner();
loop {
- if self.abort.aborted_ctrld() {
+ if self.abort_signal.aborted_ctrld() {
break;
}
let sig = self.editor.read_line(&self.prompt);
match sig {
Ok(Signal::Success(line)) => {
- self.abort.reset();
+ self.abort_signal.reset();
match self.handle(&line).await {
Ok(exit) => {
if exit {
@@ -142,11 +153,11 @@ impl Repl {
}
}
Ok(Signal::CtrlC) => {
- self.abort.set_ctrlc();
+ self.abort_signal.set_ctrlc();
println!("(To exit, press Ctrl+D or enter \".exit\")\n");
}
Ok(Signal::CtrlD) => {
- self.abort.set_ctrld();
+ self.abort_signal.set_ctrld();
break;
}
_ => {}
@@ -176,6 +187,10 @@ impl Repl {
let info = self.config.read().session_info()?;
println!("{}", info);
}
+ Some("rag") => {
+ let info = self.config.read().rag_info()?;
+ println!("{}", info);
+ }
Some(_) => unknown_command()?,
None => {
let output = self.config.read().system_info()?;
@@ -193,7 +208,7 @@ impl Repl {
},
".prompt" => match args {
Some(text) => {
- self.config.write().set_prompt(text)?;
+ self.config.write().use_prompt(text)?;
}
None => println!("Usage: .prompt <text>..."),
},
@@ -206,16 +221,19 @@ impl Repl {
text.trim(),
Some(InputContext::role(role)),
);
- ask(&self.config, self.abort.clone(), input).await?;
+ ask(&self.config, self.abort_signal.clone(), input).await?;
}
None => {
- self.config.write().set_role(args)?;
+ self.config.write().use_role(args)?;
}
},
None => println!(r#"Usage: .role <name> [text]..."#),
},
".session" => {
- self.config.write().start_session(args)?;
+ self.config.write().use_session(args)?;
+ }
+ ".rag" => {
+ Config::use_rag(&self.config, args, self.abort_signal.clone()).await?;
}
".save" => {
match args.map(|v| match v.split_once(' ') {
@@ -248,16 +266,19 @@ impl Repl {
let (files, text) = split_files_text(args);
let files = shell_words::split(files).with_context(|| "Invalid args")?;
let input = Input::new(&self.config, text, files, None)?;
- ask(&self.config, self.abort.clone(), input).await?;
+ ask(&self.config, self.abort_signal.clone(), input).await?;
}
None => println!("Usage: .file <files>... [-- <text>...]"),
},
".exit" => match args {
Some("role") => {
- self.config.write().clear_role()?;
+ self.config.write().exit_role()?;
}
Some("session") => {
- self.config.write().end_session()?;
+ self.config.write().exit_session()?;
+ }
+ Some("rag") => {
+ self.config.write().exit_rag()?;
}
Some(_) => unknown_command()?,
None => {
@@ -273,8 +294,9 @@ impl Repl {
_ => unknown_command()?,
},
None => {
- let input = Input::from_str(&self.config, line, None);
- ask(&self.config, self.abort.clone(), input).await?;
+ let mut input = Input::from_str(&self.config, line, None);
+ input.maybe_embeddings(self.abort_signal.clone()).await?;
+ ask(&self.config, self.abort_signal.clone(), input).await?;
}
}
@@ -407,7 +429,7 @@ impl Validator for ReplValidator {
}
#[async_recursion]
-async fn ask(config: &GlobalConfig, abort: AbortSignal, input: Input) -> Result<()> {
+async fn ask(config: &GlobalConfig, abort: AbortSignal, mut input: Input) -> Result<()> {
if input.is_empty() {
return Ok(());
}
@@ -417,9 +439,10 @@ async fn ask(config: &GlobalConfig, abort: AbortSignal, input: Input) -> Result<
let client = input.create_client()?;
let (output, tool_call_results) =
send_stream(&input, client.as_ref(), config, abort.clone()).await?;
+
config
.write()
- .save_message(&input, &output, &tool_call_results)?;
+ .save_message(&mut input, &output, &tool_call_results)?;
config.read().maybe_copy(&output);
if config.write().should_compress_session() {
let config = config.clone();
diff --git a/src/serve.rs b/src/serve.rs
index 2d375f8..116942a 100644
--- a/src/serve.rs
+++ b/src/serve.rs
@@ -1,11 +1,4 @@
-use crate::{
- client::{
- init_client, list_models, ChatCompletionsData, ChatCompletionsOutput, ClientConfig,
- Message, Model, ModelData, SseEvent, SseHandler,
- },
- config::{Config, GlobalConfig, Role},
- utils::create_abort_signal,
-};
+use crate::{client::*, config::*, utils::*};
use anyhow::{anyhow, bail, Result};
use bytes::Bytes;
@@ -76,7 +69,7 @@ impl Server {
let clients = config.clients.clone();
let model = config.model.clone();
let roles = config.roles.clone();
- let mut models = list_models(&config);
+ let mut models = list_chat_models(&config);
let mut default_model = model.clone();
default_model.data_mut().name = DEFAULT_MODEL_NAME.into();
models.insert(0, &default_model);
diff --git a/src/utils/abort_signal.rs b/src/utils/abort_signal.rs
index af58b35..ac93653 100644
--- a/src/utils/abort_signal.rs
+++ b/src/utils/abort_signal.rs
@@ -53,3 +53,12 @@ impl AbortSignalInner {
self.ctrld.store(true, Ordering::SeqCst);
}
}
+
+pub async fn watch_abort_signal(abort: AbortSignal) {
+ loop {
+ if abort.aborted() {
+ break;
+ }
+ tokio::time::sleep(std::time::Duration::from_millis(100)).await;
+ }
+}
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index fa67c63..95c6725 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -6,7 +6,7 @@ mod prompt_input;
mod render_prompt;
mod spinner;
-pub use self::abort_signal::{create_abort_signal, AbortSignal};
+pub use self::abort_signal::*;
pub use self::clipboard::set_text;
pub use self::command::*;
pub use self::crypto::*;
diff --git a/src/utils/spinner.rs b/src/utils/spinner.rs
index 0dd4d01..c746469 100644
--- a/src/utils/spinner.rs
+++ b/src/utils/spinner.rs
@@ -4,7 +4,10 @@ use std::{
io::{stdout, Stdout, Write},
time::Duration,
};
-use tokio::{sync::oneshot, time::interval};
+use tokio::{
+ sync::{mpsc, oneshot},
+ time::interval,
+};
pub struct Spinner {
index: usize,
@@ -23,6 +26,10 @@ impl Spinner {
}
}
+ pub fn set_message(&mut self, message: &str) {
+ self.message = format!(" {message}");
+ }
+
pub fn step(&mut self, writer: &mut Stdout) -> Result<()> {
if self.stopped {
return Ok(());
@@ -55,18 +62,38 @@ impl Spinner {
}
}
-pub async fn run_spinner(message: &str, rx: oneshot::Receiver<()>) -> Result<()> {
+pub async fn run_spinner(message: &str) -> (oneshot::Sender<()>, mpsc::UnboundedSender<String>) {
+ let message = format!(" {message}");
+ let (stop_tx, stop_rx) = oneshot::channel();
+ let (message_tx, message_rx) = mpsc::unbounded_channel();
+ tokio::spawn(run_spinner_inner(message, stop_rx, message_rx));
+ (stop_tx, message_tx)
+}
+
+async fn run_spinner_inner(
+ message: String,
+ stop_rx: oneshot::Receiver<()>,
+ mut message_rx: mpsc::UnboundedReceiver<String>,
+) -> Result<()> {
let mut writer = stdout();
- let mut spinner = Spinner::new(message);
+ let mut spinner = Spinner::new(&message);
let mut interval = interval(Duration::from_millis(50));
tokio::select! {
_ = async {
loop {
- interval.tick().await;
- let _ = spinner.step(&mut writer);
+ tokio::select! {
+ _ = interval.tick() => {
+ let _ = spinner.step(&mut writer);
+ }
+ message = message_rx.recv() => {
+ if let Some(message) = message {
+ spinner.set_message(&message);
+ }
+ }
+ }
}
} => {}
- _ = rx => {
+ _ = stop_rx => {
spinner.stop(&mut writer)?;
}
}