Files
task-6a22c713fd30e81cf315e9fe/src/compare_agent.py
T
2026-06-11 14:58:44 +00:00

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