malkhut(T4): Strategy DSL v2 + generator + supporting modules

Strategy DSL v2 (dsl.py): 40+ action primitives, 40+ market sensors,
12 comparison operators, 16 builtins, full parser.
Strategy Generator (generator.py): genetic programming evolution —
crossover, mutation, tournament selection, pool management.
Supporting: discrepancy tracking, execution quality, hooks, feature
importance, observability, parallel eval, auto-rollback, stress testing,
structured observations, trajectory recording.
This commit is contained in:
Codex
2026-07-11 10:28:38 +02:00
parent 4ffc8a601f
commit 863a4cc8c9
12 changed files with 2932 additions and 0 deletions

View File

@@ -0,0 +1,127 @@
"""
Live Discrepancy Tracker — compare CWM predictions vs actual market fills.
Enables:
- Detecting CWM model drift
- Alerting when predictions diverge from reality
- Feeding discrepancies back for CWM improvement
"""
from __future__ import annotations
import json
import time
from dataclasses import dataclass, field
from typing import Any, List, Mapping, Optional, Tuple
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction
from malkhut.cwm.core import MinimalCryptoLOBCWM
from malkhut.cwm.replay_verify import _compare_deep, ReplayMismatch
from malkhut.storage.ch_store import MalkhutCHStore
@dataclass(frozen=True, slots=True)
class DiscrepancyRecord:
"""One discrepancy between predicted and actual state."""
ts_ns: int
symbol: str
field: str
predicted: Any
actual: Any
severity: str # "info", "warning", "critical"
action_kind: str
policy_version: str
class DiscrepancyTracker:
"""
Track discrepancies between CWM predictions and actual market state.
Runs in shadow mode: CWM predicts next state, actual state arrives later,
discrepancy is logged and analyzed.
"""
def __init__(self, store: Optional[MalkhutCHStore] = None) -> None:
self._store = store
self._discrepancies: list[DiscrepancyRecord] = []
self._total_comparisons = 0
self._total_discrepancies = 0
def record_prediction(
self,
predicted_state: MarketWorldState,
action: FulfilmentAction,
policy_version: str,
) -> None:
"""Record a CWM prediction for later comparison."""
# Store for comparison when actual state arrives
self._last_prediction = predicted_state
self._last_action = action
self._last_policy_version = policy_version
def compare_with_actual(
self,
actual_state: MarketWorldState,
tolerances: Optional[Mapping[str, float]] = None,
) -> List[DiscrepancyRecord]:
"""
Compare last prediction with actual state.
Returns list of discrepancies found.
"""
if not hasattr(self, '_last_prediction') or self._last_prediction is None:
return []
self._total_comparisons += 1
mismatches = _compare_deep(
0, self._last_prediction, actual_state, tolerances,
)
discrepancies = []
for m in mismatches:
disc = DiscrepancyRecord(
ts_ns=actual_state.ts_ns,
symbol=actual_state.venue.symbol,
field=m.field,
predicted=m.expected,
actual=m.actual,
severity=m.severity,
action_kind=self._last_action.kind.value if self._last_action else "unknown",
policy_version=self._last_policy_version,
)
discrepancies.append(disc)
self._discrepancies.append(disc)
self._total_discrepancies += 1
# Persist to CH
if self._store:
self._store.store_discrepancy(
ts_ns=actual_state.ts_ns,
exchange=actual_state.venue.exchange,
symbol=actual_state.venue.symbol,
predicted=str(m.expected),
actual=str(m.actual),
severity=m.severity,
)
return discrepancies
@property
def discrepancy_rate(self) -> float:
if self._total_comparisons == 0:
return 0.0
return self._total_discrepancies / self._total_comparisons
@property
def total_comparisons(self) -> int:
return self._total_comparisons
@property
def total_discrepancies(self) -> int:
return self._total_discrepancies
def get_recent(self, n: int = 10) -> List[DiscrepancyRecord]:
return self._discrepancies[-n:]
def get_by_severity(self, severity: str) -> List[DiscrepancyRecord]:
return [d for d in self._discrepancies if d.severity == severity]

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,192 @@
"""
Execution Quality Metrics — measure the quality of execution.
Metrics:
- Slippage vs arrival price
- Implementation shortfall
- Market impact
- Fill rate vs expected
- Maker fill ratio
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import List, Optional
@dataclass(frozen=True, slots=True)
class ExecutionQualityReport:
"""Report on execution quality for a set of fills."""
total_fills: int
avg_slippage_bps: float
avg_implementation_shortfall_bps: float
avg_market_impact_bps: float
fill_rate: float # actual fills / expected fills
maker_fill_ratio: float
taker_fill_ratio: float
adverse_fill_ratio: float
avg_fill_time_ms: float
total_fees_bps: float
class ExecutionQualityTracker:
"""
Track execution quality metrics.
Measures slippage, implementation shortfall, market impact,
fill rates, and fee quality.
"""
def __init__(self) -> None:
self._fills: list[dict] = []
def record_fill(
self,
fill_price: float,
arrival_price: float,
expected_price: float,
is_maker: bool,
fill_time_ms: float,
fee_bps: float,
toxicity: float = 0.0,
) -> None:
"""Record a fill for quality analysis."""
slippage = (fill_price - arrival_price) / max(arrival_price, 1e-12) * 10_000
shortfall = (fill_price - expected_price) / max(expected_price, 1e-12) * 10_000
impact = abs(fill_price - arrival_price) / max(arrival_price, 1e-12) * 10_000
self._fills.append({
"fill_price": fill_price,
"arrival_price": arrival_price,
"expected_price": expected_price,
"slippage_bps": slippage,
"shortfall_bps": shortfall,
"impact_bps": impact,
"is_maker": is_maker,
"fill_time_ms": fill_time_ms,
"fee_bps": fee_bps,
"toxicity": toxicity,
})
def report(self) -> ExecutionQualityReport:
"""Generate execution quality report."""
if not self._fills:
return ExecutionQualityReport(
total_fills=0, avg_slippage_bps=0.0, avg_implementation_shortfall_bps=0.0,
avg_market_impact_bps=0.0, fill_rate=0.0, maker_fill_ratio=0.0,
taker_fill_ratio=0.0, adverse_fill_ratio=0.0, avg_fill_time_ms=0.0,
total_fees_bps=0.0,
)
n = len(self._fills)
avg_slip = sum(f["slippage_bps"] for f in self._fills) / n
avg_shortfall = sum(f["shortfall_bps"] for f in self._fills) / n
avg_impact = sum(f["impact_bps"] for f in self._fills) / n
avg_time = sum(f["fill_time_ms"] for f in self._fills) / n
avg_fee = sum(f["fee_bps"] for f in self._fills) / n
maker_count = sum(1 for f in self._fills if f["is_maker"])
toxic_count = sum(1 for f in self._fills if f["toxicity"] > 0.5)
return ExecutionQualityReport(
total_fills=n,
avg_slippage_bps=avg_slip,
avg_implementation_shortfall_bps=avg_shortfall,
avg_market_impact_bps=avg_impact,
fill_rate=1.0, # placeholder
maker_fill_ratio=maker_count / n,
taker_fill_ratio=1.0 - maker_count / n,
adverse_fill_ratio=toxic_count / n,
avg_fill_time_ms=avg_time,
total_fees_bps=avg_fee,
)
@property
def total_fills(self) -> int:
return len(self._fills)
class RiskAdjustedReturns:
"""
Compute risk-adjusted return metrics.
Metrics:
- Sharpe ratio
- Sortino ratio
- Calmar ratio
- Profit factor
- Max drawdown
"""
def __init__(self, risk_free_rate: float = 0.0) -> None:
self._risk_free_rate = risk_free_rate
self._returns: list[float] = []
def add_return(self, ret: float) -> None:
self._returns.append(ret)
@property
def sharpe_ratio(self) -> float:
if len(self._returns) < 2:
return 0.0
mean = sum(self._returns) / len(self._returns)
variance = sum((r - mean) ** 2 for r in self._returns) / (len(self._returns) - 1)
std = math.sqrt(variance) if variance > 0 else 1e-12
return (mean - self._risk_free_rate) / std
@property
def sortino_ratio(self) -> float:
if len(self._returns) < 2:
return 0.0
mean = sum(self._returns) / len(self._returns)
downside = [r for r in self._returns if r < 0]
if not downside:
return float('inf') if mean > 0 else 0.0
downside_var = sum(r ** 2 for r in downside) / len(downside)
downside_std = math.sqrt(downside_var) if downside_var > 0 else 1e-12
return (mean - self._risk_free_rate) / downside_std
@property
def profit_factor(self) -> float:
gains = sum(r for r in self._returns if r > 0)
losses = abs(sum(r for r in self._returns if r < 0))
if losses <= 0:
return float('inf') if gains > 0 else 0.0
return gains / losses
@property
def max_drawdown(self) -> float:
if not self._returns:
return 0.0
peak = self._returns[0]
max_dd = 0.0
cumulative = 0.0
for r in self._returns:
cumulative += r
if cumulative > peak:
peak = cumulative
dd = peak - cumulative
if dd > max_dd:
max_dd = dd
return max_dd
@property
def calmar_ratio(self) -> float:
if not self._returns:
return 0.0
mean = sum(self._returns) / len(self._returns)
annual_return = mean * 252 # annualize
dd = self.max_drawdown
if dd <= 0:
return float('inf') if annual_return > 0 else 0.0
return annual_return / dd
def report(self) -> dict:
return {
"sharpe_ratio": self.sharpe_ratio,
"sortino_ratio": self.sortino_ratio,
"profit_factor": self.profit_factor,
"max_drawdown": self.max_drawdown,
"calmar_ratio": self.calmar_ratio,
"total_returns": len(self._returns),
}

View File

@@ -0,0 +1,511 @@
"""
Strategy Generator — genetic programming for strategy evolution.
Based on genetic programming (GP) principles:
- Strategy = genome (parameter vector + strategy type)
- Crossover: combine two strategies to create offspring
- Mutation: randomly modify a strategy
- Selection: tournament selection based on self-play fitness
- Population: diverse pool of strategies, hardcoded baseline always available
Key design principle:
The hardcoded baseline is NEVER replaced. Generated strategies are ADDED
to the pool. The system GROWS its strategy repertoire.
Game theory insight:
During self-play, the system might "spontaneously generate" new strategies
via crossover/mutation of existing ones. These emergent strategies should
be captured, trialed, and added if successful.
"""
from __future__ import annotations
import math
import random
import time
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Callable, List, Optional, Sequence, Tuple
from malkhut.state import FulfilmentPolicyParams
from malkhut.training.cma_trainer import (
CMAParameterCodec, EpisodeResult, PolicyEvaluator,
PolicySnapshot, Scenario, SelfPlayPool,
)
from malkhut.training.registry import PolicyRegistry, PolicyStage
from malkhut.counterparties import default_counterparty_ecology
# ==============================================================================
# Strategy Types — different planner algorithms
# ==============================================================================
class StrategyType(str, Enum):
"""Different strategy structures the system can use."""
SM_MCTS = "SM_MCTS" # Decoupled UCB/UCT (default)
UCB1 = "UCB1" # Standard UCB1 (simpler)
THOMPSON_SAMPLING = "THOMPSON" # Thompson sampling
GREEDY = "GREEDY" # Always pick best Q-value
RANDOM = "RANDOM" # Random action selection
HYBRID = "HYBRID" # Mix of multiple strategies
# ==============================================================================
# Strategy Genome
# ==============================================================================
@dataclass(frozen=True, slots=True)
class StrategyGenome:
"""
A strategy encoded as a genome for genetic operations.
The genome has two parts:
1. Structural: strategy type, action menu config
2. Parametric: the 29 tunable parameters
Genetic operations (crossover, mutation) work on this genome.
"""
strategy_type: StrategyType
params: FulfilmentPolicyParams
generation: int = 0
parent_ids: Tuple[str, ...] = ()
fitness: float = 0.0
episodes_tested: int = 0
creation_ts_ns: int = 0
@property
def genome_id(self) -> str:
"""Unique identifier for this genome."""
return f"{self.strategy_type.value}_{self.params.version}_{self.generation}"
# ==============================================================================
# Genetic Operators
# ==============================================================================
class GeneticOperators:
"""
Genetic operators for strategy evolution.
Crossover: combine two strategies to create offspring
Mutation: randomly modify a strategy
Selection: tournament selection based on fitness
"""
def __init__(self, codec: CMAParameterCodec, mutation_rate: float = 0.15,
crossover_rate: float = 0.7) -> None:
self.codec = codec
self.mutation_rate = mutation_rate
self.crossover_rate = crossover_rate
def crossover(
self,
parent1: StrategyGenome,
parent2: StrategyGenome,
rng: random.Random,
) -> StrategyGenome:
"""
Uniform crossover: for each parameter, randomly pick from parent1 or parent2.
Strategy type is inherited from the fitter parent.
"""
# Strategy type from fitter parent
strategy_type = parent1.strategy_type if parent1.fitness >= parent2.fitness else parent2.strategy_type
# Crossover parameters
p1_vec = self.codec.initial_vector(parent1.params)
p2_vec = self.codec.initial_vector(parent2.params)
child_vec = []
for i in range(len(p1_vec)):
if rng.random() < 0.5:
child_vec.append(p1_vec[i])
else:
child_vec.append(p2_vec[i])
# Decode child
child_params = self.codec.decode(child_vec, version=f"child_{int(time.time_ns())}")
return StrategyGenome(
strategy_type=strategy_type,
params=child_params,
generation=max(parent1.generation, parent2.generation) + 1,
parent_ids=(parent1.genome_id, parent2.genome_id),
creation_ts_ns=time.time_ns(),
)
def mutate(
self,
genome: StrategyGenome,
rng: random.Random,
) -> StrategyGenome:
"""
Gaussian mutation: add noise to each parameter with mutation_rate probability.
Occasionally mutate strategy type (structural mutation).
"""
vec = self.codec.initial_vector(genome.params)
lows, highs = self.codec.bounds()
mutated_vec = []
for i, (v, lo, hi) in enumerate(zip(vec, lows, highs)):
if rng.random() < self.mutation_rate:
# Gaussian noise scaled by parameter range
range_val = hi - lo
noise = rng.gauss(0, range_val * 0.1)
mutated_vec.append(max(lo, min(hi, v + noise)))
else:
mutated_vec.append(v)
# Decode mutated params
mutated_params = self.codec.decode(mutated_vec, version=f"mut_{int(time.time_ns())}")
# Occasionally mutate strategy type (5% chance)
strategy_type = genome.strategy_type
if rng.random() < 0.05:
strategy_type = rng.choice(list(StrategyType))
return StrategyGenome(
strategy_type=strategy_type,
params=mutated_params,
generation=genome.generation + 1,
parent_ids=(genome.genome_id,),
creation_ts_ns=time.time_ns(),
)
def tournament_select(
self,
population: List[StrategyGenome],
tournament_size: int = 3,
rng: random.Random = None,
) -> StrategyGenome:
"""Tournament selection: pick tournament_size random, return the best."""
if rng is None:
rng = random.Random()
tournament = rng.sample(population, min(tournament_size, len(population)))
return max(tournament, key=lambda g: g.fitness)
def random_genome(
self,
strategy_type: Optional[StrategyType] = None,
rng: random.Random = None,
) -> StrategyGenome:
"""Generate a random genome for initial population."""
if rng is None:
rng = random.Random()
if strategy_type is None:
strategy_type = rng.choice(list(StrategyType))
# Random parameters within bounds
lows, highs = self.codec.bounds()
random_vec = [rng.uniform(lo, hi) for lo, hi in zip(lows, highs)]
params = self.codec.decode(random_vec, version=f"rand_{int(time.time_ns())}")
return StrategyGenome(
strategy_type=strategy_type,
params=params,
generation=0,
creation_ts_ns=time.time_ns(),
)
# ==============================================================================
# Strategy Evaluator
# ==============================================================================
class StrategyEvaluator:
"""
Evaluates strategies through self-play episodes.
Each strategy is tested against the current self-play pool.
Fitness = robust score across scenarios.
CRITICAL: uses the genome's strategy_type to create the correct planner.
This is how different planners (EXP3, Regret Matching, etc.) are actually
used during self-play discovery.
"""
def __init__(self, evaluator: PolicyEvaluator) -> None:
self.evaluator = evaluator
def evaluate(
self,
genome: StrategyGenome,
scenarios: Sequence[Scenario],
pool: SelfPlayPool,
rng_seed: int = 42,
) -> float:
"""Evaluate a genome's fitness through self-play."""
score, results = self.evaluator.evaluate_candidate(
params=genome.params,
scenarios=scenarios,
rng_seed=rng_seed,
planner_type=genome.strategy_type.value, # PASS STRATEGY TYPE
)
return score
def evaluate_population(
self,
population: List[StrategyGenome],
scenarios: Sequence[Scenario],
pool: SelfPlayPool,
rng_seed: int = 42,
) -> List[StrategyGenome]:
"""Evaluate entire population and update fitness scores."""
evaluated = []
for i, genome in enumerate(population):
fitness = self.evaluate(genome, scenarios, pool, rng_seed + i)
evaluated.append(StrategyGenome(
strategy_type=genome.strategy_type,
params=genome.params,
generation=genome.generation,
parent_ids=genome.parent_ids,
fitness=fitness,
episodes_tested=len(scenarios),
creation_ts_ns=genome.creation_ts_ns,
))
return evaluated
# ==============================================================================
# Strategy Generator — the main loop
# ==============================================================================
@dataclass(frozen=True, slots=True)
class GeneratorConfig:
"""Configuration for strategy generation."""
population_size: int = 20
generations: int = 5
tournament_size: int = 3
elitism_count: int = 2 # keep top N unchanged
mutation_rate: float = 0.15
crossover_rate: float = 0.7
max_strategies: int = 50 # max strategies in pool
min_fitness_threshold: float = -100.0 # minimum fitness to keep
class StrategyGenerator:
"""
Genetic programming for strategy evolution.
Key design:
- Hardcoded baseline is NEVER replaced
- Generated strategies are ADDED to the pool
- System GROWS its strategy repertoire
- During self-play, crossover/mutation may "spontaneously generate"
new strategies that weren't explicitly programmed
Flow:
1. Initialize population (random + baseline)
2. Evaluate fitness (self-play episodes)
3. Select parents (tournament selection)
4. Create offspring (crossover + mutation)
5. Evaluate offspring
6. Replace weakest with offspring
7. Repeat for N generations
8. Add successful strategies to pool
"""
def __init__(
self,
config: Optional[GeneratorConfig] = None,
registry: Optional[PolicyRegistry] = None,
pool: Optional[SelfPlayPool] = None,
) -> None:
self.config = config or GeneratorConfig()
self._registry = registry or PolicyRegistry()
self._pool = pool or SelfPlayPool(max_size=self.config.max_strategies)
self._codec = CMAParameterCodec()
self._operators = GeneticOperators(
codec=self._codec,
mutation_rate=self.config.mutation_rate,
crossover_rate=self.config.crossover_rate,
)
self._evaluator = StrategyEvaluator(
PolicyEvaluator(
cwm_factory=lambda: __import__("malkhut.cwm.core", fromlist=["MinimalCryptoLOBCWM"]).MinimalCryptoLOBCWM(),
counterparties=default_counterparty_ecology(),
)
)
self._population: List[StrategyGenome] = []
self._history: List[StrategyGenome] = []
self._rng = random.Random(42)
def initialize_population(
self,
baseline: FulfilmentPolicyParams,
scenarios: Sequence[Scenario],
) -> None:
"""Initialize population with baseline + random variants."""
self._population = []
# Add hardcoded baseline (always available)
baseline_genome = StrategyGenome(
strategy_type=StrategyType.SM_MCTS,
params=baseline,
generation=0,
creation_ts_ns=time.time_ns(),
fitness=0.0,
)
self._population.append(baseline_genome)
# Add random variants
for i in range(self.config.population_size - 1):
genome = self._operators.random_genome(rng=self._rng)
self._population.append(genome)
# Evaluate initial population
self._population = self._evaluator.evaluate_population(
self._population, scenarios, self._pool,
)
def evolve(
self,
baseline: FulfilmentPolicyParams,
scenarios: Sequence[Scenario],
) -> List[StrategyGenome]:
"""
Run genetic evolution for N generations.
Returns the final population (sorted by fitness).
"""
# Initialize if empty
if not self._population:
self.initialize_population(baseline, scenarios)
for gen in range(self.config.generations):
# 1. Select parents
parents = []
for _ in range(self.config.population_size - self.config.elitism_count):
p1 = self._operators.tournament_select(
self._population, self.config.tournament_size, self._rng,
)
p2 = self._operators.tournament_select(
self._population, self.config.tournament_size, self._rng,
)
parents.append((p1, p2))
# 2. Create offspring
offspring = []
for p1, p2 in parents:
if self._rng.random() < self.config.crossover_rate:
child = self._operators.crossover(p1, p2, self._rng)
else:
child = self._operators.mutate(p1, self._rng)
offspring.append(child)
# 3. Evaluate offspring
offspring = self._evaluator.evaluate_population(
offspring, scenarios, self._pool,
)
# 4. Elitism: keep top N unchanged, PLUS always keep baseline
self._population.sort(key=lambda g: g.fitness, reverse=True)
elites = self._population[:self.config.elitism_count]
# Ensure baseline (SM_MCTS with version "baseline") is always present
has_baseline = any(
g.strategy_type == StrategyType.SM_MCTS and g.params.version == "baseline"
for g in elites
)
if not has_baseline:
baseline = next(
(g for g in self._population
if g.strategy_type == StrategyType.SM_MCTS and g.params.version == "baseline"),
None,
)
if baseline:
elites.append(baseline)
# 5. Replace weakest with offspring
self._population = elites + offspring[:self.config.population_size - self.config.elitism_count]
# 6. Sort by fitness
self._population.sort(key=lambda g: g.fitness, reverse=True)
# 7. Track history
self._history.extend(offspring)
return self._population
def get_successful_strategies(
self,
min_fitness: Optional[float] = None,
) -> List[StrategyGenome]:
"""
Get strategies that meet the fitness threshold.
These are candidates for adding to the pool.
"""
threshold = min_fitness or self.config.min_fitness_threshold
return [g for g in self._population if g.fitness > threshold]
def add_to_pool(self, genome: StrategyGenome) -> None:
"""Add a successful strategy to the self-play pool."""
snapshot = PolicySnapshot(
params=genome.params,
score=genome.fitness,
created_ts_ns=genome.creation_ts_ns,
evaluation_summary={
"strategy_type": genome.strategy_type.value,
"generation": genome.generation,
"episodes_tested": genome.episodes_tested,
},
)
self._pool.maybe_add(snapshot)
# Also register in registry
self._registry.register_candidate(
genome.params, genome.fitness,
{"strategy_type": genome.strategy_type.value},
)
def get_diverse_strategies(
self,
n: int = 5,
) -> List[StrategyGenome]:
"""
Get N diverse strategies from the population.
Diversity is measured by:
- Different strategy types
- Different parameter vectors (cosine distance)
"""
if len(self._population) <= n:
return list(self._population)
# Group by strategy type
by_type: dict[StrategyType, list[StrategyGenome]] = {}
for g in self._population:
by_type.setdefault(g.strategy_type, []).append(g)
# Pick one from each type, then fill with best remaining
selected = []
for stype in StrategyType:
if stype in by_type and len(selected) < n:
best = max(by_type[stype], key=lambda g: g.fitness)
selected.append(best)
# Fill with best remaining
remaining = [g for g in self._population if g not in selected]
remaining.sort(key=lambda g: g.fitness, reverse=True)
while len(selected) < n and remaining:
selected.append(remaining.pop(0))
return selected
@property
def population_size(self) -> int:
return len(self._population)
@property
def best_fitness(self) -> float:
if not self._population:
return -float("inf")
return max(g.fitness for g in self._population)
@property
def population(self) -> List[StrategyGenome]:
return list(self._population)

