"""Deterministic scoring for synthetic academic-visit plans.

Execute this module on Molab, alongside the data generation and inference jobs.
The oracle enumerates the finite candidate schedules supplied with each profile.
It does not search arbitrary new actions or assume access to unlisted financing.
"""

from __future__ import annotations

import argparse
import json
import math
from collections import Counter
from itertools import permutations, product
from pathlib import Path


def read_jsonl(path):
    with Path(path).open(encoding="utf-8") as handle:
        return [json.loads(line) for line in handle if line.strip()]


def write_jsonl(path, rows):
    Path(path).parent.mkdir(parents=True, exist_ok=True)
    with Path(path).open("w", encoding="utf-8") as handle:
        for row in rows:
            handle.write(json.dumps(row, ensure_ascii=False) + "\n")


def parse_response(response):
    """Read ordinary JSON, a fenced JSON block, or a JSON object in prose."""
    if isinstance(response, dict):
        return response, None
    if not isinstance(response, str):
        return None, "response_is_not_text"
    decoder = json.JSONDecoder()
    final_plan = None
    for start, char in enumerate(response):
        if char != "{":
            continue
        try:
            value, _ = decoder.raw_decode(response[start:])
        except json.JSONDecodeError:
            continue
        if isinstance(value, dict) and isinstance(value.get("actions"), list):
            final_plan = value
    if final_plan is not None:
        return final_plan, None
    return None, "no_action_json"


def _normal_actions(plan):
    result, errors = [], []
    for index, action in enumerate(plan.get("actions", [])):
        if not isinstance(action, dict):
            errors.append("malformed_action")
            continue
        action_id, day = action.get("id"), action.get("day")
        if not isinstance(action_id, str) or not isinstance(day, int) or isinstance(day, bool):
            errors.append("action_needs_id_and_integer_day")
            continue
        result.append({"id": action_id, "day": day, "order": index})
    return result, errors


def structural_check(profile, plan):
    actions, errors = _normal_actions(plan)
    specs = {a["id"]: a for a in profile["actions"]}
    counts = Counter(a["id"] for a in actions)
    by_id = {a["id"]: a for a in actions}
    errors += ["duplicate:" + key for key, count in counts.items() if count > 1]
    errors += ["unknown:" + key for key in counts if key not in specs]
    for action in actions:
        spec = specs.get(action["id"])
        if spec is None:
            continue
        day = action["day"]
        if day < spec.get("earliest", 0) or day > spec.get("latest", profile["horizon_days"]):
            errors.append("window:" + action["id"])
        if "allowed_days" in spec and day not in spec["allowed_days"]:
            errors.append("unavailable_date:" + action["id"])
        for dep in spec.get("deps", []):
            source = by_id.get(dep["id"])
            if source is None and dep["id"] in profile.get("initial_completed", {}):
                source = {"day": profile["initial_completed"][dep["id"]], "order": -1}
            if source is None:
                errors.append("dependency_missing:" + action["id"] + ":" + dep["id"])
            elif day < source["day"] + dep.get("lag", 0):
                errors.append("dependency_time:" + action["id"] + ":" + dep["id"])
            elif day == source["day"] and source["order"] >= action["order"]:
                errors.append("dependency_order:" + action["id"] + ":" + dep["id"])
            elif dep.get("max_lag") is not None and day > source["day"] + dep["max_lag"]:
                errors.append("dependency_expired:" + action["id"] + ":" + dep["id"])
        for other in spec.get("excludes", []):
            if other in by_id:
                errors.append("incompatible:" + action["id"] + ":" + other)
    for group in profile["required_groups"]:
        if not any(key in by_id for key in group["ids"]):
            errors.append("required_stage_missing:" + group["stage"])
    return actions, sorted(set(errors))


