Files
sentiment-engine/MALKHUT/malkhut/tests/test_pipeline.py
Codex 4c239f7774 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
2026-07-11 10:46:12 +02:00

256 lines
12 KiB
Python

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