Files
task-6a1864fa8a94f887e50d46f0/agent.py
T
2026-06-04 23:10:50 +00:00

131 lines
3.7 KiB
Python
Raw Blame History

This file contains invisible Unicode characters
This file contains invisible Unicode characters that are indistinguishable to humans but may be processed differently by a computer. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
Selfcorrecting LangGraph agent.
The agent executes a task, verifies the result with an LLM judge, and retries on failure
until success or a maximum number of attempts is reached.
"""
from __future__ import annotations
from typing import TypedDict
from langchain_openai import ChatOpenAI
from langgraph.graph import StateGraph, END
from langgraph.checkpoint.memory import InMemorySaver
# ---------- State ----------
class AgentState(TypedDict):
task: str
result: str
attempts: int
status: str # pending | success | failed | max_attempts
error: str | None
max_attempts: int
# ---------- Unreliable tool ----------
import random
class UnreliableTool:
"""Simulates a tool that fails 30% of the time."""
def __call__(self, input_text: str) -> str:
if random.random() < 0.3:
raise ValueError("Simulated tool failure")
return f"Result for: {input_text}"
unreliable_tool = UnreliableTool()
# ---------- Nodes ----------
async def execute_task(state: AgentState) -> AgentState:
"""Runs the task using the unreliable tool."""
try:
result = unreliable_tool(state["task"])
state["result"] = result
state["error"] = None
except Exception as e:
state["result"] = ""
state["error"] = str(e)
return state
async def verify_result(state: AgentState) -> AgentState:
"""LLM judge that returns 'success' or 'failed' based on the result."""
llm = ChatOpenAI(temperature=0)
prompt = (
"""You are a judge. Evaluate the following result for the task:\n"
f"Task: {state['task']}\n"
f"Result: {state['result']}\n"
"If the result is correct, reply with the word 'success'.\n"
"If the result is incorrect or missing, reply with the word 'failed'."""
)
response = llm.invoke(prompt)
verdict = response.content.strip().lower()
if verdict == "success":
state["status"] = "success"
else:
state["status"] = "failed"
return state
async def handle_error(state: AgentState) -> AgentState:
"""Increment attempts and prepare for retry."""
state["attempts"] += 1
state["status"] = "pending"
return state
# ---------- Graph ----------
def create_agent_graph(max_attempts: int = 3) -> StateGraph:
graph = StateGraph(AgentState)
graph.add_node("execute_task", execute_task)
graph.add_node("verify_result", verify_result)
graph.add_node("handle_error", handle_error)
# Start state
graph.set_entry_point("execute_task")
# Transitions
graph.add_conditional_edges(
"verify_result",
lambda state: (
"END" if state["status"] == "success" else
"handle_error" if state["attempts"] < state["max_attempts"] else
"END"
),
)
graph.add_edge("handle_error", "execute_task")
# Final states
graph.add_edge("execute_task", "verify_result")
return graph
# ---------- Runner ----------
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Selfcorrecting agent demo")
parser.add_argument("task", type=str, help="Task to execute")
parser.add_argument("--max_attempts", type=int, default=3, 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_attempts,
}
graph = create_agent_graph(max_attempts=args.max_attempts)
memory = InMemorySaver()
app = graph.compile(checkpointer=memory)
final_state = app.invoke(initial_state)
print("\nFinal state:")
for k, v in final_state.items():
print(f"{k}: {v}")