"""Merge safe carried-forward planning artifacts into an incremental re-run."""
from __future__ import annotations

import argparse
import copy
import json
import re
import sys
from pathlib import Path
from typing import Any, Mapping

from .convergence_store import write_json_atomic
from .final_report_schema import load_schema_version
from .report_narrative import parse_narrative, writer_owned_data
from .json_boundary import load_owned_object
from .plan_items import (
    PlanItemContractError,
    content_hash,
    extract_plan_items,
    planning_stage_ledger,
    reverify_item_ids,
)


PREP_VERDICT_RE = re.compile(r"^P-Prep-S([1-9][0-9]*)-")
STEP_VERDICT_RE = re.compile(r"^P-Step-([1-9][0-9]*)\.")
_CHECKLIST_ID_PREFIXES = ("P-Val-", "P-Req-", "P-Rb-")


class CarryError(Exception):
    """Carry merge refused — schema drift or structural mismatch."""


def _planning(data: dict) -> dict:
    planning = data.get("implementationPlanning", {})
    if not isinstance(planning, dict):
        raise CarryError("implementationPlanning must be an object")
    return planning


def _indexed_rows(rows: object, *, key: str, label: str) -> dict:
    if not isinstance(rows, list):
        raise CarryError(f"{label} must be an array")
    indexed = {}
    for row in rows:
        if not isinstance(row, dict) or key not in row:
            raise CarryError(f"{label} contains a row without {key}")
        identity = row[key]
        if identity in indexed:
            raise CarryError(f"{label} contains duplicate {key} {identity!r}")
        indexed[identity] = row
    return indexed


def _stage_rows(planning: dict, *, snapshot: str) -> dict[int, dict]:
    rows = _indexed_rows(
        planning.get("stages", []), key="stage", label=f"{snapshot} stages",
    )
    try:
        return {int(stage): row for stage, row in rows.items()}
    except (TypeError, ValueError) as exc:
        raise CarryError(f"{snapshot} stages contains an invalid stage number") from exc


def _prep_items(planning: dict, *, snapshot: str) -> dict[str, dict]:
    preparation = planning.get("designPreparation", {})
    if not isinstance(preparation, dict):
        raise CarryError(f"{snapshot} designPreparation must be an object")
    return _indexed_rows(
        preparation.get("items", []),
        key="id",
        label=f"{snapshot} designPreparation.items",
    )


def _stage_refs(item: dict, *, label: str) -> set[int]:
    refs = item.get("stageRefs")
    if not isinstance(refs, list) or not refs:
        raise CarryError(f"{label} has invalid stageRefs")
    try:
        parsed = {int(stage) for stage in refs}
    except (TypeError, ValueError) as exc:
        raise CarryError(f"{label} has invalid stageRefs") from exc
    if any(stage < 1 for stage in parsed):
        raise CarryError(f"{label} has invalid stageRefs")
    return parsed


def _canonical(value: dict, *, ignore_carry_marker: bool = False) -> str:
    comparable = copy.deepcopy(value)
    if ignore_carry_marker:
        comparable.pop("carriedForwardFromSeq", None)
    return json.dumps(comparable, ensure_ascii=False, sort_keys=True, separators=(",", ":"))


def _prep_verdict_stage(item: dict) -> int | None:
    match = PREP_VERDICT_RE.match(str(item.get("id", "")))
    return int(match.group(1)) if match else None


def _merge_legacy_plan_items(prev_planning: dict, cur_planning: dict, prev_seq: str) -> None:
    prev_items = prev_planning.get("planBodyVerification", {}).get("planItems", [])
    cur_pbv = cur_planning.setdefault("planBodyVerification", {})
    cur_items = cur_pbv.setdefault("planItems", [])
    current_ids = {item["id"] for item in cur_items}
    for item in prev_items:
        if item["id"] not in current_ids:
            carried = copy.deepcopy(item)
            carried["carriedForwardFromSeq"] = str(prev_seq)
            cur_items.append(carried)
    cur_items.sort(key=lambda item: item["id"])


def _validate_stage_sets(
    prev_stages: dict[int, dict],
    cur_stages: dict[int, dict],
    carry_stages: set[int],
    reverify_stages: set[int],
) -> None:
    overlap = carry_stages & reverify_stages
    if overlap:
        raise CarryError(f"carry/reverify stage sets overlap: {sorted(overlap)}")
    missing_prior = carry_stages - set(prev_stages)
    if missing_prior:
        raise CarryError(f"carry stage(s) missing from prior snapshot: {sorted(missing_prior)}")
    missing_current = reverify_stages - set(cur_stages)
    if missing_current:
        raise CarryError(
            f"reverify stage(s) missing from current snapshot: {sorted(missing_current)}"
        )


