diff --git a/agent.py b/agent.py new file mode 100644 index 0000000..9282f74 --- /dev/null +++ b/agent.py @@ -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) \ No newline at end of file