use std::collections::HashMap;
use anyhow::Result;
use indicatif::{ProgressBar, ProgressStyle};
use velesdb_core::{Point, VectorCollection};
pub struct BatchImporter<'a> {
collection: &'a VectorCollection,
batch: Vec<Point>,
batch_size: usize,
pub stats: ImportAccumulator,
}
#[derive(Debug, Default)]
pub struct ImportAccumulator {
pub imported: usize,
pub errors: usize,
}
impl<'a> BatchImporter<'a> {
pub fn new(collection: &'a VectorCollection, batch_size: usize) -> Self {
Self {
collection,
batch: Vec::with_capacity(batch_size),
batch_size,
stats: ImportAccumulator::default(),
}
}
pub fn push(&mut self, point: Point) -> Result<()> {
self.batch.push(point);
self.stats.imported += 1;
if self.batch.len() >= self.batch_size {
self.collection.upsert_bulk(&self.batch)?;
self.batch.clear();
}
Ok(())
}
pub fn record_error(&mut self) {
self.stats.errors += 1;
}
pub fn flush(self) -> Result<ImportAccumulator> {
if !self.batch.is_empty() {
self.collection.upsert_bulk(&self.batch)?;
}
Ok(self.stats)
}
}
#[must_use]
pub fn create_progress_bar(total: usize, show: bool) -> ProgressBar {
if show {
let pb = ProgressBar::new(total as u64);
if let Ok(style) = ProgressStyle::default_bar().template(
"{spinner:.green} [{elapsed_precise}] [{bar:40.cyan/blue}] {pos}/{len} ({eta})",
) {
pb.set_style(style.progress_chars("#>-"));
}
pb
} else {
ProgressBar::hidden()
}
}
pub fn set_import_message(progress: &ProgressBar, total: usize, file_size: u64, show: bool) {
if show {
#[allow(clippy::cast_precision_loss)]
let size_mb = file_size as f64 / (1024.0 * 1024.0);
progress.set_message(format!("Importing {total} vectors ({size_mb:.1} MB)"));
}
}
pub fn point_payload_to_row(
id: u64,
payload: &Option<serde_json::Value>,
) -> HashMap<String, serde_json::Value> {
let mut row = HashMap::new();
row.insert("id".to_string(), serde_json::json!(id));
if let Some(serde_json::Value::Object(map)) = payload {
for (k, v) in map {
row.insert(k.clone(), v.clone());
}
}
row
}
pub fn point_payload_to_browse_row(
id: u64,
payload: &Option<serde_json::Value>,
) -> HashMap<String, serde_json::Value> {
let mut row = HashMap::new();
row.insert("id".to_string(), serde_json::json!(id));
if let Some(serde_json::Value::Object(map)) = payload {
for (k, v) in map {
row.insert(k.clone(), truncate_display_value(v));
}
}
row
}
fn truncate_display_value(value: &serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::String(s) if s.len() > 50 => {
let truncated: String = s.chars().take(47).collect();
serde_json::json!(format!("{truncated}..."))
}
other => other.clone(),
}
}
pub fn print_json(data: &serde_json::Value) -> Result<()> {
println!("{}", serde_json::to_string_pretty(data)?);
Ok(())
}
pub fn point_to_export_record(
id: u64,
vector: Option<&[f32]>,
payload: &Option<serde_json::Value>,
) -> serde_json::Value {
let mut record = serde_json::Map::new();
record.insert("id".to_string(), serde_json::json!(id));
if let Some(v) = vector {
record.insert("vector".to_string(), serde_json::json!(v));
}
if let Some(p) = payload {
record.insert("payload".to_string(), p.clone());
}
serde_json::Value::Object(record)
}
pub fn write_export_file(records: &[serde_json::Value], filename: &str) -> Result<(), String> {
let json_str = serde_json::to_string_pretty(records)
.map_err(|e| format!("Failed to serialize records: {e}"))?;
std::fs::write(filename, json_str).map_err(|e| format!("Failed to write file: {e}"))?;
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_point_payload_to_row_with_payload() {
let payload = Some(serde_json::json!({
"title": "Hello",
"score": 0.95
}));
let row = point_payload_to_row(42, &payload);
assert_eq!(row.get("id"), Some(&serde_json::json!(42)));
assert_eq!(row.get("title"), Some(&serde_json::json!("Hello")));
assert_eq!(row.get("score"), Some(&serde_json::json!(0.95)));
assert_eq!(row.len(), 3);
}
#[test]
fn test_point_payload_to_row_without_payload() {
let row = point_payload_to_row(7, &None);
assert_eq!(row.get("id"), Some(&serde_json::json!(7)));
assert_eq!(row.len(), 1);
}
#[test]
fn test_point_payload_to_browse_row_truncates() {
let long_string = "a".repeat(80);
let payload = Some(serde_json::json!({
"content": long_string,
"short": "ok"
}));
let row = point_payload_to_browse_row(1, &payload);
assert_eq!(row.get("id"), Some(&serde_json::json!(1)));
assert_eq!(row.get("short"), Some(&serde_json::json!("ok")));
let content = row.get("content").unwrap().as_str().unwrap();
assert_eq!(content.len(), 50);
assert!(content.ends_with("..."));
}
#[test]
fn test_truncate_display_value_short_string() {
let val = serde_json::json!("short text");
let result = truncate_display_value(&val);
assert_eq!(result, serde_json::json!("short text"));
}
#[test]
fn test_truncate_display_value_long_string() {
let long = "x".repeat(100);
let result = truncate_display_value(&serde_json::json!(long));
let s = result.as_str().unwrap();
assert_eq!(s.len(), 50);
assert!(s.ends_with("..."));
assert!(s.starts_with("xxxxxxx"));
}
}