#!/usr/bin/env python3
"""Validate merger-model-builder plan.json inputs.

This validator protects the input contract. It does not compute the model.
"""

from __future__ import annotations

import json
import sys
from pathlib import Path
from typing import Any

ALLOWED_EVIDENCE_LABELS = {
    "signed_agreement",
    "filed",
    "audited",
    "reviewed",
    "vdr",
    "management_case",
    "consensus",
    "financing_commitment",
    "accounting_memo",
    "tax_memo",
    "user_provided",
    "estimate",
    "assumption",
    "placeholder",
    "unsupported",
}
REQUIRED_SOURCE_CATEGORIES = {
    "financials",
    "share_count",
    "offer_terms",
    "financing",
    "synergies",
    "purchase_accounting",
}
LOW_CONFIDENCE_LABELS = {"estimate", "assumption", "placeholder", "unsupported"}
REQUIRED_TOP_LEVEL = [
    "meta",
    "source_basis",
    "periods",
    "acquirer",
    "target",
    "transaction",
    "consideration",
    "financing",
    "purchase_accounting",
    "synergies",
    "scenarios",
    "sensitivities",
]
REQUIRED_SCENARIOS = ["base", "downside", "upside"]


def load_json(path: Path) -> dict[str, Any]:
    try:
        return json.loads(path.read_text())
    except Exception as exc:
        raise ValueError(f"could not parse JSON: {exc}") from exc


def is_num(x: Any) -> bool:
    return isinstance(x, (int, float)) and not isinstance(x, bool)


def has_num(obj: dict[str, Any], key: str) -> bool:
    return key in obj and is_num(obj[key])


def check_rate(
    errors: list[str], obj: dict[str, Any], path: str, key: str, required: bool = True
) -> None:
    if key not in obj or obj.get(key) is None:
        if required:
            errors.append(f"{path}.{key} is required")
        return
    if not is_num(obj[key]) or not (0 <= float(obj[key]) <= 1):
        errors.append(f"{path}.{key} must be a number between 0 and 1")


def check_nonnegative(
    errors: list[str], obj: dict[str, Any], path: str, key: str, required: bool = True
) -> None:
    if key not in obj or obj.get(key) is None:
        if required:
            errors.append(f"{path}.{key} is required")
        return
    if not is_num(obj[key]) or float(obj[key]) < 0:
        errors.append(f"{path}.{key} must be a non-negative number")


def check_positive(
    errors: list[str], obj: dict[str, Any], path: str, key: str, required: bool = True
) -> None:
    if key not in obj or obj.get(key) is None:
        if required:
            errors.append(f"{path}.{key} is required")
        return
    if not is_num(obj[key]) or float(obj[key]) <= 0:
        errors.append(f"{path}.{key} must be a positive number")


def check_period_map(
    errors: list[str], obj: dict[str, Any], path: str, periods: list[str], required: bool = True
) -> None:
    if obj is None:
        if required:
            errors.append(f"{path} is required and must map every period to a number")
        return
    if not isinstance(obj, dict):
        errors.append(f"{path} must be an object keyed by period")
        return
    for p in periods:
        if p not in obj:
            errors.append(f"{path}.{p} is missing")
        elif not is_num(obj[p]):
            errors.append(f"{path}.{p} must be numeric")


