#!/usr/bin/env python3
"""merge_group_pr_resolver.py — merge_group 이벤트에 포함된 (pr_number, pr_head_sha) 3계층 역추적.
task-2781. fail-closed: 확정 못하면 MergeGroupPRsUnresolved."""
from __future__ import annotations

import argparse
import contextvars
import json
import re
import subprocess

DEFAULT_REPO = "JonghyukJeon/dev_workspace"

_STDERR_SNIP_MAX = 500
# ★ PR266 microfix: 모듈-레벨 mutable global sink → contextvars.
# default=None(immutable sentinel) — resolver context 밖에서는 sink가 None이라 diagnostics 누적 0(누수 없음).
# resolve_merge_group_prs 시작 시 새 list를 context에 set, 종료 시 token reset으로 복원 → 병렬/중첩 호출 격리.
_GH_DIAG_SINK: "contextvars.ContextVar[list | None]" = contextvars.ContextVar(
    "gh_diag_sink", default=None
)
_GH_DIAG_TIER: "contextvars.ContextVar[str | None]" = contextvars.ContextVar(
    "gh_diag_tier", default=None
)


def _reset_gh_diag():
    """새 diagnostics context 시작. 복원용 (sink_token, tier_token) 반환."""
    sink_token = _GH_DIAG_SINK.set([])
    tier_token = _GH_DIAG_TIER.set(None)
    return sink_token, tier_token


def _restore_gh_diag(tokens):
    """resolve 종료 시 context를 이전 상태로 복원(token reset)."""
    if not tokens:
        return
    sink_token, tier_token = tokens
    _GH_DIAG_SINK.reset(sink_token)
    _GH_DIAG_TIER.reset(tier_token)


def _record_gh_failure(api_path, rc, stderr):
    """gh 호출 실패를 현재 context sink에 누적. 시그니처 불변 원칙 위해 wrapper가 이 함수만 호출.
    ★ resolver context 밖 호출(sink=None) → 누수 없이 no-op.
    ★ PR266 microfix: stderr가 bytes/비-str이면 str로 안전 변환 → stderr_snip은 항상 str
      (subprocess text=False 등으로 bytes가 들어와도 diagnostics JSON 직렬화가 깨지지 않음)."""
    sink = _GH_DIAG_SINK.get()
    if sink is None:
        return
    if isinstance(stderr, bytes):
        stderr = stderr.decode("utf-8", errors="replace")
    elif stderr is not None and not isinstance(stderr, str):
        stderr = str(stderr)
    snip = (stderr or "")[:_STDERR_SNIP_MAX]
    sink.append({
        "tier": _GH_DIAG_TIER.get(),
        "api_path": api_path,
        "rc": rc,
        "stderr_snip": snip,
    })


class MergeGroupPRsUnresolved(Exception):
    # additive diagnostics 부착용 속성(순수 관측 — 판정/결과 불변)
    diagnostics: list
    repo: str | None
    merge_group_sha: str | None

    def __init__(self, *args):
        super().__init__(*args)
        self.diagnostics = []
        self.repo = None
        self.merge_group_sha = None


# --- gh subprocess wrappers (테스트 monkeypatch seam) ---

def _gh_api(endpoint: str, timeout: int = 30) -> tuple[int, list | dict]:
    """gh api {endpoint} 호출 → (returncode, parsed_json_or_empty).
    gemini_evidence_verify._gh_api 와 동일 패턴."""
    try:
        proc = subprocess.run(
            ["gh", "api", endpoint],
            capture_output=True, text=True, timeout=timeout
        )
        rc = proc.returncode
        if rc != 0 or not proc.stdout.strip():
            _record_gh_failure(endpoint, rc, proc.stderr)
            return rc, []
        try:
            data = json.loads(proc.stdout)
        except json.JSONDecodeError:
            _record_gh_failure(endpoint, rc, proc.stderr or "JSONDecodeError")
            return rc, []
        return rc, data
    except subprocess.TimeoutExpired:
        _record_gh_failure(endpoint, -1, "TimeoutExpired")
        return -1, []
    except Exception as e:
        _record_gh_failure(endpoint, -1, f"{type(e).__name__}: {e}")
        return -1, []


