

from __future__ import annotations

import argparse
import copy
import csv
import json
import os
import random
import re
import sys
from dataclasses import asdict, dataclass, field
from typing import Any, Callable

ROOMS = {
    "standard": {"price": 120},
    "suite": {"price": 340},
    "deluxe": {"price": 210},
}

@dataclass
class Booking:
    confirmation: str
    room_type: str
    guest_name: str
    check_in: str
    check_out: str

    def as_state(self) -> dict[str, str]:

        return {
            "room_type": self.room_type.lower().strip(),
            "guest_name": self.guest_name.lower().strip(),
            "check_in": self.check_in.strip(),
            "check_out": self.check_out.strip(),
        }

@dataclass
class ToolEvent:
    index: int
    name: str
    arguments: dict[str, Any]
    result: dict[str, Any]
    after_correction: bool

class BookingEnv:

    def __init__(self) -> None:
        self.bookings: dict[str, Booking] = {}
        self.log: list[ToolEvent] = []
        self._counter = 0
        self._correction_issued = False

    def mark_correction_issued(self) -> None:
        self._correction_issued = True

    def _record(self, name: str, arguments: dict, result: dict) -> dict:
        self.log.append(
            ToolEvent(
                index=len(self.log),
                name=name,
                arguments=arguments,
                result=result,
                after_correction=self._correction_issued,
            )
        )
        return result

    def check_availability(self, room_type: str) -> dict:
        key = (room_type or "").lower().strip()
        room = ROOMS.get(key)
        if room is None:
            result = {"error": f"Unknown room type: {room_type}"}
        else:
            result = {"available": True, "room_type": key, "price": room["price"]}
        return self._record("check_availability", {"room_type": room_type}, result)

    def book_room(
        self, room_type: str, guest_name: str, check_in: str, check_out: str
    ) -> dict:
        args = {
            "room_type": room_type,
            "guest_name": guest_name,
            "check_in": check_in,
            "check_out": check_out,
        }
        key = (room_type or "").lower().strip()
        if key not in ROOMS:
            return self._record(
                "book_room", args, {"success": False, "reason": "Unknown room type"}
            )
        self._counter += 1
        conf = f"CONF{self._counter:04d}"
        self.bookings[conf] = Booking(conf, key, guest_name, check_in, check_out)
        result = {
            "success": True,
            "confirmation": conf,
            "room_type": key,
            "guest_name": guest_name,
            "check_in": check_in,
            "check_out": check_out,
        }
        return self._record("book_room", args, result)

    def modify_booking(self, confirmation: str, **changes: str) -> dict:
        args = {"confirmation": confirmation, **changes}
        booking = self.bookings.get(confirmation)
        if booking is None:
            return self._record(
                "modify_booking", args, {"success": False, "reason": "No such booking"}
            )
        for k, v in changes.items():
            if v is None:
                continue
            if k == "room_type":
                v = str(v).lower().strip()
                if v not in ROOMS:
                    return self._record(
                        "modify_booking",
                        args,
                        {"success": False, "reason": "Unknown room type"},
                    )
            if hasattr(booking, k):
                setattr(booking, k, v)
        return self._record(
            "modify_booking", args, {"success": True, **booking.as_state()}
        )

    def cancel_booking(self, confirmation: str) -> dict:
        args = {"confirmation": confirmation}
        if confirmation not in self.bookings:
            return self._record(
                "cancel_booking", args, {"success": False, "reason": "No such booking"}
            )
        del self.bookings[confirmation]
        return self._record("cancel_booking", args, {"success": True})

    def committed_state(self) -> list[dict[str, str]]:

        return sorted(
            (b.as_state() for b in self.bookings.values()),
            key=lambda d: json.dumps(d, sort_keys=True),
        )

    def clean_run_tool_calls(self) -> int:
        return sum(1 for e in self.log if not e.after_correction)

    def recovery_tool_calls(self) -> int:
        return sum(1 for e in self.log if e.after_correction)

    def log_as_dicts(self) -> list[dict]:
        return [asdict(e) for e in self.log]