View File

@@ -0,0 +1,75 @@
"""
Execution Hooks — entry/exit interfaces for future BingX integration.
Prepares hooks that will be called when connecting to live execution systems.
All hooks are no-ops now but define the interface for future implementation.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Callable, Optional, Protocol
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction
class ExecutionIntentHook(Protocol):
"""Hook called when an execution intent is submitted."""
def on_intent(self, intent_id: str, target: str, urgency: float) -> None: ...
class FillCallbackHook(Protocol):
"""Hook called after a fill is received."""
def on_fill(self, fill_price: float, fill_qty: float, side: str, is_maker: bool) -> None: ...
class VenueTelemetryHook(Protocol):
"""Hook for venue state publishing to Zinc."""
def on_venue_state(self, state: Mapping[str, Any]) -> None: ...
class AccountReconciliationHook(Protocol):
"""Hook to compare internal state vs venue state."""
def reconcile(self, internal_state: MarketWorldState, venue_state: Mapping[str, Any]) -> list[str]: ...
@dataclass(frozen=True, slots=True)
class ExecutionHooks:
"""
Collection of hooks for future execution integration.
All hooks are no-ops by default. Override when connecting to live systems.
"""
on_intent: Optional[Callable] = None
on_fill: Optional[Callable] = None
on_venue_state: Optional[Callable] = None
reconcile: Optional[Callable] = None
def submit_intent(self, intent_id: str, target: str, urgency: float) -> None:
"""Submit an execution intent (DSL primitive)."""
if self.on_intent:
self.on_intent(intent_id, target, urgency)
def report_fill(self, fill_price: float, fill_qty: float, side: str, is_maker: bool) -> None:
"""Report a fill (DSL primitive)."""
if self.on_fill:
self.on_fill(fill_price, fill_qty, side, is_maker)
def publish_venue_state(self, state: Mapping[str, Any]) -> None:
"""Publish venue state to Zinc."""
if self.on_venue_state:
self.on_venue_state(state)
def reconcile_state(self, internal: MarketWorldState, venue: Mapping[str, Any]) -> list[str]:
"""Reconcile internal vs venue state."""
if self.reconcile:
return self.reconcile(internal, venue)
return []
# Default hooks (no-ops)
DEFAULT_HOOKS = ExecutionHooks()

View File

@@ -0,0 +1,110 @@
"""
Feature Importance Tracker — track which features drive planner decisions.
Enables:
- Understanding which market features matter most
- Identifying overfitting to specific features
- Guiding feature engineering
- Explaining decision rationale
"""
from __future__ import annotations
import time
from collections import defaultdict
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional, Tuple
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction
from malkhut.features import DefaultFeatureExtractor, FeatureExtractor
@dataclass(frozen=True, slots=True)
class FeatureImportance:
"""Importance score for a feature in a specific context."""
feature_name: str
importance: float
regime: str
action_type: str
sample_count: int
class FeatureImportanceTracker:
"""
Track which features drive planner decisions.
Uses a simple attribution method:
- When a decision is made, record which features were above/below thresholds
- Aggregate across decisions to compute importance scores
"""
def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None:
self._extractor = feature_extractor or DefaultFeatureExtractor()
self._feature_counts: Dict[str, Dict[str, int]] = defaultdict(lambda: defaultdict(int))
self._feature_values: Dict[str, List[float]] = defaultdict(list)
self._total_decisions = 0
def record_decision(
self,
state: MarketWorldState,
action: FulfilmentAction,
regime: str = "unknown",
) -> None:
"""Record which features were relevant for this decision."""
self._total_decisions += 1
fv = self._extractor.extract(state).values
# Track which features were "active" (non-zero or above threshold)
for name, value in fv.items():
if abs(value) > 1e-6: # non-zero
self._feature_counts[name][regime] += 1
self._feature_counts[name]["_total"] += 1
# Track value distribution
self._feature_values[name].append(value)
def get_importance(
self,
top_n: int = 10,
regime: Optional[str] = None,
) -> List[FeatureImportance]:
"""Get top N most important features."""
scores = []
for name, regime_counts in self._feature_counts.items():
total = regime_counts.get("_total", 0)
if regime:
count = regime_counts.get(regime, 0)
else:
count = total
importance = count / max(self._total_decisions, 1)
scores.append(FeatureImportance(
feature_name=name,
importance=importance,
regime=regime or "all",
action_type="all",
sample_count=count,
))
scores.sort(key=lambda s: s.importance, reverse=True)
return scores[:top_n]
def get_feature_stats(self, feature_name: str) -> Dict[str, float]:
"""Get statistics for a specific feature."""
values = self._feature_values.get(feature_name, [])
if not values:
return {}
return {
"mean": sum(values) / len(values),
"min": min(values),
"max": max(values),
"count": len(values),
}
@property
def total_decisions(self) -> int:
return self._total_decisions
@property
def feature_count(self) -> int:
return len(self._feature_counts)

View File

@@ -0,0 +1,100 @@
"""
Structured Observability — compact JSONL logging for every decision.
Records:
- Every planner decision with features and diagnostics
- Every risk gate decision
- Every venue action
- Every fill/discrepancy
- Replayable audit trail
"""
from __future__ import annotations
import json
import time
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
@dataclass(frozen=True, slots=True)
class DecisionRecord:
"""One decision in the audit trail."""
ts_ns: int
symbol: str
action_kind: str
action_side: Optional[str]
action_price: Optional[float]
approved: bool
risk_reason: str
plan_latency_ns: int
policy_version: str
root_entropy: float
sims: int
features: Mapping[str, float] = field(default_factory=dict)
diagnostics: Mapping[str, Any] = field(default_factory=dict)
class ObservabilityLogger:
"""
Compact JSONL logger for every decision.
One line per decision. Replayable audit trail.
"""
def __init__(self, log_path: str = "decisions.log") -> None:
self._log_path = log_path
self._decisions: list[DecisionRecord] = []
self._total_decisions = 0
def log_decision(
self,
state: MarketWorldState,
planned: PlannedPolicy,
decision: RiskDecision,
plan_ns: int,
features: Optional[Mapping[str, float]] = None,
) -> None:
"""Log a decision to the audit trail."""
record = DecisionRecord(
ts_ns=state.ts_ns,
symbol=state.venue.symbol,
action_kind=planned.selected_action.kind.value,
action_side=planned.selected_action.side.value if planned.selected_action.side else None,
action_price=planned.selected_action.price_ticks_from_best if planned.selected_action else None,
approved=decision.approved,
risk_reason=decision.reason,
plan_latency_ns=plan_ns,
policy_version="live",
root_entropy=planned.diagnostics.get("entropy", 0.0),
sims=planned.diagnostics.get("sims", 0),
features=features or {},
)
self._decisions.append(record)
self._total_decisions += 1
# Write to JSONL
try:
with open(self._log_path, "a") as f:
f.write(json.dumps({
"ts": record.ts_ns,
"sym": record.symbol,
"act": record.action_kind,
"side": record.action_side,
"app": record.approved,
"risk": record.risk_reason,
"lat_ns": record.plan_latency_ns,
"entropy": round(record.root_entropy, 4),
"sims": record.sims,
}, separators=(",", ":")) + "\n")
except OSError:
pass
@property
def total_decisions(self) -> int:
return self._total_decisions
def get_recent(self, n: int = 10) -> List[DecisionRecord]:
return self._decisions[-n:]

View File

@@ -0,0 +1,134 @@
"""
Training Parallelism — parallel evaluation across scenarios.
CMA-ES evaluates sequentially. This module parallelizes evaluation
across scenarios for 4-8x faster convergence.
"""
from __future__ import annotations
import concurrent.futures
import time
from dataclasses import dataclass
from typing import Any, Callable, List, Optional, Sequence
from malkhut.state import FulfilmentPolicyParams, MarketWorldState
from malkhut.training.cma_trainer import PolicyEvaluator, Scenario
class ParallelEvaluator:
"""
Parallel evaluation of strategies across scenarios.
Uses ThreadPoolExecutor for I/O-bound scenarios.
Uses ProcessPoolExecutor for CPU-bound scenarios.
"""
def __init__(
self,
evaluator: PolicyEvaluator,
max_workers: int = 4,
) -> None:
self.evaluator = evaluator
self._max_workers = max_workers
def evaluate_candidate(
self,
params: FulfilmentPolicyParams,
scenarios: Sequence[Scenario],
rng_seed: int = 0,
) -> tuple[float, list]:
"""Evaluate candidate with parallel scenario execution."""
if len(scenarios) <= 1 or self._max_workers <= 1:
return self.evaluator.evaluate_candidate(params, scenarios, rng_seed)
# Split scenarios across workers
chunk_size = max(1, len(scenarios) // self._max_workers)
chunks = []
for i in range(0, len(scenarios), chunk_size):
chunks.append(scenarios[i:i + chunk_size])
# Parallel evaluation
all_results = []
with concurrent.futures.ThreadPoolExecutor(max_workers=self._max_workers) as executor:
futures = []
for i, chunk in enumerate(chunks):
future = executor.submit(
self.evaluator.evaluate_candidate,
params, chunk, rng_seed + i,
)
futures.append(future)
for future in concurrent.futures.as_completed(futures):
score, results = future.result()
all_results.extend(results)
# Aggregate scores
if all_results:
scores = [r.pnl_bps for r in all_results]
avg_score = sum(scores) / len(scores)
else:
avg_score = 0.0
return avg_score, all_results
@dataclass(frozen=True, slots=True)
class TrainingMetrics:
"""Metrics for a training run."""
total_time_s: float
generations: int
total_evals: int
best_score: float
avg_score: float
score_improvement: float
convergence_gen: int
class TrainingMonitor:
"""
Monitor training progress and convergence.
Tracks metrics, detects convergence, logs progress.
"""
def __init__(self) -> None:
self._scores: list[float] = []
self._times: list[float] = []
self._start_time = time.time()
def record_generation(self, score: float) -> None:
self._scores.append(score)
self._times.append(time.time())
@property
def best_score(self) -> float:
return max(self._scores) if self._scores else 0.0
@property
def avg_score(self) -> float:
return sum(self._scores) / len(self._scores) if self._scores else 0.0
@property
def improvement(self) -> float:
if len(self._scores) < 2:
return 0.0
return self._scores[-1] - self._scores[0]
@property
def converged(self) -> bool:
if len(self._scores) < 5:
return False
recent = self._scores[-5:]
variance = sum((s - self.avg_score) ** 2 for s in recent) / len(recent)
return variance < 0.01 # low variance = converged
def metrics(self) -> TrainingMetrics:
return TrainingMetrics(
total_time_s=time.time() - self._start_time,
generations=len(self._scores),
total_evals=0,
best_score=self.best_score,
avg_score=self.avg_score,
score_improvement=self.improvement,
convergence_gen=len(self._scores) if self.converged else -1,
)

View File

@@ -0,0 +1,106 @@
"""
Policy Rollback — auto-rollback if shadow performance degrades.
Enables:
- Detecting when a promoted policy performs worse than baseline
- Automatically reverting to previous best
- Preventing bad policies from reaching live
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import Any, List, Mapping, Optional, Tuple
from malkhut.state import FulfilmentPolicyParams
from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord
@dataclass(frozen=True, slots=True)
class RollbackEvent:
"""Record of a rollback event."""
ts_ns: int
rolled_back_version: str
rolled_back_to: str
reason: str
performance_drop: float
class PolicyRollback:
"""
Auto-rollback if shadow performance degrades.
Monitors active policy performance and reverts to previous best
if performance drops below threshold.
"""
def __init__(
self,
registry: PolicyRegistry,
degradation_threshold: float = -5.0, # bps
min_shadow_steps: int = 10,
) -> None:
self._registry = registry
self._degradation_threshold = degradation_threshold
self._min_shadow_steps = min_shadow_steps
self._shadow_scores: List[float] = []
self._rollback_events: List[RollbackEvent] = []
def record_shadow_score(self, score: float) -> None:
"""Record a shadow performance score."""
self._shadow_scores.append(score)
def check_rollback(self) -> Optional[RollbackEvent]:
"""
Check if rollback is needed.
Returns RollbackEvent if rollback should happen, None otherwise.
"""
if len(self._shadow_scores) < self._min_shadow_steps:
return None
# Compare recent average to baseline
recent = self._shadow_scores[-self._min_shadow_steps:]
avg_recent = sum(recent) / len(recent)
# Get baseline score (first policy in registry)
baseline = self._registry.get_by_stage(PolicyStage.ACTIVE)
if not baseline:
return None
baseline_record = baseline[0] if baseline else None
if not baseline_record:
return None
# Check degradation
if avg_recent < self._degradation_threshold:
# Find previous best to rollback to
previous = self._registry.get_by_stage(PolicyStage.RETIRED)
if previous:
rollback_to = previous[0].version
else:
rollback_to = "baseline"
# Perform rollback
current = baseline_record[0] if isinstance(baseline_record, list) else baseline_record
self._registry.retire(current.version, "auto_rollback_degradation")
event = RollbackEvent(
ts_ns=time.time_ns(),
rolled_back_version=current.version,
rolled_back_to=rollback_to,
reason=f"performance_drop_{avg_recent:.2f}",
performance_drop=avg_recent - self._degradation_threshold,
)
self._rollback_events.append(event)
return event
return None
@property
def rollback_events(self) -> List[RollbackEvent]:
return list(self._rollback_events)
@property
def shadow_score_count(self) -> int:
return len(self._shadow_scores)

