""" Tests for tie-in fix, cognition pipeline, and regime expansion. """ import json import os import tempfile import time import pytest 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: PerformanceMatrix wired to evaluator # ══════════════════════════════════════════════════════════════════════════════ class TestTieInFix: def test_matrix_records_regime_performance(self): """PerformanceMatrix should record strategy × regime scores.""" matrix = PerformanceMatrix() matrix.record("strat_1", "flash_crash", score=10.0, pnl_bps=5.0) matrix.record("strat_1", "normal", score=15.0, pnl_bps=8.0) scores = matrix.get_scores_for_regime("flash_crash") assert len(scores) == 1 assert scores[0].strategy_id == "strat_1" def test_matrix_tracks_multiple_strategies(self): matrix = PerformanceMatrix() matrix.record("s1", "normal", score=10.0) matrix.record("s2", "normal", score=15.0) scores = matrix.get_scores_for_regime("normal") assert len(scores) == 2 assert scores[0].score == 15.0 # sorted descending def test_matrix_best_for_regime(self): matrix = PerformanceMatrix() matrix.record("s1", "flash_crash", score=5.0) matrix.record("s2", "flash_crash", score=10.0) best = matrix.get_best("flash_crash") assert best == "s2" def test_matrix_best_excludes_baseline(self): matrix = PerformanceMatrix() matrix.record("baseline", "normal", score=10.0) matrix.record("s1", "normal", score=8.0) best = matrix.get_best("normal", exclude={"baseline"}) assert best == "s1" def test_matrix_coverage(self): matrix = PerformanceMatrix() matrix.record("s1", "normal", score=10.0) matrix.record("s1", "flash_crash", score=5.0) matrix.record("s2", "normal", score=8.0) cov = matrix.get_coverage() assert cov["s1"] == 2 assert cov["s2"] == 1 # ══════════════════════════════════════════════════════════════════════════════ # COGNITION PIPELINE # ══════════════════════════════════════════════════════════════════════════════ 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 # simulate 2 seconds ago rl._tokens = 0 # force empty assert rl.acquire() # should refill and acquire 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 Source", "https://example.com", "news", 0.8) assert cat.source_count == 1 assert len(cat.get_enabled()) == 1 finally: os.unlink(path) def test_record_fetch(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) cat.record_fetch("s1", success=False) sources = cat.get_enabled() assert sources[0].fetch_count == 2 assert sources[0].error_count == 1 finally: os.unlink(path) class TestRegimeExtractor: def test_extract_flash_crash(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Bitcoin crashed 20% in 5 minutes") assert "flash_crash" in regimes def test_extract_liquidity_vacuum(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Liquidity vacuum observed on exchange") assert "liquidity_vacuum" in regimes def test_extract_normal(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Normal market conditions") assert "normal" in regimes def test_extract_multiple(self): ext = RegimeExtractor() regimes = ext.extract_regimes("Flash crash with high volatility and liquidation cascade") assert "flash_crash" in regimes # "high_volatility" might not match exactly — check for volatility-related assert any("liquid" in r or "crash" in r for r in regimes) def test_sentiment_positive(self): ext = RegimeExtractor() score = ext.extract_sentiment("Bitcoin rally surges to new highs with gains") assert score > 0 def test_sentiment_negative(self): ext = RegimeExtractor() score = ext.extract_sentiment("Flash crash panic liquidation cascade") assert score < 0 class TestCognitionPipeline: def test_add_source(self): pipe = CognitionPipeline(catalogue_path="/dev/null") pipe.add_source("s1", "Test", "https://example.com", "news", 0.8) 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_regimes = pipe.fetch_and_extract("s1", "Flash crash observed in BTC") assert "flash_crash" in new_regimes 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 # already discovered 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 # ══════════════════════════════════════════════════════════════════════════════ # REGIME EXPANSION # ══════════════════════════════════════════════════════════════════════════════ class TestRegimeExpansion: def test_generate_regimes(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=50) assert len(regimes) == 50 assert exp.generated_count == 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 # many distinct labels 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) assert scenario.max_steps == 20 def test_regime_dimensions_represented(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=100) liq_labels = set(r.liquidity.label for r in regimes) vol_labels = set(r.volatility.label for r in regimes) flow_labels = set(r.flow.label for r in regimes) assert len(liq_labels) == 4 # vacuum, thin, normal, deep assert len(vol_labels) == 4 # tight, normal, wide, extreme assert len(flow_labels) == 4 # balanced, buy_pressure, sell_pressure, toxic def test_orthogonal_to_cognition(self): """Regime expansion is orthogonal to cognition pipeline.""" from malkhut.training.cognition import CognitionPipeline pipe = CognitionPipeline(catalogue_path="/dev/null") exp = RegimeExpander() # Cognition discovers regimes from news pipe.add_source("s1", "Test", "https://example.com") pipe.fetch_and_extract("s1", "Flash crash and high volatility") cog_regimes = set(pipe.get_discovered_regimes()) # Expansion generates regimes from dimensions exp_regimes = exp.generate_regimes(max_regimes=50) exp_labels = set(r.label for r in exp_regimes) # They are different (orthogonal) — no overlap required assert len(cog_regimes) > 0 assert len(exp_labels) > 20 def test_max_regimes_capped(self): exp = RegimeExpander() regimes = exp.generate_regimes(max_regimes=200) assert len(regimes) <= 200