#!/usr/bin/env python3
"""Reproducible hypothetical routing budget; no network or model calls.

Python 3.10+, standard library only. Own material: CC BY 4.0.
Default: calculate JSON on stdout. --verify checks published artifacts.
--write-output is an editorial helper; it does not update the frozen manifest.
"""

import argparse
from collections import Counter
import csv
from decimal import Decimal, localcontext
import hashlib
import io
import json
from pathlib import Path


HERE = Path(__file__).resolve().parent
PREFIX = "decision-1-"


def read_json(path):
    return json.loads(path.read_text(encoding="utf-8"))


def decimal_text(value):
    return format(value, "f")


def calculate(assumptions):
    def number(key):
        return Decimal(assumptions[key])

    n = number("cases_per_month")
    wage = number("labor_pln_per_hour")
    tokens = number("input_tokens_per_call")
    extra = number("extra_calls_fraction")
    price = number("input_usd_per_million_tokens")
    exchange = number("hypothetical_pln_per_usd")
    manual_seconds = number("manual_seconds_per_case_including_corrections")
    review_seconds = number("ai_review_seconds_per_case")
    correction_seconds = number("correction_seconds_per_wrong_label")
    fixed = number("infrastructure_and_maintenance_pln_per_month")
    setup = number("setup_pln")
    months = number("setup_budget_months")
    completed_fraction = number("correctly_routed_fraction_after_human_work")
    if n <= 0 or n != n.to_integral_value() or wage <= 0 or months <= 0:
        raise ValueError("Cases must be a positive integer; wage and budget months must be positive.")
    if not 0 < completed_fraction <= 1:
        raise ValueError("Correctly routed fraction must be above zero and at most one.")
    if any(value < 0 for value in (tokens, extra, price, manual_seconds,
                                  review_seconds, correction_seconds, fixed, setup)) or exchange <= 0:
        raise ValueError("Costs, token counts, times and extra-call fraction cannot be negative; exchange must be positive.")

    calls = n * (1 + extra)
    input_tokens = calls * tokens
    usd = input_tokens * price / 1_000_000
    api = usd * exchange
    manual = n * manual_seconds * wage / 3600
    review = n * review_seconds * wage / 3600
    setup_monthly = setup / months
    correct_cases = n * completed_fraction
    scenarios = []
    for fraction_text in assumptions["wrong_label_fractions_before_review"]:
        fraction = Decimal(fraction_text)
        if not 0 <= fraction <= 1:
            raise ValueError("Wrong-label fractions must be in [0, 1].")
        corrections = n * fraction * correction_seconds * wage / 3600
        without_review = api + corrections + fixed + setup_monthly
        ai_total = without_review + review
        seconds_limit = (manual - without_review) * 3600 / (n * wage)
        first_month_cash = ai_total - setup_monthly + setup
        scenarios.append({
            "wrong_label_fraction_before_review": fraction,
            "wrong_labels": n * fraction,
            "api_pln": api,
            "human_review_pln": review,
            "human_correction_pln": corrections,
            "infrastructure_maintenance_pln": fixed,
            "setup_monthly_budget_pln": setup_monthly,
            "ai_monthly_budget_pln": ai_total,
            "ai_pln_per_correctly_routed_case": ai_total / correct_cases,
            "manual_monthly_pln": manual,
            "manual_pln_per_correctly_routed_case": manual / correct_cases,
            "monthly_saving_pln": manual - ai_total,
            "break_even_review_seconds_per_case": seconds_limit,
            "first_month_ai_cash_pln_if_setup_paid_upfront": first_month_cash,
        })
    return {
        "status": "hypothetical_calculation_no_model_calls",
        "case_count": n,
        "call_count_including_retries": calls,
        "input_tokens": input_tokens,
        "api_usd": usd,
        "hypothetical_api_pln": api,
        "correctly_routed_cases_assumed": correct_cases,
        "scenarios": scenarios,
    }


def json_text(data):
    return json.dumps(data, ensure_ascii=False, indent=2, default=decimal_text) + "\n"


def csv_text(result):
    buffer = io.StringIO(newline="")
    scenarios = result["scenarios"]
    writer = csv.DictWriter(buffer, fieldnames=list(scenarios[0]), lineterminator="\n")
    writer.writeheader()
    for row in scenarios:
        writer.writerow({key: decimal_text(value) for key, value in row.items()})
    return buffer.getvalue()


def input_text(corpus, labels):
    return "".join(json.dumps({
        "id": row["id"],
        "request": {
            "state": row["tekst"],
            "questions": {
                "kolejka": {
                    "type": "choice",
                    "instructions": labels["instructions"],
                    "criteria": labels["criteria"],
                }
            },
        },
    }, ensure_ascii=False) + "\n" for row in corpus)


