"""CLI — 스텝별 캐시 적중 실측 (읽기 전용).

축 A1(정적/동적 순서 재배열)의 효과 판정 SOT. 재배열 **전**과 **후**를
같은 제공사·같은 시간 간격·같은 호출 순서에서 한 번씩 재고, cached_tokens
의 변화만 재배열 효과로 인정한다. 문헌 수치는 기대값으로 쓰지 않는다.

    # 기준선 (재배열 전)
    python tools/opik_prompt_audit/cache_baseline.py \
        --since 2026-08-15T00:00:00 --label before
    # 재배열 후
    python tools/opik_prompt_audit/cache_baseline.py \
        --since 2026-08-16T00:00:00 --label after
    # 대조
    python tools/opik_prompt_audit/cache_baseline.py --compare before.json after.json

읽기만 한다 — Opik 에 쓰지 않고 저장소도 건드리지 않는다(출력 JSON 제외).
"""
import argparse
import json
import logging
import sys
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Optional

ROOT = Path(__file__).resolve().parent.parent.parent  # 저장소 루트
sys.path.insert(0, str(ROOT / "backend"))
sys.path.insert(0, str(ROOT))

logging.basicConfig(level=logging.INFO, format="%(message)s")
logger = logging.getLogger(__name__)


def _known_steps() -> set:
    """스텝 이름 사전 — trace 태그에서 스텝을 고를 때 쓴다."""
    from app.core.step_manifest import STEP_MANIFEST
    return set(STEP_MANIFEST)


def step_name(trace: Dict[str, Any], known: set,
              trace_index: Optional[Dict[str, Dict[str, Any]]] = None) -> str:
    """행(span/trace)의 스텝 이름.

    ★`trace_index` 를 주면 span 으로 보고 부모에서 얻는다(2026-08-24) —
    litellm span 이름은 절대 비지 않아 그것만 보면 부모를 영영 안 본다.

    1순위는 기존 감사 도구와 같은 `step_of` (metadata.trace_name) 다 —
    두 리포트의 스텝 이름이 갈리면 대조가 안 된다. trace_name 은 주행
    태그(run_tag) 가 있을 때만 실리므로, 없으면 태그 중 스텝 사전에 있는
    이름을 쓴다(`build_opik_metadata` 가 태그 첫 칸에 step_id 를 싣는다).
    """
    from tools.opik_prompt_audit.audit.fetch import step_of, step_of_call

    name = (step_of_call(trace, trace_index)
            if trace_index is not None else step_of(trace))
    if name in known:
        return name
    tags = trace.get("tags") or []
    hits = [t for t in tags if t in known]
    if len(hits) == 1:
        return hits[0]
    if len(hits) > 1:
        return sorted(hits)[0]
    return name


def collect(traces: List[Dict[str, Any]],
            trace_index: Optional[Dict[str, Dict[str, Any]]] = None,
            ) -> Dict[str, Any]:
    """스텝별 캐시 집계.

    한 **span** = 한 발송으로 센다(2026-08-24 — 감사 도구와 같은 단위).
    v1 호출은 trace 1 + span 1 로 남아 있어 둘을 합치면 두 배가 된다.
    usage 가 없는 행(이미지 생성 등)은 `calls_no_usage` 로 따로 세고
    토큰 합에서 뺀다.
    """
    from tools.opik_prompt_audit.audit.cache_usage import model_of, read_tokens

    known = _known_steps()
    rows: Dict[str, Dict[str, Any]] = {}
    for trace in traces:
        step = step_name(trace, known, trace_index)
        row = rows.setdefault(step, {
            "calls": 0,
            "calls_no_usage": 0,
            "calls_cache_reported": 0,
            "prompt_tokens": 0,
            "cached_tokens": 0,
            "cache_write_tokens": 0,
            "models": {},
        })
        row["calls"] += 1
        row["models"][model_of(trace)] = row["models"].get(model_of(trace), 0) + 1

        tok = read_tokens(trace)
        if tok["prompt_tokens"] is None:
            row["calls_no_usage"] += 1
            continue
        row["prompt_tokens"] += tok["prompt_tokens"]
        if tok["cached_tokens"] is not None:
            row["calls_cache_reported"] += 1
            row["cached_tokens"] += tok["cached_tokens"]
        if tok["cache_write_tokens"] is not None:
            row["cache_write_tokens"] += tok["cache_write_tokens"]

    for row in rows.values():
        billed = row["calls"] - row["calls_no_usage"]
        row["calls_with_usage"] = billed
        row["prompt_tokens_avg"] = (
            round(row["prompt_tokens"] / billed) if billed else 0)
        # 캐시 필드를 실은 호출이 하나도 없으면 적중률은 0% 가 아니라 미보고.
        row["cache_hit_rate"] = (
            round(row["cached_tokens"] / row["prompt_tokens"], 4)
            if row["calls_cache_reported"] and row["prompt_tokens"] else None)
        row["models"] = dict(
            sorted(row["models"].items(), key=lambda kv: -kv[1]))
    return dict(sorted(rows.items(), key=lambda kv: -kv[1]["prompt_tokens"]))


def _fmt_rate(rate: Optional[float]) -> str:
    return "미보고" if rate is None else f"{rate * 100:5.1f}%"


