"""Tests for the WAITING_CONTINUE restart-recovery mechanism."""

from __future__ import annotations

import os
import subprocess
from pathlib import Path
from unittest.mock import MagicMock

import pytest

from dual_agent.domain import ErrorCategory, TaskRecord, TaskState
from dual_agent.persistence import TaskStore
from dual_agent.services import Orchestrator
from dual_agent.state_machine import InvalidTransition, TRANSITIONS


# ---------------------------------------------------------- helpers

FAKE_PLAN = '{"schema_version":1,"goal":"g","acceptance_criteria":["ok"],"steps":[{"schema_version":1,"id":"s1","action":"noop","verification":""}]}'
FAKE_IMPL = '{"schema_version":1,"files_changed":[],"tests_run":[],"notes":"fake"}'
FAKE_REVIEW = '{"schema_version":1,"verdict":"PASS","issues":[],"summary":"ok"}'


def _fake_adapter(responses: dict):
    from dual_agent.adapters.fake import FakeAgentAdapter
    return FakeAgentAdapter(responses)


def _make_orchestrator(tmp_path: Path) -> Orchestrator:
    store = TaskStore(tmp_path / ".dual-agent")
    planner = _fake_adapter({"PLAN": FAKE_PLAN, "REVIEW": FAKE_REVIEW, "FINAL_VERIFY": FAKE_REVIEW})
    implementer = _fake_adapter({"PLAN_REVIEW": FAKE_REVIEW, "IMPLEMENT": FAKE_IMPL, "FIX": FAKE_IMPL})

    workspace = MagicMock()
    workspace.create.return_value = (tmp_path / "worktree", "abc123")

    test_runner = MagicMock(return_value=True)

    return Orchestrator(
        store=store,
        planner=planner,
        implementer=implementer,
        workspace=workspace,
        test_runner=test_runner,
    )


def _make_task(orch: Orchestrator, tmp_path: Path) -> TaskRecord:
    from dual_agent.domain import TaskSpec
    spec = TaskSpec(repo_path=tmp_path, goal="test goal")
    return orch.create_task(spec)


def _git_repo(path: Path) -> str:
    path.mkdir(parents=True, exist_ok=True)
    subprocess.run(["git", "init", str(path)], check=True, capture_output=True)
    subprocess.run(["git", "-C", str(path), "config", "user.email", "test@example.com"], check=True)
    subprocess.run(["git", "-C", str(path), "config", "user.name", "Test"], check=True)
    (path / "tracked.txt").write_text("baseline\n", encoding="utf-8")
    subprocess.run(["git", "-C", str(path), "add", "tracked.txt"], check=True)
    subprocess.run(["git", "-C", str(path), "commit", "-m", "baseline"], check=True, capture_output=True)
    return subprocess.run(
        ["git", "-C", str(path), "rev-parse", "HEAD"], check=True, capture_output=True, text=True
    ).stdout.strip()


def test_needs_human_checkpoints_agent_changes_but_not_later_external_changes(tmp_path):
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    worktree = tmp_path / "real-worktree"
    baseline = _git_repo(worktree)
    task = orch._save(
        task,
        state=TaskState.FIX,
        worktree_path=str(worktree),
        baseline_commit=baseline,
        last_successful_stage="FINAL_VERIFY",
    )

    (worktree / "tracked.txt").write_text("agent fix\n", encoding="utf-8")
    stopped = orch._needs_human(task, "provider report invalid", ErrorCategory.PROTOCOL)

    assert stopped.last_checkpoint is not None
    assert stopped.last_checkpoint.stage == TaskState.NEEDS_HUMAN.value
    resumed = orch.recover(task.task_id)
    assert resumed.state is TaskState.FINAL_VERIFY

    stopped_again = orch._needs_human(resumed, "provider report invalid", ErrorCategory.PROTOCOL)
    (worktree / "tracked.txt").write_text("external edit\n", encoding="utf-8")
    refused = orch.recover(task.task_id)
    assert refused.state is TaskState.NEEDS_HUMAN
    assert refused.error == "workspace changed since the last checkpoint"


