Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.12.10
rev: v0.14.10
hooks:
- id: ruff-format
- id: ruff
Expand Down
8 changes: 5 additions & 3 deletions src/ml4t/backtest/accounting/gatekeeper.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@

from collections.abc import Callable

from ..models import CommissionModel
from ..models import CommissionModel, calculate_commission
from ..types import Order, OrderSide
from .account import AccountState

Expand Down Expand Up @@ -137,7 +137,9 @@ def validate_order(self, order: Order, price: float) -> tuple[bool, str]:
# Check for position reversal (long→short or short→long)
# Delegate to policy's handle_reversal() method
if self._is_reversal(current_qty, order_qty_delta):
commission = self.commission_model.calculate(order.asset, order.quantity, price)
commission = calculate_commission(
self.commission_model, order.asset, order.quantity, price
)
return self.account.policy.handle_reversal(
asset=order.asset,
current_quantity=current_qty,
Expand All @@ -156,7 +158,7 @@ def validate_order(self, order: Order, price: float) -> tuple[bool, str]:

# This is an opening order (new position or adding to existing)
# Calculate commission to include in cost
commission = self.commission_model.calculate(order.asset, order.quantity, price)
commission = calculate_commission(self.commission_model, order.asset, order.quantity, price)

# Use buffered cash (reserves cash_buffer_pct for safety margin)
available = self._available_cash()
Expand Down
14 changes: 11 additions & 3 deletions src/ml4t/backtest/broker.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,8 @@ def from_config(
VolumeShareSlippage,
)

config._validate_for_execution()

effective_commission_type = config.commission_type
if effective_commission_type == CommissionType.NONE:
if config.commission_per_share > 0:
Expand Down Expand Up @@ -1050,11 +1052,17 @@ def flatten_all_positions(

liquidations: list[Order] = []
for asset in list(self.positions):
order = self.close_position(asset, order_type=order_type)
order = self.close_position(
asset,
order_type=order_type,
_options=SubmitOrderOptions(
eligible_in_next_bar_mode=True,
risk_exit_reason=reason,
exit_reason=ExitReason.RISK_LIQUIDATION,
),
)
if order is None:
continue
order._exit_reason = ExitReason.RISK_LIQUIDATION
order._risk_exit_reason = reason
liquidations.append(order)

return liquidations
Expand Down
47 changes: 39 additions & 8 deletions src/ml4t/backtest/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,14 +22,14 @@

from __future__ import annotations

import math
import os
from dataclasses import asdict, dataclass, field, replace
from enum import Enum
from pathlib import Path
from typing import Any

import yaml

from ml4t.specs.base import serialize_artifact_value
from ml4t.specs.market_data import FeedSpec, TimestampSemantics

Expand Down Expand Up @@ -535,7 +535,7 @@ def validate(self, warn: bool = True) -> list[str]:
"""
import warnings as _warnings

issues: list[str] = []
issues: list[str] = self._execution_validation_errors()

# Look-ahead bias warning
if self.execution_mode == ExecutionMode.SAME_BAR:
Expand Down Expand Up @@ -565,12 +565,6 @@ def validate(self, warn: bool = True) -> list[str]:
"Verify this matches your broker's actual costs."
)

if self.slippage_spread < 0:
issues.append(f"slippage_spread ({self.slippage_spread}) must be >= 0")

if any(spread < 0 for spread in self.slippage_spread_by_asset.values()):
issues.append("slippage_spread_by_asset values must all be >= 0")

if (
self.slippage_type == SlippageType.SPREAD
and self.slippage_spread == 0.0
Expand Down Expand Up @@ -641,6 +635,43 @@ def validate(self, warn: bool = True) -> list[str]:

return issues

def _execution_validation_errors(self) -> list[str]:
errors: list[str] = []
cost_fields = (
"commission_rate",
"commission_per_share",
"commission_per_trade",
"commission_minimum",
"slippage_rate",
"slippage_fixed",
"slippage_spread",
"stop_slippage_rate",
)
for field_name in cost_fields:
value = getattr(self, field_name)
try:
valid = math.isfinite(value) and value >= 0.0
except TypeError:
valid = False
if not valid:
errors.append(f"{field_name} ({value!r}) must be finite and >= 0")

for asset, spread in self.slippage_spread_by_asset.items():
try:
valid = math.isfinite(spread) and spread >= 0.0
except TypeError:
valid = False
if not valid:
errors.append(
f"slippage_spread_by_asset[{asset!r}] ({spread!r}) must be finite and >= 0"
)
return errors

def _validate_for_execution(self) -> None:
errors = self._execution_validation_errors()
if errors:
raise ValueError("Invalid BacktestConfig: " + "; ".join(errors))

def get_effective_account_settings(self) -> tuple[bool, bool]:
"""Get account settings as a tuple.

Expand Down
13 changes: 9 additions & 4 deletions src/ml4t/backtest/core/execution_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

import copy

from ..models import calculate_commission
from ..types import ExecutionMode, OrderSide, OrderStatus, OrderType, Position
from .shared import is_exit_order

Expand Down Expand Up @@ -239,8 +240,8 @@ def _validate_shadow_queue_order(
and abs(new_qty) > 1e-12
and ((current_qty > 0 and new_qty < 0) or (current_qty < 0 and new_qty > 0))
)
commission = broker.commission_model.calculate(
order.asset, order.quantity, validation_price
commission = calculate_commission(
broker.commission_model, order.asset, order.quantity, validation_price
)
multiplier = broker.get_multiplier(order.asset)

Expand Down Expand Up @@ -288,7 +289,9 @@ def _commit_shadow_queue_fill(
shadow_positions[order.asset].quantity if order.asset in shadow_positions else 0.0
)
new_qty = current_qty + qty_delta
commission = broker.commission_model.calculate(order.asset, order.quantity, fill_price)
commission = calculate_commission(
broker.commission_model, order.asset, order.quantity, fill_price
)
shadow_cash += -qty_delta * fill_price * broker.get_multiplier(order.asset) - commission

if abs(new_qty) <= 1e-12:
Expand Down Expand Up @@ -586,7 +589,9 @@ def _passes_simple_cash_check(self, order, fill_price: float) -> bool:
return True

signed_qty = order.quantity if order.side is OrderSide.BUY else -order.quantity
commission = broker.commission_model.calculate(order.asset, order.quantity, fill_price)
commission = calculate_commission(
broker.commission_model, order.asset, order.quantity, fill_price
)
projected_cash = broker.cash - signed_qty * fill_price - commission
return projected_cash >= 0.0

Expand Down
21 changes: 12 additions & 9 deletions src/ml4t/backtest/core/order_book.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from datetime import datetime

from ..models import calculate_commission
from ..types import ExecutionMode, Order, OrderSide, OrderStatus, OrderType, Position
from .shared import SubmitOrderOptions, is_exit_order

Expand Down Expand Up @@ -339,16 +340,16 @@ def _passes_submission_precheck(self, order: Order) -> bool:
if closed != 0.0:
close_cash = (-closed) * signal_price
shadow_cash += close_cash
closed_commission = broker.commission_model.calculate(
order.asset, abs(closed), signal_price
closed_commission = calculate_commission(
broker.commission_model, order.asset, abs(closed), signal_price
)
shadow_cash -= closed_commission

if opened != 0.0:
open_cash = opened * signal_price
shadow_cash -= open_cash
opened_commission = broker.commission_model.calculate(
order.asset, abs(opened), signal_price
opened_commission = calculate_commission(
broker.commission_model, order.asset, abs(opened), signal_price
)
shadow_cash -= opened_commission

Expand Down Expand Up @@ -419,8 +420,8 @@ def _passes_buying_power_check(self, order: Order) -> bool:
if closed != 0.0:
closed_value = (-closed) * signal_price
shadow_cash += closed_value
closed_commission = broker.commission_model.calculate(
order.asset, abs(closed), signal_price
closed_commission = calculate_commission(
broker.commission_model, order.asset, abs(closed), signal_price
)
shadow_cash -= closed_commission

Expand All @@ -430,8 +431,8 @@ def _passes_buying_power_check(self, order: Order) -> bool:
# This prevents credit-model inflation where short proceeds
# artificially inflate shadow cash.
shadow_cash -= abs(opened) * signal_price
opened_commission = broker.commission_model.calculate(
order.asset, abs(opened), signal_price
opened_commission = calculate_commission(
broker.commission_model, order.asset, abs(opened), signal_price
)
shadow_cash -= opened_commission

Expand Down Expand Up @@ -461,7 +462,9 @@ def _passes_margin_submission_precheck(self, order: Order, signal_price: float)
old_qty, old_price, size, signal_price
)

commission = broker.commission_model.calculate(order.asset, order.quantity, signal_price)
commission = calculate_commission(
broker.commission_model, order.asset, order.quantity, signal_price
)
available_cash = self._submission_shadow_cash
if broker.cash_buffer_pct > 0 and available_cash > 0:
available_cash *= 1.0 - broker.cash_buffer_pct
Expand Down
1 change: 0 additions & 1 deletion src/ml4t/backtest/datafeed.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,6 @@
from typing import Any

import polars as pl

from ml4t.specs.market_data import FeedSpec


Expand Down
12 changes: 11 additions & 1 deletion src/ml4t/backtest/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,8 @@ def _build_activity_metrics(self) -> dict[str, int | float]:
)
max_open_positions = max((state[5] for state in self.portfolio_state), default=0)
return {
"num_orders": len(self.broker.orders),
"num_rejected_orders": len(self.broker.get_rejected_orders()),
"num_fills": 0,
"num_rebalance_events": 0,
"unique_symbols_traded": 0,
Expand Down Expand Up @@ -335,6 +337,8 @@ def _build_activity_metrics(self) -> dict[str, int | float]:
max_open_positions = max((state[5] for state in self.portfolio_state), default=0)

return {
"num_orders": len(self.broker.orders),
"num_rejected_orders": len(self.broker.get_rejected_orders()),
"num_fills": len(fills),
"num_rebalance_events": len(rebalance_events),
"unique_symbols_traded": len(traded_symbols),
Expand All @@ -356,9 +360,14 @@ def _generate_results(self) -> BacktestResult:
trades=[],
equity_curve=[],
fills=[],
rejected_orders=self.broker.get_rejected_orders(),
predictions=self.feed.signals,
portfolio_state=[],
metrics={"skipped_bars": self._skipped_bars},
metrics={
"skipped_bars": self._skipped_bars,
"num_orders": len(self.broker.orders),
"num_rejected_orders": len(self.broker.get_rejected_orders()),
},
config=self.config,
)

Expand Down Expand Up @@ -475,6 +484,7 @@ def _generate_results(self) -> BacktestResult:
trades=all_trades, # Includes both closed and open trades
equity_curve=self.equity_curve,
fills=self.broker.fills,
rejected_orders=self.broker.get_rejected_orders(),
predictions=self.feed.signals,
portfolio_state=self.portfolio_state,
metrics=metrics,
Expand Down
Loading
Loading