def print_table(payload: Dict[str, Any]) -> None:
    w = payload["window"]
    print(f"\n창: {w['since']} ~ {w['until']}  호출 {w.get('call_count', w.get('trace_count'))}건"
          f"  (수집 {w['collected_at']})")
    head = (f"{'스텝':28} {'호출':>5} {'캐시보고':>7} {'입력토큰':>10} "
            f"{'평균':>7} {'캐시토큰':>10} {'적중률':>7}  모델")
    print(head)
    print("-" * len(head))
    total_prompt = total_cached = 0
    for step, r in payload["steps"].items():
        print(f"{step[:28]:28} {r['calls_with_usage']:5d} "
              f"{r['calls_cache_reported']:7d} {r['prompt_tokens']:10d} "
              f"{r['prompt_tokens_avg']:7d} {r['cached_tokens']:10d} "
              f"{_fmt_rate(r['cache_hit_rate']):>7}  "
              f"{next(iter(r['models']), '?')}")
        total_prompt += r["prompt_tokens"]
        total_cached += r["cached_tokens"]
    print("-" * len(head))
    reported = sum(r["calls_cache_reported"] for r in payload["steps"].values())
    rate = (total_cached / total_prompt) if (reported and total_prompt) else None
    print(f"{'합계':28} {'':5} {reported:7d} {total_prompt:10d} {'':7} "
          f"{total_cached:10d} {_fmt_rate(rate):>7}")
    if not reported:
        print("\n★ 이 창의 어떤 호출도 캐시 토큰을 싣지 않았다 — 적중률 0% 가"
              " 아니라 '기록 없음'이다.\n"
              "  재배열 전/후 대조를 하려면 먼저 usage 에 캐시 항목이 실리는지"
              " 확인해야 한다(제공사 응답 또는 기록 경로).")


def print_compare(before: Dict[str, Any], after: Dict[str, Any]) -> None:
    print(f"\nbefore: {before['label']} ({before['window']['since']})"
          f"   after: {after['label']} ({after['window']['since']})")
    head = (f"{'스텝':28} {'적중률 before':>13} {'적중률 after':>12} "
            f"{'평균 입력 before':>15} {'평균 입력 after':>14}")
    print(head)
    print("-" * len(head))
    for step in sorted(set(before["steps"]) | set(after["steps"])):
        b = before["steps"].get(step)
        a = after["steps"].get(step)
        print(f"{step[:28]:28} "
              f"{(_fmt_rate(b['cache_hit_rate']) if b else '없음'):>13} "
              f"{(_fmt_rate(a['cache_hit_rate']) if a else '없음'):>12} "
              f"{(b['prompt_tokens_avg'] if b else 0):15d} "
              f"{(a['prompt_tokens_avg'] if a else 0):14d}")
    print("\n★ 두 창의 제공사·모델·호출 순서가 같을 때만 이 차이를 재배열"
          " 효과로 읽는다.")


def main() -> None:
    ap = argparse.ArgumentParser(description="스텝별 캐시 적중 실측")
    ap.add_argument("--since", help="ISO 시작 시각 (예: 2026-08-15T00:00:00)")
    ap.add_argument("--until", default="9999")
    ap.add_argument("--label", default="baseline", help="이 창의 이름")
    ap.add_argument("--out", default=None, help="출력 JSON 경로")
    ap.add_argument("--opik-url", default=None,
                    help="OPIK_URL_OVERRIDE 대신 쓸 주소")
    ap.add_argument("--workspace", default=None)
    ap.add_argument("--project", default=None)
    ap.add_argument("--compare", nargs=2, metavar=("BEFORE", "AFTER"),
                    help="이미 수집한 JSON 두 개를 대조만 한다")
    args = ap.parse_args()

    if args.compare:
        before = json.loads(Path(args.compare[0]).read_text("utf-8"))
        after = json.loads(Path(args.compare[1]).read_text("utf-8"))
        print_compare(before, after)
        return

    if not args.since:
        raise SystemExit("--since 가 필요하다 (또는 --compare)")

    import os
    os.chdir(ROOT / "backend")  # pydantic-settings 가 .env 를 cwd 에서 찾는다
    from app.core.config import settings

    base_url = (args.opik_url or settings.opik_url_override or "").rstrip("/")
    if not base_url:
        raise SystemExit(
            "OPIK 주소가 없다 — backend/.env 의 OPIK_URL_OVERRIDE 또는 "
            "--opik-url 로 준다 (셀프 호스팅 전용 도구)")
    workspace = args.workspace or settings.opik_workspace
    project = args.project or settings.opik_project_name

    from tools.opik_prompt_audit.audit.fetch import (
        fetch_spans, fetch_trace_index, select_call_rows,
    )

    # ★세는 것은 span 이다. 캐시 항목(usage)도 span 에 실려 있어 오히려 정확해진다.
    _trace_index = fetch_trace_index(
        base_url, workspace, project, args.since, args.until)
    _raw = fetch_spans(base_url, workspace, project, args.since, args.until)
    traces = select_call_rows(_raw, _trace_index)
    payload = {
        "label": args.label,
        "window": {
            "since": args.since,
            "until": args.until,
            "project": project,
            "call_count": len(traces),
            "raw_count": len(_raw),
            "excluded_test_rows": len(_raw) - len(traces),
            # 구 JSON 을 읽는 것이 있으면 이 별칭으로 찾는다.
            "trace_count": len(traces),
            "collected_at": datetime.now().isoformat(timespec="seconds"),
        },
        "steps": collect(traces, _trace_index),
    }

    out = Path(args.out) if args.out else (
        ROOT / "artifact" / f"{datetime.now():%Y%m%d}_prompt_cache_baseline"
        / f"cache_{args.label}.json")
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(
        json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
    print_table(payload)
    print(f"\nJSON: {out}")


if __name__ == "__main__":
    main()