# ---------------------------------------------------------- state machine contract

def test_waiting_continue_not_in_transitions():
    """WAITING_CONTINUE must not appear as a source state in TRANSITIONS."""
    assert TaskState.WAITING_CONTINUE not in TRANSITIONS


def test_waiting_continue_not_terminal():
    """WAITING_CONTINUE is not a terminal state — it can be resumed."""
    from dual_agent.domain import is_terminal
    assert not is_terminal(TaskState.WAITING_CONTINUE)


# ---------------------------------------------------------- mark_waiting_after_restart

def test_mark_waiting_saves_original_state(tmp_path):
    """mark_waiting_after_restart saves the original state in resume_state then enters WAITING_CONTINUE."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    # Manually advance to PLAN so we have a non-INIT state to save
    task = orch._transition(task, "planning_started")
    assert task.state is TaskState.PLAN

    result = orch.mark_waiting_after_restart(task.task_id, task.state)

    assert result.state is TaskState.WAITING_CONTINUE
    assert result.resume_state is TaskState.PLAN


def test_mark_waiting_must_not_accept_direct_state_mutation(tmp_path):
    """Tests confirm we enter WAITING_CONTINUE only via mark_waiting_after_restart."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")

    result = orch.mark_waiting_after_restart(task.task_id, task.state)
    reloaded = orch.store.load(task.task_id)
    assert reloaded.state is TaskState.WAITING_CONTINUE
    assert reloaded.resume_state is TaskState.PLAN


# ---------------------------------------------------------- startup_recovery_scan idempotency

def test_startup_scan_skips_already_waiting(tmp_path):
    """A task already in WAITING_CONTINUE is not touched — no new event, resume_state unchanged."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")

    # First scan: enters WAITING_CONTINUE
    orch.startup_recovery_scan()
    after_first = orch.store.load(task.task_id)
    assert after_first.state is TaskState.WAITING_CONTINUE
    assert after_first.resume_state is TaskState.PLAN

    events_after_first = (orch.store.task_dir(task.task_id) / "events.jsonl").read_text()

    # Second scan: must leave the task completely unchanged
    orch.startup_recovery_scan()
    after_second = orch.store.load(task.task_id)
    assert after_second.state is TaskState.WAITING_CONTINUE
    assert after_second.resume_state is TaskState.PLAN

    events_after_second = (orch.store.task_dir(task.task_id) / "events.jsonl").read_text()
    assert events_after_first == events_after_second, "second scan must not append any events"


def test_startup_scan_idempotent_across_three_restarts(tmp_path):
    """Consecutive restarts leave resume_state pointing at the original pre-first-scan state."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")

    for _ in range(3):
        orch.startup_recovery_scan()

    final = orch.store.load(task.task_id)
    assert final.state is TaskState.WAITING_CONTINUE
    assert final.resume_state is TaskState.PLAN


