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

117 lines
3.9 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 agent with reflection and rewrite.
This implementation follows the assignment requirements:
- Draft answer node
- Reflect node that critiques the draft
- Rewrite node that updates draft based on critique
- max_rounds default 2
- CLI entry point
"""
from typing import TypedDict, Dict
import os
from langgraph.graph import StateGraph, END
from langchain_openai import ChatOpenAI
# --- State definition -----------------------------------------------------
class ReflectState(TypedDict):
question: str
draft: str
critique: str
verdict: str # "ok" | "needs_revision"
round: int
max_rounds: int
# --- LLM setup ------------------------------------------------------------
# Use environment variable for API key; fallback to dummy for local testing
llm = ChatOpenAI(model_name="gpt-4o-mini", temperature=0.2)
# --- Node definitions -----------------------------------------------------
def draft_answer(state: ReflectState) -> Dict:
"""Generate initial draft answer to the question."""
question = state["question"]
prompt = f"Write a concise answer (510 sentences) to the following question: {question}"
response = llm.invoke(prompt)
return {"draft": response.content}
def reflect(state: ReflectState) -> Dict:
"""Critique the draft.
The node returns a verdict ('ok' or 'needs_revision') and 23 concise points of improvement.
"""
draft = state["draft"]
question = state["question"]
prompt = (
f"You are a critical reviewer. Evaluate the following draft answer to the question '{question}'. "
"Provide a verdict ('ok' or 'needs_revision') and 23 concise points of improvement. "
"Respond in JSON with keys 'verdict' and 'critique'."
)
response = llm.invoke(prompt)
# Expect JSON; simple parse
import json
data = json.loads(response.content)
verdict = data.get("verdict", "needs_revision")
critique = data.get("critique", "")
return {"verdict": verdict, "critique": critique}
def rewrite(state: ReflectState) -> Dict:
"""Rewrite draft based on critique and increment round."""
draft = state["draft"]
critique = state["critique"]
round_num = state["round"] + 1
prompt = (
f"Rewrite the following draft answer to improve it based on these points: {critique}. "
f"Keep the answer concise (510 sentences)."
)
response = llm.invoke(prompt)
return {"draft": response.content, "round": round_num}
# --- Graph construction ---------------------------------------------------
builder = StateGraph(ReflectState)
builder.add_node("draft_answer", draft_answer)
builder.add_node("reflect", reflect)
builder.add_node("rewrite", rewrite)
# Connections
builder.set_entry_point("draft_answer")
builder.add_edge("draft_answer", "reflect")
# Conditional after reflect: if ok -> END, else if round < max_rounds -> rewrite, else -> END
builder.add_conditional_edges(
"reflect",
lambda x: END if x["verdict"] == "ok" else "rewrite" if x["round"] < x["max_rounds"] else END,
)
builder.add_edge("rewrite", "reflect")
graph = builder.compile()
# --- CLI ---------------------------------------------------------------
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="LangGraph reflection demo")
parser.add_argument("question", type=str, help="Question to answer")
parser.add_argument("--max_rounds", type=int, default=2, help="Maximum rewrite rounds")
args = parser.parse_args()
initial_state: ReflectState = {
"question": args.question,
"draft": "",
"critique": "",
"verdict": "",
"round": 0,
"max_rounds": args.max_rounds,
}
# Run graph
result = graph.invoke(initial_state)
print("\nFinal answer:\n", result["draft"])
print("\nCritique:\n", result["critique"])
print("\nVerdict:\n", result["verdict"])
print("\nRounds used:\n", result["round"])