def _gh_graphql(query: str, variables: dict | None = None, timeout: int = 30) -> tuple[int, dict]:
    """gh api graphql -f query=... -F var=... 형태. 실패 시 (rc, {}) 반환. 예외 시 (-1, {})."""
    try:
        cmd = ["gh", "api", "graphql", "-f", f"query={query}"]
        if variables:
            for key, val in variables.items():
                cmd += ["-F", f"{key}={val}"]
        proc = subprocess.run(
            cmd,
            capture_output=True, text=True, timeout=timeout
        )
        rc = proc.returncode
        if rc != 0 or not proc.stdout.strip():
            _record_gh_failure("graphql", rc, proc.stderr)
            return rc, {}
        try:
            data = json.loads(proc.stdout)
        except json.JSONDecodeError:
            _record_gh_failure("graphql", rc, proc.stderr or "JSONDecodeError")
            return rc, {}
        if not isinstance(data, dict):
            _record_gh_failure("graphql", rc, "non-dict response")
            return rc, {}
        return rc, data
    except subprocess.TimeoutExpired:
        _record_gh_failure("graphql", -1, "TimeoutExpired")
        return -1, {}
    except Exception as e:
        _record_gh_failure("graphql", -1, f"{type(e).__name__}: {e}")
        return -1, {}


# --- 내부 검증 헬퍼 ---

def _accept(entries: list[dict]) -> bool:
    """모든 entry가 pr_number(int>0) AND pr_head_sha(비어있지 않은 str) 이면 True.
    하나라도 불완전하면 False (부분 수용 금지, fail-closed)."""
    if not entries:
        return False
    for e in entries:
        pr_num = e.get("pr_number")
        pr_sha = e.get("pr_head_sha")
        if not isinstance(pr_num, int) or pr_num <= 0:
            return False
        if not isinstance(pr_sha, str) or not pr_sha.strip():
            return False
    return True


def _dedupe(entries: list[dict]) -> list[dict]:
    """pr_number 기준 dedupe (첫 번째 entry 우선)."""
    seen: set[int] = set()
    result: list[dict] = []
    for e in entries:
        pr_num = e.get("pr_number")
        if isinstance(pr_num, int) and pr_num not in seen:
            seen.add(pr_num)
            result.append(e)
    return result


# --- 3계층 역추적 함수들 ---

def resolve_via_api(repo: str, merge_group_sha: str, base_sha: str | None = None) -> list[dict]:
    """1차 신뢰원: GitHub GraphQL mergeQueue entries 조회로 (pr_number, pr_head_sha) 확정.
    _gh_graphql 사용. 파싱 실패/빈 결과면 [] 반환."""
    owner, name = (repo.split("/", 1) + [""])[:2]
    if not owner or not name:
        return []

    query = """
query($owner: String!, $repo: String!) {
  repository(owner: $owner, name: $repo) {
    mergeQueue {
      entries(first: 50) {
        nodes {
          pullRequest {
            number
            headRefOid
          }
        }
      }
    }
  }
}
"""
    variables = {"owner": owner, "repo": name}
    rc, data = _gh_graphql(query, variables)
    if rc != 0 or not data:
        return []

    try:
        nodes = (
            data["data"]["repository"]["mergeQueue"]["entries"]["nodes"]
        )
    except (KeyError, TypeError):
        return []

    if not isinstance(nodes, list):
        return []

    results: list[dict] = []
    for node in nodes:
        try:
            pr = node["pullRequest"]
            pr_number = pr["number"]
            pr_head_sha = pr["headRefOid"]
        except (KeyError, TypeError):
            continue
        if isinstance(pr_number, int) and pr_number > 0 and isinstance(pr_head_sha, str) and pr_head_sha.strip():
            results.append({"pr_number": pr_number, "pr_head_sha": pr_head_sha})

    return _dedupe(results)


