470 lines
19 KiB
Python
470 lines
19 KiB
Python
|
|
"""
|
||
|
|
Comprehensive tests for Discrepancy Tracker, Feature Importance, Rollback, Stress (60+ tests).
|
||
|
|
"""
|
||
|
|
import pytest
|
||
|
|
from malkhut.state import (
|
||
|
|
AccountState, FulfilmentPolicyParams, MarketWorldState, Mode,
|
||
|
|
OrderBookState, PriceLevel, Side, TradePathState, VenueRules,
|
||
|
|
)
|
||
|
|
from malkhut.actions import ActionKind, FulfilmentAction
|
||
|
|
from malkhut.training.discrepancy import DiscrepancyTracker
|
||
|
|
from malkhut.training.importance import FeatureImportanceTracker
|
||
|
|
from malkhut.training.rollback import PolicyRollback
|
||
|
|
from malkhut.training.stress import StressScenarioFactory
|
||
|
|
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
||
|
|
|
||
|
|
|
||
|
|
def _venue():
|
||
|
|
return VenueRules(exchange="bingx", symbol="BTCUSDT", tick_size=0.1, lot_size=0.001,
|
||
|
|
min_qty=0.001, min_notional=5.0, maker_fee_bps=-0.2, taker_fee_bps=0.5,
|
||
|
|
post_only_supported=True, reduce_only_supported=True,
|
||
|
|
max_orders_per_second=100, max_cancels_per_minute=120)
|
||
|
|
|
||
|
|
|
||
|
|
def _state(**kw):
|
||
|
|
tp = kw.get("trade_path")
|
||
|
|
return MarketWorldState(
|
||
|
|
ts_ns=kw.get("ts", 1_000_000_000), mode=Mode.REPLAY_NO_IMPACT, venue=_venue(),
|
||
|
|
book=OrderBookState(ts_ns=1, symbol="BTCUSDT",
|
||
|
|
bids=(PriceLevel(kw.get("bid", 50000.0), 1.0),),
|
||
|
|
asks=(PriceLevel(kw.get("ask", 50001.0), 1.0),)),
|
||
|
|
account=AccountState(ts_ns=1, equity=10000.0, wallet_balance=10000.0,
|
||
|
|
available_balance=10000.0, margin_used=0.0, total_notional=0.0),
|
||
|
|
trade_path=tp,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _baseline(**kw):
|
||
|
|
d = dict(version="baseline", ucb_c=1.414, max_sims=256, max_depth=3,
|
||
|
|
rollout_depth=3, root_temperature=0.5, min_root_entropy=0.25,
|
||
|
|
quote_offsets_ticks=(0, 1, 2), quote_size_fractions=(0.1, 0.25, 0.5),
|
||
|
|
passive_ttl_ms=200, aggressive_ttl_ms=50, maker_edge_min_bps=0.5,
|
||
|
|
cross_spread_edge_min_bps=5.0, adverse_toxicity_cancel_threshold=0.5,
|
||
|
|
queue_churn_cancel_threshold=0.5, mae_tail_cut_bps=50.0,
|
||
|
|
mfe_giveback_cut_fraction=0.5, max_time_in_loss_s=300.0,
|
||
|
|
failed_recovery_cut_count=3, recovery_velocity_min_bps_per_s=0.0,
|
||
|
|
max_symbol_notional_fraction=0.20, max_single_order_notional_fraction=0.05,
|
||
|
|
reduce_when_global_up_fraction=0.30, session_profit_lock_fraction=0.02,
|
||
|
|
w_expected_pnl=1.0, w_fill_probability=0.5, w_adverse_selection=2.0,
|
||
|
|
w_queue_priority=0.5, w_inventory_risk=1.5, w_tail_loss=5.0,
|
||
|
|
w_fee_quality=0.5, w_time_decay=0.3, w_policy_entropy=0.5,
|
||
|
|
robust_tail_weight=2.0, toxic_counterparty_weight=3.0,
|
||
|
|
low_liquidity_weight=2.0, latency_stress_weight=1.0)
|
||
|
|
d.update(kw)
|
||
|
|
return FulfilmentPolicyParams(**d)
|
||
|
|
|
||
|
|
|
||
|
|
def _action(kind=ActionKind.NOOP):
|
||
|
|
return FulfilmentAction(kind, None, None, 0, 0.0, 0)
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# DISCREPANCY TRACKER (25+ tests)
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestDiscrepancyComprehensive:
|
||
|
|
def test_multiple_predictions(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
for i in range(5):
|
||
|
|
dt.record_prediction(_state(ts=i), _action(), f"v{i}")
|
||
|
|
assert dt._last_prediction is not None
|
||
|
|
|
||
|
|
def test_compare_after_each_prediction(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
for i in range(5):
|
||
|
|
dt.record_prediction(_state(ts=i), _action(), "v1")
|
||
|
|
dt.compare_with_actual(_state(ts=i+100))
|
||
|
|
assert dt.total_comparisons == 5
|
||
|
|
|
||
|
|
def test_discrepancy_count_matches(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(ts=1), _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(_state(ts=2))
|
||
|
|
assert len(discs) == dt.total_discrepancies
|
||
|
|
|
||
|
|
def test_rate_calculation(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(ts=1), _action(), "v1")
|
||
|
|
dt.compare_with_actual(_state(ts=2))
|
||
|
|
assert dt.discrepancy_rate > 0
|
||
|
|
|
||
|
|
def test_rate_zero_when_no_comparisons(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
assert dt.discrepancy_rate == 0.0
|
||
|
|
|
||
|
|
def test_different_severities(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(ts=1), _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(_state(ts=2))
|
||
|
|
severities = set(d.severity for d in discs)
|
||
|
|
assert len(severities) > 0
|
||
|
|
|
||
|
|
def test_get_recent_limit(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
for i in range(10):
|
||
|
|
dt.record_prediction(_state(ts=i), _action(), "v1")
|
||
|
|
dt.compare_with_actual(_state(ts=i+100))
|
||
|
|
recent = dt.get_recent(3)
|
||
|
|
assert len(recent) == 3
|
||
|
|
|
||
|
|
def test_get_by_severity_empty(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
result = dt.get_by_severity("nonexistent")
|
||
|
|
assert len(result) == 0
|
||
|
|
|
||
|
|
def test_record_prediction_stores_action(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
a = _action(ActionKind.PLACE)
|
||
|
|
dt.record_prediction(_state(), a, "v1")
|
||
|
|
assert dt._last_action.kind == ActionKind.PLACE
|
||
|
|
|
||
|
|
def test_record_prediction_stores_version(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(), _action(), "my_version")
|
||
|
|
assert dt._last_policy_version == "my_version"
|
||
|
|
|
||
|
|
def test_compare_with_custom_tolerance(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
s1 = _state(ts=1)
|
||
|
|
s2 = _state(ts=1, bid=50000.01) # tiny difference
|
||
|
|
dt.record_prediction(s1, _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(s2, tolerances={"book_price": 0.1})
|
||
|
|
assert len(discs) == 0 # within tolerance
|
||
|
|
|
||
|
|
def test_compare_without_tolerance(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
s1 = _state(ts=1)
|
||
|
|
s2 = _state(ts=1, bid=50000.01)
|
||
|
|
dt.record_prediction(s1, _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(s2)
|
||
|
|
# Default tolerance is 1e-9, so 0.01 difference should be detected
|
||
|
|
assert len(discs) > 0
|
||
|
|
|
||
|
|
def test_multiple_comparisons_accumulate(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
for i in range(20):
|
||
|
|
dt.record_prediction(_state(ts=i), _action(), "v1")
|
||
|
|
dt.compare_with_actual(_state(ts=i+1000))
|
||
|
|
assert dt.total_comparisons == 20
|
||
|
|
assert dt.total_discrepancies > 0
|
||
|
|
|
||
|
|
def test_discrepancy_record_has_all_fields(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(ts=1), _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(_state(ts=2))
|
||
|
|
d = discs[0]
|
||
|
|
assert d.ts_ns > 0
|
||
|
|
assert d.symbol == "BTCUSDT"
|
||
|
|
assert len(d.field) > 0
|
||
|
|
assert d.severity in ("info", "warning", "critical")
|
||
|
|
assert d.action_kind == "NOOP"
|
||
|
|
assert d.policy_version == "v1"
|
||
|
|
|
||
|
|
def test_empty_state_comparison(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
dt.record_prediction(_state(), _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(_state())
|
||
|
|
assert len(discs) == 0 # identical states
|
||
|
|
|
||
|
|
def test_wide_book_difference(self):
|
||
|
|
dt = DiscrepancyTracker()
|
||
|
|
s1 = _state(bid=50000.0, ask=50001.0)
|
||
|
|
s2 = _state(bid=49000.0, ask=51000.0)
|
||
|
|
dt.record_prediction(s1, _action(), "v1")
|
||
|
|
discs = dt.compare_with_actual(s2)
|
||
|
|
assert len(discs) > 0
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# FEATURE IMPORTANCE (20+ tests)
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestFeatureImportanceComprehensive:
|
||
|
|
def test_empty_tracker(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
assert fit.total_decisions == 0
|
||
|
|
assert fit.feature_count == 0
|
||
|
|
|
||
|
|
def test_single_decision(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
assert fit.total_decisions == 1
|
||
|
|
assert fit.feature_count > 0
|
||
|
|
|
||
|
|
def test_importance_all_regimes(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance(regime="normal")
|
||
|
|
assert len(imp) > 0
|
||
|
|
|
||
|
|
def test_importance_no_regime_filter(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance()
|
||
|
|
assert len(imp) > 0
|
||
|
|
|
||
|
|
def test_importance_top_n_limit(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance(top_n=3)
|
||
|
|
assert len(imp) <= 3
|
||
|
|
|
||
|
|
def test_importance_sorted_descending(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(20):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance(top_n=10)
|
||
|
|
for i in range(len(imp) - 1):
|
||
|
|
assert imp[i].importance >= imp[i+1].importance
|
||
|
|
|
||
|
|
def test_feature_stats_mean(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
stats = fit.get_feature_stats("mid")
|
||
|
|
assert "mean" in stats
|
||
|
|
assert stats["mean"] > 0
|
||
|
|
|
||
|
|
def test_feature_stats_min_max(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
stats = fit.get_feature_stats("mid")
|
||
|
|
assert stats["min"] <= stats["max"]
|
||
|
|
|
||
|
|
def test_feature_stats_count(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(5):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
stats = fit.get_feature_stats("mid")
|
||
|
|
assert stats["count"] == 5
|
||
|
|
|
||
|
|
def test_feature_stats_empty(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
stats = fit.get_feature_stats("nonexistent")
|
||
|
|
assert stats == {}
|
||
|
|
|
||
|
|
def test_multiple_regimes(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
fit.record_decision(_state(), _action(), "volatile")
|
||
|
|
imp_normal = fit.get_importance(regime="normal")
|
||
|
|
imp_volatile = fit.get_importance(regime="volatile")
|
||
|
|
assert len(imp_normal) > 0
|
||
|
|
assert len(imp_volatile) > 0
|
||
|
|
|
||
|
|
def test_importance_importance_positive(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(10):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance()
|
||
|
|
for i in imp:
|
||
|
|
assert i.importance >= 0
|
||
|
|
|
||
|
|
def test_importance_sample_count(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
for _ in range(5):
|
||
|
|
fit.record_decision(_state(), _action(), "normal")
|
||
|
|
imp = fit.get_importance(top_n=1)
|
||
|
|
assert imp[0].sample_count > 0
|
||
|
|
|
||
|
|
def test_importance_regime_field(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
fit.record_decision(_state(), _action(), "choppy")
|
||
|
|
imp = fit.get_importance(regime="choppy")
|
||
|
|
assert imp[0].regime == "choppy"
|
||
|
|
|
||
|
|
def test_feature_count_increases(self):
|
||
|
|
fit = FeatureImportanceTracker()
|
||
|
|
assert fit.feature_count == 0
|
||
|
|
fit.record_decision(_state(), _action())
|
||
|
|
assert fit.feature_count > 0
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# POLICY ROLLBACK (15+ tests)
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestPolicyRollbackComprehensive:
|
||
|
|
def test_initial_state(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
rb = PolicyRollback(registry=reg)
|
||
|
|
assert rb.shadow_score_count == 0
|
||
|
|
assert len(rb.rollback_events) == 0
|
||
|
|
|
||
|
|
def test_record_many_scores(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
rb = PolicyRollback(registry=reg)
|
||
|
|
for _ in range(20):
|
||
|
|
rb.record_shadow_score(10.0)
|
||
|
|
assert rb.shadow_score_count == 20
|
||
|
|
|
||
|
|
def test_no_rollback_insufficient_data(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=10)
|
||
|
|
for _ in range(9):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
assert rb.check_rollback() is None
|
||
|
|
|
||
|
|
def test_no_rollback_no_active_policy(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
assert rb.check_rollback() is None
|
||
|
|
|
||
|
|
def test_no_rollback_performance_ok(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(10.0)
|
||
|
|
assert rb.check_rollback() is None
|
||
|
|
|
||
|
|
def test_rollback_on_degradation(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-5.0)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
event = rb.check_rollback()
|
||
|
|
assert event is not None
|
||
|
|
assert event.rolled_back_version == "v1"
|
||
|
|
|
||
|
|
def test_rollback_event_fields(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-5.0)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
event = rb.check_rollback()
|
||
|
|
assert event.ts_ns > 0
|
||
|
|
assert event.rolled_back_version == "v1"
|
||
|
|
assert event.reason.startswith("performance_drop")
|
||
|
|
|
||
|
|
def test_multiple_rollbacks(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-5.0)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
rb.check_rollback()
|
||
|
|
assert len(rb.rollback_events) == 1
|
||
|
|
|
||
|
|
def test_custom_threshold(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
reg.register_candidate(_baseline(version="v1"), score=10.0)
|
||
|
|
reg.promote("v1", PolicyStage.ACTIVE)
|
||
|
|
rb = PolicyRollback(registry=reg, min_shadow_steps=3, degradation_threshold=-100.0)
|
||
|
|
for _ in range(5):
|
||
|
|
rb.record_shadow_score(-10.0)
|
||
|
|
assert rb.check_rollback() is None # -10 > -100
|
||
|
|
|
||
|
|
def test_shadow_score_returns_list(self):
|
||
|
|
reg = PolicyRegistry()
|
||
|
|
rb = PolicyRollback(registry=reg)
|
||
|
|
rb.record_shadow_score(5.0)
|
||
|
|
rb.record_shadow_score(10.0)
|
||
|
|
assert rb.shadow_score_count == 2
|
||
|
|
|
||
|
|
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
# STRESS SCENARIOS (15+ tests)
|
||
|
|
# ══════════════════════════════════════════════════════════════════════════════
|
||
|
|
|
||
|
|
class TestStressComprehensive:
|
||
|
|
def test_all_scenarios_unique_id(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
suite = factory.build_stress_suite()
|
||
|
|
ids = [s.scenario_id for s in suite]
|
||
|
|
assert len(ids) == len(set(ids))
|
||
|
|
|
||
|
|
def test_all_scenarios_have_valid_state(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
for sc in factory.build_stress_suite():
|
||
|
|
assert sc.initial_state.book.best_bid > 0
|
||
|
|
assert sc.initial_state.book.best_ask > 0
|
||
|
|
assert sc.initial_state.account.equity > 0
|
||
|
|
|
||
|
|
def test_all_scenarios_have_positive_max_steps(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
for sc in factory.build_stress_suite():
|
||
|
|
assert sc.max_steps > 0
|
||
|
|
|
||
|
|
def test_flash_crash_thin_book(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.flash_crash()
|
||
|
|
assert sc.initial_state.book.bids[0].qty == 0.1
|
||
|
|
|
||
|
|
def test_liquidity_vacuum_extremely_thin(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.liquidity_vacuum()
|
||
|
|
assert sc.initial_state.book.bids[0].qty == 0.001
|
||
|
|
|
||
|
|
def test_extreme_vol_wide_spread(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.extreme_volatility()
|
||
|
|
spread = sc.initial_state.book.best_ask - sc.initial_state.book.best_bid
|
||
|
|
assert spread == 2000.0
|
||
|
|
|
||
|
|
def test_toxic_flood_multiple_counterparties(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.toxic_flood()
|
||
|
|
assert len(sc.counterparties) == 3
|
||
|
|
|
||
|
|
def test_choppy_tight_range(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.choppy_market()
|
||
|
|
spread = sc.initial_state.book.best_ask - sc.initial_state.book.best_bid
|
||
|
|
assert spread == 0.5
|
||
|
|
|
||
|
|
def test_weekend_low_participation(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.weekend_low_participation()
|
||
|
|
assert sc.initial_state.book.bids[0].qty == 0.2
|
||
|
|
|
||
|
|
def test_liquidation_cascade_tags(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.liquidation_cascade()
|
||
|
|
assert "liquidation_cascade" in sc.tags
|
||
|
|
assert "cascade" in sc.tags
|
||
|
|
|
||
|
|
def test_correlation_breakdown_tags(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
sc = factory.correlation_breakdown()
|
||
|
|
assert "correlation_breakdown" in sc.tags
|
||
|
|
|
||
|
|
def test_suite_count_per_symbol(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
suite = factory.build_stress_suite(symbols=("BTCUSDT",))
|
||
|
|
assert len(suite) == 8
|
||
|
|
|
||
|
|
def test_suite_multi_symbol(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
suite = factory.build_stress_suite(symbols=("BTCUSDT", "ETHUSDT", "SOLUSDT"))
|
||
|
|
assert len(suite) == 24
|
||
|
|
|
||
|
|
def test_all_tags_are_strings(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
for sc in factory.build_stress_suite():
|
||
|
|
for tag in sc.tags:
|
||
|
|
assert isinstance(tag, str)
|
||
|
|
|
||
|
|
def test_all_descriptions_non_empty(self):
|
||
|
|
factory = StressScenarioFactory()
|
||
|
|
for sc in factory.build_stress_suite():
|
||
|
|
assert len(sc.description) > 10
|
||
|
|
|
||
|
|
def test_custom_counterparties(self):
|
||
|
|
from malkhut.counterparties import ToxicTakerPolicy
|
||
|
|
factory = StressScenarioFactory(counterparties=(ToxicTakerPolicy(),))
|
||
|
|
sc = factory.flash_crash()
|
||
|
|
assert len(sc.counterparties) == 1
|