Add agent.py

This commit is contained in:
2026-06-04 16:07:11 +00:00
parent 168229b3f1
commit d8de2eccfd
+135
View File
@@ -0,0 +1,135 @@
"""
Agent that generates a research brief using LangGraph, LangChain, and Tavily.
"""
import os
import sys
import logging
from typing import List, TypedDict
from dotenv import load_dotenv
from tavily import TavilyClient
from langgraph import StateGraph, node, END, START
from langchain_openai import ChatOpenAI
from langgraph.checkpoint.sqlite import SqliteSaver
# Load environment variables
load_dotenv()
TAVILY_API_KEY = os.getenv("TAVILY_API_KEY")
if not TAVILY_API_KEY:
raise RuntimeError("TAVILY_API_KEY not set in .env")
# LLM client
llm = ChatOpenAI(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o-mini")
# Tavily client
search_client = TavilyClient(api_key=TAVILY_API_KEY)
# State definition
class BriefState(TypedDict):
topic: str
outline: List[str] | None
step_index: int
notes: List[str]
final_brief: str | None
# Node: outline
@node
def outline(state: BriefState) -> BriefState:
prompt = f"Generate a concise 45 point outline for the topic: {state['topic']}. Return list of strings."
response = llm.invoke(prompt)
outline_text = response.content.strip()
outline_items = [line.strip() for line in outline_text.split("\n") if line.strip()]
if not outline_items:
raise ValueError("Outline generation returned empty list")
state["outline"] = outline_items
state["step_index"] = 0
state["notes"] = []
return state
# Node: research_step
@node
def research_step(state: BriefState) -> BriefState:
if state["outline"] is None:
raise ValueError("Outline not initialized")
idx = state["step_index"]
if idx >= len(state["outline"]):
return state
query = state["outline"][idx]
try:
results = search_client.search(query, max_results=1)
except Exception as e:
logging.warning(f"Tavily search failed for query '{query}': {e}. Retrying once.")
try:
results = search_client.search(query, max_results=1)
except Exception as e2:
logging.error(f"Tavily search failed again for query '{query}': {e2}. Skipping note.")
results = None
note = ""
if results and results.get("results"):
content = results["results"][0].get("content", "")
summary_prompt = (
f"Summarize the following content into a 58 sentence note.\n\n{content}"
)
try:
summary = llm.invoke(summary_prompt)
note = summary.content.strip()
except Exception as e:
logging.error(f"LLM summarization failed for query '{query}': {e}. Using raw content.")
note = content[:500].strip()
if note:
state["notes"].append(note)
else:
state["notes"].append(f"No information found for: {query}")
state["step_index"] = idx + 1
return state
# Node: synthesize
@node
def synthesize(state: BriefState) -> BriefState:
notes_text = "\n\n---\n\n".join(state["notes"])
prompt = f"Combine the following notes into a coherent brief with headings. Notes: {notes_text}."
response = llm.invoke(prompt)
state["final_brief"] = response.content.strip()
return state
# Build graph
workflow = StateGraph(BriefState)
workflow.add_node("outline", outline)
workflow.add_node("research_step", research_step)
workflow.add_node("synthesize", synthesize)
workflow.set_entry_point("outline")
workflow.add_edge("outline", "research_step")
workflow.add_conditional_edges(
"research_step",
lambda state: "synthesize" if state["step_index"] >= len(state["outline"]) else "research_step",
)
workflow.add_edge("synthesize", END)
graph = workflow.compile(checkpointer=SqliteSaver.from_conn_str("sqlite:///graph_state.db"))
def main(topic: str | None = None) -> None:
if not topic:
topic = "Artificial Intelligence"
initial_state: BriefState = {
"topic": topic,
"outline": None,
"step_index": 0,
"notes": [],
"final_brief": None,
}
result = graph.invoke(initial_state)
print("\n=== Outline ===")
if result["outline"]:
for i, item in enumerate(result["outline"], 1):
print(f"{i}. {item}")
print("\n=== Step Notes ===")
for i, note in enumerate(result["notes"], 1):
print(f"Step {i}: {note}\n")
print("\n=== Final Brief ===")
print(result["final_brief"] or "No brief generated.")
if __name__ == "__main__":
topic_arg = sys.argv[1] if len(sys.argv) > 1 else None
main(topic_arg)