TOOL_SCHEMAS = [
    {
        "type": "function",
        "function": {
            "name": "check_availability",
            "description": "Check whether a specific room type is available.",
            "parameters": {
                "type": "object",
                "properties": {
                    "room_type": {
                        "type": "string",
                        "description": "standard, suite, or deluxe",
                    }
                },
                "required": ["room_type"],
            },
        },
    },
    {
        "type": "function",
        "function": {
            "name": "book_room",
            "description": "Book a room for a guest. Creates a new booking.",
            "parameters": {
                "type": "object",
                "properties": {
                    "room_type": {"type": "string"},
                    "guest_name": {"type": "string"},
                    "check_in": {"type": "string"},
                    "check_out": {"type": "string"},
                },
                "required": ["room_type", "guest_name", "check_in", "check_out"],
            },
        },
    },
    {
        "type": "function",
        "function": {
            "name": "modify_booking",
            "description": "Change fields on an existing booking, identified by confirmation code.",
            "parameters": {
                "type": "object",
                "properties": {
                    "confirmation": {"type": "string"},
                    "room_type": {"type": "string"},
                    "guest_name": {"type": "string"},
                    "check_in": {"type": "string"},
                    "check_out": {"type": "string"},
                },
                "required": ["confirmation"],
            },
        },
    },
    {
        "type": "function",
        "function": {
            "name": "cancel_booking",
            "description": "Cancel an existing booking, identified by confirmation code.",
            "parameters": {
                "type": "object",
                "properties": {"confirmation": {"type": "string"}},
                "required": ["confirmation"],
            },
        },
    },
]

def dispatch(env: BookingEnv, name: str, args: dict) -> dict:

    if name == "check_availability":
        return env.check_availability(**args)
    if name == "book_room":
        return env.book_room(**args)
    if name == "modify_booking":
        conf = args.pop("confirmation", "")
        return env.modify_booking(conf, **args)
    if name == "cancel_booking":
        return env.cancel_booking(**args)
    return {"error": f"Unknown tool: {name}"}

BASE_TASKS = [
    {"task_id": "t01", "guest_name": "alex johnson", "init_room": "standard",
     "check_in": "2026-04-05", "check_out": "2026-04-08"},
    {"task_id": "t02", "guest_name": "morgan lee", "init_room": "deluxe",
     "check_in": "2026-05-11", "check_out": "2026-05-14"},
    {"task_id": "t03", "guest_name": "sam rivera", "init_room": "standard",
     "check_in": "2026-06-01", "check_out": "2026-06-03"},
    {"task_id": "t04", "guest_name": "priya nair", "init_room": "suite",
     "check_in": "2026-07-19", "check_out": "2026-07-22"},
    {"task_id": "t05", "guest_name": "chen wei", "init_room": "deluxe",
     "check_in": "2026-08-02", "check_out": "2026-08-06"},
    {"task_id": "t06", "guest_name": "dana brooks", "init_room": "standard",
     "check_in": "2026-09-14", "check_out": "2026-09-17"},
    {"task_id": "t07", "guest_name": "omar haddad", "init_room": "suite",
     "check_in": "2026-10-03", "check_out": "2026-10-05"},
    {"task_id": "t08", "guest_name": "riley quinn", "init_room": "deluxe",
     "check_in": "2026-11-21", "check_out": "2026-11-24"},
    {"task_id": "t09", "guest_name": "nina petrova", "init_room": "standard",
     "check_in": "2026-12-08", "check_out": "2026-12-12"},
    {"task_id": "t10", "guest_name": "luis ortega", "init_room": "suite",
     "check_in": "2027-01-15", "check_out": "2027-01-18"},
]

CORRECTION_FIELDS = {
    "room_type": ["room_type"],
    "dates": ["check_in", "check_out"],
    "guest_name": ["guest_name"],
    "cancellation": [],
}

