#!/usr/bin/env python3
"""Exact population example of a threshold policy that changes its target."""

import argparse
import json
from dataclasses import dataclass
from decimal import Decimal, InvalidOperation
from fractions import Fraction
from pathlib import Path
from typing import TypeAlias


Number: TypeAlias = int | float | str | Decimal | Fraction


def exact_rational(value: Number) -> Fraction:
    """Interpret a finite decimal input as an exact rational number."""
    if isinstance(value, Fraction):
        return value
    try:
        decimal = value if isinstance(value, Decimal) else Decimal(str(value))
    except (InvalidOperation, ValueError) as error:
        raise ValueError("Numeric inputs must be finite decimals") from error
    if not decimal.is_finite():
        raise ValueError("Numeric inputs must be finite decimals")
    return Fraction(decimal)


@dataclass(frozen=True)
class Parameters:
    baseline: Number
    effect: Number
    threshold: Number
    preventive_cost: Number = 0.0

    def __post_init__(self) -> None:
        baseline = exact_rational(self.baseline)
        effect = exact_rational(self.effect)
        low = baseline - effect
        if not Fraction(0) <= low <= baseline <= Fraction(1):
            raise ValueError("Require 0 <= baseline - effect <= baseline <= 1")
        if exact_rational(self.preventive_cost) < 0:
            raise ValueError("Preventive cost must be nonnegative")


def clean(value: Number) -> float:
    """Convert an exact internal rational to a JSON number."""
    return float(exact_rational(value))


def action_from_rate(rate: Number, threshold: Number) -> int:
    """The naive controller prevents on equality."""
    return int(exact_rational(rate) >= exact_rational(threshold))


def _defect_rational(parameters: Parameters, action: int) -> Fraction:
    if action not in (0, 1):
        raise ValueError("Action must be 0 or 1")
    return exact_rational(parameters.baseline) - exact_rational(parameters.effect) * action


def _loss_rational(parameters: Parameters, action: int) -> Fraction:
    return _defect_rational(parameters, action) + (
        exact_rational(parameters.preventive_cost) * action
    )


def defect_probability(parameters: Parameters, action: int) -> float:
    return clean(_defect_rational(parameters, action))


def loss(parameters: Parameters, action: int) -> float:
    return clean(_loss_rational(parameters, action))


def trajectory(
    parameters: Parameters, initial_rate: float, rounds: int
) -> list[dict]:
    """Retrain to the exact on-policy population rate after each round."""
    if rounds < 0:
        raise ValueError("Rounds must be nonnegative")
    rate = exact_rational(initial_rate)
    records = []
    for round_index in range(rounds):
        action = action_from_rate(rate, parameters.threshold)
        probability = _defect_rational(parameters, action)
        records.append(
            {
                "round": round_index,
                "q": clean(rate),
                "action": action,
                "defect_probability": clean(probability),
                "loss": clean(_loss_rational(parameters, action)),
            }
        )
        rate = probability
    return records


def phase(parameters: Parameters, initial_rate: float) -> dict:
    """Classify the eventual exact dynamics, including threshold ties."""
    low = _defect_rational(parameters, 1)
    high = _defect_rational(parameters, 0)
    threshold = exact_rational(parameters.threshold)
    if threshold <= low:
        behavior = "eventually_fixed_action_1"
        period = 1
    elif threshold <= high:
        behavior = "two_cycle"
        period = 2
    else:
        behavior = "eventually_fixed_action_0"
        period = 1
    return {
        "low_rate_under_action_1": clean(low),
        "high_rate_under_action_0": clean(high),
        "limit_cycle_condition": "low_rate < threshold <= high_rate",
        "tie_rule": "action 1 when q >= threshold",
        "initial_rate": clean(initial_rate),
        "initial_action": action_from_rate(initial_rate, parameters.threshold),
        "eventual_behavior": behavior,
        "period": period,
    }


def long_run_naive(parameters: Parameters) -> dict:
    """Return the limiting average for the classified deterministic phase."""
    behavior = phase(parameters, parameters.baseline)["eventual_behavior"]
    if behavior == "two_cycle":
        action_rate = Fraction(1, 2)
        defects = (_defect_rational(parameters, 0) + _defect_rational(parameters, 1)) / 2
    elif behavior == "eventually_fixed_action_1":
        action_rate = Fraction(1)
        defects = _defect_rational(parameters, 1)
    else:
        action_rate = Fraction(0)
        defects = _defect_rational(parameters, 0)
    return {
        "eventual_behavior": behavior,
        "action_1_rate": clean(action_rate),
        "average_defect_probability": clean(defects),
        "average_loss": clean(
            defects + exact_rational(parameters.preventive_cost) * action_rate
        ),
    }


