"""나간 payload 가 무엇으로 채워졌나 — **전 스텝**을 같은 잣대로 잰다.

## 왜 필요한가

`card_static.py` 는 `scene_detail` 의 카드 하나만 본다. 그런데 프롬프트가
길어서 생기는 문제(문장끼리 모순되고 불필요한 것이 섞여 헛것이 늘어나는 것)는
그 스텝만의 것이 아니다. 어느 스텝이 얼마나 보내고, 그중 **매번 똑같은 것**이
얼마인지 한 잣대로 봐야 어디부터 손댈지 정할 수 있다.

## 무엇을 재나

스텝마다:

- **호출 수 · system/user 크기** — 나간 것 그대로
- **판(edition)** — system 을 sha 로 묶는다. 같은 스텝이 여러 판으로 돌았으면
  섞어 재면 안 된다
- **user 안에서 매번 똑같은 줄** — 샷이 달라도 안 바뀌는 줄은 샷별로 보낼
  이유가 없다. system 으로 올리거나 걷을 후보다
- **가장 긴 호출** — 한 번에 얼마나 밀어 넣었나

## 쓰는 법

    python tools/prompt_measure/payload_shape.py                  # 최근 주행
    python tools/prompt_measure/payload_shape.py --episode <id>    # 한 에피소드
    python tools/prompt_measure/payload_shape.py --step scene_detail

★돈이 안 든다 — 이미 쌓인 기록을 다시 읽을 뿐이다.
★Opik 이 SOT 다. 저장소의 프롬프트 파일이 아니라 **나간 것**이 진실이다.
"""
from __future__ import annotations

import argparse
import hashlib

import sys
from collections import Counter, defaultdict
from typing import Any, Dict, List, Tuple

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


def _msgs(span: Dict[str, Any]) -> List[Dict[str, Any]]:
    inp = span.get("input")
    if isinstance(inp, list):
        return [m for m in inp if isinstance(m, dict)]
    return (inp or {}).get("messages") or []


def _text(m: Dict[str, Any]) -> str:
    """메시지 하나의 글자 — 여러 조각으로 온 것도 이어 붙인다."""
    c = m.get("content")
    if isinstance(c, str):
        return c
    if isinstance(c, list):
        return "".join(p.get("text", "") for p in c
                       if isinstance(p, dict) and p.get("type") == "text")
    return ""


def _split(span: Dict[str, Any]) -> Tuple[str, str]:
    sys_parts, user_parts = [], []
    for m in _msgs(span):
        (sys_parts if m.get("role") == "system" else user_parts).append(_text(m))
    return "".join(sys_parts), "".join(user_parts)


def _step_of(span: Dict[str, Any]) -> str:
    for t in (span.get("tags") or []):
        if isinstance(t, str) and t.startswith("step:"):
            return t[5:]
    for t in (span.get("tags") or []):
        if isinstance(t, str) and t.startswith("op:"):
            return t[3:]
    return span.get("name") or "?"


def 반복되는_줄(users: List[str]) -> Tuple[int, int]:
    """user 메시지들에서 **모든 호출에 똑같이 나온 줄**의 글자 수.

    호출이 하나뿐이면 「매번 같다」를 말할 수 없으므로 0 을 돌려준다.
    """
    if len(users) < 2:
        return 0, 0
    줄집합 = [set(u.splitlines()) for u in users]
    공통 = set.intersection(*줄집합)
    공통 = {ln for ln in 공통 if ln.strip()}
    반복 = sum(len(ln) + 1 for ln in 공통)
    전체 = sum(len(u) for u in users) // len(users)
    return 반복, 전체


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--episode", default="")
    ap.add_argument("--step", default="")
    ap.add_argument("--since", default="2026-08-25T00:00:00")
    ap.add_argument("--until", default="2026-12-31T00:00:00")
    ap.add_argument("--top", type=int, default=25)
    a = ap.parse_args()

    B, W, P = opik_target()
    spans = fetch_spans(B, W, P, a.since, a.until)

    묶음: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
    for s in spans:
        md = s.get("metadata") or {}
        if a.episode and md.get("episode_id") != a.episode:
            continue
        step = _step_of(s)
        if a.step and step != a.step:
            continue
        if not _msgs(s):
            continue
        묶음[step].append(s)

    if not 묶음:
        print("해당하는 기록이 없다 — --since/--episode 를 확인할 것")
        return

    행 = []
    for step, ss in 묶음.items():
        판 = Counter()
        sys_len, user_len, users = [], [], []
        for s in ss:
            sy, us = _split(s)
            판[hashlib.sha256(sy.encode()).hexdigest()[:12]] += 1
            sys_len.append(len(sy))
            user_len.append(len(us))
            users.append(us)
        반복, user평균 = 반복되는_줄(users)
        행.append({
            "step": step, "호출": len(ss), "판수": len(판),
            "sys평균": sum(sys_len) // len(sys_len),
            "user평균": user평균 or (sum(user_len) // len(user_len)),
            "user최대": max(user_len),
            "반복": 반복,
        })
    행.sort(key=lambda r: -(r["sys평균"] + r["user평균"]))

    print(f"기록 {sum(r['호출'] for r in 행):,}건 · 스텝 {len(행)}개"
          f"{' · 에피소드 ' + a.episode[:8] if a.episode else ''}\n")
    print(f"{'스텝':<30}{'호출':>5}{'판':>4}{'system':>9}{'user':>9}"
          f"{'user최대':>9}{'매번같은줄':>11}")
    print("─" * 78)
    for r in 행[:a.top]:
        비율 = f"{r['반복'] / r['user평균'] * 100:.0f}%" if r["user평균"] else "-"
        print(f"{r['step']:<30}{r['호출']:>5}{r['판수']:>4}{r['sys평균']:>9,}"
              f"{r['user평균']:>9,}{r['user최대']:>9,}"
              f"{r['반복']:>8,} {비율:>3}")
    print("─" * 78)
    print("★ '매번같은줄' = user 안에서 호출이 달라도 안 바뀐 줄. 샷별로 보낼")
    print("  이유가 없으므로 system 으로 올리거나 걷을 후보다.")
    print("★ '판' 이 2 이상이면 그 스텝은 여러 판으로 돌았다 — 섞어 재면 안 된다.")


if __name__ == "__main__":
    main()
