summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/config/mod.rs104
-rw-r--r--src/rag/mod.rs86
-rw-r--r--src/repl/mod.rs2
3 files changed, 141 insertions, 51 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index a5b5c5b..d2d7c2f 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -511,6 +511,10 @@ impl Config {
.wrap
.clone()
.map_or_else(|| String::from("no"), |v| v.to_string());
+ let (rag_reranker_model, rag_top_k) = match self.rag.as_ref() {
+ Some(rag) => rag.get_config(),
+ None => (self.rag_reranker_model.clone(), self.rag_top_k),
+ };
let role = self.extract_role();
let mut items = vec![
("model", role.model().id()),
@@ -535,9 +539,9 @@ impl Config {
("use_tools", format_option_value(&role.use_tools())),
(
"rag_reranker_model",
- format_option_value(&self.rag_reranker_model),
+ format_option_value(&rag_reranker_model),
),
- ("rag_top_k", self.rag_top_k.to_string()),
+ ("rag_top_k", rag_top_k.to_string()),
("highlight", self.highlight.to_string()),
("light_theme", self.light_theme.to_string()),
("config_file", display_path(&Self::config_file()?)),
@@ -559,7 +563,7 @@ impl Config {
Ok(output)
}
- pub fn update(&mut self, data: &str) -> Result<()> {
+ pub fn update(config: &GlobalConfig, data: &str) -> Result<()> {
let parts: Vec<&str> = data.split_whitespace().collect();
if parts.len() != 2 {
bail!("Usage: .set <key> <value>. If value is null, unset key.");
@@ -569,62 +573,58 @@ impl Config {
match key {
"max_output_tokens" => {
let value = parse_value(value)?;
- self.set_max_output_tokens(value);
+ config.write().set_max_output_tokens(value);
}
"temperature" => {
let value = parse_value(value)?;
- self.set_temperature(value);
+ config.write().set_temperature(value);
}
"top_p" => {
let value = parse_value(value)?;
- self.set_top_p(value);
+ config.write().set_top_p(value);
}
"dry_run" => {
let value = value.parse().with_context(|| "Invalid value")?;
- self.dry_run = value;
+ config.write().dry_run = value;
}
"stream" => {
let value = value.parse().with_context(|| "Invalid value")?;
- self.stream = value;
+ config.write().stream = value;
}
"save" => {
let value = value.parse().with_context(|| "Invalid value")?;
- self.save = value;
+ config.write().save = value;
}
"rag_reranker_model" => {
- self.rag_reranker_model = if value == "null" {
- None
- } else {
- Some(value.to_string())
- }
+ let value = parse_value(value)?;
+ Self::set_rag_reranker_model(config, value)?;
}
"rag_top_k" => {
- if let Some(value) = parse_value(value)? {
- self.rag_top_k = value;
- }
+ let value = value.parse().with_context(|| "Invalid value")?;
+ Self::set_rag_top_k(config, value)?;
}
"function_calling" => {
let value = value.parse().with_context(|| "Invalid value")?;
- if value && self.functions.is_empty() {
+ if value && config.write().functions.is_empty() {
bail!("Function calling cannot be enabled because no functions are installed.")
}
- self.function_calling = value;
+ config.write().function_calling = value;
}
"use_tools" => {
let value = parse_value(value)?;
- self.set_use_tools(value);
+ config.write().set_use_tools(value);
}
"save_session" => {
let value = parse_value(value)?;
- self.set_save_session(value);
+ config.write().set_save_session(value);
}
"compress_threshold" => {
let value = parse_value(value)?;
- self.set_compress_threshold(value);
+ config.write().set_compress_threshold(value);
}
"highlight" => {
let value = value.parse().with_context(|| "Invalid value")?;
- self.highlight = value;
+ config.write().highlight = value;
}
_ => bail!("Unknown key `{key}`"),
}
@@ -668,6 +668,33 @@ impl Config {
}
}
+ pub fn set_rag_reranker_model(config: &GlobalConfig, value: Option<String>) -> Result<()> {
+ if let Some(id) = &value {
+ Model::retrieve_reranker(&config.read(), id)?;
+ }
+ let has_rag = config.read().rag.is_some();
+ match has_rag {
+ true => update_rag(config, |rag| {
+ rag.set_reranker_model(value)?;
+ Ok(())
+ })?,
+ false => config.write().rag_reranker_model = value,
+ }
+ Ok(())
+ }
+
+ pub fn set_rag_top_k(config: &GlobalConfig, value: usize) -> Result<()> {
+ let has_rag = config.read().rag.is_some();
+ match has_rag {
+ true => update_rag(config, |rag| {
+ rag.set_top_k(value)?;
+ Ok(())
+ })?,
+ false => config.write().rag_top_k = value,
+ }
+ Ok(())
+ }
+
pub fn set_wrap(&mut self, value: &str) -> Result<()> {
if value == "no" {
self.wrap = None;
@@ -1090,13 +1117,11 @@ impl Config {
}
pub async fn rebuild_rag(config: &GlobalConfig, abort_signal: AbortSignal) -> Result<()> {
- let rag_name = match config.read().rag.clone() {
- Some(v) => v.name().to_string(),
+ let mut rag = match config.read().rag.clone() {
+ Some(v) => v.as_ref().clone(),
None => bail!("No RAG"),
};
- let rag_path = config.read().rag_file(&rag_name)?;
- let mut rag = Rag::load(config, &rag_name, &rag_path)?;
- rag.rebuild(config, &rag_path, abort_signal).await?;
+ rag.rebuild(config, abort_signal).await?;
config.write().rag = Some(Arc::new(rag));
Ok(())
}
@@ -1120,20 +1145,20 @@ impl Config {
text: &str,
abort_signal: AbortSignal,
) -> Result<String> {
- let (top_k, min_score_vector_search, min_score_keyword_search) = {
+ let (reranker_model, top_k) = rag.get_config();
+ let (min_score_vector_search, min_score_keyword_search, rag_min_score_rerank) = {
let config = config.read();
(
- config.rag_top_k,
config.rag_min_score_vector_search,
config.rag_min_score_keyword_search,
+ config.rag_min_score_rerank,
)
};
- let rerank = match config.read().rag_reranker_model.clone() {
+ let rerank = match reranker_model {
Some(reranker_model_id) => {
- let min_score = config.read().rag_min_score_rerank;
let rerank_model = Model::retrieve_reranker(&config.read(), &reranker_model_id)?;
let rerank_client = init_client(config, Some(rerank_model))?;
- Some((rerank_client, min_score))
+ Some((rerank_client, rag_min_score_rerank))
}
None => None,
};
@@ -2056,3 +2081,16 @@ fn complete_option_bool(value: Option<bool>) -> Vec<String> {
None => vec!["true".to_string(), "false".to_string()],
}
}
+
+fn update_rag<F>(config: &GlobalConfig, f: F) -> Result<()>
+where
+ F: FnOnce(&mut Rag) -> Result<()>,
+{
+ let mut rag = match config.read().rag.clone() {
+ Some(v) => v.as_ref().clone(),
+ None => bail!("No RAG"),
+ };
+ f(&mut rag)?;
+ config.write().rag = Some(Arc::new(rag));
+ Ok(())
+}
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index aa1911e..f6a511d 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -22,13 +22,13 @@ use std::collections::HashMap;
use std::{fmt::Debug, io::BufReader, path::Path};
pub struct Rag {
+ config: GlobalConfig,
name: String,
path: String,
embedding_model: Model,
hnsw: Hnsw<'static, f32, DistCosine>,
bm25: BM25<DocumentId>,
data: RagData,
- embedding_client: Box<dyn Client>,
}
impl Debug for Rag {
@@ -42,6 +42,20 @@ impl Debug for Rag {
}
}
+impl Clone for Rag {
+ fn clone(&self) -> Self {
+ Self {
+ config: self.config.clone(),
+ name: self.name.clone(),
+ path: self.path.clone(),
+ embedding_model: self.embedding_model.clone(),
+ hnsw: self.data.build_hnsw(),
+ bm25: self.bm25.clone(),
+ data: self.data.clone(),
+ }
+ }
+}
+
impl Rag {
pub async fn init(
config: &GlobalConfig,
@@ -51,8 +65,18 @@ impl Rag {
abort_signal: AbortSignal,
) -> Result<Self> {
debug!("init rag: {name}");
- let (embedding_model, chunk_size, chunk_overlap) = Self::config(config)?;
- let data = RagData::new(embedding_model.id(), chunk_size, chunk_overlap);
+ let (embedding_model, chunk_size, chunk_overlap) = Self::create_config(config)?;
+ let (reranker_model, top_k) = {
+ let config = config.read();
+ (config.rag_reranker_model.clone(), config.rag_top_k)
+ };
+ let data = RagData::new(
+ embedding_model.id(),
+ chunk_size,
+ chunk_overlap,
+ reranker_model,
+ top_k,
+ );
let mut rag = Self::create(config, name, save_path, data)?;
let mut paths = doc_paths.to_vec();
if paths.is_empty() {
@@ -71,8 +95,7 @@ impl Rag {
bail!("Aborted!")
},
};
- if !rag.is_temp() {
- rag.save(save_path)?;
+ if rag.save()? {
println!("✨ Saved rag to '{}'", save_path.display());
}
Ok(rag)
@@ -90,15 +113,14 @@ impl Rag {
let hnsw = data.build_hnsw();
let bm25 = data.build_bm25();
let embedding_model = Model::retrieve_embedding(&config.read(), &data.embedding_model)?;
- let embedding_client = init_client(config, Some(embedding_model.clone()))?;
let rag = Rag {
+ config: config.clone(),
name: name.to_string(),
path: path.display().to_string(),
data,
embedding_model,
hnsw,
bm25,
- embedding_client,
};
Ok(rag)
}
@@ -106,7 +128,6 @@ impl Rag {
pub async fn rebuild(
&mut self,
config: &GlobalConfig,
- save_path: &Path,
abort_signal: AbortSignal,
) -> Result<()> {
debug!("rebuild rag: {}", self.name);
@@ -123,14 +144,13 @@ impl Rag {
bail!("Aborted!")
},
};
- if !self.is_temp() {
- self.save(save_path)?;
- println!("✨ Saved rag to '{}'", save_path.display());
+ if self.save()? {
+ println!("✨ Saved rag to '{}'", self.path);
}
Ok(())
}
- pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> {
+ pub fn create_config(config: &GlobalConfig) -> Result<(Model, usize, usize)> {
let (embedding_model_id, chunk_size, chunk_overlap) = {
let config = config.read();
(
@@ -194,12 +214,32 @@ impl Rag {
Ok((embedding_model, chunk_size, chunk_overlap))
}
- pub fn save(&self, path: &Path) -> Result<()> {
+ pub fn get_config(&self) -> (Option<String>, usize) {
+ (self.data.reranker_model.clone(), self.data.top_k)
+ }
+
+ pub fn set_reranker_model(&mut self, reranker_model: Option<String>) -> Result<()> {
+ self.data.reranker_model = reranker_model;
+ self.save()?;
+ Ok(())
+ }
+
+ pub fn set_top_k(&mut self, top_k: usize) -> Result<()> {
+ self.data.top_k = top_k;
+ self.save()?;
+ Ok(())
+ }
+
+ pub fn save(&self) -> Result<bool> {
+ if self.is_temp() {
+ return Ok(false);
+ }
+ let path = Path::new(&self.path);
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(())
+ Ok(true)
}
pub fn export(&self) -> Result<String> {
@@ -219,6 +259,8 @@ impl Rag {
"embedding_model": self.embedding_model.id(),
"chunk_size": self.data.chunk_size,
"chunk_overlap": self.data.chunk_overlap,
+ "reranker_model": self.data.reranker_model,
+ "top_k": self.data.top_k,
"document_paths": self.data.document_paths,
"files": files,
});
@@ -490,6 +532,7 @@ impl Rag {
data: EmbeddingsData,
spinner: Option<Spinner>,
) -> Result<EmbeddingsOutput> {
+ let embedding_client = init_client(&self.config, Some(self.embedding_model.clone()))?;
let EmbeddingsData { texts, query } = data;
let size = match self.embedding_model.max_input_tokens() {
Some(max_input_tokens) => {
@@ -513,8 +556,7 @@ impl Rag {
texts: texts.to_vec(),
query,
};
- let chunk_output = self
- .embedding_client
+ let chunk_output = embedding_client
.embeddings(chunk_data)
.await
.context("Failed to create embedding")?;
@@ -529,6 +571,8 @@ pub struct RagData {
pub embedding_model: String,
pub chunk_size: usize,
pub chunk_overlap: usize,
+ pub reranker_model: Option<String>,
+ pub top_k: usize,
pub next_file_id: FileId,
pub document_paths: Vec<String>,
pub files: IndexMap<FileId, RagFile>,
@@ -549,11 +593,19 @@ impl Debug for RagData {
}
impl RagData {
- pub fn new(embedding_model: String, chunk_size: usize, chunk_overlap: usize) -> Self {
+ pub fn new(
+ embedding_model: String,
+ chunk_size: usize,
+ chunk_overlap: usize,
+ reranker_model: Option<String>,
+ top_k: usize,
+ ) -> Self {
Self {
embedding_model,
chunk_size,
chunk_overlap,
+ reranker_model,
+ top_k,
next_file_id: 0,
document_paths: Default::default(),
files: Default::default(),
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index eec43e8..b6c71a9 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -386,7 +386,7 @@ impl Repl {
}
".set" => match args {
Some(args) => {
- self.config.write().update(args)?;
+ Config::update(&self.config, args)?;
}
_ => {
println!("Usage: .set <key> <value>...")