use std::collections::HashMap;
#[derive(Debug, Clone)]
pub enum RebalancingTrigger {
Calendar {
frequency_days: u32,
},
Threshold {
drift_pct: f64,
},
BandBased {
inner_band: f64,
outer_band: f64,
},
Hybrid {
max_days: u32,
threshold_pct: f64,
},
}
#[derive(Debug, Clone, Default)]
pub struct RebalancingCost {
pub transaction_costs: f64,
pub market_impact: f64,
pub tax_cost: f64,
pub total: f64,
}
#[derive(Debug, Clone)]
pub struct HoldingDrift {
pub symbol: String,
pub target_weight: f64,
pub current_weight: f64,
pub drift_pct: f64,
pub requires_rebalance: bool,
}
#[derive(Debug, Clone)]
pub struct RebalancingPlan {
pub trades: Vec<(String, f64, bool)>,
pub total_turnover: f64,
pub estimated_cost: RebalancingCost,
pub trigger_reason: String,
}
pub struct RebalancingEngine;
impl RebalancingEngine {
pub fn compute_drifts(
targets: &HashMap<String, f64>,
current_values: &HashMap<String, f64>,
) -> Vec<HoldingDrift> {
let total_value: f64 = current_values.values().sum();
if total_value <= 0.0 {
return targets
.iter()
.map(|(sym, &tw)| HoldingDrift {
symbol: sym.clone(),
target_weight: tw,
current_weight: 0.0,
drift_pct: -tw,
requires_rebalance: false,
})
.collect();
}
let mut drifts = Vec::with_capacity(targets.len());
for (sym, &target_weight) in targets {
let value = current_values.get(sym).copied().unwrap_or(0.0);
let current_weight = value / total_value;
let drift_pct = current_weight - target_weight;
drifts.push(HoldingDrift {
symbol: sym.clone(),
target_weight,
current_weight,
drift_pct,
requires_rebalance: false, });
}
drifts
}
pub fn should_rebalance(
drifts: &[HoldingDrift],
trigger: &RebalancingTrigger,
days_since_last: u32,
) -> bool {
match trigger {
RebalancingTrigger::Calendar { frequency_days } => {
days_since_last >= *frequency_days
}
RebalancingTrigger::Threshold { drift_pct } => drifts
.iter()
.any(|d| d.drift_pct.abs() >= *drift_pct),
RebalancingTrigger::BandBased { outer_band, .. } => drifts
.iter()
.any(|d| d.drift_pct.abs() >= *outer_band),
RebalancingTrigger::Hybrid {
max_days,
threshold_pct,
} => {
days_since_last >= *max_days
|| drifts.iter().any(|d| d.drift_pct.abs() >= *threshold_pct)
}
}
}
pub fn generate_plan(
drifts: &[HoldingDrift],
portfolio_value: f64,
trigger: &RebalancingTrigger,
days_since_last: u32,
) -> RebalancingPlan {
let trigger_reason = Self::trigger_reason(trigger, drifts, days_since_last);
let mut trades: Vec<(String, f64, bool)> = Vec::new();
let mut total_turnover = 0.0_f64;
for drift in drifts {
let notional = drift.drift_pct.abs() * portfolio_value;
if notional < 1e-8 {
continue;
}
let is_buy = drift.drift_pct < 0.0; trades.push((drift.symbol.clone(), notional, is_buy));
total_turnover += notional;
}
total_turnover /= portfolio_value.max(1.0);
RebalancingPlan {
trades,
total_turnover,
estimated_cost: RebalancingCost::default(),
trigger_reason,
}
}
pub fn estimate_cost(
plan: &RebalancingPlan,
bid_ask_bps: f64,
tax_rate: f64,
_holding_days: u32,
) -> RebalancingCost {
let total_notional: f64 = plan.trades.iter().map(|(_, n, _)| n).sum();
let transaction_costs = total_notional * bid_ask_bps / 10_000.0;
let market_impact = transaction_costs * 0.5; let tax_cost = total_notional * tax_rate * 0.02; let total = transaction_costs + market_impact + tax_cost;
RebalancingCost {
transaction_costs,
market_impact,
tax_cost,
total,
}
}
pub fn net_benefit(
plan: &RebalancingPlan,
tracking_error_reduction: f64,
annual_alpha_bps: f64,
) -> f64 {
let alpha_gain = tracking_error_reduction * annual_alpha_bps / 10_000.0;
alpha_gain - plan.estimated_cost.total
}
pub fn tax_lot_optimization(
plan: &RebalancingPlan,
unrealized_gains: &HashMap<String, f64>,
) -> RebalancingPlan {
let mut sells: Vec<(String, f64, bool)> = plan
.trades
.iter()
.filter(|(_, _, is_buy)| !is_buy)
.cloned()
.collect();
sells.sort_by(|(sym_a, _, _), (sym_b, _, _)| {
let ga = unrealized_gains.get(sym_a).copied().unwrap_or(0.0);
let gb = unrealized_gains.get(sym_b).copied().unwrap_or(0.0);
ga.partial_cmp(&gb).unwrap_or(std::cmp::Ordering::Equal)
});
let buys: Vec<(String, f64, bool)> = plan
.trades
.iter()
.filter(|(_, _, is_buy)| *is_buy)
.cloned()
.collect();
let mut trades = sells;
trades.extend(buys);
RebalancingPlan {
trades,
total_turnover: plan.total_turnover,
estimated_cost: plan.estimated_cost.clone(),
trigger_reason: plan.trigger_reason.clone(),
}
}
pub fn minimum_variance_rebalance(
targets: &HashMap<String, f64>,
current: &HashMap<String, f64>,
cov_matrix: &[Vec<f64>],
) -> RebalancingPlan {
let symbols: Vec<&String> = targets.keys().collect();
let n = symbols.len();
let total_value: f64 = current.values().sum();
let mut trades: Vec<(String, f64, bool)> = Vec::new();
let mut total_turnover = 0.0;
for (i, sym) in symbols.iter().enumerate() {
let target_w = targets.get(*sym).copied().unwrap_or(0.0);
let current_v = current.get(*sym).copied().unwrap_or(0.0);
let current_w = if total_value > 0.0 {
current_v / total_value
} else {
0.0
};
let mv: f64 = if i < cov_matrix.len() {
cov_matrix[i].iter().sum::<f64>() / n.max(1) as f64
} else {
1.0
};
let drift = target_w - current_w;
let scale = (1.0 / mv.abs().max(1e-10)).min(2.0);
let notional = drift.abs() * total_value * scale;
if notional < 1e-8 {
continue;
}
let is_buy = drift > 0.0;
trades.push(((*sym).clone(), notional, is_buy));
total_turnover += notional;
}
total_turnover /= total_value.max(1.0);
RebalancingPlan {
trades,
total_turnover,
estimated_cost: RebalancingCost::default(),
trigger_reason: "MinimumVariance".to_string(),
}
}
fn trigger_reason(
trigger: &RebalancingTrigger,
drifts: &[HoldingDrift],
days_since_last: u32,
) -> String {
match trigger {
RebalancingTrigger::Calendar { frequency_days } => {
format!("Calendar: {} days elapsed (freq={})", days_since_last, frequency_days)
}
RebalancingTrigger::Threshold { drift_pct } => {
let max_drift = drifts
.iter()
.map(|d| d.drift_pct.abs())
.fold(0.0_f64, f64::max);
format!("Threshold: max drift {:.4} >= {:.4}", max_drift, drift_pct)
}
RebalancingTrigger::BandBased { inner_band, outer_band } => {
format!(
"BandBased: inner={:.4} outer={:.4}",
inner_band, outer_band
)
}
RebalancingTrigger::Hybrid { max_days: _, threshold_pct } => {
format!(
"Hybrid: days={} threshold={:.4}",
days_since_last, threshold_pct
)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_targets() -> HashMap<String, f64> {
[("AAPL".into(), 0.40), ("MSFT".into(), 0.60)]
.iter()
.cloned()
.collect()
}
fn sample_values() -> HashMap<String, f64> {
[("AAPL".into(), 350.0_f64), ("MSFT".into(), 650.0)]
.iter()
.cloned()
.collect()
}
#[test]
fn test_compute_drifts_balanced() {
let targets = sample_targets();
let values = sample_values();
let drifts = RebalancingEngine::compute_drifts(&targets, &values);
for d in &drifts {
if d.symbol == "AAPL" {
assert!((d.drift_pct - (-0.05)).abs() < 1e-9);
}
}
}
#[test]
fn test_should_rebalance_threshold() {
let targets = sample_targets();
let values = sample_values();
let drifts = RebalancingEngine::compute_drifts(&targets, &values);
let trigger = RebalancingTrigger::Threshold { drift_pct: 0.03 };
assert!(RebalancingEngine::should_rebalance(&drifts, &trigger, 0));
let trigger_high = RebalancingTrigger::Threshold { drift_pct: 0.10 };
assert!(!RebalancingEngine::should_rebalance(&drifts, &trigger_high, 0));
}
#[test]
fn test_generate_plan() {
let targets = sample_targets();
let values = sample_values();
let drifts = RebalancingEngine::compute_drifts(&targets, &values);
let trigger = RebalancingTrigger::Threshold { drift_pct: 0.03 };
let plan = RebalancingEngine::generate_plan(&drifts, 1_000.0, &trigger, 0);
assert!(!plan.trades.is_empty());
assert!(plan.total_turnover > 0.0);
}
}