Files
task-6a1d75d1fd30e81cf3126af8/main.py
T
2026-06-05 14:32:20 +00:00

103 lines
3.3 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.
import os
import re
import asyncio
from typing import TypedDict, Annotated
from langchain_openai import ChatOpenAI
from langgraph.graph import StateGraph, START, END
from langgraph.graph.message import add_messages
# LLM configuration (OpenRouter)
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,
)
# ---------- State definition ----------
class ReflectState(TypedDict):
question: str
draft: str
critique: str
verdict: str # ok | needs_revision
round: int
max_rounds: int
# ---------- Node implementations ----------
async def draft_answer(state: ReflectState) -> dict:
prompt = (
f"Write a concise answer (510 sentences) to the following question:\n\n"
f"Question: {state['question']}"
)
response = await llm.ainvoke([{"role": "user", "content": prompt}])
draft = response.content.strip()
return {"draft": draft, "round": 0}
async def reflect(state: ReflectState) -> dict:
prompt = (
f"You are a critical reviewer. Evaluate the following draft answer for completeness, specificity, and lack of filler.\n\n"
f"Draft: {state['draft']}\n\n"
f"Provide a verdict (ok or needs_revision) and 23 bullet points of critique."
)
response = await llm.ainvoke([{"role": "user", "content": prompt}])
text = response.content.strip()
verdict_match = re.search(r"(ok|needs_revision)", text, re.IGNORECASE)
verdict = verdict_match.group(1).lower() if verdict_match else "needs_revision"
return {"critique": text, "verdict": verdict}
async def rewrite(state: ReflectState) -> dict:
prompt = (
f"Rewrite the draft answer incorporating the following critique. Keep the answer concise (510 sentences).\n\n"
f"Critique: {state['critique']}\n\n"
f"Original Draft: {state['draft']}"
)
response = await llm.ainvoke([{"role": "user", "content": prompt}])
new_draft = response.content.strip()
return {"draft": new_draft, "round": state['round'] + 1}
# ---------- Graph construction ----------
builder = StateGraph(ReflectState)
builder.add_node("draft_answer", draft_answer)
builder.add_node("reflect", reflect)
builder.add_node("rewrite", rewrite)
builder.set_entry_point("draft_answer")
builder.add_edge("draft_answer", "reflect")
builder.add_conditional_edges(
"reflect",
lambda x: x["verdict"],
{
"ok": END,
"needs_revision": "rewrite",
},
)
builder.add_edge("rewrite", "reflect")
# Limit rounds
async def limit_rounds(state: ReflectState) -> str:
if state["round"] >= state["max_rounds"] and state["verdict"] == "needs_revision":
return END
return "reflect"
builder.add_conditional_edges("rewrite", limit_rounds, {"reflect": "reflect", END: END})
graph = builder.compile()
# ---------- Demo execution ----------
async def main():
question = "Объясни студенту разницу между tool и resource в MCP."
initial_state: ReflectState = {
"question": question,
"draft": "",
"critique": "",
"verdict": "",
"round": 0,
"max_rounds": 2,
}
final_state = await graph.ainvoke(initial_state)
print("\nFinal Answer:\n", final_state["draft"])
if __name__ == "__main__":
asyncio.run(main())