summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/model.rs11
-rw-r--r--src/config/bot.rs2
-rw-r--r--src/config/mod.rs30
-rw-r--r--src/config/session.rs2
-rw-r--r--src/rag/mod.rs110
5 files changed, 113 insertions, 42 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index d555232..56421bf 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,5 +1,5 @@
use super::{
- list_chat_models,
+ list_chat_models, list_embedding_models,
message::{Message, MessageContent},
EmbeddingsData,
};
@@ -43,13 +43,20 @@ impl Model {
.collect()
}
- pub fn retrieve(config: &Config, model_id: &str) -> Result<Self> {
+ pub fn retrieve_chat(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_chat_models(config), model_id) {
Some(v) => Ok(v),
None => bail!("Invalid model '{model_id}'"),
}
}
+ pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result<Self> {
+ match Self::find(&list_embedding_models(config), model_id) {
+ Some(v) => Ok(v),
+ None => bail!("Invalid model '{model_id}'"),
+ }
+ }
+
pub fn find(models: &[&Self], model_id: &str) -> Option<Self> {
let mut model = None;
let (client_name, model_name) = match model_id.split_once(':') {
diff --git a/src/config/bot.rs b/src/config/bot.rs
index e02cce7..098c8d7 100644
--- a/src/config/bot.rs
+++ b/src/config/bot.rs
@@ -49,7 +49,7 @@ impl Bot {
let model = {
let config = config.read();
match bot_config.model_id.as_ref() {
- Some(model_id) => Model::retrieve(&config, model_id)?,
+ Some(model_id) => Model::retrieve_chat(&config, model_id)?,
None => config.current_model().clone(),
}
};
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 571bf2c..a905ae8 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -61,13 +61,19 @@ 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 RAG_TEMPLATE: &str = r#"Answer the question based only on the provided context:
-<context>
-__CONTEXT__
-</context>
+const RAG_TEMPLATE: &str = r#"Use the following context as your learned knowledge, inside <context></context> XML tags.
+ <context>
+ __CONTEXT__
+ </context>
-Question: __INPUT__
-"#;
+ When answer to user:
+ - If you don't know, just say that you don't know.
+ - If you don't know when you are not sure, ask for clarification.
+ Avoid mentioning that you obtained the information from the context.
+ And answer according to the language of the user's question.
+
+ Given the context information, answer the query.
+ Query: __INPUT__"#;
const LEFT_PROMPT: &str = "{color.green}{?session {?bot {bot}#}{session}{?role /}}{!session {?bot {bot}}}{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}";
@@ -96,7 +102,9 @@ pub struct Config {
pub dangerously_functions_filter: Option<FunctionsFilter>,
pub bot_prelude: Option<String>,
pub bots: Vec<BotConfig>,
- pub embedding_model: Option<String>,
+ pub rag_embedding_model: Option<String>,
+ pub rag_chunk_size: Option<usize>,
+ pub rag_chunk_overlap: Option<usize>,
pub rag_top_k: usize,
pub rag_template: Option<String>,
pub compress_threshold: usize,
@@ -147,7 +155,9 @@ impl Default for Config {
dangerously_functions_filter: None,
bot_prelude: None,
bots: vec![],
- embedding_model: None,
+ rag_embedding_model: None,
+ rag_chunk_size: None,
+ rag_chunk_overlap: None,
rag_top_k: 4,
rag_template: None,
compress_threshold: 4000,
@@ -616,7 +626,7 @@ impl Config {
}
pub fn set_model(&mut self, model_id: &str) -> Result<()> {
- let model = Model::retrieve(self, model_id)?;
+ let model = Model::retrieve_chat(self, model_id)?;
match self.role_like_mut() {
Some(role_like) => role_like.set_model(&model),
None => {
@@ -682,7 +692,7 @@ impl Config {
match role.model_id() {
Some(model_id) => {
if self.model.id() != model_id {
- let model = Model::retrieve(self, model_id)?;
+ let model = Model::retrieve_chat(self, model_id)?;
role.set_model(&model);
}
}
diff --git a/src/config/session.rs b/src/config/session.rs
index 38855cc..7400f30 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -73,7 +73,7 @@ impl Session {
let mut session: Self =
serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?;
- session.model = Model::retrieve(config, &session.model_id)?;
+ session.model = Model::retrieve_chat(config, &session.model_id)?;
session.name = name.to_string();
session.path = Some(path.display().to_string());
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index b968232..b7ac6cd 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -50,13 +50,8 @@ impl Rag {
doc_paths: &[String],
abort_signal: AbortSignal,
) -> Result<Self> {
- if !*IS_STDOUT_TERMINAL {
- bail!("An interactive shell is required to initialize rag.")
- }
debug!("init rag: {name}");
- let model = select_embedding_model(config)?;
- let chunk_size = set_chunk_size(&model)?;
- let chunk_overlap = chunk_size / 20;
+ let (model, chunk_size, chunk_overlap) = Self::config(config)?;
let data = RagData::new(&model.id(), chunk_size, chunk_overlap);
let mut rag = Self::create(config, name, save_path, data)?;
let mut paths = doc_paths.to_vec();
@@ -92,7 +87,7 @@ impl Rag {
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 model = Model::retrieve_embedding(&config.read(), &data.model)?;
let client = init_client(config, Some(model.clone()))?;
let rag = Rag {
client,
@@ -105,6 +100,68 @@ impl Rag {
Ok(rag)
}
+ pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> {
+ let (embedding_model, chunk_size, chunk_overlap) = {
+ let config = config.read();
+ (
+ config.rag_embedding_model.clone(),
+ config.rag_chunk_size,
+ config.rag_chunk_overlap,
+ )
+ };
+ let model_id = match embedding_model {
+ Some(value) => {
+ println!("Select embedding model: {value}");
+ value
+ }
+ None => {
+ let models = list_embedding_models(&config.read());
+ if models.is_empty() {
+ bail!("No available embedding model");
+ }
+ if *IS_STDOUT_TERMINAL {
+ select_embedding_model(&models)?
+ } else {
+ let value = models[0].id();
+ println!("Select embedding model: {value}");
+ value
+ }
+ }
+ };
+ let model = Model::retrieve_embedding(&config.read(), &model_id)?;
+ let chunk_size = match chunk_size {
+ Some(value) => {
+ println!("Set chunk size: {value}");
+ value
+ }
+ None => {
+ if *IS_STDOUT_TERMINAL {
+ set_chunk_size(&model)?
+ } else {
+ let value = model.default_chunk_size();
+ println!("Set chunk size: {value}");
+ value
+ }
+ }
+ };
+ let chunk_overlap = match chunk_overlap {
+ Some(value) => {
+ println!("Set chunk overlay: {value}");
+ value
+ }
+ None => {
+ let value = chunk_size / 20;
+ if *IS_STDOUT_TERMINAL {
+ set_chunk_overlay(value)?
+ } else {
+ println!("Set chunk overlay: {value}");
+ value
+ }
+ }
+ };
+ Ok((model, chunk_size, chunk_overlap))
+ }
+
pub fn save(&self, path: &Path) -> Result<()> {
ensure_parent_exists(path)?;
let mut file = std::fs::File::create(path)?;
@@ -392,27 +449,10 @@ fn document_text(file_path: &str, document: &RagDocument) -> String {
)
}
-fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result<Model> {
- let model = Model::find(&list_embedding_models(config), 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 select_embedding_model(models: &[&Model]) -> Result<String> {
+ let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect();
+ let model_id = Select::new("Select embedding model:", model_ids).prompt()?;
+ Ok(model_id)
}
fn set_chunk_size(model: &Model) -> Result<usize> {
@@ -437,6 +477,20 @@ fn set_chunk_size(model: &Model) -> Result<usize> {
value.parse().map_err(|_| anyhow!("Invalid chunk_size"))
}
+fn set_chunk_overlay(default_value: usize) -> Result<usize> {
+ let value = Text::new("Set chunk overlay:")
+ .with_default(&default_value.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_overlay"))
+}
+
fn add_doc_paths() -> Result<Vec<String>> {
let text = Text::new("Add document paths:")
.with_validator(required!("This field is required"))