summaryrefslogtreecommitdiffstats
path: root/src/rag/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-08 19:50:02 +0800
committerGitHub <noreply@github.com>2024-09-08 19:50:02 +0800
commita56d5f2ddfcfb290c308f2b1a3464c3bebc340de (patch)
tree96d7bbb0c45c2f1edead44f38524039ddf92cb87 /src/rag/mod.rs
parent89554e0d4ebcc66442456f07403be2a70a50386b (diff)
downloadaichat-a56d5f2ddfcfb290c308f2b1a3464c3bebc340de.tar.gz
feat: save rag in YAML instead of bin (#848)
Diffstat (limited to 'src/rag/mod.rs')
-rw-r--r--src/rag/mod.rs23
1 files changed, 13 insertions, 10 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index f6a511d..e03b3f1 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -8,18 +8,17 @@ use crate::utils::*;
mod bm25;
mod loader;
+mod serde_vectors;
mod splitter;
-use anyhow::bail;
-use anyhow::{anyhow, Context, Result};
+use anyhow::{anyhow, bail, Context, Result};
use hnsw_rs::prelude::*;
use indexmap::{IndexMap, IndexSet};
use inquire::{required, validator::Validation, Confirm, Select, Text};
use path_absolutize::Absolutize;
use serde::{Deserialize, Serialize};
use serde_json::json;
-use std::collections::HashMap;
-use std::{fmt::Debug, io::BufReader, path::Path};
+use std::{collections::HashMap, fmt::Debug, fs, path::Path};
pub struct Rag {
config: GlobalConfig,
@@ -103,9 +102,8 @@ impl Rag {
pub fn load(config: &GlobalConfig, name: &str, path: &Path) -> Result<Self> {
let err = || format!("Failed to load rag '{name}' at '{}'", path.display());
- 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)?;
+ let content = fs::read_to_string(path).with_context(err)?;
+ let data: RagData = serde_yaml::from_str(&content).with_context(err)?;
Self::create(config, name, path, data)
}
@@ -236,9 +234,13 @@ impl Rag {
}
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))?;
+
+ let content = serde_yaml::to_string(&self.data)
+ .with_context(|| format!("Failed to serde rag '{}'", self.name))?;
+ fs::write(path, content).with_context(|| {
+ format!("Failed to save rag '{}' to '{}'", self.name, path.display())
+ })?;
+
Ok(true)
}
@@ -576,6 +578,7 @@ pub struct RagData {
pub next_file_id: FileId,
pub document_paths: Vec<String>,
pub files: IndexMap<FileId, RagFile>,
+ #[serde(with = "serde_vectors")]
pub vectors: IndexMap<DocumentId, Vec<f32>>,
}