wedb_embed 0.1.0

Embedded Kvrocks-compatible storage engine for WeDb
Documentation
use crate::error::{Error, Result};
use crate::search::meta::{DistanceMetric, VectorType};
use sonic_rs::JsonValueTrait;
use sonic_rs::prelude::*;
use std::str;

/// 浮点数可排序字符串编码与解码(按位变换保证前缀扫描有序,IEEE 754 标准映射)
#[inline]
pub fn encode_sortable_f64(val: f64) -> String {
    let bits = val.to_bits();
    let sortable = if (bits & (1 << 63)) != 0 {
        !bits
    } else {
        bits ^ (1 << 63)
    };
    format!("{sortable:016x}")
}

#[inline]
pub fn decode_sortable_f64(hex_str: &str) -> Option<f64> {
    let sortable = u64::from_str_radix(hex_str, 16).ok()?;
    let bits = if (sortable & (1 << 63)) != 0 {
        sortable ^ (1 << 63)
    } else {
        !sortable
    };
    Some(f64::from_bits(bits))
}

/// 有符号 64 位整数可排序编码与解码
#[inline]
pub fn encode_sortable_i64(val: i64) -> String {
    let unsigned = (val as u64) ^ (1 << 63);
    format!("{unsigned:016x}")
}

#[inline]
pub fn decode_sortable_i64(hex_str: &str) -> Option<i64> {
    let unsigned = u64::from_str_radix(hex_str, 16).ok()?;
    Some((unsigned ^ (1 << 63)) as i64)
}

/// 向量距离与相似度度量计算(对标 Apache Kvrocks ComputeSimilarity,单次迭代聚合)
#[inline]
pub fn compute_vector_distance(v1: &[f64], v2: &[f64], metric: DistanceMetric) -> Result<f64> {
    if v1.len() != v2.len() {
        return Err(Error::invalid_data(format!(
            "vector dimension mismatch: {} vs {}",
            v1.len(),
            v2.len()
        )));
    }
    if v1.is_empty() {
        return Err(Error::invalid_data("empty vector is invalid"));
    }

    match metric {
        DistanceMetric::L2 => {
            let sum: f64 = v1
                .iter()
                .zip(v2.iter())
                .map(|(&a, &b)| {
                    let diff = a - b;
                    diff * diff
                })
                .sum();
            Ok(sum.sqrt())
        }
        DistanceMetric::IP => {
            let dot: f64 = v1.iter().zip(v2.iter()).map(|(&a, &b)| a * b).sum();
            // 内积度量下,数值越大距离越小,故取负值使得升序排序与余弦/L2一致
            Ok(-dot)
        }
        DistanceMetric::Cosine => {
            let (dot, norm1, norm2) = v1.iter().zip(v2.iter()).fold(
                (0.0f64, 0.0f64, 0.0f64),
                |(dot_acc, n1_acc, n2_acc), (&a, &b)| {
                    (dot_acc + a * b, n1_acc + a * a, n2_acc + b * b)
                },
            );
            if norm1 <= 0.0 || norm2 <= 0.0 {
                return Ok(1.0);
            }
            let sim = dot / (norm1.sqrt() * norm2.sqrt());
            let sim_clamped = sim.clamp(-1.0, 1.0);
            Ok(1.0 - sim_clamped)
        }
    }
}

/// 从二进制字节数组或 JSON 数组中解析浮点向量(零冗余单次解析)
pub fn parse_vector_from_slice(bytes: &[u8], vector_type: VectorType) -> Result<Vec<f64>> {
    let elem_size = vector_type.byte_size();
    if bytes.len().is_multiple_of(elem_size) && !bytes.is_empty() {
        match vector_type {
            VectorType::Float64 => {
                let count = bytes.len() / 8;
                let mut vec = Vec::with_capacity(count);
                for i in 0..count {
                    let chunk = &bytes[i * 8..(i + 1) * 8];
                    if let Ok(arr) = chunk.try_into() {
                        let val = f64::from_le_bytes(arr);
                        vec.push(val);
                    }
                }
                if vec.len() == count {
                    return Ok(vec);
                }
            }
            VectorType::Float32 => {
                let count = bytes.len() / 4;
                let mut vec = Vec::with_capacity(count);
                for i in 0..count {
                    let chunk = &bytes[i * 4..(i + 1) * 4];
                    if let Ok(arr) = chunk.try_into() {
                        let val = f32::from_le_bytes(arr) as f64;
                        vec.push(val);
                    }
                }
                if vec.len() == count {
                    return Ok(vec);
                }
            }
        }
    }

    // 尝试 JSON 浮点数组解析
    if let Ok(json_v) = sonic_rs::from_slice::<sonic_rs::Value>(bytes)
        && let Some(arr) = json_v.as_array()
    {
        let mut vec = Vec::with_capacity(arr.len());
        for item in arr {
            if let Some(n) = item.as_f64() {
                vec.push(n);
            }
        }
        if !vec.is_empty() {
            return Ok(vec);
        }
    }

    // 尝试逗号分隔字符串解析
    if let Ok(s) = str::from_utf8(bytes) {
        let trimmed = s.trim().trim_start_matches('[').trim_end_matches(']');
        let mut vec = Vec::new();
        for part in trimmed.split(',') {
            if let Ok(num) = part.trim().parse::<f64>() {
                vec.push(num);
            }
        }
        if !vec.is_empty() {
            return Ok(vec);
        }
    }

    Err(Error::invalid_data("invalid vector byte format or length"))
}