use crate::error::FinError;
use crate::types::{NanoTimestamp, Price, Quantity, Side, Symbol};
use rust_decimal::Decimal;
#[derive(Debug, Clone)]
pub struct OpenTrade {
pub symbol: Symbol,
pub side: Side,
pub quantity: Quantity,
pub entry_price: Price,
pub entry_fair_value: Price,
pub entry_fee: Decimal,
pub opened_at: NanoTimestamp,
}
impl OpenTrade {
pub fn new(
symbol: Symbol,
side: Side,
quantity: Quantity,
entry_price: Price,
entry_fair_value: Price,
entry_fee: Decimal,
opened_at: NanoTimestamp,
) -> Self {
Self {
symbol,
side,
quantity,
entry_price,
entry_fair_value,
entry_fee,
opened_at,
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PnlEvent {
pub symbol: Symbol,
pub side: Side,
pub quantity: Decimal,
pub entry_price: Decimal,
pub exit_price: Decimal,
pub opened_at: NanoTimestamp,
pub closed_at: NanoTimestamp,
pub direction_alpha: Decimal,
pub timing_alpha: Decimal,
pub slippage_cost: Decimal,
pub fee_cost: Decimal,
pub realized_pnl: Decimal,
}
#[derive(Debug, Default)]
pub struct PnlAttributor {
open_trades: std::collections::HashMap<String, OpenTrade>,
pub total_realized_pnl: Decimal,
pub closed_trade_count: usize,
}
impl PnlAttributor {
pub fn new() -> Self {
Self::default()
}
pub fn open_trade(&mut self, trade_id: impl Into<String>, trade: OpenTrade) {
self.open_trades.insert(trade_id.into(), trade);
}
pub fn close_trade(
&mut self,
trade_id: &str,
exit_price: Price,
exit_fair_value: Price,
exit_fee: Decimal,
closed_at: NanoTimestamp,
) -> Result<PnlEvent, FinError> {
let trade = self
.open_trades
.remove(trade_id)
.ok_or_else(|| FinError::InvalidInput(format!("trade '{trade_id}' not found")))?;
let qty = trade.quantity.value();
let entry_p = trade.entry_price.value();
let exit_p = exit_price.value();
let entry_fair = trade.entry_fair_value.value();
let exit_fair = exit_fair_value.value();
let raw_pnl = match trade.side {
Side::Bid => (exit_p - entry_p) * qty,
Side::Ask => (entry_p - exit_p) * qty,
};
let direction_alpha = match trade.side {
Side::Bid => (exit_fair - entry_fair) * qty,
Side::Ask => (entry_fair - exit_fair) * qty,
};
let entry_slip = match trade.side {
Side::Bid => (entry_p - entry_fair) * qty,
Side::Ask => (entry_fair - entry_p) * qty,
};
let exit_slip = match trade.side {
Side::Bid => (exit_fair - exit_p) * qty,
Side::Ask => (exit_p - exit_fair) * qty,
};
let slippage_cost = (entry_slip + exit_slip).max(Decimal::ZERO);
let fee_cost = (trade.entry_fee + exit_fee).max(Decimal::ZERO);
let realized_pnl = raw_pnl - fee_cost;
let timing_alpha = Decimal::ZERO;
self.total_realized_pnl += realized_pnl;
self.closed_trade_count += 1;
Ok(PnlEvent {
symbol: trade.symbol,
side: trade.side,
quantity: qty,
entry_price: entry_p,
exit_price: exit_p,
opened_at: trade.opened_at,
closed_at,
direction_alpha,
timing_alpha,
slippage_cost,
fee_cost,
realized_pnl,
})
}
pub fn unrealized_pnl(&self, trade_id: &str, current_price: Decimal) -> Option<Decimal> {
let trade = self.open_trades.get(trade_id)?;
let qty = trade.quantity.value();
let entry_p = trade.entry_price.value();
let upnl = match trade.side {
Side::Bid => (current_price - entry_p) * qty,
Side::Ask => (entry_p - current_price) * qty,
};
Some(upnl - trade.entry_fee)
}
pub fn open_trade_count(&self) -> usize {
self.open_trades.len()
}
pub fn has_open_trade(&self, trade_id: &str) -> bool {
self.open_trades.contains_key(trade_id)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal_macros::dec;
fn sym() -> Symbol {
Symbol::new("AAPL").unwrap()
}
fn ts(n: i64) -> NanoTimestamp {
NanoTimestamp::new(n)
}
#[test]
fn test_long_profitable_trade() {
let mut attr = PnlAttributor::new();
let trade = OpenTrade::new(
sym(),
Side::Bid,
Quantity::new(dec!(10)).unwrap(),
Price::new(dec!(100)).unwrap(),
Price::new(dec!(100)).unwrap(),
dec!(0.50),
ts(1000),
);
attr.open_trade("t1", trade);
assert!(attr.has_open_trade("t1"));
let event = attr
.close_trade(
"t1",
Price::new(dec!(110)).unwrap(),
Price::new(dec!(110)).unwrap(),
dec!(0.50),
ts(2000),
)
.unwrap();
assert_eq!(event.realized_pnl, dec!(99));
assert_eq!(event.fee_cost, dec!(1.00));
assert_eq!(event.direction_alpha, dec!(100));
assert!(!attr.has_open_trade("t1"));
assert_eq!(attr.closed_trade_count, 1);
}
#[test]
fn test_short_profitable_trade() {
let mut attr = PnlAttributor::new();
let trade = OpenTrade::new(
sym(),
Side::Ask,
Quantity::new(dec!(5)).unwrap(),
Price::new(dec!(200)).unwrap(),
Price::new(dec!(200)).unwrap(),
dec!(0.25),
ts(1000),
);
attr.open_trade("t2", trade);
let event = attr
.close_trade(
"t2",
Price::new(dec!(190)).unwrap(),
Price::new(dec!(190)).unwrap(),
dec!(0.25),
ts(3000),
)
.unwrap();
assert_eq!(event.realized_pnl, dec!(49.50));
assert_eq!(event.direction_alpha, dec!(50));
}
#[test]
fn test_slippage_computed() {
let mut attr = PnlAttributor::new();
let trade = OpenTrade::new(
sym(),
Side::Bid,
Quantity::new(dec!(1)).unwrap(),
Price::new(dec!(100.10)).unwrap(), Price::new(dec!(100)).unwrap(),
Decimal::ZERO,
ts(1000),
);
attr.open_trade("t3", trade);
let event = attr
.close_trade(
"t3",
Price::new(dec!(105)).unwrap(),
Price::new(dec!(105.05)).unwrap(), Decimal::ZERO,
ts(2000),
)
.unwrap();
assert_eq!(event.slippage_cost, dec!(0.15));
}
#[test]
fn test_unrealized_pnl() {
let mut attr = PnlAttributor::new();
let trade = OpenTrade::new(
sym(),
Side::Bid,
Quantity::new(dec!(10)).unwrap(),
Price::new(dec!(50)).unwrap(),
Price::new(dec!(50)).unwrap(),
dec!(1.00),
ts(1000),
);
attr.open_trade("t4", trade);
let upnl = attr.unrealized_pnl("t4", dec!(55)).unwrap();
assert_eq!(upnl, dec!(49));
}
#[test]
fn test_close_unknown_trade_errors() {
let mut attr = PnlAttributor::new();
let err = attr
.close_trade(
"nonexistent",
Price::new(dec!(100)).unwrap(),
Price::new(dec!(100)).unwrap(),
Decimal::ZERO,
ts(1000),
)
.unwrap_err();
assert!(matches!(err, FinError::InvalidInput(_)));
}
}