use std::path::PathBuf;
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::RwLock;
use vecboost::config::model::{EngineType, ModelConfig, Precision};
use vecboost::engine::{AnyEngine, InferenceEngine};
use vecboost::error::VecboostError;
#[derive(Debug, Clone, PartialEq)]
pub enum TestMode {
Mock,
Light,
Full,
}
impl TestMode {
pub fn from_env() -> Self {
match std::env::var("TEST_MODE").as_deref() {
Ok("mock") => TestMode::Mock,
Ok("light") | Ok("real") => TestMode::Light,
Ok("full") => TestMode::Full,
_ => TestMode::Mock, }
}
pub fn is_mock(&self) -> bool {
matches!(self, TestMode::Mock)
}
#[allow(dead_code)]
pub fn is_real(&self) -> bool {
matches!(self, TestMode::Light | TestMode::Full)
}
}
pub fn get_test_model_config() -> ModelConfig {
let mode = TestMode::from_env();
match mode {
TestMode::Mock => ModelConfig::default(),
TestMode::Light => ModelConfig {
name: "bge-small-en-v1.5".to_string(),
engine_type: EngineType::Candle,
model_path: PathBuf::from("models/bge-small-en-v1.5"),
tokenizer_path: Some(PathBuf::from("models/bge-small-en-v1.5-tokenizer")),
device: vecboost::config::model::DeviceType::Cpu,
max_batch_size: 16,
pooling_mode: None,
expected_dimension: Some(384),
memory_limit_bytes: None,
oom_fallback_enabled: true,
model_sha256: None,
},
TestMode::Full => ModelConfig {
name: "bge-m3".to_string(),
engine_type: EngineType::Candle,
model_path: PathBuf::from("models/bge-m3"),
tokenizer_path: Some(PathBuf::from("models/bge-m3-tokenizer")),
device: vecboost::config::model::DeviceType::Cpu,
max_batch_size: 8,
pooling_mode: None,
expected_dimension: Some(1024),
memory_limit_bytes: None,
oom_fallback_enabled: true,
model_sha256: None,
},
}
}
#[derive(Clone)]
pub struct MockEngine {
dimension: usize,
}
impl MockEngine {
pub fn new(dimension: usize) -> Self {
Self { dimension }
}
fn generate_embedding(&self, text: &str) -> Vec<f32> {
let mut embedding = vec![0.0; self.dimension];
let bytes = text.as_bytes();
let mut hash: u64 = 1469598103934665603;
for &byte in bytes {
hash ^= byte as u64;
hash = hash.wrapping_mul(1099511628211);
}
let mut state = hash;
for val in embedding.iter_mut() {
state = state.wrapping_mul(1664525).wrapping_add(1013904223);
let float_val = (state as f32 / u32::MAX as f32) * 2.0 - 1.0;
*val = float_val;
}
let norm: f32 = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for val in embedding.iter_mut() {
*val /= norm;
}
}
embedding
}
}
#[async_trait]
impl InferenceEngine for MockEngine {
fn embed(&self, text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(self.generate_embedding(text))
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
let embeddings: Vec<Vec<f32>> = texts.iter().map(|t| self.generate_embedding(t)).collect();
Ok(embeddings)
}
fn precision(&self) -> &Precision {
&Precision::Fp32
}
fn supports_mixed_precision(&self) -> bool {
false
}
async fn try_fallback_to_cpu(&mut self, _config: &ModelConfig) -> Result<(), VecboostError> {
Ok(())
}
}
pub struct RealTestEngine {
real_engine: Option<AnyEngine>,
mock_engine: MockEngine,
use_fallback: bool,
expected_dimension: usize,
}
#[allow(dead_code)]
impl RealTestEngine {
pub fn new() -> Self {
let mode = TestMode::from_env();
let config = get_test_model_config();
let expected_dimension = config.expected_dimension.unwrap_or(384);
if mode.is_mock() {
tracing::info!("Using Mock engine (TEST_MODE=mock)");
Self {
real_engine: None,
mock_engine: MockEngine::new(expected_dimension),
use_fallback: true,
expected_dimension,
}
} else {
match AnyEngine::new(&config, config.engine_type.clone(), Precision::Fp32) {
Ok(engine) => {
tracing::info!(
"Using real engine: {} (dimension={})",
config.name,
expected_dimension
);
Self {
real_engine: Some(engine),
mock_engine: MockEngine::new(expected_dimension),
use_fallback: false,
expected_dimension,
}
}
Err(e) => {
tracing::warn!(
"Failed to initialize real engine: {}. Falling back to mock.",
e
);
Self {
real_engine: None,
mock_engine: MockEngine::new(expected_dimension),
use_fallback: true,
expected_dimension,
}
}
}
}
}
pub fn with_dimension(dimension: usize) -> Self {
Self {
real_engine: None,
mock_engine: MockEngine::new(dimension),
use_fallback: true,
expected_dimension: dimension,
}
}
pub fn is_using_real_engine(&self) -> bool {
self.real_engine.is_some() && !self.use_fallback
}
pub fn is_using_fallback(&self) -> bool {
self.use_fallback
}
pub fn engine_info(&self) -> &str {
if self.use_fallback { "mock" } else { "real" }
}
}
impl Default for RealTestEngine {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
#[allow(clippy::collapsible_if)]
impl InferenceEngine for RealTestEngine {
fn embed(&self, text: &str) -> Result<Vec<f32>, VecboostError> {
if let Some(ref engine) = self.real_engine
&& !self.use_fallback
{
match engine.embed(text) {
Ok(embedding) => {
if embedding.len() == self.expected_dimension {
return Ok(embedding);
}
tracing::warn!(
"Engine returned dimension {}, expected {}. Using fallback.",
embedding.len(),
self.expected_dimension
);
}
Err(e) => {
tracing::warn!("Real engine embed failed: {}. Using fallback.", e);
}
}
}
Ok(self.mock_engine.generate_embedding(text))
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
if let Some(ref engine) = self.real_engine
&& !self.use_fallback
{
match engine.embed_batch(texts) {
Ok(embeddings) => {
#[allow(clippy::collapsible_if)]
if let Some(first) = embeddings.first() {
if first.len() == self.expected_dimension {
return Ok(embeddings);
}
}
tracing::warn!(
"Engine returned unexpected dimension. Expected {}. Using fallback.",
self.expected_dimension
);
}
Err(e) => {
tracing::warn!("Real engine embed_batch failed: {}. Using fallback.", e);
}
}
}
let embeddings: Vec<Vec<f32>> = texts
.iter()
.map(|t| self.mock_engine.generate_embedding(t))
.collect();
Ok(embeddings)
}
fn precision(&self) -> &Precision {
if self.use_fallback {
&Precision::Fp32
} else if let Some(ref engine) = self.real_engine {
engine.precision()
} else {
&Precision::Fp32
}
}
fn supports_mixed_precision(&self) -> bool {
if self.use_fallback {
false
} else if let Some(ref engine) = self.real_engine {
engine.supports_mixed_precision()
} else {
false
}
}
fn is_fallback_triggered(&self) -> bool {
self.use_fallback
}
async fn try_fallback_to_cpu(&mut self, config: &ModelConfig) -> Result<(), VecboostError> {
if self.use_fallback {
return Ok(()); }
if let Some(ref mut engine) = self.real_engine {
match engine.try_fallback_to_cpu(config).await {
Ok(()) => {
self.use_fallback = false;
tracing::info!("Successfully fell back to CPU");
Ok(())
}
Err(e) => {
tracing::warn!("Failed to fallback to CPU: {}. Using mock fallback.", e);
self.use_fallback = true;
Ok(())
}
}
} else {
Ok(())
}
}
}
pub fn create_test_engine()
-> Result<Arc<RwLock<dyn InferenceEngine + Send + Sync>>, Box<dyn std::error::Error>> {
let mode = TestMode::from_env();
if mode.is_mock() {
let engine = RealTestEngine::with_dimension(1024);
Ok(Arc::new(RwLock::new(engine)))
} else {
let engine = RealTestEngine::new();
Ok(Arc::new(RwLock::new(engine)))
}
}
#[allow(dead_code)]
pub fn create_test_engine_with_dimension(
dimension: usize,
) -> Result<Arc<RwLock<dyn InferenceEngine + Send + Sync>>, Box<dyn std::error::Error>> {
let engine = RealTestEngine::with_dimension(dimension);
Ok(Arc::new(RwLock::new(engine)))
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_mock_engine_basic() {
let engine = RealTestEngine::with_dimension(384);
let result = engine.embed("Hello world").unwrap();
assert_eq!(result.len(), 384);
assert!(result.iter().all(|&x| x.is_finite()));
let norm: f32 = result.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5);
}
#[tokio::test]
async fn test_mock_engine_determinism() {
let engine = RealTestEngine::with_dimension(384);
let result1 = engine.embed("Hello world").unwrap();
let result2 = engine.embed("Hello world").unwrap();
assert_eq!(result1, result2);
}
#[tokio::test]
async fn test_mock_engine_batch() {
let engine = RealTestEngine::with_dimension(384);
let texts = vec![
"Hello world".to_string(),
"Machine learning".to_string(),
"Artificial intelligence".to_string(),
];
let results = engine.embed_batch(&texts).unwrap();
assert_eq!(results.len(), 3);
for result in &results {
assert_eq!(result.len(), 384);
}
}
#[test]
fn test_test_mode_from_env() {
unsafe {
std::env::remove_var("TEST_MODE");
}
assert_eq!(TestMode::from_env(), TestMode::Mock);
unsafe {
std::env::set_var("TEST_MODE", "mock");
}
assert_eq!(TestMode::from_env(), TestMode::Mock);
unsafe {
std::env::set_var("TEST_MODE", "light");
}
assert_eq!(TestMode::from_env(), TestMode::Light);
unsafe {
std::env::set_var("TEST_MODE", "full");
}
assert_eq!(TestMode::from_env(), TestMode::Full);
unsafe {
std::env::remove_var("TEST_MODE");
}
}
}