_ALT_ROOM = {"standard": "suite", "deluxe": "suite", "suite": "deluxe"}

def build_task(base: dict, kind: str) -> dict:

    t = copy.deepcopy(base)
    initial = (
        f"I'd like to book a room for {t['guest_name'].title()}, checking in "
        f"{t['check_in']} and checking out {t['check_out']}. "
        f"A {t['init_room']} room please."
    )

    goal = {
        "room_type": t["init_room"],
        "guest_name": t["guest_name"],
        "check_in": t["check_in"],
        "check_out": t["check_out"],
    }

    if kind == "room_type":
        new_room = _ALT_ROOM[t["init_room"]]
        goal["room_type"] = new_room
        correction = (
            f"Actually, I wanted the {new_room}, not the {t['init_room']}. "
            "Can you fix that?"
        )
    elif kind == "dates":
        goal["check_in"] = t["check_in"][:-2] + "20"
        goal["check_out"] = t["check_out"][:-2] + "23"
        correction = (
            f"Sorry, the dates are wrong. It should be {goal['check_in']} to "
            f"{goal['check_out']}."
        )
    elif kind == "guest_name":
        goal["guest_name"] = "jordan taylor"
        correction = "The booking should be under Jordan Taylor, not that name."
    elif kind == "cancellation":
        correction = "Actually, please cancel the whole thing. I don't need it."
    else:
        raise ValueError(f"unknown correction kind {kind}")

    goal_state = [] if kind == "cancellation" else [goal]

    return {
        **t,
        "correction_kind": kind,
        "initial_request": initial,
        "correction_text": correction,
        "goal": goal,
        "goal_state": goal_state,
        "corrected_fields": CORRECTION_FIELDS[kind],
    }

def all_tasks(kinds: list[str] | None = None) -> list[dict]:
    kinds = kinds or list(CORRECTION_FIELDS)
    return [build_task(b, k) for b in BASE_TASKS for k in kinds]

@dataclass
class ToolCall:
    id: str
    name: str
    arguments: dict[str, Any]

@dataclass
class ModelReply:

    tool_calls: list[ToolCall]
    text: str | None

class RealClient:

    def __init__(self, model: str = "gpt-4o", temperature: float = 1.0) -> None:
        from openai import OpenAI  # imported lazily so the mock path needs nothing

        self.client = OpenAI(api_key=os.environ["OPENAI_API_KEY"])
        self.model = model
        self.temperature = temperature

    def step(self, messages: list[dict], tools: list[dict]) -> ModelReply:
        resp = self.client.chat.completions.create(
            model=self.model,
            messages=messages,
            tools=tools,
            tool_choice="auto",
            temperature=self.temperature,
        )
        msg = resp.choices[0].message
        calls = []
        for tc in msg.tool_calls or []:
            try:
                args = json.loads(tc.function.arguments)
            except json.JSONDecodeError:
                args = {}
            calls.append(ToolCall(tc.id, tc.function.name, args))
        return ModelReply(tool_calls=calls, text=msg.content)

Behaviour = Callable[[dict], ModelReply]

