"""cache_baseline 도구의 usage 판독 — "미보고"와 "0" 을 가른다.

이 구분이 도구의 전부다. 캐시 항목이 없는 호출을 0 으로 세면 적중률 0% 라는
거짓 기준선이 남고, 재배열 전/후 대조가 그 위에서 이뤄진다.
"""
import sys
from pathlib import Path

import pytest

ROOT = Path(__file__).resolve().parents[3]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

from tools.opik_prompt_audit.audit.cache_usage import (  # noqa: E402
    flatten_usage, read_tokens,
)
from tools.opik_prompt_audit.cache_baseline import collect  # noqa: E402


def _trace(usage, step="entity_t2i", model="gemini/gemini-3.1-pro-preview"):
    return {
        "usage": usage,
        "tags": [step],
        "metadata": {"model": model, "trace_name": f"프 > 에 > {step}"},
    }


def test_flatten_handles_nested_and_dotted_forms():
    nested = {"prompt_tokens": 10, "prompt_tokens_details": {"cached_tokens": 4}}
    dotted = {"prompt_tokens": 10, "prompt_tokens_details.cached_tokens": 4}
    assert flatten_usage(nested) == flatten_usage(dotted)


@pytest.mark.parametrize("key", [
    "prompt_tokens_details.cached_tokens",
    "original_usage.prompt_tokens_details.cached_tokens",
    "cache_read_input_tokens",
    "cached_content_token_count",
])
def test_reads_provider_specific_cache_names(key):
    tok = read_tokens(_trace({"prompt_tokens": 100, key: 40}))
    assert tok == {
        "prompt_tokens": 100, "cached_tokens": 40, "cache_write_tokens": None}


def test_missing_cache_field_is_none_not_zero():
    assert read_tokens(_trace({"prompt_tokens": 100}))["cached_tokens"] is None
    assert read_tokens(_trace({"prompt_tokens": 100, "prompt_tokens_details.cached_tokens": 0}))[
        "cached_tokens"] == 0


def test_hit_rate_is_none_when_nothing_reported():
    rows = collect([_trace({"prompt_tokens": 100}), _trace({"prompt_tokens": 300})])
    row = rows["entity_t2i"]
    assert row["calls_with_usage"] == 2
    assert row["calls_cache_reported"] == 0
    assert row["cache_hit_rate"] is None      # 0% 가 아니다
    assert row["prompt_tokens_avg"] == 200


def test_hit_rate_counts_only_reported_calls():
    rows = collect([
        _trace({"prompt_tokens": 100, "prompt_tokens_details.cached_tokens": 40}),
        _trace({"prompt_tokens": 100, "prompt_tokens_details.cached_tokens": 0}),
    ])
    row = rows["entity_t2i"]
    assert row["calls_cache_reported"] == 2
    assert row["cached_tokens"] == 40
    assert row["cache_hit_rate"] == 0.2


def test_traces_without_usage_are_excluded_from_averages():
    rows = collect([_trace({"prompt_tokens": 100}), _trace(None)])
    row = rows["entity_t2i"]
    assert row["calls"] == 2
    assert row["calls_no_usage"] == 1
    assert row["prompt_tokens_avg"] == 100


def test_step_falls_back_to_tag_when_trace_name_missing():
    trace = _trace({"prompt_tokens": 10})
    trace["metadata"].pop("trace_name")
    trace["name"] = "chat.completion"
    rows = collect([trace])
    assert "entity_t2i" in rows