def _merge_stage_rows(
    cur_planning: dict,
    prev_stages: dict[int, dict],
    cur_stages: dict[int, dict],
    carry_stages: set[int],
    reverify_stages: set[int],
) -> None:
    merged = [copy.deepcopy(prev_stages[stage]) for stage in carry_stages]
    merged.extend(copy.deepcopy(cur_stages[stage]) for stage in reverify_stages)
    merged.sort(key=lambda row: int(row["stage"]))
    cur_planning["stages"] = merged


def _merge_prep_items(
    prev_planning: dict,
    cur_planning: dict,
    carry_stages: set[int],
    reverify_stages: set[int],
) -> None:
    prev_items = _prep_items(prev_planning, snapshot="prior")
    cur_items = _prep_items(cur_planning, snapshot="current")
    scoped_stages = carry_stages | reverify_stages
    carried: dict[str, dict] = {}
    for item_id, item in prev_items.items():
        refs = _stage_refs(item, label=f"prior PREP item {item_id}")
        if refs & carry_stages and refs & reverify_stages:
            raise CarryError(f"prior PREP item {item_id} crosses carry and reverify stages")
        if not refs <= scoped_stages:
            raise CarryError(f"prior PREP item {item_id} falls outside carry/reverify scope")
        if refs <= carry_stages:
            duplicate = cur_items.get(item_id)
            if duplicate is not None and _canonical(duplicate) != _canonical(item):
                raise CarryError(f"carry-owned PREP item {item_id} conflicts with current")
            carried[item_id] = copy.deepcopy(item)

    retained: dict[str, dict] = {}
    for item_id, item in cur_items.items():
        refs = _stage_refs(item, label=f"current PREP item {item_id}")
        if refs & carry_stages and refs & reverify_stages:
            raise CarryError(f"current PREP item {item_id} crosses carry and reverify stages")
        if not refs <= scoped_stages:
            raise CarryError(f"current PREP item {item_id} falls outside carry/reverify scope")
        if refs & carry_stages and not refs & reverify_stages:
            previous = prev_items.get(item_id)
            if previous is None:
                raise CarryError(f"current PREP item {item_id} is a carry-stage scope leak")
            previous_refs = _stage_refs(
                previous,
                label=f"prior PREP item {item_id}",
            )
            if previous_refs != refs:
                raise CarryError(
                    f"current PREP item {item_id} changes stage ownership"
                )
            continue
        retained[item_id] = copy.deepcopy(item)
    retained.update(carried)

    preparation = cur_planning.setdefault("designPreparation", {})
    preparation["items"] = [retained[item_id] for item_id in sorted(retained)]


def _merge_non_prep_verdicts(
    prev_items: dict[str, dict],
    cur_items: dict[str, dict],
    prev_seq: str,
) -> list[dict]:
    merged = {
        item_id: copy.deepcopy(item)
        for item_id, item in cur_items.items()
        if _prep_verdict_stage(item) is None
    }
    for item_id, item in prev_items.items():
        if _prep_verdict_stage(item) is None and item_id not in merged:
            carried = copy.deepcopy(item)
            carried["carriedForwardFromSeq"] = str(prev_seq)
            merged[item_id] = carried
    return list(merged.values())


def _merge_prep_verdicts(
    prev_items: dict[str, dict],
    cur_items: dict[str, dict],
    carry_stages: set[int],
    reverify_stages: set[int],
    prev_seq: str,
) -> list[dict]:
    merged = []
    scoped_stages = carry_stages | reverify_stages
    for item_id, item in cur_items.items():
        stage = _prep_verdict_stage(item)
        if stage is None:
            continue
        if stage not in scoped_stages:
            raise CarryError(f"current PREP verdict {item_id} falls outside carry/reverify scope")
        if stage in carry_stages:
            previous = prev_items.get(item_id)
            if previous is None:
                raise CarryError(f"current PREP verdict {item_id} is a carry-stage scope leak")
            if _canonical(previous, ignore_carry_marker=True) != _canonical(
                item, ignore_carry_marker=True,
            ):
                raise CarryError(f"carry-owned PREP verdict {item_id} conflicts with current")
        else:
            merged.append(copy.deepcopy(item))

    for item_id, item in prev_items.items():
        stage = _prep_verdict_stage(item)
        if stage is None:
            continue
        if stage not in scoped_stages:
            raise CarryError(f"prior PREP verdict {item_id} falls outside carry/reverify scope")
        if stage in carry_stages:
            carried = copy.deepcopy(item)
            carried["carriedForwardFromSeq"] = str(prev_seq)
            merged.append(carried)
    return merged


