#!/usr/bin/env python3
"""Mol2Mat research calculation: a published three-excited-state kinetic model.

Only Python's standard library is required. Run this file to write analysis.json.
No parameters below are new measurements made by Mol2Mat.

The source rates are rounded computational values from Shizu & Kaji (2024),
Table 1, DOI 10.1038/s41467-024-49069-4. Experimental observations are kept
separately in primary-observations.json. Optical excitation is initially S1;
electrical excitation's 1:3 S1:T1 split is an explicit idealized assumption.

dN/dt = -M N; integral_0^infty N dt = M^-1 N(0).
For a constant low-density generation rate G with the same state fractions,
N_steady/G equals this integrated residence vector. No time-stepping is used.
"""

from copy import deepcopy
from math import isfinite
from pathlib import Path
import json

SOURCE = "https://www.nature.com/articles/s41467-024-49069-4"

# [S1, T1, T2]. Every rate has units s^-1. The primary source gives these
# computational values to roughly two significant figures.
MATERIALS = json.loads((Path(__file__).resolve().parent / 'rates.json').read_text())


def solve(matrix, vector):
    """Small dense linear solve with partial pivoting; no external packages."""
    n = len(vector)
    augmented = [list(row) + [float(value)] for row, value in zip(matrix, vector)]
    for column in range(n):
        pivot = max(range(column, n), key=lambda row: abs(augmented[row][column]))
        if augmented[pivot][column] == 0:
            raise ValueError("Non-decaying population: finite residence does not exist")
        augmented[column], augmented[pivot] = augmented[pivot], augmented[column]
        divisor = augmented[column][column]
        augmented[column] = [value / divisor for value in augmented[column]]
        for row in range(n):
            if row == column:
                continue
            factor = augmented[row][column]
            augmented[row] = [a - factor*b for a, b in zip(augmented[row], augmented[column])]
    return [row[-1] for row in augmented]


def calculate(material, excitation="electrical", extra_triplet_loss=0.0,
              t2_return_enabled=True, outcoupling=0.25, charge_balance=1.0):
    """Return normalized single-exciton fate and residence; density is low.

    extra_triplet_loss is an independently imposed first-order sink applied
    equally to T1 and T2. It is NOT a fitted TTA constant, oxygen concentration,
    current density, temperature, or device luminance. T2-off removes only
    T2->S1; it is a mathematical channel-deletion experiment, not a molecule.
    """
    if excitation not in {"optical", "electrical"}:
        raise ValueError("excitation must be optical or electrical")
    if not isfinite(extra_triplet_loss) or extra_triplet_loss < 0:
        raise ValueError("extra_triplet_loss must be finite and non-negative")
    if not 0 <= outcoupling <= 1 or not 0 <= charge_balance <= 1:
        raise ValueError("outcoupling and charge_balance must be within [0,1]")
    rates = deepcopy(MATERIALS[material] if isinstance(material, str) else material)
    radiative, nonradiative, transfer = (rates[key] for key in ("radiative", "nonradiative", "transitions"))
    if any(not isfinite(x) or x < 0 for x in radiative + nonradiative + sum(transfer, [])):
        raise ValueError("Every rate must be finite and non-negative")
    if not t2_return_enabled:
        transfer[2][0] = 0.0
    added_loss = [0.0, extra_triplet_loss, extra_triplet_loss]
    matrix = [[0.0] * 3 for _ in range(3)]
    for origin in range(3):
        matrix[origin][origin] = radiative[origin] + nonradiative[origin] + added_loss[origin] + sum(transfer[origin])
        for destination in range(3):
            if destination != origin:
                matrix[destination][origin] = -transfer[origin][destination]
    initial = [1.0, 0.0, 0.0] if excitation == "optical" else [0.25, 0.75, 0.0]
    # Eliminate T2 analytically. Positive sums avoid catastrophic cancellation
    # between ~1e12 s^-1 forward/backward triplet rates and slow ground decay.
    loss = [radiative[i] + nonradiative[i] + added_loss[i] for i in range(3)]
    c = loss[2] + transfer[2][0] + transfer[2][1]
    if c <= 0:
        raise ValueError("T2 has no exit; finite residence does not exist")
    singlet_loss = loss[0] + transfer[0][2] * loss[2] / c
    triplet_loss = loss[1] + transfer[1][2] * loss[2] / c
    s_to_t = transfer[0][1] + transfer[0][2] * transfer[2][1] / c
    t_to_s = transfer[1][0] + transfer[1][2] * transfer[2][0] / c
    determinant = singlet_loss*triplet_loss + singlet_loss*t_to_s + s_to_t*triplet_loss
    if determinant <= 0:
        raise ValueError("Non-decaying population: finite residence does not exist")
    us = ((triplet_loss+t_to_s)*initial[0] + t_to_s*initial[1]) / determinant
    ut = (s_to_t*initial[0] + (singlet_loss+s_to_t)*initial[1]) / determinant
    residence = [us, ut, (transfer[0][2]*us + transfer[1][2]*ut) / c]
    yields = {
        "fluorescence": radiative[0] * residence[0],
        "phosphorescence": radiative[1] * residence[1] + radiative[2] * residence[2],
        "intrinsic_nonradiative_loss": sum(k * t for k, t in zip(nonradiative, residence)),
        "added_triplet_loss": sum(k * t for k, t in zip(added_loss, residence)),
    }
    photons = yields["fluorescence"] + yields["phosphorescence"]
    # A flux-integral can exceed one because the same excitation revisits states.
    return {
        "fate": yields,
        "probability_sum": sum(yields.values()),
        "relative_balance_residual": max(abs(sum(matrix[i][j]*residence[j] for j in range(3))-initial[i]) / max(1.0,sum(abs(matrix[i][j]*residence[j]) for j in range(3))) for i in range(3)),
        "radiative_probability": photons,
        "residence_s": {state: value for state, value in zip(("S1", "T1", "T2"), residence)},
        "triplet_residence_s": residence[1] + residence[2],
        "total_residence_s": sum(residence),
        "expected_risc_events": transfer[1][0]*residence[1] + transfer[2][0]*residence[2],
        "expected_t2_to_s1_events": transfer[2][0]*residence[2],
        "external_photons_per_injected_pair_scenario": photons * outcoupling * charge_balance if excitation == "electrical" else None,
        "scope": "Linear, low-density published-rate model; no bimolecular annihilation, charge transport, morphology, optical-stack calculation, or degradation.",
    }


