#!/usr/bin/env python3
"""Dependency-free protocol test for verifiable tool-calling agents.

This is a deterministic state-machine experiment, not an LLM benchmark and not
a reproduction of ReAct/BFCL/ToolSandbox scores.  It makes one subtle failure
observable: a tool may commit a side effect and then time out.  Retrying with a
new idempotency key can duplicate the side effect even when the final answer is
correct; reusing the same key plus checking final state avoids that false pass.
"""

from __future__ import annotations

import argparse
import hashlib
import json
from dataclasses import asdict, dataclass, field
from typing import Any


SEED = 60
EXPECTED_SAMPLE_ID = "sample-007"
EXPECTED_VALUE = 7.4


@dataclass(frozen=True)
class ToolResult:
    ok: bool
    code: str
    observation: dict[str, Any]
    retryable: bool = False
    committed: bool = False


@dataclass
class LabState:
    samples: dict[str, str] = field(
        default_factory=lambda: {"S-7": EXPECTED_SAMPLE_ID}
    )
    assay_runs: list[dict[str, Any]] = field(default_factory=list)
    idempotency_cache: dict[str, dict[str, Any]] = field(default_factory=dict)


class LabEnvironment:
    """Small typed tool environment with injectable failure semantics."""

    def __init__(self, failure_mode: str = "none") -> None:
        self.state = LabState()
        self.failure_mode = failure_mode
        self.failure_consumed = False
        self.trace: list[dict[str, Any]] = []

    def _record(self, call: dict[str, Any], result: ToolResult) -> ToolResult:
        self.trace.append(
            {
                "step": len(self.trace) + 1,
                "call": call,
                "result": asdict(result),
                "assay_run_count": len(self.state.assay_runs),
            }
        )
        return result

    def execute(self, call: dict[str, Any]) -> ToolResult:
        if not isinstance(call, dict) or set(call) != {"name", "arguments"}:
            return self._record(
                call,
                ToolResult(False, "INVALID_CALL", {"error": "expected name+arguments"}),
            )
        name = call["name"]
        arguments = call["arguments"]
        if name not in {"find_sample", "run_assay", "read_assay"}:
            return self._record(
                call,
                ToolResult(False, "UNKNOWN_TOOL", {"error": f"unknown tool: {name}"}),
            )
        if not isinstance(arguments, dict):
            return self._record(
                call,
                ToolResult(False, "INVALID_ARGUMENT", {"error": "arguments must be an object"}),
            )
        if name == "find_sample":
            return self._find_sample(call, arguments)
        if name == "run_assay":
            return self._run_assay(call, arguments)
        return self._read_assay(call, arguments)

    def _find_sample(self, call: dict[str, Any], args: dict[str, Any]) -> ToolResult:
        if set(args) != {"label"} or not isinstance(args.get("label"), str):
            return self._record(
                call,
                ToolResult(False, "INVALID_ARGUMENT", {"error": "label must be a string"}),
            )
        sample_id = self.state.samples.get(args["label"])
        if sample_id is None:
            return self._record(
                call,
                ToolResult(False, "NOT_FOUND", {"error": "sample not found"}),
            )
        return self._record(call, ToolResult(True, "OK", {"sample_id": sample_id}))

    def _run_assay(self, call: dict[str, Any], args: dict[str, Any]) -> ToolResult:
        required = {"sample_id", "idempotency_key"}
        if set(args) != required or not all(isinstance(args.get(k), str) for k in required):
            return self._record(
                call,
                ToolResult(
                    False,
                    "INVALID_ARGUMENT",
                    {"error": "sample_id and idempotency_key must be strings"},
                ),
            )
        if args["sample_id"] != EXPECTED_SAMPLE_ID:
            return self._record(
                call,
                ToolResult(False, "NOT_FOUND", {"error": "sample_id not found"}),
            )
        key = args["idempotency_key"]
        if key in self.state.idempotency_cache:
            cached = self.state.idempotency_cache[key]
            return self._record(
                call,
                ToolResult(True, "IDEMPOTENT_REPLAY", cached, committed=True),
            )
        if self.failure_mode == "timeout_before_commit" and not self.failure_consumed:
            self.failure_consumed = True
            return self._record(
                call,
                ToolResult(
                    False,
                    "TIMEOUT_BEFORE_COMMIT",
                    {"error": "no state change"},
                    retryable=True,
                ),
            )
        run = {
            "run_id": f"run-{len(self.state.assay_runs) + 1:03d}",
            "sample_id": args["sample_id"],
            "value": EXPECTED_VALUE,
            "idempotency_key": key,
        }
        self.state.assay_runs.append(run)
        observation = {"run_id": run["run_id"], "value": run["value"]}
        self.state.idempotency_cache[key] = observation
        if self.failure_mode == "timeout_after_commit" and not self.failure_consumed:
            self.failure_consumed = True
            return self._record(
                call,
                ToolResult(
                    False,
                    "TIMEOUT_UNKNOWN_COMMIT",
                    {"error": "caller cannot know whether commit happened"},
                    retryable=True,
                    committed=True,
                ),
            )
        return self._record(call, ToolResult(True, "OK", observation, committed=True))

    def _read_assay(self, call: dict[str, Any], args: dict[str, Any]) -> ToolResult:
        if set(args) != {"sample_id"} or not isinstance(args.get("sample_id"), str):
            return self._record(
                call,
                ToolResult(False, "INVALID_ARGUMENT", {"error": "sample_id must be a string"}),
            )
        runs = [r for r in self.state.assay_runs if r["sample_id"] == args["sample_id"]]
        if not runs:
            return self._record(
                call,
                ToolResult(False, "NOT_FOUND", {"error": "no assay result"}),
            )
        return self._record(
            call,
            ToolResult(True, "OK", {"value": runs[-1]["value"], "runs": len(runs)}),
        )


