1use std::{fmt::Debug, sync::Arc};
19
20use rust_decimal::Decimal;
21
22use crate::{
23 instruments::Instrument,
24 types::{Money, Price, Quantity},
25};
26
27pub trait MarginModel: Send + Sync {
29 #[must_use]
31 fn name(&self) -> &'static str;
32
33 fn calculate_initial_margin(
39 &self,
40 instrument: &dyn Instrument,
41 quantity: Quantity,
42 price: Price,
43 leverage: Decimal,
44 use_quote_for_inverse: Option<bool>,
45 ) -> anyhow::Result<Money>;
46
47 fn calculate_maintenance_margin(
53 &self,
54 instrument: &dyn Instrument,
55 quantity: Quantity,
56 price: Price,
57 leverage: Decimal,
58 use_quote_for_inverse: Option<bool>,
59 ) -> anyhow::Result<Money>;
60}
61
62#[derive(Clone)]
64pub struct MarginModelHandle(Arc<dyn MarginModel>);
65
66impl MarginModelHandle {
67 #[must_use]
69 pub fn new<T>(model: T) -> Self
70 where
71 T: MarginModel + 'static,
72 {
73 Self(Arc::new(model))
74 }
75
76 #[must_use]
78 pub fn from_arc(model: Arc<dyn MarginModel>) -> Self {
79 Self(model)
80 }
81}
82
83impl Debug for MarginModelHandle {
84 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
85 f.debug_tuple(stringify!(MarginModelHandle))
86 .field(&"<dyn MarginModel>")
87 .finish()
88 }
89}
90
91impl MarginModel for MarginModelHandle {
92 fn name(&self) -> &'static str {
93 self.0.name()
94 }
95
96 fn calculate_initial_margin(
97 &self,
98 instrument: &dyn Instrument,
99 quantity: Quantity,
100 price: Price,
101 leverage: Decimal,
102 use_quote_for_inverse: Option<bool>,
103 ) -> anyhow::Result<Money> {
104 self.0.calculate_initial_margin(
105 instrument,
106 quantity,
107 price,
108 leverage,
109 use_quote_for_inverse,
110 )
111 }
112
113 fn calculate_maintenance_margin(
114 &self,
115 instrument: &dyn Instrument,
116 quantity: Quantity,
117 price: Price,
118 leverage: Decimal,
119 use_quote_for_inverse: Option<bool>,
120 ) -> anyhow::Result<Money> {
121 self.0.calculate_maintenance_margin(
122 instrument,
123 quantity,
124 price,
125 leverage,
126 use_quote_for_inverse,
127 )
128 }
129}
130
131#[derive(Debug, Clone)]
133pub enum MarginModelAny {
134 Standard(StandardMarginModel),
135 Leveraged(LeveragedMarginModel),
136}
137
138impl MarginModel for MarginModelAny {
139 fn name(&self) -> &'static str {
140 match self {
141 Self::Standard(model) => model.name(),
142 Self::Leveraged(model) => model.name(),
143 }
144 }
145
146 fn calculate_initial_margin(
147 &self,
148 instrument: &dyn Instrument,
149 quantity: Quantity,
150 price: Price,
151 leverage: Decimal,
152 use_quote_for_inverse: Option<bool>,
153 ) -> anyhow::Result<Money> {
154 match self {
155 Self::Standard(m) => m.calculate_initial_margin(
156 instrument,
157 quantity,
158 price,
159 leverage,
160 use_quote_for_inverse,
161 ),
162 Self::Leveraged(m) => m.calculate_initial_margin(
163 instrument,
164 quantity,
165 price,
166 leverage,
167 use_quote_for_inverse,
168 ),
169 }
170 }
171
172 fn calculate_maintenance_margin(
173 &self,
174 instrument: &dyn Instrument,
175 quantity: Quantity,
176 price: Price,
177 leverage: Decimal,
178 use_quote_for_inverse: Option<bool>,
179 ) -> anyhow::Result<Money> {
180 match self {
181 Self::Standard(m) => m.calculate_maintenance_margin(
182 instrument,
183 quantity,
184 price,
185 leverage,
186 use_quote_for_inverse,
187 ),
188 Self::Leveraged(m) => m.calculate_maintenance_margin(
189 instrument,
190 quantity,
191 price,
192 leverage,
193 use_quote_for_inverse,
194 ),
195 }
196 }
197}
198
199impl Default for MarginModelAny {
200 fn default() -> Self {
201 Self::Leveraged(LeveragedMarginModel)
202 }
203}
204
205impl Default for MarginModelHandle {
206 fn default() -> Self {
207 MarginModelAny::default().into()
208 }
209}
210
211impl From<MarginModelAny> for MarginModelHandle {
212 fn from(model: MarginModelAny) -> Self {
213 Self::new(model)
214 }
215}
216
217fn margin_currency(
219 instrument: &dyn Instrument,
220 use_quote_for_inverse: bool,
221) -> anyhow::Result<crate::types::Currency> {
222 if instrument.is_inverse() && !use_quote_for_inverse {
223 instrument.base_currency().ok_or_else(|| {
224 anyhow::anyhow!(
225 "Inverse instrument {} has no base currency",
226 instrument.id()
227 )
228 })
229 } else {
230 Ok(instrument.quote_currency())
231 }
232}
233
234#[derive(Debug, Clone, Copy)]
240#[cfg_attr(
241 feature = "python",
242 pyo3::pyclass(module = "nautilus_trader.model", from_py_object)
243)]
244#[cfg_attr(
245 feature = "python",
246 pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.model")
247)]
248pub struct StandardMarginModel;
249
250impl MarginModel for StandardMarginModel {
251 fn name(&self) -> &'static str {
252 "standard"
253 }
254
255 fn calculate_initial_margin(
256 &self,
257 instrument: &dyn Instrument,
258 quantity: Quantity,
259 price: Price,
260 _leverage: Decimal,
261 use_quote_for_inverse: Option<bool>,
262 ) -> anyhow::Result<Money> {
263 let use_quote = use_quote_for_inverse.unwrap_or(false);
264 let notional = instrument.try_calculate_notional_value(quantity, price, Some(use_quote))?;
265 let margin = notional
268 .as_decimal()
269 .abs()
270 .checked_mul(instrument.margin_init())
271 .ok_or_else(|| anyhow::anyhow!("initial margin calculation overflow"))?;
272 let currency = margin_currency(instrument, use_quote)?;
273 Money::from_decimal(margin, currency).map_err(Into::into)
274 }
275
276 fn calculate_maintenance_margin(
277 &self,
278 instrument: &dyn Instrument,
279 quantity: Quantity,
280 price: Price,
281 _leverage: Decimal,
282 use_quote_for_inverse: Option<bool>,
283 ) -> anyhow::Result<Money> {
284 let use_quote = use_quote_for_inverse.unwrap_or(false);
285 let notional = instrument.try_calculate_notional_value(quantity, price, Some(use_quote))?;
286 let margin = notional
287 .as_decimal()
288 .abs()
289 .checked_mul(instrument.margin_maint())
290 .ok_or_else(|| anyhow::anyhow!("maintenance margin calculation overflow"))?;
291 let currency = margin_currency(instrument, use_quote)?;
292 Money::from_decimal(margin, currency).map_err(Into::into)
293 }
294}
295
296#[derive(Debug, Clone, Copy)]
302#[cfg_attr(
303 feature = "python",
304 pyo3::pyclass(module = "nautilus_trader.model", from_py_object)
305)]
306#[cfg_attr(
307 feature = "python",
308 pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.model")
309)]
310pub struct LeveragedMarginModel;
311
312impl MarginModel for LeveragedMarginModel {
313 fn name(&self) -> &'static str {
314 "leveraged"
315 }
316
317 fn calculate_initial_margin(
318 &self,
319 instrument: &dyn Instrument,
320 quantity: Quantity,
321 price: Price,
322 leverage: Decimal,
323 use_quote_for_inverse: Option<bool>,
324 ) -> anyhow::Result<Money> {
325 if leverage <= Decimal::ZERO {
326 anyhow::bail!("Invalid leverage {leverage} for {}", instrument.id());
327 }
328 let use_quote = use_quote_for_inverse.unwrap_or(false);
329 let notional = instrument.try_calculate_notional_value(quantity, price, Some(use_quote))?;
330 let margin = notional
331 .as_decimal()
332 .abs()
333 .checked_div(leverage)
334 .and_then(|adjusted| adjusted.checked_mul(instrument.margin_init()))
335 .ok_or_else(|| anyhow::anyhow!("initial margin calculation overflow"))?;
336 let currency = margin_currency(instrument, use_quote)?;
337 Money::from_decimal(margin, currency).map_err(Into::into)
338 }
339
340 fn calculate_maintenance_margin(
341 &self,
342 instrument: &dyn Instrument,
343 quantity: Quantity,
344 price: Price,
345 leverage: Decimal,
346 use_quote_for_inverse: Option<bool>,
347 ) -> anyhow::Result<Money> {
348 if leverage <= Decimal::ZERO {
349 anyhow::bail!("Invalid leverage {leverage} for {}", instrument.id());
350 }
351 let use_quote = use_quote_for_inverse.unwrap_or(false);
352 let notional = instrument.try_calculate_notional_value(quantity, price, Some(use_quote))?;
353 let margin = notional
354 .as_decimal()
355 .abs()
356 .checked_div(leverage)
357 .and_then(|adjusted| adjusted.checked_mul(instrument.margin_maint()))
358 .ok_or_else(|| anyhow::anyhow!("maintenance margin calculation overflow"))?;
359 let currency = margin_currency(instrument, use_quote)?;
360 Money::from_decimal(margin, currency).map_err(Into::into)
361 }
362}
363
364#[cfg(test)]
365mod tests {
366 use rstest::rstest;
367 use rust_decimal::Decimal;
368 use rust_decimal_macros::dec;
369 use ustr::Ustr;
370
371 use super::*;
372 use crate::{
373 enums::AssetClass,
374 identifiers::{InstrumentId, Symbol},
375 instruments::{
376 CryptoPerpetual, FuturesSpread, Instrument, stubs::crypto_perpetual_ethusdt,
377 },
378 types::{Currency, Price, Quantity},
379 };
380
381 struct FixedMarginModel {
382 initial: Money,
383 maintenance: Money,
384 }
385
386 impl MarginModel for FixedMarginModel {
387 fn name(&self) -> &'static str {
388 "fixed"
389 }
390
391 fn calculate_initial_margin(
392 &self,
393 _instrument: &dyn Instrument,
394 _quantity: Quantity,
395 _price: Price,
396 _leverage: Decimal,
397 _use_quote_for_inverse: Option<bool>,
398 ) -> anyhow::Result<Money> {
399 Ok(self.initial)
400 }
401
402 fn calculate_maintenance_margin(
403 &self,
404 _instrument: &dyn Instrument,
405 _quantity: Quantity,
406 _price: Price,
407 _leverage: Decimal,
408 _use_quote_for_inverse: Option<bool>,
409 ) -> anyhow::Result<Money> {
410 Ok(self.maintenance)
411 }
412 }
413
414 fn ethusdt() -> CryptoPerpetual {
415 crypto_perpetual_ethusdt()
416 }
417
418 #[rstest]
419 fn test_leveraged_initial_margin() {
420 let model = LeveragedMarginModel;
421 let instrument = ethusdt();
422 let quantity = Quantity::from("10.000");
423 let price = Price::from("5000.00");
424 let leverage = dec!(10);
425
426 let margin = model
427 .calculate_initial_margin(&instrument, quantity, price, leverage, None)
428 .unwrap();
429
430 let expected = Decimal::from(50000) / leverage * instrument.margin_init();
433 assert_eq!(margin.as_decimal(), expected);
434 assert_eq!(margin.currency, Currency::USDT());
435 }
436
437 #[rstest]
438 fn test_standard_ignores_leverage() {
439 let model = StandardMarginModel;
440 let instrument = ethusdt();
441 let quantity = Quantity::from("10.000");
442 let price = Price::from("5000.00");
443
444 let margin_low = model
445 .calculate_initial_margin(&instrument, quantity, price, dec!(2), None)
446 .unwrap();
447 let margin_high = model
448 .calculate_initial_margin(&instrument, quantity, price, dec!(100), None)
449 .unwrap();
450
451 assert_eq!(margin_low, margin_high);
453 }
454
455 fn negative_price_spread() -> FuturesSpread {
459 FuturesSpread::builder()
460 .instrument_id(InstrumentId::from("ESM4-ESU4.GLBX"))
461 .raw_symbol(Symbol::from("ESM4-ESU4"))
462 .asset_class(AssetClass::Index)
463 .underlying(Ustr::from("ES"))
464 .strategy_type(Ustr::from("EQ"))
465 .activation_ns(1_000.into())
466 .expiration_ns(2_000.into())
467 .currency(Currency::USD())
468 .price_precision(2)
469 .price_increment(Price::from("0.01"))
470 .multiplier(Quantity::from(50))
471 .lot_size(Quantity::from(1))
472 .margin_init(dec!(0.01))
473 .margin_maint(dec!(0.02))
474 .ts_event(1.into())
475 .ts_init(2.into())
476 .build()
477 .unwrap()
478 }
479
480 #[rstest]
481 fn test_standard_margin_is_positive_for_a_negative_price() {
482 let model = StandardMarginModel;
483 let instrument = negative_price_spread();
484 let quantity = Quantity::from(2);
485 let positive = Price::from("2.00");
486 let negative = Price::from("-2.00");
487
488 let initial = model
489 .calculate_initial_margin(&instrument, quantity, negative, dec!(1), None)
490 .unwrap();
491 let maintenance = model
492 .calculate_maintenance_margin(&instrument, quantity, negative, dec!(1), None)
493 .unwrap();
494
495 assert_eq!(initial.as_decimal(), dec!(2));
497 assert_eq!(maintenance.as_decimal(), dec!(4));
498 assert_eq!(
500 initial,
501 model
502 .calculate_initial_margin(&instrument, quantity, positive, dec!(1), None)
503 .unwrap()
504 );
505 }
506
507 #[rstest]
508 fn test_leveraged_margin_is_positive_for_a_negative_price() {
509 let model = LeveragedMarginModel;
510 let instrument = negative_price_spread();
511 let quantity = Quantity::from(2);
512 let negative = Price::from("-2.00");
513 let leverage = dec!(10);
514
515 let initial = model
516 .calculate_initial_margin(&instrument, quantity, negative, leverage, None)
517 .unwrap();
518 let maintenance = model
519 .calculate_maintenance_margin(&instrument, quantity, negative, leverage, None)
520 .unwrap();
521
522 assert_eq!(initial.as_decimal(), dec!(0.2));
524 assert_eq!(maintenance.as_decimal(), dec!(0.4));
525 }
526
527 #[rstest]
528 #[case::zero(Decimal::ZERO)]
529 #[case::negative(dec!(-1))]
530 fn test_leveraged_rejects_non_positive_leverage(#[case] leverage: Decimal) {
531 let model = LeveragedMarginModel;
532 let instrument = ethusdt();
533 let quantity = Quantity::from("1.000");
534 let price = Price::from("5000.00");
535 let expected = format!("Invalid leverage {leverage} for {}", instrument.id());
536
537 let initial = model.calculate_initial_margin(&instrument, quantity, price, leverage, None);
538 let maintenance =
539 model.calculate_maintenance_margin(&instrument, quantity, price, leverage, None);
540
541 assert_eq!(initial.unwrap_err().to_string(), expected);
542 assert_eq!(maintenance.unwrap_err().to_string(), expected);
543 }
544
545 #[rstest]
546 fn test_leveraged_margin_decimal_overflow_returns_error() {
547 let model = LeveragedMarginModel;
548 let instrument = ethusdt();
549 let quantity = Quantity::from("1.000");
550 let price = Price::from("5000.00");
551 let leverage = Decimal::new(1, 28);
552
553 let initial = model.calculate_initial_margin(&instrument, quantity, price, leverage, None);
554 let maintenance =
555 model.calculate_maintenance_margin(&instrument, quantity, price, leverage, None);
556
557 assert_eq!(
558 initial.unwrap_err().to_string(),
559 "initial margin calculation overflow"
560 );
561 assert_eq!(
562 maintenance.unwrap_err().to_string(),
563 "maintenance margin calculation overflow"
564 );
565 }
566
567 #[rstest]
568 fn test_margin_model_any_default_is_leveraged() {
569 let model = MarginModelAny::default();
570 let instrument = ethusdt();
571 let quantity = Quantity::from("10.000");
572 let price = Price::from("5000.00");
573 let leverage = dec!(10);
574
575 let initial = model
576 .calculate_initial_margin(&instrument, quantity, price, leverage, None)
577 .unwrap();
578 let maintenance = model
579 .calculate_maintenance_margin(&instrument, quantity, price, leverage, None)
580 .unwrap();
581
582 assert!(matches!(model, MarginModelAny::Leveraged(_)));
584 assert_eq!(model.name(), "leveraged");
585 assert_eq!(
586 initial.as_decimal(),
587 Decimal::from(5000) * instrument.margin_init()
588 );
589 assert_eq!(
590 maintenance.as_decimal(),
591 Decimal::from(5000) * instrument.margin_maint()
592 );
593 }
594
595 #[rstest]
596 fn test_margin_model_any_standard_dispatches_to_the_standard_model() {
597 let model = MarginModelAny::Standard(StandardMarginModel);
598 let instrument = ethusdt();
599 let quantity = Quantity::from("10.000");
600 let price = Price::from("5000.00");
601 let leverage = dec!(10);
602
603 let initial = model
604 .calculate_initial_margin(&instrument, quantity, price, leverage, None)
605 .unwrap();
606 let maintenance = model
607 .calculate_maintenance_margin(&instrument, quantity, price, leverage, None)
608 .unwrap();
609
610 assert_eq!(model.name(), "standard");
612 assert_eq!(
613 initial.as_decimal(),
614 Decimal::from(50000) * instrument.margin_init()
615 );
616 assert_eq!(
617 maintenance.as_decimal(),
618 Decimal::from(50000) * instrument.margin_maint()
619 );
620 assert_eq!(initial.currency, Currency::USDT());
621 }
622
623 #[rstest]
624 fn test_margin_model_handle_calls_custom_model() {
625 let initial = Money::from("12.34 USDT");
626 let maintenance = Money::from("5.67 USDT");
627 let model: Arc<dyn MarginModel> = Arc::new(FixedMarginModel {
628 initial,
629 maintenance,
630 });
631 let handle = MarginModelHandle::from_arc(model);
632 let cloned_handle = handle.clone();
633 drop(handle);
634 let instrument = ethusdt();
635
636 let initial_result = cloned_handle
637 .calculate_initial_margin(
638 &instrument,
639 Quantity::from("1.000"),
640 Price::from("5000.00"),
641 dec!(10),
642 None,
643 )
644 .unwrap();
645 let maintenance_result = cloned_handle
646 .calculate_maintenance_margin(
647 &instrument,
648 Quantity::from("1.000"),
649 Price::from("5000.00"),
650 dec!(10),
651 None,
652 )
653 .unwrap();
654
655 assert_eq!(cloned_handle.name(), "fixed");
656 assert_eq!(initial_result, initial);
657 assert_eq!(maintenance_result, maintenance);
658 }
659
660 #[rstest]
661 fn test_maintenance_margin() {
662 let model = LeveragedMarginModel;
663 let instrument = ethusdt();
664 let quantity = Quantity::from("10.000");
665 let price = Price::from("5000.00");
666 let leverage = dec!(10);
667
668 let margin = model
669 .calculate_maintenance_margin(&instrument, quantity, price, leverage, None)
670 .unwrap();
671
672 let expected = Decimal::from(50000) / leverage * instrument.margin_maint();
673 assert_eq!(margin.as_decimal(), expected);
674 }
675}