def reference_residence(material, excitation, extra_triplet_loss, t2_return_enabled):
    """Independent 60-digit Gaussian solve of the original 3x3 balance matrix."""
    from decimal import Decimal, localcontext
    with localcontext() as ctx:
        ctx.prec = 60
        rates = deepcopy(MATERIALS[material])
        transfer = rates["transitions"]
        if not t2_return_enabled:
            transfer[2][0] = 0.0
        transfer = [[Decimal(str(value)) for value in row] for row in transfer]
        loss = [Decimal(str(rates["radiative"][i])) + Decimal(str(rates["nonradiative"][i])) + (Decimal(str(extra_triplet_loss)) if i else 0) for i in range(3)]
        initial = [Decimal(1), Decimal(0), Decimal(0)] if excitation == "optical" else [Decimal("0.25"), Decimal("0.75"), Decimal(0)]
        a = [[(loss[i] + sum(transfer[i])) if i == j else -transfer[j][i] for j in range(3)] + [initial[i]] for i in range(3)]
        for column in range(3):
            pivot = max(range(column, 3), key=lambda row: abs(a[row][column]))
            a[column], a[pivot] = a[pivot], a[column]
            divisor = a[column][column]
            a[column] = [value/divisor for value in a[column]]
            for row in range(3):
                if row != column:
                    factor = a[row][column]
                    a[row] = [left-factor*right for left,right in zip(a[row], a[column])]
        return [float(row[-1]) for row in a]


def validate():
    checks = []
    def check(name, condition):
        if not condition:
            raise AssertionError(name)
        checks.append(name)
    for key, rates in MATERIALS.items():
        optical = calculate(key, "optical")
        # Printed rate constants are rounded; 0.5 percentage point tolerance.
        check(f"{key}: reproduce rounded published computed PLQY", abs(optical["radiative_probability"] - rates["published_computed_plqy"]) < 0.005)
        previous = 1.0
        for loss in [0, 1, 10, 100, 1000, 10000, 1e5, 1e6, 1e7, 1e8]:
            for excitation in ["optical", "electrical"]:
                for t2 in [False, True]:
                    result = calculate(key, excitation, loss, t2)
                    reference = reference_residence(key, excitation, loss, t2)
                    check(f"{key}/{excitation}/{loss}/T2={t2}: independent 60-digit matrix solution", all(abs(value-reference[i]) < 1e-10*max(abs(reference[i]),1e-30) for i,value in enumerate(result["residence_s"].values())))
                    check(f"{key}/{excitation}/{loss}/T2={t2}: conservation", abs(result["probability_sum"]-1) < 1e-12)
                    check(f"{key}/{excitation}/{loss}/T2={t2}: matrix balance", result["relative_balance_residual"] < 1e-12)
                    check(f"{key}/{excitation}/{loss}/T2={t2}: non-negative states", all(value >= 0 for value in result["residence_s"].values()))
                    check(f"{key}/{excitation}/{loss}/T2={t2}: bounded fates", all(-1e-12 <= value <= 1+1e-12 for value in result["fate"].values()))
            current = calculate(key, "electrical", loss)["radiative_probability"]
            check(f"{key}/{loss}: quenching monotonicity", current <= previous + 1e-8)
            previous = current
    # Analytic limiting cases independent of the source dataset.
    dark_triplets = {"radiative": [1e8, 0, 0], "nonradiative": [0, 1e5, 1e5], "transitions": [[0,0,0],[0,0,0],[0,0,0]]}
    check("Fluorescent ideal singlet-only electrical limit is 25%", abs(calculate(dark_triplets)["radiative_probability"]-0.25) < 1e-12)
    check("Fluorescent ideal optical limit is 100%", abs(calculate(dark_triplets,"optical")["radiative_probability"]-1) < 1e-12)
    fully_recycled = {"radiative": [1e8,0,0], "nonradiative": [0,0,0], "transitions": [[0,1e8,0],[1e6,0,0],[1e6,0,0]]}
    check("Lossless recycling gives 100% internal photons", abs(calculate(fully_recycled)["radiative_probability"]-1) < 1e-12)
    no_photons = {"radiative": [0,0,0], "nonradiative": [1e5,1e5,1e5], "transitions": [[0,0,0],[0,0,0],[0,0,0]]}
    check("Zero radiative channels give zero photons", calculate(no_photons)["radiative_probability"] == 0)
    return checks


