use super::AnyEngine;
use crate::config::model::{EngineType, ModelConfig, Precision};
use crate::error::VecboostError;
pub struct EngineFactory;
impl EngineFactory {
pub fn create(
engine_type: EngineType,
config: &ModelConfig,
) -> Result<AnyEngine, VecboostError> {
let precision = Precision::Fp32;
AnyEngine::new(config, engine_type, precision)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::model::{DeviceType, ModelConfig};
use std::path::PathBuf;
fn test_config() -> ModelConfig {
ModelConfig {
name: "test-factory".to_string(),
engine_type: EngineType::Candle,
model_path: PathBuf::from("/nonexistent/model"),
tokenizer_path: None,
device: DeviceType::Cpu,
max_batch_size: 32,
pooling_mode: None,
expected_dimension: Some(1024),
memory_limit_bytes: None,
oom_fallback_enabled: true,
model_sha256: None,
}
}
#[test]
fn test_create_candle_returns_error_for_missing_model() {
let config = test_config();
let result = EngineFactory::create(EngineType::Candle, &config);
assert!(result.is_err());
}
#[cfg(feature = "onnx")]
#[test]
fn test_create_onnx_engine_missing_model() {
let mut config = test_config();
config.engine_type = EngineType::Onnx;
config.model_path = PathBuf::from("/nonexistent/onnx/model");
let result = EngineFactory::create(EngineType::Onnx, &config);
assert!(
result.is_err(),
"ONNX engine should return error for missing model path"
);
}
#[test]
fn test_engine_type_display() {
assert_eq!(EngineType::Candle.to_string(), "candle");
}
#[test]
fn test_removed_engine_types_return_error() {
let json = "\"tensorrt\"";
let result: Result<EngineType, _> = serde_json::from_str(json);
assert!(
result.is_err(),
"tensorrt should not deserialize after removal"
);
let json = "\"openvino\"";
let result: Result<EngineType, _> = serde_json::from_str(json);
assert!(
result.is_err(),
"openvino should not deserialize after removal"
);
}
}