Add src/brief.py

This commit is contained in:
2026-06-05 12:37:38 +00:00
parent 7410948700
commit 850f1c6c15
+140
View File
@@ -0,0 +1,140 @@
from __future__ import annotations
from typing import List, Optional
from langchain_openai import ChatOpenAI
from langchain_tavily import TavilySearch
from langgraph.graph import StateGraph, END
from langchain_core.messages import HumanMessage
import json
import re
import os
from dotenv import load_dotenv
# Load environment variables for OpenAI and Tavily
load_dotenv()
# ---------------------------------------------------------------------------
# State definition
# ---------------------------------------------------------------------------
class BriefState(dict):
"""Typed dictionary representing the agent state.
Attributes
----------
topic: str
The research topic supplied by the user.
outline: List[str] | None
45 point outline generated by the LLM.
step_index: int
Current index in the outline being processed.
notes: List[str]
Collected notes for each outline point.
final_brief: str | None
The synthesized brief.
"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
# Ensure all keys exist for type safety
self.setdefault("topic", "")
self.setdefault("outline", None)
self.setdefault("step_index", 0)
self.setdefault("notes", [])
self.setdefault("final_brief", None)
# ---------------------------------------------------------------------------
# Node implementations
# ---------------------------------------------------------------------------
def outline_node(state: BriefState) -> BriefState:
"""Generate a concise outline for the topic.
The LLM is instructed to return a JSON array of strings. The output
is parsed and stored in ``state['outline']``.
"""
llm = ChatOpenAI()
prompt = (
f"Create a concise 45 point outline for a research brief on the topic: "
f"{state['topic']}. Return a JSON array of strings."
)
response = llm.invoke([HumanMessage(content=prompt)])
text = response.content if hasattr(response, "content") else str(response)
text = text.strip()
# Try to parse JSON array
try:
outline: List[str] = json.loads(text)
except Exception:
# Fallback: extract first JSON array from the text
m = re.search(r"\[.*\]", text, re.S)
outline = json.loads(m.group(0)) if m else []
state["outline"] = outline
state["step_index"] = 0
state["notes"] = []
print("\n[Outline] Generated outline:")
for i, point in enumerate(outline, 1):
print(f"{i}. {point}")
return state
def research_step_node(state: BriefState) -> BriefState:
"""Search the web for the current outline point and produce a note.
The node performs a single Tavily search and then asks the LLM to
summarize the results into a short note (58 sentences).
"""
if state["step_index"] >= len(state["outline"]):
return state
point = state["outline"][state["step_index"]]
print(f"\n[Research Step {state['step_index'] + 1}] Searching for: {point}")
tavily = TavilySearch()
search_results = tavily.invoke({"query": point, "max_results": 3})
llm = ChatOpenAI()
summary_prompt = (
f"Summarize the following search results into a concise note (58 sentences) for the research brief:\n{search_results}"
)
summary = llm.invoke([HumanMessage(content=summary_prompt)])
note = summary.content if hasattr(summary, "content") else str(summary)
note = note.strip()
state["notes"].append(f"{note}")
state["step_index"] += 1
print(f"[Note] {note[:80]}...")
return state
def synthesize_node(state: BriefState) -> BriefState:
"""Combine all collected notes into the final brief."""
print("\n[Synthesize] Combining notes into final brief.")
brief = f"Research Brief on {state['topic']}:\n\n" + "\n".join(state["notes"])
state["final_brief"] = brief
print("\n[Final Brief]\n" + brief)
return state
# ---------------------------------------------------------------------------
# Graph construction
# ---------------------------------------------------------------------------
def build_graph(max_rounds: int = 2) -> StateGraph:
"""Return a LangGraph that implements the research brief workflow.
Parameters
----------
max_rounds: int
Maximum number of research steps (unused here but kept for
compatibility with the original task description).
"""
graph = StateGraph(BriefState)
graph.add_node("outline", outline_node)
graph.add_node("research", research_step_node)
graph.add_node("synthesize", synthesize_node)
graph.set_entry_point("outline")
graph.add_conditional_edges(
"outline",
lambda state: "research" if state["outline"] else "synthesize",
)
graph.add_conditional_edges(
"research",
lambda state: "research" if state["step_index"] < len(state["outline"]) else "synthesize",
)
graph.add_edge("synthesize", END)
return graph
# Export for external use
__all__ = ["BriefState", "build_graph"]