"""Evaluate a deterministic LangGraph order agent without external services.

这个示例展示如何在没有真实 LLM、数据库或外部 API 的情况下评测 Agent：
用 LangGraph 搭一个确定性“查订单”工作流，再用小型数据集验证答案和工具调用是否符合预期。
"""

from dataclasses import dataclass
from typing import TypedDict

from langgraph.graph import END, START, StateGraph


class AgentState(TypedDict):
    """图在节点之间传递的状态结构。"""

    # 用户问题本身。这个确定性示例里不会解析问题，只用于说明真实 Agent 状态通常会保留原问题。
    question: str
    # 要查询的订单号。
    order_id: str
    # Agent 生成或工具返回的最终答案。
    answer: str
    # 记录被调用过的工具名，用于评测“是否调用了正确工具”。
    tools: list[str]


# 用内存字典模拟订单服务，保证评测可复现且不依赖外部系统。
ORDERS = {"A1042": "paid", "A1043": "shipped"}


def lookup(state: AgentState) -> dict:
    """模拟调用订单查询工具，并把结果写回图状态。"""

    # 未命中的订单显式返回 not_found，避免测试里出现 None 或异常分支的不确定性。
    status = ORDERS.get(state["order_id"], "not_found")

    # LangGraph 节点可以只返回要更新的字段；这里同时更新答案和工具调用轨迹。
    return {"answer": status, "tools": ["orders_get_by_id"]}


# 构建一个最小状态图：START -> lookup -> END。
builder = StateGraph(AgentState)
builder.add_node("lookup", lookup)
builder.add_edge(START, "lookup")
builder.add_edge("lookup", END)
agent = builder.compile()


@dataclass(frozen=True)
class Example:
    """单条评测样本。

    frozen=True 让样本不可变，避免评测过程中被误改。
    """

    question: str
    order_id: str
    expected_answer: str
    expected_tool: str


# 离线 golden dataset：每条样本都有输入和期望输出。
DATASET = [
    Example("What is the order status?", "A1042", "paid", "orders_get_by_id"),
    Example("Has this order shipped?", "A1043", "shipped", "orders_get_by_id"),
    Example("Find this order", "A9999", "not_found", "orders_get_by_id"),
]


def evaluate(example: Example) -> dict[str, bool]:
    """运行一条样本，并分别评测答案和工具调用。"""

    result = agent.invoke(
        {
            "question": example.question,
            "order_id": example.order_id,
            "answer": "",
            "tools": [],
        }
    )
    return {
        # 答案正确性：业务输出是否匹配期望。
        "answer_correct": result["answer"] == example.expected_answer,
        # 工具正确性：是否调用了期望的订单查询工具。
        "tool_correct": result["tools"] == [example.expected_tool],
    }


if __name__ == "__main__":
    # 逐条运行数据集，得到每条样本的多维评测结果。
    rows = [evaluate(example) for example in DATASET]

    # 只有一条样本的所有指标都为 True，才算该样本通过。
    passed = sum(all(row.values()) for row in rows)
    print(rows)

    # pass_rate 是最简单的汇总指标；真实评测通常还会输出失败样本明细。
    print(f"pass_rate={passed / len(rows):.0%}")
