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