"""LangGraph agent with retry until success or max_attempts. This module implements: - AgentState TypedDict - An unreliable tool that fails 30% of the time - Nodes: execute_task, verify_result, handle_error - StateGraph with cycle per plan - CLI to run a single task """ from typing import TypedDict, Literal, Any import random import os import sys # LLM imports from langchain_openai import ChatOpenAI from langgraph.graph import StateGraph, START, END from langgraph.checkpoint.memory import MemorySaver # 1. Состояние графа class AgentState(TypedDict): task: str result: str attempts: int status: Literal["pending", "success", "failed", "max_attempts"] error: str | None max_attempts: int # 2. Tool: unreliable_tool def unreliable_tool(input: str) -> str: """Simulate an unreliable tool that fails 30% of the time.""" if random.random() < 0.3: raise ValueError("Simulated tool failure") # Example task: compute sum of two numbers in the input string try: parts = input.split() nums = [int(p) for p in parts if p.isdigit()] return str(sum(nums)) except Exception as e: raise ValueError(f"Tool error: {e}") # 3. LLM for verification llm = ChatOpenAI( model="gpt-4o-mini", # placeholder, user can set via env temperature=0, base_url=os.getenv("OPENAI_API_BASE", "https://api.openai.com/v1"), api_key=os.getenv("OPENAI_API_KEY"), ) # 4. Node: execute_task async def execute_task(state: AgentState) -> AgentState: try: result = unreliable_tool(state["task"]) state["result"] = result state["error"] = None except Exception as e: state["result"] = "" state["error"] = str(e) return state # 5. Node: verify_result async def verify_result(state: AgentState) -> AgentState: # Use LLM to decide success or failed prompt = f"Given the result '{state['result']}', is this a correct answer? Respond with 'success' or 'failed'." try: resp = await llm.ainvoke(prompt) # Extract keyword verdict = "failed" if "success" in resp.content.lower(): verdict = "success" except Exception as e: verdict = "failed" state["error"] = f"LLM error: {e}" state["status"] = verdict return state # 6. Node: handle_error async def handle_error(state: AgentState) -> AgentState: state["attempts"] += 1 if state["attempts"] >= state["max_attempts"]: state["status"] = "max_attempts" else: state["status"] = "failed" return state # 7. Graph definition builder = StateGraph(AgentState) builder.add_node("execute_task", execute_task) builder.add_node("verify_result", verify_result) builder.add_node("handle_error", handle_error) builder.set_entry_point("execute_task") builder.add_edge("execute_task", "verify_result") builder.add_conditional_edges( "verify_result", lambda state: state["status"], { "success": END, "failed": "handle_error", "max_attempts": END, }, ) builder.add_edge("handle_error", "execute_task") graph = builder.compile(checkpointer=MemorySaver()) # 8. CLI if __name__ == "__main__": import argparse parser = argparse.ArgumentParser(description="Run a single task with retry until success or max attempts.") parser.add_argument("task", type=str, help="Task string to feed to unreliable_tool") parser.add_argument("--max", type=int, default=5, help="Maximum retry attempts") args = parser.parse_args() initial_state: AgentState = { "task": args.task, "result": "", "attempts": 0, "status": "pending", "error": None, "max_attempts": args.max, } result = graph.invoke(initial_state) print("\nFinal state:") for k, v in result.items(): print(f"{k}: {v}")