def validate(plan: dict[str, Any]) -> tuple[list[str], list[str]]:
    errors: list[str] = []
    warnings: list[str] = []

    for key in REQUIRED_TOP_LEVEL:
        if key not in plan:
            errors.append(f"top-level key '{key}' is required")

    meta = plan.get("meta", {}) if isinstance(plan.get("meta", {}), dict) else {}
    for key in ["deal_name", "acquirer", "target", "currency", "units", "accounting_basis"]:
        if not meta.get(key):
            errors.append(f"meta.{key} is required")
    if not (meta.get("valuation_date") or meta.get("announcement_date")):
        errors.append("meta.valuation_date or meta.announcement_date is required")
    if meta.get("accounting_basis") not in {"us_gaap", "ifrs", "local_gaap", "unknown", None}:
        errors.append("meta.accounting_basis must be one of us_gaap, ifrs, local_gaap, unknown")

    periods = plan.get("periods", [])
    if not isinstance(periods, list) or not periods:
        errors.append("periods must be a non-empty list")
        periods = []
    else:
        periods = [str(p) for p in periods]

    # Source basis.
    source_basis = plan.get("source_basis", [])
    if not isinstance(source_basis, list) or not source_basis:
        errors.append("source_basis must be a non-empty list")
        source_basis = []
    present_categories = set()
    for i, src in enumerate(source_basis):
        if not isinstance(src, dict):
            errors.append(f"source_basis[{i}] must be an object")
            continue
        for key in ["id", "label", "description", "date", "categories"]:
            if key not in src or src.get(key) in (None, ""):
                errors.append(f"source_basis[{i}].{key} is required")
        label = src.get("label")
        if label not in ALLOWED_EVIDENCE_LABELS:
            errors.append(f"source_basis[{i}].label '{label}' is not allowed")
        if label in LOW_CONFIDENCE_LABELS:
            warnings.append(f"source_basis[{i}] uses low-confidence evidence label '{label}'")
        categories = src.get("categories", [])
        if not isinstance(categories, list) or not categories:
            errors.append(f"source_basis[{i}].categories must be a non-empty list")
        else:
            present_categories.update(str(c) for c in categories)
    missing_categories = sorted(REQUIRED_SOURCE_CATEGORIES - present_categories)
    if missing_categories:
        errors.append(
            "source_basis is missing required categories: " + ", ".join(missing_categories)
        )

    acq = plan.get("acquirer", {}) if isinstance(plan.get("acquirer", {}), dict) else {}
    check_positive(errors, acq, "acquirer", "share_price")
    check_positive(errors, acq, "acquirer", "diluted_shares")
    check_nonnegative(errors, acq, "acquirer", "cash")
    check_nonnegative(errors, acq, "acquirer", "debt")
    check_rate(errors, acq, "acquirer", "tax_rate")
    if periods:
        check_period_map(errors, acq.get("net_income"), "acquirer.net_income", periods)
        if "eps" in acq:
            check_period_map(errors, acq.get("eps"), "acquirer.eps", periods, required=False)

    tgt = plan.get("target", {}) if isinstance(plan.get("target", {}), dict) else {}
    check_positive(errors, tgt, "target", "diluted_shares")
    if plan.get("transaction", {}).get("equity_purchase_price") is None:
        check_positive(errors, tgt, "target", "offer_price")
    check_nonnegative(errors, tgt, "target", "cash")
    check_nonnegative(errors, tgt, "target", "debt")
    if not has_num(tgt, "book_equity"):
        errors.append("target.book_equity is required and must be numeric")
    check_rate(errors, tgt, "target", "tax_rate", required=False)
    if periods:
        check_period_map(errors, tgt.get("net_income"), "target.net_income", periods)
        check_period_map(errors, tgt.get("standalone_ebitda"), "target.standalone_ebitda", periods)

    tx = plan.get("transaction", {}) if isinstance(plan.get("transaction", {}), dict) else {}
    if tx.get("equity_purchase_price") is not None:
        check_positive(errors, tx, "transaction", "equity_purchase_price")
    else:
        if tx.get("offer_price") is not None:
            check_positive(errors, tx, "transaction", "offer_price")
    check_nonnegative(errors, tx, "transaction", "required_min_cash")
    check_rate(errors, tx, "transaction", "tax_rate")
    fees = tx.get("fees", {}) if isinstance(tx.get("fees", {}), dict) else {}
    for key in ["transaction_fees", "financing_fees", "equity_issuance_fees"]:
        check_nonnegative(errors, fees, "transaction.fees", key)
    for key, val in fees.items():
        if val is not None and (not is_num(val) or float(val) < 0):
            errors.append(f"transaction.fees.{key} must be non-negative")

    cons = plan.get("consideration", {}) if isinstance(plan.get("consideration", {}), dict) else {}
    for key in ["cash_percent", "stock_percent", "other_percent"]:
        check_rate(errors, cons, "consideration", key)
    if all(
        key in cons and is_num(cons[key])
        for key in ["cash_percent", "stock_percent", "other_percent"]
    ):
        mix_sum = (
            float(cons["cash_percent"])
            + float(cons["stock_percent"])
            + float(cons["other_percent"])
        )
        if abs(mix_sum - 1.0) > 0.0001 and not tx.get("allow_unbalanced_consideration_mix"):
            errors.append(f"consideration mix must sum to 1.0; got {mix_sum:.6f}")
        if float(cons["stock_percent"]) > 0 and not (
            has_num(acq, "share_price") or cons.get("exchange_ratio") is not None
        ):
            errors.append(
                "stock consideration requires acquirer.share_price or consideration.exchange_ratio"
            )
    if cons.get("exchange_ratio") is not None and (
        not is_num(cons.get("exchange_ratio")) or float(cons.get("exchange_ratio")) <= 0
    ):
        errors.append("consideration.exchange_ratio must be positive when provided")

    fin = plan.get("financing", {}) if isinstance(plan.get("financing", {}), dict) else {}
    check_nonnegative(errors, fin, "financing", "new_debt")
    check_nonnegative(errors, fin, "financing", "cash_used")
    check_rate(errors, fin, "financing", "debt_interest_rate")
    check_rate(errors, fin, "financing", "lost_cash_interest_rate")
    check_positive(errors, fin, "financing", "fee_amortization_years")
    if "use_target_cash" not in fin:
        errors.append("financing.use_target_cash is required")

    ppa = (
        plan.get("purchase_accounting", {})
        if isinstance(plan.get("purchase_accounting", {}), dict)
        else {}
    )
    for key in [
        "target_book_equity",
        "existing_goodwill",
        "ppe_step_up",
        "inventory_step_up",
        "nci_fair_value",
        "previously_held_interest_fair_value",
        "contingent_consideration_fair_value",
    ]:
        if key == "target_book_equity":
            if not has_num(ppa, key):
                errors.append(f"purchase_accounting.{key} is required and must be numeric")
        else:
            check_nonnegative(errors, ppa, "purchase_accounting", key)
    if float(ppa.get("ppe_step_up", 0) or 0) > 0:
        check_positive(errors, ppa, "purchase_accounting", "ppe_step_up_life")
    check_rate(errors, ppa, "purchase_accounting", "deferred_tax_rate")
    if ppa.get("deferred_tax_liability") is not None:
        check_nonnegative(errors, ppa, "purchase_accounting", "deferred_tax_liability")
    intangibles = ppa.get("intangible_assets", [])
    if not isinstance(intangibles, list):
        errors.append("purchase_accounting.intangible_assets must be a list")
    else:
        for i, asset in enumerate(intangibles):
            if not isinstance(asset, dict):
                errors.append(f"purchase_accounting.intangible_assets[{i}] must be an object")
                continue
            if not asset.get("name"):
                errors.append(f"purchase_accounting.intangible_assets[{i}].name is required")
            check_nonnegative(
                errors, asset, f"purchase_accounting.intangible_assets[{i}]", "fair_value"
            )
            if float(asset.get("fair_value", 0) or 0) > 0:
                check_positive(
                    errors,
                    asset,
                    f"purchase_accounting.intangible_assets[{i}]",
                    "amortization_years",
                )

    syn = plan.get("synergies", {}) if isinstance(plan.get("synergies", {}), dict) else {}
    if periods:
        for key in ["cost_synergies", "revenue_synergies", "dis_synergies", "integration_costs"]:
            check_period_map(errors, syn.get(key), f"synergies.{key}", periods)
    check_rate(errors, syn, "synergies", "revenue_synergy_margin")
    check_rate(errors, syn, "synergies", "tax_rate")
    if syn.get("realization_basis") not in {
        "immediate",
        "phased",
        "run_rate",
        "probability_weighted",
        "unknown",
    }:
        errors.append(
            "synergies.realization_basis must be one of immediate, phased, run_rate, probability_weighted, unknown"
        )
    if "integration_costs_excluded_from_adjusted_eps" not in syn:
        errors.append("synergies.integration_costs_excluded_from_adjusted_eps is required")

    scenarios = plan.get("scenarios", {}) if isinstance(plan.get("scenarios", {}), dict) else {}
    base_rate = float(fin.get("debt_interest_rate", 0) or 0)
    for name in REQUIRED_SCENARIOS:
        if name not in scenarios:
            errors.append(f"scenarios.{name} is required")
            continue
        sc = scenarios[name]
        if not isinstance(sc, dict):
            errors.append(f"scenarios.{name} must be an object")
            continue
        for key in [
            "acquirer_net_income_factor",
            "target_net_income_factor",
            "target_ebitda_factor",
            "synergy_factor",
            "dis_synergy_factor",
            "integration_cost_factor",
            "purchase_price_factor",
            "share_price_factor",
        ]:
            if key not in sc or not is_num(sc[key]) or float(sc[key]) < 0:
                errors.append(f"scenarios.{name}.{key} must be a non-negative number")
        for key in ["debt_rate_delta", "tax_rate_delta"]:
            if key not in sc or not is_num(sc[key]):
                errors.append(f"scenarios.{name}.{key} must be numeric")
        if is_num(sc.get("debt_rate_delta")) and base_rate + float(sc.get("debt_rate_delta")) < 0:
            errors.append(f"scenarios.{name}.debt_rate_delta drives debt interest rate below zero")
        if sc.get("cash_percent_override") is not None:
            if not is_num(sc.get("cash_percent_override")) or not (
                0 <= float(sc.get("cash_percent_override")) <= 1
            ):
                errors.append(
                    f"scenarios.{name}.cash_percent_override must be null or a rate between 0 and 1"
                )

    sens = plan.get("sensitivities", {}) if isinstance(plan.get("sensitivities", {}), dict) else {}
    for key in [
        "synergy_factors",
        "debt_rate_deltas",
        "cash_percentages",
        "premium_factors",
        "tax_rates",
        "share_price_factors",
    ]:
        arr = sens.get(key)
        if not isinstance(arr, list) or not arr:
            errors.append(f"sensitivities.{key} must be a non-empty numeric list")
        else:
            for i, val in enumerate(arr):
                if not is_num(val):
                    errors.append(f"sensitivities.{key}[{i}] must be numeric")
                elif key in {"cash_percentages", "tax_rates"} and not (0 <= float(val) <= 1):
                    errors.append(f"sensitivities.{key}[{i}] must be between 0 and 1")
                elif (
                    key in {"synergy_factors", "premium_factors", "share_price_factors"}
                    and float(val) < 0
                ):
                    errors.append(f"sensitivities.{key}[{i}] must be non-negative")
                elif key == "debt_rate_deltas" and base_rate + float(val) < 0:
                    errors.append(f"sensitivities.{key}[{i}] drives debt interest rate below zero")

    return errors, warnings


def main(argv: list[str]) -> int:
    if len(argv) != 2:
        print("Usage: python3 scripts/validate_plan.py path/to/plan.json", file=sys.stderr)
        return 1
    path = Path(argv[1])
    if not path.exists():
        print(f"INVALID PLAN\n- file does not exist: {path}", file=sys.stderr)
        return 1
    try:
        plan = load_json(path)
        errors, warnings = validate(plan)
    except Exception as exc:
        print(f"INVALID PLAN\n- {exc}", file=sys.stderr)
        return 1
    if errors:
        print("INVALID PLAN")
        for err in errors:
            print(f"- {err}")
        if warnings:
            print("WARNINGS")
            for warn in warnings:
                print(f"- {warn}")
        return 1
    print("VALID PLAN")
    if warnings:
        print("WARNINGS")
        for warn in warnings:
            print(f"- {warn}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main(sys.argv))
