556 lines
23 KiB
Python
556 lines
23 KiB
Python
|
|
"""
|
||
|
|
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)
|