class MockClient:

    def __init__(self, profile: str = "mixed", seed: int | None = None) -> None:
        self.profile = profile
        self.rng = random.Random(seed)

    @staticmethod
    def _call(name: str, args: dict, n: int = 0) -> ToolCall:
        return ToolCall(id=f"call_{name}_{n}", name=name, arguments=args)

    def _resolve_profile(self, ctx: dict) -> str:
        if self.profile != "mixed":
            return self.profile
        lateness = {"pre_tool": 0, "post_check": 1, "post_commit": 2}[ctx["timing"]]
        weights = [
            (0.85, 0.05, 0.05, 0.05),
            (0.60, 0.18, 0.12, 0.10),
            (0.35, 0.35, 0.15, 0.15),
        ][lateness]
        return self.rng.choices(
            ["clean", "silent", "over", "partial"], weights=weights
        )[0]

    def step(self, messages: list[dict], tools: list[dict], ctx: dict) -> ModelReply:

        task = ctx["task"]
        done = ctx["done"]  # list of tool names already called
        corrected = ctx["corrected"]

        if not corrected:
            if "check_availability" not in done:
                return ModelReply(
                    [self._call("check_availability", {"room_type": task["init_room"]})],
                    None,
                )
            if "book_room" not in done:
                return ModelReply(
                    [
                        self._call(
                            "book_room",
                            {
                                "room_type": task["init_room"],
                                "guest_name": task["guest_name"],
                                "check_in": task["check_in"],
                                "check_out": task["check_out"],
                            },
                        )
                    ],
                    None,
                )
            return ModelReply([], f"Booked the {task['init_room']} room. All set.")

        mode = ctx.setdefault("_mode", self._resolve_profile(ctx))
        goal = task["goal"]
        is_cancellation = task["goal_state"] == []

        stale = [
            c for c in ctx["stale_confirmations"] if c in ctx["open_confirmations"]
        ]

        if mode == "silent":
            claim = (
                "cancelled that for you"
                if is_cancellation
                else f"switched that to the {goal['room_type']} for you"
            )
            return ModelReply([], f"Done, I have {claim}.")

        if stale:
            return ModelReply(
                [self._call("cancel_booking", {"confirmation": stale[0]})], None
            )

        if is_cancellation:
            return ModelReply([], "Done, the booking has been cancelled.")

        if mode == "partial":
            return ModelReply([], "I have released the old room.")

        if "book_room" not in ctx["post_calls"]:
            args = {
                "room_type": goal["room_type"],
                "guest_name": goal["guest_name"],
                "check_in": goal["check_in"],
                "check_out": goal["check_out"],
            }
            if mode == "over":
                args["check_out"] = "9999-12-31"  # changed without being asked
            return ModelReply([self._call("book_room", args, 1)], None)
        return ModelReply([], f"All set, you are now booked in the {goal['room_type']}.")

SYSTEM_PROMPT = (
    "You are a hotel booking assistant. Help guests check availability and book "
    "rooms. Use the check_availability tool before booking. If a guest corrects "
    "an earlier request, make sure the final state of their booking matches what "
    "they actually asked for. Be helpful and efficient."
)

TIMINGS = ["pre_tool", "post_check", "post_commit"]

MAX_STEPS = 12

@dataclass
class RunRecord:
    task_id: str
    timing: str
    entanglement: str
    correction_kind: str
    repeat: int
    committed_state: list[dict]
    goal_state: list[dict]
    final_text: str
    tool_calls_before: int
    tool_calls_after: int
    turns_after: int
    tokens_after: int
    log: list[dict] = field(default_factory=list)
    hit_step_limit: bool = False

