#!/usr/bin/env python3
"""두 판(system sha)의 발송문을 줄 단위로 댄다 — **무엇이 되돌아왔는가**.

`edition_hash.py` 는 판이 몇 개인지와 각 판의 크기를 준다. 그것만으로는
「다이어트가 되돌아왔다」까지만 알고 **무엇이** 되돌렸는지는 모른다.
이 도구는 두 판의 system 을 가져와 한쪽에만 있는 줄을 보여 준다.

★파일이 아니라 **실제로 나간 문장**을 댄다. 팩을 읽어 대조하면 조립·주입
 뒤의 실물과 다르다(이 저장소가 여러 번 겪은 자리).

usage:
  edition_diff.py <sha_before> <sha_after> [--step scene_detail] [--full]
  edition_diff.py --list            # 판 목록만
"""
import argparse
import difflib
import hashlib
import sys
from collections import defaultdict

sys.path.insert(0, "/Users/manta/Documents/Projects/TheRoad-I1/scratchpad")
from _opik_env import opik_target  # noqa: E402

from tools.opik_prompt_audit.audit.fetch import fetch_spans  # noqa: E402

WINDOW = ("2026-08-24T00:00:00", "2026-08-27T00:00:00")


def systems_by_sha(step: str):
    base, ws, proj = opik_target()
    out = defaultdict(list)
    for sp in fetch_spans(base, ws, proj, *WINDOW):
        tags = sp.get("tags") or []
        if f"op:{step}" not in tags and f"step:{step}" not in tags:
            continue
        inp = sp.get("input")
        msgs = inp if isinstance(inp, list) else (inp or {}).get("messages")
        for m in (msgs or []):
            if not isinstance(m, dict) or m.get("role") != "system":
                continue
            c = m.get("content")
            if not isinstance(c, str):
                continue
            sha = hashlib.sha256(c.encode("utf-8")).hexdigest()[:12]
            out[sha].append((sp.get("start_time"), c))
    return out


def main():
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("before", nargs="?", default="")
    ap.add_argument("after", nargs="?", default="")
    ap.add_argument("--step", default="scene_detail")
    ap.add_argument("--list", action="store_true")
    ap.add_argument("--full", action="store_true",
                    help="줄을 자르지 않고 전부 보여 준다")
    args = ap.parse_args()

    by = systems_by_sha(args.step)
    if args.list or not (args.before and args.after):
        print(f"── {args.step} 판 {len(by)}개")
        for sha, rows in sorted(by.items(), key=lambda kv: kv[1][0][0]):
            print(f"   {sha}  {len(rows[0][1]):>7,}자  {len(rows):3d}컷  "
                  f"{rows[0][0][:19]}")
        return 0

    for sha in (args.before, args.after):
        if sha not in by:
            print(f"★ sha {sha} 를 못 찾았다 — --list 로 확인할 것")
            return 1
    a = by[args.before][0][1].splitlines()
    b = by[args.after][0][1].splitlines()
    print(f"── {args.step}  {args.before}({len(by[args.before][0][1]):,}자) "
          f"→ {args.after}({len(by[args.after][0][1]):,}자)")
    add = rm = 0
    add_ch = rm_ch = 0
    for ln in difflib.unified_diff(a, b, lineterm="", n=0):
        if ln.startswith("+++") or ln.startswith("---") or ln.startswith("@@"):
            continue
        if ln.startswith("+"):
            add += 1
            add_ch += len(ln) - 1
            print(f"  ＋ {ln[1:] if args.full else ln[1:151]}")
        elif ln.startswith("-"):
            rm += 1
            rm_ch += len(ln) - 1
            print(f"  － {ln[1:] if args.full else ln[1:151]}")
    print(f"\n── 더해진 줄 {add}개 {add_ch:,}자 · 빠진 줄 {rm}개 {rm_ch:,}자 "
          f"· 순증 {add_ch - rm_ch:+,}자")
    return 0


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