From a56d5f2ddfcfb290c308f2b1a3464c3bebc340de Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 8 Sep 2024 19:50:02 +0800 Subject: feat: save rag in YAML instead of bin (#848) --- src/rag/mod.rs | 23 +++++++++++++---------- 1 file changed, 13 insertions(+), 10 deletions(-) (limited to 'src/rag/mod.rs') 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 { 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, pub files: IndexMap, + #[serde(with = "serde_vectors")] pub vectors: IndexMap>, } -- cgit v1.2.3