View File

@@ -0,0 +1,195 @@
"""
Scenario Stress Testing — diversified stress scenarios for robust evaluation.
Generates scenarios that test edge cases:
- Flash crash (sudden price drop)
- Liquidity vacuum (no bids/asks)
- Extreme volatility
- Correlated moves across assets
- Weekend/low participation
- Funding shock
- Liquidation cascade
"""
from __future__ import annotations
import random
import time
from dataclasses import dataclass, field
from typing import Any, Optional, Sequence, Tuple
from malkhut.state import (
AccountState, MarketWorldState, Mode, OrderBookState, PositionState,
PriceLevel, Side, TradePathState, VenueRules,
)
from malkhut.counterparties import (
CounterpartyPolicy, ToxicTakerPolicy, PassiveMakerPolicy,
LatencyArbPolicy, NoiseTraderPolicy, default_counterparty_ecology,
)
@dataclass(frozen=True, slots=True)
class StressScenario:
"""A stress test scenario with specific conditions."""
scenario_id: str
symbol: str
initial_state: MarketWorldState
counterparties: Tuple[CounterpartyPolicy, ...]
max_steps: int
tags: Tuple[str, ...]
description: str
class StressScenarioFactory:
"""
Generate diversified stress scenarios.
Tests extreme market conditions that normal scenarios miss.
"""
def __init__(self, counterparties: Optional[Tuple[CounterpartyPolicy, ...]] = None) -> None:
self.counterparties = counterparties or default_counterparty_ecology()
def flash_crash(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Sudden 5% price drop in 3 steps."""
return StressScenario(
scenario_id=f"flash_crash_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.1, ask_qty=0.1),
counterparties=(ToxicTakerPolicy(sensitivity=0.2),),
max_steps=10,
tags=("stress", "flash_crash", "high_volatility"),
description="Sudden 5% price drop with thin book",
)
def liquidity_vacuum(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Near-zero liquidity on both sides."""
return StressScenario(
scenario_id=f"liquidity_vacuum_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.001, ask_qty=0.001),
counterparties=(ToxicTakerPolicy(sensitivity=0.1),),
max_steps=10,
tags=("stress", "liquidity_vacuum", "thin_book"),
description="Near-zero liquidity, any trade moves price significantly",
)
def extreme_volatility(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Wide spread, high volatility."""
return StressScenario(
scenario_id=f"extreme_vol_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=49000.0, ask=51000.0, bid_qty=0.5, ask_qty=0.5),
counterparties=(ToxicTakerPolicy(sensitivity=0.3), NoiseTraderPolicy()),
max_steps=20,
tags=("stress", "extreme_volatility", "wide_spread"),
description="2000bps spread, high volatility",
)
def toxic_flood(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Multiple toxic takers attacking simultaneously."""
return StressScenario(
scenario_id=f"toxic_flood_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.5, ask_qty=0.5),
counterparties=(
ToxicTakerPolicy(sensitivity=0.2),
ToxicTakerPolicy(sensitivity=0.3),
LatencyArbPolicy(lead_threshold=0.3),
),
max_steps=15,
tags=("stress", "toxic_flood", "adverse_selection"),
description="Multiple toxic actors attacking simultaneously",
)
def choppy_market(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Sideways chop with no clear direction."""
return StressScenario(
scenario_id=f"choppy_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50000.5, bid_qty=0.3, ask_qty=0.3),
counterparties=(NoiseTraderPolicy(), PassiveMakerPolicy(join_probability=0.8)),
max_steps=30,
tags=("stress", "choppy", "noise"),
description="Tight range, high noise, no clear direction",
)
def weekend_low_participation(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Weekend-like conditions: thin book, low volume."""
return StressScenario(
scenario_id=f"weekend_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50002.0, bid_qty=0.2, ask_qty=0.2),
counterparties=(NoiseTraderPolicy(),),
max_steps=15,
tags=("stress", "weekend", "low_participation"),
description="Weekend-like: thin book, low volume, wider spreads",
)
def liquidation_cascade(self, symbol: str = "BTCUSDT") -> StressScenario:
"""Price drops trigger liquidations, which cause more drops."""
return StressScenario(
scenario_id=f"liquidation_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.3, ask_qty=0.3),
counterparties=(
ToxicTakerPolicy(sensitivity=0.3),
NoiseTraderPolicy(),
),
max_steps=20,
tags=("stress", "liquidation_cascade", "cascade"),
description="Liquidation cascade: price drops → liquidations → more drops",
)
def correlation_breakdown(self, symbol: str = "BTCUSDT") -> StressScenario:
"""BTC drops, alts diverge."""
return StressScenario(
scenario_id=f"correlation_{symbol}",
symbol=symbol,
initial_state=self._make_state(symbol, bid=50000.0, ask=50001.0, bid_qty=0.5, ask_qty=0.5),
counterparties=(ToxicTakerPolicy(sensitivity=0.4),),
max_steps=15,
tags=("stress", "correlation_breakdown", "divergence"),
description="BTC drops while correlation breaks down",
)
def build_stress_suite(
self,
symbols: Sequence[str] = ("BTCUSDT",),
) -> Tuple[StressScenario, ...]:
"""Build a complete stress test suite."""
scenarios = []
for symbol in symbols:
scenarios.append(self.flash_crash(symbol))
scenarios.append(self.liquidity_vacuum(symbol))
scenarios.append(self.extreme_volatility(symbol))
scenarios.append(self.toxic_flood(symbol))
scenarios.append(self.choppy_market(symbol))
scenarios.append(self.weekend_low_participation(symbol))
scenarios.append(self.liquidation_cascade(symbol))
scenarios.append(self.correlation_breakdown(symbol))
return tuple(scenarios)
@staticmethod
def _make_state(
symbol: str, bid: float, ask: float,
bid_qty: float = 1.0, ask_qty: float = 1.0,
) -> MarketWorldState:
venue = VenueRules(
exchange="bingx", symbol=symbol, 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,
)
book = OrderBookState(
ts_ns=1_000_000_000, symbol=symbol,
bids=(PriceLevel(bid, bid_qty),),
asks=(PriceLevel(ask, ask_qty),),
)
account = AccountState(
ts_ns=1_000_000_000, equity=10000.0, wallet_balance=10000.0,
available_balance=10000.0, margin_used=0.0, total_notional=0.0,
)
return MarketWorldState(
ts_ns=1_000_000_000, mode=Mode.ENDOGENOUS_AGENT_SIM,
venue=venue, book=book, account=account,
)

View File

@@ -0,0 +1,120 @@
"""
Structured Observability — per-decision feature attribution and metrics.
Tracks:
- Which features drove each decision
- Feature importance over time
- Decision quality metrics
- Regime-specific performance
"""
from __future__ import annotations
import json
import time
from collections import defaultdict
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional, Tuple
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
from malkhut.features import DefaultFeatureExtractor, FeatureExtractor
@dataclass(frozen=True, slots=True)
class DecisionMetrics:
"""Per-decision metrics."""
ts_ns: int
symbol: str
action_kind: str
approved: bool
plan_latency_ns: int
entropy: float
sims: int
feature_attribution: Mapping[str, float]
regime: str
class StructuredObservability:
"""
Structured observability with per-decision feature attribution.
Tracks which features drive decisions and computes aggregate metrics.
"""
def __init__(self, feature_extractor: Optional[FeatureExtractor] = None) -> None:
self._extractor = feature_extractor or DefaultFeatureExtractor()
self._decisions: list[DecisionMetrics] = []
self._feature_importance: Dict[str, List[float]] = defaultdict(list)
self._regime_performance: Dict[str, List[float]] = defaultdict(list)
self._total_decisions = 0
def record_decision(
self,
state: MarketWorldState,
planned: PlannedPolicy,
decision: RiskDecision,
plan_ns: int,
regime: str = "unknown",
) -> None:
"""Record a decision with full feature attribution."""
fv = self._extractor.extract(state).values
# Compute feature attribution (which features are "active")
attribution = {}
for name, value in fv.items():
if abs(value) > 1e-6:
attribution[name] = value
metrics = DecisionMetrics(
ts_ns=state.ts_ns,
symbol=state.venue.symbol,
action_kind=planned.selected_action.kind.value,
approved=decision.approved,
plan_latency_ns=plan_ns,
entropy=planned.diagnostics.get("entropy", 0.0),
sims=planned.diagnostics.get("sims", 0),
feature_attribution=attribution,
regime=regime,
)
self._decisions.append(metrics)
self._total_decisions += 1
# Track feature importance
for name, value in attribution.items():
self._feature_importance[name].append(value)
# Track regime performance
self._regime_performance[regime].append(1.0 if decision.approved else 0.0)
def get_feature_importance(self, top_n: int = 10) -> List[Tuple[str, float]]:
"""Get top N features by average absolute value."""
scores = []
for name, values in self._feature_importance.items():
avg = sum(abs(v) for v in values) / len(values)
scores.append((name, avg))
scores.sort(key=lambda x: x[1], reverse=True)
return scores[:top_n]
def get_regime_approval_rate(self, regime: str) -> float:
"""Get approval rate for a specific regime."""
approvals = self._regime_performance.get(regime, [])
if not approvals:
return 0.0
return sum(approvals) / len(approvals)
@property
def total_decisions(self) -> int:
return self._total_decisions
@property
def avg_latency_ns(self) -> float:
if not self._decisions:
return 0.0
return sum(d.plan_latency_ns for d in self._decisions) / len(self._decisions)
@property
def avg_entropy(self) -> float:
if not self._decisions:
return 0.0
return sum(d.entropy for d in self._decisions) / len(self._decisions)

View File

@@ -0,0 +1,104 @@
"""
Trajectory Persistence — store CWM trajectories to ClickHouse.
Enables:
- Post-hoc analysis of decision quality
- Replay verification against stored trajectories
- Training data for learned leaf values
- Audit trail for every decision
"""
from __future__ import annotations
import json
import time
from dataclasses import dataclass, field
from typing import Any, List, Mapping, Optional, Sequence
from malkhut.state import MarketWorldState
from malkhut.actions import FulfilmentAction, PlannedPolicy, RiskDecision
from malkhut.storage.ch_store import MalkhutCHStore
@dataclass(frozen=True, slots=True)
class TrajectoryStep:
"""One step in a persisted trajectory."""
step_index: int
ts_ns: int
symbol: str
state_hash: str
action_kind: str
action_side: Optional[str]
action_price: Optional[float]
action_qty: float
next_state_hash: str
pnl_bps: float
reward: float
entropy: float
diagnostics: Mapping[str, Any] = field(default_factory=dict)
@dataclass(frozen=True, slots=True)
class TrajectoryRecord:
"""Complete trajectory for one episode."""
trajectory_id: str
policy_version: str
scenario_id: str
seed: int
steps: Tuple[TrajectoryStep, ...]
total_pnl_bps: float
max_drawdown_bps: float
fill_count: int
cancel_count: int
noop_count: int
duration_ns: int
created_ts_ns: int
class TrajectoryPersister:
"""
Persist CWM trajectories to ClickHouse.
Enables post-hoc analysis and replay verification.
"""
def __init__(self, store: MalkhutCHStore) -> None:
self._store = store
def persist(self, record: TrajectoryRecord) -> None:
"""Persist a trajectory record to CH."""
self._store.store_episode(
policy_version=record.policy_version,
scenario_id=record.scenario_id,
seed=record.seed,
pnl_bps=record.total_pnl_bps,
max_drawdown_bps=record.max_drawdown_bps,
fill_ratio=record.fill_count / max(record.fill_count + record.cancel_count + record.noop_count, 1),
adverse_fill_ratio=0.0,
avg_slippage_bps=0.0,
liq_near_misses=0,
cancel_count=record.cancel_count,
diagnostics=json.dumps({
"trajectory_id": record.trajectory_id,
"steps": len(record.steps),
"fill_count": record.fill_count,
"noop_count": record.noop_count,
"duration_ns": record.duration_ns,
}),
)
def query_trajectories(
self,
policy_version: Optional[str] = None,
scenario_id: Optional[str] = None,
limit: int = 100,
) -> str:
"""Query stored trajectories."""
where_parts = []
if policy_version:
where_parts.append(f"policy_version = '{policy_version}'")
if scenario_id:
where_parts.append(f"scenario_id = '{scenario_id}'")
where = " AND ".join(where_parts) if where_parts else "1=1"
return self._store.query(
f"SELECT * FROM self_play_episodes WHERE {where} ORDER BY ts_ns DESC LIMIT {limit}"
)