"""trace 의 usage 에서 프롬프트 토큰·캐시 적중 토큰을 읽는다 (읽기 전용).

Opik 은 usage 를 점 표기로 평평하게 눕혀 저장한다(`original_usage.prompt_
tokens` 등). 제공사마다 캐시 토큰의 이름이 다르므로 아래 후보를 순서대로
찾는다. **어느 후보도 없으면 "0" 이 아니라 "미보고"다** — 둘을 섞으면
적중률 0% 라는 거짓 기준선이 남는다.
"""
from typing import Any, Dict, Optional, Tuple

# 캐시로 재사용된 입력 토큰의 제공사별 이름 (우선순위 순).
CACHED_READ_KEYS = (
    "prompt_tokens_details.cached_tokens",   # OpenAI / litellm 표준
    "cache_read_input_tokens",               # Anthropic
    "cached_content_token_count",            # Gemini 원본
    "cached_tokens",
)
# 캐시에 새로 써 넣은 입력 토큰 (Anthropic 계열만 별도 청구).
CACHE_WRITE_KEYS = (
    "cache_creation_input_tokens",
    "cache_creation.ephemeral_5m_input_tokens",
)
PROMPT_TOKEN_KEYS = ("prompt_tokens", "input_tokens")

_ORIGINAL_PREFIX = "original_usage."


def flatten_usage(usage: Any) -> Dict[str, Any]:
    """중첩 dict 도 점 표기로 눕힌다 — Opik 저장 형태와 SDK 원형 둘 다 수용."""
    out: Dict[str, Any] = {}

    def walk(node: Any, prefix: str) -> None:
        if isinstance(node, dict):
            for k, v in node.items():
                walk(v, f"{prefix}{k}.")
            return
        if prefix:
            out[prefix[:-1]] = node

    walk(usage if isinstance(usage, dict) else {}, "")
    return out


def _pick(flat: Dict[str, Any], candidates: Tuple[str, ...]) -> Optional[int]:
    """후보 이름 중 먼저 잡히는 정수 값. `original_usage.` 접두는 벗겨 본다."""
    normalized: Dict[str, Any] = {}
    for key, value in flat.items():
        name = key[len(_ORIGINAL_PREFIX):] if key.startswith(_ORIGINAL_PREFIX) else key
        # 접두 없는 쪽이 있으면 그것을 남긴다 (Opik 정규화 값 우선).
        if name not in normalized or not key.startswith(_ORIGINAL_PREFIX):
            normalized[name] = value
    for cand in candidates:
        value = normalized.get(cand)
        if isinstance(value, bool):
            continue
        if isinstance(value, (int, float)):
            return int(value)
    return None


def read_tokens(trace: Dict[str, Any]) -> Dict[str, Optional[int]]:
    """한 trace 의 토큰 값 — 없는 항목은 None(미보고) 으로 남긴다."""
    flat = flatten_usage(trace.get("usage"))
    return {
        "prompt_tokens": _pick(flat, PROMPT_TOKEN_KEYS),
        "cached_tokens": _pick(flat, CACHED_READ_KEYS),
        "cache_write_tokens": _pick(flat, CACHE_WRITE_KEYS),
    }


def model_of(trace: Dict[str, Any]) -> str:
    md = trace.get("metadata") or {}
    if isinstance(md, dict) and md.get("model"):
        return str(md["model"])
    return "?"
