import argparse
import json
import re
from datetime import datetime, timezone
from pathlib import Path


ID_PATTERN = re.compile(r"\b(CHG|TRK|MEM|OBS|SES)-(\d{3,})\b")
STATE_PREFIXES = ("CHG", "TRK", "MEM", "OBS", "SES")
# Planning-chain IDs live inside state/artifacts documents and have no
# allocating tool, so a counter left behind their maxima would collide on the
# next planning session (found via the PRD-010/ARCH-010/IMP-011 drift, REQ-107).
PLANNING_ID_PATTERN = re.compile(
    r"\b(BR-REQ|BR-DEC|BR-Q|BR-RISK|PRD-REQ|PRD-NFR|PRD-ACC|ARCH-COMP|ARCH-DEC"
    r"|ARCH-IF|ARCH-RISK|IMP-PHASE|IMP-TASK|IMP-VAL)-(\d+)\b")
EXCLUDED_DIR_NAMES = {".git", ".mypy_cache", ".pytest_cache", ".ruff_cache", ".venv", "__pycache__", "node_modules"}
EXCLUDED_RELATIVE_PREFIXES = (
    ".prd_plugin/local/",
    ".prd_plugin/inbox/",
    ".prd_plugin/outbox/",
    ".prd_plugin/mailboxes/",
)
CLAIM_SOURCE_PREFIXES = (
    "docs/evidence/",
    "docs/traceability/",
    ".prd_plugin/state/sessions/",
)
TEXT_CLAIM_SOURCE_SUFFIXES = frozenset({".json", ".jsonl", ".md", ".txt"})
TIMESTAMP_SUFFIXES = (".json", ".jsonl")


def _relative_posix(root, path):
    return path.relative_to(root).as_posix()


def _is_excluded(root, path):
    relative = _relative_posix(root, path)
    if any(relative.startswith(prefix) for prefix in EXCLUDED_RELATIVE_PREFIXES):
        return True
    return any(part in EXCLUDED_DIR_NAMES for part in path.parts)


def _read_json(path):
    return json.loads(path.read_text(encoding="utf-8-sig"))


def _walk_values(value):
    if isinstance(value, dict):
        for child in value.values():
            yield from _walk_values(child)
    elif isinstance(value, list):
        for child in value:
            yield from _walk_values(child)
    else:
        yield value


def _parse_time(value):
    if not isinstance(value, str) or "T" not in value:
        return None
    try:
        parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
    except ValueError:
        return None
    if parsed.tzinfo is None:
        parsed = parsed.replace(tzinfo=timezone.utc)
    return parsed.astimezone(timezone.utc)


def _add_ids_from_value(value, ids_by_prefix):
    if isinstance(value, dict):
        identifier = value.get("id")
        if isinstance(identifier, str):
            match = ID_PATTERN.fullmatch(identifier.strip())
            if match:
                ids_by_prefix.setdefault(match.group(1), set()).add(identifier.strip())
        for child in value.values():
            _add_ids_from_value(child, ids_by_prefix)
    elif isinstance(value, list):
        for child in value:
            _add_ids_from_value(child, ids_by_prefix)


def _collect_canonical_ids(root):
    ids_by_prefix = {prefix: set() for prefix in STATE_PREFIXES}
    state_dir = root / ".prd_plugin" / "state"
    if state_dir.exists():
        for path in sorted(state_dir.rglob("*.json")):
            if _is_excluded(root, path):
                continue
            _add_ids_from_value(_read_json(path), ids_by_prefix)
        for path in sorted(state_dir.rglob("*.jsonl")):
            if _is_excluded(root, path):
                continue
            for line in path.read_text(encoding="utf-8-sig").splitlines():
                if line.strip():
                    _add_ids_from_value(json.loads(line), ids_by_prefix)
    return ids_by_prefix


def _claim_source_files(root):
    candidates = []
    for prefix in CLAIM_SOURCE_PREFIXES:
        base = root / prefix
        if base.exists():
            candidates.extend(
                path
                for path in base.rglob("*")
                if path.is_file() and path.suffix.lower() in TEXT_CLAIM_SOURCE_SUFFIXES
            )
    return sorted(path for path in candidates if not _is_excluded(root, path))


def _timestamp_files(root):
    return sorted(
        path
        for path in root.rglob("*")
        if path.is_file()
        and path.suffix in TIMESTAMP_SUFFIXES
        and not _is_excluded(root, path)
    )