def test_startup_scan_skips_terminal_tasks(tmp_path):
    """Terminal tasks (DONE, NEEDS_HUMAN) are left untouched by the startup scan."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    # Drive to NEEDS_HUMAN
    from dual_agent.domain import ErrorCategory
    done_task = orch._needs_human(task, "forced", ErrorCategory.OPERATOR)
    assert done_task.state is TaskState.NEEDS_HUMAN

    results = orch.startup_recovery_scan()
    assert all(t.state is TaskState.NEEDS_HUMAN for t in results if t.task_id == done_task.task_id)


def test_startup_scan_terminates_recorded_provider_before_waiting(tmp_path, monkeypatch):
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    task = orch._save(task, active_run_id="old-run", active_run_pid=4242, active_run_pids=[4242])
    alive = {4242}

    monkeypatch.setattr("dual_agent.services.is_running", lambda pid: pid in alive)

    def stop(pid):
        alive.discard(pid)
        return True

    monkeypatch.setattr("dual_agent.services.terminate_group", stop)
    orch.startup_recovery_scan()

    recovered = orch.store.load(task.task_id)
    assert recovered.state is TaskState.WAITING_CONTINUE
    assert recovered.resume_state is TaskState.PLAN
    assert recovered.active_run_id is None
    assert recovered.active_run_pid is None
    assert recovered.active_run_pids == []
    assert not alive


def test_resume_refuses_while_recorded_provider_is_alive(tmp_path, monkeypatch):
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    waiting = orch.mark_waiting_after_restart(task.task_id, task.state)
    orch._save(waiting, active_run_id="old-run", active_run_pid=4242, active_run_pids=[4242])
    monkeypatch.setattr("dual_agent.services.is_running", lambda pid: pid == 4242)

    with pytest.raises(InvalidTransition, match="still running"):
        orch.resume(task.task_id)

    assert orch.store.load(task.task_id).state is TaskState.WAITING_CONTINUE

# ---------------------------------------------------------- _advance with WAITING_CONTINUE

def test_advance_returns_unchanged_for_waiting(tmp_path):
    """_advance must return the task as-is when state is WAITING_CONTINUE."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    call_count = [0]
    original_agent_stage = orch._agent_stage

    def counting_agent_stage(t):
        call_count[0] += 1
        return original_agent_stage(t)

    orch._agent_stage = counting_agent_stage

    result = orch.advance(task.task_id)
    assert result.state is TaskState.WAITING_CONTINUE
    assert call_count[0] == 0, "_agent_stage must not be called in WAITING_CONTINUE"


def test_run_to_completion_stops_at_waiting(tmp_path):
    """run_to_completion treats WAITING_CONTINUE as a stop condition like approval gates."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    result = orch.run_to_completion(task.task_id)
    assert result.state is TaskState.WAITING_CONTINUE


# ---------------------------------------------------------- resume

def test_resume_restores_state(tmp_path):
    """resume must restore task.state from resume_state without touching TRANSITIONS."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    assert task.state is TaskState.PLAN

    orch.mark_waiting_after_restart(task.task_id, task.state)
    waiting = orch.store.load(task.task_id)
    assert waiting.state is TaskState.WAITING_CONTINUE
    assert waiting.resume_state is TaskState.PLAN

    resumed = orch.resume(task.task_id)
    assert resumed.state is TaskState.PLAN
    assert resumed.resume_state is None


def test_resume_clears_resume_state(tmp_path):
    """After resume, resume_state is None (consumed)."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    orch.resume(task.task_id)
    reloaded = orch.store.load(task.task_id)
    assert reloaded.resume_state is None


def test_resume_raises_when_not_waiting(tmp_path):
    """resume on a non-WAITING_CONTINUE task raises InvalidTransition (→ 409)."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    assert task.state is TaskState.INIT

    with pytest.raises(InvalidTransition):
        orch.resume(task.task_id)


def test_resume_raises_when_no_resume_state(tmp_path):
    """resume with resume_state=None raises InvalidTransition (→ 409)."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    # Force WAITING_CONTINUE without a resume_state (edge case / corrupt record)
    orch._save(task, state=TaskState.WAITING_CONTINUE, resume_state=None)

    with pytest.raises(InvalidTransition):
        orch.resume(task.task_id)


# ---------------------------------------------------------- zero provider calls

def test_advance_zero_provider_calls_in_waiting(tmp_path):
    """No adapter.run call happens while a task sits in WAITING_CONTINUE."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    planner_calls = []
    implementer_calls = []
    original_planner_run = orch.planner.run
    original_impl_run = orch.implementer.run

    def spy_planner(req, **kw):
        planner_calls.append(req)
        return original_planner_run(req, **kw)

    def spy_impl(req, **kw):
        implementer_calls.append(req)
        return original_impl_run(req, **kw)

    orch.planner.run = spy_planner
    orch.implementer.run = spy_impl

    for _ in range(3):
        orch.advance(task.task_id)

    assert len(planner_calls) == 0
    assert len(implementer_calls) == 0


# ---------------------------------------------------------- needs_human from WAITING_CONTINUE

