Skip to main content

fin_primitives/execution/
mod.rs

1//! Execution cost estimation and turnover optimization.
2//!
3//! ## Overview
4//!
5//! This module provides:
6//! - [`ExecutionCost`]: breakdown of round-trip execution costs (commission, spread, market impact).
7//! - [`CostParams`]: parameters for the cost model (commission rate, spread, impact coefficient).
8//! - [`CostModel`]: estimates execution cost for a given order.
9//! - [`TurnoverOptimizer`]: finds trades minimizing cost while tracking a target portfolio.
10//! - [`Trade`]: a single rebalancing trade with estimated cost.
11
12use std::collections::HashMap;
13
14// ─── ExecutionCost ────────────────────────────────────────────────────────────
15
16/// Full breakdown of estimated round-trip execution cost for one order.
17#[derive(Debug, Clone)]
18pub struct ExecutionCost {
19    /// Commission paid to the broker, in USD.
20    pub commission_usd: f64,
21    /// Half-spread cost (buy at ask, sell at bid), in USD.
22    pub spread_cost_usd: f64,
23    /// Estimated market impact cost (price concession due to order size), in USD.
24    pub market_impact_usd: f64,
25    /// Sum of all cost components, in USD.
26    pub total_cost_usd: f64,
27    /// Total cost expressed as basis points of notional: `total_cost_usd / notional * 10_000`.
28    pub cost_bps: f64,
29}
30
31// ─── CostParams ───────────────────────────────────────────────────────────────
32
33/// Parameters controlling the execution cost model.
34#[derive(Debug, Clone)]
35pub struct CostParams {
36    /// Commission per share (USD/share).
37    pub commission_per_share: f64,
38    /// One-way bid-ask spread in basis points.
39    pub spread_bps: f64,
40    /// Almgren-Chriss impact coefficient `η` (dimensionless).
41    pub impact_coefficient: f64,
42    /// Average daily volume for the instrument (shares/day).
43    pub avg_daily_volume: f64,
44}
45
46// ─── CostModel ────────────────────────────────────────────────────────────────
47
48/// Estimates round-trip execution cost for a given order.
49///
50/// ## Formula
51///
52/// ```text
53/// commission_usd  = commission_per_share * shares
54/// spread_cost_usd = (spread_bps / 10_000) * notional_usd
55/// impact_bps      = impact_coefficient * sqrt(shares / avg_daily_volume) * 10_000
56/// market_impact_usd = (impact_bps / 10_000) * notional_usd
57/// total_cost_usd  = commission_usd + spread_cost_usd + market_impact_usd
58/// cost_bps        = total_cost_usd / notional_usd * 10_000
59/// ```
60pub struct CostModel;
61
62impl CostModel {
63    /// Estimates round-trip execution cost for an order of `shares` at `price`.
64    ///
65    /// `notional_usd` = `shares * price` (passed explicitly to avoid rounding differences).
66    pub fn estimate(notional_usd: f64, shares: f64, price: f64, params: &CostParams) -> ExecutionCost {
67        let _ = price; // price is implicit in notional/shares
68
69        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// ─── Trade ────────────────────────────────────────────────────────────────────
98
99/// Direction of a rebalancing trade.
100#[derive(Debug, Clone, PartialEq, Eq)]
101pub enum TradeDirection {
102    /// Buy (increase position / add weight).
103    Buy,
104    /// Sell (decrease position / reduce weight).
105    Sell,
106}
107
108/// A single rebalancing trade recommended by [`TurnoverOptimizer`].
109#[derive(Debug, Clone)]
110pub struct Trade {
111    /// Instrument symbol.
112    pub symbol: String,
113    /// Buy or sell.
114    pub direction: TradeDirection,
115    /// Absolute change in portfolio weight.
116    pub weight_change: f64,
117    /// Estimated one-way cost in basis points.
118    pub estimated_cost_bps: f64,
119}
120
121// ─── TurnoverOptimizer ────────────────────────────────────────────────────────
122
123/// Finds the minimal set of trades that moves a portfolio from `current` weights
124/// to `target` weights while keeping total execution cost low.
125///
126/// ## Algorithm
127///
128/// 1. For each symbol in the union of current and target weights, compute the
129///    weight delta `Δw = target - current`.
130/// 2. If `|Δw| < tolerance`, skip (already within band).
131/// 3. Otherwise add a [`Trade`] with the estimated cost for a unit-notional order
132///    sized proportionally to `|Δw|`.
133/// 4. Trades are sorted by `|Δw|` descending (largest rebalance first).
134pub struct TurnoverOptimizer;
135
136impl TurnoverOptimizer {
137    /// Generate the minimal set of trades to rebalance from `current` to `target`.
138    ///
139    /// `tolerance` is the minimum absolute weight change that warrants a trade
140    /// (e.g. 0.005 = 50 bps). Smaller changes are ignored to avoid excessive
141    /// round-trip cost.
142    ///
143    /// `cost_params` are used to estimate the cost of each trade assuming a
144    /// unit notional of $1, with `shares = |Δw| / price` where price is set
145    /// to 1.0 for weight-space estimation.
146    pub fn optimize(
147        current: &HashMap<String, f64>,
148        target: &HashMap<String, f64>,
149        cost_params: &CostParams,
150        tolerance: f64,
151    ) -> Vec<Trade> {
152        // Collect all symbols from both maps.
153        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            // Estimate cost: treat weight change as fraction of notional = 1 USD.
174            // shares = |delta| * 1 USD / $1 per share = |delta|
175            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        // Sort by weight change descending (largest rebalance first).
188        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// ─── tests ────────────────────────────────────────────────────────────────────
194
195#[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    // ── CostModel ──────────────────────────────────────────────────────────
209
210    #[test]
211    fn commission_computed_correctly() {
212        let params = default_params();
213        let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, &params);
214        // commission = 0.005 * 1000 = 5.0
215        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, &params);
222        // spread = 5 bps * 10000 = 5.0 USD
223        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        // impact_bps = 0.1 * sqrt(1000 / 1_000_000) * 10_000 = 0.1 * 0.03162 * 10_000 = 31.62
229        let params = default_params();
230        let cost = CostModel::estimate(10_000.0, 1_000.0, 10.0, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
277        let cost_large = CostModel::estimate(100_000.0, 10_000.0, 10.0, &params);
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, &params_lo);
288        let hi = CostModel::estimate(10_000.0, 1_000.0, 10.0, &params_hi);
289        assert!(hi.spread_cost_usd > lo.spread_cost_usd);
290    }
291
292    // ── TurnoverOptimizer ──────────────────────────────────────────────────
293
294    #[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(&current, &target, &params, 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(&current, &target, &params, 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(&current, &target, &params, 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(&current, &target, &params, 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(&current, &target, &params, 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),  // +0.4
355            ("B".to_string(), 0.1),  // -0.4
356            ("C".to_string(), 0.4),  // +0.2
357        ].into();
358        let params = default_params();
359        let trades = TurnoverOptimizer::optimize(&current, &target, &params, 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(&current, &target, &params, 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(&current, &target, &params, 0.005);
393        // All four have |delta| > 0.005
394        assert_eq!(trades.len(), 4);
395    }
396
397    #[test]
398    fn exact_tolerance_boundary_included() {
399        // `tolerance` is documented as the minimum change that warrants a trade and
400        // only `|delta| < tolerance` is skipped, so delta == tolerance trades.
401        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(&current, &target, &params, 0.005);
405        assert_eq!(trades.len(), 1, "delta exactly = tolerance should trade");
406        // Just inside the band is skipped.
407        let target: HashMap<String, f64> = [("X".to_string(), 0.0049)].into();
408        let trades = TurnoverOptimizer::optimize(&current, &target, &params, 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, &params);
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, &params);
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(&current, &target, &params, 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(&current, &target, &params, 0.005);
447        assert!(trades.is_empty());
448    }
449}