def oracle_conditional_policy(parameters: Parameters) -> dict:
    """Choose using both action-specific risks; favor prevention on a tie."""
    action = int(
        exact_rational(parameters.preventive_cost) <= exact_rational(parameters.effect)
    )
    return {
        "privilege": "oracle knows both action-specific defect probabilities",
        "tie_rule": "choose action 1 when its loss equals action 0 loss",
        "action": action,
        "defect_probability": clean(defect_probability(parameters, action)),
        "loss": clean(loss(parameters, action)),
    }


def cost_comparison(cost: float) -> dict:
    parameters = Parameters(0.8, 0.6, 0.5, cost)
    return {
        "preventive_cost": clean(cost),
        "always_action_0": {
            "defect_probability": clean(defect_probability(parameters, 0)),
            "loss": clean(loss(parameters, 0)),
        },
        "always_action_1": {
            "defect_probability": clean(defect_probability(parameters, 1)),
            "loss": clean(loss(parameters, 1)),
        },
        "naive_threshold_long_run": long_run_naive(parameters),
        "oracle_conditional_policy": oracle_conditional_policy(parameters),
    }


def phase_case(
    name: str, baseline: float, effect: float, threshold: float, initial_rate: float
) -> dict:
    parameters = Parameters(baseline, effect, threshold)
    return {
        "name": name,
        "parameters": {
            "baseline": clean(baseline),
            "effect": clean(effect),
            "threshold": clean(threshold),
            "initial_rate": clean(initial_rate),
        },
        "phase": phase(parameters, initial_rate),
        "trajectory": trajectory(parameters, initial_rate, 6),
    }


def affine_probability(intercept: float, slope: float, action: int) -> float:
    return clean(exact_rational(intercept) + exact_rational(slope) * action)


def observational_equivalence() -> dict:
    """Two mechanisms agreeing under action 1 but disagreeing under action 0."""
    worlds = (
        ("no_action_effect", 0.8, 0.0),
        ("positive_action_response", 0.2, 0.6),
    )
    return {
        "deployed_action": 1,
        "observed_on_policy_probability": 0.8,
        "worlds": [
            {
                "name": name,
                "mechanism": f"p(a) = {intercept} + {slope}a",
                "p_action_0": clean(affine_probability(intercept, slope, 0)),
                "p_action_1": clean(affine_probability(intercept, slope, 1)),
            }
            for name, intercept, slope in worlds
        ],
        "boundary": (
            "Observing only action 1 cannot distinguish these mechanisms; "
            "the positive-response example is distinct from the preventive factory model."
        ),
    }


def results() -> dict:
    baseline = Parameters(0.8, 0.6, 0.5)
    return {
        "schema_version": 1,
        "status": "exact population illustration, not empirical validation",
        "primary_anchor": "https://proceedings.mlr.press/v119/perdomo20a.html",
        "assumptions": [
            "Independent factory batches with exact population rates; no sampling noise",
            "Two actions: 0 is no prevention and 1 is prevention",
            "Fixed response p(a) = baseline - effect * action with no carryover",
            "Boundary decisions and trajectories use exact decimal-rational arithmetic",
            "q is the last observed on-policy rate, not untreated risk",
            "The naive threshold controller ignores preventive cost",
            (
                "The oracle comparison knows both action-specific risks and is not "
                "a learning benchmark"
            ),
        ],
        "baseline_parameters": {
            "baseline": 0.8,
            "effect": 0.6,
            "threshold": 0.5,
            "initial_rate": 0.8,
        },
        "baseline_phase": phase(baseline, 0.8),
        "baseline_trajectory": trajectory(baseline, 0.8, 8),
        "cost_comparisons": [cost_comparison(cost) for cost in (0.0, 0.2, 0.6, 0.8)],
        "phase_cases": [
            phase_case("cycle_start_high", 0.8, 0.6, 0.5, 0.8),
            phase_case("cycle_start_low", 0.8, 0.6, 0.5, 0.2),
            phase_case("initial_tie", 0.8, 0.6, 0.5, 0.5),
            phase_case("threshold_at_low_rate", 0.8, 0.6, 0.2, 0.2),
            phase_case("threshold_at_high_rate", 0.8, 0.6, 0.8, 0.8),
            phase_case("threshold_below_interval", 0.8, 0.6, 0.1, 0.8),
            phase_case("threshold_above_interval", 0.8, 0.6, 0.9, 0.8),
            phase_case("no_action_effect", 0.8, 0.0, 0.5, 0.8),
        ],
        "observational_equivalence": observational_equivalence(),
    }


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path, help="Optional JSON output path")
    arguments = parser.parse_args()
    rendered = json.dumps(results(), indent=2) + "\n"
    if arguments.output is None:
        print(rendered, end="")
    else:
        arguments.output.parent.mkdir(parents=True, exist_ok=True)
        arguments.output.write_text(rendered, encoding="utf-8")


if __name__ == "__main__":
    main()
