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:
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user