#!/usr/bin/env python3
"""Build a static CogGym EML directory from an EML v1.5 build manifest."""

from __future__ import annotations

import argparse
import copy
import hashlib
import itertools
import json
import random
from pathlib import Path
from typing import Any, Iterable

BUILDER_VERSION = "1.0.0"
MODEL_CARD_FIELDS = {
    "name",
    "version",
    "description",
    "intended_use",
    "provenance",
    "contract",
}


def load_json(path: Path) -> dict[str, Any]:
    try:
        return json.loads(path.read_text(encoding="utf-8"))
    except json.JSONDecodeError as error:
        raise ValueError(f"{path}: {error.msg}") from error


def sha256(path: Path) -> str:
    return hashlib.sha256(path.read_bytes()).hexdigest()


def write_json(path: Path, value: Any) -> None:
    path.write_text(json.dumps(value, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")


def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None:
    with path.open("w", encoding="utf-8") as destination:
        for record in records:
            destination.write(json.dumps(record, separators=(",", ":"), ensure_ascii=False) + "\n")


def validate_model_card(unit: dict[str, Any]) -> None:
    unit_id = unit.get("id", "<missing-id>")
    card = unit.get("model_card")
    if not isinstance(card, dict):
        raise ValueError(f"{unit_id}: trial primitive requires a model_card object")
    missing = sorted(MODEL_CARD_FIELDS.difference(card))
    if missing:
        raise ValueError(f"{unit_id}: model_card is missing fields {missing}")
    contract = card.get("contract", {})
    missing_contract = sorted({"stimulus", "query", "feedback"}.difference(contract))
    if missing_contract:
        raise ValueError(
            f"{unit_id}: model_card.contract is missing fields {missing_contract}"
        )


def expand_unit(unit: dict[str, Any]) -> list[dict[str, Any]]:
    validate_model_card(unit)
    unit_id = unit["id"]
    stimuli = unit["stimuli"]
    queries = unit["queries"]
    generation = unit["generation"]
    axes = generation.get("vary", [])
    axis_names = [axis["stimulus"] for axis in axes]
    held = generation.get("hold", {})
    query_names = generation.get("queries", list(queries))
    template = generation.get("id_template", "{unit}")

    if len(set(axis_names)) != len(axis_names):
        raise ValueError(f"{unit_id}: generation.vary contains duplicate stimuli")

    condition_lists: list[list[str]] = []
    for axis in axes:
        stimulus_name = axis["stimulus"]
        if stimulus_name not in stimuli:
            raise ValueError(f"{unit_id}: unknown varied stimulus {stimulus_name!r}")
        conditions = axis["conditions"]
        unknown = [name for name in conditions if name not in stimuli[stimulus_name]]
        if unknown:
            raise ValueError(f"{unit_id}: {stimulus_name!r} has no conditions {unknown!r}")
        condition_lists.append(conditions)

    assignments = itertools.product(*condition_lists) if axes else [()]
    expanded: list[dict[str, Any]] = []

    for values in assignments:
        selected = dict(zip(axis_names, values))
        try:
            trial_id = template.format_map({"unit": unit_id, **selected})
        except KeyError as error:
            raise ValueError(f"{unit_id}: unknown ID placeholder {error.args[0]!r}") from error

        concrete_stimuli = []
        for stimulus_name, conditions in stimuli.items():
            condition_name = selected.get(stimulus_name, held.get(stimulus_name, "base"))
            if condition_name not in conditions:
                raise ValueError(
                    f"{unit_id}: {stimulus_name!r} has no condition {condition_name!r}"
                )
            concrete_stimuli.append(copy.deepcopy(conditions[condition_name]))

        try:
            concrete_queries = [copy.deepcopy(queries[name]) for name in query_names]
        except KeyError as error:
            raise ValueError(f"{unit_id}: unknown query {error.args[0]!r}") from error

        trial: dict[str, Any] = {
            "id": trial_id,
            "stimuli": concrete_stimuli,
            "queries": concrete_queries,
        }
        for inherited in ("feedback", "delay"):
            if inherited in unit:
                trial[inherited] = copy.deepcopy(unit[inherited])
        expanded.append(trial)

    return expanded


def expand_units(units: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], dict[str, list[str]]]:
    trials: list[dict[str, Any]] = []
    ids_by_unit: dict[str, list[str]] = {}
    seen_ids: set[str] = set()

    for unit in units:
        if unit["id"] in ids_by_unit:
            raise ValueError(f"duplicate trial-unit ID: {unit['id']}")
        unit_trials = expand_unit(unit)
        unit_ids = [trial["id"] for trial in unit_trials]
        within_unit = {trial_id for trial_id in unit_ids if unit_ids.count(trial_id) > 1}
        duplicates = seen_ids.intersection(unit_ids).union(within_unit)
        if duplicates:
            raise ValueError(f"duplicate generated trial IDs: {sorted(duplicates)}")
        seen_ids.update(unit_ids)
        ids_by_unit[unit["id"]] = unit_ids
        trials.extend(unit_trials)

    return trials, ids_by_unit


def chunked(values: list[str], size: int) -> list[list[str]]:
    if size <= 0:
        raise ValueError("block_size must be greater than zero")
    return [values[index : index + size] for index in range(0, len(values), size)]


def materialize_flow(
    sequencing: dict[str, Any], ids_by_unit: dict[str, list[str]], seed: int
) -> list[dict[str, Any]]:
    flow = []
    sequence_index = 0

    for condition_spec in sequencing.get("conditions", []):
        condition: dict[str, Any] = {"sequences": []}
        if condition_spec.get("condition"):
            condition["condition"] = condition_spec["condition"]
        prefix = condition_spec.get("prefix", [])

        for sequence_spec in condition_spec["sequences"]:
            trial_ids: list[str] = []
            for unit_id in sequence_spec["units"]:
                if unit_id not in ids_by_unit:
                    raise ValueError(f"unknown trial-unit ID in sequence: {unit_id}")
                trial_ids.extend(ids_by_unit[unit_id])

            order = sequence_spec.get("order", "generated")
            if order == "reverse":
                trial_ids.reverse()
            elif order == "shuffle":
                random.Random(seed + sequence_index).shuffle(trial_ids)
            elif order != "generated":
                raise ValueError(f"unsupported sequence order: {order}")

            block_size = sequence_spec.get("block_size", len(trial_ids) or 1)
            blocks = ([copy.deepcopy(prefix)] if prefix else []) + chunked(trial_ids, block_size)
            condition["sequences"].append(
                {"seq_id": sequence_spec["seq_id"], "blocks": blocks}
            )
            sequence_index += 1
        flow.append(condition)

    return flow


def build(card_path: Path, output_override: Path | None = None) -> Path:
    card_path = card_path.resolve()
    card = load_json(card_path)
    if card.get("card_type") != "eml-build-card" or card.get("eml_version") != "1.5":
        raise ValueError("expected an EML v1.5 eml-build-card")

    generator = card["generator"]
    if generator.get("entrypoint") != Path(__file__).name:
        raise ValueError("build card entrypoint does not match this builder")
    if generator.get("version") != BUILDER_VERSION:
        raise ValueError(
            f"build card requests generator {generator.get('version')!r}; "
            f"this builder is {BUILDER_VERSION!r}"
        )

    source_path = (card_path.parent / card["source"]).resolve()
    source = load_json(source_path)
    required_metadata = {
        "experimentName",
        "description",
        "paperDOI",
        "taskType",
        "responseType",
        "paper-title",
        "citation",
        "year",
        "authors",
        "alias",
    }
    missing_metadata = sorted(required_metadata.difference(source.get("metadata", {})))
    if missing_metadata:
        raise ValueError(f"missing required metadata fields: {missing_metadata}")

    outputs = card["outputs"]
    output_dir = (
        output_override.resolve()
        if output_override
        else (card_path.parent / outputs["directory"]).resolve()
    )
    output_dir.mkdir(parents=True, exist_ok=True)

    trials, ids_by_unit = expand_units(source.get("trial_units", []))
    instructions = source.get("instructions", [])
    instruction_ids = [instruction["id"] for instruction in instructions]
    if len(instruction_ids) != len(set(instruction_ids)):
        raise ValueError("instruction IDs must be unique")
    seed = int(card.get("reproducibility", {}).get("seed", 0))
    experiment_flow = materialize_flow(source.get("sequencing", {}), ids_by_unit, seed)
    valid_component_ids = set(instruction_ids).union(trial["id"] for trial in trials)
    referenced_ids = {
        component_id
        for condition in experiment_flow
        for sequence in condition["sequences"]
        for block in sequence["blocks"]
        for component_id in block
    }
    unknown_ids = sorted(referenced_ids.difference(valid_component_ids))
    if unknown_ids:
        raise ValueError(f"experimentFlow references unknown IDs: {unknown_ids}")

    config = copy.deepcopy(source["metadata"])
    config["stimuli_count"] = len(trials)
    if "trial_layout_config" in source:
        config["trial_layout_config"] = copy.deepcopy(source["trial_layout_config"])
    config["experimentFlow"] = experiment_flow

    config_path = output_dir / outputs["config"]
    trials_path = output_dir / outputs["trials"]
    instructions_path = output_dir / outputs["instructions"]
    readme_path = output_dir / outputs["readme"]
    provenance_path = output_dir / outputs["provenance"]

    write_json(config_path, config)
    write_jsonl(trials_path, trials)
    write_jsonl(instructions_path, instructions)
    sequence_ids = [
        sequence["seq_id"]
        for condition in experiment_flow
        for sequence in condition["sequences"]
    ]
    readme_path.write_text(
        "\n".join(
            [
                f"# {card['name']}",
                "",
                "This runnable EML directory was generated from an EML v1.5 build manifest and self-describing trial primitives.",
                "",
                f"- Concrete trials: {len(trials)}",
                f"- Explicit sequences: {len(sequence_ids)}",
                f"- Generator: {generator['entrypoint']} {generator['version']}",
                f"- Reproducibility seed: {seed}",
                "- Build details and file hashes: `build_provenance.json`",
                "",
            ]
        ),
        encoding="utf-8",
    )

    generated_paths = [config_path, trials_path, instructions_path, readme_path]
    provenance = {
        "card_type": card["card_type"],
        "eml_version": card["eml_version"],
        "generator": {
            **copy.deepcopy(generator),
            "sha256": sha256(Path(__file__).resolve()),
        },
        "seed": seed,
        "source": {"path": card["source"], "sha256": sha256(source_path)},
        "build_card_sha256": sha256(card_path),
        "primitive_model_cards": [
            {
                "unit_id": unit["id"],
                "name": unit["model_card"]["name"],
                "version": unit["model_card"]["version"],
            }
            for unit in source.get("trial_units", [])
        ],
        "outputs": {
            path.name: {"sha256": sha256(path)} for path in generated_paths
        },
        "generated_trial_ids": [trial["id"] for trial in trials],
        "generated_sequence_ids": sequence_ids,
    }
    write_json(provenance_path, provenance)
    return output_dir


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("card", type=Path, help="path to eml_build_card.json")
    parser.add_argument("--out", type=Path, help="optional output-directory override")
    args = parser.parse_args()
    output = build(args.card, args.out)
    print(output)


if __name__ == "__main__":
    main()
