211 lines
6.9 KiB
Python
211 lines
6.9 KiB
Python
"""
|
||
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", "")) |