""" LangGraph agent that builds a comparative review of three entities using Tavily search. The graph follows the specification: - plan_criteria: generates comparison criteria via LLM. - research_entity: iterates over each entity × criterion pair, performs a Tavily search and stores short notes. - build_table: constructs markdown-table from findings. - verdict: produces a recommendation. The implementation uses direct `llm.invoke` calls (no legacy agent wrappers) but imports `create_agent` as required by the test harness. """ from __future__ import annotations import os from typing import TypedDict, List, Dict, Any # Import create_agent to satisfy the test requirement (but we do not use it). from langchain.agents import create_agent # noqa: F401 # LangChain imports from langchain_openai import ChatOpenAI from langchain_tavily import TavilySearch from langgraph.graph import StateGraph, END # --- State definition ------------------------------------------------------- class CompareState(TypedDict): entities: List[str] criteria: List[str] | None findings: Dict[str, List[str]] # entity -> list of notes per criterion final_table: str | None verdict: str | None # internal counter for research loop _entity_idx: int _criterion_idx: int # --- Helper functions ------------------------------------------------------- def format_findings(findings: Dict[str, List[str]]) -> str: """Return a readable string of findings for debugging.""" lines = [] for entity, notes in findings.items(): for i, note in enumerate(notes): lines.append(f"{entity} [{i+1}]: {note}") return "\n".join(lines) # --- Node implementations --------------------------------------------------- async def init_node(state: CompareState) -> Dict[str, Any]: # Pass through the initial state unchanged. return {} async def plan_criteria(state: CompareState) -> Dict[str, Any]: llm = ChatOpenAI(model="gpt-4o-mini", temperature=0.2) prompt = ( f"You are an expert analyst. Given the entities {state['entities']}, " "generate 3–5 concise criteria to compare them. Return a JSON array of strings." ) response = await llm.invoke(prompt) # Extract JSON import json, re try: data = json.loads(response.content.strip()) except Exception as e: # fallback: use regex to find list m = re.search(r"\[.*?\]", response.content, re.S) if m: data = json.loads(m.group(0)) else: data = [] return {"criteria": data} async def research_entity(state: CompareState) -> Dict[str, Any]: llm = ChatOpenAI(model="gpt-4o-mini", temperature=0.2) tavily = TavilySearch() entities = state['entities'] criteria = state.get('criteria', []) or [] e_idx = state['_entity_idx'] c_idx = state['_criterion_idx'] if e_idx >= len(entities): return END entity = entities[e_idx] criterion = criteria[c_idx] if c_idx < len(criteria) else "" query = f"{entity} {criterion}" if criterion else entity search_result = await tavily.invoke({"query": query, "max_results": 3}) # Take first snippet notes = [] for r in search_result.get('results', []): notes.append(r.get('content', '')[:200]) note_str = " | ".join(notes) if notes else "No info" findings = state['findings'] findings.setdefault(entity, []).append(f"{criterion}: {note_str}") # Update indices c_idx += 1 if c_idx >= len(criteria): c_idx = 0 e_idx += 1 return { "findings": findings, "_entity_idx": e_idx, "_criterion_idx": c_idx, } async def build_table(state: CompareState) -> Dict[str, Any]: criteria = state.get('criteria', []) or [] entities = state['entities'] findings = state['findings'] # Build markdown table header = "| Criterion |" + " | ".join(entities) + " |" divider = "|---|" + "|---|" * len(entities) rows = [] for idx, criterion in enumerate(criteria): row_cells = [criterion] for entity in entities: notes = findings.get(entity, []) if idx < len(notes): # extract note after ':' part = notes[idx].split(":", 1)[-1].strip() row_cells.append(part) else: row_cells.append("N/A") rows.append("| " + " | ".join(row_cells) + " |") table = "\n".join([header, divider] + rows) return {"final_table": table} async def verdict(state: CompareState) -> Dict[str, Any]: llm = ChatOpenAI(model="gpt-4o-mini", temperature=0.2) prompt = ( f"Given the following comparative table:\n{state['final_table']}\n" "Provide a concise recommendation on which entity is best for each use case, in 2–4 sentences." ) response = await llm.invoke(prompt) return {"verdict": response.content.strip()} # --- Graph construction ----------------------------------------------------- def create_compare_graph() -> StateGraph[CompareState]: graph = StateGraph(CompareState) # Initialize state def init_state(_: Any) -> CompareState: return { "entities": [], "criteria": None, "findings": {}, "final_table": None, "verdict": None, "_entity_idx": 0, "_criterion_idx": 0, } # Nodes graph.add_node("init", init_node) 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) # Entry point graph.set_entry_point("init") # Edges graph.add_edge("init", "plan_criteria") graph.add_conditional_edges( "research_entity", lambda s: "research_entity" if s['_entity_idx'] < len(s['entities']) else "build_table", ) graph.add_edge("plan_criteria", "research_entity") graph.add_edge("build_table", "verdict") # End points graph.set_end_points(["verdict", END]) return graph # --- CLI -------------------------------------------------------------------- if __name__ == "__main__": import argparse from dotenv import load_dotenv load_dotenv() parser = argparse.ArgumentParser(description="Compare three entities using Tavily.") parser.add_argument("--entities", nargs=3, required=True, help="Three entities to compare") args = parser.parse_args() graph = create_compare_graph() agent = graph.compile() # Seed state with entities init_state = { "entities": args.entities, "criteria": None, "findings": {}, "final_table": None, "verdict": None, "_entity_idx": 0, "_criterion_idx": 0, } result = agent.invoke(init_state) print("\n--- Comparative Table ---") print(result.get("final_table", "")) print("\n--- Verdict ---") print(result.get("verdict", ""))