def _merge_stage_aware(
    prev_planning: dict,
    cur_planning: dict,
    prev_seq: str,
    carry_stages: set[int],
    reverify_stages: set[int],
) -> None:
    prev_stages = _stage_rows(prev_planning, snapshot="prior")
    cur_stages = _stage_rows(cur_planning, snapshot="current")
    _validate_stage_sets(prev_stages, cur_stages, carry_stages, reverify_stages)
    _merge_stage_rows(
        cur_planning, prev_stages, cur_stages, carry_stages, reverify_stages,
    )
    _merge_prep_items(prev_planning, cur_planning, carry_stages, reverify_stages)

    prev_pbv = prev_planning.get("planBodyVerification", {})
    cur_pbv = cur_planning.setdefault("planBodyVerification", {})
    prev_items = _indexed_rows(
        prev_pbv.get("planItems", []), key="id", label="prior planItems",
    )
    cur_items = _indexed_rows(
        cur_pbv.get("planItems", []), key="id", label="current planItems",
    )
    merged = _merge_non_prep_verdicts(prev_items, cur_items, prev_seq)
    merged.extend(_merge_prep_verdicts(
        prev_items, cur_items, carry_stages, reverify_stages, prev_seq,
    ))
    merged.sort(key=lambda item: item["id"])
    cur_pbv["planItems"] = merged


def merge_carried_forward(
    prev: dict,
    cur: dict,
    prev_seq: str = "previous",
    carry_stages: set[int] | None = None,
    reverify_stages: set[int] | None = None,
) -> dict:
    if prev.get("schemaVersion") != cur.get("schemaVersion"):
        raise CarryError(
            f"schemaVersion drift {prev.get('schemaVersion')!r} != {cur.get('schemaVersion')!r}; "
            "cannot carry forward — caller must fall back to full"
        )
    prev_planning = _planning(prev)
    cur_planning = cur.setdefault("implementationPlanning", {})
    if (carry_stages is None) != (reverify_stages is None):
        raise CarryError("carry_stages and reverify_stages must be provided together")
    if carry_stages is None:
        _merge_legacy_plan_items(prev_planning, cur_planning, prev_seq)
    else:
        _merge_stage_aware(
            prev_planning, cur_planning, prev_seq, carry_stages, reverify_stages,
        )
    return cur


def _plan_item_stage(item: Mapping[str, Any]) -> int | None:
    item_id = str(item.get("id") or "")
    for pattern in (PREP_VERDICT_RE, STEP_VERDICT_RE):
        match = pattern.match(item_id)
        if match:
            return int(match.group(1))
    return None


def _checklist_id(item_id: str) -> bool:
    return item_id.startswith(_CHECKLIST_ID_PREFIXES)


def _item_stages(item: Mapping[str, Any]) -> list[int]:
    scope = item.get("stageScope")
    if not isinstance(scope, list):
        return []
    return [
        value for value in scope
        if isinstance(value, int) and not isinstance(value, bool)
    ]


def _stamp_carried(previous: dict, current: Mapping[str, Any], prev_seq: str) -> dict:
    carried = copy.deepcopy(previous)
    carried["carriedForwardFromSeq"] = str(prev_seq)
    current_hash = current.get("contentHash")
    if isinstance(current_hash, str) and current_hash:
        carried["verifiedContentHash"] = current_hash
        carried["contentHash"] = current_hash
    return carried


def _carry_unchanged_checklists(
    prev: Mapping[str, Any],
    narrative: Mapping[str, Any],
    state_items: dict[str, dict],
    prev_items: dict[str, dict],
    *,
    prev_seq: str,
    reverify_stages: set[int],
) -> None:
    """본문이 같은 P-Val / P-Req / P-Rb 판정을 이월한다.

    단계가 없으면 매 재실행 라운드 1 큐에 남아 이미 통과한 줄이 다시 반대되고
    새 C 가 열린다. 추출 해시가 같으면 이전 표를 가져오고, 재검증 스테이지에
    걸린 행과 본문이 바뀐 행은 그대로 둔다.
    """
    try:
        prev_extracted = {
            item["id"]: item
            for item in extract_plan_items(_planning(dict(prev)))
            if isinstance(item.get("id"), str)
        }
        cur_extracted = {
            item["id"]: item
            for item in extract_plan_items(_planning(dict(narrative)))
            if isinstance(item.get("id"), str)
        }
    except (CarryError, PlanItemContractError, KeyError, TypeError):
        return
    for item_id, current in list(state_items.items()):
        if not _checklist_id(item_id):
            continue
        stages = _item_stages(current)
        if stages and set(stages) & reverify_stages:
            continue
        previous = prev_items.get(item_id)
        if previous is None:
            continue
        prev_row = prev_extracted.get(item_id)
        cur_row = cur_extracted.get(item_id)
        if prev_row is None or cur_row is None:
            continue
        if content_hash(prev_row) != content_hash(cur_row):
            continue
        state_items[item_id] = _stamp_carried(previous, current, prev_seq)


