from __future__ import annotations

import json
import signal
import subprocess
from pathlib import Path
from unittest.mock import MagicMock, call, patch

import pytest

from adw_modules import agent_pi
from adw_modules.data_types import PiRequest


@pytest.fixture(autouse=True)
def clear_catalog_cache():
    agent_pi._pi_catalog.cache_clear()
    yield
    agent_pi._pi_catalog.cache_clear()


def test_pi_catalog_parses_merged_catalog_compact_counts_and_account(monkeypatch) -> None:
    env = {"PI_CODING_AGENT_DIR": "/accounts/work"}
    completed = subprocess.CompletedProcess(
        ["pi", "--list-models"],
        0,
        "Provider Model Context\n"
        "anthropic claude-sonnet 272K\n"
        "openai gpt-5.6-terra 1.5M\n"
        "invalid row\n"
        "broken bad-count nope\n",
        "",
    )
    operator_env = MagicMock(return_value=env)
    cleanup = MagicMock()
    run = MagicMock(return_value=completed)
    monkeypatch.setattr(agent_pi, "operator_env", operator_env)
    monkeypatch.setattr(agent_pi, "cleanup_operator_env", cleanup)
    monkeypatch.setattr(agent_pi.subprocess, "run", run)

    assert agent_pi._pi_catalog("work") == [
        ("anthropic", "claude-sonnet", 272_000),
        ("openai", "gpt-5.6-terra", 1_500_000),
    ]

    assert operator_env.mock_calls == [call("work")]
    assert cleanup.mock_calls == [call(env)]
    assert run.mock_calls == [call(
        [agent_pi.PI_PATH, "--list-models"],
        capture_output=True,
        text=True,
        timeout=30,
        env=env,
        check=False,
    )]


@pytest.mark.parametrize("outcome", [
    subprocess.CompletedProcess(["pi", "--list-models"], 1, "", "failed"),
    OSError("missing pi"),
    subprocess.TimeoutExpired("pi", 30),
])
def test_pi_catalog_cleans_account_environment_on_failures(monkeypatch, outcome) -> None:
    env = {"PI_CODING_AGENT_DIR": "/accounts/failure"}
    cleanup = MagicMock()
    monkeypatch.setattr(agent_pi, "operator_env", MagicMock(return_value=env))
    monkeypatch.setattr(agent_pi, "cleanup_operator_env", cleanup)
    if isinstance(outcome, subprocess.CompletedProcess):
        monkeypatch.setattr(agent_pi.subprocess, "run", MagicMock(return_value=outcome))
    else:
        def raise_outcome(*_args, **_kwargs):
            raise outcome

        monkeypatch.setattr(agent_pi.subprocess, "run", raise_outcome)

    assert agent_pi._pi_catalog("failure") == []

    assert cleanup.mock_calls == [call(env)]


def test_resolve_model_uses_exact_fully_qualified_match(monkeypatch) -> None:
    monkeypatch.setattr(agent_pi, "_pi_catalog", lambda account: [
        ("openai", "gpt-5.6-terra", 272_000),
        ("openrouter", "gpt-5.6-terra", 272_000),
    ])

    assert agent_pi.resolve_model("openai/gpt-5.6-terra", "acct") == (
        "openai", "gpt-5.6-terra",
    )


def test_resolve_model_uses_unique_short_name(monkeypatch) -> None:
    monkeypatch.setattr(agent_pi, "_pi_catalog", lambda account: [
        ("anthropic", "claude-sonnet-4-5", 200_000),
        ("openai", "gpt-5.6-terra", 272_000),
    ])

    assert agent_pi.resolve_model("sonnet-4-5", "acct") == (
        "anthropic", "claude-sonnet-4-5",
    )


@pytest.mark.parametrize(("pattern", "message"), [
    ("gpt-5.6-terra", "ambiguous"),
    ("missing", "not found"),
])
def test_resolve_model_rejects_ambiguous_and_missing_patterns(monkeypatch, pattern, message) -> None:
    monkeypatch.setattr(agent_pi, "_pi_catalog", lambda account: [
        ("openai", "gpt-5.6-terra", 272_000),
        ("openrouter", "gpt-5.6-terra", 272_000),
    ])

    with pytest.raises(ValueError, match=message):
        agent_pi.resolve_model(pattern, "acct")


def test_context_window_uses_account_registry_then_catalog_fallback(monkeypatch, tmp_path: Path) -> None:
    account_dir = tmp_path / "account"
    account_dir.mkdir()
    (account_dir / "models.json").write_text(json.dumps({
        "providers": {
            "openai": {"models": [{"id": "registered", "contextWindow": 123_456}]},
        },
    }))
    env = {"PI_CODING_AGENT_DIR": str(account_dir)}
    cleanup = MagicMock()
    catalog = MagicMock(return_value=[("openai", "fallback", 272_000)])
    monkeypatch.setattr(agent_pi, "operator_env", MagicMock(return_value=env))
    monkeypatch.setattr(agent_pi, "cleanup_operator_env", cleanup)
    monkeypatch.setattr(agent_pi, "_pi_catalog", catalog)

    assert agent_pi.context_window("openai", "registered", "acct") == 123_456
    assert agent_pi.context_window("openai", "fallback", "acct") == 272_000
    assert agent_pi.context_window("openai", "missing", "acct") == 0

    assert cleanup.mock_calls == [call(env), call(env), call(env)]
    assert catalog.mock_calls == [call("acct"), call("acct")]


