Add src/compare_agent.py

This commit is contained in:
2026-06-11 14:58:44 +00:00
parent 36491d2a2e
commit e6e5b86cd6
+211
View File
@@ -0,0 +1,211 @@
"""
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", ""))