Files
2026-06-26 20:50:07 +08:00

79 lines
2.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Gate primitives — boolean masks over bars (doc 04 Rule 6).
A gate is a callable that takes the bars DataFrame and returns a boolean
``numpy.ndarray`` (``True`` = entry allowed on that bar). Gates never open or
close trades; they only mask the signal array the caller hands to the engine.
"""
from __future__ import annotations
from typing import Protocol, runtime_checkable
import numpy as np
import pandas as pd
@runtime_checkable
class Gate(Protocol):
"""A callable producing a boolean mask over bars (``True`` = allow entry)."""
def __call__(self, bars: pd.DataFrame) -> np.ndarray: ...
def time_of_day_gate(
bars: pd.DataFrame,
*,
start_hour: int,
end_hour: int,
timezone: str | None = None,
) -> np.ndarray:
"""Allow entries only within ``[start_hour, end_hour)`` (hours, 024).
Useful for sessions that only trade London/NY open. ``timezone`` is
applied to ``bars['timestamp']`` if given; otherwise the timestamp's
existing tz is used (or naive local time).
"""
ts = pd.to_datetime(bars["timestamp"])
if timezone is not None:
ts = ts.dt.tz_localize(None).dt.tz_localize(timezone) if ts.dt.tz is None else ts.dt.tz_convert(timezone)
hours = ts.dt.hour
if start_hour <= end_hour:
mask = (hours >= start_hour) & (hours < end_hour)
else:
# Wrap past midnight, e.g. 22 → 6.
mask = (hours >= start_hour) | (hours < end_hour)
return mask.to_numpy()
def regime_gate(
bars: pd.DataFrame,
*,
trend_filter: np.ndarray,
direction: int,
) -> np.ndarray:
"""Allow entries only when ``trend_filter`` agrees with ``direction``.
``trend_filter`` is a +1/-1 array (e.g. from an EMA slope or ADX sign).
``direction=+1`` keeps bars where the trend is up; ``-1`` keeps downtrend.
"""
tf = np.asarray(trend_filter)
mask = tf == direction
return mask
def exhaustion_gate(
rsi_arr: np.ndarray,
*,
overbought: float = 70.0,
oversold: float = 30.0,
) -> np.ndarray:
"""Block entries when RSI is in the exhaustion zone for the direction.
Returns ``True`` where entry is *allowed* (i.e. not exhausted). Block longs
when ``rsi >= overbought`` and shorts when ``rsi <= oversold`` — combine
with the directional signal in the caller.
"""
rsi_arr = np.asarray(rsi_arr, dtype=float)
allow = (rsi_arr < overbought) & (rsi_arr > oversold)
allow = np.where(np.isnan(rsi_arr), False, allow)
return allow