from __future__ import annotations

import json
from pathlib import Path
from unittest.mock import patch

import pytest

from adw_modules import agent_pi
from adw_modules.data_types import PiRequest


def _request(tmp_path: Path, session_id: str) -> PiRequest:
    return PiRequest(
        prompt="Build the feature",
        system_prompt="System",
        model="gpt/sol-web-pro",
        thinking="off",
        session_id=session_id,
        session_dir=str(tmp_path / "sessions"),
        raw_output_path=str(tmp_path / f"{session_id}.jsonl"),
        cwd=str(tmp_path),
    )


class InlineThread:
    def __init__(self, *, target, daemon: bool) -> None:
        self.target = target

    def start(self) -> None:
        self.target()

    def join(self, timeout=None) -> None:
        return None

    def is_alive(self) -> bool:
        return False


class FakeProcess:
    def __init__(self, stdout: list[str], stderr: list[str] | None = None, *, returncode: int = 0) -> None:
        self.pid = 1234
        self.stdout = iter(stdout)
        self.stderr = iter(stderr or [])
        self.returncode = returncode

    def poll(self):
        return self.returncode

    def wait(self, timeout=None):
        return self.returncode


def test_gpt_attempt_emits_client_session_and_estimated_usage(tmp_path: Path) -> None:
    message = json.dumps({
        "type": "message_end",
        "message": {
            "role": "assistant",
            "content": [{"type": "text", "text": "final"}],
            "usage": {
                "input": 20,
                "output": 40,
                "totalTokens": 60,
                "usage_estimated": True,
                "billing_status": "unavailable",
                "max_tokens": 16000,
            },
            "stopReason": "stop",
        },
    }) + "\n"
    events: list[dict] = []

    with patch.object(agent_pi, "resolve_model", return_value=("gpt", "sol-web-pro")), \
         patch.object(agent_pi, "context_window", return_value=120_000):
        result = agent_pi.run(
            _request(tmp_path, "session-a"),
            on_event=events.append,
            popen_factory=lambda *_args, **_kwargs: FakeProcess([message]),
            thread_factory=InlineThread,
        )

    start = next(event for event in events if event["type"] == "agent_attempt_start")
    end = next(event for event in events if event["type"] == "agent_attempt_end")
    assert start["clientSessionId"] == result.client_session_id
    assert end["usage"]["usage_estimated"] is True
    assert end["usage"]["billing_status"] == "unavailable"
    assert end["usage"]["max_tokens"] == 16000


def test_gpt_provider_failure_keeps_typed_cap_metadata(tmp_path: Path) -> None:
    stderr = json.dumps({
        "provider_failure": {
            "kind": "capped",
            "detail": "usage cap hit",
            "retry_after_seconds": 900,
            "resume_at": "2026-08-11T09:00:00Z",
        },
    }) + "\n"
    events: list[dict] = []

    with patch.object(agent_pi, "resolve_model", return_value=("gpt", "sol-web-pro")), \
         patch.object(agent_pi, "context_window", return_value=120_000):
        with pytest.raises(RuntimeError, match="provider failure \\[capped\\]"):
            agent_pi.run(
                _request(tmp_path, "session-b"),
                on_event=events.append,
                popen_factory=lambda *_args, **_kwargs: FakeProcess([], [stderr], returncode=1),
                thread_factory=InlineThread,
            )

    end = next(event for event in events if event["type"] == "agent_attempt_end")
    assert end["providerFailure"] == {
        "kind": "capped",
        "detail": "usage cap hit",
        "retry_after_seconds": 900,
        "resume_at": "2026-08-11T09:00:00Z",
    }
