from __future__ import annotations

import ast
import json
from contextlib import contextmanager
from pathlib import Path
from unittest.mock import MagicMock

import pytest

import adw_build_test
from adw_modules import agent_pi, agents, permissions
from adw_modules.data_types import (AgentConfig, BuildOutput, ConfigDefaults,
                                    ObservabilityConfig, Phase, PhaseParams,
                                    PiResult, PromptEngineering, QualityResult,
                                    SSSFConfig)


class _WorkflowRun:
    def __init__(self, tmp_path: Path, cfg: SSSFConfig) -> None:
        self.cfg = cfg
        self.adw_id = "gate-retry"
        self.engineer = "engineer"
        self.repo_root = tmp_path
        self.session_dir = tmp_path / "session"
        self.context_handoff_dir = self.session_dir / "context"
        self.session_dir.mkdir()
        self.agent_map: dict[str, dict] = {}
        self.tracer = MagicMock()
        self.tracer.event.return_value = "agent-start"
        self.console = MagicMock()
        self.phase_params: list[PhaseParams] = []
        self.add_usage = MagicMock()
        self.claim_paths = MagicMock()
        self.save_agent_map = MagicMock()

    @contextmanager
    def phase(self, params: PhaseParams):
        self.phase_params.append(params)
        phase = Phase(
            phase_id=f"phase-{len(self.phase_params)}",
            task_id=params.task_id,
            adw_id=self.adw_id,
            seq=len(self.phase_params),
            params=params,
        )
        yield _WorkflowPhase(self, phase)

    def finish(self, **_kwargs) -> int:
        return 0


class _WorkflowPhase:
    def __init__(self, run: _WorkflowRun, phase: Phase) -> None:
        self.run = run
        self.phase = phase

    def log(self, **_kwargs) -> None:
        pass

    def call(self, call):
        return agents.execute(self.run, self.phase, call, "baseline")


def _config(tmp_path: Path, pi_spawn: str = "local") -> SSSFConfig:
    system = tmp_path / "system.md"
    user = tmp_path / "user.md"
    system.write_text("Build system")
    user.write_text("{{prompt}}")
    return SSSFConfig(
        defaults=ConfigDefaults(data_dir=str(tmp_path / "data"), pi_spawn=pi_spawn),
        observability=ObservabilityConfig(db=str(tmp_path / "sssf.db")),
        agents=[AgentConfig(
            name="builder",
            prompt_engineering=PromptEngineering(system=str(system), user=str(user)),
        )],
    )


def _output(*changed_files: str) -> PiResult:
    return PiResult(text=json.dumps(BuildOutput(
        status="success",
        summary="implemented",
        changed_files=list(changed_files),
    ).model_dump()))


def _run_workflow(tmp_path: Path, monkeypatch, results: list[PiResult], pi_spawn: str = "local"):
    cfg = _config(tmp_path, pi_spawn)
    run = _WorkflowRun(tmp_path, cfg)
    requests = []
    spawn_paths = []

    def fake_pi(request, **_kwargs):
        requests.append(request)
        spawn_paths.append(_kwargs.get("spawn_path"))
        return results.pop(0)

    monkeypatch.chdir(tmp_path)
    monkeypatch.setattr(adw_build_test.agents, "load_config", lambda _path, _preset=None: cfg)
    monkeypatch.setattr(adw_build_test.agents, "validate", lambda *_args: None)
    monkeypatch.setattr(adw_build_test.session, "ensure", lambda *_args, **_kwargs: run)
    monkeypatch.setattr(adw_build_test.quality, "run_tests",
                        lambda _run: QualityResult(passed=True))
    monkeypatch.setattr(agent_pi, "run", fake_pi)
    monkeypatch.setattr(permissions, "snapshot", lambda _run: object())
    monkeypatch.setattr(permissions, "enforce", lambda *_args: [])
    return run, requests, spawn_paths


def test_initial_build_corrects_claim_gate_in_same_session(tmp_path: Path, monkeypatch) -> None:
    (tmp_path / "fixed.py").write_text("fixed\n")
    results = [_output("missing.py"), _output("fixed.py")]
    run, requests, _ = _run_workflow(tmp_path, monkeypatch, results)

    assert adw_build_test.main("implement it") == 0

    build = next(params for params in run.phase_params if params.name == "build")
    assert build.retries == 1
    assert len(requests) == 2
    assert requests[0].session_id == requests[1].session_id
    assert "missing.py" in requests[1].prompt
    assert "claimed changed file does not exist" in requests[1].prompt


def test_agent_requests_receive_configured_account(tmp_path: Path, monkeypatch) -> None:
    (tmp_path / "fixed.py").write_text("fixed\n")
    run, requests, _ = _run_workflow(tmp_path, monkeypatch, [_output("fixed.py")])
    run.cfg.account = "acct"
    monkeypatch.setattr(adw_build_test.utils, "operator_env", lambda *_args: {})

    assert adw_build_test.main("implement it") == 0

    assert [request.account for request in requests] == ["acct"]


def test_initial_build_gate_correction_is_bounded(tmp_path: Path, monkeypatch) -> None:
    results = [_output("missing.py"), _output("still-missing.py")]
    _run, requests, _ = _run_workflow(tmp_path, monkeypatch, results)

    with pytest.raises(agents.GateFailure, match="after 2 attempt"):
        adw_build_test.main("implement it")

    assert len(requests) == 2


def test_remote_pi_spawn_passes_the_installed_wrapper(tmp_path: Path, monkeypatch) -> None:
    (tmp_path / "fixed.py").write_text("fixed\n")
    run, _requests, spawn_paths = _run_workflow(
        tmp_path, monkeypatch, [_output("fixed.py")], pi_spawn="remote"
    )

    assert adw_build_test.main("implement it") == 0

    assert spawn_paths == [str(agents.PI_REMOTE)]
    assert run.cfg.defaults.pi_spawn == "remote"


def test_every_production_gated_agent_phase_allows_one_correction() -> None:
    factory_dir = Path(adw_build_test.__file__).parent
    invalid = []

    for path in sorted(factory_dir.glob("adw_*.py")):
        tree = ast.parse(path.read_text())
        for node in ast.walk(tree):
            if not isinstance(node, ast.With):
                continue
            phase_calls = [
                item.context_expr for item in node.items
                if isinstance(item.context_expr, ast.Call)
                and isinstance(item.context_expr.func, ast.Attribute)
                and item.context_expr.func.attr == "phase"
            ]
            if not phase_calls:
                continue
            params = phase_calls[0].args[0]
            if not (isinstance(params, ast.Call)
                    and isinstance(params.func, ast.Name)
                    and params.func.id == "PhaseParams"):
                continue
            gated = any(
                isinstance(call, ast.Call)
                and isinstance(call.func, ast.Name)
                and call.func.id == "AgentCall"
                and any(keyword.arg == "gates" for keyword in call.keywords)
                for call in ast.walk(node)
            )
            if not gated:
                continue
            keywords = {keyword.arg: keyword.value for keyword in params.keywords}
            retry = keywords.get("retries")
            if not (isinstance(retry, ast.Constant) and retry.value == 1):
                invalid.append(f"{path.name}:{node.lineno}")

    assert invalid == []
