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