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.
99 lines
3.6 KiB
Python
99 lines
3.6 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 import CompanyService, FinancialService, PriceService, make_oracle_client
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
|
|
async def run_feature_builder(
|
|
run_id: str,
|
|
start_date: str | None = None,
|
|
end_date: str | None = None,
|
|
) -> 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)
|
|
company_svc = CompanyService(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()
|
|
|
|
stmt = select(Event).where(Event.status == "pending")
|
|
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))
|
|
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,
|
|
company_service=company_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())
|
|
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_feature_builder(args.run_id, start_date=args.start_date, end_date=args.end_date))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|