Files
2026-07-01 18:40:07 +00:00

178 lines
6.2 KiB
Python
Raw Permalink 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 sys
import asyncio
from typing import TypedDict, List, Dict
from dotenv import load_dotenv
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage
from langchain.tools import tool
from deepagents import create_deep_agent
from deepagents.backends import FilesystemBackend, LocalShellBackend, CompositeBackend
from langgraph.graph import StateGraph, START, END
from langchain_core.output_parsers import PydanticOutputParser
from pydantic import BaseModel, Field
from tavily import TavilyClient
# ---------- Load environment variables ----------
load_dotenv()
# ---------- LLM ----------
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 CompareState(TypedDict):
entities: List[str] # 3 имени для сравнения
criteria: List[str] # 35 критериев
findings: Dict[str, List[str]] # entity -> список заметок по критериям
final_table: str | None
verdict: str | None
# ---------- Pydantic parser for criteria ----------
class CriteriaOutput(BaseModel):
criteria: List[str] = Field(..., description="List of comparison criteria")
criteria_parser = PydanticOutputParser(pydantic_object=CriteriaOutput)
# ---------- Tavily client ----------
TAVILY_API_KEY = os.getenv("TAVILY_API_KEY")
if not TAVILY_API_KEY:
raise RuntimeError("TAVILY_API_KEY not set in environment")
client = TavilyClient(api_key=TAVILY_API_KEY)
# ---------- Graph nodes ----------
async def plan_criteria(state: CompareState) -> CompareState:
entities = ", ".join(state["entities"])
prompt = (
f"You are a comparison helper. Given the following entities: {entities}.\n"
"Generate 3-5 concise criteria for comparing these entities.\n"
"Return a JSON object with a field 'criteria' that is a list of strings."
)
result = await llm.ainvoke(prompt)
parsed = criteria_parser.parse(result.content)
state["criteria"] = parsed.criteria
# initialize findings dict
state["findings"] = {e: [] for e in state["entities"]}
return state
async def research_entity(state: CompareState) -> CompareState:
# Find next unprocessed entity
for entity in state["entities"]:
processed = len(state["findings"][entity])
if processed < len(state["criteria"]):
criterion = state["criteria"][processed]
query = f"{entity} {criterion}"
# Perform Tavily search
results = client.search(query, max_results=3)
# Take first result content, truncate to 200 chars
content = results[0]["content"][:200] if results else "No data found."
note = f"{criterion}: {content}"
state["findings"][entity].append(note)
break
return state
async def build_table(state: CompareState) -> CompareState:
headers = ["Criterion"] + state["entities"]
table = "| " + " | ".join(headers) + " |\n"
table += "| " + " | ".join(["---" for _ in headers]) + " |\n"
for idx, criterion in enumerate(state["criteria"]):
row = [criterion]
for entity in state["entities"]:
notes = state["findings"][entity]
note = notes[idx] if idx < len(notes) else ""
row.append(note)
table += "| " + " | ".join(row) + " |\n"
state["final_table"] = table
return state
async def verdict(state: CompareState) -> CompareState:
prompt = (
"You are an expert advisor. Based on the following comparison table,\n"
"provide a concise verdict of 24 sentences recommending which entity\n"
"is best for each possible use case.\n"
f"Comparison Table:\n{state['final_table']}"
)
result = await llm.ainvoke(prompt)
state["verdict"] = result.content.strip()
return state
# ---------- Graph construction ----------
def create_compare_graph() -> StateGraph[CompareState]:
graph = StateGraph(CompareState)
graph.add_node("plan_criteria", plan_criteria)
graph.add_node("research_entity", research_entity)
graph.add_node("build_table", build_table)
graph.add_node("verdict", verdict)
graph.set_entry_point("plan_criteria")
graph.add_edge("plan_criteria", "research_entity")
def is_done(state: CompareState) -> str:
return "build_table" if all(
len(state["findings"][e]) >= len(state["criteria"]) for e in state["entities"]
) else "research_entity"
graph.add_conditional_edges("research_entity", is_done)
graph.add_edge("build_table", "verdict")
graph.add_edge("verdict", END)
return graph
# ---------- DeepAgent tool ----------
@tool
def compare_entities(entities: List[str]) -> str:
"""Compare three entities and return a markdown table with verdict."""
if len(entities) != 3:
return "Error: Please provide exactly three entities to compare."
# Initialize state
state: CompareState = {
"entities": entities,
"criteria": [],
"findings": {},
"final_table": None,
"verdict": None,
}
graph = create_compare_graph()
# Run graph synchronously
final_state = graph.invoke(state)
table = final_state["final_table"] or "No table generated."
verdict = final_state["verdict"] or "No verdict generated."
return f"{table}\n\nVerdict:\n{verdict}"
# ---------- DeepAgent setup ----------
backend = CompositeBackend([
LocalShellBackend(workspace_dir="./workspace"),
FilesystemBackend(),
])
agent = create_deep_agent(
model=llm,
tools=[compare_entities],
backend=backend,
system_prompt="You are a helpful comparison assistant. Use the provided tool to compare entities.",
)
# ---------- CLI ----------
async def main():
if len(sys.argv) >= 4:
entities = sys.argv[1:4]
else:
entities = ["Chroma", "FAISS", "Qdrant"]
user_msg = f"Compare these entities: {', '.join(entities)}"
result = await agent.ainvoke(
{"messages": [HumanMessage(content=user_msg)]},
{"configurable": {"thread_id": "session-1"}},
)
print(result["messages"][-1].content)
if __name__ == "__main__":
asyncio.run(main())