from __future__ import annotations from datetime import datetime import numpy as np import pandas as pd from .types import DiagnosticResult, OptionSnapshot def _validate_calls_df(calls_df: pd.DataFrame) -> pd.DataFrame: required = {"expiry", "strike", "mid"} if not required.issubset(calls_df.columns): missing = sorted(required - set(calls_df.columns)) raise ValueError(f"Missing required columns: {missing}") return calls_df.sort_values(["expiry", "strike"]).reset_index(drop=True) def check_monotonicity(calls_df: pd.DataFrame, tol: float = 1e-9) -> DiagnosticResult: df = _validate_calls_df(calls_df) violations = 0 comparisons = 0 for _, grp in df.groupby("expiry"): mids = grp["mid"].to_numpy(dtype=float) diffs = np.diff(mids) comparisons += int(len(diffs)) violations += int(np.sum(diffs > tol)) violation_rate = float(violations / comparisons) if comparisons else 0.0 return DiagnosticResult( name="monotonicity", passed=violations == 0, violations=violations, comparisons=comparisons, violation_rate=violation_rate, details="Call prices should be non-increasing in strike.", ) def check_convexity(calls_df: pd.DataFrame, tol: float = 1e-8) -> DiagnosticResult: df = _validate_calls_df(calls_df) violations = 0 comparisons = 0 for _, grp in df.groupby("expiry"): g = grp.sort_values("strike") k = g["strike"].to_numpy(dtype=float) c = g["mid"].to_numpy(dtype=float) if len(k) < 3: continue dk = np.diff(k) valid = dk > 0 if not np.all(valid): keep = np.concatenate(([True], valid)) k = k[keep] c = c[keep] if len(k) < 3: continue dk = np.diff(k) slope_left = (c[1:-1] - c[:-2]) / (k[1:-1] - k[:-2]) slope_right = (c[2:] - c[1:-1]) / (k[2:] - k[1:-1]) comparisons += int(len(slope_right)) violations += int(np.sum((slope_right - slope_left) < -tol)) violation_rate = float(violations / comparisons) if comparisons else 0.0 return DiagnosticResult( name="convexity", passed=violations == 0, violations=violations, comparisons=comparisons, violation_rate=violation_rate, details="Call strike slopes should be non-decreasing (uneven-grid convexity).", ) def check_calendar(calls_df: pd.DataFrame, tol: float = 1e-8) -> DiagnosticResult: df = _validate_calls_df(calls_df) strikes = sorted(set(df["strike"].tolist())) violations = 0 comparisons = 0 for strike in strikes: s = df[df["strike"] == strike].sort_values("expiry") if len(s) < 2: continue mids = s["mid"].to_numpy(dtype=float) diffs = np.diff(mids) comparisons += int(len(diffs)) violations += int(np.sum(diffs < -tol)) violation_rate = float(violations / comparisons) if comparisons else 0.0 return DiagnosticResult( name="calendar", passed=violations == 0, violations=violations, comparisons=comparisons, violation_rate=violation_rate, details="Call prices should be non-decreasing with maturity at same strike.", ) def run_all_checks(snapshot: OptionSnapshot) -> list[DiagnosticResult]: calls = snapshot.options[snapshot.options["option_type"] == "call"].copy() return [ check_monotonicity(calls), check_convexity(calls), check_calendar(calls), ] def run_checks_by_expiry(snapshot: OptionSnapshot) -> pd.DataFrame: calls = snapshot.options[snapshot.options["option_type"] == "call"].copy() rows: list[dict[str, object]] = [] for expiry, grp in calls.groupby("expiry"): mini = OptionSnapshot( ticker=snapshot.ticker, snapshot_time=snapshot.snapshot_time, spot=snapshot.spot, options=grp.copy(), ) for d in run_all_checks(mini): rows.append( { "expiry": pd.to_datetime(expiry), "check": d.name, "passed": d.passed, "violations": d.violations, "comparisons": d.comparisons, "violation_rate": d.violation_rate, "details": d.details, } ) return pd.DataFrame(rows).sort_values(["expiry", "check"]).reset_index(drop=True) def select_best_quality_expiry(snapshot: OptionSnapshot) -> datetime: per_expiry = run_checks_by_expiry(snapshot) if per_expiry.empty: raise ValueError("No per-expiry diagnostics available") scores = ( per_expiry.groupby("expiry")["violation_rate"] .mean() .sort_values(ascending=True) ) return pd.to_datetime(scores.index[0]).to_pydatetime()