Files
task_6a1d75dbfd30e81cf3126b…/agent.py
T
2026-06-04 16:07:11 +00:00

135 lines
4.4 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.
"""
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)