def _events(profile, actions):
    """Cash receipts and payments occur only on their stipulated dates."""
    specs = {a["id"]: a for a in profile["actions"]}
    events = []
    for action in actions:
        spec = specs.get(action["id"])
        if spec is None:
            events.append({"day": action["day"], "order": action["order"] * 1000,
                           "action": action["id"], "net": 0, "spent": 0, "kind": "action"})
            continue
        events.append({"day": action["day"], "order": action["order"] * 1000,
                       "action": action["id"], "net": -spec.get("cost", 0),
                       "spent": spec.get("cost", 0), "kind": "action"})
        for index, effect in enumerate(spec.get("cash_effects", [])):
            effect_day = action["day"] + effect["offset"]
            if effect_day > profile["horizon_days"]:
                continue
            value = effect["amount"]
            if effect["offset"] == 0:
                event_order = action["order"] * 1000 + index + 1
            else:
                # Date-level accounting: later receipts arrive at day start;
                # later living payments fall at day end. Offset-zero effects
                # are part of the originating step and follow that step.
                event_order = (-9000 if value >= 0 else 900000) + index
            events.append({"day": effect_day, "order": event_order,
                           "action": action["id"], "net": value, "spent": max(0, -value),
                           "kind": effect.get("label", "cash_effect")})
    for index, effect in enumerate(profile.get("external_cash_events", [])):
        value = effect["amount"]
        events.append({"day": effect["day"], "order": -10000 + index,
                       "action": None, "net": value, "spent": max(0, -value),
                       "kind": effect.get("label", "external_cash")})
    return sorted(events, key=lambda event: (event["day"], event["order"]))


def cash_requirement(profile, actions):
    """Analytic minimum initial reserve for one fully specified schedule."""
    specs = {a["id"]: a for a in profile["actions"]}
    events = _events(profile, actions)
    balance, minimum_balance = 0, 0
    daily_low = [0] * (profile["horizon_days"] + 1)
    cursor = 0
    for day in range(len(daily_low)):
        low = balance
        while cursor < len(events) and events[cursor]["day"] == day:
            balance += events[cursor]["net"]
            low = min(low, balance)
            minimum_balance = min(minimum_balance, balance)
            cursor += 1
        daily_low[day] = low
    needed = -minimum_balance
    hold_requirements = []
    for action in actions:
        hold = specs.get(action["id"], {}).get("hold")
        if not hold:
            continue
        start, end = action["day"] - hold["days"] + 1, action["day"]
        if start < 0 or end >= len(daily_low):
            return None, ["hold_window_outside_plan:" + action["id"]]
        required = max(0, hold["amount"] - min(daily_low[start:end + 1]))
        needed = max(needed, required)
        hold_requirements.append({"id": action["id"], "amount": hold["amount"],
                                  "start": start, "end": end, "minimum_initial_cash": required})
    return int(math.ceil(max(0, needed))), hold_requirements


def execute(profile, actions, initial_cash):
    """Replay the submitted schedule until its first impossible action/payment.

    Spending before failure is the execution loss measure. It is not necessarily
    irrecoverable loss: nonrefundable loss is recorded separately from explicit
    action refund terms. Missed or invalid actions do not receive imagined funds.
    """
    specs = {a["id"]: a for a in profile["actions"]}
    scheduled = {a["id"]: a for a in actions}
    completed = dict(profile.get("initial_completed", {}))
    active, seen = set(completed), set()
    balance, spent, nonrefundable = initial_cash, 0, 0
    history = [initial_cash] * (profile["horizon_days"] + 1)
    action_at = { (a["day"], a["order"] * 1000): a for a in actions }
    cursor, failure = 0, None
    events = _events(profile, actions)
    if any(action["day"] < 0 for action in actions):
        failure = {"kind": "date_before_plan", "day": 0}
    for day in range(len(history)):
        if failure:
            break
        daily_low = balance
        while cursor < len(events) and events[cursor]["day"] == day:
            event = events[cursor]
            action = action_at.get((day, event["order"])) if event["kind"] == "action" else None
            if action is not None:
                key, spec = action["id"], specs.get(action["id"])
                if spec is None:
                    failure = {"kind": "unknown_action", "action": key, "day": day}
                    break
                if key in seen:
                    failure = {"kind": "duplicate", "action": key, "day": day}
                    break
                seen.add(key)
                if not spec.get("earliest", 0) <= day <= spec.get("latest", profile["horizon_days"]):
                    failure = {"kind": "window", "action": key, "day": day}
                    break
                if "allowed_days" in spec and day not in spec["allowed_days"]:
                    failure = {"kind": "unavailable_date", "action": key, "day": day}
                    break
                for dep in spec.get("deps", []):
                    previous = completed.get(dep["id"])
                    if previous is None or day < previous + dep.get("lag", 0):
                        failure = {"kind": "dependency", "action": key, "day": day}
                        break
                    if dep.get("max_lag") is not None and day > previous + dep["max_lag"]:
                        failure = {"kind": "dependency_expired", "action": key, "day": day}
                        break
                if failure:
                    break
                if any(other in scheduled for other in spec.get("excludes", [])):
                    failure = {"kind": "incompatible", "action": key, "day": day}
                    break
                hold = spec.get("hold")
                if hold:
                    start = day - hold["days"] + 1
                    prior = history[start:day] + [daily_low]
                    if start < 0 or min(prior) < hold["amount"]:
                        failure = {"kind": "proof_hold", "action": key, "day": day,
                                   "required_stock": hold["amount"]}
                        break
            if event["action"] is not None and action is None and event["action"] not in active:
                cursor += 1
                continue
            if balance + event["net"] < 0:
                failure = {"kind": "cash_payment", "action": event["action"], "day": day,
                           "payment": event["spent"], "cash_before": balance}
                break
            balance += event["net"]
            spent += event["spent"]
            daily_low = min(daily_low, balance)
            if action is not None:
                key, spec = action["id"], specs[action["id"]]
                completed[key] = day
                active.add(key)
                nonrefundable += spec.get("cost", 0) * (1 - spec.get("refundable_fraction", 0))
            elif event["spent"]:
                nonrefundable += event["spent"]
            cursor += 1
        history[day] = daily_low
        if failure:
            break
    if failure is None:
        missing = [group["stage"] for group in profile["required_groups"]
                   if not any(key in completed for key in group["ids"])]
        if missing:
            failure = {"kind": "incomplete", "stages": missing, "day": profile["horizon_days"]}
    return {"execution_completed": failure is None, "failure": failure,
            "spending_before_failure": 0 if failure is None else spent,
            "nonrefundable_before_failure": 0 if failure is None else nonrefundable,
            "cash_at_stop": balance, "completed_action_count": len(completed)}


