Update src/graph.py
This commit is contained in:
+10
-6
@@ -21,7 +21,7 @@ llm = ChatOpenAI(model="gpt-4o-mini", temperature=0)
|
|||||||
agent = create_agent(model=llm, tools=[])
|
agent = create_agent(model=llm, tools=[])
|
||||||
|
|
||||||
# --- Node implementations -----------------------------------------------
|
# --- Node implementations -----------------------------------------------
|
||||||
async def draft_answer(state: CodeReviewState) -> Dict[str, str]:
|
async def draft_review(state: CodeReviewState) -> Dict[str, str]:
|
||||||
code = state["code"]
|
code = state["code"]
|
||||||
prompt = (
|
prompt = (
|
||||||
"You are a senior Python developer.\n"
|
"You are a senior Python developer.\n"
|
||||||
@@ -71,23 +71,27 @@ async def agent_node(state: CodeReviewState) -> Dict[str, str]:
|
|||||||
|
|
||||||
# --- Graph construction ---------------------------------------------------
|
# --- Graph construction ---------------------------------------------------
|
||||||
builder = StateGraph(CodeReviewState)
|
builder = StateGraph(CodeReviewState)
|
||||||
builder.add_node("draft_answer", draft_answer)
|
builder.add_node("draft_review", draft_review)
|
||||||
builder.add_node("reflect", reflect)
|
builder.add_node("reflect", reflect)
|
||||||
builder.add_node("rewrite", rewrite)
|
builder.add_node("rewrite", rewrite)
|
||||||
builder.add_node("agent_node", agent_node)
|
builder.add_node("agent_node", agent_node)
|
||||||
|
|
||||||
builder.set_entry_point("draft_answer")
|
builder.set_entry_point("draft_review")
|
||||||
|
# After drafting, always go to reflect
|
||||||
builder.add_conditional_edges(
|
builder.add_conditional_edges(
|
||||||
"draft_answer",
|
"draft_review",
|
||||||
lambda x: "reflect" if True else None,
|
lambda x: "reflect" if True else None,
|
||||||
)
|
)
|
||||||
builder.add_edge("reflect", "agent_node") # compliance step
|
# From reflect to compliance step
|
||||||
|
builder.add_edge("reflect", "agent_node")
|
||||||
|
# Conditional rewrite based on verdict and round
|
||||||
builder.add_conditional_edges(
|
builder.add_conditional_edges(
|
||||||
"agent_node",
|
"agent_node",
|
||||||
lambda x: "rewrite" if x["verdict"] == "needs_revision" and x["round"] < x["max_rounds"] else None,
|
lambda x: "rewrite" if x["verdict"] == "needs_revision" and x["round"] < x["max_rounds"] else None,
|
||||||
)
|
)
|
||||||
|
# From rewrite back to reflect
|
||||||
builder.add_edge("rewrite", "reflect")
|
builder.add_edge("rewrite", "reflect")
|
||||||
# Final edge to END
|
# Final edge to END when verdict ok or max rounds reached
|
||||||
builder.add_conditional_edges(
|
builder.add_conditional_edges(
|
||||||
"reflect",
|
"reflect",
|
||||||
lambda x: "END" if x["verdict"] == "ok" or x["round"] >= x["max_rounds"] else None,
|
lambda x: "END" if x["verdict"] == "ok" or x["round"] >= x["max_rounds"] else None,
|
||||||
|
|||||||
Reference in New Issue
Block a user