1use std::collections::HashMap;
7
8#[derive(Debug, Clone, PartialEq)]
10pub struct TaxLot {
11 pub lot_id: String,
13 pub symbol: String,
15 pub quantity: f64,
17 pub cost_basis: f64,
19 pub acquired_date: u64,
21}
22
23#[derive(Debug, Clone, PartialEq)]
25pub enum TaxMethod {
26 Fifo,
28 Lifo,
30 SpecificLot(String),
32 MinTax,
34 AverageCost,
36}
37
38#[derive(Debug, Clone, PartialEq)]
40pub struct RealizedGain {
41 pub lot_id: String,
43 pub quantity: f64,
45 pub proceeds: f64,
47 pub cost_basis: f64,
49 pub gain: f64,
51 pub holding_period_days: u64,
53 pub is_long_term: bool,
55}
56
57#[derive(Debug, Clone, PartialEq)]
60pub struct WashSale {
61 pub sold_lot_id: String,
63 pub repurchased_lot_id: String,
65 pub disallowed_loss: f64,
67}
68
69#[derive(Debug, Default)]
71pub struct TaxLotManager {
72 lots: HashMap<String, Vec<TaxLot>>,
74}
75
76impl TaxLotManager {
77 pub fn new() -> Self {
79 Self::default()
80 }
81
82 pub fn acquire(&mut self, lot: TaxLot) {
84 self.lots.entry(lot.symbol.clone()).or_default().push(lot);
85 }
86
87 pub fn dispose(
92 &mut self,
93 symbol: &str,
94 quantity: f64,
95 proceeds: f64,
96 date: u64,
97 method: &TaxMethod,
98 ) -> Vec<RealizedGain> {
99 let lots = match self.lots.get_mut(symbol) {
100 Some(v) if !v.is_empty() => v,
101 _ => return Vec::new(),
102 };
103
104 let order: Vec<usize> = match method {
106 TaxMethod::Fifo => (0..lots.len()).collect(),
107 TaxMethod::Lifo => (0..lots.len()).rev().collect(),
108 TaxMethod::SpecificLot(id) => lots
109 .iter()
110 .enumerate()
111 .filter(|(_, l)| &l.lot_id == id)
112 .map(|(i, _)| i)
113 .collect(),
114 TaxMethod::MinTax => {
115 let mut indexed: Vec<(usize, f64)> = lots
117 .iter()
118 .enumerate()
119 .map(|(i, l)| {
120 let price_per_share = if quantity > 0.0 {
121 proceeds / quantity
122 } else {
123 0.0
124 };
125 let cost_per_share = if l.quantity > 0.0 {
126 l.cost_basis / l.quantity
127 } else {
128 0.0
129 };
130 (i, price_per_share - cost_per_share)
131 })
132 .collect();
133 indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
134 indexed.into_iter().map(|(i, _)| i).collect()
135 }
136 TaxMethod::AverageCost => (0..lots.len()).collect(),
137 };
138
139 let mut remaining = quantity;
140 let mut gains: Vec<RealizedGain> = Vec::new();
141
142 if matches!(method, TaxMethod::AverageCost) {
143 let total_shares: f64 = lots.iter().map(|l| l.quantity).sum();
145 let total_basis: f64 = lots.iter().map(|l| l.cost_basis).sum();
146 let avg_cost_per_share = if total_shares > 0.0 {
147 total_basis / total_shares
148 } else {
149 0.0
150 };
151 let disposed = remaining.min(total_shares);
152 let allocated_basis = avg_cost_per_share * disposed;
153 let lot_proceeds = (proceeds / quantity) * disposed;
154 let gain_val = lot_proceeds - allocated_basis;
155 let earliest_lot = lots.iter().min_by_key(|l| l.acquired_date);
157 let (lot_id, holding_period_days) = earliest_lot
158 .map(|l| {
159 let hp = date.saturating_sub(l.acquired_date);
160 (l.lot_id.clone(), hp)
161 })
162 .unwrap_or_default();
163 gains.push(RealizedGain {
164 lot_id,
165 quantity: disposed,
166 proceeds: lot_proceeds,
167 cost_basis: allocated_basis,
168 gain: gain_val,
169 holding_period_days,
170 is_long_term: holding_period_days > 365,
171 });
172 let fraction = disposed / total_shares;
174 for lot in lots.iter_mut() {
175 lot.cost_basis -= lot.cost_basis * fraction;
176 lot.quantity -= lot.quantity * fraction;
177 }
178 lots.retain(|l| l.quantity > 1e-10);
179 return gains;
180 }
181
182 let mut consumed: Vec<usize> = Vec::new();
184
185 for idx in order {
186 if remaining <= 1e-10 {
187 break;
188 }
189 let lot = &mut lots[idx];
191 let taken = remaining.min(lot.quantity);
192 let lot_basis = (lot.cost_basis / lot.quantity) * taken;
193 let lot_proceeds = (proceeds / quantity) * taken;
194 let gain_val = lot_proceeds - lot_basis;
195 let holding_period_days = date.saturating_sub(lot.acquired_date);
196 gains.push(RealizedGain {
197 lot_id: lot.lot_id.clone(),
198 quantity: taken,
199 proceeds: lot_proceeds,
200 cost_basis: lot_basis,
201 gain: gain_val,
202 holding_period_days,
203 is_long_term: holding_period_days > 365,
204 });
205 lot.quantity -= taken;
206 lot.cost_basis -= lot_basis;
207 remaining -= taken;
208 if lot.quantity <= 1e-10 {
209 consumed.push(idx);
210 }
211 }
212
213 let mut consumed_sorted = consumed;
215 consumed_sorted.sort_unstable();
216 consumed_sorted.dedup();
217 for idx in consumed_sorted.into_iter().rev() {
218 lots.remove(idx);
219 }
220
221 gains
222 }
223
224 pub fn detect_wash_sales(
235 gains: &[RealizedGain],
236 recent_acquisitions: &[TaxLot],
237 window_days: u64,
238 ) -> Vec<WashSale> {
239 let mut wash_sales = Vec::new();
240 for gain in gains {
241 if gain.gain >= 0.0 {
242 continue; }
244 for acq in recent_acquisitions {
245 if acq.lot_id != gain.lot_id {
256 let _ = window_days; wash_sales.push(WashSale {
262 sold_lot_id: gain.lot_id.clone(),
263 repurchased_lot_id: acq.lot_id.clone(),
264 disallowed_loss: gain.gain.abs(),
265 });
266 }
267 }
268 }
269 wash_sales
270 }
271
272 pub fn unrealized_gain(&self, symbol: &str, current_price: f64) -> f64 {
274 self.lots
275 .get(symbol)
276 .map(|lots| {
277 lots.iter()
278 .map(|l| current_price * l.quantity - l.cost_basis)
279 .sum()
280 })
281 .unwrap_or(0.0)
282 }
283
284 pub fn total_cost_basis(&self, symbol: &str) -> f64 {
286 self.lots
287 .get(symbol)
288 .map(|lots| lots.iter().map(|l| l.cost_basis).sum())
289 .unwrap_or(0.0)
290 }
291
292 pub fn lots_for_symbol(&self, symbol: &str) -> Vec<&TaxLot> {
294 self.lots
295 .get(symbol)
296 .map(|lots| lots.iter().collect())
297 .unwrap_or_default()
298 }
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304
305 fn make_lot(id: &str, sym: &str, qty: f64, basis: f64, day: u64) -> TaxLot {
306 TaxLot {
307 lot_id: id.to_string(),
308 symbol: sym.to_string(),
309 quantity: qty,
310 cost_basis: basis,
311 acquired_date: day,
312 }
313 }
314
315 #[test]
316 fn test_acquire_and_lots_for_symbol() {
317 let mut mgr = TaxLotManager::new();
318 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1500.0, 100));
319 mgr.acquire(make_lot("L2", "AAPL", 5.0, 800.0, 200));
320 let lots = mgr.lots_for_symbol("AAPL");
321 assert_eq!(lots.len(), 2);
322 }
323
324 #[test]
325 fn test_total_cost_basis() {
326 let mut mgr = TaxLotManager::new();
327 mgr.acquire(make_lot("L1", "TSLA", 10.0, 1000.0, 100));
328 mgr.acquire(make_lot("L2", "TSLA", 5.0, 500.0, 200));
329 assert!((mgr.total_cost_basis("TSLA") - 1500.0).abs() < 1e-9);
330 }
331
332 #[test]
333 fn test_unrealized_gain() {
334 let mut mgr = TaxLotManager::new();
335 mgr.acquire(make_lot("L1", "MSFT", 10.0, 2000.0, 100));
336 assert!((mgr.unrealized_gain("MSFT", 250.0) - 500.0).abs() < 1e-9);
338 }
339
340 #[test]
341 fn test_fifo_dispose() {
342 let mut mgr = TaxLotManager::new();
343 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1000.0, 100));
345 mgr.acquire(make_lot("L2", "AAPL", 10.0, 1200.0, 200));
346 let gains = mgr.dispose("AAPL", 15.0, 2250.0, 400, &TaxMethod::Fifo);
348 assert_eq!(gains.len(), 2);
349 assert!((gains[0].proceeds - 1500.0).abs() < 1e-6);
351 assert!((gains[0].gain - 500.0).abs() < 1e-6);
352 assert!((gains[1].proceeds - 750.0).abs() < 1e-6);
354 assert!((gains[1].gain - 150.0).abs() < 1e-6);
355 let remaining = mgr.lots_for_symbol("AAPL");
357 assert_eq!(remaining.len(), 1);
358 assert_eq!(remaining[0].lot_id, "L2");
359 assert!((remaining[0].quantity - 5.0).abs() < 1e-9);
360 }
361
362 #[test]
363 fn test_lifo_dispose() {
364 let mut mgr = TaxLotManager::new();
365 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1000.0, 100));
366 mgr.acquire(make_lot("L2", "AAPL", 10.0, 1200.0, 200));
367 let gains = mgr.dispose("AAPL", 15.0, 1950.0, 300, &TaxMethod::Lifo);
369 assert_eq!(gains.len(), 2);
371 assert_eq!(gains[0].lot_id, "L2");
372 assert_eq!(gains[1].lot_id, "L1");
373 }
374
375 #[test]
376 fn test_specific_lot_dispose() {
377 let mut mgr = TaxLotManager::new();
378 mgr.acquire(make_lot("L1", "GOOG", 10.0, 10000.0, 100));
379 mgr.acquire(make_lot("L2", "GOOG", 10.0, 12000.0, 200));
380 let gains = mgr.dispose(
381 "GOOG",
382 5.0,
383 6000.0,
384 300,
385 &TaxMethod::SpecificLot("L2".to_string()),
386 );
387 assert_eq!(gains.len(), 1);
388 assert_eq!(gains[0].lot_id, "L2");
389 assert!((gains[0].cost_basis - 6000.0).abs() < 1e-6);
390 assert!((gains[0].gain).abs() < 1e-6); }
392
393 #[test]
394 fn test_average_cost_dispose() {
395 let mut mgr = TaxLotManager::new();
396 mgr.acquire(make_lot("L1", "ETF", 10.0, 100.0, 100));
400 mgr.acquire(make_lot("L2", "ETF", 10.0, 200.0, 200));
401 let gains = mgr.dispose("ETF", 10.0, 200.0, 300, &TaxMethod::AverageCost);
402 assert_eq!(gains.len(), 1);
403 assert!((gains[0].cost_basis - 150.0).abs() < 1e-6);
405 assert!((gains[0].gain - 50.0).abs() < 1e-6);
406 }
407
408 #[test]
409 fn test_min_tax_prefers_loss_lots() {
410 let mut mgr = TaxLotManager::new();
411 mgr.acquire(make_lot("L1", "XYZ", 10.0, 200.0, 100));
414 mgr.acquire(make_lot("L2", "XYZ", 10.0, 50.0, 200));
415 let gains = mgr.dispose("XYZ", 10.0, 150.0, 300, &TaxMethod::MinTax);
416 assert_eq!(gains[0].lot_id, "L1");
418 assert!(gains[0].gain < 0.0);
419 }
420
421 #[test]
422 fn test_long_term_classification() {
423 let mut mgr = TaxLotManager::new();
424 mgr.acquire(make_lot("L1", "BOND", 10.0, 1000.0, 0));
425 let gains = mgr.dispose("BOND", 10.0, 1200.0, 400, &TaxMethod::Fifo);
427 assert!(gains[0].is_long_term);
428 assert_eq!(gains[0].holding_period_days, 400);
429 }
430
431 #[test]
432 fn test_detect_wash_sales() {
433 let gains = vec![RealizedGain {
434 lot_id: "L1".to_string(),
435 quantity: 10.0,
436 proceeds: 900.0,
437 cost_basis: 1000.0,
438 gain: -100.0,
439 holding_period_days: 10,
440 is_long_term: false,
441 }];
442 let recent = vec![TaxLot {
443 lot_id: "L2".to_string(),
444 symbol: "AAPL".to_string(),
445 quantity: 10.0,
446 cost_basis: 950.0,
447 acquired_date: 5,
448 }];
449 let wash_sales = TaxLotManager::detect_wash_sales(&gains, &recent, 30);
450 assert_eq!(wash_sales.len(), 1);
451 assert_eq!(wash_sales[0].sold_lot_id, "L1");
452 assert_eq!(wash_sales[0].repurchased_lot_id, "L2");
453 assert!((wash_sales[0].disallowed_loss - 100.0).abs() < 1e-9);
454 }
455
456 #[test]
457 fn test_no_wash_sale_on_gain() {
458 let gains = vec![RealizedGain {
459 lot_id: "L1".to_string(),
460 quantity: 10.0,
461 proceeds: 1100.0,
462 cost_basis: 1000.0,
463 gain: 100.0,
464 holding_period_days: 10,
465 is_long_term: false,
466 }];
467 let recent = vec![TaxLot {
468 lot_id: "L2".to_string(),
469 symbol: "AAPL".to_string(),
470 quantity: 10.0,
471 cost_basis: 950.0,
472 acquired_date: 5,
473 }];
474 let wash_sales = TaxLotManager::detect_wash_sales(&gains, &recent, 30);
475 assert!(wash_sales.is_empty());
476 }
477
478 #[test]
479 fn test_empty_symbol() {
480 let mut mgr = TaxLotManager::new();
481 let gains = mgr.dispose("NOTHING", 10.0, 1000.0, 100, &TaxMethod::Fifo);
482 assert!(gains.is_empty());
483 assert_eq!(mgr.total_cost_basis("NOTHING"), 0.0);
484 assert_eq!(mgr.unrealized_gain("NOTHING", 100.0), 0.0);
485 }
486}