def oracle(profile):
    if profile.get("available_routes"):
        specs = {action["id"]: action for action in profile["actions"]}
        candidates = []
        combinations = 0
        for route in profile["available_routes"]:
            ids = route["actions"]
            if any(not specs[key].get("allowed_days") for key in ids):
                raise ValueError("Finite-route oracle requires allowed_days for every action: " + profile["id"])
            for days in product(*(specs[key]["allowed_days"] for key in ids)):
                by_day = {}
                for key, day in zip(ids, days):
                    by_day.setdefault(day, []).append(key)
                date_order = sorted(by_day)
                for orders in product(*(permutations(by_day[day]) for day in date_order)):
                    combinations += 1
                    candidate = {"route_id": route["id"], "actions": [
                        {"id": key, "day": day} for day, order in zip(date_order, orders) for key in order]}
                    _, errors = structural_check(profile, candidate)
                    if not errors:
                        candidates.append(candidate)
        scope = "all supplied available routes, all permitted dates, and all within-day action orders"
    else:
        candidates = profile["oracle_candidates"]
        combinations = len(candidates)
        scope = "authored candidate schedules only"
    valid = []
    for candidate in candidates:
        actions, errors = structural_check(profile, candidate)
        required, holds = cash_requirement(profile, actions)
        if errors or required is None:
            raise ValueError(f"Invalid authored oracle schedule for {profile['id']}: {errors or holds}")
        replay = execute(profile, actions, required)
        if not replay["execution_completed"]:
            raise ValueError(f"Oracle cash calculation disagrees with execution for {profile['id']}: {replay['failure']}")
        valid.append((required, candidate))
    if not valid:
        raise ValueError("No oracle candidates for " + profile["id"])
    required, candidate = min(valid, key=lambda pair: pair[0])
    by_route = {}
    for route_required, route_candidate in valid:
        route_id = route_candidate.get("route_id", "authored_reference")
        by_route.setdefault(route_id, []).append((route_required, route_candidate))
    route_minima = []
    for route_id, route_candidates in sorted(by_route.items()):
        route_required, route_candidate = min(route_candidates, key=lambda pair: pair[0])
        route_minima.append({"route_id": route_id, "minimum_initial_cash": route_required,
                             "candidate": route_candidate,
                             "enumerated_candidate_count": len(route_candidates)})
    return {"minimum_initial_cash": required, "candidate": candidate,
            "enumerated_candidate_count": len(valid), "enumerated_combinations": combinations,
            "scope": scope, "route_minima": route_minima}


