malkhut(tests): 1140 test functions across 46 test files
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
This commit is contained in:
255
MALKHUT/malkhut/tests/test_pipeline.py
Normal file
255
MALKHUT/malkhut/tests/test_pipeline.py
Normal file
@@ -0,0 +1,255 @@
|
||||
"""
|
||||
Tests for training pipeline and logging.
|
||||
|
||||
Verifies:
|
||||
- Pipeline runs bounded iterations
|
||||
- Early stopping on convergence
|
||||
- Time budget respected
|
||||
- Eval budget respected
|
||||
- Logger produces compact JSONL
|
||||
- Events are complete and observable
|
||||
- Full pipeline flow: train → promote → reload
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
import pytest
|
||||
from malkhut.state import FulfilmentPolicyParams
|
||||
from malkhut.training.pipeline import (
|
||||
TrainingPipeline, PipelineConfig, PipelineResult,
|
||||
TrainingLogger, TrainingEvent,
|
||||
)
|
||||
from malkhut.training.registry import PolicyRegistry, PolicyStage
|
||||
|
||||
|
||||
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. TRAINING LOGGER
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
class TestTrainingLogger:
|
||||
def test_log_creates_file(self):
|
||||
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||
path = f.name
|
||||
try:
|
||||
logger = TrainingLogger(log_path=path)
|
||||
logger.log(TrainingEvent(
|
||||
timestamp_ns=time.time_ns(), event_type="test",
|
||||
policy_version="v1", score=10.0,
|
||||
))
|
||||
assert os.path.exists(path)
|
||||
with open(path) as f:
|
||||
lines = f.readlines()
|
||||
assert len(lines) == 1
|
||||
record = json.loads(lines[0])
|
||||
assert record["type"] == "test"
|
||||
assert record["ver"] == "v1"
|
||||
assert record["score"] == 10.0
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
def test_log_run_start(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_run_start(generation=0, budget_evals=14)
|
||||
assert logger.event_count == 1
|
||||
assert logger.get_events()[0].event_type == "run_start"
|
||||
|
||||
def test_log_generation(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_generation(generation=1, best_score=5.0, mean_score=3.0, evals=14, improvement=2.0)
|
||||
assert logger.event_count == 1
|
||||
e = logger.get_events()[0]
|
||||
assert e.score == 5.0
|
||||
assert e.details["mean"] == 3.0
|
||||
|
||||
def test_log_candidate(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_candidate("v1", 10.0, 1)
|
||||
assert logger.get_events()[0].policy_version == "v1"
|
||||
|
||||
def test_log_promote(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_promote("v1", "CANDIDATE", "ACTIVE", "promoted")
|
||||
e = logger.get_events()[0]
|
||||
assert e.event_type == "promote"
|
||||
assert e.details["from"] == "CANDIDATE"
|
||||
assert e.details["to"] == "ACTIVE"
|
||||
|
||||
def test_log_reject(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_reject("v1", "tail risk")
|
||||
assert logger.get_events()[0].event_type == "reject"
|
||||
|
||||
def test_log_reload(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_reload("v2", "v1")
|
||||
e = logger.get_events()[0]
|
||||
assert e.details["old"] == "v1"
|
||||
|
||||
def test_log_run_end(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log_run_end(generation=5, total_evals=70, best_score=8.0, duration_s=120.0)
|
||||
e = logger.get_events()[0]
|
||||
assert e.details["duration_s"] == 120.0
|
||||
|
||||
def test_filter_by_event_type(self):
|
||||
logger = TrainingLogger(log_path="/dev/null")
|
||||
logger.log(TrainingEvent(timestamp_ns=1, event_type="start"))
|
||||
logger.log(TrainingEvent(timestamp_ns=2, event_type="generation"))
|
||||
logger.log(TrainingEvent(timestamp_ns=3, event_type="start"))
|
||||
starts = logger.get_events(event_type="start")
|
||||
assert len(starts) == 2
|
||||
|
||||
def test_jsonl_compact(self):
|
||||
with tempfile.NamedTemporaryFile(suffix=".log", delete=False) as f:
|
||||
path = f.name
|
||||
try:
|
||||
logger = TrainingLogger(log_path=path)
|
||||
logger.log(TrainingEvent(
|
||||
timestamp_ns=1234567890, event_type="test",
|
||||
policy_version="v1", score=10.0, generation=1, evals=14,
|
||||
details={"mean": 5.0},
|
||||
))
|
||||
with open(path) as f:
|
||||
line = f.readline()
|
||||
# Compact: no spaces after separators
|
||||
assert ",\"score\":10.0" in line
|
||||
assert "\"mean\":5.0" in line
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 2. PIPELINE CONFIG
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
class TestPipelineConfig:
|
||||
def test_default_config(self):
|
||||
cfg = PipelineConfig()
|
||||
assert cfg.max_generations == 10
|
||||
assert cfg.max_evals_per_generation == 14
|
||||
assert cfg.max_time_s == 300.0
|
||||
assert cfg.patience == 3
|
||||
|
||||
def test_custom_config(self):
|
||||
cfg = PipelineConfig(max_generations=5, patience=2)
|
||||
assert cfg.max_generations == 5
|
||||
assert cfg.patience == 2
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 3. PIPELINE RUN
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
class TestTrainingPipeline:
|
||||
def test_pipeline_returns_result(self):
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert isinstance(result, PipelineResult)
|
||||
assert result.generations_run >= 1
|
||||
assert result.total_evals > 0
|
||||
assert result.duration_s > 0
|
||||
|
||||
def test_pipeline_respects_max_generations(self):
|
||||
config = PipelineConfig(max_generations=2, max_evals_per_generation=3, max_time_s=60)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert result.generations_run <= 2
|
||||
|
||||
def test_pipeline_respects_time_budget(self):
|
||||
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, max_time_s=2.0)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
# Should stop within reasonable time
|
||||
assert result.generations_run <= 10 # not all 100
|
||||
assert result.duration_s < 60.0
|
||||
|
||||
def test_pipeline_early_stopping(self):
|
||||
config = PipelineConfig(max_generations=100, max_evals_per_generation=3, patience=1, max_time_s=60)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert result.generations_run <= 100
|
||||
|
||||
def test_pipeline_logs_events(self):
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
events = pipeline.logger.get_events()
|
||||
assert len(events) > 0
|
||||
event_types = [e.event_type for e in events]
|
||||
assert "run_start" in event_types
|
||||
assert "run_end" in event_types
|
||||
assert "generation" in event_types
|
||||
|
||||
def test_pipeline_has_best_score(self):
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert isinstance(result.best_score, float)
|
||||
|
||||
def test_pipeline_registry_has_records(self):
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30, auto_promote=True)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert pipeline.registry.record_count > 0
|
||||
|
||||
def test_pipeline_result_events_match_logger(self):
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
assert len(result.events) == pipeline.logger.event_count
|
||||
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
# 4. FULL OBSERVABLE FLOW
|
||||
# ══════════════════════════════════════════════════════════════════════════════
|
||||
|
||||
class TestFullObservableFlow:
|
||||
def test_train_log_promote_activate(self):
|
||||
"""Full observable flow: train → log → promote → activate → engine loads."""
|
||||
from malkhut.engine import FulfilmentEngine
|
||||
|
||||
config = PipelineConfig(max_generations=1, max_evals_per_generation=7, max_time_s=30)
|
||||
pipeline = TrainingPipeline(config=config, log_path="/dev/null")
|
||||
|
||||
# Engine starts with default
|
||||
engine = FulfilmentEngine(registry=pipeline.registry)
|
||||
|
||||
# Run pipeline
|
||||
result = pipeline.run(incumbent=_baseline(version="init"), symbols=("BTCUSDT",))
|
||||
|
||||
# Verify logging happened
|
||||
assert len(result.events) > 0
|
||||
|
||||
# Verify registry has promoted policies
|
||||
active = pipeline.registry.load_active()
|
||||
if active:
|
||||
# Hot reload engine
|
||||
engine.hot_reload_policy(active)
|
||||
assert engine.params_provider().version == active.version
|
||||
|
||||
engine.close()
|
||||
Reference in New Issue
Block a user