Add agent.py
This commit is contained in:
@@ -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 4–5 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 5–8 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)
|
||||
Reference in New Issue
Block a user