def _recompute_dispatch_queue(
    narrative: Mapping[str, Any],
    state_items: dict[str, dict],
) -> list[str]:
    try:
        extracted = extract_plan_items(_planning(dict(narrative)))
    except (CarryError, PlanItemContractError, KeyError, TypeError):
        return []
    by_id = {
        item["id"]: item
        for item in extracted
        if isinstance(item.get("id"), str)
    }
    previous_hashes = {
        item_id: content_hash(by_id[item_id])
        for item_id, row in state_items.items()
        if row.get("carriedForwardFromSeq") and item_id in by_id
    }
    if not previous_hashes:
        return []
    ledger = planning_stage_ledger(_planning(dict(narrative)))
    return reverify_item_ids(extracted, previous_hashes, ledger)


def _writer_stage_rows(data: Mapping[str, Any], *, snapshot: str) -> dict[int, dict]:
    planning = _planning(dict(writer_owned_data(data)))
    return _stage_rows(planning, snapshot=snapshot)


def merge_v3_plan_state(
    prev: dict[str, Any],
    narrative: dict[str, Any],
    state: dict[str, Any],
    *,
    prev_seq: str,
    carry_stages: set[int],
    reverify_stages: set[int],
) -> dict[str, Any]:
    """이전 판정 중 이월 스테이지 소유분과 본문이 같은 체크리스트를 복사한다."""
    prev_stages = _writer_stage_rows(prev, snapshot="prior")
    cur_stages = _stage_rows(_planning(narrative), snapshot="current")
    _validate_stage_sets(prev_stages, cur_stages, carry_stages, reverify_stages)
    for stage in sorted(carry_stages):
        current = cur_stages.get(stage)
        if current is None:
            raise CarryError(f"carry stage {stage} is missing from current narrative")
        if _canonical(prev_stages[stage]) != _canonical(current):
            raise CarryError(f"carry stage {stage} changed in current narrative")

    prev_items = _indexed_rows(
        _planning(prev).get("planBodyVerification", {}).get("planItems", []),
        key="id",
        label="prior planItems",
    )
    state_items = _indexed_rows(
        state.get("planBodyVerification", {}).get("planItems", []),
        key="id",
        label="current state planItems",
    )
    for item_id, current in state_items.items():
        stage = _plan_item_stage(current)
        if stage not in carry_stages:
            continue
        previous = prev_items.get(item_id)
        if previous is None:
            raise CarryError(f"carry-owned plan item {item_id} is missing from prior report")
        for field in ("subject", "sourceSection"):
            if current.get(field) != previous.get(field):
                raise CarryError(f"carry-owned plan item {item_id} changed {field}")
        state_items[item_id] = _stamp_carried(previous, current, prev_seq)

    prior_carried_ids = {
        item_id
        for item_id, item in prev_items.items()
        if _plan_item_stage(item) in carry_stages
    }
    missing = sorted(prior_carried_ids - set(state_items))
    if missing:
        raise CarryError(f"current state omitted carry-owned plan items: {missing}")
    _carry_unchanged_checklists(
        prev,
        narrative,
        state_items,
        prev_items,
        prev_seq=prev_seq,
        reverify_stages=reverify_stages,
    )
    pbv = state.setdefault("planBodyVerification", {})
    pbv["planItems"] = [
        state_items[item_id] for item_id in sorted(state_items)
    ]
    queue = _recompute_dispatch_queue(narrative, state_items)
    if queue:
        pbv["dispatchQueue"] = queue
    return state


def _sync_prepared_dispatch_queue(state_path: Path, merged: Mapping[str, Any]) -> None:
    """이월 뒤 프롬프트 큐를 상태 파일과 맞춘다.

    `plan-items prompt` 는 prepare 가 쓴 `plan-items-*.json` 을 읽는다. 시드가
    그 파일을 먼저 쓰므로, 이월이 큐를 줄인 뒤에는 형제 파일을 같이 고쳐야
    워커가 이미 통과한 P-Val 을 다시 보지 않는다.
    """
    name = state_path.name
    prefix = "plan-body-verification-"
    if not name.startswith(prefix):
        return
    prepared = state_path.with_name("plan-items-" + name.removeprefix(prefix))
    if not prepared.is_file():
        return
    queue = (
        merged.get("planBodyVerification", {}).get("dispatchQueue")
        if isinstance(merged.get("planBodyVerification"), Mapping)
        else None
    )
    if not isinstance(queue, list):
        return
    try:
        envelope = load_owned_object(prepared, artifact="prepared plan items")
    except (OSError, ValueError):
        return
    if not isinstance(envelope, dict):
        return
    envelope["dispatchQueue"] = queue
    write_json_atomic(prepared, envelope)