def trace_digest(trace: list[dict[str, Any]]) -> str:
    payload = json.dumps(trace, sort_keys=True, separators=(",", ":"))
    return hashlib.sha256(payload.encode("utf-8")).hexdigest()[:12]


def verify(env: LabEnvironment, answer: dict[str, Any] | None) -> dict[str, Any]:
    answer_correct = answer == {"sample_id": EXPECTED_SAMPLE_ID, "value": EXPECTED_VALUE}
    exactly_one_side_effect = len(env.state.assay_runs) == 1
    state_correct = exactly_one_side_effect and all(
        run["sample_id"] == EXPECTED_SAMPLE_ID for run in env.state.assay_runs
    )
    return {
        "answer_correct": answer_correct,
        "state_correct": state_correct,
        "side_effect_count": len(env.state.assay_runs),
        "task_pass": answer_correct and state_correct,
    }


def run_timeout_episode(reuse_idempotency_key: bool) -> dict[str, Any]:
    env = LabEnvironment("timeout_after_commit")
    found = env.execute({"name": "find_sample", "arguments": {"label": "S-7"}})
    sample_id = found.observation["sample_id"]
    base_key = "task-060-assay"
    first = env.execute(
        {
            "name": "run_assay",
            "arguments": {"sample_id": sample_id, "idempotency_key": base_key},
        }
    )
    assert first.code == "TIMEOUT_UNKNOWN_COMMIT" and first.retryable
    retry_key = base_key if reuse_idempotency_key else base_key + "-retry"
    env.execute(
        {
            "name": "run_assay",
            "arguments": {"sample_id": sample_id, "idempotency_key": retry_key},
        }
    )
    read = env.execute({"name": "read_assay", "arguments": {"sample_id": sample_id}})
    answer = {"sample_id": sample_id, "value": read.observation["value"]}
    result = verify(env, answer)
    result.update(
        {
            "episode": "verified_retry" if reuse_idempotency_key else "blind_retry",
            "tool_calls": len(env.trace),
            "trace_digest": trace_digest(env.trace),
        }
    )
    return result


def run_schema_repair_episode() -> dict[str, Any]:
    env = LabEnvironment()
    bad = env.execute({"name": "find_sample", "arguments": {"label": 7}})
    assert bad.code == "INVALID_ARGUMENT" and not bad.retryable
    repaired = env.execute({"name": "find_sample", "arguments": {"label": "S-7"}})
    result = {
        "episode": "schema_repair",
        "invalid_calls": 1,
        "recovered": repaired.ok,
        "side_effect_count": len(env.state.assay_runs),
        "tool_calls": len(env.trace),
        "task_pass": repaired.ok and len(env.state.assay_runs) == 0,
        "trace_digest": trace_digest(env.trace),
    }
    return result


def run_loop_guard_episode(step_limit: int = 4) -> dict[str, Any]:
    env = LabEnvironment()
    for _ in range(step_limit):
        result = env.execute({"name": "invent_tool", "arguments": {}})
        assert result.code == "UNKNOWN_TOOL"
    return {
        "episode": "loop_guard",
        "force_terminated": True,
        "tool_calls": len(env.trace),
        "side_effect_count": len(env.state.assay_runs),
        "task_pass": len(env.trace) == step_limit and len(env.state.assay_runs) == 0,
        "trace_digest": trace_digest(env.trace),
    }


def run(check_only: bool) -> None:
    print(f"seed={SEED} deterministic_state_machine=true")
    episodes = [
        run_timeout_episode(reuse_idempotency_key=True),
        run_timeout_episode(reuse_idempotency_key=False),
        run_schema_repair_episode(),
        run_loop_guard_episode(),
    ]
    for episode in episodes:
        print(json.dumps(episode, ensure_ascii=False, sort_keys=True))

    verified, blind, repaired, guarded = episodes
    assert verified["task_pass"] and verified["side_effect_count"] == 1
    assert blind["answer_correct"] and not blind["state_correct"]
    assert blind["side_effect_count"] == 2 and not blind["task_pass"]
    assert repaired["task_pass"] and repaired["recovered"]
    assert guarded["task_pass"] and guarded["force_terminated"]
    if check_only:
        print("check-only: PASS")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--check-only", action="store_true")
    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()
    run(args.check_only)
