""" Phase 0 Extensive Tests — tie-in, cognition pipeline, regime expansion. Covers: 1. Tie-in: PerformanceMatrix recording during evaluation 2. RateLimiter: token bucket mechanics 3. SourceCatalogue: add, fetch, error tracking 4. RegimeExtractor: keyword matching, sentiment 5. CognitionPipeline: full pipeline flow 6. RegimeExpander: dimension combinations 7. Integration: cognition + expansion + matrix + selector """ import json import os import tempfile import time import pytest from malkhut.state import FulfilmentPolicyParams from malkhut.training.cognition import ( RateLimiter, SourceCatalogue, RegimeExtractor, CognitionPipeline, ) from malkhut.training.regime_expansion import ( RegimeExpander, ExpandedRegime, LIQUIDITY_DIMS, VOLATILITY_DIMS, FLOW_DIMS, STRUCTURE_DIMS, ) from malkhut.training.selector import PerformanceMatrix, MarketRegime # ══════════════════════════════════════════════════════════════════════════════ # TIE-IN FIX (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestTieInFix: def test_record_single_regime(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) assert len(m.get_scores_for_regime("normal")) == 1 def test_record_multiple_strategies(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s2", "normal", score=15.0) scores = m.get_scores_for_regime("normal") assert len(scores) == 2 assert scores[0].score == 15.0 def test_record_multiple_regimes(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s1", "flash_crash", score=5.0) assert len(m.get_scores_for_regime("normal")) == 1 assert len(m.get_scores_for_regime("flash_crash")) == 1 def test_best_for_regime(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s2", "normal", score=15.0) assert m.get_best("normal") == "s2" def test_best_excludes_baseline(self): m = PerformanceMatrix() m.record("baseline", "normal", score=10.0) m.record("s1", "normal", score=8.0) assert m.get_best("normal", exclude={"baseline"}) == "s1" def test_best_returns_none_when_empty(self): m = PerformanceMatrix() assert m.get_best("nonexistent") is None def test_coverage_tracking(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s1", "flash_crash", score=5.0) m.record("s2", "normal", score=8.0) cov = m.get_coverage() assert cov["s1"] == 2 assert cov["s2"] == 1 def test_ema_update(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s1", "normal", score=20.0) scores = m.get_scores_for_regime("normal") assert scores[0].episodes == 2 # EMA: 0.3 * 20 + 0.7 * 10 = 13 assert scores[0].score == pytest.approx(13.0, abs=0.1) def test_total_entries(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s2", "flash_crash", score=5.0) assert m.total_entries == 2 def test_regime_list(self): m = PerformanceMatrix() m.record("s1", "normal", score=10.0) m.record("s1", "flash_crash", score=5.0) regimes = m.get_regimes_for_strategy("s1") assert len(regimes) == 2 # ══════════════════════════════════════════════════════════════════════════════ # RATE LIMITER (8 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestRateLimiter: def test_acquire_within_burst(self): rl = RateLimiter(requests_per_minute=60, burst_size=5) for _ in range(5): assert rl.acquire() def test_acquire_exceeds_burst(self): rl = RateLimiter(requests_per_minute=60, burst_size=2) assert rl.acquire() assert rl.acquire() assert not rl.acquire() def test_token_refill(self): rl = RateLimiter(requests_per_minute=60, burst_size=1) rl.acquire() rl._last_refill = time.time() - 2 rl._tokens = 0 assert rl.acquire() def test_wait_succeeds(self): rl = RateLimiter(requests_per_minute=60, burst_size=1) rl.acquire() rl._tokens = 0 rl._last_refill = time.time() - 2 assert rl.wait(timeout_s=5) def test_burst_size_respected(self): rl = RateLimiter(requests_per_minute=60, burst_size=3) for _ in range(3): assert rl.acquire() assert not rl.acquire() def test_rpm_affects_refill(self): rl = RateLimiter(requests_per_minute=120, burst_size=1) rl.acquire() rl._tokens = 0 rl._last_refill = time.time() - 1 # 120 RPM = 2/sec, so 1 second should refill 2 tokens assert rl.acquire() def test_multiple_instances(self): rl1 = RateLimiter(requests_per_minute=30, burst_size=2) rl2 = RateLimiter(requests_per_minute=60, burst_size=5) assert rl1.acquire() assert rl2.acquire() def test_concurrent_acquire(self): rl = RateLimiter(requests_per_minute=60, burst_size=10) results = [rl.acquire() for _ in range(15)] assert sum(results) == 10 # ══════════════════════════════════════════════════════════════════════════════ # SOURCE CATALOGUE (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestSourceCatalogue: def test_add_source(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat = SourceCatalogue(path) cat.add_source("s1", "Test", "https://example.com") assert cat.source_count == 1 finally: os.unlink(path) def test_record_fetch_success(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat = SourceCatalogue(path) cat.add_source("s1", "Test", "https://example.com") cat.record_fetch("s1", success=True) assert cat.get_enabled()[0].fetch_count == 1 assert cat.get_enabled()[0].error_count == 0 finally: os.unlink(path) def test_record_fetch_error(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat = SourceCatalogue(path) cat.add_source("s1", "Test", "https://example.com") cat.record_fetch("s1", success=False) assert cat.get_enabled()[0].error_count == 1 finally: os.unlink(path) def test_get_by_type(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat = SourceCatalogue(path) cat.add_source("s1", "News", "https://news.com", "news") cat.add_source("s2", "Data", "https://data.com", "data") assert len(cat.get_by_type("news")) == 1 assert len(cat.get_by_type("data")) == 1 finally: os.unlink(path) def test_persistence(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat1 = SourceCatalogue(path) cat1.add_source("s1", "Test", "https://example.com") cat1.record_fetch("s1", success=True) # Reload cat2 = SourceCatalogue(path) assert cat2.source_count == 1 assert cat2.get_enabled()[0].fetch_count == 1 finally: os.unlink(path) def test_source_count(self): with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: cat = SourceCatalogue(path) cat.add_source("s1", "A", "https://a.com") cat.add_source("s2", "B", "https://b.com") assert cat.source_count == 2 finally: os.unlink(path) # ══════════════════════════════════════════════════════════════════════════════ # REGIME EXTRACTOR (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestRegimeExtractor: def test_flash_crash(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Bitcoin crashed 20% in 5 minutes") assert "flash_crash" in regimes def test_liquidity_vacuum(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Liquidity vacuum on exchange") assert "liquidity_vacuum" in regimes def test_normal(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Normal market conditions") assert "normal" in regimes def test_multiple_regimes(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Flash crash with liquidation cascade") assert "flash_crash" in regimes assert "liquidation" in regimes def test_sentiment_positive(self): ext = RegimeExtractor() score = ext.extract_sentiment("Bitcoin rally surges to new highs") assert score > 0 def test_sentiment_negative(self): ext = RegimeExtractor() score = ext.extract_sentiment("Flash crash panic liquidation") assert score < 0 def test_sentiment_neutral(self): ext = RegimeExtractor() score = ext.extract_sentiment("Market trading normally") assert score == 0.0 def test_empty_text(self): ext = RegimeExtractor() regimes = ext.extract_regimes("") assert regimes == ["normal"] def test_whale_activity(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Whale order detected on exchange") assert "whale_activity" in regimes def test_stop_hunting(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Stop hunt observed in BTC") assert "stop_hunting" in regimes # ══════════════════════════════════════════════════════════════════════════════ # COGNITION PIPELINE (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestCognitionPipeline: def test_add_source(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") assert pipe._catalogue.source_count == 1 def test_fetch_and_extract(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") new = pipe.fetch_and_extract("s1", "Flash crash observed in BTC") assert "flash_crash" in new assert pipe._total_fetched == 1 def test_deduplication(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Flash crash observed") new = pipe.fetch_and_extract("s1", "Flash crash again") assert len(new) == 0 def test_discovered_regimes(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Flash crash observed") regimes = pipe.get_discovered_regimes() assert "flash_crash" in regimes def test_source_stats(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Test text") stats = pipe.get_source_stats() assert stats["total_sources"] == 1 assert stats["total_fetched"] == 1 def test_seed_default_sources(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.seed_default_sources() assert pipe._catalogue.source_count == 8 def test_rate_limiting(self): pipe = CognitionPipeline(catalogue_path="/dev/null", rate_limit_rpm=1) pipe.add_source("s1", "Test", "https://example.com") # First fetch should succeed pipe.fetch_and_extract("s1", "Test") # Second should be rate limited (returns empty) new = pipe.fetch_and_extract("s1", "Test") # May or may not be rate limited depending on timing assert isinstance(new, list) def test_multiple_sources(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "News", "https://news.com") pipe.add_source("s2", "Data", "https://data.com") pipe.fetch_and_extract("s1", "Flash crash") pipe.fetch_and_extract("s2", "High volatility") assert pipe._total_fetched == 2 def test_discovered_count(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Flash crash and liquidation") assert pipe.discovered_regime_count >= 2 def test_perm_run_capable(self): """Pipeline should not crash after many iterations.""" pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") for i in range(20): pipe.fetch_and_extract("s1", f"Flash crash iteration {i}") assert pipe._total_fetched >= 1 # at least some fetched # ══════════════════════════════════════════════════════════════════════════════ # REGIME EXPANSION (10 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestRegimeExpansion: def test_generate_regimes(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=50) assert len(regimes) == 50 def test_regime_ids_unique(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=100) ids = [r.regime_id for r in regimes] assert len(ids) == len(set(ids)) def test_regime_labels_diverse(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=50) labels = [r.label for r in regimes] assert len(set(labels)) > 20 def test_regime_to_scenario(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=5) from malkhut.training.cma_trainer import Scenario scenario = exp.regime_to_scenario(regimes[0]) assert isinstance(scenario, Scenario) def test_dimension_coverage(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=100) liq = set(r.liquidity.label for r in regimes) vol = set(r.volatility.label for r in regimes) flow = set(r.flow.label for r in regimes) assert len(liq) == 4 assert len(vol) == 4 assert len(flow) == 4 def test_orthogonal_to_cognition(self): from malkhut.training.cognition import CognitionPipeline pipe = CognitionPipeline(catalogue_path="/dev/null") exp = RegimeExpander() pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Flash crash") exp_regimes = exp.generate_regimes(max_regimes=50) assert len(exp_regimes) > 0 def test_max_regimes_capped(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=200) assert len(regimes) <= 200 def test_regime_counterparties(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=5) for r in regimes: assert len(r.structure.counterparties) > 0 def test_regime_book_state(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=5) for r in regimes: assert r.bid > 0 assert r.ask > 0 assert r.ask > r.bid # ══════════════════════════════════════════════════════════════════════════════ # INTEGRATION (5 tests) # ══════════════════════════════════════════════════════════════════════════════ class TestPhase0Integration: def test_cognition_to_matrix(self): """Cognition findings should be recordable to matrix.""" pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") regimes = pipe.fetch_and_extract("s1", "Flash crash and liquidation") matrix = PerformanceMatrix() for regime in regimes: matrix.record("test_strategy", regime, score=10.0) assert matrix.total_entries >= 1 def test_expansion_to_scenarios(self): """Regime expansion should produce evaluatable scenarios.""" exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=10) from malkhut.training.cma_trainer import Scenario for r in regimes: scenario = exp.regime_to_scenario(r) assert isinstance(scenario, Scenario) assert scenario.max_steps > 0 def test_full_pipeline_flow(self): """Cognition → expansion → matrix → selector.""" # 1. Cognition discovers regimes pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com") cog_regimes = pipe.fetch_and_extract("s1", "Flash crash") # 2. Expansion generates regimes exp = RegimeExpander() exp_regimes = exp.generate_regimes(max_regimes=10) # 3. Matrix records performance matrix = PerformanceMatrix() for r in cog_regimes: matrix.record("s1", r, score=10.0) for r in exp_regimes: matrix.record("s2", r.label, score=8.0) # 4. Selector queries matrix from malkhut.training.selector import StrategySelector selector = StrategySelector(matrix=matrix, min_episodes_for_selection=1) strategies = {"s1": _baseline(), "s2": _baseline(version="v2")} result = selector.select(_state(), strategies) # Should return a valid strategy (may be fallback if matrix has too few entries) assert result.strategy_id in strategies or result.strategy_id == "baseline" def test_rate_limiting_prevents_abuse(self): """Pipeline should respect rate limits.""" pipe = CognitionPipeline(catalogue_path="/dev/null", rate_limit_rpm=2) pipe.add_source("s1", "Test", "https://example.com") results = [] for _ in range(5): new = pipe.fetch_and_extract("s1", "Flash crash") results.append(len(new)) # Some should be rate limited (return empty) assert sum(results) < 5 def test_source_persistence(self): """Sources should persist across pipeline instances.""" with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: path = f.name try: pipe1 = CognitionPipeline(catalogue_path=path) pipe1.add_source("s1", "Test", "https://example.com") pipe1.fetch_and_extract("s1", "Flash crash") pipe2 = CognitionPipeline(catalogue_path=path) assert pipe2._catalogue.source_count == 1 finally: os.unlink(path) 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 _state(): from malkhut.state import AccountState, MarketWorldState, Mode, OrderBookState, PriceLevel, VenueRules return MarketWorldState( ts_ns=1, mode=Mode.REPLAY_NO_IMPACT, venue=_venue(), book=OrderBookState(ts_ns=1, symbol="BTCUSDT", bids=(PriceLevel(50000.0, 1.0),), asks=(PriceLevel(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), ) def _venue(): from malkhut.state import VenueRules 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)