def _missing_claim_findings(root, ids_by_prefix):
    findings = []
    seen = set()
    for path in _claim_source_files(root):
        text = path.read_text(encoding="utf-8-sig")
        for match in ID_PATTERN.finditer(text):
            identifier = match.group(0)
            prefix = match.group(1)
            key = (_relative_posix(root, path), identifier)
            if key in seen or identifier in ids_by_prefix.get(prefix, set()):
                continue
            seen.add(key)
            findings.append(
                {
                    "id": "CONS-ERR-001",
                    "severity": "error",
                    "path": _relative_posix(root, path),
                    "summary": f"{identifier} is claimed but no canonical {prefix} record exists.",
                    "missing_id": identifier,
                    "required_action": f"Create the canonical {identifier} record, remove the claim, or mark the work incomplete.",
                }
            )
    return findings


def _zero_pad(root):
    """Honor ids.zero_pad from config (default 3) so repos configured with a
    different width don't get false counter-gap errors for every valid ID."""
    config_path = root / ".prd_plugin" / "config.json"
    try:
        config = _read_json(config_path) if config_path.exists() else {}
        pad = (config.get("ids") or {}).get("zero_pad")
        return pad if isinstance(pad, int) and pad > 0 else 3
    except Exception:
        return 3


def _registry_counter_findings(root, ids_by_prefix):
    registry_path = root / ".prd_plugin" / "ids" / "registry.json"
    if not registry_path.exists():
        return []
    registry = _read_json(registry_path)
    next_values = registry.get("next", {})
    if not isinstance(next_values, dict):
        return []

    pad = _zero_pad(root)
    findings = []
    for prefix in STATE_PREFIXES:
        next_value = next_values.get(prefix)
        if not isinstance(next_value, int) or next_value <= 1:
            continue
        expected = {f"{prefix}-{number:0{pad}d}" for number in range(1, next_value)}
        missing = sorted(expected - ids_by_prefix.get(prefix, set()))
        if missing:
            findings.append(
                {
                    "id": "CONS-ERR-002",
                    "severity": "error",
                    "path": _relative_posix(root, registry_path),
                    "summary": f"Registry counter for {prefix} skips missing canonical records.",
                    "prefix": prefix,
                    "next": next_value,
                    "missing_ids": missing,
                    "required_action": "Add the missing canonical records or repair the registry counter before claiming completion.",
                }
            )
    return findings


def _planning_counter_findings(root):
    """Registry counters must stay ahead of every planning ID already used in
    state/artifacts, or the next allocation collides with an approved chain."""
    registry_path = root / ".prd_plugin" / "ids" / "registry.json"
    artifacts_dir = root / ".prd_plugin" / "state" / "artifacts"
    if not registry_path.exists() or not artifacts_dir.exists():
        return []
    next_values = _read_json(registry_path).get("next", {})
    if not isinstance(next_values, dict):
        return []

    maxima = {}
    for path in sorted(artifacts_dir.rglob("*.json")):
        if _is_excluded(root, path):
            continue
        text = path.read_text(encoding="utf-8-sig")
        for prefix, number in PLANNING_ID_PATTERN.findall(text):
            maxima[prefix] = max(maxima.get(prefix, 0), int(number))

    findings = []
    for prefix, max_used in sorted(maxima.items()):
        next_value = next_values.get(prefix)
        if isinstance(next_value, int) and next_value <= max_used:
            findings.append(
                {
                    "id": "CONS-ERR-004",
                    "severity": "error",
                    "path": _relative_posix(root, registry_path),
                    "summary": (f"Registry counter for {prefix} (next {next_value}) is at or behind "
                                f"{prefix}-{max_used} already used in planning artifacts; the next "
                                "allocation would collide."),
                    "prefix": prefix,
                    "next": next_value,
                    "max_used": max_used,
                    "required_action": f"Advance the registry counter for {prefix} beyond {max_used}.",
                }
            )
    return findings


def _future_timestamp_findings(root, now):
    findings = []
    for path in _timestamp_files(root):
        if path.suffix == ".json":
            values = _walk_values(_read_json(path))
        else:
            parsed_lines = [
                json.loads(line)
                for line in path.read_text(encoding="utf-8-sig").splitlines()
                if line.strip()
            ]
            values = _walk_values(parsed_lines)
        for value in values:
            parsed = _parse_time(value)
            if parsed and parsed > now:
                findings.append(
                    {
                        "id": "CONS-ERR-003",
                        "severity": "error",
                        "path": _relative_posix(root, path),
                        "summary": f"Future timestamp {value} is later than the validation time.",
                        "timestamp": value,
                        "required_action": "Use the actual validation time or mark the artifact as planned instead of completed.",
                    }
                )
                break
    return findings


