Skip to main content

fin_primitives/rebalancing/
mod.rs

1//! Portfolio rebalancing: drift calculation, threshold and calendar triggers,
2//! trade generation, turnover estimation, and tax-aware rebalancing.
3//!
4//! [`Rebalancer`] is a stateless utility — all methods take explicit snapshots of
5//! positions and targets so callers maintain full control over state.
6
7/// Target allocation band for a single asset.
8#[derive(Debug, Clone, PartialEq)]
9pub struct TargetAllocation {
10    /// Ticker symbol.
11    pub symbol: String,
12    /// Ideal portfolio weight (0.0–1.0).
13    pub target_weight: f64,
14    /// Lower drift band — weight below this triggers rebalancing.
15    pub min_weight: f64,
16    /// Upper drift band — weight above this triggers rebalancing.
17    pub max_weight: f64,
18}
19
20/// Current snapshot of a position's market value and weight.
21#[derive(Debug, Clone, PartialEq)]
22pub struct PortfolioPosition {
23    /// Ticker symbol.
24    pub symbol: String,
25    /// Current market value of the position.
26    pub market_value: f64,
27    /// Current portfolio weight (market_value / total_portfolio_value).
28    pub current_weight: f64,
29}
30
31/// Drift metrics for a single asset.
32#[derive(Debug, Clone, PartialEq)]
33pub struct RebalanceDrift {
34    /// Ticker symbol.
35    pub symbol: String,
36    /// Current portfolio weight.
37    pub current_weight: f64,
38    /// Target portfolio weight.
39    pub target_weight: f64,
40    /// Signed drift: current_weight − target_weight.
41    pub drift: f64,
42    /// Absolute drift (|drift|).
43    pub abs_drift: f64,
44}
45
46/// Condition that triggers a rebalance.
47#[derive(Debug, Clone, PartialEq)]
48pub enum RebalanceTrigger {
49    /// Rebalance if any asset's absolute drift exceeds `max_drift`.
50    ThresholdBreach(f64),
51    /// Rebalance if `days_since_last` >= `interval_days`.
52    CalendarBased(u32),
53    /// Rebalance if either the threshold OR the calendar condition is met.
54    BothConditions,
55}
56
57/// Direction of a rebalance trade.
58#[derive(Debug, Clone, PartialEq, Eq)]
59pub enum TradeDirection {
60    /// Buy additional shares to increase allocation.
61    Buy,
62    /// Sell shares to decrease allocation.
63    Sell,
64}
65
66/// A single trade required to move toward target allocation.
67#[derive(Debug, Clone, PartialEq)]
68pub struct RebalanceTrade {
69    /// Ticker symbol.
70    pub symbol: String,
71    /// Whether to buy or sell.
72    pub direction: TradeDirection,
73    /// Dollar amount of the trade (always positive).
74    pub amount: f64,
75    /// Target weight this trade is working toward.
76    pub target_weight: f64,
77}
78
79/// Stateless portfolio rebalancer.
80pub struct Rebalancer;
81
82impl Rebalancer {
83    /// Compute the drift of each position relative to its target allocation.
84    ///
85    /// Positions without a matching target are assigned a target weight of 0.0.
86    pub fn compute_drift(
87        positions: &[PortfolioPosition],
88        targets: &[TargetAllocation],
89    ) -> Vec<RebalanceDrift> {
90        positions
91            .iter()
92            .map(|pos| {
93                let target_weight = targets
94                    .iter()
95                    .find(|t| t.symbol == pos.symbol)
96                    .map(|t| t.target_weight)
97                    .unwrap_or(0.0);
98                let drift = pos.current_weight - target_weight;
99                RebalanceDrift {
100                    symbol: pos.symbol.clone(),
101                    current_weight: pos.current_weight,
102                    target_weight,
103                    drift,
104                    abs_drift: drift.abs(),
105                }
106            })
107            .collect()
108    }
109
110    /// Determine whether a rebalance should be triggered.
111    ///
112    /// - [`RebalanceTrigger::ThresholdBreach`]: true if any asset's `abs_drift` exceeds the threshold.
113    /// - [`RebalanceTrigger::CalendarBased`]: true if `days_since_last >= interval_days`.
114    /// - [`RebalanceTrigger::BothConditions`]: true if either sub-condition fires.
115    pub fn should_rebalance(
116        drift: &[RebalanceDrift],
117        trigger: &RebalanceTrigger,
118        days_since_last: u32,
119    ) -> bool {
120        match trigger {
121            RebalanceTrigger::ThresholdBreach(max_drift) => {
122                drift.iter().any(|d| d.abs_drift > *max_drift)
123            }
124            RebalanceTrigger::CalendarBased(interval_days) => {
125                days_since_last >= *interval_days
126            }
127            RebalanceTrigger::BothConditions => {
128                // "Both" in common parlance means "either condition" (union trigger).
129                let threshold_hit = drift.iter().any(|d| d.abs_drift > 0.05);
130                let calendar_hit = days_since_last >= 90;
131                threshold_hit || calendar_hit
132            }
133        }
134    }
135
136    /// Generate the trades needed to bring each position to its target weight.
137    ///
138    /// Proportional rebalancing: each asset's target dollar value = target_weight × total_value.
139    /// Assets without a current position are also included (new buys).
140    pub fn generate_trades(
141        positions: &[PortfolioPosition],
142        targets: &[TargetAllocation],
143        total_value: f64,
144    ) -> Vec<RebalanceTrade> {
145        let mut trades: Vec<RebalanceTrade> = Vec::new();
146
147        for target in targets {
148            let current_value = positions
149                .iter()
150                .find(|p| p.symbol == target.symbol)
151                .map(|p| p.market_value)
152                .unwrap_or(0.0);
153            let desired_value = target.target_weight * total_value;
154            let diff = desired_value - current_value;
155
156            if diff.abs() < 1e-6 {
157                continue;
158            }
159
160            trades.push(RebalanceTrade {
161                symbol: target.symbol.clone(),
162                direction: if diff > 0.0 {
163                    TradeDirection::Buy
164                } else {
165                    TradeDirection::Sell
166                },
167                amount: diff.abs(),
168                target_weight: target.target_weight,
169            });
170        }
171
172        // Also handle positions with no matching target (sell down to zero).
173        for pos in positions {
174            let has_target = targets.iter().any(|t| t.symbol == pos.symbol);
175            if !has_target && pos.market_value > 1e-6 {
176                trades.push(RebalanceTrade {
177                    symbol: pos.symbol.clone(),
178                    direction: TradeDirection::Sell,
179                    amount: pos.market_value,
180                    target_weight: 0.0,
181                });
182            }
183        }
184
185        trades
186    }
187
188    /// Estimate portfolio turnover as a fraction of total portfolio value.
189    ///
190    /// Turnover = sum(trade amounts) / total_value.
191    /// Values above 1.0 are possible for large rebalances.
192    pub fn estimated_turnover(trades: &[RebalanceTrade], total_value: f64) -> f64 {
193        if total_value <= 0.0 {
194            return 0.0;
195        }
196        let total_traded: f64 = trades.iter().map(|t| t.amount).sum();
197        total_traded / total_value
198    }
199
200    /// Tax-aware rebalancing: prefer selling positions with losses (negative gain)
201    /// before selling positions with gains, to minimize realized tax liability.
202    ///
203    /// `gains` is a slice of `(symbol, unrealized_gain)` pairs. Symbols with negative
204    /// gains are prioritized for selling; symbols with positive gains are sold last.
205    pub fn tax_aware_rebalance(
206        positions: &[PortfolioPosition],
207        targets: &[TargetAllocation],
208        gains: &[(String, f64)],
209    ) -> Vec<RebalanceTrade> {
210        let total_value: f64 = positions.iter().map(|p| p.market_value).sum();
211        let mut trades = Self::generate_trades(positions, targets, total_value);
212
213        // Separate sells by loss vs gain status.
214        let gain_map: std::collections::HashMap<&str, f64> = gains
215            .iter()
216            .map(|(s, g)| (s.as_str(), *g))
217            .collect();
218
219        // Sort: sells — loss positions first (smallest/most negative gain first),
220        //       then gains; buys keep their natural order.
221        trades.sort_by(|a, b| {
222            // Both buys: equal priority.
223            if a.direction == TradeDirection::Buy && b.direction == TradeDirection::Buy {
224                return std::cmp::Ordering::Equal;
225            }
226            // Sells before buys in the ordering (buys come after).
227            if a.direction == TradeDirection::Sell && b.direction == TradeDirection::Buy {
228                return std::cmp::Ordering::Less;
229            }
230            if a.direction == TradeDirection::Buy && b.direction == TradeDirection::Sell {
231                return std::cmp::Ordering::Greater;
232            }
233            // Both sells: sort by gain ascending (losses first).
234            let ga = gain_map.get(a.symbol.as_str()).copied().unwrap_or(0.0);
235            let gb = gain_map.get(b.symbol.as_str()).copied().unwrap_or(0.0);
236            ga.partial_cmp(&gb).unwrap_or(std::cmp::Ordering::Equal)
237        });
238
239        trades
240    }
241}
242
243#[cfg(test)]
244mod tests {
245    use super::*;
246
247    fn pos(sym: &str, mv: f64, w: f64) -> PortfolioPosition {
248        PortfolioPosition {
249            symbol: sym.to_string(),
250            market_value: mv,
251            current_weight: w,
252        }
253    }
254
255    fn tgt(sym: &str, tw: f64) -> TargetAllocation {
256        TargetAllocation {
257            symbol: sym.to_string(),
258            target_weight: tw,
259            min_weight: tw - 0.05,
260            max_weight: tw + 0.05,
261        }
262    }
263
264    #[test]
265    fn test_compute_drift_basic() {
266        let positions = vec![
267            pos("AAPL", 6000.0, 0.60),
268            pos("MSFT", 4000.0, 0.40),
269        ];
270        let targets = vec![tgt("AAPL", 0.50), tgt("MSFT", 0.50)];
271        let drift = Rebalancer::compute_drift(&positions, &targets);
272        assert_eq!(drift.len(), 2);
273        let aapl = drift.iter().find(|d| d.symbol == "AAPL").unwrap();
274        assert!((aapl.drift - 0.10).abs() < 1e-9);
275        assert!((aapl.abs_drift - 0.10).abs() < 1e-9);
276        let msft = drift.iter().find(|d| d.symbol == "MSFT").unwrap();
277        assert!((msft.drift - (-0.10)).abs() < 1e-9);
278    }
279
280    #[test]
281    fn test_should_rebalance_threshold() {
282        let drift = vec![RebalanceDrift {
283            symbol: "AAPL".to_string(),
284            current_weight: 0.60,
285            target_weight: 0.50,
286            drift: 0.10,
287            abs_drift: 0.10,
288        }];
289        assert!(Rebalancer::should_rebalance(
290            &drift,
291            &RebalanceTrigger::ThresholdBreach(0.05),
292            0
293        ));
294        assert!(!Rebalancer::should_rebalance(
295            &drift,
296            &RebalanceTrigger::ThresholdBreach(0.15),
297            0
298        ));
299    }
300
301    #[test]
302    fn test_should_rebalance_calendar() {
303        let drift: Vec<RebalanceDrift> = Vec::new();
304        assert!(Rebalancer::should_rebalance(
305            &drift,
306            &RebalanceTrigger::CalendarBased(90),
307            90
308        ));
309        assert!(!Rebalancer::should_rebalance(
310            &drift,
311            &RebalanceTrigger::CalendarBased(90),
312            89
313        ));
314    }
315
316    #[test]
317    fn test_generate_trades_balanced() {
318        // Already balanced — no trades.
319        let positions = vec![
320            pos("AAPL", 5000.0, 0.50),
321            pos("MSFT", 5000.0, 0.50),
322        ];
323        let targets = vec![tgt("AAPL", 0.50), tgt("MSFT", 0.50)];
324        let trades = Rebalancer::generate_trades(&positions, &targets, 10000.0);
325        assert!(trades.is_empty());
326    }
327
328    #[test]
329    fn test_generate_trades_unbalanced() {
330        let positions = vec![
331            pos("AAPL", 7000.0, 0.70),
332            pos("MSFT", 3000.0, 0.30),
333        ];
334        let targets = vec![tgt("AAPL", 0.50), tgt("MSFT", 0.50)];
335        let trades = Rebalancer::generate_trades(&positions, &targets, 10000.0);
336        let aapl_trade = trades.iter().find(|t| t.symbol == "AAPL").unwrap();
337        let msft_trade = trades.iter().find(|t| t.symbol == "MSFT").unwrap();
338        assert_eq!(aapl_trade.direction, TradeDirection::Sell);
339        assert!((aapl_trade.amount - 2000.0).abs() < 1e-6);
340        assert_eq!(msft_trade.direction, TradeDirection::Buy);
341        assert!((msft_trade.amount - 2000.0).abs() < 1e-6);
342    }
343
344    #[test]
345    fn test_sell_untargeted_position() {
346        let positions = vec![
347            pos("AAPL", 5000.0, 0.50),
348            pos("JUNK", 5000.0, 0.50),
349        ];
350        let targets = vec![tgt("AAPL", 1.0)];
351        let trades = Rebalancer::generate_trades(&positions, &targets, 10000.0);
352        let junk_trade = trades.iter().find(|t| t.symbol == "JUNK").unwrap();
353        assert_eq!(junk_trade.direction, TradeDirection::Sell);
354        assert!((junk_trade.amount - 5000.0).abs() < 1e-6);
355    }
356
357    #[test]
358    fn test_estimated_turnover() {
359        let trades = vec![
360            RebalanceTrade {
361                symbol: "AAPL".to_string(),
362                direction: TradeDirection::Sell,
363                amount: 1000.0,
364                target_weight: 0.50,
365            },
366            RebalanceTrade {
367                symbol: "MSFT".to_string(),
368                direction: TradeDirection::Buy,
369                amount: 1000.0,
370                target_weight: 0.50,
371            },
372        ];
373        let turnover = Rebalancer::estimated_turnover(&trades, 10000.0);
374        assert!((turnover - 0.20).abs() < 1e-9);
375    }
376
377    #[test]
378    fn test_estimated_turnover_zero_value() {
379        let trades: Vec<RebalanceTrade> = Vec::new();
380        assert_eq!(Rebalancer::estimated_turnover(&trades, 0.0), 0.0);
381    }
382
383    #[test]
384    fn test_tax_aware_rebalance_sells_losses_first() {
385        let positions = vec![
386            pos("AAPL", 4000.0, 0.40),
387            pos("MSFT", 6000.0, 0.60),
388        ];
389        let targets = vec![tgt("AAPL", 0.50), tgt("MSFT", 0.50)];
390        let gains = vec![
391            ("AAPL".to_string(), -500.0),  // loss — should sell first
392            ("MSFT".to_string(), 1000.0),  // gain — sell last
393        ];
394        let trades = Rebalancer::tax_aware_rebalance(&positions, &targets, &gains);
395        // MSFT should be sold; AAPL should be bought.
396        let sell = trades.iter().find(|t| t.direction == TradeDirection::Sell).unwrap();
397        assert_eq!(sell.symbol, "MSFT");
398    }
399
400    #[test]
401    fn test_compute_drift_missing_target() {
402        let positions = vec![pos("AAPL", 5000.0, 1.0)];
403        let targets: Vec<TargetAllocation> = Vec::new();
404        let drift = Rebalancer::compute_drift(&positions, &targets);
405        assert_eq!(drift[0].target_weight, 0.0);
406        assert!((drift[0].drift - 1.0).abs() < 1e-9);
407    }
408}