From 82a46efbd1e54cee6d632db91a4d32f11a0709d8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=D0=94=D0=B0=D0=BD=D0=B8=D0=B8=D0=BB=20=D0=92=D0=B8=D0=BA?= =?UTF-8?q?=D1=82=D0=BE=D1=80=D0=BE=D0=B2?= Date: Tue, 30 Jun 2026 07:57:38 +0000 Subject: [PATCH] add: main.py --- main.py | 207 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 207 insertions(+) create mode 100644 main.py diff --git a/main.py b/main.py new file mode 100644 index 0000000..0460c6c --- /dev/null +++ b/main.py @@ -0,0 +1,207 @@ +import os +import argparse +import asyncio +from typing import TypedDict, List, Dict, Optional + +from dotenv import load_dotenv +from langchain_openai import ChatOpenAI +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() +OPENAI_API_KEY = os.getenv("OPENAI_API_KEY") +TAVILY_API_KEY = os.getenv("TAVILY_API_KEY") + +# LLM configuration +llm = ChatOpenAI( + model="openai/gpt-oss-20b:free", + base_url="https://openrouter.ai/api/v1", + api_key=OPENAI_API_KEY, + temperature=0.0, +) + +# Tavily client +tavily_client = TavilyClient(api_key=TAVILY_API_KEY) + +# State definition +class CompareState(TypedDict): + entities: List[str] + criteria: List[str] + findings: Dict[str, List[str]] + final_table: Optional[str] + verdict: Optional[str] + +# Pydantic models for parsing +class CriteriaOutput(BaseModel): + criteria: List[str] = Field(description="List of comparison criteria") + +class VerdictOutput(BaseModel): + verdict: str = Field(description="Verdict text") + +criteria_parser = PydanticOutputParser(pydantic_object=CriteriaOutput) +verdict_parser = PydanticOutputParser(pydantic_object=VerdictOutput) + +# Node: plan_criteria +def plan_criteria(state: CompareState) -> CompareState: + prompt = ( + f"Given the following entities: {', '.join(state['entities'])}. " + "Generate 3 to 5 distinct criteria for comparing these entities. " + "Return a JSON object with a field 'criteria' that is a list of strings." + ) + response = llm.invoke(prompt) + parsed = criteria_parser.parse(response.content) + state["criteria"] = parsed.criteria + # Initialize findings dict + state["findings"] = {entity: [] for entity in state["entities"]} + return state + +# Node: research_entity +def research_entity(state: CompareState) -> CompareState: + # Find first entity with missing notes + for entity in state["entities"]: + if len(state["findings"][entity]) < len(state["criteria"]): + idx = len(state["findings"][entity]) + criterion = state["criteria"][idx] + query = f"{entity} {criterion}" + # Perform Tavily search + results = tavily_client.search(query) + # Take first result snippet + if results and results[0].snippet: + note = results[0].snippet.strip() + else: + note = f"No relevant info found for {entity} on {criterion}." + state["findings"][entity].append(note) + break + return state + +# Node: check_done +def check_done(state: CompareState) -> str: + for entity in state["entities"]: + if len(state["findings"][entity]) < len(state["criteria"]): + return "continue" + return "done" + +# Node: build_table +def build_table(state: CompareState) -> CompareState: + header = ["Criterion"] + state["entities"] + rows = [] + for criterion in state["criteria"]: + row = [criterion] + for entity in state["entities"]: + notes = state["findings"][entity] + idx = state["criteria"].index(criterion) + note = notes[idx] if idx < len(notes) else "" + row.append(note) + rows.append(row) + # Build markdown table + table_lines = ["| " + " | ".join(header) + " |"] + table_lines.append("|" + "|".join(["---"] * len(header)) + "|") + for row in rows: + table_lines.append("| " + " | ".join(row) + " |") + table_md = "\n".join(table_lines) + state["final_table"] = table_md + return state + +# Node: verdict +def verdict(state: CompareState) -> CompareState: + prompt = ( + f"Here is a markdown table comparing the entities:\n\n{state['final_table']}\n\n" + "Based on this table, provide a concise verdict recommending which entity " + "is best suited for a typical vector database use case. " + "Return a JSON object with a field 'verdict' that is a string." + ) + response = llm.invoke(prompt) + parsed = verdict_parser.parse(response.content) + state["verdict"] = parsed.verdict + return state + +# Build LangGraph +graph = StateGraph(CompareState) +graph.add_node("plan_criteria", plan_criteria) +graph.add_node("research_entity", research_entity) +graph.add_node("check_done", check_done) +graph.add_node("build_table", build_table) +graph.add_node("verdict", verdict) + +graph.add_edge(START, "plan_criteria") +graph.add_edge("plan_criteria", "research_entity") +graph.add_edge("research_entity", "check_done") +graph.add_conditional_edges( + "check_done", + lambda x: x, + { + "continue": "research_entity", + "done": "build_table", + }, +) +graph.add_edge("build_table", "verdict") +graph.add_edge("verdict", END) + +compiled_graph = graph.compile() + +# Tool: compare_entities +@tool +def compare_entities(query: str) -> str: + """ + Compare three entities based on user query. + Expected format: "Compare 3 vector DBs: Chroma, FAISS, Qdrant" + """ + # Extract entities after colon + if ":" in query: + parts = query.split(":", 1) + entities_part = parts[1] + else: + entities_part = query + entities = [e.strip() for e in entities_part.split(",") if e.strip()] + if len(entities) != 3: + return "Please provide exactly three entities separated by commas." + initial_state: CompareState = { + "entities": entities, + "criteria": [], + "findings": {}, + "final_table": None, + "verdict": None, + } + final_state = compiled_graph.invoke(initial_state) + table = final_state["final_table"] or "" + verdict_text = final_state["verdict"] or "" + return f"{table}\n\nVerdict:\n{verdict_text}" + +# 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 assistant that can compare three entities.", +) + +# CLI +def main(): + parser = argparse.ArgumentParser(description="Compare three entities.") + parser.add_argument( + "--query", + type=str, + default="Compare 3 vector DBs: Chroma, FAISS, Qdrant", + help="Comparison query in the format 'Compare 3 vector DBs: A, B, C'", + ) + args = parser.parse_args() + async def run(): + result = await agent.ainvoke( + {"messages": [{"role": "user", "content": args.query}]}, + {"configurable": {"thread_id": "session-1"}}, + ) + print(result["messages"][-1]["content"]) + asyncio.run(run()) + +if __name__ == "__main__": + main() \ No newline at end of file