def corpus_rows():
    with (HERE / (PREFIX + "reklamacje.csv")).open(encoding="utf-8", newline="") as file:
        return list(csv.DictReader(file))


def require(condition, message):
    if not condition:
        raise ValueError(message)


def verify(result):
    require(result["api_usd"] == Decimal("0.6615"), "Published API USD no longer matches.")
    require(result["call_count_including_retries"] == 10500, "Retry denominator changed.")
    require(result["input_tokens"] == 15750000, "Published token assumption changed.")
    require(result["correctly_routed_cases_assumed"] == 10000, "Outcome denominator changed.")
    base = result["scenarios"][0]
    require(base["manual_monthly_pln"] == 5000, "Manual baseline changed.")
    require(base["human_correction_pln"] == 600, "Base correction cost changed.")
    require(base["break_even_review_seconds_per_case"] == Decimal("17.384124"), "Break-even threshold changed.")
    # Independent unit check: one second of review costs N * hourly wage / 3600.
    for row in result["scenarios"]:
        review_at_limit = row["break_even_review_seconds_per_case"] * Decimal(10000) / 60
        non_review = sum(row[key] for key in (
            "api_pln", "human_correction_pln", "infrastructure_maintenance_pln", "setup_monthly_budget_pln"))
        require(abs(non_review + review_at_limit - 5000) < Decimal("0.000000000001"), "Break-even identity failed.")
    require((HERE / (PREFIX + "wynik.json")).read_text(encoding="utf-8") == json_text(result), "Saved JSON differs from calculation.")
    require((HERE / (PREFIX + "koszty.csv")).read_text(encoding="utf-8") == csv_text(result), "Saved CSV differs from calculation.")
    labels = read_json(HERE / (PREFIX + "etykiety.json"))
    corpus = corpus_rows()
    counts = Counter(row["oczekiwana_etykieta"] for row in corpus)
    require(len(corpus) == 30 and len({row["id"] for row in corpus}) == 30, "Corpus must contain 30 unique IDs.")
    require(counts == Counter({label: 6 for label in labels["criteria"]}), "Expected six examples per label.")
    require(all(row["tekst"] and row["uzasadnienie"] for row in corpus), "Empty case or rationale.")
    require((HERE / (PREFIX + "wejscia.jsonl")).read_text(encoding="utf-8") == input_text(corpus, labels), "Model inputs differ or contain extra fields.")
    manifest = read_json(HERE / (PREFIX + "manifest.json"))
    expected_files = {PREFIX + name for name in (
        "kalkulacja.py", "zalozenia.json", "koszty.csv", "wynik.json", "etykiety.json",
        "reklamacje.csv", "wejscia.jsonl", "protokol.md")}
    require(set(manifest["sha256"]) == expected_files, "Manifest file inventory changed.")
    for filename, expected in manifest["sha256"].items():
        actual = hashlib.sha256((HERE / filename).read_bytes()).hexdigest()
        require(actual == expected, "SHA-256 differs: " + filename)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--assumptions", type=Path, default=HERE / (PREFIX + "zalozenia.json"))
    parser.add_argument("--verify", action="store_true", help="Verify original published bundle only.")
    parser.add_argument("--write-output", action="store_true", help="Editorial: regenerate output files, never the manifest.")
    args = parser.parse_args()
    if args.verify and (args.write_output or args.assumptions.resolve() != (HERE / (PREFIX + "zalozenia.json")).resolve()):
        parser.error("--verify checks the original bundle and cannot be combined with custom assumptions or writing.")
    with localcontext() as context:
        context.prec = 40
        result = calculate(read_json(args.assumptions))
        if args.write_output:
            if args.assumptions.resolve() != (HERE / (PREFIX + "zalozenia.json")).resolve():
                parser.error("--write-output is restricted to original assumptions; use stdout for custom scenarios.")
            (HERE / (PREFIX + "wynik.json")).write_text(json_text(result), encoding="utf-8")
            (HERE / (PREFIX + "koszty.csv")).write_text(csv_text(result), encoding="utf-8")
            (HERE / (PREFIX + "wejscia.jsonl")).write_text(input_text(corpus_rows(), read_json(HERE / (PREFIX + "etykiety.json"))), encoding="utf-8")
        if args.verify:
            verify(result)
            print("OK: exact arithmetic, break-even identity, saved results, 30 cases / 5 labels, answer-free inputs and SHA-256 manifest. No model calls.")
        else:
            print(json_text(result), end="")


if __name__ == "__main__":
    main()
