1#[derive(Debug, Clone, PartialEq)]
9pub struct TargetAllocation {
10 pub symbol: String,
12 pub target_weight: f64,
14 pub min_weight: f64,
16 pub max_weight: f64,
18}
19
20#[derive(Debug, Clone, PartialEq)]
22pub struct PortfolioPosition {
23 pub symbol: String,
25 pub market_value: f64,
27 pub current_weight: f64,
29}
30
31#[derive(Debug, Clone, PartialEq)]
33pub struct RebalanceDrift {
34 pub symbol: String,
36 pub current_weight: f64,
38 pub target_weight: f64,
40 pub drift: f64,
42 pub abs_drift: f64,
44}
45
46#[derive(Debug, Clone, PartialEq)]
48pub enum RebalanceTrigger {
49 ThresholdBreach(f64),
51 CalendarBased(u32),
53 BothConditions,
55}
56
57#[derive(Debug, Clone, PartialEq, Eq)]
59pub enum TradeDirection {
60 Buy,
62 Sell,
64}
65
66#[derive(Debug, Clone, PartialEq)]
68pub struct RebalanceTrade {
69 pub symbol: String,
71 pub direction: TradeDirection,
73 pub amount: f64,
75 pub target_weight: f64,
77}
78
79pub struct Rebalancer;
81
82impl Rebalancer {
83 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 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 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 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 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 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 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 let gain_map: std::collections::HashMap<&str, f64> = gains
215 .iter()
216 .map(|(s, g)| (s.as_str(), *g))
217 .collect();
218
219 trades.sort_by(|a, b| {
222 if a.direction == TradeDirection::Buy && b.direction == TradeDirection::Buy {
224 return std::cmp::Ordering::Equal;
225 }
226 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 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 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), ("MSFT".to_string(), 1000.0), ];
394 let trades = Rebalancer::tax_aware_rebalance(&positions, &targets, &gains);
395 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}