def test_cancel_from_waiting_continue(tmp_path):
    """cancel() on a WAITING_CONTINUE task must transition to NEEDS_HUMAN without raising."""
    orch = _make_orchestrator(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    cancelled = orch.cancel(task.task_id)
    assert cancelled.state is TaskState.NEEDS_HUMAN


# ---------------------------------------------------------- API integration

def _make_client(tmp_path: Path):
    from starlette.testclient import TestClient
    from dual_agent.api import create_app
    orch = _make_orchestrator(tmp_path)
    return TestClient(create_app(orch)), orch


def test_resume_api_404_unknown_task(tmp_path):
    """POST /tasks/{id}/resume returns 404 for an unknown task id."""
    client, _ = _make_client(tmp_path)
    r = client.post("/tasks/no-such-task-id/resume")
    assert r.status_code == 404


def test_resume_api_409_when_not_waiting(tmp_path):
    """POST /tasks/{id}/resume returns 409 when task is not in WAITING_CONTINUE."""
    client, orch = _make_client(tmp_path)
    task = _make_task(orch, tmp_path)
    r = client.post(f"/tasks/{task.task_id}/resume")
    assert r.status_code == 409


def test_resume_api_recovers_and_runs_to_completion(tmp_path):
    """POST /tasks/{id}/resume restores state and immediately runs to a terminal state."""
    client, orch = _make_client(tmp_path)
    task = _make_task(orch, tmp_path)
    task = orch._transition(task, "planning_started")
    orch.mark_waiting_after_restart(task.task_id, task.state)

    r = client.post(f"/tasks/{task.task_id}/resume")
    assert r.status_code == 200

    reloaded = orch.store.load(task.task_id)
    assert reloaded.state is TaskState.DONE
    assert reloaded.resume_state is None


def test_resume_api_409_when_no_resume_state(tmp_path):
    """POST /tasks/{id}/resume returns 409 when WAITING_CONTINUE but resume_state is None."""
    client, orch = _make_client(tmp_path)
    task = _make_task(orch, tmp_path)
    orch._save(task, state=TaskState.WAITING_CONTINUE, resume_state=None)

    r = client.post(f"/tasks/{task.task_id}/resume")
    assert r.status_code == 409


# ---------------------------------------------------------- CLI permission checks

def test_env_file_0600_is_allowed(tmp_path, monkeypatch):
    """_check_env_file_permissions does not raise when the file is exactly 0600."""
    env_file = tmp_path / "env"
    env_file.write_text("KEY=value")
    env_file.chmod(0o600)
    monkeypatch.setenv("DUAL_AGENT_ENV_FILE", str(env_file))

    from dual_agent.cli import _check_env_file_permissions
    _check_env_file_permissions()


def test_env_file_too_open_raises_systemexit(tmp_path, monkeypatch):
    """_check_env_file_permissions raises SystemExit when permissions are wider than 0600."""
    env_file = tmp_path / "env"
    env_file.write_text("KEY=value")
    env_file.chmod(0o644)
    monkeypatch.setenv("DUAL_AGENT_ENV_FILE", str(env_file))

    from dual_agent.cli import _check_env_file_permissions
    with pytest.raises(SystemExit):
        _check_env_file_permissions()


def test_env_file_missing_is_skipped(tmp_path, monkeypatch):
    """_check_env_file_permissions is a no-op when the referenced file does not exist."""
    monkeypatch.setenv("DUAL_AGENT_ENV_FILE", str(tmp_path / "does-not-exist"))
    from dual_agent.cli import _check_env_file_permissions
    _check_env_file_permissions()


def test_env_file_var_unset_is_skipped(monkeypatch):
    """_check_env_file_permissions is a no-op when DUAL_AGENT_ENV_FILE is not set."""
    monkeypatch.delenv("DUAL_AGENT_ENV_FILE", raising=False)
    from dual_agent.cli import _check_env_file_permissions
    _check_env_file_permissions()


@pytest.mark.skipif(os.name == "nt", reason="chmod 0o604 semantics are POSIX-specific")
def test_env_file_0604_raises_systemexit(tmp_path, monkeypatch):
    """0604 (other-readable) also triggers the permission guard."""
    env_file = tmp_path / "env"
    env_file.write_text("KEY=value")
    env_file.chmod(0o604)
    monkeypatch.setenv("DUAL_AGENT_ENV_FILE", str(env_file))

    from dual_agent.cli import _check_env_file_permissions
    with pytest.raises(SystemExit):
        _check_env_file_permissions()


# ---------------------------------------------------------- Web static / behavior

def test_web_waiting_continue_label():
    """INDEX_HTML contains the Chinese '等待继续' label mapped to WAITING_CONTINUE."""
    from dual_agent.web import INDEX_HTML
    assert "WAITING_CONTINUE" in INDEX_HTML
    assert "等待继续" in INDEX_HTML


def test_web_step_button_hidden_when_waiting():
    """'推进一步' button is in the falsy branch of the waiting conditional."""
    from dual_agent.web import INDEX_HTML
    # Direct pattern check — avoids ambiguity with the static help-text occurrence
    assert "(waiting ? '' : '<button onclick=\"step()\">推进一步" in INDEX_HTML


def test_web_runall_button_hidden_when_waiting():
    """'自动跑到底' button is in the falsy branch of the waiting conditional."""
    from dual_agent.web import INDEX_HTML
    assert "(waiting ? '' : '<button onclick=\"runAll()\">自动跑到底" in INDEX_HTML


def test_web_resume_button_shown_when_waiting():
    """'继续执行' button appears and calls the resume API when state is WAITING_CONTINUE."""
    from dual_agent.web import INDEX_HTML
    # Python \'  in a triple-quoted source string → just ' in the actual string,
    # so the pattern in INDEX_HTML is act('resume') without backslashes.
    assert "继续执行" in INDEX_HTML
    assert "(waiting ? '<button class=\"gate\" onclick=\"act('resume')\">" in INDEX_HTML


# ---------------------------------------------------------- systemd unit static

def _service_sections() -> dict:
    """Parse deploy/dual-agent.service into {section: {key: value}}."""
    service_path = Path(__file__).parent.parent / "deploy" / "dual-agent.service"
    sections: dict = {}
    current: str | None = None
    for raw_line in service_path.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip()
        if not line or line.startswith("#"):
            continue
        if line.startswith("[") and line.endswith("]"):
            current = line[1:-1]
            sections.setdefault(current, {})
        elif "=" in line and current is not None:
            key, _, value = line.partition("=")
            sections[current].setdefault(key.strip(), value.strip().rstrip("\\").strip())
    return sections


def test_systemd_restart_sec_in_service():
    """RestartSec=10 is present in [Service]."""
    assert _service_sections().get("Service", {}).get("RestartSec") == "10"


def test_systemd_start_limit_in_unit_not_service():
    """StartLimitIntervalSec=300 and StartLimitBurst=5 must be in [Unit], not [Service]."""
    sections = _service_sections()
    unit = sections.get("Unit", {})
    svc = sections.get("Service", {})
    assert unit.get("StartLimitIntervalSec") == "300", "StartLimitIntervalSec must be in [Unit]"
    assert unit.get("StartLimitBurst") == "5", "StartLimitBurst must be in [Unit]"
    assert "StartLimitIntervalSec" not in svc, "StartLimitIntervalSec must not be in [Service]"
    assert "StartLimitBurst" not in svc, "StartLimitBurst must not be in [Service]"


def test_systemd_environment_file_present():
    """EnvironmentFile directive is in [Service] for primary credential loading."""
    assert "EnvironmentFile" in _service_sections().get("Service", {})


def test_systemd_exec_start_serve():
    """ExecStart must invoke the 'serve' command."""
    service_path = Path(__file__).parent.parent / "deploy" / "dual-agent.service"
    content = service_path.read_text(encoding="utf-8")
    assert "ExecStart=" in content
    assert "serve" in content





