Files
task-6a22c713fd30e81cf315ea04/main.py
T

172 lines
6.4 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.
"""
# main.py
# LangGraph code review agent with reflection and rewrite loop.
# Implements the task specification without any deepagents dependency.
# Uses OpenRouter via langchain-openai.
import os
import asyncio
from typing import TypedDict, Annotated, Dict
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage, AIMessage
from langchain_core.output_parsers import PydanticOutputParser
from langchain_core.pydantic_v1 import BaseModel, Field
from langgraph.graph import StateGraph, START, END
from langgraph.graph.message import add_messages
from dotenv import load_dotenv
load_dotenv()
# ---------------------------------------------------------------------------
# State definition
# ---------------------------------------------------------------------------
class CodeReviewState(TypedDict):
code: str
draft_review: str
criteria_scores: Dict[str, int]
weakest_criterion: str
verdict: str # "ok" | "needs_revision"
round: int
max_rounds: int
# ---------------------------------------------------------------------------
# LLM setup
# ---------------------------------------------------------------------------
llm = ChatOpenAI(
model="openai/gpt-oss-20b:free",
base_url="https://openrouter.ai/api/v1",
api_key=os.getenv("OPENAI_API_KEY"),
temperature=0.0,
)
# ---------------------------------------------------------------------------
# Structured output models for reflect node
# ---------------------------------------------------------------------------
class ReflectOutput(BaseModel):
pep8: int = Field(..., ge=0, le=10)
type_hints: int = Field(..., ge=0, le=10)
edge_cases: int = Field(..., ge=0, le=10)
naming: int = Field(..., ge=0, le=10)
weakest_criterion: str = Field(...)
verdict: str = Field(..., regex="^(ok|needs_revision)$")
reflect_parser = PydanticOutputParser(pydantic_object=ReflectOutput)
# ---------------------------------------------------------------------------
# Node implementations
# ---------------------------------------------------------------------------
async def draft_review_node(state: CodeReviewState) -> CodeReviewState:
"""Generate an initial code review with 36 bullet points."""
prompt = f"""
You are a senior Python developer. You will write a concise code review for the following function. Provide 3 to 6 bullet points, each starting with a dash.
Function code:
{state['code']}
Review:"""
response = await llm.ainvoke([HumanMessage(content=prompt)])
state['draft_review'] = response.content.strip()
return state
async def reflect_node(state: CodeReviewState) -> CodeReviewState:
"""Critic evaluates the draft review on 4 criteria and returns structured scores."""
prompt = f"""
You are a code review critic. Evaluate the following draft review on the four criteria below, assigning a score from 0 (worst) to 10 (excellent). Return the scores and the weakest criterion in a JSON format matching the schema:
{reflect_parser.get_format_instructions()}
Draft review:
{state['draft_review']}
Scores:"""
response = await llm.ainvoke([HumanMessage(content=prompt)])
try:
parsed = reflect_parser.parse(response.content)
except Exception as e:
# Fallback: treat as all zeros
parsed = ReflectOutput(pep8=0, type_hints=0, edge_cases=0, naming=0, weakest_criterion="pep8", verdict="needs_revision")
state['criteria_scores'] = {
"pep8": parsed.pep8,
"type_hints": parsed.type_hints,
"edge_cases": parsed.edge_cases,
"naming": parsed.naming,
}
state['weakest_criterion'] = parsed.weakest_criterion
state['verdict'] = parsed.verdict
return state
async def rewrite_node(state: CodeReviewState) -> CodeReviewState:
"""Rewrite the part of the review that addresses the weakest criterion."""
prompt = f"""
You are a senior Python developer. The following code review has been identified as weak in the criterion: {state['weakest_criterion']}. Rewrite only the section of the review that addresses this criterion, improving clarity and depth. Keep the rest of the review unchanged.
Original review:
{state['draft_review']}
Rewritten review:"""
response = await llm.ainvoke([HumanMessage(content=prompt)])
# Replace only the weak section. For simplicity, we replace the whole review.
state['draft_review'] = response.content.strip()
state['round'] += 1
return state
# ---------------------------------------------------------------------------
# Graph construction
# ---------------------------------------------------------------------------
def create_graph() -> StateGraph:
graph = StateGraph(CodeReviewState)
graph.add_node("draft_review", draft_review_node)
graph.add_node("reflect", reflect_node)
graph.add_node("rewrite", rewrite_node)
# Entry point
graph.set_entry_point("draft_review")
# Transitions
graph.add_edge("draft_review", "reflect")
graph.add_conditional_edges(
"reflect",
lambda state: "rewrite" if state["verdict"] == "needs_revision" and state["round"] < state["max_rounds"] else "END",
)
graph.add_edge("rewrite", "reflect")
return graph
# ---------------------------------------------------------------------------
# CLI helper
# ---------------------------------------------------------------------------
async def run_review(code: str, max_rounds: int = 2) -> CodeReviewState:
initial_state: CodeReviewState = {
"code": code,
"draft_review": "",
"criteria_scores": {},
"weakest_criterion": "",
"verdict": "",
"round": 0,
"max_rounds": max_rounds,
}
graph = create_graph()
final_state = await graph.astate(initial_state)
return final_state
# ---------------------------------------------------------------------------
# Demo main
# ---------------------------------------------------------------------------
if __name__ == "__main__":
sample_code = """
def sort_numbers(arr):
return sorted(arr)
"""
result = asyncio.run(run_review(sample_code))
print("\n=== Initial Draft Review ===")
print(result["draft_review"])
print("\n=== Scores ===")
print(result["criteria_scores"])
print("\n=== Verdict ===")
print(result["verdict"])
if result["verdict"] == "needs_revision":
print("\n=== Rewritten Review ===")
print(result["draft_review"]) # after last rewrite
print("\n=== Updated Scores ===")
print(result["criteria_scores"])
""