diff --git a/ml/src/model_registry.rs b/ml/src/model_registry.rs index 447d4fd3d..b5adc93a6 100644 --- a/ml/src/model_registry.rs +++ b/ml/src/model_registry.rs @@ -697,6 +697,310 @@ pub struct RegistryStatistics { pub earliest_training_date: Option>, } +/// A rollback event recorded in the audit log +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RollbackEvent { + /// Model name that was rolled back + pub model_name: String, + /// Version that was active before the rollback + pub from_version: String, + /// Version that became active after the rollback + pub to_version: String, + /// Timestamp of the rollback + pub timestamp: DateTime, + /// Optional reason for the rollback + pub reason: Option, +} + +/// Version entry stored per model in the in-memory registry +#[derive(Debug, Clone)] +struct VersionEntry { + /// All registered versions (ordered by registration time) + versions: Vec, + /// Index into `versions` for the currently active version + active_index: usize, +} + +/// In-memory model registry with version tracking and rollback support. +/// +/// This registry does not require a database. It stores all model versions +/// in memory, tracks which version is "active" for each model name, and +/// records rollback events in an audit log. +#[derive(Debug, Clone)] +pub struct InMemoryModelRegistry { + /// Per-model version entries keyed by model name (e.g. "dqn", "ppo") + models: Arc>>, + /// Audit log of rollback events + rollback_log: Arc>>, +} + +impl InMemoryModelRegistry { + /// Create a new empty in-memory model registry. + pub fn new() -> Self { + Self { + models: Arc::new(RwLock::new(HashMap::new())), + rollback_log: Arc::new(RwLock::new(Vec::new())), + } + } + + /// Register a new model version. + /// + /// The first version registered for a given model name automatically becomes + /// the active version. Subsequent registrations are stored but do not change + /// the active version (use [`rollback_to_version`] or [`promote_version`] for that). + /// + /// # Arguments + /// + /// * `model_name` - Logical model name (e.g. "dqn", "ppo") + /// * `metadata` - Full version metadata + /// + /// # Errors + /// + /// Returns `MLError::ModelError` if a version with the same version string + /// is already registered for this model name. + pub async fn register_version( + &self, + model_name: &str, + metadata: ModelVersionMetadata, + ) -> MLResult<()> { + let mut models = self.models.write().await; + let entry = models + .entry(model_name.to_string()) + .or_insert_with(|| VersionEntry { + versions: Vec::new(), + active_index: 0, + }); + + // Check for duplicate version strings + let version_str = metadata.version.clone(); + for existing in &entry.versions { + if existing.version == version_str { + return Err(MLError::ModelError(format!( + "Version {} already registered for model {}", + version_str, model_name + ))); + } + } + + entry.versions.push(metadata); + + // First version auto-becomes active (active_index is already 0) + // Subsequent versions do not change active_index + + tracing::info!( + model_name = model_name, + version = version_str.as_str(), + total_versions = entry.versions.len(), + "Registered model version" + ); + + Ok(()) + } + + /// Roll back a model to a previously registered version. + /// + /// This sets the specified version as the active version and records + /// a rollback event in the audit log. + /// + /// # Arguments + /// + /// * `model_name` - Logical model name + /// * `target_version` - Semantic version string to roll back to + /// * `reason` - Optional reason for the rollback + /// + /// # Errors + /// + /// Returns `MLError::ModelNotFound` if the model name or target version + /// is not found in the registry. + pub async fn rollback_to_version( + &self, + model_name: &str, + target_version: &str, + reason: Option, + ) -> MLResult<()> { + let mut models = self.models.write().await; + let entry = models.get_mut(model_name).ok_or_else(|| { + MLError::ModelNotFound(format!("Model {} not found in registry", model_name)) + })?; + + // Find the target version index + let target_index = entry + .versions + .iter() + .position(|v| v.version == target_version) + .ok_or_else(|| { + MLError::ModelNotFound(format!( + "Version {} not found for model {}", + target_version, model_name + )) + })?; + + let from_version = entry + .versions + .get(entry.active_index) + .map(|v| v.version.clone()) + .unwrap_or_default(); + + if entry.active_index == target_index { + tracing::warn!( + model_name = model_name, + version = target_version, + "Rollback requested to already-active version (no-op)" + ); + return Ok(()); + } + + entry.active_index = target_index; + + tracing::info!( + model_name = model_name, + from_version = from_version.as_str(), + to_version = target_version, + reason = reason.as_deref().unwrap_or("none"), + "Rolled back model version" + ); + + // Record rollback event + let event = RollbackEvent { + model_name: model_name.to_string(), + from_version, + to_version: target_version.to_string(), + timestamp: Utc::now(), + reason, + }; + // Drop models lock before acquiring rollback_log lock to avoid deadlock + drop(models); + self.rollback_log.write().await.push(event); + + Ok(()) + } + + /// Promote a version to be the active version (same as rollback but with + /// clearer semantics for forward version changes). + /// + /// # Errors + /// + /// Returns `MLError::ModelNotFound` if the model or version is not found. + pub async fn promote_version( + &self, + model_name: &str, + target_version: &str, + ) -> MLResult<()> { + self.rollback_to_version(model_name, target_version, Some("promoted".to_string())) + .await + } + + /// List all registered versions for a model, ordered by registration time. + /// + /// # Arguments + /// + /// * `model_name` - Logical model name + /// + /// # Returns + /// + /// A vector of `(version_string, is_active)` tuples. + /// + /// # Errors + /// + /// Returns `MLError::ModelNotFound` if the model name is not found. + pub async fn list_versions( + &self, + model_name: &str, + ) -> MLResult> { + let models = self.models.read().await; + let entry = models.get(model_name).ok_or_else(|| { + MLError::ModelNotFound(format!("Model {} not found in registry", model_name)) + })?; + + let result = entry + .versions + .iter() + .enumerate() + .map(|(i, v)| (v.version.clone(), i == entry.active_index)) + .collect(); + + Ok(result) + } + + /// Get the active version metadata for a model. + /// + /// # Errors + /// + /// Returns `MLError::ModelNotFound` if the model is not found or has no versions. + pub async fn get_active_version( + &self, + model_name: &str, + ) -> MLResult { + let models = self.models.read().await; + let entry = models.get(model_name).ok_or_else(|| { + MLError::ModelNotFound(format!("Model {} not found in registry", model_name)) + })?; + + entry + .versions + .get(entry.active_index) + .cloned() + .ok_or_else(|| { + MLError::ModelNotFound(format!("No versions registered for model {}", model_name)) + }) + } + + /// Get a specific version's metadata for a model. + /// + /// # Errors + /// + /// Returns `MLError::ModelNotFound` if the model or version is not found. + pub async fn get_version( + &self, + model_name: &str, + version: &str, + ) -> MLResult { + let models = self.models.read().await; + let entry = models.get(model_name).ok_or_else(|| { + MLError::ModelNotFound(format!("Model {} not found in registry", model_name)) + })?; + + entry + .versions + .iter() + .find(|v| v.version == version) + .cloned() + .ok_or_else(|| { + MLError::ModelNotFound(format!( + "Version {} not found for model {}", + version, model_name + )) + }) + } + + /// Get all model names in the registry. + pub async fn list_models(&self) -> Vec { + self.models.read().await.keys().cloned().collect() + } + + /// Get the rollback audit log. + pub async fn get_rollback_log(&self) -> Vec { + self.rollback_log.read().await.clone() + } + + /// Get rollback events for a specific model. + pub async fn get_rollback_log_for_model(&self, model_name: &str) -> Vec { + self.rollback_log + .read() + .await + .iter() + .filter(|e| e.model_name == model_name) + .cloned() + .collect() + } +} + +impl Default for InMemoryModelRegistry { + fn default() -> Self { + Self::new() + } +} + #[cfg(test)] mod tests { use super::*; @@ -747,4 +1051,300 @@ mod tests { assert_eq!(retrieved.model_id, "dqn-test-v1.0.0"); assert_eq!(retrieved.version, "1.0.0"); } + + // ---- In-memory registry tests (no database required) ---- + + fn make_version(model_id: &str, version: &str) -> ModelVersionMetadata { + ModelVersionMetadata::new( + model_id.to_string(), + ModelType::DQN, + version.to_string(), + "test_data".to_string(), + format!("s3://models/{}/{}/", model_id, version), + ) + } + + #[tokio::test] + async fn test_inmemory_register_and_get_active() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + registry.register_version("dqn", v1).await.unwrap(); + + let active = registry.get_active_version("dqn").await.unwrap(); + assert_eq!(active.version, "1.0.0"); + } + + #[tokio::test] + async fn test_inmemory_first_version_is_active() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + + // First registered version stays active + let active = registry.get_active_version("dqn").await.unwrap(); + assert_eq!(active.version, "1.0.0"); + } + + #[tokio::test] + async fn test_inmemory_duplicate_version_rejected() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v1_dup = make_version("dqn-v1-dup", "1.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + let result = registry.register_version("dqn", v1_dup).await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_inmemory_list_versions() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + let v3 = make_version("dqn-v3", "3.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + registry.register_version("dqn", v3).await.unwrap(); + + let versions = registry.list_versions("dqn").await.unwrap(); + assert_eq!(versions.len(), 3); + assert_eq!(versions.first().map(|v| v.0.as_str()), Some("1.0.0")); + assert_eq!(versions.first().map(|v| v.1), Some(true)); // active + assert_eq!(versions.get(1).map(|v| v.0.as_str()), Some("2.0.0")); + assert_eq!(versions.get(1).map(|v| v.1), Some(false)); // not active + assert_eq!(versions.get(2).map(|v| v.0.as_str()), Some("3.0.0")); + assert_eq!(versions.get(2).map(|v| v.1), Some(false)); // not active + } + + #[tokio::test] + async fn test_inmemory_list_versions_unknown_model() { + let registry = InMemoryModelRegistry::new(); + let result = registry.list_versions("nonexistent").await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_inmemory_rollback_to_version() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + let v3 = make_version("dqn-v3", "3.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + registry.register_version("dqn", v3).await.unwrap(); + + // Promote to v3 first + registry.promote_version("dqn", "3.0.0").await.unwrap(); + let active = registry.get_active_version("dqn").await.unwrap(); + assert_eq!(active.version, "3.0.0"); + + // Roll back to v1 + registry + .rollback_to_version("dqn", "1.0.0", Some("regression in v3".to_string())) + .await + .unwrap(); + + let active = registry.get_active_version("dqn").await.unwrap(); + assert_eq!(active.version, "1.0.0"); + + // Check rollback log + let log = registry.get_rollback_log().await; + // Two events: promote to v3 and rollback to v1 + assert_eq!(log.len(), 2); + + let last = log.get(1); + assert!(last.is_some()); + if let Some(event) = last { + assert_eq!(event.model_name, "dqn"); + assert_eq!(event.from_version, "3.0.0"); + assert_eq!(event.to_version, "1.0.0"); + assert_eq!(event.reason.as_deref(), Some("regression in v3")); + } + } + + #[tokio::test] + async fn test_inmemory_rollback_unknown_model() { + let registry = InMemoryModelRegistry::new(); + let result = registry + .rollback_to_version("nonexistent", "1.0.0", None) + .await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_inmemory_rollback_unknown_version() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + registry.register_version("dqn", v1).await.unwrap(); + + let result = registry + .rollback_to_version("dqn", "99.0.0", None) + .await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_inmemory_rollback_to_same_version_is_noop() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + registry.register_version("dqn", v1).await.unwrap(); + + // Rolling back to already-active version succeeds silently + registry + .rollback_to_version("dqn", "1.0.0", None) + .await + .unwrap(); + + // No rollback event recorded for no-op + let log = registry.get_rollback_log().await; + assert!(log.is_empty()); + } + + #[tokio::test] + async fn test_inmemory_get_version() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + + let retrieved = registry.get_version("dqn", "2.0.0").await.unwrap(); + assert_eq!(retrieved.version, "2.0.0"); + assert_eq!(retrieved.model_id, "dqn-v2"); + } + + #[tokio::test] + async fn test_inmemory_get_version_not_found() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + registry.register_version("dqn", v1).await.unwrap(); + + let result = registry.get_version("dqn", "99.0.0").await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn test_inmemory_list_models() { + let registry = InMemoryModelRegistry::new(); + + let dqn = make_version("dqn-v1", "1.0.0"); + let ppo = make_version("ppo-v1", "1.0.0"); + + registry.register_version("dqn", dqn).await.unwrap(); + registry.register_version("ppo", ppo).await.unwrap(); + + let mut models = registry.list_models().await; + models.sort(); + assert_eq!(models, vec!["dqn", "ppo"]); + } + + #[tokio::test] + async fn test_inmemory_rollback_log_per_model() { + let registry = InMemoryModelRegistry::new(); + + let dqn_v1 = make_version("dqn-v1", "1.0.0"); + let dqn_v2 = make_version("dqn-v2", "2.0.0"); + let ppo_v1 = make_version("ppo-v1", "1.0.0"); + let ppo_v2 = make_version("ppo-v2", "2.0.0"); + + registry.register_version("dqn", dqn_v1).await.unwrap(); + registry.register_version("dqn", dqn_v2).await.unwrap(); + registry.register_version("ppo", ppo_v1).await.unwrap(); + registry.register_version("ppo", ppo_v2).await.unwrap(); + + // Roll back both + registry + .rollback_to_version("dqn", "2.0.0", None) + .await + .unwrap(); + registry + .rollback_to_version("ppo", "2.0.0", None) + .await + .unwrap(); + + // Filter by model + let dqn_log = registry.get_rollback_log_for_model("dqn").await; + assert_eq!(dqn_log.len(), 1); + assert_eq!(dqn_log.first().map(|e| e.model_name.as_str()), Some("dqn")); + + let ppo_log = registry.get_rollback_log_for_model("ppo").await; + assert_eq!(ppo_log.len(), 1); + assert_eq!(ppo_log.first().map(|e| e.model_name.as_str()), Some("ppo")); + } + + #[tokio::test] + async fn test_inmemory_default_trait() { + let registry = InMemoryModelRegistry::default(); + let models = registry.list_models().await; + assert!(models.is_empty()); + } + + #[tokio::test] + async fn test_inmemory_multiple_rollbacks() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + let v3 = make_version("dqn-v3", "3.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + registry.register_version("dqn", v3).await.unwrap(); + + // v1 -> v3 -> v2 -> v1 -> v3 + registry.promote_version("dqn", "3.0.0").await.unwrap(); + registry + .rollback_to_version("dqn", "2.0.0", None) + .await + .unwrap(); + registry + .rollback_to_version("dqn", "1.0.0", None) + .await + .unwrap(); + registry.promote_version("dqn", "3.0.0").await.unwrap(); + + let active = registry.get_active_version("dqn").await.unwrap(); + assert_eq!(active.version, "3.0.0"); + + let log = registry.get_rollback_log().await; + assert_eq!(log.len(), 4); + } + + #[tokio::test] + async fn test_inmemory_version_list_reflects_active_after_rollback() { + let registry = InMemoryModelRegistry::new(); + + let v1 = make_version("dqn-v1", "1.0.0"); + let v2 = make_version("dqn-v2", "2.0.0"); + + registry.register_version("dqn", v1).await.unwrap(); + registry.register_version("dqn", v2).await.unwrap(); + + // Initially v1 is active + let versions = registry.list_versions("dqn").await.unwrap(); + assert_eq!(versions.first().map(|v| v.1), Some(true)); + assert_eq!(versions.get(1).map(|v| v.1), Some(false)); + + // Promote v2 + registry.promote_version("dqn", "2.0.0").await.unwrap(); + + let versions = registry.list_versions("dqn").await.unwrap(); + assert_eq!(versions.first().map(|v| v.1), Some(false)); + assert_eq!(versions.get(1).map(|v| v.1), Some(true)); + } } diff --git a/services/trading_agent_service/src/allocation.rs b/services/trading_agent_service/src/allocation.rs index a90729e82..19b3c7493 100644 --- a/services/trading_agent_service/src/allocation.rs +++ b/services/trading_agent_service/src/allocation.rs @@ -75,6 +75,51 @@ impl PortfolioAllocator { } } + /// Allocate capital across assets using a correlation matrix + /// + /// Like [`allocate`](Self::allocate), but accepts an N x N correlation matrix + /// to build a full covariance matrix for mean-variance optimization. + /// Only meaningful when the allocation method is `MeanVariance` or `MLOptimized`; + /// other methods ignore the correlation matrix. + /// + /// # Arguments + /// * `assets` - Asset information (returns, volatility, ML scores) + /// * `total_capital` - Total capital to allocate + /// * `correlations` - N x N correlation matrix (must be symmetric, 1.0 on diagonal) + /// + /// # Errors + /// Returns an error if the correlation matrix dimensions do not match the asset count. + pub fn allocate_with_correlations( + &self, + assets: &[AssetInfo], + total_capital: Decimal, + correlations: &DMatrix, + ) -> Result> { + if assets.is_empty() { + return Ok(HashMap::new()); + } + + match &self.method { + AllocationMethod::MeanVariance { lambda } => { + self.mean_variance_with_corr(assets, total_capital, *lambda, Some(correlations)) + } + AllocationMethod::MLOptimized => { + // Use ML scores as expected returns, then apply correlated mean-variance + let ml_assets: Vec = assets + .iter() + .map(|a| { + let mut asset = a.clone(); + asset.expected_return = a.ml_score; + asset + }) + .collect(); + self.mean_variance_with_corr(&ml_assets, total_capital, 1.0, Some(correlations)) + } + // Other methods don't use correlations — delegate to standard allocate + _ => self.allocate(assets, total_capital), + } + } + /// Strategy 1: Equal Weight (Baseline) /// /// Allocates capital equally across all assets (1/N portfolio). @@ -126,35 +171,69 @@ impl PortfolioAllocator { /// /// # Arguments /// * `lambda` - Risk aversion parameter (higher = more conservative) + /// * `correlations` - Optional N x N correlation matrix. When `None`, assumes + /// independent assets (diagonal covariance). When provided, builds full + /// covariance: `Sigma[i][j] = corr[i][j] * vol_i * vol_j`. fn mean_variance( &self, assets: &[AssetInfo], total_capital: Decimal, lambda: f64, + ) -> Result> { + self.mean_variance_with_corr(assets, total_capital, lambda, None) + } + + /// Mean-Variance optimization with optional correlation matrix. + /// + /// When `correlations` is `Some`, builds the full covariance matrix from the + /// correlation matrix and per-asset volatilities. Falls back to diagonal + /// covariance if the correlation matrix is ill-conditioned. + fn mean_variance_with_corr( + &self, + assets: &[AssetInfo], + total_capital: Decimal, + lambda: f64, + correlations: Option<&DMatrix>, ) -> Result> { let n = assets.len(); // Expected returns vector let mu = DVector::from_vec(assets.iter().map(|a| a.expected_return).collect()); - // Covariance matrix — currently diagonal (independent assets). - // - // A diagonal covariance matrix assumes zero correlation between all asset pairs, - // which is a simplification. For full Markowitz optimization with correlation - // support, the following steps are needed: - // - // 1. Accept a `correlations: Option<&DMatrix>` parameter (N x N correlation matrix) - // 2. When provided, build full covariance: Sigma[i][j] = corr[i][j] * vol_i * vol_j - // 3. Ensure the correlation matrix is symmetric positive-definite (Cholesky check) - // 4. Fall back to diagonal if the matrix is ill-conditioned (det < epsilon) - // - // nalgebra's DMatrix is already available in this crate, so the implementation - // is straightforward once historical return data is available to estimate - // pairwise correlations (e.g., via a rolling Pearson correlation window). - let mut sigma = DMatrix::zeros(n, n); - for (i, asset) in assets.iter().enumerate() { - sigma[(i, i)] = asset.volatility.powi(2); - } + // Build covariance matrix + let mut sigma = if let Some(corr) = correlations { + // Validate dimensions + if corr.nrows() != n || corr.ncols() != n { + anyhow::bail!( + "Correlation matrix dimensions ({}, {}) do not match asset count {}", + corr.nrows(), + corr.ncols(), + n + ); + } + // Build full covariance: Sigma[i][j] = corr[i][j] * vol_i * vol_j + let mut cov = DMatrix::zeros(n, n); + for i in 0..n { + let vol_i = assets.get(i).map(|a| a.volatility).unwrap_or(0.0); + for j in 0..n { + let vol_j = assets.get(j).map(|a| a.volatility).unwrap_or(0.0); + let corr_ij = corr.get((i, j)).copied().unwrap_or(0.0); + if let Some(cell) = cov.get_mut((i, j)) { + *cell = corr_ij * vol_i * vol_j; + } + } + } + cov + } else { + // Diagonal covariance (independent assets) + let mut cov = DMatrix::zeros(n, n); + for (i, asset) in assets.iter().enumerate() { + if let Some(cell) = cov.get_mut((i, i)) { + *cell = asset.volatility.powi(2); + } + } + cov + }; // Add small regularization to diagonal for numerical stability for i in 0..n { @@ -642,4 +721,146 @@ mod tests { } } } + + /// Identity correlation matrix (diagonal = 1.0) should produce the same result + /// as the default diagonal covariance path (no correlations). + #[test] + fn test_mean_variance_identity_correlation_matches_diagonal() { + let allocator = PortfolioAllocator::new(AllocationMethod::MeanVariance { lambda: 2.0 }); + let assets = create_test_assets(); + let total_capital = Decimal::from(100_000); + let n = assets.len(); + + // Identity correlation matrix + let identity = DMatrix::identity(n, n); + + let alloc_diagonal = allocator.allocate(&assets, total_capital).unwrap(); + let alloc_identity = allocator + .allocate_with_correlations(&assets, total_capital, &identity) + .unwrap(); + + // Both should produce identical allocations + for asset in &assets { + let diag_val = alloc_diagonal.get(&asset.symbol).unwrap(); + let ident_val = alloc_identity.get(&asset.symbol).unwrap(); + let diff = (*diag_val - *ident_val).abs(); + assert!( + diff < Decimal::from_f64_retain(0.01).unwrap(), + "Symbol {} differs: diagonal={}, identity={}", + asset.symbol, + diag_val, + ident_val, + ); + } + } + + /// When two assets are highly correlated, the optimizer should allocate + /// differently compared to the uncorrelated (diagonal) case. + #[test] + fn test_correlated_allocation_differs_from_diagonal() { + let allocator = PortfolioAllocator::new(AllocationMethod::MeanVariance { lambda: 2.0 }); + let assets = create_test_assets(); // ES, NQ, ZN + let total_capital = Decimal::from(100_000); + let n = assets.len(); + + // High correlation between ES and NQ (both equity futures), low with ZN (bonds) + let corr_data = vec![ + 1.0, 0.90, 0.10, // ES row + 0.90, 1.0, 0.10, // NQ row + 0.10, 0.10, 1.0, // ZN row + ]; + let corr = DMatrix::from_row_slice(n, n, &corr_data); + + let alloc_diagonal = allocator.allocate(&assets, total_capital).unwrap(); + let alloc_correlated = allocator + .allocate_with_correlations(&assets, total_capital, &corr) + .unwrap(); + + // Correlated allocation should differ from diagonal + let mut any_differs = false; + for asset in &assets { + let diag_val = alloc_diagonal.get(&asset.symbol).unwrap(); + let corr_val = alloc_correlated.get(&asset.symbol).unwrap(); + if (*diag_val - *corr_val).abs() > Decimal::from_f64_retain(1.0).unwrap() { + any_differs = true; + } + } + assert!( + any_differs, + "Correlated allocation should differ from diagonal allocation" + ); + + // With high ES-NQ correlation, ZN (diversifier) should get relatively more weight + // compared to the diagonal case + let zn_diag = alloc_diagonal.get("ZN.FUT").unwrap(); + let zn_corr = alloc_correlated.get("ZN.FUT").unwrap(); + assert!( + zn_corr > zn_diag, + "ZN (uncorrelated diversifier) should get more weight with correlations: corr={}, diag={}", + zn_corr, + zn_diag, + ); + } + + /// Correlation matrix with wrong dimensions should return an error. + #[test] + fn test_invalid_correlation_matrix_dimensions() { + let allocator = PortfolioAllocator::new(AllocationMethod::MeanVariance { lambda: 2.0 }); + let assets = create_test_assets(); // 3 assets + let total_capital = Decimal::from(100_000); + + // 2x2 matrix for 3 assets — wrong dimensions + let bad_corr = DMatrix::identity(2, 2); + let result = allocator.allocate_with_correlations(&assets, total_capital, &bad_corr); + assert!(result.is_err(), "Should fail with mismatched dimensions"); + let err_msg = format!("{}", result.unwrap_err()); + assert!( + err_msg.contains("do not match"), + "Error should mention dimension mismatch: {}", + err_msg + ); + + // 4x4 matrix for 3 assets — also wrong + let bad_corr_large = DMatrix::identity(4, 4); + let result = allocator.allocate_with_correlations(&assets, total_capital, &bad_corr_large); + assert!(result.is_err(), "Should fail with oversized dimensions"); + } + + /// Non-square correlation matrix should also fail. + #[test] + fn test_non_square_correlation_matrix() { + let allocator = PortfolioAllocator::new(AllocationMethod::MeanVariance { lambda: 2.0 }); + let assets = create_test_assets(); + let total_capital = Decimal::from(100_000); + + // 3x2 matrix — not square + let bad_corr = DMatrix::zeros(3, 2); + let result = allocator.allocate_with_correlations(&assets, total_capital, &bad_corr); + assert!(result.is_err(), "Should fail with non-square matrix"); + } + + /// Allocate with correlations on non-MeanVariance methods should delegate + /// to standard allocate (correlations ignored). + #[test] + fn test_correlations_ignored_for_equal_weight() { + let allocator = PortfolioAllocator::new(AllocationMethod::EqualWeight); + let assets = create_test_assets(); + let total_capital = Decimal::from(100_000); + let n = assets.len(); + + let corr = DMatrix::identity(n, n); + let alloc_std = allocator.allocate(&assets, total_capital).unwrap(); + let alloc_corr = allocator + .allocate_with_correlations(&assets, total_capital, &corr) + .unwrap(); + + for asset in &assets { + assert_eq!( + alloc_std.get(&asset.symbol), + alloc_corr.get(&asset.symbol), + "EqualWeight should ignore correlations for {}", + asset.symbol + ); + } + } }