1use std::collections::HashMap;
13
14#[derive(Debug, Clone)]
18pub struct ExecutionCost {
19 pub commission_usd: f64,
21 pub spread_cost_usd: f64,
23 pub market_impact_usd: f64,
25 pub total_cost_usd: f64,
27 pub cost_bps: f64,
29}
30
31#[derive(Debug, Clone)]
35pub struct CostParams {
36 pub commission_per_share: f64,
38 pub spread_bps: f64,
40 pub impact_coefficient: f64,
42 pub avg_daily_volume: f64,
44}
45
46pub struct CostModel;
61
62impl CostModel {
63 pub fn estimate(notional_usd: f64, shares: f64, price: f64, params: &CostParams) -> ExecutionCost {
67 let _ = price; let commission_usd = params.commission_per_share * shares;
70
71 let spread_cost_usd = (params.spread_bps / 10_000.0) * notional_usd;
72
73 let impact_bps = if params.avg_daily_volume > 0.0 {
74 params.impact_coefficient * (shares / params.avg_daily_volume).sqrt() * 10_000.0
75 } else {
76 0.0
77 };
78 let market_impact_usd = (impact_bps / 10_000.0) * notional_usd;
79
80 let total_cost_usd = commission_usd + spread_cost_usd + market_impact_usd;
81 let cost_bps = if notional_usd.abs() > 0.0 {
82 total_cost_usd / notional_usd * 10_000.0
83 } else {
84 0.0
85 };
86
87 ExecutionCost {
88 commission_usd,
89 spread_cost_usd,
90 market_impact_usd,
91 total_cost_usd,
92 cost_bps,
93 }
94 }
95}
96
97#[derive(Debug, Clone, PartialEq, Eq)]
101pub enum TradeDirection {
102 Buy,
104 Sell,
106}
107
108#[derive(Debug, Clone)]
110pub struct Trade {
111 pub symbol: String,
113 pub direction: TradeDirection,
115 pub weight_change: f64,
117 pub estimated_cost_bps: f64,
119}
120
121pub struct TurnoverOptimizer;
135
136impl TurnoverOptimizer {
137 pub fn optimize(
147 current: &HashMap<String, f64>,
148 target: &HashMap<String, f64>,
149 cost_params: &CostParams,
150 tolerance: f64,
151 ) -> Vec<Trade> {
152 let mut symbols: Vec<String> = current
154 .keys()
155 .chain(target.keys())
156 .cloned()
157 .collect::<std::collections::HashSet<_>>()
158 .into_iter()
159 .collect();
160 symbols.sort();
161
162 let mut trades: Vec<Trade> = Vec::new();
163
164 for symbol in &symbols {
165 let cur = current.get(symbol).copied().unwrap_or(0.0);
166 let tgt = target.get(symbol).copied().unwrap_or(0.0);
167 let delta = tgt - cur;
168
169 if delta.abs() < tolerance {
170 continue;
171 }
172
173 let notional = delta.abs();
176 let shares = delta.abs();
177 let cost = CostModel::estimate(notional, shares, 1.0, cost_params);
178
179 trades.push(Trade {
180 symbol: symbol.clone(),
181 direction: if delta > 0.0 { TradeDirection::Buy } else { TradeDirection::Sell },
182 weight_change: delta.abs(),
183 estimated_cost_bps: cost.cost_bps,
184 });
185 }
186
187 trades.sort_by(|a, b| b.weight_change.partial_cmp(&a.weight_change).unwrap_or(std::cmp::Ordering::Equal));
189 trades
190 }
191}
192
193#[cfg(test)]
196mod tests {
197 use super::*;
198
199 fn default_params() -> CostParams {
200 CostParams {
201 commission_per_share: 0.005,
202 spread_bps: 5.0,
203 impact_coefficient: 0.1,
204 avg_daily_volume: 1_000_000.0,
205 }
206 }
207
208 #[test]
211 fn commission_computed_correctly() {
212 let params = default_params();
213 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
214 assert!((cost.commission_usd - 5.0).abs() < 1e-9, "commission={}", cost.commission_usd);
216 }
217
218 #[test]
219 fn spread_cost_computed_correctly() {
220 let params = default_params();
221 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
222 assert!((cost.spread_cost_usd - 5.0).abs() < 1e-9, "spread={}", cost.spread_cost_usd);
224 }
225
226 #[test]
227 fn market_impact_formula() {
228 let params = default_params();
230 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
231 let expected_impact_bps = 0.1 * (1_000.0f64 / 1_000_000.0).sqrt() * 10_000.0;
232 let expected_impact_usd = expected_impact_bps / 10_000.0 * 10_000.0;
233 assert!((cost.market_impact_usd - expected_impact_usd).abs() < 1e-6,
234 "impact={} expected={}", cost.market_impact_usd, expected_impact_usd);
235 }
236
237 #[test]
238 fn total_cost_is_sum_of_components() {
239 let params = default_params();
240 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
241 let expected = cost.commission_usd + cost.spread_cost_usd + cost.market_impact_usd;
242 assert!((cost.total_cost_usd - expected).abs() < 1e-9);
243 }
244
245 #[test]
246 fn cost_bps_equals_total_over_notional() {
247 let params = default_params();
248 let notional = 50_000.0;
249 let cost = CostModel::estimate(notional, 5_000.0, 10.0, ¶ms);
250 let expected_bps = cost.total_cost_usd / notional * 10_000.0;
251 assert!((cost.cost_bps - expected_bps).abs() < 1e-9);
252 }
253
254 #[test]
255 fn zero_shares_zero_cost() {
256 let params = default_params();
257 let cost = CostModel::estimate(0.0, 0.0, 10.0, ¶ms);
258 assert_eq!(cost.commission_usd, 0.0);
259 assert_eq!(cost.spread_cost_usd, 0.0);
260 assert_eq!(cost.market_impact_usd, 0.0);
261 assert_eq!(cost.total_cost_usd, 0.0);
262 assert_eq!(cost.cost_bps, 0.0);
263 }
264
265 #[test]
266 fn zero_adv_zero_impact() {
267 let mut params = default_params();
268 params.avg_daily_volume = 0.0;
269 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
270 assert_eq!(cost.market_impact_usd, 0.0);
271 }
272
273 #[test]
274 fn impact_increases_with_shares() {
275 let params = default_params();
276 let cost_small = CostModel::estimate(1_000.0, 100.0, 10.0, ¶ms);
277 let cost_large = CostModel::estimate(100_000.0, 10_000.0, 10.0, ¶ms);
278 assert!(cost_large.market_impact_usd > cost_small.market_impact_usd);
279 }
280
281 #[test]
282 fn higher_spread_higher_cost() {
283 let mut params_lo = default_params();
284 let mut params_hi = default_params();
285 params_lo.spread_bps = 1.0;
286 params_hi.spread_bps = 20.0;
287 let lo = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms_lo);
288 let hi = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms_hi);
289 assert!(hi.spread_cost_usd > lo.spread_cost_usd);
290 }
291
292 #[test]
295 fn no_trades_when_within_tolerance() {
296 let current: HashMap<String, f64> = [("AAPL".to_string(), 0.3), ("MSFT".to_string(), 0.7)].into();
297 let target: HashMap<String, f64> = [("AAPL".to_string(), 0.301), ("MSFT".to_string(), 0.699)].into();
298 let params = default_params();
299 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
300 assert!(trades.is_empty(), "expected no trades, got {}", trades.len());
301 }
302
303 #[test]
304 fn trades_generated_for_large_deltas() {
305 let current: HashMap<String, f64> = [("AAPL".to_string(), 0.2)].into();
306 let target: HashMap<String, f64> = [("AAPL".to_string(), 0.5)].into();
307 let params = default_params();
308 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
309 assert_eq!(trades.len(), 1);
310 assert_eq!(trades[0].symbol, "AAPL");
311 assert_eq!(trades[0].direction, TradeDirection::Buy);
312 assert!((trades[0].weight_change - 0.3).abs() < 1e-9);
313 }
314
315 #[test]
316 fn sell_direction_for_reduce() {
317 let current: HashMap<String, f64> = [("SPY".to_string(), 0.6)].into();
318 let target: HashMap<String, f64> = [("SPY".to_string(), 0.3)].into();
319 let params = default_params();
320 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
321 assert_eq!(trades.len(), 1);
322 assert_eq!(trades[0].direction, TradeDirection::Sell);
323 assert!((trades[0].weight_change - 0.3).abs() < 1e-9);
324 }
325
326 #[test]
327 fn new_position_is_buy() {
328 let current: HashMap<String, f64> = HashMap::new();
329 let target: HashMap<String, f64> = [("GLD".to_string(), 0.1)].into();
330 let params = default_params();
331 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
332 assert_eq!(trades.len(), 1);
333 assert_eq!(trades[0].direction, TradeDirection::Buy);
334 }
335
336 #[test]
337 fn liquidate_position_is_sell() {
338 let current: HashMap<String, f64> = [("TLT".to_string(), 0.25)].into();
339 let target: HashMap<String, f64> = HashMap::new();
340 let params = default_params();
341 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
342 assert_eq!(trades.len(), 1);
343 assert_eq!(trades[0].direction, TradeDirection::Sell);
344 }
345
346 #[test]
347 fn trades_sorted_by_weight_change_desc() {
348 let current: HashMap<String, f64> = [
349 ("A".to_string(), 0.1),
350 ("B".to_string(), 0.5),
351 ("C".to_string(), 0.2),
352 ].into();
353 let target: HashMap<String, f64> = [
354 ("A".to_string(), 0.5), ("B".to_string(), 0.1), ("C".to_string(), 0.4), ].into();
358 let params = default_params();
359 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
360 assert_eq!(trades.len(), 3);
361 for i in 0..trades.len() - 1 {
362 assert!(trades[i].weight_change >= trades[i + 1].weight_change);
363 }
364 }
365
366 #[test]
367 fn estimated_cost_bps_non_negative() {
368 let current: HashMap<String, f64> = [("X".to_string(), 0.0)].into();
369 let target: HashMap<String, f64> = [("X".to_string(), 0.1)].into();
370 let params = default_params();
371 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
372 for t in &trades {
373 assert!(t.estimated_cost_bps >= 0.0, "cost_bps={}", t.estimated_cost_bps);
374 }
375 }
376
377 #[test]
378 fn multiple_symbols_multi_trade() {
379 let current: HashMap<String, f64> = [
380 ("AAPL".to_string(), 0.25),
381 ("MSFT".to_string(), 0.25),
382 ("GOOG".to_string(), 0.25),
383 ("AMZN".to_string(), 0.25),
384 ].into();
385 let target: HashMap<String, f64> = [
386 ("AAPL".to_string(), 0.4),
387 ("MSFT".to_string(), 0.1),
388 ("GOOG".to_string(), 0.35),
389 ("AMZN".to_string(), 0.15),
390 ].into();
391 let params = default_params();
392 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
393 assert_eq!(trades.len(), 4);
395 }
396
397 #[test]
398 fn exact_tolerance_boundary_included() {
399 let current: HashMap<String, f64> = [("X".to_string(), 0.0)].into();
402 let target: HashMap<String, f64> = [("X".to_string(), 0.005)].into();
403 let params = default_params();
404 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
405 assert_eq!(trades.len(), 1, "delta exactly = tolerance should trade");
406 let target: HashMap<String, f64> = [("X".to_string(), 0.0049)].into();
408 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
409 assert!(trades.is_empty(), "delta below tolerance should be skipped");
410 }
411
412 #[test]
413 fn cost_model_all_fields_populated() {
414 let params = default_params();
415 let cost = CostModel::estimate(10_000.0, 500.0, 20.0, ¶ms);
416 assert!(cost.commission_usd > 0.0);
417 assert!(cost.spread_cost_usd > 0.0);
418 assert!(cost.market_impact_usd > 0.0);
419 assert!(cost.total_cost_usd > 0.0);
420 assert!(cost.cost_bps > 0.0);
421 }
422
423 #[test]
424 fn impact_coefficient_zero_no_impact() {
425 let mut params = default_params();
426 params.impact_coefficient = 0.0;
427 let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, ¶ms);
428 assert_eq!(cost.market_impact_usd, 0.0);
429 }
430
431 #[test]
432 fn trade_weight_change_is_absolute() {
433 let current: HashMap<String, f64> = [("X".to_string(), 0.5)].into();
434 let target: HashMap<String, f64> = [("X".to_string(), 0.2)].into();
435 let params = default_params();
436 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
437 assert_eq!(trades.len(), 1);
438 assert!(trades[0].weight_change > 0.0, "weight_change should be positive");
439 }
440
441 #[test]
442 fn empty_portfolios_no_trades() {
443 let current: HashMap<String, f64> = HashMap::new();
444 let target: HashMap<String, f64> = HashMap::new();
445 let params = default_params();
446 let trades = TurnoverOptimizer::optimize(¤t, &target, ¶ms, 0.005);
447 assert!(trades.is_empty());
448 }
449}