"""A LangGraph short-term memory demo with two turns in one thread.

这个示例演示 LangGraph 的短期记忆：
同一个 thread_id 下，多次 invoke 会通过 checkpointer 共享历史消息，
因此第二轮可以看到第一轮写入的消息。
"""

from langchain_core.messages import AIMessage, HumanMessage
from langgraph.checkpoint.memory import InMemorySaver
from langgraph.graph import END, START, MessagesState, StateGraph


def reply(state: MessagesState) -> dict:
    """根据当前线程里的历史消息生成一个确定性回复。"""

    # MessagesState["messages"] 会包含当前线程已累计的消息。
    # 这里只统计 HumanMessage，忽略 AIMessage，便于展示“用户说了几轮”。
    human_messages = [
        message for message in state["messages"] if isinstance(message, HumanMessage)
    ]

    # 取最新一条用户消息，用在回复中。
    latest = human_messages[-1].content
    return {
        "messages": [
            AIMessage(content=f"本线程已收到 {len(human_messages)} 条用户消息；最新消息：{latest}")
        ]
    }


# 图结构很简单：收到消息后只执行 reply 节点，然后结束。
builder = StateGraph(MessagesState)
builder.add_node("reply", reply)
builder.add_edge(START, "reply")
builder.add_edge("reply", END)

# InMemorySaver 负责按 thread_id 保存消息状态。
# 生产环境可替换成数据库型 checkpointer。
graph = builder.compile(checkpointer=InMemorySaver())


if __name__ == "__main__":
    # 两次调用使用同一个 thread_id，因此第二次能读到第一次的消息。
    config = {"configurable": {"thread_id": "demo-thread"}}

    # 第一轮写入“我叫 Elaine”。
    graph.invoke({"messages": [HumanMessage("我叫 Elaine")]}, config)

    # 第二轮继续同一个线程，reply 能看到两条 HumanMessage。
    result = graph.invoke({"messages": [HumanMessage("这是第几条消息？")]}, config)
    print(result["messages"][-1].content)
    # => 本线程已收到 2 条用户消息；最新消息：这是第几条消息？
