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

222 lines
9.6 KiB
Python
Raw Normal View History

"""
Tests for policy registry and the learning/improvement loop.
Verifies that:
- Trained policies are persisted to CH
- Registry manages lifecycle (CANDIDATE → ACTIVE)
- Engine loads active policy from registry
- Hot-reload via control plane works
- Full training → registry → engine flow works
"""
import pytest
from malkhut.state import FulfilmentPolicyParams
from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord
from malkhut.training.cma_trainer import (
CMAParameterCodec, PolicySnapshot, SelfPlayPool,
)
from malkhut.storage.ch_store import MalkhutCHStore
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)
# ══════════════════════════════════════════════════════════════════════════════
# 1. POLICY REGISTRY
# ══════════════════════════════════════════════════════════════════════════════
class TestPolicyRegistry:
def test_register_candidate(self):
reg = PolicyRegistry()
record = reg.register_candidate(_baseline(version="v1"), score=10.0)
assert record.stage == PolicyStage.CANDIDATE
assert record.version == "v1"
assert record.score == 10.0
def test_promote_through_stages(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.BACKTESTED, "tests passed")
reg.promote("v1", PolicyStage.SELF_PLAY_CONFIRMED, "pool hardened")
reg.promote("v1", PolicyStage.SHADOW, "shadow verified")
reg.promote("v1", PolicyStage.ACTIVE, "promoted")
record = reg.get_record("v1")
assert record.stage == PolicyStage.ACTIVE
def test_load_active(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.ACTIVE)
params = reg.load_active()
assert params is not None
assert params.version == "v1"
def test_no_active_returns_none(self):
reg = PolicyRegistry()
assert reg.load_active() is None
def test_reject(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.reject("v1", "tail risk too high")
record = reg.get_record("v1")
assert record.stage == PolicyStage.REJECTED
def test_retire(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.ACTIVE)
reg.retire("v1", "replaced by v2")
record = reg.get_record("v1")
assert record.stage == PolicyStage.RETIRED
def test_only_one_active(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.ACTIVE)
reg.register_candidate(_baseline(version="v2"), score=12.0)
reg.promote("v2", PolicyStage.ACTIVE)
# v1 should be retired when v2 becomes active
# (or we just have two ACTIVE — the load_active returns latest)
params = reg.load_active()
assert params.version == "v2"
def test_get_by_stage(self):
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.register_candidate(_baseline(version="v2"), score=12.0)
reg.promote("v1", PolicyStage.ACTIVE)
candidates = reg.get_by_stage(PolicyStage.CANDIDATE)
assert len(candidates) == 1
assert candidates[0].version == "v2"
def test_record_count(self):
reg = PolicyRegistry()
assert reg.record_count == 0
reg.register_candidate(_baseline(version="v1"), score=10.0)
assert reg.record_count == 1
def test_unknown_version_raises(self):
reg = PolicyRegistry()
with pytest.raises(KeyError):
reg.promote("nonexistent", PolicyStage.ACTIVE)
# ══════════════════════════════════════════════════════════════════════════════
# 2. ENGINE-REGISTRY INTEGRATION
# ══════════════════════════════════════════════════════════════════════════════
class TestEngineRegistryIntegration:
def test_engine_loads_active_policy(self):
from malkhut.engine import FulfilmentEngine
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.ACTIVE)
engine = FulfilmentEngine(registry=reg)
params = engine.params_provider()
assert params.version == "v1"
engine.close()
def test_engine_hot_reload(self):
from malkhut.engine import FulfilmentEngine
engine = FulfilmentEngine()
engine.hot_reload_policy(_baseline(version="v2"))
params = engine.params_provider()
assert params.version == "v2"
engine.close()
def test_engine_hot_reload_via_registry(self):
from malkhut.engine import FulfilmentEngine
reg = PolicyRegistry()
reg.register_candidate(_baseline(version="v1"), score=10.0)
reg.promote("v1", PolicyStage.ACTIVE)
engine = FulfilmentEngine(registry=reg)
# Initially v1
assert engine.params_provider().version == "v1"
# Register v2 and promote
reg.register_candidate(_baseline(version="v2"), score=12.0)
reg.promote("v2", PolicyStage.ACTIVE)
# Hot reload to v2
engine.hot_reload_policy(reg.load_active())
assert engine.params_provider().version == "v2"
engine.close()
# ══════════════════════════════════════════════════════════════════════════════
# 3. FULL LEARNING LOOP
# ══════════════════════════════════════════════════════════════════════════════
class TestFullLearningLoop:
def test_train_register_promote_activate(self):
"""Full loop: train → register → promote → activate → engine loads."""
from malkhut.engine import FulfilmentEngine
from malkhut.training.cma_trainer import CMAESTrainer, PolicyEvaluator
from malkhut.counterparties import default_counterparty_ecology
from malkhut.cwm.core import MinimalCryptoLOBCWM
from malkhut.training.cma_trainer import ScenarioFactory
# Setup
codec = CMAParameterCodec()
evaluator = PolicyEvaluator(
cwm_factory=lambda: MinimalCryptoLOBCWM(),
counterparties=default_counterparty_ecology(),
)
pool = SelfPlayPool(max_size=5)
registry = PolicyRegistry()
trainer = CMAESTrainer(codec=codec, evaluator=evaluator, pool=pool)
# Register baseline as initial active
registry.register_candidate(_baseline(version="init"), score=0.0)
registry.promote("init", PolicyStage.ACTIVE)
# Engine starts with init policy
engine = FulfilmentEngine(registry=registry)
assert engine.params_provider().version == "init"
# Train
factory = ScenarioFactory()
scenarios = factory.build_suite(symbols=("BTCUSDT",), steps_per_scenario=3)
best = trainer.train(
incumbent=_baseline(version="init"),
scenarios=scenarios,
budget_evals=14,
seed=42,
)
# Register trained candidate
registry.register_candidate(best.params, score=best.score)
# Promote through pipeline
registry.promote(best.params.version, PolicyStage.BACKTESTED, "tests passed")
registry.promote(best.params.version, PolicyStage.SELF_PLAY_CONFIRMED, "pool hardened")
registry.promote(best.params.version, PolicyStage.ACTIVE, "promoted")
# Hot reload engine to new policy
engine.hot_reload_policy(registry.load_active())
assert engine.params_provider().version == best.params.version
engine.close()