use crate::traits::{
Dataset, DatasetGenerator, DatasetTraitError, DatasetTraitResult, GeneratorConfig,
InMemoryDataset,
};
use scirs2_core::random::{Distribution, Random, RandNormal};
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use thiserror::Error;
#[inline]
fn gen_normal_value(rng: &mut Random, mean: f64, std: f64) -> f64 {
let dist = RandNormal::new(mean, std).expect("operation should succeed");
dist.sample(rng)
}
#[derive(Error, Debug)]
pub enum PluginError {
#[error("Generator not found: {0}")]
GeneratorNotFound(String),
#[error("Plugin already registered: {0}")]
AlreadyRegistered(String),
#[error("Plugin registration failed: {0}")]
RegistrationFailed(String),
#[error("Plugin validation failed: {0}")]
ValidationFailed(String),
#[error("Dynamic loading error: {0}")]
DynamicLoading(String),
#[error("API version mismatch: expected {expected}, got {actual}")]
ApiVersionMismatch { expected: String, actual: String },
#[error("Plugin dependency missing: {0}")]
DependencyMissing(String),
}
pub type PluginResult<T> = Result<T, PluginError>;
#[derive(Debug, Clone)]
pub struct PluginMetadata {
pub name: String,
pub version: String,
pub description: String,
pub author: String,
pub api_version: String,
pub dependencies: Vec<String>,
pub capabilities: Vec<String>,
pub tags: Vec<String>,
}
impl Default for PluginMetadata {
fn default() -> Self {
Self {
name: "unknown".to_string(),
version: "0.1.0".to_string(),
description: "No description".to_string(),
author: "Unknown".to_string(),
api_version: "1.0.0".to_string(),
dependencies: Vec::new(),
capabilities: Vec::new(),
tags: Vec::new(),
}
}
}
pub trait PluginGenerator: Send + Sync {
fn metadata(&self) -> PluginMetadata;
fn generate(&self, config: GeneratorConfig) -> DatasetTraitResult<InMemoryDataset>;
fn validate_config(&self, config: &GeneratorConfig) -> DatasetTraitResult<()> {
let _ = config;
Ok(())
}
fn parameter_schema(&self) -> HashMap<String, ParameterInfo> {
HashMap::new()
}
fn can_handle(&self, config: &GeneratorConfig) -> bool {
let _ = config;
true
}
fn initialize(&mut self) -> PluginResult<()> {
Ok(())
}
fn cleanup(&mut self) -> PluginResult<()> {
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct ParameterInfo {
pub name: String,
pub description: String,
pub parameter_type: ParameterType,
pub required: bool,
pub default_value: Option<String>,
pub constraints: Vec<ParameterConstraint>,
}
#[derive(Debug, Clone)]
pub enum ParameterType {
Integer { min: Option<i64>, max: Option<i64> },
Float { min: Option<f64>, max: Option<f64> },
String { pattern: Option<String> },
Boolean,
IntegerArray,
FloatArray,
Enum { values: Vec<String> },
}
#[derive(Debug, Clone)]
pub enum ParameterConstraint {
Range { min: f64, max: f64 },
Length { min: usize, max: usize },
Pattern(String),
Custom(String),
}
pub struct PluginRegistry {
generators: Arc<RwLock<HashMap<String, Box<dyn PluginGenerator>>>>,
metadata_cache: Arc<RwLock<HashMap<String, PluginMetadata>>>,
hooks: Arc<RwLock<Vec<Box<dyn PluginHook>>>>,
}
impl PluginRegistry {
pub fn new() -> Self {
Self {
generators: Arc::new(RwLock::new(HashMap::new())),
metadata_cache: Arc::new(RwLock::new(HashMap::new())),
hooks: Arc::new(RwLock::new(Vec::new())),
}
}
pub fn register<G>(&self, generator: G) -> PluginResult<()>
where
G: PluginGenerator + 'static,
{
let metadata = generator.metadata();
let name = metadata.name.clone();
if metadata.api_version != "1.0.0" {
return Err(PluginError::ApiVersionMismatch {
expected: "1.0.0".to_string(),
actual: metadata.api_version,
});
}
{
let generators = self.generators.read().expect("operation should succeed");
if generators.contains_key(&name) {
return Err(PluginError::AlreadyRegistered(name));
}
}
self.validate_dependencies(&metadata.dependencies)?;
self.run_registration_hooks(&metadata)?;
{
let mut generators = self.generators.write().expect("operation should succeed");
let mut metadata_cache = self.metadata_cache.write().expect("operation should succeed");
generators.insert(name.clone(), Box::new(generator));
metadata_cache.insert(name.clone(), metadata);
}
Ok(())
}
pub fn unregister(&self, name: &str) -> PluginResult<()> {
let mut generators = self.generators.write().expect("operation should succeed");
let mut metadata_cache = self.metadata_cache.write().expect("operation should succeed");
if let Some(mut generator) = generators.remove(name) {
generator.cleanup()?;
metadata_cache.remove(name);
Ok(())
} else {
Err(PluginError::GeneratorNotFound(name.to_string()))
}
}
pub fn get(&self, name: &str) -> Option<Box<dyn PluginGenerator>> {
let generators = self.generators.read().expect("operation should succeed");
None }
pub fn has_generator(&self, name: &str) -> bool {
let generators = self.generators.read().expect("operation should succeed");
generators.contains_key(name)
}
pub fn list_generators(&self) -> Vec<String> {
let generators = self.generators.read().expect("operation should succeed");
generators.keys().cloned().collect()
}
pub fn get_metadata(&self, name: &str) -> Option<PluginMetadata> {
let metadata_cache = self.metadata_cache.read().expect("operation should succeed");
metadata_cache.get(name).cloned()
}
pub fn list_metadata(&self) -> Vec<PluginMetadata> {
let metadata_cache = self.metadata_cache.read().expect("operation should succeed");
metadata_cache.values().cloned().collect()
}
pub fn generate(
&self,
name: &str,
config: GeneratorConfig,
) -> DatasetTraitResult<InMemoryDataset> {
let generators = self.generators.read().expect("operation should succeed");
if let Some(generator) = generators.get(name) {
generator.validate_config(&config)?;
generator.generate(config)
} else {
Err(DatasetTraitError::Configuration(format!(
"Generator not found: {}",
name
)))
}
}
pub fn find_by_capability(&self, capability: &str) -> Vec<String> {
let metadata_cache = self.metadata_cache.read().expect("operation should succeed");
metadata_cache
.iter()
.filter(|(_, meta)| meta.capabilities.contains(&capability.to_string()))
.map(|(name, _)| name.clone())
.collect()
}
pub fn find_by_tag(&self, tag: &str) -> Vec<String> {
let metadata_cache = self.metadata_cache.read().expect("operation should succeed");
metadata_cache
.iter()
.filter(|(_, meta)| meta.tags.contains(&tag.to_string()))
.map(|(name, _)| name.clone())
.collect()
}
fn validate_dependencies(&self, dependencies: &[String]) -> PluginResult<()> {
let generators = self.generators.read().expect("operation should succeed");
for dep in dependencies {
if !generators.contains_key(dep) {
return Err(PluginError::DependencyMissing(dep.clone()));
}
}
Ok(())
}
fn run_registration_hooks(&self, metadata: &PluginMetadata) -> PluginResult<()> {
let hooks = self.hooks.read().expect("operation should succeed");
for hook in hooks.iter() {
hook.on_registration(metadata)?;
}
Ok(())
}
pub fn add_hook<H>(&self, hook: H)
where
H: PluginHook + 'static,
{
let mut hooks = self.hooks.write().expect("operation should succeed");
hooks.push(Box::new(hook));
}
pub fn clear_hooks(&self) {
let mut hooks = self.hooks.write().expect("operation should succeed");
hooks.clear();
}
}
impl Default for PluginRegistry {
fn default() -> Self {
Self::new()
}
}
pub trait PluginHook: Send + Sync {
fn on_registration(&self, metadata: &PluginMetadata) -> PluginResult<()>;
fn on_unregistration(&self, name: &str) -> PluginResult<()> {
let _ = name;
Ok(())
}
fn on_generation_start(
&self,
generator_name: &str,
config: &GeneratorConfig,
) -> PluginResult<()> {
let _ = (generator_name, config);
Ok(())
}
fn on_generation_complete(
&self,
generator_name: &str,
dataset: &InMemoryDataset,
) -> PluginResult<()> {
let _ = (generator_name, dataset);
Ok(())
}
}
pub struct CustomLinearGenerator;
impl PluginGenerator for CustomLinearGenerator {
fn metadata(&self) -> PluginMetadata {
PluginMetadata {
name: "custom_linear".to_string(),
version: "1.0.0".to_string(),
description: "Generates linear datasets with custom patterns".to_string(),
author: "Example Author".to_string(),
api_version: "1.0.0".to_string(),
dependencies: vec![],
capabilities: vec!["regression".to_string(), "linear".to_string()],
tags: vec!["custom".to_string(), "linear".to_string()],
}
}
fn generate(&self, config: GeneratorConfig) -> DatasetTraitResult<InMemoryDataset> {
use scirs2_core::ndarray::{Array1, Array2};
use scirs2_core::random::Random;
let mut rng = match config.random_state {
Some(seed) => Random::new_with_seed(seed),
None => Random::new(),
};
let slope = config
.get_parameter("slope")
.and_then(|v| match v {
crate::traits::ConfigValue::Float(s) => Some(*s),
_ => None,
})
.unwrap_or(1.0);
let mut features = Array2::<f64>::zeros((config.n_samples, config.n_features));
for mut row in features.rows_mut() {
for val in row.iter_mut() {
*val = rng.random_range(-10.0..10.0);
}
}
let targets: Array1<f64> = features
.rows()
.into_iter()
.map(|row| row.sum() * slope + gen_normal_value(&mut rng, 0.0, 0.1))
.collect();
let mut metadata = std::collections::HashMap::new();
metadata.insert("generator".to_string(), "custom_linear".to_string());
metadata.insert("slope".to_string(), slope.to_string());
Ok(crate::traits::InMemoryDataset::with_metadata(
features,
Some(targets),
metadata,
))
}
fn parameter_schema(&self) -> HashMap<String, ParameterInfo> {
let mut schema = HashMap::new();
schema.insert(
"slope".to_string(),
ParameterInfo {
name: "slope".to_string(),
description: "Linear slope coefficient".to_string(),
parameter_type: ParameterType::Float {
min: Some(-10.0),
max: Some(10.0),
},
required: false,
default_value: Some("1.0".to_string()),
constraints: vec![ParameterConstraint::Range {
min: -10.0,
max: 10.0,
}],
},
);
schema
}
fn validate_config(&self, config: &GeneratorConfig) -> DatasetTraitResult<()> {
if config.n_samples == 0 || config.n_features == 0 {
return Err(DatasetTraitError::Configuration(
"n_samples and n_features must be > 0".to_string(),
));
}
if let Some(slope_val) = config.get_parameter("slope") {
match slope_val {
crate::traits::ConfigValue::Float(slope) => {
if !(-10.0..=10.0).contains(slope) {
return Err(DatasetTraitError::Configuration(
"slope must be between -10.0 and 10.0".to_string(),
));
}
}
_ => {
return Err(DatasetTraitError::Configuration(
"slope parameter must be a float".to_string(),
));
}
}
}
Ok(())
}
}
pub struct LoggingHook;
impl PluginHook for LoggingHook {
fn on_registration(&self, metadata: &PluginMetadata) -> PluginResult<()> {
println!("Plugin registered: {} v{}", metadata.name, metadata.version);
Ok(())
}
fn on_unregistration(&self, name: &str) -> PluginResult<()> {
println!("Plugin unregistered: {}", name);
Ok(())
}
fn on_generation_start(
&self,
generator_name: &str,
config: &GeneratorConfig,
) -> PluginResult<()> {
println!(
"Starting generation with {}: {} samples x {} features",
generator_name, config.n_samples, config.n_features
);
Ok(())
}
fn on_generation_complete(
&self,
generator_name: &str,
dataset: &InMemoryDataset,
) -> PluginResult<()> {
println!(
"Completed generation with {}: {} samples generated",
generator_name,
dataset.n_samples()
);
Ok(())
}
}
pub struct PluginManager {
registry: PluginRegistry,
}
impl PluginManager {
pub fn new() -> Self {
Self {
registry: PluginRegistry::new(),
}
}
pub fn registry(&self) -> &PluginRegistry {
&self.registry
}
pub fn load_builtin_plugins(&self) -> PluginResult<()> {
self.registry.register(CustomLinearGenerator)?;
self.registry.add_hook(LoggingHook);
Ok(())
}
pub fn discover_plugins(&self, _plugin_dir: &str) -> PluginResult<Vec<String>> {
Ok(vec![])
}
pub fn validate_all(&self) -> PluginResult<Vec<String>> {
let mut failed = Vec::new();
let generators = self.registry.list_generators();
for generator_name in generators {
if let Some(metadata) = self.registry.get_metadata(&generator_name) {
if let Err(_) = self.registry.validate_dependencies(&metadata.dependencies) {
failed.push(generator_name);
}
}
}
if failed.is_empty() {
Ok(vec![])
} else {
Err(PluginError::ValidationFailed(format!(
"Failed plugins: {:?}",
failed
)))
}
}
pub fn create_test_config(&self, generator_name: &str) -> Option<GeneratorConfig> {
if let Some(metadata) = self.registry.get_metadata(generator_name) {
let mut config = GeneratorConfig::new(100, 4);
config = config.with_random_state(42);
if metadata.capabilities.contains(&"regression".to_string()) {
config.set_parameter("noise".to_string(), 0.1);
}
Some(config)
} else {
None
}
}
}
impl Default for PluginManager {
fn default() -> Self {
Self::new()
}
}
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use crate::traits::ConfigValue;
#[test]
fn test_plugin_registration() {
let registry = PluginRegistry::new();
assert!(registry.register(CustomLinearGenerator).is_ok());
assert!(registry.has_generator("custom_linear"));
assert_eq!(registry.list_generators(), vec!["custom_linear"]);
let metadata = registry.get_metadata("custom_linear").expect("operation should succeed");
assert_eq!(metadata.name, "custom_linear");
assert_eq!(metadata.version, "1.0.0");
}
#[test]
fn test_custom_generator() {
let generator = CustomLinearGenerator;
let metadata = generator.metadata();
assert_eq!(metadata.name, "custom_linear");
assert!(metadata.capabilities.contains(&"regression".to_string()));
let schema = generator.parameter_schema();
assert!(schema.contains_key("slope"));
let mut config = GeneratorConfig::new(50, 3);
config.set_parameter("slope".to_string(), 2.0);
config = config.with_random_state(42);
assert!(generator.validate_config(&config).is_ok());
let dataset = generator.generate(config).expect("operation should succeed");
assert_eq!(dataset.n_samples(), 50);
assert_eq!(dataset.n_features(), 3);
assert!(dataset.has_targets());
}
#[test]
fn test_plugin_hooks() {
let registry = PluginRegistry::new();
registry.add_hook(LoggingHook);
assert!(registry.register(CustomLinearGenerator).is_ok());
}
#[test]
fn test_plugin_manager() {
let manager = PluginManager::new();
assert!(manager.load_builtin_plugins().is_ok());
assert!(manager.validate_all().is_ok());
let config = manager.create_test_config("custom_linear");
assert!(config.is_some());
let config = config.expect("operation should succeed");
assert_eq!(config.n_samples, 100);
assert_eq!(config.n_features, 4);
}
#[test]
fn test_capability_and_tag_search() {
let registry = PluginRegistry::new();
registry.register(CustomLinearGenerator).expect("operation should succeed");
let regression_generators = registry.find_by_capability("regression");
assert!(regression_generators.contains(&"custom_linear".to_string()));
let custom_generators = registry.find_by_tag("custom");
assert!(custom_generators.contains(&"custom_linear".to_string()));
}
#[test]
fn test_parameter_validation() {
let generator = CustomLinearGenerator;
let mut valid_config = GeneratorConfig::new(100, 5);
valid_config.set_parameter("slope".to_string(), 2.0);
assert!(generator.validate_config(&valid_config).is_ok());
let mut invalid_config = GeneratorConfig::new(100, 5);
invalid_config.set_parameter("slope".to_string(), 15.0);
assert!(generator.validate_config(&invalid_config).is_err());
let invalid_dims = GeneratorConfig::new(0, 5);
assert!(generator.validate_config(&invalid_dims).is_err());
}
}