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 = if date >= l.acquired_date {
160 date - l.acquired_date
161 } else {
162 0
163 };
164 (l.lot_id.clone(), hp)
165 })
166 .unwrap_or_default();
167 gains.push(RealizedGain {
168 lot_id,
169 quantity: disposed,
170 proceeds: lot_proceeds,
171 cost_basis: allocated_basis,
172 gain: gain_val,
173 holding_period_days,
174 is_long_term: holding_period_days > 365,
175 });
176 let fraction = disposed / total_shares;
178 for lot in lots.iter_mut() {
179 lot.cost_basis -= lot.cost_basis * fraction;
180 lot.quantity -= lot.quantity * fraction;
181 }
182 lots.retain(|l| l.quantity > 1e-10);
183 return gains;
184 }
185
186 let mut consumed: Vec<usize> = Vec::new();
188
189 for idx in order {
190 if remaining <= 1e-10 {
191 break;
192 }
193 let lot = &mut lots[idx];
195 let taken = remaining.min(lot.quantity);
196 let lot_basis = (lot.cost_basis / lot.quantity) * taken;
197 let lot_proceeds = (proceeds / quantity) * taken;
198 let gain_val = lot_proceeds - lot_basis;
199 let holding_period_days = if date >= lot.acquired_date {
200 date - lot.acquired_date
201 } else {
202 0
203 };
204 gains.push(RealizedGain {
205 lot_id: lot.lot_id.clone(),
206 quantity: taken,
207 proceeds: lot_proceeds,
208 cost_basis: lot_basis,
209 gain: gain_val,
210 holding_period_days,
211 is_long_term: holding_period_days > 365,
212 });
213 lot.quantity -= taken;
214 lot.cost_basis -= lot_basis;
215 remaining -= taken;
216 if lot.quantity <= 1e-10 {
217 consumed.push(idx);
218 }
219 }
220
221 let mut consumed_sorted = consumed;
223 consumed_sorted.sort_unstable();
224 consumed_sorted.dedup();
225 for idx in consumed_sorted.into_iter().rev() {
226 lots.remove(idx);
227 }
228
229 gains
230 }
231
232 pub fn detect_wash_sales(
243 gains: &[RealizedGain],
244 recent_acquisitions: &[TaxLot],
245 window_days: u64,
246 ) -> Vec<WashSale> {
247 let mut wash_sales = Vec::new();
248 for gain in gains {
249 if gain.gain >= 0.0 {
250 continue; }
252 for acq in recent_acquisitions {
253 if acq.lot_id != gain.lot_id {
264 let _ = window_days; wash_sales.push(WashSale {
270 sold_lot_id: gain.lot_id.clone(),
271 repurchased_lot_id: acq.lot_id.clone(),
272 disallowed_loss: gain.gain.abs(),
273 });
274 }
275 }
276 }
277 wash_sales
278 }
279
280 pub fn unrealized_gain(&self, symbol: &str, current_price: f64) -> f64 {
282 self.lots
283 .get(symbol)
284 .map(|lots| {
285 lots.iter()
286 .map(|l| current_price * l.quantity - l.cost_basis)
287 .sum()
288 })
289 .unwrap_or(0.0)
290 }
291
292 pub fn total_cost_basis(&self, symbol: &str) -> f64 {
294 self.lots
295 .get(symbol)
296 .map(|lots| lots.iter().map(|l| l.cost_basis).sum())
297 .unwrap_or(0.0)
298 }
299
300 pub fn lots_for_symbol(&self, symbol: &str) -> Vec<&TaxLot> {
302 self.lots
303 .get(symbol)
304 .map(|lots| lots.iter().collect())
305 .unwrap_or_default()
306 }
307}
308
309#[cfg(test)]
310mod tests {
311 use super::*;
312
313 fn make_lot(id: &str, sym: &str, qty: f64, basis: f64, day: u64) -> TaxLot {
314 TaxLot {
315 lot_id: id.to_string(),
316 symbol: sym.to_string(),
317 quantity: qty,
318 cost_basis: basis,
319 acquired_date: day,
320 }
321 }
322
323 #[test]
324 fn test_acquire_and_lots_for_symbol() {
325 let mut mgr = TaxLotManager::new();
326 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1500.0, 100));
327 mgr.acquire(make_lot("L2", "AAPL", 5.0, 800.0, 200));
328 let lots = mgr.lots_for_symbol("AAPL");
329 assert_eq!(lots.len(), 2);
330 }
331
332 #[test]
333 fn test_total_cost_basis() {
334 let mut mgr = TaxLotManager::new();
335 mgr.acquire(make_lot("L1", "TSLA", 10.0, 1000.0, 100));
336 mgr.acquire(make_lot("L2", "TSLA", 5.0, 500.0, 200));
337 assert!((mgr.total_cost_basis("TSLA") - 1500.0).abs() < 1e-9);
338 }
339
340 #[test]
341 fn test_unrealized_gain() {
342 let mut mgr = TaxLotManager::new();
343 mgr.acquire(make_lot("L1", "MSFT", 10.0, 2000.0, 100));
344 assert!((mgr.unrealized_gain("MSFT", 250.0) - 500.0).abs() < 1e-9);
346 }
347
348 #[test]
349 fn test_fifo_dispose() {
350 let mut mgr = TaxLotManager::new();
351 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1000.0, 100));
353 mgr.acquire(make_lot("L2", "AAPL", 10.0, 1200.0, 200));
354 let gains = mgr.dispose("AAPL", 15.0, 2250.0, 400, &TaxMethod::Fifo);
356 assert_eq!(gains.len(), 2);
357 assert!((gains[0].proceeds - 1500.0).abs() < 1e-6);
359 assert!((gains[0].gain - 500.0).abs() < 1e-6);
360 assert!((gains[1].proceeds - 750.0).abs() < 1e-6);
362 assert!((gains[1].gain - 150.0).abs() < 1e-6);
363 let remaining = mgr.lots_for_symbol("AAPL");
365 assert_eq!(remaining.len(), 1);
366 assert_eq!(remaining[0].lot_id, "L2");
367 assert!((remaining[0].quantity - 5.0).abs() < 1e-9);
368 }
369
370 #[test]
371 fn test_lifo_dispose() {
372 let mut mgr = TaxLotManager::new();
373 mgr.acquire(make_lot("L1", "AAPL", 10.0, 1000.0, 100));
374 mgr.acquire(make_lot("L2", "AAPL", 10.0, 1200.0, 200));
375 let gains = mgr.dispose("AAPL", 15.0, 1950.0, 300, &TaxMethod::Lifo);
377 assert_eq!(gains.len(), 2);
379 assert_eq!(gains[0].lot_id, "L2");
380 assert_eq!(gains[1].lot_id, "L1");
381 }
382
383 #[test]
384 fn test_specific_lot_dispose() {
385 let mut mgr = TaxLotManager::new();
386 mgr.acquire(make_lot("L1", "GOOG", 10.0, 10000.0, 100));
387 mgr.acquire(make_lot("L2", "GOOG", 10.0, 12000.0, 200));
388 let gains = mgr.dispose(
389 "GOOG",
390 5.0,
391 6000.0,
392 300,
393 &TaxMethod::SpecificLot("L2".to_string()),
394 );
395 assert_eq!(gains.len(), 1);
396 assert_eq!(gains[0].lot_id, "L2");
397 assert!((gains[0].cost_basis - 6000.0).abs() < 1e-6);
398 assert!((gains[0].gain).abs() < 1e-6); }
400
401 #[test]
402 fn test_average_cost_dispose() {
403 let mut mgr = TaxLotManager::new();
404 mgr.acquire(make_lot("L1", "ETF", 10.0, 100.0, 100));
408 mgr.acquire(make_lot("L2", "ETF", 10.0, 200.0, 200));
409 let gains = mgr.dispose("ETF", 10.0, 200.0, 300, &TaxMethod::AverageCost);
410 assert_eq!(gains.len(), 1);
411 assert!((gains[0].cost_basis - 150.0).abs() < 1e-6);
413 assert!((gains[0].gain - 50.0).abs() < 1e-6);
414 }
415
416 #[test]
417 fn test_min_tax_prefers_loss_lots() {
418 let mut mgr = TaxLotManager::new();
419 mgr.acquire(make_lot("L1", "XYZ", 10.0, 200.0, 100));
422 mgr.acquire(make_lot("L2", "XYZ", 10.0, 50.0, 200));
423 let gains = mgr.dispose("XYZ", 10.0, 150.0, 300, &TaxMethod::MinTax);
424 assert_eq!(gains[0].lot_id, "L1");
426 assert!(gains[0].gain < 0.0);
427 }
428
429 #[test]
430 fn test_long_term_classification() {
431 let mut mgr = TaxLotManager::new();
432 mgr.acquire(make_lot("L1", "BOND", 10.0, 1000.0, 0));
433 let gains = mgr.dispose("BOND", 10.0, 1200.0, 400, &TaxMethod::Fifo);
435 assert!(gains[0].is_long_term);
436 assert_eq!(gains[0].holding_period_days, 400);
437 }
438
439 #[test]
440 fn test_detect_wash_sales() {
441 let gains = vec![RealizedGain {
442 lot_id: "L1".to_string(),
443 quantity: 10.0,
444 proceeds: 900.0,
445 cost_basis: 1000.0,
446 gain: -100.0,
447 holding_period_days: 10,
448 is_long_term: false,
449 }];
450 let recent = vec![TaxLot {
451 lot_id: "L2".to_string(),
452 symbol: "AAPL".to_string(),
453 quantity: 10.0,
454 cost_basis: 950.0,
455 acquired_date: 5,
456 }];
457 let wash_sales = TaxLotManager::detect_wash_sales(&gains, &recent, 30);
458 assert_eq!(wash_sales.len(), 1);
459 assert_eq!(wash_sales[0].sold_lot_id, "L1");
460 assert_eq!(wash_sales[0].repurchased_lot_id, "L2");
461 assert!((wash_sales[0].disallowed_loss - 100.0).abs() < 1e-9);
462 }
463
464 #[test]
465 fn test_no_wash_sale_on_gain() {
466 let gains = vec![RealizedGain {
467 lot_id: "L1".to_string(),
468 quantity: 10.0,
469 proceeds: 1100.0,
470 cost_basis: 1000.0,
471 gain: 100.0,
472 holding_period_days: 10,
473 is_long_term: false,
474 }];
475 let recent = vec![TaxLot {
476 lot_id: "L2".to_string(),
477 symbol: "AAPL".to_string(),
478 quantity: 10.0,
479 cost_basis: 950.0,
480 acquired_date: 5,
481 }];
482 let wash_sales = TaxLotManager::detect_wash_sales(&gains, &recent, 30);
483 assert!(wash_sales.is_empty());
484 }
485
486 #[test]
487 fn test_empty_symbol() {
488 let mut mgr = TaxLotManager::new();
489 let gains = mgr.dispose("NOTHING", 10.0, 1000.0, 100, &TaxMethod::Fifo);
490 assert!(gains.is_empty());
491 assert_eq!(mgr.total_cost_basis("NOTHING"), 0.0);
492 assert_eq!(mgr.unrealized_gain("NOTHING", 100.0), 0.0);
493 }
494}