feat(trading_agent): correlation matrix support in Markowitz allocation

Add optional correlation matrix parameter to mean-variance optimization.
When provided, builds full covariance matrix (Sigma[i][j] = corr[i][j] *
vol_i * vol_j) instead of diagonal-only. Existing API unchanged — callers
pass None by default. New allocate_with_correlations() public method for
correlated optimization. Five new tests: identity-matches-diagonal,
correlated-differs-from-diagonal, invalid dimensions, non-square matrix,
and non-MeanVariance delegation.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-23 10:33:14 +01:00
parent 940aae71b1
commit 18e00fff12
2 changed files with 839 additions and 18 deletions

View File

@@ -697,6 +697,310 @@ pub struct RegistryStatistics {
pub earliest_training_date: Option<DateTime<Utc>>,
}
/// 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<Utc>,
/// Optional reason for the rollback
pub reason: Option<String>,
}
/// Version entry stored per model in the in-memory registry
#[derive(Debug, Clone)]
struct VersionEntry {
/// All registered versions (ordered by registration time)
versions: Vec<ModelVersionMetadata>,
/// 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<RwLock<HashMap<String, VersionEntry>>>,
/// Audit log of rollback events
rollback_log: Arc<RwLock<Vec<RollbackEvent>>>,
}
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<String>,
) -> 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<Vec<(String, bool)>> {
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<ModelVersionMetadata> {
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<ModelVersionMetadata> {
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<String> {
self.models.read().await.keys().cloned().collect()
}
/// Get the rollback audit log.
pub async fn get_rollback_log(&self) -> Vec<RollbackEvent> {
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<RollbackEvent> {
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));
}
}

View File

@@ -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<f64>,
) -> Result<HashMap<String, Decimal>> {
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<AssetInfo> = 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<HashMap<String, Decimal>> {
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<f64>>,
) -> Result<HashMap<String, Decimal>> {
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<f64>>` 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
);
}
}
}