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

"""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))