Skip to main content

fin_primitives/pnl/
mod.rs

1//! Streaming P&L attribution: decomposes realized P&L into alpha and cost components.
2//!
3//! ## Responsibility
4//! Streaming P&L attribution: tracks realized and unrealized P&L per trade and
5//! decomposes it into direction alpha, timing alpha, slippage cost, and fee cost.
6//!
7//! ## Guarantees
8//! - All arithmetic uses `rust_decimal::Decimal`; no floating-point drift
9//! - `PnlAttributor::close_trade` emits a `PnlEvent` with fully decomposed components
10//! - Slippage cost is always non-negative (it is a cost, not a gain)
11//! - Fee cost is always non-negative
12//!
13//! ## NOT Responsible For
14//! - Order routing or execution
15//! - Position risk checks (see `risk` module)
16
17use crate::error::FinError;
18use crate::types::{NanoTimestamp, Price, Quantity, Side, Symbol};
19use rust_decimal::Decimal;
20
21/// An open trade leg awaiting closure.
22#[derive(Debug, Clone)]
23pub struct OpenTrade {
24    /// Instrument traded.
25    pub symbol: Symbol,
26    /// Direction of the opening fill.
27    pub side: Side,
28    /// Size of the opening fill.
29    pub quantity: Quantity,
30    /// Execution price of the opening fill.
31    pub entry_price: Price,
32    /// Theoretical fair value at entry time (e.g. mid-price).
33    ///
34    /// Used to compute slippage: `(entry_price - fair_value).abs() * qty`.
35    pub entry_fair_value: Price,
36    /// Commission charged on entry.
37    pub entry_fee: Decimal,
38    /// When the trade was opened.
39    pub opened_at: NanoTimestamp,
40}
41
42impl OpenTrade {
43    /// Creates a new `OpenTrade`.
44    pub fn new(
45        symbol: Symbol,
46        side: Side,
47        quantity: Quantity,
48        entry_price: Price,
49        entry_fair_value: Price,
50        entry_fee: Decimal,
51        opened_at: NanoTimestamp,
52    ) -> Self {
53        Self {
54            symbol,
55            side,
56            quantity,
57            entry_price,
58            entry_fair_value,
59            entry_fee,
60            opened_at,
61        }
62    }
63}
64
65/// Decomposed P&L event emitted when a trade is closed.
66///
67/// All components are signed from the perspective of the trader:
68/// positive = profit, negative = loss.
69///
70/// ## Decomposition identity
71/// `total_pnl ≈ direction_alpha + timing_alpha - slippage_cost - fee_cost`
72#[derive(Debug, Clone)]
73#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
74pub struct PnlEvent {
75    /// Instrument.
76    pub symbol: Symbol,
77    /// Side of the opening leg.
78    pub side: Side,
79    /// Size closed.
80    pub quantity: Decimal,
81    /// Price at which the trade was entered.
82    pub entry_price: Decimal,
83    /// Price at which the trade was exited.
84    pub exit_price: Decimal,
85    /// When the trade opened.
86    pub opened_at: NanoTimestamp,
87    /// When the trade closed.
88    pub closed_at: NanoTimestamp,
89
90    /// **Direction alpha**: profit from the raw directional move,
91    /// measured at fair values (no slippage, no fees).
92    ///
93    /// `long: (exit_fair - entry_fair) * qty`
94    /// `short: (entry_fair - exit_fair) * qty`
95    pub direction_alpha: Decimal,
96
97    /// **Timing alpha**: additional P&L gained from entering/exiting
98    /// at a favourable time relative to a passive mid strategy.
99    ///
100    /// Currently zero-valued (reserved for future bar-level attribution).
101    pub timing_alpha: Decimal,
102
103    /// **Slippage cost**: execution cost from crossing the spread
104    /// (difference between execution price and fair value).
105    ///
106    /// Always `>= 0`.
107    pub slippage_cost: Decimal,
108
109    /// **Fee cost**: total commissions paid on both legs.
110    ///
111    /// Always `>= 0`.
112    pub fee_cost: Decimal,
113
114    /// Net realized P&L: raw directional P&L minus costs.
115    ///
116    /// `= (exit_price - entry_price) * qty` (for long)
117    /// minus total fees.
118    pub realized_pnl: Decimal,
119}
120
121/// Streaming P&L attributor.
122///
123/// Open a trade with [`PnlAttributor::open_trade`], close it with
124/// [`PnlAttributor::close_trade`] and receive a [`PnlEvent`] with
125/// fully decomposed attribution.
126///
127/// # Example
128/// ```rust
129/// use fin_primitives::pnl::{PnlAttributor, OpenTrade};
130/// use fin_primitives::types::{Symbol, Side, Price, Quantity, NanoTimestamp};
131/// use rust_decimal_macros::dec;
132///
133/// let mut attr = PnlAttributor::new();
134/// let sym = Symbol::new("AAPL").unwrap();
135/// let ts = NanoTimestamp::new(1_000_000_000);
136/// let trade = OpenTrade::new(
137///     sym.clone(), Side::Bid,
138///     Quantity::new(dec!(10)).unwrap(),
139///     Price::new(dec!(100)).unwrap(),
140///     Price::new(dec!(100.05)).unwrap(),  // fair value slightly above
141///     dec!(0.10),
142///     ts,
143/// );
144/// attr.open_trade("t1", trade);
145/// let event = attr.close_trade(
146///     "t1",
147///     Price::new(dec!(105)).unwrap(),
148///     Price::new(dec!(104.95)).unwrap(),
149///     dec!(0.10),
150///     NanoTimestamp::new(2_000_000_000),
151/// ).unwrap();
152/// assert!(event.realized_pnl > dec!(0));
153/// ```
154#[derive(Debug, Default)]
155pub struct PnlAttributor {
156    open_trades: std::collections::HashMap<String, OpenTrade>,
157    /// Cumulative realized P&L across all closed trades.
158    pub total_realized_pnl: Decimal,
159    /// Count of closed trades.
160    pub closed_trade_count: usize,
161}
162
163impl PnlAttributor {
164    /// Creates an empty `PnlAttributor`.
165    pub fn new() -> Self {
166        Self::default()
167    }
168
169    /// Registers an open trade under `trade_id`.
170    ///
171    /// If a trade with the same id already exists it is silently replaced.
172    pub fn open_trade(&mut self, trade_id: impl Into<String>, trade: OpenTrade) {
173        self.open_trades.insert(trade_id.into(), trade);
174    }
175
176    /// Closes the trade identified by `trade_id` and returns a [`PnlEvent`].
177    ///
178    /// # Errors
179    /// - [`FinError::InvalidInput`] if `trade_id` is not found.
180    /// - [`FinError::ArithmeticOverflow`] on checked-arithmetic failure (extremely unlikely
181    ///   with normal price magnitudes).
182    pub fn close_trade(
183        &mut self,
184        trade_id: &str,
185        exit_price: Price,
186        exit_fair_value: Price,
187        exit_fee: Decimal,
188        closed_at: NanoTimestamp,
189    ) -> Result<PnlEvent, FinError> {
190        let trade = self
191            .open_trades
192            .remove(trade_id)
193            .ok_or_else(|| FinError::InvalidInput(format!("trade '{trade_id}' not found")))?;
194
195        let qty = trade.quantity.value();
196        let entry_p = trade.entry_price.value();
197        let exit_p = exit_price.value();
198        let entry_fair = trade.entry_fair_value.value();
199        let exit_fair = exit_fair_value.value();
200
201        // Raw realized P&L (signed, before fees)
202        let raw_pnl = match trade.side {
203            Side::Bid => (exit_p - entry_p) * qty,
204            Side::Ask => (entry_p - exit_p) * qty,
205        };
206
207        // Direction alpha: fair-value move in the direction of the trade
208        let direction_alpha = match trade.side {
209            Side::Bid => (exit_fair - entry_fair) * qty,
210            Side::Ask => (entry_fair - exit_fair) * qty,
211        };
212
213        // Slippage: cost of crossing the spread on entry and exit
214        // entry slippage: long paid above fair, short sold below fair
215        let entry_slip = match trade.side {
216            Side::Bid => (entry_p - entry_fair) * qty,
217            Side::Ask => (entry_fair - entry_p) * qty,
218        };
219        let exit_slip = match trade.side {
220            Side::Bid => (exit_fair - exit_p) * qty,
221            Side::Ask => (exit_p - exit_fair) * qty,
222        };
223        // Slippage cost is the total execution disadvantage vs fair value
224        let slippage_cost = (entry_slip + exit_slip).max(Decimal::ZERO);
225
226        let fee_cost = (trade.entry_fee + exit_fee).max(Decimal::ZERO);
227        let realized_pnl = raw_pnl - fee_cost;
228
229        // Timing alpha is currently zero; reserved for bar-level decomposition
230        let timing_alpha = Decimal::ZERO;
231
232        self.total_realized_pnl += realized_pnl;
233        self.closed_trade_count += 1;
234
235        Ok(PnlEvent {
236            symbol: trade.symbol,
237            side: trade.side,
238            quantity: qty,
239            entry_price: entry_p,
240            exit_price: exit_p,
241            opened_at: trade.opened_at,
242            closed_at,
243            direction_alpha,
244            timing_alpha,
245            slippage_cost,
246            fee_cost,
247            realized_pnl,
248        })
249    }
250
251    /// Returns unrealized P&L for a trade given the current fair value.
252    ///
253    /// Returns `None` if `trade_id` is not found.
254    pub fn unrealized_pnl(&self, trade_id: &str, current_price: Decimal) -> Option<Decimal> {
255        let trade = self.open_trades.get(trade_id)?;
256        let qty = trade.quantity.value();
257        let entry_p = trade.entry_price.value();
258        let upnl = match trade.side {
259            Side::Bid => (current_price - entry_p) * qty,
260            Side::Ask => (entry_p - current_price) * qty,
261        };
262        Some(upnl - trade.entry_fee)
263    }
264
265    /// Returns the number of currently open trades.
266    pub fn open_trade_count(&self) -> usize {
267        self.open_trades.len()
268    }
269
270    /// Returns `true` if `trade_id` is an open trade.
271    pub fn has_open_trade(&self, trade_id: &str) -> bool {
272        self.open_trades.contains_key(trade_id)
273    }
274}
275
276#[cfg(test)]
277mod tests {
278    use super::*;
279    use rust_decimal_macros::dec;
280
281    fn sym() -> Symbol {
282        Symbol::new("AAPL").unwrap()
283    }
284    fn ts(n: i64) -> NanoTimestamp {
285        NanoTimestamp::new(n)
286    }
287
288    #[test]
289    fn test_long_profitable_trade() {
290        let mut attr = PnlAttributor::new();
291        let trade = OpenTrade::new(
292            sym(),
293            Side::Bid,
294            Quantity::new(dec!(10)).unwrap(),
295            Price::new(dec!(100)).unwrap(),
296            Price::new(dec!(100)).unwrap(),
297            dec!(0.50),
298            ts(1000),
299        );
300        attr.open_trade("t1", trade);
301        assert!(attr.has_open_trade("t1"));
302        let event = attr
303            .close_trade(
304                "t1",
305                Price::new(dec!(110)).unwrap(),
306                Price::new(dec!(110)).unwrap(),
307                dec!(0.50),
308                ts(2000),
309            )
310            .unwrap();
311        // raw = (110-100)*10 = 100, fee = 1.00, net = 99
312        assert_eq!(event.realized_pnl, dec!(99));
313        assert_eq!(event.fee_cost, dec!(1.00));
314        assert_eq!(event.direction_alpha, dec!(100));
315        assert!(!attr.has_open_trade("t1"));
316        assert_eq!(attr.closed_trade_count, 1);
317    }
318
319    #[test]
320    fn test_short_profitable_trade() {
321        let mut attr = PnlAttributor::new();
322        let trade = OpenTrade::new(
323            sym(),
324            Side::Ask,
325            Quantity::new(dec!(5)).unwrap(),
326            Price::new(dec!(200)).unwrap(),
327            Price::new(dec!(200)).unwrap(),
328            dec!(0.25),
329            ts(1000),
330        );
331        attr.open_trade("t2", trade);
332        let event = attr
333            .close_trade(
334                "t2",
335                Price::new(dec!(190)).unwrap(),
336                Price::new(dec!(190)).unwrap(),
337                dec!(0.25),
338                ts(3000),
339            )
340            .unwrap();
341        // raw = (200-190)*5 = 50, fees=0.50, net=49.50
342        assert_eq!(event.realized_pnl, dec!(49.50));
343        assert_eq!(event.direction_alpha, dec!(50));
344    }
345
346    #[test]
347    fn test_slippage_computed() {
348        let mut attr = PnlAttributor::new();
349        let trade = OpenTrade::new(
350            sym(),
351            Side::Bid,
352            Quantity::new(dec!(1)).unwrap(),
353            Price::new(dec!(100.10)).unwrap(), // paid 0.10 above fair
354            Price::new(dec!(100)).unwrap(),
355            Decimal::ZERO,
356            ts(1000),
357        );
358        attr.open_trade("t3", trade);
359        let event = attr
360            .close_trade(
361                "t3",
362                Price::new(dec!(105)).unwrap(),
363                Price::new(dec!(105.05)).unwrap(), // exited 0.05 below fair
364                Decimal::ZERO,
365                ts(2000),
366            )
367            .unwrap();
368        // entry slip = (100.10-100)*1 = 0.10
369        // exit slip = (105.05-105)*1 = 0.05
370        assert_eq!(event.slippage_cost, dec!(0.15));
371    }
372
373    #[test]
374    fn test_unrealized_pnl() {
375        let mut attr = PnlAttributor::new();
376        let trade = OpenTrade::new(
377            sym(),
378            Side::Bid,
379            Quantity::new(dec!(10)).unwrap(),
380            Price::new(dec!(50)).unwrap(),
381            Price::new(dec!(50)).unwrap(),
382            dec!(1.00),
383            ts(1000),
384        );
385        attr.open_trade("t4", trade);
386        let upnl = attr.unrealized_pnl("t4", dec!(55)).unwrap();
387        // (55-50)*10 - 1.00 = 49
388        assert_eq!(upnl, dec!(49));
389    }
390
391    #[test]
392    fn test_close_unknown_trade_errors() {
393        let mut attr = PnlAttributor::new();
394        let err = attr
395            .close_trade(
396                "nonexistent",
397                Price::new(dec!(100)).unwrap(),
398                Price::new(dec!(100)).unwrap(),
399                Decimal::ZERO,
400                ts(1000),
401            )
402            .unwrap_err();
403        assert!(matches!(err, FinError::InvalidInput(_)));
404    }
405}