def resolve_via_event_context(event: dict) -> list[dict]:
    """2차 보조: event.get("merge_group") 의 head_sha/base_sha/head_ref 로 단일 PR 보조 확인.
    head_ref "gh-readonly-queue/main/pr-<N>-<sha>" 에서 pr 번호 추출 시도.
    ★ 주의: event_context 만으로 pr_head_sha(원본 head)를 확정하기 어려우면 [] 반환(fail-closed 유도).
      단, event 에 명시적 pr_head_sha 정보가 있으면 사용. 없으면 []."""
    if not isinstance(event, dict):
        return []

    mg = event.get("merge_group")
    if not isinstance(mg, dict):
        return []

    head_ref = mg.get("head_ref") or ""

    # head_ref 패턴: "gh-readonly-queue/main/pr-<N>-<sha>"
    # 여기서 <sha>는 base_sha이지 PR의 original head_sha가 아님
    # → pr_head_sha를 event에서 직접 확정할 수 없으므로 [] 반환 (fail-closed)
    # 단, event에 명시적으로 pr_head_sha 필드가 있는 경우만 사용
    pr_head_sha = mg.get("pr_head_sha")
    if not isinstance(pr_head_sha, str) or not pr_head_sha.strip():
        # pr_head_sha 확정 불가 → fail-closed
        return []

    # pr_number 추출 시도
    pr_number = mg.get("pr_number")
    if not isinstance(pr_number, int) or pr_number <= 0:
        # head_ref에서 파싱 시도: gh-readonly-queue/<base>/pr-<N>-<sha>
        m = re.search(r"/pr-(\d+)-", head_ref)
        if m:
            try:
                pr_number = int(m.group(1))
            except (ValueError, IndexError):
                return []
        else:
            return []

    if pr_number > 0 and pr_head_sha.strip():
        return [{"pr_number": pr_number, "pr_head_sha": pr_head_sha}]
    return []


def resolve_via_ref_and_trailers(
    repo: str,
    head_ref: str | None,
    base_sha: str | None,
    merge_group_sha: str | None,
) -> list[dict]:
    """3차 fallback only: head_ref 의 "pr-<N>-<base_sha>" 정규식 파싱으로 PR 번호 추출,
    그 PR 번호로 _gh_api(f"repos/{repo}/pulls/{N}") 조회하여 head.sha=pr_head_sha 확보.
    base_sha..merge_group_sha commit trailer "(#N)" 파싱도 보조. 확정 실패 시 []."""
    results: list[dict] = []
    candidate_pr_numbers: list[int] = []

    # head_ref에서 PR 번호 추출: gh-readonly-queue/<base>/pr-<N>-<sha>
    if head_ref:
        m = re.search(r"/pr-(\d+)-", head_ref)
        if m:
            try:
                n = int(m.group(1))
                if n > 0:
                    candidate_pr_numbers.append(n)
            except (ValueError, IndexError):
                pass

    # commit trailers "(#N)" 파싱 보조
    if base_sha and merge_group_sha and repo:
        try:
            proc = subprocess.run(
                ["gh", "api", f"repos/{repo}/compare/{base_sha}...{merge_group_sha}"],
                capture_output=True, text=True, timeout=30
            )
            if proc.returncode == 0 and proc.stdout.strip():
                compare_data = json.loads(proc.stdout)
                commits = compare_data.get("commits", [])
                for commit in commits:
                    msg = commit.get("commit", {}).get("message", "")
                    # "(#N)" 패턴 파싱
                    for m in re.finditer(r"\(#(\d+)\)", msg):
                        try:
                            n = int(m.group(1))
                            if n > 0 and n not in candidate_pr_numbers:
                                candidate_pr_numbers.append(n)
                        except (ValueError, IndexError):
                            pass
        except Exception:
            pass

    # 각 PR 번호에 대해 API로 head.sha 확보
    for pr_number in candidate_pr_numbers:
        rc, pr_data = _gh_api(f"repos/{repo}/pulls/{pr_number}")
        if rc != 0 or not isinstance(pr_data, dict):
            continue
        try:
            pr_head_sha = pr_data["head"]["sha"]
        except (KeyError, TypeError):
            continue
        if isinstance(pr_head_sha, str) and pr_head_sha.strip():
            results.append({"pr_number": pr_number, "pr_head_sha": pr_head_sha})

    return _dedupe(results)