def _tracking_branch_findings(root):
    """Validate the branch/promoted-state handshake used by parallel agents."""
    branch_dir = root / ".prd_plugin" / "state" / "tracking-branches"
    if not branch_dir.exists():
        return []
    tracking_path = root / ".prd_plugin" / "state" / "tracking.json"
    tracking_data = _read_json(tracking_path) if tracking_path.exists() else {"records": []}
    records = {
        record.get("id"): record
        for record in tracking_data.get("records", [])
        if isinstance(record, dict) and isinstance(record.get("id"), str)
    }
    findings = []
    seen = set()
    for path in sorted(branch_dir.glob("*.json")):
        relative = _relative_posix(root, path)
        try:
            branch = _read_json(path)
        except Exception as exc:
            findings.append({
                "id": "CONS-ERR-TBR-001", "severity": "error", "path": relative,
                "summary": f"Tracking branch is not valid JSON: {exc}",
                "required_action": "Repair or reject the branch before promotion.",
            })
            continue
        branch_id = branch.get("id")
        basic_errors = []
        if not isinstance(branch_id, str) or not re.fullmatch(r"DBR-\d{3,}", branch_id):
            basic_errors.append("missing or invalid DBR id")
        elif path.stem != branch_id:
            basic_errors.append("file name does not match branch id")
        elif branch_id in seen:
            basic_errors.append("duplicate branch id")
        seen.add(branch_id)
        if branch.get("kind") != "tracking":
            basic_errors.append("kind is not tracking")
        if not isinstance(branch.get("owner"), str) or not branch.get("owner"):
            basic_errors.append("owner is missing")
        if branch.get("state") not in {"observed", "experimenting", "promoted", "rejected", "retired"}:
            basic_errors.append("state is invalid")
        deltas = branch.get("deltas")
        if not isinstance(deltas, list) or not deltas or not all(
            isinstance(item, dict) and re.fullmatch(r"DBR-DELTA-\d{3,}", str(item.get("id", "")))
            for item in deltas
        ):
            basic_errors.append("DBR-DELTA identity is missing or invalid")
        proposed = branch.get("proposed")
        if not isinstance(proposed, dict):
            basic_errors.append("proposed tracking delta is missing or invalid")
        else:
            notes = proposed.get("notes")
            linked_ids = proposed.get("linked_ids")
            if not isinstance(notes, list) or not all(isinstance(note, str) and note for note in notes):
                basic_errors.append("proposed.notes is not an array of non-empty strings")
            if not isinstance(linked_ids, list) or not all(
                isinstance(linked_id, str) and re.fullmatch(r"[A-Z][A-Z0-9]*(?:-[A-Z0-9]+)*-\d+", linked_id)
                for linked_id in linked_ids
            ):
                basic_errors.append("proposed.linked_ids contains an invalid record ID")
        if basic_errors:
            findings.append({
                "id": "CONS-ERR-TBR-002", "severity": "error", "path": relative,
                "summary": "; ".join(basic_errors),
                "required_action": "Repair the branch metadata before it is merged or promoted.",
            })
            continue

        canonical = [record for record in records.values() if branch_id in (record.get("tracking_branch_ids") or [])]
        if branch.get("state") == "promoted":
            promotion = branch.get("promotion") or {}
            merge_id = promotion.get("id")
            tracking_id = promotion.get("tracking_id")
            record = records.get(tracking_id)
            merge_matches = record and any(
                item.get("branch_id") == branch_id and item.get("merge_id") == merge_id
                for item in (record.get("tracking_branch_merges") or []) if isinstance(item, dict)
            )
            if (
                not re.fullmatch(r"DBR-MERGE-\d{3,}", str(merge_id or ""))
                or not record
                or branch_id not in (record.get("tracking_branch_ids") or [])
                or not merge_matches
            ):
                findings.append({
                    "id": "CONS-ERR-TBR-003", "severity": "error", "path": relative,
                    "summary": "Promoted tracking branch does not match canonical TRK and DBR-MERGE provenance.",
                    "required_action": "Repair the canonical tracking link or roll the branch back to an unpromoted state.",
                })
        elif canonical:
            findings.append({
                "id": "CONS-ERR-TBR-004", "severity": "error", "path": relative,
                "summary": "Unpromoted tracking branch is already referenced by canonical tracking state.",
                "required_action": "Complete idempotent promotion or remove the unsupported canonical reference.",
            })
    return findings


