"""Does the library do what it says? Three tests, one per module — the shapes from Days 2-5."""

from __future__ import annotations

import pytest

from llmutils import chunk, repair_json, retry, timed


def test_chunk_shares_overlap_and_keeps_the_last_partial_chunk() -> None:
    text = " ".join(f"t{i}" for i in range(10))

    out = chunk(text, 4, 1)
    assert len(out) == 3, out                       # stride 3 over 10 tokens
    assert out[0] == ["t0", "t1", "t2", "t3"]
    assert out[1][0] == out[0][3], out              # consecutive chunks share `overlap`
    assert out[-1] == ["t6", "t7", "t8", "t9"]

    # a short tail is still a chunk, never dropped
    assert [len(c) for c in chunk(" ".join(f"t{i}" for i in range(6)), 4, 1)] == [4, 3]
    assert chunk("   ", 4, 1) == []

    for size, overlap in [(0, 0), (3, 3), (3, 4), (3, -1)]:
        with pytest.raises(ValueError):
            chunk(text, size, overlap)


def test_repair_json_fixes_fences_quotes_commas_and_a_truncated_tail() -> None:
    assert repair_json('{"a": 1}') == {"a": 1}                       # valid passes through
    assert repair_json('```json\n{"a": [1, 2]}\n```') == {"a": [1, 2]}
    assert repair_json("{'a': 1, 'b': [2, 3]}") == {"a": 1, "b": [2, 3]}
    assert repair_json('{"a": 1, "b": [2, 3,],}') == {"a": 1, "b": [2, 3]}
    assert repair_json('{"a": [1, 2') == {"a": [1, 2]}               # truncated mid-array
    assert repair_json('{"a": "unfinis') == {"a": "unfinis"}         # truncated mid-string

    with pytest.raises(ValueError):
        repair_json("I could not produce JSON for that.")


def test_retry_re_raises_after_n_attempts_and_timed_reports_seconds() -> None:
    waits: list[float] = []
    attempts = 0

    @retry(3, backoff=0.1, on=(ValueError,), sleep=waits.append)
    def always_fails() -> str:
        nonlocal attempts
        attempts += 1
        raise ValueError("nope")

    with pytest.raises(ValueError):
        always_fails()
    assert attempts == 3                    # exactly `times` calls, not times + 1
    assert waits == [0.1, 0.2]              # one wait between attempts, doubling

    waits.clear()
    attempts = 0

    @retry(3, backoff=0.1, on=(ValueError,), sleep=waits.append)
    def wrong_error() -> str:
        nonlocal attempts
        attempts += 1
        raise KeyError("k")

    with pytest.raises(KeyError):
        wrong_error()
    assert attempts == 1 and waits == []    # a type not in `on` is not retried

    seen: list[tuple[str, float]] = []

    @timed(lambda name, seconds: seen.append((name, seconds)))
    def work(n: int) -> int:
        """double it"""
        return n * 2

    assert work(21) == 42
    assert work.__name__ == "work" and work.__doc__ == "double it"   # functools.wraps
    assert len(seen) == 1
    name, seconds = seen[0]
    assert name == "work" and isinstance(seconds, float) and seconds >= 0.0