def resolve_merge_group_prs(
    repo: str,
    merge_group_sha: str,
    base_sha: str | None = None,
    head_ref: str | None = None,
    event: dict | None = None,
) -> list[dict]:
    """순서: resolve_via_api → resolve_via_event_context → resolve_via_ref_and_trailers.
    각 tier 결과를 _accept()로 검증: 모든 entry가 pr_number(int>0) AND pr_head_sha(비어있지 않은 str)이면 accept.
    하나라도 불완전하면 그 tier는 거부(fall through) — 부분 수용 금지(fail-closed).
    세 tier 모두 실패하면 raise MergeGroupPRsUnresolved.
    ★ 1차(API)가 accept되면 이후 tier 절대 호출 안 함(1차 신뢰원 우선)."""
    _diag_tokens = _reset_gh_diag()
    try:
        tier_diags: list = []

        def _run_tier(name, invoked, fn):
            """tier 실행 + diagnostics 누적. (entries, accepted) 반환."""
            if not invoked:
                tier_diags.append({"tier": name, "invoked": False, "entry_count": 0,
                                   "accepted": False, "reject_reason": "not_invoked", "gh": []})
                return [], False
            _GH_DIAG_TIER.set(name)
            gh_start = len(_GH_DIAG_SINK.get() or [])
            try:
                entries = fn()
                exc = None
            except Exception as e:  # 방어적: tier 함수는 통상 raise 안 함(내부 try/except). 예외 시 fail-closed 기록.
                entries = []
                exc = e
            gh_slice = list((_GH_DIAG_SINK.get() or [])[gh_start:])
            accepted = _accept(entries)
            if accepted:
                reject_reason = None
            elif exc is not None:
                reject_reason = "exception"
            elif not entries:
                reject_reason = "empty"
            else:
                reject_reason = "incomplete_entry"
            tier_diags.append({
                "tier": name,
                "invoked": True,
                "entry_count": len(entries) if isinstance(entries, list) else 0,
                "accepted": accepted,
                "reject_reason": reject_reason,
                "gh": gh_slice,
            })
            return entries, accepted

        # 1차: GraphQL API
        tier1, ok1 = _run_tier("api", True, lambda: resolve_via_api(repo, merge_group_sha, base_sha))
        if ok1:
            return _dedupe(tier1)

        # 2차: event context (event 있을 때만 invoked)
        tier2, ok2 = _run_tier("event", event is not None,
                               lambda: resolve_via_event_context(event if event is not None else {}))
        if ok2:
            return _dedupe(tier2)

        # 3차: ref + trailers fallback
        tier3, ok3 = _run_tier("ref", True,
                               lambda: resolve_via_ref_and_trailers(repo, head_ref, base_sha, merge_group_sha))
        if ok3:
            return _dedupe(tier3)

        exc = MergeGroupPRsUnresolved(
            f"Failed to resolve PRs for merge_group_sha={merge_group_sha!r} "
            f"in repo={repo!r} after all 3 tiers."
        )
        exc.diagnostics = tier_diags   # list of tier dicts (context 밖으로 materialize됨)
        exc.repo = repo
        exc.merge_group_sha = merge_group_sha
        raise exc
    finally:
        _restore_gh_diag(_diag_tokens)


def main() -> int:
    ap = argparse.ArgumentParser(
        description="merge_group_pr_resolver — merge_group 이벤트 PR 역추적"
    )
    ap.add_argument("--merge-group-sha", required=True, help="merge_group head SHA")
    ap.add_argument("--base-sha", default=None, help="merge_group base SHA")
    ap.add_argument("--head-ref", default=None, help="merge_group head_ref")
    ap.add_argument("--repo", default=DEFAULT_REPO, help="OWNER/REPO")
    ap.add_argument("--event-json", default=None, help="GitHub event JSON 파일 경로")
    args = ap.parse_args()

    event: dict | None = None
    if args.event_json:
        try:
            import pathlib
            event = json.loads(pathlib.Path(args.event_json).read_text(encoding="utf-8"))
        except Exception as e:
            print(json.dumps({"error": f"event-json load failed: {e}"}))
            return 3

    try:
        prs = resolve_merge_group_prs(
            repo=args.repo,
            merge_group_sha=args.merge_group_sha,
            base_sha=args.base_sha,
            head_ref=args.head_ref,
            event=event,
        )
        print(json.dumps({"resolved_prs": prs}, ensure_ascii=False))
        return 0
    except MergeGroupPRsUnresolved as e:
        print(json.dumps({"error": "MERGE_GROUP_PRS_UNRESOLVED", "detail": str(e)}))
        return 3


if __name__ == "__main__":
    import sys
    sys.exit(main())
