#!/usr/bin/env python3
"""Audit completed RFO cycles; this does not evaluate molecular quality."""
import argparse
from collections import Counter
import json
import math
from pathlib import Path


def audit(records, expected):
    if expected < 1:
        raise ValueError("expected cycle count must be positive")
    if not isinstance(records, list) or not records:
        raise ValueError("records must be a nonempty JSON list")
    if len(records) != expected:
        raise ValueError(f"expected {expected} cycles, found {len(records)}")
    methods = Counter()
    sequences = []
    seconds = 0.0
    for index, row in enumerate(records):
        if not isinstance(row, dict):
            raise ValueError(f"record {index}: expected an object")
        cycle = row.get("cycle")
        if type(cycle) is not int or cycle != index:
            raise ValueError(f"record {index}: expected zero-based cycle {index}")
        method = row.get("method")
        if method not in ("RF3", "Boltz"):
            raise ValueError(f"cycle {index}: unknown method")
        sequence = row.get("mpnn_sequence")
        if not isinstance(sequence, str) or not sequence:
            raise ValueError(f"cycle {index}: missing MPNN sequence")
        if set(sequence) - set("ACDEFGHIKLMNPQRSTVWY"):
            raise ValueError(f"cycle {index}: nonstandard amino acid")
        if sequences and len(sequence) != len(sequences[0]):
            raise ValueError(f"cycle {index}: chain-A length changed")
        elapsed = row.get("wall_time_sec")
        if (type(elapsed) not in (int, float)
                or not math.isfinite(elapsed) or elapsed < 0):
            raise ValueError(f"cycle {index}: invalid wall_time_sec")
        methods[method] += 1
        sequences.append(sequence)
        seconds += elapsed
    return {
        "completed_cycles": len(records),
        "methods": dict(sorted(methods.items())),
        "chain_a_length": len(sequences[0]),
        "unique_handoff_sequences": len(set(sequences)),
        "recorded_cycle_seconds": round(seconds, 2),
        "quality_evaluated": False,
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("records", type=Path)
    parser.add_argument("--expected-cycles", required=True, type=int)
    args = parser.parse_args()
    try:
        records = json.loads(args.records.read_text(encoding="utf-8"))
        result = audit(records, args.expected_cycles)
    except (OSError, UnicodeError, ValueError) as exc:
        parser.exit(1, f"ERROR: {exc}\n")
    print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
