You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
435 lines
14 KiB
Python
435 lines
14 KiB
Python
"""Build, rank and filter Candidate objects from raw Parquet row dicts."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
from typing import Any
|
|
from zoneinfo import ZoneInfo
|
|
|
|
from libs.backtest.domain import (
|
|
Candidate,
|
|
EventTypeProfile,
|
|
SignalConfig,
|
|
StrategyEngineConfig,
|
|
UniverseConfig,
|
|
)
|
|
from libs.common.logging import get_logger
|
|
|
|
logger = get_logger(__name__)
|
|
_UTC = ZoneInfo("UTC")
|
|
|
|
|
|
def build_candidate(
|
|
row: dict[str, Any],
|
|
strategy_engine: StrategyEngineConfig | None = None,
|
|
) -> Candidate | None:
|
|
"""Build a Candidate from a raw Parquet row dict.
|
|
|
|
Returns None (logged as skip) if:
|
|
- event_timestamp is null/missing
|
|
- entry_price_est is null/zero
|
|
- execution_date is null/missing
|
|
"""
|
|
event_id = row.get("event_id", "")
|
|
|
|
# Strict: no silent substitution for event_timestamp
|
|
raw_ts = row.get("event_timestamp")
|
|
if raw_ts is None:
|
|
logger.warning("skip_candidate_no_timestamp", event_id=event_id)
|
|
return None
|
|
|
|
# Normalise to timezone-aware datetime
|
|
if isinstance(raw_ts, str):
|
|
try:
|
|
event_timestamp = dt.datetime.fromisoformat(raw_ts)
|
|
except ValueError:
|
|
logger.warning("skip_candidate_bad_timestamp", event_id=event_id, raw=raw_ts)
|
|
return None
|
|
elif isinstance(raw_ts, dt.datetime):
|
|
event_timestamp = raw_ts
|
|
else:
|
|
logger.warning("skip_candidate_unknown_timestamp_type", event_id=event_id)
|
|
return None
|
|
|
|
if event_timestamp.tzinfo is None:
|
|
event_timestamp = event_timestamp.replace(tzinfo=_UTC)
|
|
|
|
# reaction_date
|
|
raw_react = row.get("reaction_date")
|
|
reaction_date = _parse_date(raw_react)
|
|
|
|
# execution_date (mapped from Parquet entry_date) or reaction date for close-entry engines
|
|
raw_exec_date = row.get("execution_date") or row.get("entry_date")
|
|
execution_date = _parse_date(raw_exec_date)
|
|
if execution_date is None:
|
|
if strategy_engine and strategy_engine.entry_timing_policy == "reaction_close":
|
|
execution_date = reaction_date
|
|
else:
|
|
logger.warning("skip_candidate_no_exec_date", event_id=event_id)
|
|
return None
|
|
|
|
if reaction_date is None:
|
|
reaction_date = execution_date
|
|
|
|
event_date = _parse_date(row.get("event_date")) or event_timestamp.date()
|
|
timing_class = _classify_timing_class(event_date, reaction_date)
|
|
|
|
trade_direction = _resolve_trade_direction(row)
|
|
if strategy_engine is not None and not _matches_strategy_engine(
|
|
row=row,
|
|
strategy_engine=strategy_engine,
|
|
event_type=str(row.get("event_type", "")),
|
|
timing_class=timing_class,
|
|
trade_direction=trade_direction,
|
|
):
|
|
return None
|
|
|
|
if strategy_engine and strategy_engine.entry_timing_policy == "reaction_close":
|
|
execution_date = reaction_date
|
|
entry_price_est = row.get("event_close") or row.get("entry_price_est")
|
|
if not entry_price_est:
|
|
logger.debug(
|
|
"skip_candidate_no_event_close",
|
|
event_id=event_id,
|
|
engine_id=strategy_engine.engine_id,
|
|
)
|
|
return None
|
|
else:
|
|
entry_price_est = row.get("entry_price") or row.get("entry_price_est")
|
|
|
|
if not entry_price_est:
|
|
logger.warning("skip_candidate_no_entry_price", event_id=event_id)
|
|
return None
|
|
entry_price_est = float(entry_price_est)
|
|
if entry_price_est <= 0:
|
|
logger.warning("skip_candidate_zero_entry_price", event_id=event_id)
|
|
return None
|
|
|
|
score = float(row.get("score", 0.0))
|
|
avg_dollar_volume = float(row.get("avg_dollar_volume", 0.0))
|
|
atr_14_raw = row.get("atr_14")
|
|
atr_14 = float(atr_14_raw) if atr_14_raw is not None else None
|
|
|
|
# Classify score bucket
|
|
score_bucket = _classify_score_bucket(score)
|
|
|
|
return Candidate(
|
|
event_id=event_id,
|
|
symbol=str(row.get("symbol", row.get("ticker", ""))),
|
|
issuer_id=row.get("issuer_id"),
|
|
score=score,
|
|
sector=str(row.get("sector") or "UNKNOWN"),
|
|
event_type=str(row.get("event_type", "")),
|
|
event_timestamp=event_timestamp,
|
|
event_date=event_date,
|
|
filing_time_bucket=str(row.get("filing_time_bucket", "unknown")),
|
|
timing_class=timing_class,
|
|
reaction_date=reaction_date,
|
|
execution_date=execution_date,
|
|
entry_price_est=entry_price_est,
|
|
avg_dollar_volume=avg_dollar_volume,
|
|
atr_14=atr_14,
|
|
score_bucket=score_bucket,
|
|
engine_id=strategy_engine.engine_id if strategy_engine else "default",
|
|
entry_timing_policy=(
|
|
strategy_engine.entry_timing_policy
|
|
if strategy_engine
|
|
else "next_open"
|
|
),
|
|
shadow_only=strategy_engine.shadow_only if strategy_engine else False,
|
|
engine_max_holding_days=(
|
|
strategy_engine.max_holding_days
|
|
if strategy_engine
|
|
else None
|
|
),
|
|
engine_risk_budget_pct=(
|
|
strategy_engine.engine_risk_budget_pct
|
|
if strategy_engine
|
|
else 1.0
|
|
),
|
|
engine_target_atr_multiplier=(
|
|
strategy_engine.target_atr_multiplier_override
|
|
if strategy_engine
|
|
else None
|
|
),
|
|
engine_target_1_fraction=(
|
|
strategy_engine.target_1_fraction_override
|
|
if strategy_engine
|
|
else None
|
|
),
|
|
engine_trailing_model=(
|
|
strategy_engine.trailing_model_override
|
|
if strategy_engine
|
|
else None
|
|
),
|
|
engine_trailing_warmup_days=(
|
|
strategy_engine.trailing_warmup_days_override
|
|
if strategy_engine
|
|
else None
|
|
),
|
|
trade_direction=trade_direction,
|
|
features={k: v for k, v in row.items() if k not in _RESERVED_KEYS},
|
|
)
|
|
|
|
|
|
_RESERVED_KEYS = {
|
|
"event_id", "symbol", "ticker", "issuer_id", "score", "sector",
|
|
"event_type", "event_timestamp", "event_date", "filing_time_bucket", "reaction_date",
|
|
"entry_date", "execution_date", "entry_price", "entry_price_est",
|
|
"event_close", "avg_dollar_volume", "atr_14", "score_bucket", "trade_direction",
|
|
}
|
|
|
|
|
|
def _parse_date(raw: Any) -> dt.date | None:
|
|
if isinstance(raw, dt.datetime):
|
|
return raw.date()
|
|
if isinstance(raw, dt.date):
|
|
return raw
|
|
if isinstance(raw, str):
|
|
try:
|
|
return dt.date.fromisoformat(raw)
|
|
except ValueError:
|
|
return None
|
|
return None
|
|
|
|
|
|
def _classify_timing_class(event_date: dt.date | None, reaction_date: dt.date) -> str:
|
|
if event_date is None:
|
|
return "unknown"
|
|
if reaction_date == event_date:
|
|
return "same_day"
|
|
if reaction_date > event_date:
|
|
return "after_close"
|
|
return "unknown"
|
|
|
|
|
|
def _resolve_trade_direction(row: dict[str, Any]) -> str:
|
|
raw_direction = str(row.get("trade_direction", "")).lower()
|
|
if raw_direction in {"long", "short"}:
|
|
return raw_direction
|
|
|
|
reaction = row.get("reaction_day_return")
|
|
if reaction is not None:
|
|
try:
|
|
return "short" if float(reaction) < 0 else "long"
|
|
except (TypeError, ValueError):
|
|
pass
|
|
return "long"
|
|
|
|
|
|
def _matches_strategy_engine(
|
|
row: dict[str, Any],
|
|
strategy_engine: StrategyEngineConfig,
|
|
event_type: str,
|
|
timing_class: str,
|
|
trade_direction: str,
|
|
) -> bool:
|
|
if strategy_engine.event_types and event_type not in strategy_engine.event_types:
|
|
return False
|
|
|
|
if strategy_engine.timing_class != "any" and timing_class != strategy_engine.timing_class:
|
|
return False
|
|
|
|
if strategy_engine.direction == "long_only" and trade_direction != "long":
|
|
return False
|
|
if strategy_engine.direction == "short_only" and trade_direction != "short":
|
|
return False
|
|
|
|
reaction_day_return = _safe_float(row.get("reaction_day_return"))
|
|
if (
|
|
strategy_engine.reaction_day_return_min is not None
|
|
and reaction_day_return is not None
|
|
and reaction_day_return < strategy_engine.reaction_day_return_min
|
|
):
|
|
return False
|
|
if (
|
|
strategy_engine.reaction_day_return_max is not None
|
|
and reaction_day_return is not None
|
|
and reaction_day_return > strategy_engine.reaction_day_return_max
|
|
):
|
|
return False
|
|
|
|
gap_size = _safe_float(row.get("gap_size"))
|
|
if (
|
|
strategy_engine.gap_size_min is not None
|
|
and gap_size is not None
|
|
and gap_size < strategy_engine.gap_size_min
|
|
):
|
|
return False
|
|
if (
|
|
strategy_engine.gap_size_max is not None
|
|
and gap_size is not None
|
|
and gap_size > strategy_engine.gap_size_max
|
|
):
|
|
return False
|
|
|
|
if (
|
|
strategy_engine.entry_timing_policy == "reaction_close"
|
|
and row.get("event_close") in (None, 0, 0.0, "")
|
|
):
|
|
return False
|
|
|
|
return True
|
|
|
|
|
|
def _classify_score_bucket(score: float) -> str:
|
|
if score >= 0.8:
|
|
return "high"
|
|
if score >= 0.6:
|
|
return "medium_high"
|
|
if score >= 0.4:
|
|
return "medium"
|
|
if score >= 0.2:
|
|
return "medium_low"
|
|
return "low"
|
|
|
|
|
|
def _safe_float(raw: Any) -> float | None:
|
|
try:
|
|
return float(raw)
|
|
except (TypeError, ValueError):
|
|
return None
|
|
|
|
|
|
def rank_candidates(candidates: list[Candidate]) -> list[Candidate]:
|
|
"""Sort by score DESC, avg_dollar_volume DESC, symbol ASC (stable, deterministic)."""
|
|
return sorted(candidates, key=lambda c: (-c.score, -c.avg_dollar_volume, c.symbol))
|
|
|
|
|
|
def filter_by_universe(
|
|
candidates: list[Candidate],
|
|
config: UniverseConfig,
|
|
) -> list[Candidate]:
|
|
"""Apply universe filters: min_price, min_avg_dollar_volume, exchange."""
|
|
filtered = []
|
|
for c in candidates:
|
|
if c.entry_price_est < config.min_price:
|
|
continue
|
|
if c.avg_dollar_volume < config.min_avg_dollar_volume:
|
|
continue
|
|
filtered.append(c)
|
|
return filtered
|
|
|
|
|
|
def filter_by_score(
|
|
candidates: list[Candidate],
|
|
score_threshold: float,
|
|
) -> list[Candidate]:
|
|
return [c for c in candidates if c.score >= score_threshold]
|
|
|
|
|
|
def truncate_candidates(
|
|
candidates: list[Candidate],
|
|
max_per_day: int,
|
|
) -> list[Candidate]:
|
|
return candidates[:max_per_day]
|
|
|
|
|
|
def filter_by_event_type(
|
|
candidates: list[Candidate],
|
|
profiles: dict[str, EventTypeProfile],
|
|
) -> list[Candidate]:
|
|
"""Filter out candidates whose event_type is unknown, disabled, or below per-type threshold.
|
|
|
|
Default-deny: if profiles dict is non-empty and event_type is not in profiles,
|
|
the candidate is skipped (unknown event types are blocked).
|
|
"""
|
|
if not profiles:
|
|
return candidates
|
|
filtered = []
|
|
for c in candidates:
|
|
profile = profiles.get(c.event_type)
|
|
if profile is None:
|
|
logger.debug("skip_unknown_event_type", symbol=c.symbol, event_type=c.event_type)
|
|
continue
|
|
if not profile.enabled:
|
|
logger.debug("skip_disabled_event_type", symbol=c.symbol, event_type=c.event_type)
|
|
continue
|
|
if profile.score_threshold_override is not None:
|
|
if c.score < profile.score_threshold_override:
|
|
logger.debug(
|
|
"skip_event_type_score",
|
|
symbol=c.symbol,
|
|
event_type=c.event_type,
|
|
score=c.score,
|
|
threshold=profile.score_threshold_override,
|
|
)
|
|
continue
|
|
filtered.append(c)
|
|
return filtered
|
|
|
|
|
|
def select_candidates(
|
|
raw_rows: list[dict[str, Any]],
|
|
universe_config: UniverseConfig,
|
|
signal_config: SignalConfig,
|
|
event_type_profiles: dict[str, EventTypeProfile] | None = None,
|
|
strategy_engine: StrategyEngineConfig | None = None,
|
|
) -> list[Candidate]:
|
|
"""Full selection pipeline: build → filter → rank → truncate."""
|
|
candidates = []
|
|
for row in raw_rows:
|
|
prepared_row = _prepare_row_for_strategy_engine(
|
|
row,
|
|
signal_config=signal_config,
|
|
strategy_engine=strategy_engine,
|
|
)
|
|
c = build_candidate(prepared_row, strategy_engine=strategy_engine)
|
|
if c is not None:
|
|
candidates.append(c)
|
|
|
|
candidates = filter_by_universe(candidates, universe_config)
|
|
candidates = filter_by_score(
|
|
candidates,
|
|
_resolve_score_threshold(signal_config, strategy_engine),
|
|
)
|
|
if event_type_profiles:
|
|
candidates = filter_by_event_type(candidates, event_type_profiles)
|
|
candidates = rank_candidates(candidates)
|
|
candidates = truncate_candidates(candidates, signal_config.max_candidates_per_day)
|
|
return candidates
|
|
|
|
|
|
def _resolve_score_threshold(
|
|
signal_config: SignalConfig,
|
|
strategy_engine: StrategyEngineConfig | None,
|
|
) -> float:
|
|
if strategy_engine and strategy_engine.score_threshold_override is not None:
|
|
return strategy_engine.score_threshold_override
|
|
return signal_config.score_threshold
|
|
|
|
|
|
def _prepare_row_for_strategy_engine(
|
|
row: dict[str, Any],
|
|
signal_config: SignalConfig,
|
|
strategy_engine: StrategyEngineConfig | None,
|
|
) -> dict[str, Any]:
|
|
if strategy_engine is None or signal_config.scoring_model != "pead":
|
|
return row
|
|
|
|
reaction_threshold = (
|
|
strategy_engine.pead_reaction_threshold_override
|
|
if strategy_engine.pead_reaction_threshold_override is not None
|
|
else signal_config.pead_reaction_threshold
|
|
)
|
|
volume_threshold = (
|
|
strategy_engine.pead_volume_threshold_override
|
|
if strategy_engine.pead_volume_threshold_override is not None
|
|
else signal_config.pead_volume_threshold
|
|
)
|
|
|
|
if (
|
|
reaction_threshold == signal_config.pead_reaction_threshold
|
|
and volume_threshold == signal_config.pead_volume_threshold
|
|
):
|
|
return row
|
|
|
|
from libs.backtest.scoring import compute_pead_score
|
|
|
|
prepared = dict(row)
|
|
prepared["score"] = compute_pead_score(
|
|
prepared,
|
|
reaction_threshold=reaction_threshold,
|
|
volume_threshold=volume_threshold,
|
|
)
|
|
return prepared
|