use enum_dispatch::enum_dispatch;
use indexmap::IndexMap;
use nautilus_core::correctness::{CorrectnessResult, CorrectnessResultExt, FAILED};
use serde::{Deserialize, Serialize};
use crate::{
accounts::{Account, BettingAccount, CashAccount, MarginAccount, WalletAccount},
enums::{AccountType, LiquiditySide},
events::{AccountState, OrderFilled},
identifiers::AccountId,
instruments::InstrumentAny,
position::Position,
types::{AccountBalance, Currency, Money, Price, Quantity},
};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[enum_dispatch(Account)]
pub enum AccountAny {
Margin(MarginAccount),
Cash(CashAccount),
Betting(BettingAccount),
Wallet(WalletAccount),
}
impl AccountAny {
#[must_use]
pub fn id(&self) -> AccountId {
match self {
Self::Margin(margin) => margin.id,
Self::Cash(cash) => cash.id,
Self::Betting(betting) => betting.id,
Self::Wallet(wallet) => wallet.id,
}
}
#[must_use]
pub fn last_event(&self) -> Option<AccountState> {
match self {
Self::Margin(margin) => margin.last_event(),
Self::Cash(cash) => cash.last_event(),
Self::Betting(betting) => betting.last_event(),
Self::Wallet(wallet) => wallet.last_event(),
}
}
#[must_use]
pub fn events(&self) -> Vec<AccountState> {
match self {
Self::Margin(margin) => margin.events(),
Self::Cash(cash) => cash.events(),
Self::Betting(betting) => betting.events(),
Self::Wallet(wallet) => wallet.events(),
}
}
pub fn apply(&mut 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
);
match self {
Self::Margin(margin) => margin.apply(event),
Self::Cash(cash) => cash.apply(event),
Self::Betting(betting) => betting.apply(event),
Self::Wallet(wallet) => wallet.apply(event),
}
}
pub fn set_calculate_account_state(&mut self, calculate_account_state: bool) {
match self {
Self::Margin(margin) => margin.base.calculate_account_state = calculate_account_state,
Self::Cash(cash) => cash.base.calculate_account_state = calculate_account_state,
Self::Betting(betting) => {
betting.base.calculate_account_state = calculate_account_state;
}
Self::Wallet(wallet) => {
wallet.base.calculate_account_state = calculate_account_state;
}
}
}
#[must_use]
pub fn balances(&self) -> IndexMap<Currency, AccountBalance> {
match self {
Self::Margin(margin) => margin.balances(),
Self::Cash(cash) => cash.balances(),
Self::Betting(betting) => betting.balances(),
Self::Wallet(wallet) => wallet.balances(),
}
}
#[must_use]
pub fn balances_locked(&self) -> IndexMap<Currency, Money> {
match self {
Self::Margin(margin) => margin.balances_locked(),
Self::Cash(cash) => cash.balances_locked(),
Self::Betting(betting) => betting.balances_locked(),
Self::Wallet(wallet) => wallet.balances_locked(),
}
}
#[must_use]
pub fn base_currency(&self) -> Option<Currency> {
match self {
Self::Margin(margin) => margin.base_currency(),
Self::Cash(cash) => cash.base_currency(),
Self::Betting(betting) => betting.base_currency(),
Self::Wallet(wallet) => wallet.base_currency(),
}
}
pub fn from_events(events: &[AccountState]) -> anyhow::Result<Self> {
let Some((init_event, remaining_events)) = events.split_first() else {
anyhow::bail!("No account events provided to create `AccountAny`");
};
let mut account = Self::from_state_checked(init_event.clone())?;
for event in remaining_events {
account.apply(event.clone())?;
}
Ok(account)
}
pub fn calculate_pnls(
&self,
instrument: &InstrumentAny,
fill: &OrderFilled,
position: Option<Position>,
) -> anyhow::Result<Vec<Money>> {
match self {
Self::Margin(margin) => margin.calculate_pnls(instrument, fill, position),
Self::Cash(cash) => cash.calculate_pnls(instrument, fill, position),
Self::Betting(betting) => betting.calculate_pnls(instrument, fill, position),
Self::Wallet(wallet) => wallet.calculate_pnls(instrument, fill, position),
}
}
pub fn calculate_commission(
&self,
instrument: &InstrumentAny,
last_qty: Quantity,
last_px: Price,
liquidity_side: LiquiditySide,
use_quote_for_inverse: Option<bool>,
) -> anyhow::Result<Money> {
match self {
Self::Margin(margin) => margin.calculate_commission(
instrument,
last_qty,
last_px,
liquidity_side,
use_quote_for_inverse,
),
Self::Cash(cash) => cash.calculate_commission(
instrument,
last_qty,
last_px,
liquidity_side,
use_quote_for_inverse,
),
Self::Betting(betting) => betting.calculate_commission(
instrument,
last_qty,
last_px,
liquidity_side,
use_quote_for_inverse,
),
Self::Wallet(wallet) => wallet.calculate_commission(
instrument,
last_qty,
last_px,
liquidity_side,
use_quote_for_inverse,
),
}
}
#[must_use]
pub fn balance(&self, currency: Option<Currency>) -> Option<&AccountBalance> {
match self {
Self::Margin(margin) => margin.balance(currency),
Self::Cash(cash) => cash.balance(currency),
Self::Betting(betting) => betting.balance(currency),
Self::Wallet(wallet) => wallet.balance(currency),
}
}
}
impl AccountAny {
pub fn try_from_state(event: AccountState) -> Result<Self, &'static str> {
Self::from_state_checked(event).map_err(|_| "Invalid wallet account state")
}
fn from_state_checked(event: AccountState) -> CorrectnessResult<Self> {
match event.account_type {
AccountType::Margin => Ok(Self::Margin(MarginAccount::new(event, false))),
AccountType::Cash => Ok(Self::Cash(CashAccount::new(event, false, false))),
AccountType::Betting => Ok(Self::Betting(BettingAccount::new(event, false))),
AccountType::Wallet => Ok(Self::Wallet(WalletAccount::new_checked(event, false)?)),
}
}
}
impl From<AccountState> for AccountAny {
fn from(event: AccountState) -> Self {
Self::from_state_checked(event).expect_display(FAILED)
}
}
impl PartialEq for AccountAny {
fn eq(&self, other: &Self) -> bool {
self.id() == other.id()
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use crate::{
accounts::{Account, AccountAny},
events::{AccountState, account::stubs::*},
identifiers::AccountId,
};
#[rstest]
fn test_from_events_empty_returns_error() {
let events: Vec<AccountState> = vec![];
let result = AccountAny::from_events(&events);
assert_eq!(
result.unwrap_err().to_string(),
"No account events provided to create `AccountAny`"
);
}
#[rstest]
fn test_from_events_single_cash_event(cash_account_state: AccountState) {
let result = AccountAny::from_events(&[cash_account_state]);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Cash(_)));
}
#[rstest]
fn test_from_events_rejects_different_account(cash_account_state: AccountState) {
let mut different_account = cash_account_state.clone();
different_account.account_id = AccountId::from("OTHER-001");
let result = AccountAny::from_events(&[cash_account_state, different_account]);
assert_eq!(
result.unwrap_err().to_string(),
"Account event had a different account ID: expected SIM-001, received OTHER-001"
);
}
#[rstest]
fn test_from_events_single_margin_event(margin_account_state: AccountState) {
let result = AccountAny::from_events(&[margin_account_state]);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Margin(_)));
}
#[rstest]
fn test_try_from_state_cash(cash_account_state: AccountState) {
let result: Result<AccountAny, &'static str> =
AccountAny::try_from_state(cash_account_state);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Cash(_)));
}
#[rstest]
fn test_try_from_state_margin(margin_account_state: AccountState) {
let result = AccountAny::try_from_state(margin_account_state);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Margin(_)));
}
#[rstest]
fn test_try_from_state_betting(betting_account_state: AccountState) {
let result = AccountAny::try_from_state(betting_account_state);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Betting(_)));
}
#[rstest]
fn test_try_from_state_wallet(wallet_account_state: AccountState) {
let result = AccountAny::try_from_state(wallet_account_state);
assert!(result.is_ok());
assert!(matches!(result.unwrap(), AccountAny::Wallet(_)));
}
#[rstest]
fn test_try_from_state_invalid_wallet_returns_static_error() {
let result: Result<AccountAny, &'static str> =
AccountAny::try_from_state(invalid_wallet_state());
assert_eq!(result.unwrap_err(), "Invalid wallet account state");
}
#[rstest]
fn test_from_events_wallet_applies_sequence(
wallet_account_state: AccountState,
wallet_account_state_changed: AccountState,
) {
let result = AccountAny::from_events(&[wallet_account_state, wallet_account_state_changed]);
assert!(result.is_ok());
let account = result.unwrap();
assert!(matches!(account, AccountAny::Wallet(_)));
assert_eq!(account.event_count(), 2);
}
#[rstest]
fn test_from_events_wallet_rejects_negative_initial_balance() {
let result = AccountAny::from_events(&[invalid_wallet_state()]);
assert!(result.is_err());
assert_eq!(
result.unwrap_err().to_string(),
"Wallet account balance total was negative"
);
}
#[rstest]
#[case::cash(cash_account_state(), "Cash")]
#[case::margin(margin_account_state(), "Margin")]
#[case::betting(betting_account_state(), "Betting")]
#[case::wallet(wallet_account_state(), "Wallet")]
fn test_serde_round_trip_preserves_variant_payload(
#[case] state: AccountState,
#[case] expected_variant: &str,
) {
let account = AccountAny::try_from_state(state).unwrap();
let value = serde_json::to_value(&account).unwrap();
let object = value.as_object().unwrap();
assert_eq!(object.len(), 1);
assert!(object.contains_key(expected_variant));
let deserialized: AccountAny = serde_json::from_value(value).unwrap();
assert_eq!(deserialized.id(), account.id());
assert_eq!(deserialized.events(), account.events());
assert_eq!(deserialized.balances(), account.balances());
}
#[rstest]
#[case::cash(include_str!("../../test_data/account_legacy_cash.json"), "Cash")]
#[case::margin(include_str!("../../test_data/account_legacy_margin.json"), "Margin")]
#[case::betting(include_str!("../../test_data/account_legacy_betting.json"), "Betting")]
fn test_deserializes_legacy_payload(#[case] json: &str, #[case] expected_variant: &str) {
let account: AccountAny = serde_json::from_str(json).unwrap();
let variant = match &account {
AccountAny::Cash(_) => "Cash",
AccountAny::Margin(_) => "Margin",
AccountAny::Betting(_) => "Betting",
AccountAny::Wallet(_) => "Wallet",
};
assert_eq!(variant, expected_variant);
assert_eq!(account.event_count(), 1);
}
fn invalid_wallet_state() -> AccountState {
AccountState::new(
AccountId::from("WALLET-001"),
crate::enums::AccountType::Wallet,
vec![crate::types::AccountBalance::new(
crate::types::Money::from("-1 ETH"),
crate::types::Money::from("0 ETH"),
crate::types::Money::from("-1 ETH"),
)],
vec![],
true,
crate::identifiers::stubs::uuid4(),
0.into(),
0.into(),
None,
)
}
}