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.
116 lines
4.3 KiB
Python
116 lines
4.3 KiB
Python
"""One-time script: rebuild features + labels for cross-midnight misclassified events.
|
|
|
|
Cross-midnight events: filed_at ET date > event_date AND hour < 10.
|
|
These were stored with ftb=pre_market but should be post_market.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import datetime as dt
|
|
|
|
from sqlalchemy import delete, select, text
|
|
|
|
from libs.common.logging import configure_logging, get_logger
|
|
from libs.db.models import Event, EventLabel, FeatureSnapshot, SymbolMaster
|
|
from libs.db.session import get_session
|
|
from libs.features.builder import build_features_for_event
|
|
from libs.labeler.label_generator import generate_labels
|
|
from libs.oracle_client import CompanyService, FinancialService, PriceService, make_oracle_client
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
async def rebuild_crossmidnight(dry_run: bool = False) -> None:
|
|
configure_logging()
|
|
async with get_session() as session, make_oracle_client() as oracle:
|
|
price_svc = PriceService(oracle)
|
|
financial_svc = FinancialService(oracle)
|
|
company_svc = CompanyService(oracle)
|
|
|
|
# Find cross-midnight events
|
|
result = await session.execute(text("""
|
|
SELECT e.event_id
|
|
FROM events e
|
|
WHERE e.filed_at_utc IS NOT NULL
|
|
AND DATE(e.filed_at_utc AT TIME ZONE 'America/New_York') > e.event_date
|
|
AND EXTRACT(HOUR FROM (e.filed_at_utc AT TIME ZONE 'America/New_York')) < 10
|
|
AND e.status = 'valid'
|
|
ORDER BY e.event_date
|
|
"""))
|
|
event_ids = [r[0] for r in result.fetchall()]
|
|
logger.info("cross_midnight_events_found", count=len(event_ids))
|
|
|
|
if dry_run:
|
|
logger.info("dry_run_mode_no_changes")
|
|
return
|
|
|
|
stats = {"rebuilt": 0, "label_rebuilt": 0, "errors": 0}
|
|
|
|
for i, eid in enumerate(event_ids):
|
|
try:
|
|
# Fetch event + symbol
|
|
evt_result = await session.execute(
|
|
select(Event, SymbolMaster)
|
|
.join(SymbolMaster, Event.symbol_id == SymbolMaster.symbol_id, isouter=True)
|
|
.where(Event.event_id == eid)
|
|
)
|
|
row = evt_result.one_or_none()
|
|
if row is None:
|
|
continue
|
|
event, symbol = row
|
|
|
|
# Delete old feature snapshots
|
|
await session.execute(
|
|
delete(FeatureSnapshot).where(FeatureSnapshot.event_id == eid)
|
|
)
|
|
|
|
# Delete old labels
|
|
await session.execute(
|
|
delete(EventLabel).where(EventLabel.event_id == eid)
|
|
)
|
|
|
|
await session.flush()
|
|
|
|
# Rebuild features
|
|
snapshots = await build_features_for_event(
|
|
session, event, price_svc,
|
|
financial_service=financial_svc,
|
|
company_service=company_svc,
|
|
)
|
|
if snapshots is None:
|
|
logger.warning("feature_rebuild_failed", event_id=eid)
|
|
stats["errors"] += 1
|
|
continue
|
|
|
|
stats["rebuilt"] += 1
|
|
|
|
# Rebuild labels (both conventions)
|
|
ticker = symbol.ticker if symbol else None
|
|
if ticker:
|
|
for convention in ("next_open_after_reaction_close", "reaction_close"):
|
|
lbl = await generate_labels(
|
|
session, event, price_svc, ticker, entry_convention=convention
|
|
)
|
|
if lbl is not None:
|
|
session.add(lbl)
|
|
await session.flush()
|
|
stats["label_rebuilt"] += 1
|
|
|
|
if (i + 1) % 50 == 0:
|
|
logger.info("rebuild_progress", done=i + 1, total=len(event_ids), **stats)
|
|
await session.commit()
|
|
|
|
except Exception as exc:
|
|
logger.error("rebuild_error", event_id=eid, error=str(exc))
|
|
stats["errors"] += 1
|
|
await session.rollback()
|
|
|
|
await session.commit()
|
|
logger.info("rebuild_done", total=len(event_ids), **stats)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
import sys
|
|
dry_run = "--dry-run" in sys.argv
|
|
asyncio.run(rebuild_crossmidnight(dry_run=dry_run))
|