CWM (103): core mechanics, exhaustive edge cases, numba, exchange mechanics Replay (118): exhaustive verification, microstructure, trajectory Training (190): asset classification, phase0 extensive, pipeline, exhaustive DSL (102): v2 syntax, expanded, new features ASEx (33): validate-before-mutate, single-writer Planner (48): MCTS, alternatives, hooks Counterparties (19): 9 adversarial agent policies Clock (30): event-driven reactor BingX (28): venue adapter IPC (8): Zinc SHM Storage (9): ClickHouse Risk (4): hard invariants State (17): frozen dataclass invariants Integration: E2E, concurrency, sync/async seams, hypothesis, fuzz, adversarial
255 lines
11 KiB
Python
255 lines
11 KiB
Python
"""
|
||
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
|