def _approx_tokens(text: str) -> int:

    return max(1, len(text) // 4)

def run_trial(
    client,
    task: dict,
    timing: str,
    entanglement: str,
    correction_kind: str,
    repeat: int,
    mock: bool = True,
) -> RunRecord:
    env = BookingEnv()
    messages: list[dict[str, Any]] = [
        {"role": "system", "content": SYSTEM_PROMPT},
        {"role": "user", "content": task["initial_request"]},
    ]

    ctx = {
        "task": task,
        "timing": timing,
        "done": [],
        "post_calls": [],
        "corrected": False,
        "open_confirmations": [],
        "stale_confirmations": [],
    }

    injected = False
    turns_after = 0
    tokens_after = 0
    hit_limit = True

    def should_inject() -> bool:
        if injected:
            return False
        if timing == "pre_tool":
            return len(ctx["done"]) == 0
        if timing == "post_check":
            return "check_availability" in ctx["done"] and "book_room" not in ctx["done"]
        if timing == "post_commit":
            return "book_room" in ctx["done"]
        raise ValueError(f"unknown timing {timing}")

    for _ in range(MAX_STEPS):
        if should_inject():
            messages.append({"role": "user", "content": task["correction_text"]})
            env.mark_correction_issued()
            ctx["corrected"] = True
            ctx["stale_confirmations"] = list(ctx["open_confirmations"])
            injected = True

        if mock:
            reply = client.step(messages, TOOL_SCHEMAS, ctx)
        else:
            reply = client.step(messages, TOOL_SCHEMAS)

        if ctx["corrected"]:
            turns_after += 1

        if reply.tool_calls:
            messages.append(
                {
                    "role": "assistant",
                    "tool_calls": [
                        {
                            "id": tc.id,
                            "type": "function",
                            "function": {
                                "name": tc.name,
                                "arguments": json.dumps(tc.arguments),
                            },
                        }
                        for tc in reply.tool_calls
                    ],
                }
            )
            for tc in reply.tool_calls:
                result = dispatch(env, tc.name, dict(tc.arguments))
                ctx["done"].append(tc.name)
                if ctx["corrected"]:
                    ctx["post_calls"].append(tc.name)
                    tokens_after += _approx_tokens(json.dumps(result))
                if tc.name == "book_room" and result.get("success"):
                    ctx["open_confirmations"].append(result["confirmation"])
                if tc.name == "cancel_booking" and result.get("success"):
                    conf = tc.arguments.get("confirmation")
                    if conf in ctx["open_confirmations"]:
                        ctx["open_confirmations"].remove(conf)
                messages.append(
                    {
                        "role": "tool",
                        "tool_call_id": tc.id,
                        "content": json.dumps(result),
                    }
                )
            continue

        final_text = reply.text or ""
        messages.append({"role": "assistant", "content": final_text})
        if ctx["corrected"]:
            tokens_after += _approx_tokens(final_text)
        hit_limit = False
        break
    else:
        final_text = ""

    if hit_limit:
        final_text = ""

    return RunRecord(
        task_id=task["task_id"],
        timing=timing,
        entanglement=entanglement,
        correction_kind=correction_kind,
        repeat=repeat,
        committed_state=env.committed_state(),
        goal_state=task["goal_state"],
        final_text=final_text,
        tool_calls_before=env.clean_run_tool_calls(),
        tool_calls_after=env.recovery_tool_calls(),
        turns_after=turns_after,
        tokens_after=tokens_after,
        log=env.log_as_dicts(),
        hit_step_limit=hit_limit,
    )

_SUCCESS_CUES = [
    r"\ball set\b",
    r"\bdone\b",
    r"\bswitched\b",
    r"\bchanged\b",
    r"\bupdated\b",
    r"\bmoved\b",
    r"\bbooked\b",
    r"\bconfirmed\b",
    r"\bhave (?:now )?(?:been )?(?:re)?booked\b",
    r"\btaken care of\b",
    r"\bsorted\b",
]

_NEGATION = [
    r"\bcould not\b",
    r"\bcouldn't\b",
    r"\bunable\b",
    r"\bfailed\b",
    r"\bwas not able\b",
    r"\bcannot\b",
    r"\bno longer available\b",
]

def claims_success_rules(final_text: str, goal: dict) -> bool:

    if not final_text:
        return False
    t = final_text.lower()
    if any(re.search(p, t) for p in _NEGATION):
        return False
    cue = any(re.search(p, t) for p in _SUCCESS_CUES)
    mentions_goal = str(goal.get("room_type", "")).lower() in t
    return bool(cue and (mentions_goal or "book" in t))

def claims_success_judge(final_text: str, correction_text: str, client) -> bool:

    prompt = (
        "A user asked an assistant to make this change:\n"
        f"{correction_text}\n\n"
        "The assistant's final message was:\n"
        f"{final_text}\n\n"
        "Does the final message assert that the requested change was completed? "
        "Answer with exactly one word: YES or NO."
    )
    reply = client.step(
        [{"role": "user", "content": prompt}],
        tools=[],
    )
    return (reply.text or "").strip().upper().startswith("YES")

@dataclass
class Metrics:
    recovery_success: int  # R
    cost_tools: int  # component of C
    cost_turns: int
    cost_tokens: int
    normalized_cost: float  # Chat
    silent_inconsistency: int  # I
    overcorrection: int  # O
    outcome: str
    claim_rules: int
    claim_disagreement: int  # 1 when the two detectors disagree

def _state_equal(committed: list[dict], goal: list[dict]) -> bool:
    return committed == goal

def _goal_fields_met(committed: list[dict], goal: list[dict], fields: list[str]) -> bool:

    if len(committed) != 1 or len(goal) != 1:
        return False
    return all(committed[0].get(f) == goal[0].get(f) for f in fields)

def _extra_fields_altered(
    committed: list[dict], goal: list[dict], corrected_fields: list[str]
) -> bool:
    if len(committed) != 1 or len(goal) != 1:
        return False
    others = [k for k in goal[0] if k not in corrected_fields]
    return any(committed[0].get(k) != goal[0].get(k) for k in others)

def compute(record, corrected_fields: list[str], claim_judge: int | None = None) -> Metrics:

    committed = record.committed_state
    goal = record.goal_state
    goal_dict = goal[0] if goal else {}

    success = int(_state_equal(committed, goal))

    claim_r = int(claims_success_rules(record.final_text, goal_dict))
    claim = claim_r if claim_judge is None else int(claim_judge)
    disagreement = 0 if claim_judge is None else int(claim_r != int(claim_judge))

    silent = int(bool(claim) and not success)

    over = int(
        (not success)
        and _goal_fields_met(committed, goal, corrected_fields)
        and _extra_fields_altered(committed, goal, corrected_fields)
    )

    if success:
        outcome = "clean_recovery"
    elif over:
        outcome = "overcorrection"
    elif silent:
        outcome = "silent_inconsistency"
    elif len(committed) == 0 and not claim:
        outcome = "partial_recovery"
    else:
        outcome = "failure"

    denom = max(1, record.tool_calls_before)
    return Metrics(
        recovery_success=success,
        cost_tools=record.tool_calls_after,
        cost_turns=record.turns_after,
        cost_tokens=record.tokens_after,
        normalized_cost=record.tool_calls_after / denom,
        silent_inconsistency=silent,
        overcorrection=over,
        outcome=outcome,
        claim_rules=claim_r,
        claim_disagreement=disagreement,
    )

CSV_FIELDS = [
    "task_id", "correction_kind", "timing", "entanglement", "repeat",
    "recovery_success", "cost_tools", "cost_turns", "cost_tokens",
    "normalized_cost", "silent_inconsistency", "overcorrection", "outcome",
    "claim_rules", "claim_disagreement", "hit_step_limit",
]

def entanglement_for(timing: str) -> str:

    return "committed" if timing == "post_commit" else "inspected"

def run_grid(
    client,
    kinds: list[str],
    repeats: int,
    mock: bool,
    out_path: str,
    traces_path: str | None = None,
    progress: bool = True,
) -> list[dict]:

    tasks = all_tasks(kinds)
    rows: list[dict] = []
    traces: list[dict] = []
    total = len(tasks) * len(TIMINGS) * repeats
    n = 0

    for task in tasks:
        for timing in TIMINGS:
            for rep in range(repeats):
                rec = run_trial(
                    client=client,
                    task=task,
                    timing=timing,
                    entanglement=entanglement_for(timing),
                    correction_kind=task["correction_kind"],
                    repeat=rep,
                    mock=mock,
                )
                m = compute(rec, task["corrected_fields"])
                rows.append({
                    "task_id": rec.task_id,
                    "correction_kind": rec.correction_kind,
                    "timing": rec.timing,
                    "entanglement": rec.entanglement,
                    "repeat": rec.repeat,
                    "hit_step_limit": int(rec.hit_step_limit),
                    **asdict(m),
                })
                if traces_path:
                    traces.append({
                        "task_id": rec.task_id, "timing": rec.timing,
                        "repeat": rec.repeat, "outcome": m.outcome,
                        "committed_state": rec.committed_state,
                        "goal_state": rec.goal_state,
                        "final_text": rec.final_text,
                        "log": rec.log,
                    })
                n += 1
                if progress and n % 200 == 0:
                    print(f"  {n}/{total} runs", file=sys.stderr)

    with open(out_path, "w", newline="") as f:
        w = csv.DictWriter(f, fieldnames=CSV_FIELDS)
        w.writeheader()
        for r in rows:
            w.writerow({k: r.get(k, "") for k in CSV_FIELDS})

    if traces_path:
        with open(traces_path, "w") as f:
            json.dump(traces, f, indent=1)

    print(f"wrote {len(rows)} runs to {out_path}")
    return rows

LATENESS = {"pre_tool": 0, "post_check": 1, "post_commit": 2}

def _np_pd():
    import numpy as np
    import pandas as pd
    return np, pd

def ci95(x):

    np, _ = _np_pd()
    x = x.dropna()
    n = len(x)
    if n < 2:
        return float("nan")
    return 1.96 * x.std(ddof=1) / np.sqrt(n)

def summarize(df, by: list[str]):
    _, pd = _np_pd()
    df = df.copy()
    df["silent_excl"] = (df["outcome"] == "silent_inconsistency").astype(int)
    g = df.groupby(by, observed=True)
    return pd.DataFrame({
        "n": g.size(),
        "success": g["recovery_success"].mean(),
        "success_ci": g["recovery_success"].apply(ci95),
        "cost_tools": g["cost_tools"].mean(),
        "cost_tools_ci": g["cost_tools"].apply(ci95),
        "norm_cost": g["normalized_cost"].mean(),
        "silent": g["silent_inconsistency"].mean(),
        "silent_ci": g["silent_inconsistency"].apply(ci95),
        "silent_excl": g["silent_excl"].mean(),
        "over": g["overcorrection"].mean(),
    }).round(3)

def slope(df, xcol: str, ycol: str) -> tuple[float, float]:

    np, _ = _np_pd()
    x = df[xcol].to_numpy(dtype=float)
    y = df[ycol].to_numpy(dtype=float)
    if x.std() == 0:
        return float("nan"), float("nan")
    return float(np.polyfit(x, y, 1)[0]), float(np.corrcoef(x, y)[0, 1])

def cohens_h(p1: float, p2: float) -> float:

    np, _ = _np_pd()
    return float(2 * np.arcsin(np.sqrt(max(p1, 0))) - 2 * np.arcsin(np.sqrt(max(p2, 0))))

def analyze(path: str) -> int:
    _, pd = _np_pd()
    df = pd.read_csv(path)
    df["lateness"] = df["timing"].map(LATENESS)
    line = "=" * 68

    print(line + "\nOVERALL\n" + line)
    print(f"runs: {len(df)}")
    print(f"recovery success rate   : {df['recovery_success'].mean():.3f}")
    print(f"silent inconsistency    : {df['silent_inconsistency'].mean():.3f}")
    print(f"overcorrection rate     : {df['overcorrection'].mean():.3f}")
    print(f"mean recovery tool calls: {df['cost_tools'].mean():.2f}")
    if df["hit_step_limit"].sum():
        print(f"WARNING: {int(df['hit_step_limit'].sum())} runs hit the step limit")
    print("\noutcome distribution:")
    print(df["outcome"].value_counts(normalize=True).round(3).to_string())

    print("\n" + line + "\nH1  lateness: cost should rise, success should fall\n" + line)
    by_t = summarize(df, ["timing"]).reindex(TIMINGS)
    print(by_t.to_string())
    b_cost, r_cost = slope(df, "lateness", "cost_tools")
    b_succ, r_succ = slope(df, "lateness", "recovery_success")
    print(f"\ncost slope per step later   : {b_cost:+.3f}  (r = {r_cost:+.3f})")
    print(f"success slope per step later: {b_succ:+.3f}  (r = {r_succ:+.3f})")
    print("H1 supported" if b_cost > 0 and b_succ < 0 else "H1 not supported by these data")

    print("\n" + line + "\nH2  entanglement: committed should cost more than inspected\n" + line)
    print(summarize(df, ["entanglement"]).to_string())
    if {"committed", "inspected"} <= set(df["entanglement"].unique()):
        c = df[df.entanglement == "committed"]
        i = df[df.entanglement == "inspected"]
        h = cohens_h(i["recovery_success"].mean(), c["recovery_success"].mean())
        print(f"\ncost difference (committed - inspected): "
              f"{c['cost_tools'].mean() - i['cost_tools'].mean():+.2f} tool calls")
        print(f"success difference: "
              f"{c['recovery_success'].mean() - i['recovery_success'].mean():+.3f}"
              f"  (Cohen's h = {h:+.3f})")
        ok = (c["cost_tools"].mean() > i["cost_tools"].mean()
              and c["recovery_success"].mean() < i["recovery_success"].mean())
        print("H2 supported" if ok else "H2 not supported by these data")

    print("\n" + line + "\nH3  silent inconsistency: nonzero, and rising with lateness\n" + line)
    print(by_t[["silent", "silent_ci", "silent_excl", "n"]].to_string())
    b_sil, r_sil = slope(df, "lateness", "silent_inconsistency")
    rate = df["silent_inconsistency"].mean()
    print(f"\nsilent inconsistency slope: {b_sil:+.3f}  (r = {r_sil:+.3f})")
    print(f"overall rate: {rate:.3f}")
    print("H3 supported" if rate > 0 and b_sil > 0 else "H3 not supported by these data")

    print("\n" + line + "\nBY CORRECTION KIND  (robustness of the effects)\n" + line)
    print(summarize(df, ["correction_kind"]).to_string())

    if df["claim_disagreement"].sum():
        print(f"\nclaim detector disagreements: {int(df['claim_disagreement'].sum())} "
              "runs flagged for manual inspection")
    return 0

def main(argv: list[str] | None = None) -> int:
    ap = argparse.ArgumentParser(
        description="Correction recovery harness (single file edition)."
    )
    sub = ap.add_subparsers(dest="cmd", required=True)

    r = sub.add_parser("run", help="sweep the condition grid")
    r.add_argument("--mock", action="store_true",
                   help="run offline with the scripted mock, no API key needed")
    r.add_argument("--profile", default="mixed",
                   help="mock behaviour: mixed, clean, silent, over, partial")
    r.add_argument("--model", default="gpt-4o")
    r.add_argument("--repeats", type=int, default=20)
    r.add_argument("--kinds", default="room_type,dates,guest_name,cancellation")
    r.add_argument("--out", default="results.csv")
    r.add_argument("--traces", default=None, help="optional path for full JSON traces")
    r.add_argument("--seed", type=int, default=0)

    a = sub.add_parser("analyze", help="run the hypothesis tests on a results CSV")
    a.add_argument("path")

    args = ap.parse_args(argv)

    if args.cmd == "analyze":
        return analyze(args.path)

    kinds = [k.strip() for k in args.kinds.split(",") if k.strip()]
    if args.mock:
        client = MockClient(profile=args.profile, seed=args.seed)
    else:
        if "OPENAI_API_KEY" not in os.environ:
            print("OPENAI_API_KEY not set. Use --mock for an offline run.",
                  file=sys.stderr)
            return 1
        client = RealClient(model=args.model)

    run_grid(client, kinds, args.repeats, args.mock, args.out, args.traces)
    return 0

if __name__ == "__main__":
    raise SystemExit(main())