def score_response(profile, row, oracle_result):
    plan, parse_error = parse_response(row.get("response"))
    result = {key: row.get(key) for key in ("id", "profile_id", "B", "model", "seed", "condition")}
    result.update({"oracle_min_cash": oracle_result["minimum_initial_cash"],
                   "oracle_feasible": row["B"] >= oracle_result["minimum_initial_cash"],
                   "parse_error": parse_error})
    if parse_error:
        result.update({"stage_coverage": 0, "structurally_valid": False,
                       "full_schedule_feasible": False, "model_min_cash": None,
                       "startup_cash_excess": None, "structural_errors": [parse_error],
                       "execution_completed": False, "failure": {"kind": "unparseable"},
                       "spending_before_failure": 0, "nonrefundable_before_failure": 0})
        return result
    actions, errors = structural_check(profile, plan)
    # Surface coverage asks whether the named stages were listed, independently
    # of whether their dates can be executed or even parsed as integer days.
    selected = {a["id"] for a in plan["actions"]
                if isinstance(a, dict) and isinstance(a.get("id"), str)}
    coverage = sum(any(key in selected for key in group["ids"]) for group in profile["required_groups"])
    model_min_cash, hold_data = cash_requirement(profile, actions) if not errors else (None, [])
    if model_min_cash is None and not errors:
        errors += hold_data
    excess = None if model_min_cash is None else model_min_cash - oracle_result["minimum_initial_cash"]
    result.update({"stage_coverage": coverage / len(profile["required_groups"]),
                   "structurally_valid": not errors, "structural_errors": errors,
                   "model_min_cash": model_min_cash, "startup_cash_excess": excess,
                   "full_schedule_feasible": not errors and model_min_cash <= row["B"],
                   "hold_requirements": hold_data})
    result.update(execute(profile, actions, row["B"]))
    return result


def summarize(scores):
    groups = {}
    for score in scores:
        key = (score["model"], score["condition"], score["B"])
        groups.setdefault(key, []).append(score)
    output = []
    for (model, condition, buffer), rows in sorted(groups.items(), key=lambda item: str(item[0])):
        feasible = [r for r in rows if r["oracle_feasible"]]
        valid = [r for r in feasible if r["startup_cash_excess"] is not None]
        mean = lambda values: sum(values) / len(values) if values else None
        output.append({"model": model, "condition": condition, "B": buffer, "n": len(rows),
                       "mean_stage_coverage": mean([r["stage_coverage"] for r in rows]),
                       "full_schedule_feasible_rate": mean([r["full_schedule_feasible"] for r in rows]),
                       "oracle_feasible_n": len(feasible), "cash_excess_valid_n": len(valid),
                       "oracle_feasible_success_rate": mean([r["full_schedule_feasible"] for r in feasible]),
                       "mean_startup_cash_excess_valid": mean([r["startup_cash_excess"] for r in valid]),
                       "mean_spending_before_failure_oracle_feasible": mean([r["spending_before_failure"] for r in feasible]),
                       "mean_nonrefundable_before_failure_oracle_feasible": mean([r["nonrefundable_before_failure"] for r in feasible]),
                       "failure_counts": dict(Counter((r.get("failure") or {}).get("kind", "success") for r in rows))})
    return output


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--profiles", default="data/profiles.jsonl")
    parser.add_argument("--responses", required=True)
    parser.add_argument("--out", default="results/scores.jsonl")
    parser.add_argument("--summary", default="results/summary.json")
    args = parser.parse_args()
    profiles = {p["id"]: p for p in read_jsonl(args.profiles)}
    oracle_results = {key: oracle(profile) for key, profile in profiles.items()}
    scores = [score_response(profiles[row["profile_id"]], row, oracle_results[row["profile_id"]])
              for row in read_jsonl(args.responses)]
    write_jsonl(args.out, scores)
    Path(args.summary).parent.mkdir(parents=True, exist_ok=True)
    Path(args.summary).write_text(json.dumps(summarize(scores), ensure_ascii=False, indent=2), encoding="utf-8")
    print(json.dumps({"scored": len(scores), "profiles": len(profiles)}, ensure_ascii=False))


if __name__ == "__main__":
    main()
