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

118 lines
3.0 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.
"""
LangGraph Reflection Agent
Usage:
python main.py "Your question"
"""
import sys
from typing import TypedDict, Dict
# State definition
class ReflectState(TypedDict):
question: str
draft: str
critique: str
verdict: str # 'ok' or 'needs_revision'
round: int
max_rounds: int
# Node functions
async def draft_answer(state: Dict) -> Dict:
from langchain_openai import ChatOpenAI
llm = ChatOpenAI(model="gpt-3.5-turbo")
prompt = f"Write a concise answer (510 sentences) to the following question:\n\n{state['question']}"
response = await llm.invoke(prompt)
state["draft"] = response.content
return state
async def reflect(state: Dict) -> Dict:
from langchain_openai import ChatOpenAI
llm = ChatOpenAI(model="gpt-3.5-turbo")
prompt = (
f"You are a critic evaluating the draft answer for completeness, specificity, and lack of filler.\n"
f"Draft: {state['draft']}\n"
"Provide verdict (ok or needs_revision) and 23 bullet points of critique."
)
response = await llm.invoke(prompt)
# Simple parsing
text = response.content.strip()
if "needs_revision" in text.lower():
state["verdict"] = "needs_revision"
else:
state["verdict"] = "ok"
state["critique"] = text
return state
async def rewrite(state: Dict) -> Dict:
from langchain_openai import ChatOpenAI
llm = ChatOpenAI(model="gpt-3.5-turbo")
prompt = (
f"Rewrite the draft answer incorporating the following critique:\n"
f"Critique: {state['critique']}\n"
"Provide a revised concise answer (510 sentences)."
)
response = await llm.invoke(prompt)
state["draft"] = response.content
state["round"] += 1
return state
# Build graph
from langgraph.graph import StateGraph
graph_builder = StateGraph(ReflectState)
graph_builder.add_node("draft_answer", draft_answer)
graph_builder.add_node("reflect", reflect)
graph_builder.add_node("rewrite", rewrite)
# Connections
start_edge = "draft_answer"
end_edge = None # will be set in condition
def should_end(state: Dict) -> str:
if state["verdict"] == "ok":
return "END"
if state["round"] >= state.get("max_rounds", 2):
return "END"
return "rewrite"
# Add edges with condition
from langgraph.graph import END
graph_builder.set_entry_point(start_edge)
graph_builder.add_conditional_edges(
"reflect",
should_end,
{
"rewrite": "rewrite",
"END": END,
},
)
# rewrite -> reflect
graph_builder.add_edge("rewrite", "reflect")
graph = graph_builder.compile()
if __name__ == "__main__":
if len(sys.argv) < 2:
print("Usage: python main.py \"Your question\"")
sys.exit(1)
question = sys.argv[1]
initial_state: ReflectState = {
"question": question,
"draft": "",
"critique": "",
"verdict": "",
"round": 0,
"max_rounds": 2,
}
result = graph.invoke(initial_state)
print("\n--- Final Answer ---")
print(result["draft"])
print("\n--- Critique ---")
print(result["critique"])