Files
brojs-task-6a1864fa8a94f887…/main.py
T

103 lines
3.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
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.
Run with:
python main.py "Вычисли 2+2"
The agent will execute the task, verify the result via an LLM judge, and retry up to max_attempts.
"""
import os
from typing import TypedDict, Dict, Any
from langgraph.graph import StateGraph
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage
MAX_ATTEMPTS = 5
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY")
if not OPENAI_API_KEY:
raise RuntimeError("Please set OPENAI_API_KEY environment variable.")
llm = ChatOpenAI(model="gpt-3.5-turbo", temperature=0)
class AgentState(TypedDict):
task: str
result: str | None
attempts: int
status: str # pending | success | failed | max_attempts
error: str | None
max_attempts: int
async def execute_task(state: AgentState) -> Dict[str, Any]:
try:
import random
if random.random() < 0.3:
raise ValueError("Simulated tool error")
result = f"Result for: {state['task']}"
return {"result": result, "error": None, "status": "pending"}
except Exception as e:
return {"result": None, "error": str(e), "status": "failed"}
async def verify_result(state: AgentState) -> Dict[str, Any]:
if state.get("error"):
return {"status": "failed"}
prompt = (
f"You are a judge. The task was: {state['task']}\n"
f"The result returned by the tool is: {state['result']}\n"
"Determine if this result satisfies the task. Respond with only one word: 'success' or 'failed'."
)
response = await llm.agenerate([HumanMessage(content=prompt)])
verdict = response.generations[0][0].content.strip().lower()
return {"status": verdict if verdict in {"success", "failed"} else "failed"}
async def handle_error(state: AgentState) -> Dict[str, Any]:
new_attempts = state["attempts"] + 1
if new_attempts >= state["max_attempts"]:
return {"attempts": new_attempts, "status": "max_attempts"}
return {"attempts": new_attempts, "status": "pending", "error": None, "result": None}
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 x: x["status"],
{"success": "END_success", "failed": "handle_error"},
)
builder.add_edge("handle_error", "execute_task")
builder.add_conditional_edges(
"handle_error",
lambda x: x["status"],
{"max_attempts": "END_max"},
)
builder.set_finish_nodes(["END_success", "END_max"])
graph = builder.compile()
async def run_task(task: str) -> AgentState:
initial_state: AgentState = {
"task": task,
"result": None,
"attempts": 0,
"status": "pending",
"error": None,
"max_attempts": MAX_ATTEMPTS,
}
return await graph.ainvoke(initial_state)
if __name__ == "__main__":
import sys, asyncio
if len(sys.argv) < 2:
print("Usage: python main.py '<task>'")
sys.exit(1)
task_text = sys.argv[1]
final = asyncio.run(run_task(task_text))
print(f"Task: {final['task']}")
print(f"Attempts: {final['attempts']}\n")
if final.get("status") == "success":
print(f"Result: {final['result']}\nStatus: success")
else:
print(f"Failed after {final['attempts']} attempts. Status: {final.get('status')}\nError: {final.get('error')}")