"""Artifact save/load and model metadata for Brettapps/trifecta-bro-v1.""" from __future__ import annotations import json from dataclasses import asdict, dataclass from pathlib import Path from typing import Any from .predictor import Race, Runner, TrifectaPredictor MODEL_ID = "Brettapps/trifecta-bro/v1" MODEL_CARD_VERSION = "1.0.0" @dataclass class ModelArtifact: model_id: str version: str method: str feature_weights: dict[str, float] notes: str # The "learned" configuration of the rule-based model. Documented explicitly so # the published artifact is interpretable and reproducible. ARTIFACT = ModelArtifact( model_id=MODEL_ID, version=MODEL_CARD_VERSION, method="multi-factor rule-based scoring (0-100)", feature_weights={ "recent_form": 25.0, "overall_win_pct": 20.0, "overall_place_pct": 10.0, "track_strike": 10.0, "distance_strike": 8.0, "condition_strike": 8.0, "barrier_draw": 5.0, "career_prize": 5.0, }, notes=( "Open-source multi-factor trifecta scorer. No supervised training " "(no historical race results available in data/results/). Scores are " "deterministic given runner form/stat inputs. Future v2 will train a " "gradient-boosted model on observed outcomes." ), ) def save_artifact(out_dir: str | Path) -> Path: out_dir = Path(out_dir) out_dir.mkdir(parents=True, exist_ok=True) path = out_dir / "model_artifact.json" path.write_text(json.dumps(asdict(ARTIFACT), indent=2)) return path def load_artifact(path: str | Path) -> ModelArtifact: data = json.loads(Path(path).read_text()) return ModelArtifact(**data) def predict_race(race: Race) -> dict[str, Any]: """Convenience: run the packaged predictor on a Race.""" return TrifectaPredictor().predict(race) def race_from_payload(payload: dict[str, Any]) -> Race: """Build a Race from the Trifecta-Bro predictions JSON shape.""" form = payload.get("form", {}) runners = [ Runner( number=r.get("number", 0), name=r.get("name", ""), jockey=r.get("jockey", ""), trainer=r.get("trainer", ""), weight=r.get("weight"), barrier=r.get("barrier"), form=r.get("form", ""), last20Starts=r.get("last20Starts", ""), careerPrizeMoney=r.get("careerPrizeMoney", "$0"), scratched=bool(r.get("scratched", False)), stats=r.get("stats", {}), ) for r in form.get("runners", []) ] return Race( date=payload.get("date", ""), track=payload.get("track", ""), track_slug=payload.get("track_slug", ""), race_number=str(payload.get("race_number", "")), race_name=payload.get("race_name", ""), distance=payload.get("distance", ""), condition=payload.get("condition", ""), weather=payload.get("weather", ""), race_class=payload.get("race_class", ""), start_time=payload.get("start_time", ""), prize_money=payload.get("prize_money", ""), number_of_runners=int(payload.get("number_of_runners", 0) or 0), runners=runners, )