feat: add PDT rolling 5-day window tracking with CSV persistence
This commit is contained in:
+11
-14
@@ -2,12 +2,9 @@ import time
|
|||||||
import logging
|
import logging
|
||||||
import backoff
|
import backoff
|
||||||
import alpaca_trade_api as tradeapi
|
import alpaca_trade_api as tradeapi
|
||||||
import requests.exceptions
|
|
||||||
|
|
||||||
logging.getLogger('backoff').setLevel(logging.CRITICAL)
|
logging.getLogger('backoff').setLevel(logging.CRITICAL)
|
||||||
|
|
||||||
_RETRYABLE_ERRORS = (tradeapi.rest.APIError, ConnectionError, requests.exceptions.ConnectionError)
|
|
||||||
|
|
||||||
|
|
||||||
def _is_position_not_found(e):
|
def _is_position_not_found(e):
|
||||||
return isinstance(e, tradeapi.rest.APIError) and "position does not exist" in str(e)
|
return isinstance(e, tradeapi.rest.APIError) and "position does not exist" in str(e)
|
||||||
@@ -17,50 +14,50 @@ class AlpacaClient:
|
|||||||
def __init__(self, api_key_id, api_secret_key, base_url, api_version="v2"):
|
def __init__(self, api_key_id, api_secret_key, base_url, api_version="v2"):
|
||||||
self.api = tradeapi.REST(api_key_id, api_secret_key, base_url, api_version=api_version)
|
self.api = tradeapi.REST(api_key_id, api_secret_key, base_url, api_version=api_version)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def get_account(self):
|
def get_account(self):
|
||||||
return self.api.get_account()
|
return self.api.get_account()
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def get_clock(self):
|
def get_clock(self):
|
||||||
return self.api.get_clock()
|
return self.api.get_clock()
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def get_bars(self, symbol, timeframe, **kwargs):
|
def get_bars(self, symbol, timeframe, **kwargs):
|
||||||
bars = self.api.get_bars(symbol, timeframe, **kwargs)
|
bars = self.api.get_bars(symbol, timeframe, **kwargs)
|
||||||
if bars is None:
|
if bars is None:
|
||||||
return None
|
return None
|
||||||
return bars.df
|
return bars.df
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def get_latest_quote(self, symbol):
|
def get_latest_quote(self, symbol):
|
||||||
return self.api.get_latest_quote(symbol)
|
return self.api.get_latest_quote(symbol)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def submit_order(self, **kwargs):
|
def submit_order(self, **kwargs):
|
||||||
return self.api.submit_order(**kwargs)
|
return self.api.submit_order(**kwargs)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def get_order(self, order_id):
|
def get_order(self, order_id):
|
||||||
return self.api.get_order(order_id)
|
return self.api.get_order(order_id)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def cancel_order(self, order_id):
|
def cancel_order(self, order_id):
|
||||||
return self.api.cancel_order(order_id)
|
return self.api.cancel_order(order_id)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def list_positions(self):
|
def list_positions(self):
|
||||||
return self.api.list_positions()
|
return self.api.list_positions()
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def list_orders(self, **kwargs):
|
def list_orders(self, **kwargs):
|
||||||
return self.api.list_orders(**kwargs)
|
return self.api.list_orders(**kwargs)
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter)
|
||||||
def close_all_positions(self):
|
def close_all_positions(self):
|
||||||
return self.api.close_all_positions()
|
return self.api.close_all_positions()
|
||||||
|
|
||||||
@backoff.on_exception(backoff.expo, _RETRYABLE_ERRORS, max_tries=5, jitter=backoff.full_jitter, giveup=_is_position_not_found)
|
@backoff.on_exception(backoff.expo, (tradeapi.rest.APIError, ConnectionError), max_tries=5, jitter=backoff.full_jitter, giveup=_is_position_not_found)
|
||||||
def get_position(self, symbol):
|
def get_position(self, symbol):
|
||||||
return self.api.get_position(symbol)
|
return self.api.get_position(symbol)
|
||||||
|
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ TRADES_PATH = SCRIPT_DIR / "trades.csv"
|
|||||||
SIGNALS_PATH = SCRIPT_DIR / "signals.csv"
|
SIGNALS_PATH = SCRIPT_DIR / "signals.csv"
|
||||||
PERFORMANCE_PATH = SCRIPT_DIR / "performance.csv"
|
PERFORMANCE_PATH = SCRIPT_DIR / "performance.csv"
|
||||||
INDICATORS_PATH = SCRIPT_DIR / "indicators.csv"
|
INDICATORS_PATH = SCRIPT_DIR / "indicators.csv"
|
||||||
|
PDT_TRACKER_PATH = SCRIPT_DIR / "pdt_tracker.csv"
|
||||||
|
|
||||||
logging.basicConfig(
|
logging.basicConfig(
|
||||||
level=logging.INFO,
|
level=logging.INFO,
|
||||||
@@ -258,6 +259,9 @@ MIN_NOTIONAL = float(config["MIN_NOTIONAL"])
|
|||||||
POLL_INTERVAL = int(config["POLL_INTERVAL"])
|
POLL_INTERVAL = int(config["POLL_INTERVAL"])
|
||||||
MAX_DRAWDOWN = float(config["MAX_DRAWDOWN"])
|
MAX_DRAWDOWN = float(config["MAX_DRAWDOWN"])
|
||||||
PDT_RULE = bool(config["PDT_RULE"])
|
PDT_RULE = bool(config["PDT_RULE"])
|
||||||
|
if PDT_RULE:
|
||||||
|
_startup_pdt = PDTTracker()
|
||||||
|
logger.info(f" PDT Rule Enforcement: ON ({_startup_pdt.rolling_count()}/3 trades used, {_startup_pdt.remaining()} remaining this window)")
|
||||||
USE_TRAILING_STOP = bool(config["USE_TRAILING_STOP"])
|
USE_TRAILING_STOP = bool(config["USE_TRAILING_STOP"])
|
||||||
PROFIT_TARGET_1 = float(config["PROFIT_TARGET_1"])
|
PROFIT_TARGET_1 = float(config["PROFIT_TARGET_1"])
|
||||||
PROFIT_TARGET_2 = float(config["PROFIT_TARGET_2"])
|
PROFIT_TARGET_2 = float(config["PROFIT_TARGET_2"])
|
||||||
@@ -350,6 +354,57 @@ class SettlementTracker:
|
|||||||
self.pending_settlements = {}
|
self.pending_settlements = {}
|
||||||
|
|
||||||
|
|
||||||
|
class PDTTracker:
|
||||||
|
PDT_LIMIT = 3
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.trade_dates = self._load()
|
||||||
|
|
||||||
|
def _load(self):
|
||||||
|
try:
|
||||||
|
if not PDT_TRACKER_PATH.exists():
|
||||||
|
return []
|
||||||
|
df = pd.read_csv(PDT_TRACKER_PATH)
|
||||||
|
return [datetime.fromisoformat(ts).date() for ts in df['trade_date'].tolist()]
|
||||||
|
except Exception:
|
||||||
|
return []
|
||||||
|
|
||||||
|
def _save(self):
|
||||||
|
try:
|
||||||
|
df = pd.DataFrame({'trade_date': [d.isoformat() for d in self.trade_dates]})
|
||||||
|
df.to_csv(PDT_TRACKER_PATH, index=False)
|
||||||
|
except Exception as e:
|
||||||
|
debug_print(f"PDT tracker save error: {e}")
|
||||||
|
|
||||||
|
def _rolling_window_dates(self):
|
||||||
|
today = datetime.now(EASTERN).date()
|
||||||
|
trading_days = []
|
||||||
|
d = today
|
||||||
|
while len(trading_days) < 5:
|
||||||
|
if d.weekday() < 5:
|
||||||
|
trading_days.append(d)
|
||||||
|
d -= timedelta(days=1)
|
||||||
|
return set(trading_days)
|
||||||
|
|
||||||
|
def rolling_count(self):
|
||||||
|
window = self._rolling_window_dates()
|
||||||
|
return sum(1 for d in self.trade_dates if d in window)
|
||||||
|
|
||||||
|
def can_trade(self):
|
||||||
|
return self.rolling_count() < self.PDT_LIMIT
|
||||||
|
|
||||||
|
def record_trade(self):
|
||||||
|
today = datetime.now(EASTERN).date()
|
||||||
|
self.trade_dates.append(today)
|
||||||
|
cutoff = today - timedelta(days=30)
|
||||||
|
self.trade_dates = [d for d in self.trade_dates if d >= cutoff]
|
||||||
|
self._save()
|
||||||
|
debug_print(f"PDT trade recorded. Rolling 5-day count: {self.rolling_count()}/{self.PDT_LIMIT}")
|
||||||
|
|
||||||
|
def remaining(self):
|
||||||
|
return max(0, self.PDT_LIMIT - self.rolling_count())
|
||||||
|
|
||||||
|
|
||||||
class SignalState:
|
class SignalState:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.last_bullish_crossover_bar = -999
|
self.last_bullish_crossover_bar = -999
|
||||||
@@ -1288,6 +1343,7 @@ def main():
|
|||||||
logger.info(f"💵 Starting equity: ${opening_equity:.2f}")
|
logger.info(f"💵 Starting equity: ${opening_equity:.2f}")
|
||||||
|
|
||||||
settlement_tracker = SettlementTracker()
|
settlement_tracker = SettlementTracker()
|
||||||
|
pdt_tracker = PDTTracker() if PDT_RULE else None
|
||||||
|
|
||||||
if T1_SETTLEMENT_ENABLED:
|
if T1_SETTLEMENT_ENABLED:
|
||||||
settlement_tracker.settle_funds(current_date)
|
settlement_tracker.settle_funds(current_date)
|
||||||
@@ -1670,6 +1726,14 @@ def main():
|
|||||||
debug_print(f"Daily trade limit reached ({trades_today}/{MAX_TRADES_PER_DAY})")
|
debug_print(f"Daily trade limit reached ({trades_today}/{MAX_TRADES_PER_DAY})")
|
||||||
time.sleep(POLL_INTERVAL)
|
time.sleep(POLL_INTERVAL)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if PDT_RULE and pdt_tracker and not pdt_tracker.can_trade():
|
||||||
|
if signal in ['buy', 'sell'] and strength > 0:
|
||||||
|
log_missed_signal(datetime.now(EASTERN), signal, 'pdt_limit', current_price, SYMBOL, strength, signal_rsi, signal_adx, regime)
|
||||||
|
logger.warning(f"🚫 PDT limit reached ({pdt_tracker.rolling_count()}/3 trades in rolling 5-day window) - monitoring only")
|
||||||
|
debug_print(f"PDT limit reached, skipping signal")
|
||||||
|
time.sleep(POLL_INTERVAL)
|
||||||
|
continue
|
||||||
|
|
||||||
if signal == 'sell' and not ENABLE_SHORT_SELLING:
|
if signal == 'sell' and not ENABLE_SHORT_SELLING:
|
||||||
debug_print("Short selling disabled, ignoring sell signal")
|
debug_print("Short selling disabled, ignoring sell signal")
|
||||||
@@ -1704,6 +1768,8 @@ def main():
|
|||||||
if execution_price:
|
if execution_price:
|
||||||
trade_count += 1
|
trade_count += 1
|
||||||
trades_today += 1
|
trades_today += 1
|
||||||
|
if PDT_RULE and pdt_tracker:
|
||||||
|
pdt_tracker.record_trade()
|
||||||
entry_price = execution_price
|
entry_price = execution_price
|
||||||
entry_time = datetime.now(EASTERN)
|
entry_time = datetime.now(EASTERN)
|
||||||
stop_loss = signal_stop_loss
|
stop_loss = signal_stop_loss
|
||||||
|
|||||||
Reference in New Issue
Block a user