Files
sentiment-engine/MALKHUT/malkhut/tests/test_phase0_extensive.py

556 lines
23 KiB
Python
Raw Normal View History

"""
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)