""" LangGraph comparison agent demo. Usage: python -m src.main """ import os from typing import TypedDict, List, Dict from langgraph.graph import StateGraph from langchain_openai import ChatOpenAI from langchain_tavily import TavilySearch from dotenv import load_dotenv load_dotenv() # ---------- State ---------- class CompareState(TypedDict): entities: List[str] criteria: List[str] findings: Dict[str, List[str]] # entity -> list of notes per criterion final_table: str | None verdict: str | None # ---------- Nodes ---------- llm = ChatOpenAI(temperature=0) search_tool = TavilySearch(api_key=os.getenv("TAVILY_API_KEY")) async def plan_criteria(state: CompareState) -> CompareState: entities_str = ", ".join(state["entities"]) prompt = ( f"You are a helpful assistant. Given the following entities: {entities_str}. " "Generate 3-5 concise criteria for comparing them. Return a JSON array of strings." ) response = await llm.agenerate([{"role": "user", "content": prompt}]) # parse JSON import json try: criteria = json.loads(response.generations[0][0].text.strip()) except Exception as e: raise ValueError(f"Failed to parse criteria: {e}") state["criteria"] = criteria return state async def research_entity(state: CompareState) -> CompareState: # find next unprocessed entity-criterion pair for ent in state["entities"]: if ent not in state["findings"]: state["findings"][ent] = [] for crit in state["criteria"]: if any(crit.lower() in note.lower() for note in state["findings"][ent]): continue # perform search query = f"{ent} {crit}" result = await search_tool.ainvoke({"query": query, "max_results": 1}) snippet = result.get("results", [{}])[0].get("snippet", "") note = f"{crit}: {snippet}" if snippet else f"{crit}: no info" state["findings"][ent].append(note) return state # all processed return state async def build_table(state: CompareState) -> CompareState: rows = [] for crit in state["criteria"]: row = [crit] for ent in state["entities"]: notes = state["findings"][ent] note = next((n for n in notes if n.startswith(crit)), "") row.append(note) rows.append(row) # markdown table header = " | ".join(["Criterion"] + state["entities"]) sep = " | ".join([":---:"] * (len(state["entities"]) + 1)) table_rows = [f"{row[0]} | {' | '.join(row[1:])}" for row in rows] state["final_table"] = f"{header}\n{sep}\n" + "\n".join(table_rows) return state async def verdict(state: CompareState) -> CompareState: prompt = ( f"Based on the following table:\n{state['final_table']}\n" "Provide a short recommendation (2-4 sentences) on which entity is best for each use case." ) response = await llm.agenerate([{"role": "user", "content": prompt}]) state["verdict"] = response.generations[0][0].text.strip() return state # ---------- Graph ---------- builder = StateGraph(CompareState) builder.add_node("plan_criteria", plan_criteria) builder.add_node("research_entity", research_entity) builder.add_node("build_table", build_table) builder.add_node("verdict", verdict) builder.set_entry_point("plan_criteria") # transition logic builder.add_conditional_edges( "plan_criteria", lambda _: "research_entity", ) builder.add_conditional_edges( "research_entity", lambda state: "build_table" if all(ent in state["findings"] and len(state["findings"][ent]) == len(state["criteria"]) for ent in state["entities"]) else "research_entity", ) builder.add_edge("build_table", "verdict") builder.add_edge("verdict", "END") app = builder.compile() # ---------- Demo ---------- if __name__ == "__main__": import argparse parser = argparse.ArgumentParser(description="Compare three entities.") parser.add_argument("entities", nargs=3, help="Three entities to compare") args = parser.parse_args() init_state: CompareState = { "entities": list(args.entities), "criteria": [], "findings": {}, "final_table": None, "verdict": None, } result = app.invoke(init_state) print("\n=== Criteria ===") print(result["criteria"]) print("\n=== Table ===") print(result["final_table"]) print("\n=== Verdict ===") print(result["verdict"])