add main
This commit is contained in:
@@ -0,0 +1,87 @@
|
|||||||
|
import random
|
||||||
|
from typing import TypedDict, Dict
|
||||||
|
from langgraph.graph import StateGraph
|
||||||
|
from langchain_openai import ChatOpenAI
|
||||||
|
from langchain_core.messages import HumanMessage, AIMessage
|
||||||
|
|
||||||
|
# Define 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
|
||||||
|
async def unreliable_tool(task: str) -> str:
|
||||||
|
if random.random() < 0.3:
|
||||||
|
raise ValueError("Tool failure")
|
||||||
|
# simple eval for arithmetic
|
||||||
|
try:
|
||||||
|
return str(eval(task))
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"Eval error: {e}")
|
||||||
|
|
||||||
|
# Execute task node
|
||||||
|
async def execute_task(state: AgentState) -> Dict[str, str]:
|
||||||
|
try:
|
||||||
|
result = await unreliable_tool(state["task"])
|
||||||
|
return {"result": result, "error": None}
|
||||||
|
except Exception as e:
|
||||||
|
return {"result": "", "error": str(e)}
|
||||||
|
|
||||||
|
# Verify result node using LLM-as-judge
|
||||||
|
llm = ChatOpenAI(temperature=0)
|
||||||
|
async def verify_result(state: AgentState) -> Dict[str, str]:
|
||||||
|
prompt = f"\nTask: {state['task']}\nResult: {state['result']}\nIs this correct? Answer with 'success' or 'failed'."
|
||||||
|
msg = await llm.ainvoke([HumanMessage(content=prompt)])
|
||||||
|
verdict = msg.content.strip().lower()
|
||||||
|
if "success" in verdict:
|
||||||
|
return {"status": "success"}
|
||||||
|
else:
|
||||||
|
return {"status": "failed"}
|
||||||
|
|
||||||
|
# Handle error / retry node
|
||||||
|
async def handle_error(state: AgentState) -> Dict[str, str]:
|
||||||
|
attempts = state["attempts"] + 1
|
||||||
|
if attempts >= state["max_attempts"]:
|
||||||
|
return {"status": "max_attempts", "attempts": attempts}
|
||||||
|
else:
|
||||||
|
return {"attempts": attempts, "status": "pending"}
|
||||||
|
|
||||||
|
# Build graph
|
||||||
|
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_conditional_edges(
|
||||||
|
"execute_task",
|
||||||
|
lambda x: "verify_result" if x["error"] is None else "handle_error",
|
||||||
|
)
|
||||||
|
builder.add_conditional_edges(
|
||||||
|
"verify_result",
|
||||||
|
lambda x: "end_success" if x.get("status") == "success" else (
|
||||||
|
"end_max_attempts" if x.get("status") == "max_attempts" else "handle_error"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
builder.add_edge("handle_error", "execute_task")
|
||||||
|
|
||||||
|
builder.set_finish_nodes(["end_success", "end_max_attempts"])
|
||||||
|
|
||||||
|
graph = builder.compile()
|
||||||
|
|
||||||
|
# Run demo
|
||||||
|
if __name__ == "__main__":
|
||||||
|
initial_state: AgentState = {
|
||||||
|
"task": "2+2",
|
||||||
|
"result": "",
|
||||||
|
"attempts": 0,
|
||||||
|
"status": "pending",
|
||||||
|
"error": None,
|
||||||
|
"max_attempts": 5,
|
||||||
|
}
|
||||||
|
result = graph.invoke(initial_state)
|
||||||
|
print("Final state:", result)
|
||||||
Reference in New Issue
Block a user