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.

89 lines
3.0 KiB
Python

"""Feature Builder: compute market + event features for pending 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, JobRun
from libs.db.session import get_session
from libs.features.builder import build_features_for_event
from libs.oracle_client.client import make_oracle_client
from libs.oracle_client.financial import FinancialService
from libs.oracle_client.price import PriceService
logger = get_logger(__name__)
async def run_feature_builder(run_id: str) -> dict[str, int]:
stats = {"seen": 0, "built": 0, "skipped": 0, "errors": 0}
async with make_oracle_client() as client:
price_svc = PriceService(client)
financial_svc = FinancialService(client)
async with get_session() as session:
job = JobRun(
job_run_id=uuid.UUID(run_id),
job_name="feature_builder",
source_name="oracle",
run_date=dt.date.today(),
status="running",
)
session.add(job)
await session.flush()
result = await session.execute(
select(Event).where(Event.status == "pending")
)
events = result.scalars().all()
stats["seen"] = len(events)
for event in events:
try:
snapshots = await build_features_for_event(
session, event, price_svc, financial_service=financial_svc
)
if snapshots is None:
event.status = "rejected"
event.updated_at_utc = dt.datetime.now(tz=dt.UTC)
stats["skipped"] += 1
else:
event.status = "valid"
event.updated_at_utc = dt.datetime.now(tz=dt.UTC)
stats["built"] += 1
except Exception as exc:
logger.error("feature_build_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["built"]
job.error_count = stats["errors"]
logger.info("feature_builder_done", **stats)
return stats
def main() -> None:
parser = argparse.ArgumentParser(description="Feature Builder")
parser.add_argument("--run-id", default=new_job_run_id())
args = parser.parse_args()
settings = get_settings()
configure_logging(settings.log_level)
bind_job_run_id(args.run_id)
asyncio.run(run_feature_builder(args.run_id))
if __name__ == "__main__":
main()