add main.py
This commit is contained in:
@@ -0,0 +1,159 @@
|
||||
"""
|
||||
Main entry point for the Human‑in‑the‑loop LangGraph example.
|
||||
|
||||
The program demonstrates a simple graph with one node that pauses execution and asks the user to confirm an action. The pause is implemented using ``langgraph``'s custom interrupt mechanism. After the user answers, the graph resumes and prints the final state.
|
||||
|
||||
Requirements:
|
||||
pip install -r requirements.txt
|
||||
|
||||
Run:
|
||||
python main.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import TypedDict, List, Dict, Any
|
||||
|
||||
import questionary
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.constants import START
|
||||
from langgraph.types import interrupt, Command
|
||||
from langgraph.checkpoint.memory import InMemorySaver
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. Define the graph state
|
||||
# ---------------------------------------------------------------------------
|
||||
class GraphState(TypedDict):
|
||||
"""Minimal state used by the example.
|
||||
|
||||
* ``human_value`` – value supplied by the user during the interrupt.
|
||||
* ``foo`` – placeholder for any initial data that might be needed later.
|
||||
"""
|
||||
|
||||
human_value: str | None
|
||||
foo: int
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. Node that triggers an interrupt
|
||||
# ---------------------------------------------------------------------------
|
||||
async def ask_user(state: GraphState) -> GraphState:
|
||||
"""Node that pauses execution and asks the user for confirmation.
|
||||
|
||||
The node returns a ``Command`` with an ``interrupt`` payload. When the graph
|
||||
receives this command it stops, yields control to the caller, and waits for
|
||||
a resume command.
|
||||
"""
|
||||
# Build interrupt payload – a simple dictionary that will be shown to the user.
|
||||
payload: Dict[str, Any] = {
|
||||
"type": "confirm",
|
||||
"question": "Do you want to continue?",
|
||||
"options": ["yes", "no"],
|
||||
}
|
||||
# Trigger interrupt – execution stops here until a resume command is sent.
|
||||
return Command(resume=interrupt(payload))
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Node that receives the user's answer and updates state
|
||||
# ---------------------------------------------------------------------------
|
||||
async def process_answer(state: GraphState) -> GraphState:
|
||||
"""Node executed after the graph resumes.
|
||||
|
||||
The ``state`` will contain the resume payload under ``__interrupt__``.
|
||||
We extract the user response, store it in ``human_value`` and return the
|
||||
updated state.
|
||||
"""
|
||||
# ``__interrupt__`` is a list of interrupt objects; we only use the first.
|
||||
interrupt_obj = state.get("__interrupt__", [])[0]
|
||||
payload: Dict[str, Any] = interrupt_obj.value # type: ignore[assignment]
|
||||
|
||||
# The user response will be added to the payload under ``answer``.
|
||||
answer = payload.get("answer")
|
||||
if not isinstance(answer, str):
|
||||
raise ValueError("Interrupt payload missing 'answer' field")
|
||||
|
||||
state["human_value"] = answer
|
||||
return state
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. Build the graph
|
||||
# ---------------------------------------------------------------------------
|
||||
def build_graph() -> StateGraph:
|
||||
"""Create and compile a simple graph with an interrupt node.
|
||||
|
||||
The graph consists of two nodes: ``ask_user`` (which pauses) and
|
||||
``process_answer`` (which continues after the user responds).
|
||||
"""
|
||||
graph = StateGraph(GraphState)
|
||||
graph.add_node("ask", ask_user)
|
||||
graph.add_node("answer", process_answer)
|
||||
graph.set_entry_point("ask")
|
||||
graph.add_edge(START, "ask")
|
||||
graph.add_edge("ask", "answer")
|
||||
|
||||
# Use an in‑memory checkpoint so we can resume after the interrupt.
|
||||
return graph.compile(checkpointer=InMemorySaver())
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 5. Main loop that runs the graph and handles interrupts
|
||||
# ---------------------------------------------------------------------------
|
||||
async def run_graph() -> None:
|
||||
"""Execute the graph, handle the custom interrupt, and resume execution.
|
||||
|
||||
The function streams the graph output, looks for ``__interrupt__`` chunks,
|
||||
prompts the user with *questionary*, then resumes the graph with the
|
||||
selected answer.
|
||||
"""
|
||||
from asyncio import run # Imported lazily to keep top‑level imports minimal.
|
||||
|
||||
graph = build_graph()
|
||||
config = {"configurable": {"thread_id": "demo-thread-1"}}
|
||||
|
||||
# Initial state – ``foo`` can be any value; it is not used in this demo.
|
||||
init_state: GraphState = {"human_value": None, "foo": 42}
|
||||
|
||||
# Stream the graph until completion.
|
||||
async for chunk in graph.stream(init_state, config):
|
||||
if "__interrupt__" in chunk:
|
||||
interrupt_obj = chunk["__interrupt__"][0]
|
||||
payload: Dict[str, Any] = interrupt_obj.value # type: ignore[assignment]
|
||||
|
||||
print("\n--- Human‑in‑the‑loop interrupt received ---")
|
||||
print(f"Type: {payload.get('type')}")
|
||||
print(f"Question: {payload.get('question')}\n")
|
||||
|
||||
# Ask the user for a choice.
|
||||
answer = questionary.select(
|
||||
"Choose an option:",
|
||||
choices=payload.get("options", []),
|
||||
).ask()
|
||||
|
||||
if answer is None:
|
||||
raise RuntimeError("User cancelled the prompt")
|
||||
|
||||
# Attach the answer to the payload and resume.
|
||||
payload["answer"] = answer
|
||||
print(f"\n> User selected: {answer}\n")
|
||||
await graph.stream(Command(resume=payload), config)
|
||||
else:
|
||||
# Normal output – just print it.
|
||||
if "messages" in chunk:
|
||||
for msg in chunk["messages"]:
|
||||
print(msg.content, end="", flush=True)
|
||||
|
||||
# After the stream ends, fetch the final state from the checkpoint.
|
||||
final_state = graph.checkpointer.get_state(config)
|
||||
print("\n--- Final state ---")
|
||||
print(final_state)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 6. Entry point
|
||||
# ---------------------------------------------------------------------------
|
||||
if __name__ == "__main__":
|
||||
import asyncio
|
||||
|
||||
try:
|
||||
asyncio.run(run_graph())
|
||||
except KeyboardInterrupt:
|
||||
print("\nInterrupted by user.")
|
||||
"""
|
||||
Reference in New Issue
Block a user