use anyhow::Result;
use std::fs::{File, Permissions};
use std::io::Write;
use std::os::unix::fs::PermissionsExt;
use tempfile::TempDir;
use turboprop::backends::gguf::{validate_gguf_file, GGUFBackend, GGUFEmbeddingModel};
use turboprop::backends::huggingface::{validation, HuggingFaceBackend};
use turboprop::embeddings::EmbeddingConfig;
use turboprop::models::{
EmbeddingBackend, EmbeddingModel, ModelInfo, ModelInfoConfig, ModelManager,
};
use turboprop::types::{CachePath, ModelBackend, ModelName, ModelType};
#[tokio::test]
async fn test_invalid_model_selection() -> Result<()> {
let fake_model_config = EmbeddingConfig::with_model("nonexistent/model");
let mut generator = turboprop::embeddings::MockEmbeddingGenerator::new(fake_model_config);
let test_texts = vec!["test text".to_string()];
let result = generator.embed_batch(&test_texts);
assert!(
result.is_ok(),
"Mock generator should handle invalid model names"
);
Ok(())
}
#[tokio::test]
async fn test_corrupted_cache_handling() -> Result<()> {
let temp_dir = TempDir::new()?;
let manager = ModelManager::new_with_defaults(temp_dir.path());
let model_name = ModelName::from("sentence-transformers/all-MiniLM-L6-v2");
let fake_model_path = manager.get_model_path(&model_name);
std::fs::create_dir_all(&fake_model_path)?;
std::fs::write(fake_model_path.join("invalid_file"), "corrupted data")?;
std::fs::write(fake_model_path.join("another_invalid"), "more corruption")?;
assert!(!manager.is_model_cached(&model_name));
let stats = manager.get_cache_stats();
assert!(
stats.is_ok(),
"Cache stats should handle corrupted directories"
);
Ok(())
}
#[test]
fn test_gguf_validation_errors() {
let temp_dir = TempDir::new().unwrap();
let nonexistent = temp_dir.path().join("nonexistent.gguf");
let result = validate_gguf_file(&nonexistent);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("File does not exist"));
let wrong_ext = temp_dir.path().join("model.bin");
File::create(&wrong_ext).unwrap();
let result = validate_gguf_file(&wrong_ext);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("does not have .gguf extension"));
let too_small = temp_dir.path().join("small.gguf");
let mut file = File::create(&too_small).unwrap();
file.write_all(b"tiny").unwrap();
let result = validate_gguf_file(&too_small);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("too small to be a valid GGUF model"));
let invalid_magic = temp_dir.path().join("invalid_magic.gguf");
let mut file = File::create(&invalid_magic).unwrap();
file.write_all(b"FAKE").unwrap(); file.write_all(&[1, 0, 0, 0]).unwrap(); file.write_all(&[0, 0, 0, 0, 0, 0, 0, 0]).unwrap(); let result = validate_gguf_file(&invalid_magic);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("Invalid GGUF magic header"));
let unsupported_version = temp_dir.path().join("unsupported.gguf");
let mut file = File::create(&unsupported_version).unwrap();
file.write_all(b"GGUF").unwrap(); file.write_all(&[255, 255, 255, 255]).unwrap(); file.write_all(&[0, 0, 0, 0]).unwrap(); let result = validate_gguf_file(&unsupported_version);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("Unsupported GGUF version"));
}
#[test]
fn test_gguf_model_loading_errors() {
let temp_dir = TempDir::new().unwrap();
let backend = GGUFBackend::new().unwrap();
let invalid_model_info = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("sentence-transformer-model"),
description: "Invalid model type for GGUF backend".to_string(),
dimensions: 384,
size_bytes: 1000,
model_type: ModelType::SentenceTransformer, backend: ModelBackend::FastEmbed,
download_url: None,
local_path: None,
});
let result = backend.load_model(&invalid_model_info);
assert!(result.is_err());
let error_msg = result.err().unwrap().to_string();
assert!(error_msg.contains("does not support model type"));
let invalid_path = temp_dir.path().join("nonexistent.gguf");
let model_info_with_path = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("test-model.gguf"),
description: "Model with invalid path".to_string(),
dimensions: 768,
size_bytes: 1000,
model_type: ModelType::GGUF,
backend: ModelBackend::Candle,
download_url: None,
local_path: Some(invalid_path.clone()),
});
let result = GGUFEmbeddingModel::load_from_path(&invalid_path, &model_info_with_path);
assert!(result.is_err());
let error_msg = result.unwrap_err().to_string();
assert!(error_msg.contains("File does not exist"));
}
#[test]
fn test_gguf_embedding_invalid_inputs() -> Result<()> {
let model = GGUFEmbeddingModel::new("test-model".to_string(), 768, candle_core::Device::Cpu)?;
let texts_with_empty = vec![
"Valid text".to_string(),
"".to_string(), "Another valid text".to_string(),
];
let result = model.embed(&texts_with_empty);
assert!(result.is_err());
let error_msg = result.unwrap_err().to_string();
assert!(error_msg.contains("Empty text found at index 1"));
let very_long_text = "word ".repeat(10000); let long_texts = vec![very_long_text];
let _result = model.embed(&long_texts);
Ok(())
}
#[test]
fn test_huggingface_validation_errors() {
let empty_name = ModelName::new("");
let result = validation::validate_model_name(&empty_name);
assert!(result.is_err());
assert!(result
.unwrap_err()
.to_string()
.contains("Model name cannot be empty"));
let invalid_formats = vec![
"no-slash-name", "/model-name", "organization/", "org//model", "org/model/extra/parts", ];
for invalid_name in invalid_formats {
let model_name = ModelName::new(invalid_name);
let result = validation::validate_model_name(&model_name);
assert!(
result.is_err(),
"Should fail for invalid name: {}",
invalid_name
);
}
let invalid_chars = vec![
"org/model@name", "org/model name", "org/model#name", "org/model%name", "org/model&name", ];
for invalid_name in invalid_chars {
let model_name = ModelName::new(invalid_name);
let result = validation::validate_model_name(&model_name);
assert!(
result.is_err(),
"Should fail for invalid characters: {}",
invalid_name
);
assert!(result
.unwrap_err()
.to_string()
.contains("invalid characters"));
}
}
#[test]
fn test_cache_directory_validation_errors() {
let nonexistent = CachePath::new("/nonexistent/path/that/does/not/exist");
let result = validation::validate_cache_directory(&nonexistent);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("does not exist"));
let temp_dir = TempDir::new().unwrap();
let temp_file = temp_dir.path().join("not_a_directory.txt");
std::fs::write(&temp_file, "test content").unwrap();
let file_as_cache = CachePath::new(&temp_file);
let result = validation::validate_cache_directory(&file_as_cache);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("not a directory"));
}
#[test]
#[cfg(unix)]
fn test_cache_permission_errors() {
let temp_dir = TempDir::new().unwrap();
let readonly_dir = temp_dir.path().join("readonly");
std::fs::create_dir_all(&readonly_dir).unwrap();
let mut perms = std::fs::metadata(&readonly_dir).unwrap().permissions();
perms.set_mode(0o444); std::fs::set_permissions(&readonly_dir, perms).unwrap();
let result = validation::validate_cache_permissions(&readonly_dir);
if result.is_err() {
let error_msg = result.unwrap_err().to_string();
assert!(
error_msg.contains("Cannot write to cache directory")
|| error_msg.contains("read-only")
|| error_msg.contains("permission"),
"Expected permission-related error, got: {}",
error_msg
);
}
let restore_perms = Permissions::from_mode(0o755);
let _ = std::fs::set_permissions(&readonly_dir, restore_perms);
}
#[test]
fn test_model_info_validation_errors() {
let invalid_name = ModelInfo::new(ModelInfoConfig {
name: ModelName::from(""),
description: "Test model".to_string(),
dimensions: 384,
size_bytes: 1000,
model_type: ModelType::SentenceTransformer,
backend: ModelBackend::FastEmbed,
download_url: None,
local_path: None,
});
let result = invalid_name.validate();
assert!(result.is_err());
assert!(result.unwrap_err().contains("Model name cannot be empty"));
let invalid_description = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("valid/model"),
description: "".to_string(),
dimensions: 384,
size_bytes: 1000,
model_type: ModelType::SentenceTransformer,
backend: ModelBackend::FastEmbed,
download_url: None,
local_path: None,
});
let result = invalid_description.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("Model description cannot be empty"));
let invalid_dimensions = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("valid/model"),
description: "Valid description".to_string(),
dimensions: 0,
size_bytes: 1000,
model_type: ModelType::SentenceTransformer,
backend: ModelBackend::FastEmbed,
download_url: None,
local_path: None,
});
let result = invalid_dimensions.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("Model dimensions must be greater than 0"));
let invalid_size = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("valid/model"),
description: "Valid description".to_string(),
dimensions: 384,
size_bytes: 0,
model_type: ModelType::SentenceTransformer,
backend: ModelBackend::FastEmbed,
download_url: None,
local_path: None,
});
let result = invalid_size.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("Model size must be greater than 0"));
let invalid_url = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("valid/model"),
description: "Valid description".to_string(),
dimensions: 384,
size_bytes: 1000,
model_type: ModelType::GGUF,
backend: ModelBackend::Candle,
download_url: Some("invalid-url".to_string()),
local_path: None,
});
let result = invalid_url.validate();
assert!(result.is_err());
assert!(result
.unwrap_err()
.contains("Download URL must be a valid HTTP/HTTPS URL"));
let invalid_path = ModelInfo::new(ModelInfoConfig {
name: ModelName::from("valid/model"),
description: "Valid description".to_string(),
dimensions: 384,
size_bytes: 1000,
model_type: ModelType::GGUF,
backend: ModelBackend::Candle,
download_url: None,
local_path: Some("/nonexistent/path/to/model.gguf".into()),
});
let result = invalid_path.validate();
assert!(result.is_err());
assert!(result.unwrap_err().contains("Local path does not exist"));
}
#[tokio::test]
async fn test_network_failure_handling() -> Result<()> {
let temp_dir = TempDir::new()?;
let manager = ModelManager::new_with_defaults(temp_dir.path());
let invalid_url_model = ModelInfo::gguf_model(
ModelName::from("invalid-download.gguf"),
"Model with invalid download URL".to_string(),
768,
1000,
"https://definitely-does-not-exist.invalid/model.gguf".to_string(),
);
let result = manager.download_gguf_model(&invalid_url_model).await;
assert!(result.is_err(), "Should fail with invalid URL");
let backend = HuggingFaceBackend::new()?;
let invalid_model_name = ModelName::new("definitely-does-not/exist");
let cache_dir = CachePath::new(temp_dir.path());
let _result = backend
.load_qwen3_model(&invalid_model_name, &cache_dir)
.await;
Ok(())
}
#[tokio::test]
async fn test_embedding_generator_errors() -> Result<()> {
let invalid_config = EmbeddingConfig::with_model("invalid/model").with_batch_size(0);
let mut mock_generator = turboprop::embeddings::MockEmbeddingGenerator::new(invalid_config);
let test_texts = vec!["test".to_string()];
let result = mock_generator.embed_batch(&test_texts);
assert!(result.is_ok());
Ok(())
}
#[tokio::test]
async fn test_concurrent_cache_access() -> Result<()> {
let temp_dir = TempDir::new()?;
let manager = ModelManager::new_with_defaults(temp_dir.path());
manager.init_cache()?;
let mut handles = Vec::new();
for i in 0..5 {
let manager_clone = ModelManager::new_with_defaults(temp_dir.path());
let handle = tokio::spawn(async move {
let model_name = ModelName::from(format!("test/model-{}", i));
let _is_cached = manager_clone.is_model_cached(&model_name);
let _cache_path = manager_clone.get_model_path(&model_name);
let _stats = manager_clone.get_cache_stats().unwrap();
});
handles.push(handle);
}
for handle in handles {
handle.await?;
}
Ok(())
}
#[test]
fn test_filesystem_error_handling() {
let manager = ModelManager::new_with_defaults("/root/inaccessible");
let _stats_result = manager.get_cache_stats();
let _clear_result = manager.clear_cache();
}
#[test]
fn test_model_name_edge_cases() {
let very_long_name = format!("{}/{}", "a".repeat(100), "b".repeat(100));
let long_model_name = ModelName::new(&very_long_name);
let _result = validation::validate_model_name(&long_model_name);
let boundary_cases = vec![
"org/-model", "org/model-", "org/_model", "org/model_", "org/.model", "org/model.", ];
for case in boundary_cases {
let model_name = ModelName::new(case);
let _result = validation::validate_model_name(&model_name);
}
}