def _request(tmp_path: Path, *, model: str = "gpt/sol-web-pro") -> PiRequest:
    return PiRequest(
        prompt="prompt",
        system_prompt="system",
        model=model,
        session_id="session-1",
        session_dir=str(tmp_path / "sessions"),
        raw_output_path=str(tmp_path / "raw.jsonl"),
        cwd=str(tmp_path),
    )


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

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

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

    def is_alive(self) -> bool:
        return False


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

    def poll(self) -> int | None:
        return self.returncode

    def wait(self, timeout: float | None = None) -> int:
        return self.returncode


def test_derive_client_session_id_is_stable_and_valid() -> None:
    derived = agent_pi.derive_client_session_id("session-1")
    assert derived == agent_pi.derive_client_session_id("session-1")
    assert derived != agent_pi.derive_client_session_id("session-2")
    assert agent_pi.CLIENT_SESSION_PATTERN.fullmatch(derived)


def test_blank_session_id_fails_before_spawn() -> None:
    with pytest.raises(ValueError, match="non-empty string"):
        agent_pi.derive_client_session_id("   ")


def test_run_sets_child_client_session_env(tmp_path: Path) -> None:
    captured_env: dict[str, str] = {}
    message = json.dumps({
        "type": "message_end",
        "message": {
            "role": "assistant",
            "content": [{"type": "text", "text": "{\"status\":\"success\",\"summary\":\"ok\"}"}],
            "usage": {"input": 1, "output": 1, "totalTokens": 2},
            "stopReason": "stop",
        },
    }) + "\n"

    def popen_factory(_cmd, **kwargs):
        captured_env.update(kwargs["env"])
        return FakeProcess([message])

    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),
            popen_factory=popen_factory,
            thread_factory=InlineThread,
        )

    assert captured_env["OVERDECK_PI_CLIENT_SESSION"] == result.client_session_id
    assert result.client_session_id is not None
    assert result.provider == "gpt"


def test_gpt_usage_is_estimated_and_not_billed(tmp_path: Path) -> None:
    message = json.dumps({
        "type": "message_end",
        "message": {
            "role": "assistant",
            "content": [{"type": "text", "text": "done"}],
            "usage": {
                "input": 10,
                "output": 20,
                "totalTokens": 30,
                "usage_estimated": True,
                "billing_status": "unavailable",
                "max_tokens": 16000,
                "cost": {"total": 9.99, "input": 4.5, "output": 5.49},
            },
            "stopReason": "stop",
        },
    }) + "\n"

    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),
            popen_factory=lambda *_args, **_kwargs: FakeProcess([message]),
            thread_factory=InlineThread,
        )

    assert result.tokens == 30
    assert result.cost == 0.0
    assert result.usage.usage_estimated is True
    assert result.usage.billing_status == "unavailable"
    assert result.max_tokens == 16_000


def test_provider_failure_is_extracted_from_stderr_json(tmp_path: Path) -> None:
    stderr = json.dumps({
        "providerFailure": {
            "kind": "capped",
            "detail": "daily limit reached",
            "retry_after_seconds": 900,
            "resume_at": "2026-08-11T09:00:00Z",
        },
    }) + "\n"
    receipts: 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),
                popen_factory=lambda *_args, **_kwargs: FakeProcess([], [stderr], returncode=7),
                thread_factory=InlineThread,
                on_attempt_end=lambda _attempt, info: receipts.append(info),
            )

    assert receipts[0]["provider_failure"] == {
        "kind": "capped",
        "detail": "daily limit reached",
        "retry_after_seconds": 900,
        "resume_at": "2026-08-11T09:00:00Z",
    }


def test_zero_exit_message_error_is_a_typed_provider_failure(tmp_path: Path) -> None:
    failure = (
        "Codex error: Your input exceeds the context window of this model. "
        "Please adjust your input and try again."
    )
    message = json.dumps({
        "type": "message_end",
        "message": {
            "role": "assistant",
            "content": [],
            "usage": {"input": 10, "output": 2, "totalTokens": 12},
            "stopReason": "error",
            "errorMessage": failure,
        },
    }) + "\n"
    receipts: list[dict] = []

    with patch.object(agent_pi, "resolve_model", return_value=("openai-codex", "gpt-5.3-codex-spark")), \
         patch.object(agent_pi, "context_window", return_value=128_000):
        with pytest.raises(RuntimeError, match="provider failure \\[context_overflow\\]"):
            agent_pi.run(
                _request(tmp_path),
                popen_factory=lambda *_args, **_kwargs: FakeProcess([message]),
                thread_factory=InlineThread,
                on_attempt_end=lambda _attempt, info: receipts.append(info),
            )

    assert receipts[0]["returncode"] == 0
    assert receipts[0]["provider_failure"] == {
        "kind": "context_overflow",
        "detail": failure,
        "retry_after_seconds": None,
        "resume_at": None,
    }


def test_cancelled_exit_records_cancelled_provider_failure(tmp_path: Path) -> None:
    receipts: 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="pi exited -15"):
            agent_pi.run(
                _request(tmp_path),
                popen_factory=lambda *_args, **_kwargs: FakeProcess([], [], returncode=-signal.SIGTERM),
                thread_factory=InlineThread,
                on_attempt_end=lambda _attempt, info: receipts.append(info),
            )

    assert receipts[0]["provider_failure"] == {
        "kind": "cancelled",
        "detail": "pi exited -15: ",
        "retry_after_seconds": None,
        "resume_at": None,
    }
