use ahash::AHashMap;
use indexmap::IndexMap;
use nautilus_core::{
DurationNanos, UnixNanos,
correctness::{
CorrectnessError, CorrectnessResult, FAILED, check_equal, check_predicate_false,
check_predicate_true,
},
};
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use crate::{
enums::{AccountType, LiquiditySide, OrderSide},
events::{AccountState, OrderFilled},
identifiers::{AccountId, InstrumentId},
instruments::{Instrument, InstrumentAny},
position::Position,
types::{AccountBalance, Currency, Money, Price, Quantity},
};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[cfg_attr(
feature = "python",
pyo3::pyclass(module = "nautilus_trader.model", from_py_object)
)]
pub struct BaseAccount {
pub id: AccountId,
pub account_type: AccountType,
pub base_currency: Option<Currency>,
pub calculate_account_state: bool,
pub events: Vec<AccountState>,
pub commissions: AHashMap<Currency, Money>,
pub balances: IndexMap<Currency, AccountBalance>,
pub balances_starting: IndexMap<Currency, Money>,
}
impl BaseAccount {
#[must_use]
pub fn new(event: AccountState, calculate_account_state: bool) -> Self {
let mut balances_starting: IndexMap<Currency, Money> = IndexMap::new();
let mut balances: IndexMap<Currency, AccountBalance> = IndexMap::new();
event.balances.iter().for_each(|balance| {
balances_starting.insert(balance.currency, balance.total);
balances.insert(balance.currency, *balance);
});
Self {
id: event.account_id,
account_type: event.account_type,
base_currency: event.base_currency,
calculate_account_state,
events: vec![event],
commissions: AHashMap::new(),
balances,
balances_starting,
}
}
#[must_use]
pub(crate) fn clone_without_events(&self) -> Self {
Self {
id: self.id,
account_type: self.account_type,
base_currency: self.base_currency,
calculate_account_state: self.calculate_account_state,
events: Vec::new(),
commissions: self.commissions.clone(),
balances: self.balances.clone(),
balances_starting: self.balances_starting.clone(),
}
}
#[must_use]
pub fn base_balance(&self, currency: Option<Currency>) -> Option<&AccountBalance> {
let currency = currency
.or(self.base_currency)
.expect("Currency must be specified");
self.balances.get(¤cy)
}
#[must_use]
pub fn base_balance_total(&self, currency: Option<Currency>) -> Option<Money> {
self.base_balance(currency).map(|balance| balance.total)
}
#[must_use]
pub fn base_balances_total(&self) -> IndexMap<Currency, Money> {
self.balances
.iter()
.map(|(currency, balance)| (*currency, balance.total))
.collect()
}
#[must_use]
pub fn base_balance_free(&self, currency: Option<Currency>) -> Option<Money> {
self.base_balance(currency).map(|balance| balance.free)
}
#[must_use]
pub fn base_balances_free(&self) -> IndexMap<Currency, Money> {
self.balances
.iter()
.map(|(currency, balance)| (*currency, balance.free))
.collect()
}
#[must_use]
pub fn base_balance_locked(&self, currency: Option<Currency>) -> Option<Money> {
self.base_balance(currency).map(|balance| balance.locked)
}
#[must_use]
pub fn base_balances_locked(&self) -> IndexMap<Currency, Money> {
self.balances
.iter()
.map(|(currency, balance)| (*currency, balance.locked))
.collect()
}
#[must_use]
pub fn base_last_event(&self) -> Option<AccountState> {
self.events.last().cloned()
}
pub fn update_balances(&mut self, balances: &[AccountBalance]) {
for balance in balances {
self.balances.insert(balance.currency, *balance);
}
}
pub fn update_commissions(&mut self, commission: Money) {
self.try_update_commissions(commission)
.expect("commission total exceeded Money bounds");
}
pub fn try_update_commissions(&mut self, commission: Money) -> anyhow::Result<()> {
let commission = commission.normalized();
if commission.is_zero() {
return Ok(());
}
let currency = commission.currency;
let total = self
.commissions
.get(¤cy)
.copied()
.map_or(Some(commission), |total| total.checked_add(commission))
.ok_or_else(|| anyhow::anyhow!("{currency} commission total exceeds Money bounds"))?;
self.commissions.insert(currency, total);
Ok(())
}
#[must_use]
pub fn commission(&self, currency: &Currency) -> Option<Money> {
self.commissions.get(currency).copied()
}
#[must_use]
pub fn commissions(&self) -> AHashMap<Currency, Money> {
self.commissions.clone()
}
pub(crate) fn check_event_account_id(&self, event: &AccountState) -> anyhow::Result<()> {
anyhow::ensure!(
event.account_id == self.id,
"Account event had a different account ID: expected {}, received {}",
self.id,
event.account_id
);
Ok(())
}
pub fn base_apply(&mut self, event: AccountState) {
check_equal(&event.account_id, &self.id, "event.account_id", "self.id").expect(FAILED);
self.update_balances(&event.balances);
self.events.push(event);
}
pub fn base_purge_account_events(&mut self, ts_now: UnixNanos, lookback_secs: u64) {
let Ok(lookback_ns) = DurationNanos::try_from_secs(lookback_secs) else {
log::warn!(
"Cannot purge account events: lookback_secs {lookback_secs} is not representable in `u64` nanoseconds"
);
return;
};
let purge_cutoff = ts_now.checked_sub(lookback_ns);
let mut retained_events = Vec::new();
for event in &self.events {
if purge_cutoff.is_none_or(|cutoff| event.ts_event > cutoff) {
retained_events.push(event.clone());
}
}
if retained_events.is_empty() && !self.events.is_empty() {
retained_events.push(self.events.last().expect("events not empty").clone());
}
self.events = retained_events;
}
pub fn base_calculate_balance_locked(
&self,
instrument: &InstrumentAny,
side: OrderSide,
quantity: Quantity,
price: Price,
use_quote_for_inverse: Option<bool>,
) -> anyhow::Result<Money> {
let base_currency = instrument
.base_currency()
.unwrap_or(instrument.quote_currency());
let quote_currency = instrument.quote_currency();
let amount = match side {
OrderSide::Buy => instrument
.try_calculate_notional_value(quantity, price, use_quote_for_inverse)?
.as_decimal()
.max(Decimal::ZERO),
OrderSide::Sell => quantity.as_decimal(),
};
if instrument.is_inverse() && !use_quote_for_inverse.unwrap_or(false) {
Ok(Money::from_decimal(amount, base_currency)?)
} else {
let currency = match side {
OrderSide::Buy => quote_currency,
OrderSide::Sell => base_currency,
};
Ok(Money::from_decimal(amount, currency)?)
}
}
pub fn base_calculate_pnls(
&self,
instrument: &InstrumentAny,
fill: &OrderFilled,
_position: Option<Position>,
) -> anyhow::Result<Vec<Money>> {
let mut pnls: IndexMap<Currency, Money> = IndexMap::new();
let base_currency = instrument.base_currency();
let fill_qty = fill.last_qty;
let notional = instrument.try_calculate_notional_value(fill_qty, fill.last_px, None)?;
if fill.order_side == OrderSide::Buy {
if let (Some(base_currency_value), None) = (base_currency, self.base_currency) {
pnls.insert(
base_currency_value,
Money::from_decimal(fill_qty.as_decimal(), base_currency_value)?,
);
}
pnls.insert(notional.currency, -notional);
} else {
if let (Some(base_currency_value), None) = (base_currency, self.base_currency) {
pnls.insert(
base_currency_value,
-Money::from_decimal(fill_qty.as_decimal(), base_currency_value)?,
);
}
pnls.insert(notional.currency, notional);
}
Ok(pnls.into_values().collect())
}
pub fn base_calculate_commission(
&self,
instrument: &InstrumentAny,
last_qty: Quantity,
last_px: Price,
liquidity_side: LiquiditySide,
use_quote_for_inverse: Option<bool>,
) -> anyhow::Result<Money> {
anyhow::ensure!(
liquidity_side != LiquiditySide::NoLiquiditySide,
"Invalid `LiquiditySide`: {liquidity_side}"
);
let notional =
instrument.try_calculate_notional_value(last_qty, last_px, use_quote_for_inverse)?;
let rate = match liquidity_side {
LiquiditySide::Maker => instrument.maker_fee(),
LiquiditySide::Taker => instrument.taker_fee(),
LiquiditySide::NoLiquiditySide => {
anyhow::bail!("Invalid `LiquiditySide`: {liquidity_side}")
}
};
let commission = notional
.as_decimal()
.checked_mul(rate)
.ok_or_else(|| anyhow::anyhow!("commission calculation overflow"))?;
Ok(Money::from_decimal(commission, notional.currency)?)
}
}
pub(crate) fn update_balance_locked(
balances: &mut IndexMap<Currency, AccountBalance>,
balances_locked: &mut AHashMap<(InstrumentId, Currency), Money>,
instrument_id: InstrumentId,
locked: Money,
) -> anyhow::Result<()> {
anyhow::ensure!(
!locked.is_negative(),
"locked balance was negative: {locked}"
);
let currency = locked.currency;
let key = (instrument_id, currency);
let Some(current_balance) = balances.get(¤cy).copied() else {
balances_locked.insert(key, locked);
return Ok(());
};
anyhow::ensure!(
current_balance.currency.precision == currency.precision,
"Cannot update {currency} reservation: precision {} differed from balance precision {}",
currency.precision,
current_balance.currency.precision
);
let previous = balances_locked.insert(key, locked);
match balance_from_locks(current_balance, balances_locked) {
Ok(balance) => {
balances.insert(currency, balance);
Ok(())
}
Err(e) => {
match previous {
Some(previous) => balances_locked.insert(key, previous),
None => balances_locked.remove(&key),
};
Err(e.into())
}
}
}
pub(crate) fn clear_balance_locked(
balances: &mut IndexMap<Currency, AccountBalance>,
balances_locked: &mut AHashMap<(InstrumentId, Currency), Money>,
instrument_id: InstrumentId,
) {
let currencies_to_recalc: Vec<Currency> = balances_locked
.keys()
.filter(|(id, _)| *id == instrument_id)
.map(|(_, currency)| *currency)
.collect();
for currency in ¤cies_to_recalc {
balances_locked.remove(&(instrument_id, *currency));
}
for currency in currencies_to_recalc {
recalculate_balance(balances, balances_locked, currency);
}
}
pub(crate) fn recalculate_balance(
balances: &mut IndexMap<Currency, AccountBalance>,
balances_locked: &AHashMap<(InstrumentId, Currency), Money>,
currency: Currency,
) {
let current_balance = if let Some(balance) = balances.get(¤cy) {
*balance
} else {
log::debug!("Cannot recalculate balance when no current balance for {currency}");
return;
};
let new_balance = match balance_from_locks(current_balance, balances_locked) {
Ok(balance) => balance,
Err(e) => {
log::error!(
"Cannot recalculate {currency} balance from reservations: {e}; using a non-spendable balance"
);
non_spendable_balance(current_balance)
}
};
balances.insert(currency, new_balance);
}
pub(crate) fn balance_from_locks(
current_balance: AccountBalance,
balances_locked: &AHashMap<(InstrumentId, Currency), Money>,
) -> CorrectnessResult<AccountBalance> {
let currency = current_balance.currency;
let mut locked_total = Money::zero(currency);
for locked in balances_locked
.values()
.filter(|locked| locked.currency == currency)
{
check_predicate_false(
locked.is_negative(),
&format!("locked balance was negative: {locked}"),
)?;
check_predicate_true(
locked.currency.precision == currency.precision,
&format!(
"locked balance precision {} differed from balance precision {} for {currency}",
locked.currency.precision, currency.precision
),
)?;
let reservation = if current_balance.total.is_negative() {
*locked
} else {
(*locked).min(current_balance.total - locked_total)
};
locked_total = locked_total.checked_add(reservation).ok_or_else(|| {
CorrectnessError::PredicateViolation {
message: format!("derived locked balance exceeded Money bounds for {currency}"),
}
})?;
}
let free = current_balance
.total
.checked_sub(locked_total)
.ok_or_else(|| CorrectnessError::PredicateViolation {
message: format!(
"derived free balance exceeded Money bounds for total {} and locked {locked_total}",
current_balance.total
),
})?;
AccountBalance::new_checked(current_balance.total, locked_total, free)
}
fn non_spendable_balance(current_balance: AccountBalance) -> AccountBalance {
let zero = Money::zero(current_balance.currency);
let (locked, free) = if current_balance.total.is_negative() {
(zero, current_balance.total)
} else {
(current_balance.total, zero)
};
AccountBalance {
currency: current_balance.currency,
total: current_balance.total,
locked,
free,
}
}
#[cfg(all(test, feature = "test-support"))]
mod tests {
use rstest::rstest;
use super::*;
use crate::{events::account::stubs::cash_account_state, types::money::MONEY_RAW_MAX};
#[rstest]
fn test_base_purge_account_events_retains_latest_when_all_purged() {
use crate::{
enums::AccountType,
events::account::stubs::cash_account_state,
identifiers::stubs::{account_id, uuid4},
types::{Currency, stubs::stub_account_balance},
};
let mut account = BaseAccount::new(cash_account_state(), true);
let event1 = AccountState::new(
account_id(),
AccountType::Cash,
vec![stub_account_balance()],
vec![],
true,
uuid4(),
UnixNanos::from(100_000_000),
UnixNanos::from(100_000_000),
Some(Currency::USD()),
);
let event2 = AccountState::new(
account_id(),
AccountType::Cash,
vec![stub_account_balance()],
vec![],
true,
uuid4(),
UnixNanos::from(200_000_000),
UnixNanos::from(200_000_000),
Some(Currency::USD()),
);
let event3 = AccountState::new(
account_id(),
AccountType::Cash,
vec![stub_account_balance()],
vec![],
true,
uuid4(),
UnixNanos::from(300_000_000),
UnixNanos::from(300_000_000),
Some(Currency::USD()),
);
account.base_apply(event1);
account.base_apply(event2);
account.base_apply(event3.clone());
assert_eq!(account.events.len(), 4);
account.base_purge_account_events(UnixNanos::from(1_000_000_000), 0);
assert_eq!(account.events.len(), 1);
assert_eq!(account.events[0].ts_event, event3.ts_event);
assert_eq!(account.base_last_event().unwrap().ts_event, event3.ts_event);
}
#[rstest]
fn test_base_purge_account_events_retains_all_for_overflowing_lookback() {
let mut account = BaseAccount::new(cash_account_state(), true);
let mut event = cash_account_state();
event.ts_event = UnixNanos::from(1);
account.base_apply(event);
account.base_purge_account_events(UnixNanos::from(u64::MAX), u64::MAX);
assert_eq!(account.events.len(), 2);
}
#[rstest]
fn test_base_purge_account_events_retains_future_event_without_overflow() {
let mut event = cash_account_state();
event.ts_event = UnixNanos::from(u64::MAX - 1);
let mut account = BaseAccount::new(event, true);
account.base_purge_account_events(UnixNanos::from(u64::MAX), 60);
assert_eq!(account.events.len(), 1);
}
#[rstest]
#[should_panic(
expected = r#"lhs_param: "event.account_id", rhs_param: "self.id", lhs: "OTHER-001", rhs: "SIM-001""#
)]
fn test_base_apply_panics_on_different_account_id() {
let mut account = BaseAccount::new(cash_account_state(), true);
let mut event = cash_account_state();
event.account_id = AccountId::from("OTHER-001");
account.base_apply(event);
}
fn usd_balances(total: &str) -> IndexMap<Currency, AccountBalance> {
let total = Money::from(total);
let mut balances = IndexMap::new();
balances.insert(
Currency::USD(),
AccountBalance::new(total, Money::zero(Currency::USD()), total),
);
balances
}
fn mismatched_usd() -> Currency {
Currency::new(
"USD",
Currency::USD().precision + 1,
840,
"United States dollar",
crate::enums::CurrencyType::Fiat,
)
}
#[rstest]
#[case::observed_currency(true)]
#[case::unobserved_currency(false)]
fn test_update_balance_locked_rejects_negative_without_mutation(#[case] observed: bool) {
let mut balances = if observed {
usd_balances("1000 USD")
} else {
IndexMap::new()
};
let balances_before = balances.clone();
let mut balances_locked = AHashMap::new();
let instrument_id = InstrumentId::from("AUD/USD.SIM");
let error = update_balance_locked(
&mut balances,
&mut balances_locked,
instrument_id,
Money::from("-1 USD"),
)
.unwrap_err();
assert_eq!(error.to_string(), "locked balance was negative: -1.00 USD");
assert!(balances_locked.is_empty());
assert_eq!(balances, balances_before);
}
#[rstest]
fn test_update_balance_locked_restores_prior_reservation_on_failure() {
let usd = Currency::USD();
let mut balances = usd_balances("1000 USD");
let stale_key = (InstrumentId::from("EUR/USD.SIM"), mismatched_usd());
let stale = Money::from_decimal(Decimal::from(10), mismatched_usd()).unwrap();
let mut balances_locked = AHashMap::from([(stale_key, stale)]);
let instrument_id = InstrumentId::from("AUD/USD.SIM");
let result = update_balance_locked(
&mut balances,
&mut balances_locked,
instrument_id,
Money::from("100 USD"),
);
assert!(result.is_err());
assert_eq!(balances_locked, AHashMap::from([(stale_key, stale)]));
assert_eq!(balances, usd_balances("1000 USD"));
assert_eq!(balances[&usd].free, Money::from("1000 USD"));
}
#[rstest]
#[case::positive_total("1000 USD", "1000 USD", "0 USD")]
#[case::negative_total("-1000 USD", "0 USD", "-1000 USD")]
fn test_recalculate_balance_degrades_to_non_spendable_for_invalid_reservation(
#[case] total: &str,
#[case] expected_locked: &str,
#[case] expected_free: &str,
) {
use crate::{enums::CurrencyType, types::Currency};
let usd = Currency::USD();
let total = Money::from(total);
let mut balances = IndexMap::new();
balances.insert(usd, AccountBalance::new(total, Money::zero(usd), total));
let mismatched_usd = Currency::new(
"USD",
usd.precision + 1,
840,
"United States dollar",
CurrencyType::Fiat,
);
let mut balances_locked = AHashMap::new();
balances_locked.insert(
(InstrumentId::from("AUD/USD.SIM"), mismatched_usd),
Money::from_decimal(Decimal::from(100), mismatched_usd).unwrap(),
);
recalculate_balance(&mut balances, &balances_locked, usd);
let balance = balances.get(&usd).expect("balance should be retained");
assert_eq!(balance.total, total);
assert_eq!(balance.locked, Money::from(expected_locked));
assert_eq!(balance.free, Money::from(expected_free));
}
#[rstest]
fn test_update_commissions_sub_canonical_raw_skipped() {
use crate::{
events::account::stubs::cash_account_state,
types::{Currency, Money},
};
let mut account = BaseAccount::new(cash_account_state(), true);
let usd = Currency::USD();
account.update_commissions(Money::from_raw(1, usd));
assert!(account.commission(&usd).is_none());
}
#[rstest]
fn test_try_update_commissions_overflow_preserves_total() {
let mut account = BaseAccount::new(cash_account_state(), true);
let usd = Currency::USD();
let maximum = Money::from_raw(MONEY_RAW_MAX, usd);
account.try_update_commissions(maximum).unwrap();
let result = account.try_update_commissions(Money::from("0.01 USD"));
assert!(result.is_err());
assert_eq!(account.commission(&usd), Some(maximum));
}
}