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.

155 lines
5.6 KiB
Python

"""Label Generator: compute forward-return labels for all valid events."""
from __future__ import annotations
import argparse
import asyncio
import datetime as dt
import uuid
from sqlalchemy import select
from libs.common.config import get_settings
from libs.common.ids import new_job_run_id
from libs.common.logging import bind_job_run_id, configure_logging, get_logger
from libs.db.models import Event, EventLabel, JobRun, SymbolMaster
from libs.db.session import get_session
from libs.labeler.label_generator import LABEL_VERSION, generate_labels
from libs.oracle_client import PriceService, make_oracle_client
logger = get_logger(__name__)
async def run_label_generator(
run_id: str,
entry_convention: str = "next_open_after_reaction_close",
event_id_filter: str | None = None,
start_date: str | None = None,
end_date: str | None = None,
) -> dict[str, int]:
stats = {"seen": 0, "labeled": 0, "skipped": 0, "errors": 0}
async with get_session() as session, make_oracle_client() as oracle:
price_svc = PriceService(oracle)
job = JobRun(
job_run_id=uuid.UUID(run_id),
job_name="label_generator",
source_name="oracle",
run_date=dt.date.today(),
status="running",
)
session.add(job)
await session.flush()
# Query: events with status=valid (or specific event_id)
stmt = select(Event, SymbolMaster).join(
SymbolMaster, Event.symbol_id == SymbolMaster.symbol_id, isouter=True
).where(Event.status == "valid")
if event_id_filter:
stmt = stmt.where(Event.event_id == event_id_filter)
if start_date:
stmt = stmt.where(Event.event_date >= dt.date.fromisoformat(start_date))
if end_date:
stmt = stmt.where(Event.event_date <= dt.date.fromisoformat(end_date))
result = await session.execute(stmt.order_by(Event.event_date, Event.event_id))
rows = result.all()
stats["seen"] = len(rows)
for event, symbol in rows:
if symbol is None:
logger.warning("label_no_symbol", event_id=event.event_id)
stats["skipped"] += 1
continue
# Skip if label already exists. Regenerate when:
# - status == "pending" (market has since closed)
# - status == "unavailable" AND entry_date is null/future (recover
# from prior runs that 404'd on a future window — see
# label_price_pending_future_window flow in label_generator).
existing = await session.execute(
select(EventLabel).where(
EventLabel.event_id == event.event_id,
EventLabel.entry_convention == entry_convention,
EventLabel.label_version == LABEL_VERSION,
)
)
existing_label = existing.scalar_one_or_none()
if existing_label is not None:
should_regenerate = existing_label.label_status == "pending" or (
existing_label.label_status == "unavailable"
and (
existing_label.entry_date is None
or existing_label.entry_date >= dt.date.today()
)
)
if not should_regenerate:
stats["skipped"] += 1
continue
# Stale pending or recoverable unavailable — delete and regenerate
await session.delete(existing_label)
await session.flush()
try:
label = await generate_labels(
session=session,
event=event,
price_svc=price_svc,
ticker=symbol.ticker,
entry_convention=entry_convention,
)
session.add(label)
await session.flush()
stats["labeled"] += 1
logger.info(
"label_created",
event_id=event.event_id,
label_status=label.label_status,
)
except Exception as exc:
logger.error("label_error", event_id=event.event_id, error=str(exc))
stats["errors"] += 1
job.status = "succeeded" if stats["errors"] == 0 else "partial"
job.finished_at_utc = dt.datetime.now(tz=dt.UTC)
job.records_seen = stats["seen"]
job.records_written = stats["labeled"]
job.error_count = stats["errors"]
logger.info("label_generator_done", **stats)
return stats
def main() -> None:
parser = argparse.ArgumentParser(description="Event Label Generator")
parser.add_argument("--run-id", default=new_job_run_id())
parser.add_argument(
"--entry-convention",
default="next_open_after_reaction_close",
choices=["next_open_after_reaction_close", "reaction_close"],
help="Entry price convention",
)
parser.add_argument("--event-id", default=None, help="Process a single event by ID")
parser.add_argument("--start-date", default=None, metavar="YYYY-MM-DD")
parser.add_argument("--end-date", default=None, metavar="YYYY-MM-DD")
args = parser.parse_args()
settings = get_settings()
configure_logging(settings.log_level)
bind_job_run_id(args.run_id)
asyncio.run(
run_label_generator(
run_id=args.run_id,
entry_convention=args.entry_convention,
event_id_filter=args.event_id,
start_date=args.start_date,
end_date=args.end_date,
)
)
if __name__ == "__main__":
main()