Add src/deep_agents_from_scratch/deep_agent.py

This commit is contained in:
2026-06-05 10:31:47 +00:00
parent 6bb349579a
commit fd30ec8d17
+194
View File
@@ -0,0 +1,194 @@
"""Deep Agent implementation based on LangGraph StateGraph.
The agent follows the architecture described in the original
`deep-agents-from-scratch` notebook but is rewritten to satisfy the
current project constraints:
* All imports use the modern LangChain modules.
* The graph is fully connected the final node writes the virtual
files to disk.
* A lightweight inmemory virtual file system is used during the
conversation.
* The search tool uses Tavily and returns only a minimal summary.
* Summarization is performed by a small model (GPT4omini).
The module exposes a single helper ``build_deep_agent_graph`` which
returns a readytorun :class:`StateGraph` instance.
"""
from __future__ import annotations
from pathlib import Path
from typing import Dict, List, Any
# Modern LangChain imports
from langchain_core.messages import BaseMessage, HumanMessage
from langgraph.types import StateGraph
# Local imports the tools are defined in ``research_tools.py``
from .research_tools import (
tavily_search,
think_tool,
summarize_webpage_content,
)
from .state import DeepAgentState
# ---------------------------------------------------------------------------
# Helper: virtual file system
# ---------------------------------------------------------------------------
class VirtualFileSystem:
"""Simple inmemory file system used by the graph.
The state already contains a ``files`` mapping, but the graph
interacts with this helper to keep the node implementations
clean.
"""
def __init__(self, files: Dict[str, str]):
self._files = files
def write(self, filename: str, content: str) -> None:
self._files[filename] = content
def read(self, filename: str) -> str | None:
return self._files.get(filename)
def list(self) -> List[str]:
return list(self._files.keys())
def dump_to_disk(self, root: Path) -> None:
root.mkdir(parents=True, exist_ok=True)
for name, data in self._files.items():
(root / name).write_text(data, encoding="utf-8")
# ---------------------------------------------------------------------------
# Graph nodes
# ---------------------------------------------------------------------------
async def think_node(state: DeepAgentState) -> Dict[str, Any]:
"""Ask the agent to decide what to search next.
The node receives the current conversation and returns a new
query string.
"""
messages = state.messages
# Find the last user message
last_user = None
for msg in reversed(messages):
if msg.role == "user":
last_user = msg.content
break
if not last_user:
last_user = ""
# The tool expects a plain string
query = await think_tool.invoke(last_user)
return {"query": query}
async def search_node(state: DeepAgentState) -> Dict[str, Any]:
"""Perform a web search using the tavily tool.
The node receives the ``query`` from the previous node and returns
the raw search results.
"""
query = state.query
if not query:
return {"search_results": []}
results = await tavily_search.invoke(query)
return {"search_results": results}
async def summarize_node(state: DeepAgentState) -> Dict[str, Any]:
"""Summarize each search result and store the content in the virtual FS.
The node writes full webpage content to the virtual file system and
returns a list of summary strings.
"""
results = state.search_results
summaries: List[str] = []
vfs = VirtualFileSystem(state.files)
for i, res in enumerate(results):
content = res.get("content", "")
if not content:
continue
summary_obj = await summarize_webpage_content.invoke(content)
filename = f"search_result_{i+1}.md"
vfs.write(filename, content)
summaries.append(summary_obj["summary"])
return {"summaries": summaries}
async def write_file_node(state: DeepAgentState) -> Dict[str, Any]:
"""Write all virtual files to disk and add a final message to the chat.
The node is the final node of the graph. It writes the files to a
directory named ``output`` relative to the current working
directory.
"""
vfs = VirtualFileSystem(state.files)
output_dir = Path("output")
vfs.dump_to_disk(output_dir)
final_msg = f"Files written to {output_dir.resolve()}"
# Append the final message to the conversation history
return {"messages": [HumanMessage(content=final_msg)]}
# ---------------------------------------------------------------------------
# Graph construction
# ---------------------------------------------------------------------------
def build_deep_agent_graph() -> StateGraph[DeepAgentState]:
"""Build and return the StateGraph.
The graph follows the sequence:
1. think → 2. search → 3. summarize → 4. write_file
"""
graph = StateGraph(DeepAgentState)
# Add nodes
graph.add_node("think", think_node)
graph.add_node("search", search_node)
graph.add_node("summarize", summarize_node)
graph.add_node("write_file", write_file_node)
# Connect nodes
graph.set_entry_point("think")
graph.add_edge("think", "search")
graph.add_edge("search", "summarize")
graph.add_edge("summarize", "write_file")
return graph
# ---------------------------------------------------------------------------
# Helper to run the graph with an initial user message
# ---------------------------------------------------------------------------
async def run_agent(user_input: str) -> str:
"""Convenience wrapper that runs the graph for a single turn.
Parameters
----------
user_input: str
The user message to start the conversation.
Returns
-------
str
The final message produced by the graph.
"""
graph = build_deep_agent_graph()
# Initialise state
state: DeepAgentState = DeepAgentState(
messages=[HumanMessage(content=user_input)],
files={},
query=None,
search_results=[],
summaries=[],
)
# Run the graph ``invoke`` returns the final state
final_state = await graph.invoke(state)
# The final message was appended to the conversation
return final_state.messages[-1].content
# ---------------------------------------------------------------------------
# CLI entry point
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import asyncio
user_msg = input("Enter your research question: ")
result = asyncio.run(run_agent(user_msg))
print(result)