def _status_for(findings):
    return "error" if any(finding["severity"] == "error" for finding in findings) else "ok"


def _extra_work_findings(root, ids_by_prefix):
    """Detect evidence/session files that reference multiple canonical IDs.

    A claim-source file that references 2+ canonical IDs is treated as a
    positive signal: the session did extra work that landed in canonical
    records. Per the D:\\Projects\\Improve Tier 4 corpus, 0/21 sessions
    were "aligned" because the bar treats extra work as drift. v0.5.32
    inverts that: extra work is informative, not an error.
    """
    findings = []
    seen_files = set()
    for path in _claim_source_files(root):
        text = path.read_text(encoding="utf-8-sig")
        referenced_ids: set[tuple[str, str]] = set()
        for match in ID_PATTERN.finditer(text):
            identifier = match.group(0)
            prefix = match.group(1)
            if identifier in ids_by_prefix.get(prefix, set()):
                referenced_ids.add((prefix, identifier))
        if len(referenced_ids) >= 2 and (root, path) not in seen_files:
            seen_files.add((root, path))
            findings.append(
                {
                    "id": "CONS-INFO-001",
                    "severity": "info",
                    "path": _relative_posix(root, path),
                    "summary": (
                        f"File references {len(referenced_ids)} canonical IDs across "
                        f"{len({p for p, _ in referenced_ids})} prefix(es); treated as a positive "
                        f"signal of extra work."
                    ),
                    "referenced_prefixes": sorted({p for p, _ in referenced_ids}),
                    "referenced_id_count": len(referenced_ids),
                    "required_action": "No action required; this is an informational signal.",
                }
            )
    return findings


def build_consistency_report(repo_root=".", now_text=None):
    root = Path(repo_root).resolve()
    now = _parse_time(now_text) if now_text else datetime.now(timezone.utc)
    if now is None:
        raise ValueError("--now must be an ISO timestamp")
    ids_by_prefix = _collect_canonical_ids(root)
    findings = []
    findings.extend(_missing_claim_findings(root, ids_by_prefix))
    findings.extend(_registry_counter_findings(root, ids_by_prefix))
    findings.extend(_planning_counter_findings(root))
    findings.extend(_tracking_branch_findings(root))
    findings.extend(_future_timestamp_findings(root, now))
    findings.extend(_extra_work_findings(root, ids_by_prefix))
    return {
        "status": _status_for(findings),
        "repo_root": str(root),
        "summary": {
            "errors": sum(1 for finding in findings if finding["severity"] == "error"),
            "warnings": sum(1 for finding in findings if finding["severity"] == "warning"),
            "infos": sum(1 for finding in findings if finding["severity"] == "info"),
        },
        "canonical_ids": {prefix: sorted(values) for prefix, values in ids_by_prefix.items()},
        "findings": findings,
    }


def format_markdown(report):
    lines = [
        "# PRD Plugin State Consistency Check",
        "",
        f"Status: `{report['status']}`",
        f"Repo: `{report['repo_root']}`",
        "",
        "## Findings",
        "",
    ]
    if not report["findings"]:
        lines.append("- None")
    else:
        for finding in report["findings"]:
            lines.append(
                f"- `{finding['id']}` `{finding['severity']}` ({finding['path']}): "
                f"{finding['summary']} Required action: {finding['required_action']}"
            )
    return "\n".join(lines) + "\n"


def main(argv=None):
    parser = argparse.ArgumentParser(description="Check PRD Plugin canonical state consistency.")
    parser.add_argument("--repo-root", default=".", help="Repository root to inspect.")
    parser.add_argument("--now", help="Validation time as an ISO timestamp. Defaults to current time.")
    parser.add_argument("--format", choices=("markdown", "json"), default="markdown")
    parser.add_argument("--output", help="Optional output path.")
    args = parser.parse_args(argv)

    report = build_consistency_report(args.repo_root, args.now)
    content = json.dumps(report, indent=2) + "\n" if args.format == "json" else format_markdown(report)
    if args.output:
        output = Path(args.output)
        output.parent.mkdir(parents=True, exist_ok=True)
        output.write_text(content, encoding="utf-8")
    else:
        print(content, end="")
    return 1 if report["status"] == "error" else 0


if __name__ == "__main__":
    raise SystemExit(main())
