"""MeteredLLM 成本仪表单测:计数 / 估价 / 分 tag / 滑动窗口 / 线程安全。"""

from __future__ import annotations

import threading

from genesis.runtime.metering import (
    DEFAULT_TAG,
    PRICE_IN_YUAN_PER_KTOKEN,
    PRICE_OUT_YUAN_PER_KTOKEN,
    WINDOW_SECONDS,
    MeteredLLM,
)


class FakeLLM:
    """桩客户端:记录收到的参数,返回固定文本。"""

    def __init__(self, reply: str = "好" * 2000) -> None:  # 2000 字 → 1000 token
        self.reply = reply
        self.seen_kwargs: list[dict] = []

    def complete(self, messages: list[dict], **kwargs) -> str:
        self.seen_kwargs.append(kwargs)
        return self.reply


class FakeClock:
    """手动拨表,测滑动窗口。"""

    def __init__(self) -> None:
        self.now = 0.0

    def __call__(self) -> float:
        return self.now


def _msgs(n_chars: int) -> list[dict]:
    return [{"role": "user", "content": "问" * n_chars}]


def test_calls_and_cost_estimate():
    fake = FakeLLM(reply="答" * 20000)                      # 输出 10000 token
    m = MeteredLLM(fake)
    out = m.complete(_msgs(20000))                          # 输入 10000 token
    assert out == fake.reply                                # 返回值透传
    s = m.stats()
    assert s["calls"] == 1
    expect = (10000 * PRICE_IN_YUAN_PER_KTOKEN + 10000 * PRICE_OUT_YUAN_PER_KTOKEN) / 1000
    assert s["est_cost_yuan"] == round(expect, 2)           # 0.01 + 0.02 = 0.03
    assert s["est_cost_yuan"] == 0.03


def test_tag_breakdown_and_tag_not_forwarded():
    fake = FakeLLM()
    m = MeteredLLM(fake)
    m.complete(_msgs(100), _tag="L1")
    m.complete(_msgs(100), _tag="narrator", temperature=0.3)
    m.complete(_msgs(100))                                  # 缺省归入 L2
    s = m.stats()
    assert s["calls"] == 3
    assert s["by_tag"]["L1"]["calls"] == 1
    assert s["by_tag"]["narrator"]["calls"] == 1
    assert s["by_tag"][DEFAULT_TAG]["calls"] == 1
    assert all(t["cost"] > 0 for t in s["by_tag"].values())
    # _tag 必须被 pop,其余 kwargs 正常透传
    assert all("_tag" not in kw for kw in fake.seen_kwargs)
    assert fake.seen_kwargs[1] == {"temperature": 0.3}


def test_sliding_window_calls_last_minute():
    clock = FakeClock()
    m = MeteredLLM(FakeLLM(), time_fn=clock)
    for _ in range(3):
        m.complete(_msgs(10))
    assert m.stats()["calls_last_minute"] == 3
    clock.now = WINDOW_SECONDS / 2                          # 窗口内再来 2 次
    m.complete(_msgs(10))
    m.complete(_msgs(10))
    assert m.stats()["calls_last_minute"] == 5
    clock.now = WINDOW_SECONDS + 1                          # 前 3 次滑出窗口
    assert m.stats()["calls_last_minute"] == 2
    assert m.stats()["calls"] == 5                          # 总数不受窗口影响


def test_thread_safety_two_threads():
    m = MeteredLLM(FakeLLM(reply="答" * 200))
    n_per_thread = 200

    def worker(tag: str) -> None:
        for _ in range(n_per_thread):
            m.complete(_msgs(200), _tag=tag)

    threads = [threading.Thread(target=worker, args=(t,)) for t in ("L1", "L2")]
    for t in threads:
        t.start()
    for t in threads:
        t.join()
    s = m.stats()
    assert s["calls"] == 2 * n_per_thread                   # 无丢计数
    assert s["by_tag"]["L1"]["calls"] == n_per_thread
    assert s["by_tag"]["L2"]["calls"] == n_per_thread
    per_call = (100 * PRICE_IN_YUAN_PER_KTOKEN + 100 * PRICE_OUT_YUAN_PER_KTOKEN) / 1000
    assert s["est_cost_yuan"] == round(2 * n_per_thread * per_call, 2)


def test_attribute_passthrough():
    fake = FakeLLM()
    fake.model = "deepseek-flash"
    m = MeteredLLM(fake)
    assert m.model == "deepseek-flash"