def build_results():
    results = {
        "title": "Can faster triplet conversion guarantee a better OLED?",
        "source": SOURCE,
        "data_class": "Mol2Mat reproduction of published computed rates, plus explicitly imposed sensitivity scenarios",
        "method": "Exact integrated solution M u = N0; u is residence per initial exciton, and k*u is the corresponding integrated event count.",
        "source_rate_precision": "Rounded Table 1 values, mainly two significant figures; outputs are not more accurate than their inputs.",
        "initial_population": {"optical": [1,0,0], "electrical": [0.25,0.75,0]},
        "controls": {
            "extra_triplet_loss": {"unit":"s^-1", "default":0, "range":[0,1000000], "meaning":"Assumed additional first-order loss from both triplet states"},
            "outcoupling": {"default":0.25, "range":[0.1,0.5], "meaning":"Illustrative external extraction fraction; not an optical-stack result"},
            "t2_return_enabled": {"default":True, "meaning":"False deletes T2->S1 as a mathematical counterfactual only"},
        },
        "rates": MATERIALS,
        "baselines": {},
        "loss_sweep": [],
        "validation": {"passed": True, "checks": validate()},
    }
    for key in MATERIALS:
        results["baselines"][key] = {
            "optical":calculate(key,"optical"),
            "electrical":calculate(key),
            "electrical_without_t2_return":calculate(key,t2_return_enabled=False),
            "electrical_added_loss_10000":calculate(key,extra_triplet_loss=10000),
        }
    for exponent in range(0,61):
        loss = 0.0 if exponent == 0 else 10 ** (exponent/10)
        results["loss_sweep"].append({"extra_triplet_loss_s-1":loss, "materials":{
            key:{"yield":calculate(key,extra_triplet_loss=loss)["radiative_probability"],"triplet_residence_s":calculate(key,extra_triplet_loss=loss)["triplet_residence_s"]}
            for key in MATERIALS
        }})
    return results


def build_analysis():
    import csv
    root = Path(__file__).resolve().parent
    records = json.loads((root / "records.json").read_text())
    model = build_results()
    curves = {key: [] for key in "ABCD"}
    excluded = {key: [] for key in "ABCD"}
    rows = list(csv.DictReader((root / "raw/device-eqe.csv").open()))
    for row in rows:
        try:
            x, y = float(row["luminance_cd_m2"]), float(row["eqe_percent"])
            reason = "below_1_cd_m2" if x < 1 else "outside_0_to_100_percent" if not 0 <= y <= 100 else ""
        except ValueError:
            x, y = row["luminance_cd_m2"], row["eqe_percent"]
            reason = "non_numeric_or_missing"
        point = {"luminance": x, "eqe": y, "sourceRow": int(row["source_row"])}
        if reason:
            point["reason"] = reason
            excluded[row["device"]].append(point)
        else:
            curves[row["device"]].append(point)
        assert reason == row["reason"], "Raw row classification mismatch"
    for curve in curves.values():
        curve.sort(key=lambda point: point["luminance"])
    comparisons = []
    for material in records["materials"]:
        if material["modelId"]:
            computed = model["baselines"][material["modelId"]]["optical"]["radiative_probability"]
            comparisons.append({"id":material["id"], "computedOpticalPlqy":computed,"measuredFilmPlqy":material["filmPlqy"],"differencePercentagePoints":100*(computed-material["filmPlqy"])})
    return {"schemaVersion":1, **records, "deviceCurves":curves,
        "curveProtocol":{**records["curveProtocol"],"excluded":excluded},
        "modelBaselines":model["baselines"],"quenchingSweep":model["loss_sweep"],
        "modelMeasurementComparison":comparisons,
        "validation":{"passed":True,"modelChecks":len(model["validation"]["checks"]),
        "sourceRows":len(rows),"displayedCurvePoints":sum(len(x) for x in curves.values()),
        "checks":model["validation"]["checks"]}}

if __name__ == "__main__":
    root = Path(__file__).resolve().parent
    analysis = build_analysis()
    (root / "analysis.json").write_text(json.dumps(analysis, indent=2, ensure_ascii=False)+"\n")
    print(json.dumps({"modelChecks":analysis["validation"]["modelChecks"],"sourceRows":analysis["validation"]["sourceRows"],"displayedCurvePoints":analysis["validation"]["displayedCurvePoints"]},indent=2))