def _parse_stage_csv(value: str | None, *, option: str) -> set[int] | None:
    if value is None:
        return None
    try:
        return {int(token.strip()) for token in value.split(",") if token.strip()}
    except ValueError as exc:
        raise CarryError(f"{option} contains an invalid stage number") from exc


def _load_data(path: str, *, option: str) -> dict:
    """Read one data.json as a CarryError-reporting operation. Unreadable or
    malformed input means the carry cannot happen, which is what CarryError
    already communicates — without this the caller gets a traceback instead of
    the documented `carry refused:` refusal it knows how to fall back from.
    """
    try:
        data = load_owned_object(Path(path), artifact="incremental carry report")
    except (OSError, ValueError) as exc:
        raise CarryError(f"{option} could not be read: {exc}") from exc
    if not isinstance(data, dict):
        raise CarryError(f"{option} must contain a JSON object")
    return data


def main(argv: list[str]) -> int:
    ap = argparse.ArgumentParser(prog="okstra incremental-carry")
    ap.add_argument("--prev-data", required=True, help="prior run final-report data.json")
    current = ap.add_mutually_exclusive_group(required=True)
    current.add_argument("--cur-data", help="historical v2 current data.json")
    current.add_argument("--cur-narrative", help="contract v3 current narrative Markdown")
    ap.add_argument("--state", help="contract v3 convergence-owned plan state")
    ap.add_argument("--prev-seq", required=True, help="prior run seq, tagged onto carried items")
    ap.add_argument("--carry-stages", default=None, help="comma-separated carried stage numbers")
    ap.add_argument("--reverify-stages", default=None, help="comma-separated reverified stage numbers")
    ap.add_argument("--out", help="historical v2 merged data.json path")
    ap.add_argument("--out-state", help="contract v3 merged plan state path")
    args = ap.parse_args(argv)

    try:
        prev = _load_data(args.prev_data, option="--prev-data")
        carry_stages = _parse_stage_csv(args.carry_stages, option="--carry-stages")
        reverify_stages = _parse_stage_csv(
            args.reverify_stages, option="--reverify-stages",
        )
        if args.cur_narrative:
            if not args.state or not args.out_state or args.out:
                raise CarryError(
                    "--cur-narrative requires --state and --out-state, and forbids --out"
                )
            if carry_stages is None or reverify_stages is None:
                raise CarryError("contract v3 carry requires both stage sets")
            try:
                narrative = parse_narrative(
                    Path(args.cur_narrative).read_text(encoding="utf-8"),
                    load_schema_version("3.0"),
                )
            except (OSError, UnicodeError, ValueError) as exc:
                raise CarryError(f"--cur-narrative could not be read: {exc}") from exc
            state = _load_data(args.state, option="--state")
            merged = merge_v3_plan_state(
                prev,
                narrative,
                state,
                prev_seq=args.prev_seq,
                carry_stages=carry_stages,
                reverify_stages=reverify_stages,
            )
            output = Path(args.out_state)
        else:
            if not args.out or args.out_state or args.state:
                raise CarryError(
                    "--cur-data requires --out, and forbids --state and --out-state"
                )
            cur = _load_data(args.cur_data, option="--cur-data")
            merged = merge_carried_forward(
                prev,
                cur,
                prev_seq=args.prev_seq,
                carry_stages=carry_stages,
                reverify_stages=reverify_stages,
            )
            output = Path(args.out)
    except (CarryError, KeyError, TypeError) as exc:
        print(f"carry refused: {exc}", file=sys.stderr)
        return 1

    try:
        write_json_atomic(output, merged)
        _sync_prepared_dispatch_queue(output, merged)
    except OSError as exc:
        print(f"carry refused: --out could not be written: {exc}", file=sys.stderr)
        return 1
    carried = sum(
        1
        for item in merged.get("implementationPlanning", {})
        .get("planBodyVerification", {})
        .get("planItems", [])
        if item.get("carriedForwardFromSeq") == str(args.prev_seq)
    )
    print(f"merged {carried} carried-forward plan item(s) from seq {args.prev_seq} -> {output}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main(sys.argv[1:]))
