mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 15:42:25 +02:00
Merge branch 'main' into v1
This commit is contained in:
File diff suppressed because one or more lines are too long
@@ -0,0 +1,812 @@
|
||||
# Multi-agent supervisor
|
||||
|
||||
[**Supervisor**](../../concepts/multi_agent.md#supervisor) is a multi-agent architecture where **specialized** agents are coordinated by a central **supervisor agent**. The supervisor agent controls all communication flow and task delegation, making decisions about which agent to invoke based on the current context and task requirements.
|
||||
|
||||
In this tutorial, you will build a supervisor system with two agents — a research and a math expert. By the end of the tutorial you will:
|
||||
|
||||
1. Build specialized research and math agents
|
||||
2. Build a supervisor for orchestrating them with the prebuilt [`langgraph-supervisor`](https://langchain-ai.github.io/langgraph/agents/multi-agent/#supervisor)
|
||||
3. Build a supervisor from scratch
|
||||
4. Implement advanced task delegation
|
||||
|
||||

|
||||
|
||||
## Setup
|
||||
|
||||
First, let's install required packages and set our API keys
|
||||
|
||||
```python
|
||||
%%capture --no-stderr
|
||||
%pip install -U langgraph langgraph-supervisor langchain-tavily "langchain[openai]"
|
||||
```
|
||||
|
||||
```python
|
||||
import getpass
|
||||
import os
|
||||
|
||||
|
||||
def _set_if_undefined(var: str):
|
||||
if not os.environ.get(var):
|
||||
os.environ[var] = getpass.getpass(f"Please provide your {var}")
|
||||
|
||||
|
||||
_set_if_undefined("OPENAI_API_KEY")
|
||||
_set_if_undefined("TAVILY_API_KEY")
|
||||
```
|
||||
|
||||
!!! tip
|
||||
Sign up for LangSmith to quickly spot issues and improve the performance of your LangGraph projects. [LangSmith](https://docs.smith.langchain.com) lets you use trace data to debug, test, and monitor your LLM apps built with LangGraph.
|
||||
|
||||
## 1. Create worker agents
|
||||
|
||||
First, let's create our specialized worker agents — research agent and math agent:
|
||||
|
||||
* Research agent will have access to a web search tool using [Tavily API](https://tavily.com/)
|
||||
* Math agent will have access to simple math tools (`add`, `multiply`, `divide`)
|
||||
|
||||
### Research agent
|
||||
|
||||
For web search, we will use `TavilySearch` tool from `langchain-tavily`:
|
||||
|
||||
```python
|
||||
from langchain_tavily import TavilySearch
|
||||
|
||||
web_search = TavilySearch(max_results=3)
|
||||
web_search_results = web_search.invoke("who is the mayor of NYC?")
|
||||
|
||||
print(web_search_results["results"][0]["content"])
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Find events, attractions, deals, and more at nyctourism.com Skip Main Navigation Menu The Official Website of the City of New York Text Size Powered by Translate SearchSearch Primary Navigation The official website of NYC Home NYC Resources NYC311 Office of the Mayor Events Connect Jobs Search Office of the Mayor | Mayor's Bio | City of New York Secondary Navigation MayorBiographyNewsOfficials Eric L. Adams 110th Mayor of New York City Mayor Eric Adams has served the people of New York City as an NYPD officer, State Senator, Brooklyn Borough President, and now as the 110th Mayor of the City of New York. Mayor Eric Adams has served the people of New York City as an NYPD officer, State Senator, Brooklyn Borough President, and now as the 110th Mayor of the City of New York. He gave voice to a diverse coalition of working families in all five boroughs and is leading the fight to bring back New York City's economy, reduce inequality, improve public safety, and build a stronger, healthier city that delivers for all New Yorkers. As the representative of one of the nation's largest counties, Eric fought tirelessly to grow the local economy, invest in schools, reduce inequality, improve public safety, and advocate for smart policies and better government that delivers for all New Yorkers.
|
||||
```
|
||||
|
||||
To create individual worker agents, we will use LangGraph's prebuilt [agent](../../agents/agents.md).
|
||||
|
||||
```python
|
||||
from langgraph.prebuilt import create_react_agent
|
||||
|
||||
research_agent = create_react_agent(
|
||||
model="openai:gpt-4.1",
|
||||
tools=[web_search],
|
||||
prompt=(
|
||||
"You are a research agent.\n\n"
|
||||
"INSTRUCTIONS:\n"
|
||||
"- Assist ONLY with research-related tasks, DO NOT do any math\n"
|
||||
"- After you're done with your tasks, respond to the supervisor directly\n"
|
||||
"- Respond ONLY with the results of your work, do NOT include ANY other text."
|
||||
),
|
||||
name="research_agent",
|
||||
)
|
||||
```
|
||||
|
||||
Let's [run the agent](../../agents/run_agents.md) to verify that it behaves as expected.
|
||||
|
||||
!!! note "We'll use `pretty_print_messages` helper to render the streamed agent outputs nicely"
|
||||
|
||||
```python
|
||||
from langchain_core.messages import convert_to_messages
|
||||
|
||||
|
||||
def pretty_print_message(message, indent=False):
|
||||
pretty_message = message.pretty_repr(html=True)
|
||||
if not indent:
|
||||
print(pretty_message)
|
||||
return
|
||||
|
||||
indented = "\n".join("\t" + c for c in pretty_message.split("\n"))
|
||||
print(indented)
|
||||
|
||||
|
||||
def pretty_print_messages(update, last_message=False):
|
||||
is_subgraph = False
|
||||
if isinstance(update, tuple):
|
||||
ns, update = update
|
||||
# skip parent graph updates in the printouts
|
||||
if len(ns) == 0:
|
||||
return
|
||||
|
||||
graph_id = ns[-1].split(":")[0]
|
||||
print(f"Update from subgraph {graph_id}:")
|
||||
print("\n")
|
||||
is_subgraph = True
|
||||
|
||||
for node_name, node_update in update.items():
|
||||
update_label = f"Update from node {node_name}:"
|
||||
if is_subgraph:
|
||||
update_label = "\t" + update_label
|
||||
|
||||
print(update_label)
|
||||
print("\n")
|
||||
|
||||
messages = convert_to_messages(node_update["messages"])
|
||||
if last_message:
|
||||
messages = messages[-1:]
|
||||
|
||||
for m in messages:
|
||||
pretty_print_message(m, indent=is_subgraph)
|
||||
print("\n")
|
||||
```
|
||||
|
||||
```python
|
||||
from langchain_core.messages import convert_to_messages
|
||||
|
||||
|
||||
def pretty_print_message(message, indent=False):
|
||||
pretty_message = message.pretty_repr(html=True)
|
||||
if not indent:
|
||||
print(pretty_message)
|
||||
return
|
||||
|
||||
indented = "\n".join("\t" + c for c in pretty_message.split("\n"))
|
||||
print(indented)
|
||||
|
||||
|
||||
def pretty_print_messages(update, last_message=False):
|
||||
is_subgraph = False
|
||||
if isinstance(update, tuple):
|
||||
ns, update = update
|
||||
# skip parent graph updates in the printouts
|
||||
if len(ns) == 0:
|
||||
return
|
||||
|
||||
graph_id = ns[-1].split(":")[0]
|
||||
print(f"Update from subgraph {graph_id}:")
|
||||
print("\n")
|
||||
is_subgraph = True
|
||||
|
||||
for node_name, node_update in update.items():
|
||||
update_label = f"Update from node {node_name}:"
|
||||
if is_subgraph:
|
||||
update_label = "\t" + update_label
|
||||
|
||||
print(update_label)
|
||||
print("\n")
|
||||
|
||||
messages = convert_to_messages(node_update["messages"])
|
||||
if last_message:
|
||||
messages = messages[-1:]
|
||||
|
||||
for m in messages:
|
||||
pretty_print_message(m, indent=is_subgraph)
|
||||
print("\n")
|
||||
```
|
||||
|
||||
```python
|
||||
for chunk in research_agent.stream(
|
||||
{"messages": [{"role": "user", "content": "who is the mayor of NYC?"}]}
|
||||
):
|
||||
pretty_print_messages(chunk)
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: research_agent
|
||||
Tool Calls:
|
||||
tavily_search (call_U748rQhQXT36sjhbkYLSXQtJ)
|
||||
Call ID: call_U748rQhQXT36sjhbkYLSXQtJ
|
||||
Args:
|
||||
query: current mayor of New York City
|
||||
search_depth: basic
|
||||
|
||||
|
||||
Update from node tools:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: tavily_search
|
||||
|
||||
{"query": "current mayor of New York City", "follow_up_questions": null, "answer": null, "images": [], "results": [{"title": "List of mayors of New York City - Wikipedia", "url": "https://en.wikipedia.org/wiki/List_of_mayors_of_New_York_City", "content": "The mayor of New York City is the chief executive of the Government of New York City, as stipulated by New York City's charter.The current officeholder, the 110th in the sequence of regular mayors, is Eric Adams, a member of the Democratic Party.. During the Dutch colonial period from 1624 to 1664, New Amsterdam was governed by the Director of Netherland.", "score": 0.9039154, "raw_content": null}, {"title": "Office of the Mayor | Mayor's Bio | City of New York - NYC.gov", "url": "https://www.nyc.gov/office-of-the-mayor/bio.page", "content": "Mayor Eric Adams has served the people of New York City as an NYPD officer, State Senator, Brooklyn Borough President, and now as the 110th Mayor of the City of New York. He gave voice to a diverse coalition of working families in all five boroughs and is leading the fight to bring back New York City's economy, reduce inequality, improve", "score": 0.8405867, "raw_content": null}, {"title": "Eric Adams - Wikipedia", "url": "https://en.wikipedia.org/wiki/Eric_Adams", "content": "Eric Leroy Adams (born September 1, 1960) is an American politician and former police officer who has served as the 110th mayor of New York City since 2022. Adams was an officer in the New York City Transit Police and then the New York City Police Department (```
|
||||
```
|
||||
|
||||
### Math agent
|
||||
|
||||
For math agent tools we will use [vanilla Python functions](../../how-tos/tool-calling.md#define-a-tool):
|
||||
|
||||
```python
|
||||
def add(a: float, b: float):
|
||||
"""Add two numbers."""
|
||||
return a + b
|
||||
|
||||
|
||||
def multiply(a: float, b: float):
|
||||
"""Multiply two numbers."""
|
||||
return a * b
|
||||
|
||||
|
||||
def divide(a: float, b: float):
|
||||
"""Divide two numbers."""
|
||||
return a / b
|
||||
|
||||
|
||||
math_agent = create_react_agent(
|
||||
model="openai:gpt-4.1",
|
||||
tools=[add, multiply, divide],
|
||||
prompt=(
|
||||
"You are a math agent.\n\n"
|
||||
"INSTRUCTIONS:\n"
|
||||
"- Assist ONLY with math-related tasks\n"
|
||||
"- After you're done with your tasks, respond to the supervisor directly\n"
|
||||
"- Respond ONLY with the results of your work, do NOT include ANY other text."
|
||||
),
|
||||
name="math_agent",
|
||||
)
|
||||
```
|
||||
|
||||
Let's run the math agent:
|
||||
|
||||
```python
|
||||
for chunk in math_agent.stream(
|
||||
{"messages": [{"role": "user", "content": "what's (3 + 5) x 7"}]}
|
||||
):
|
||||
pretty_print_messages(chunk)
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: math_agent
|
||||
Tool Calls:
|
||||
add (call_p6OVLDHB4LyCNCxPOZzWR15v)
|
||||
Call ID: call_p6OVLDHB4LyCNCxPOZzWR15v
|
||||
Args:
|
||||
a: 3
|
||||
b: 5
|
||||
|
||||
|
||||
Update from node tools:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: add
|
||||
|
||||
8.0
|
||||
|
||||
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: math_agent
|
||||
Tool Calls:
|
||||
multiply (call_EoaWHMLFZAX4AkajQCtZvbli)
|
||||
Call ID: call_EoaWHMLFZAX4AkajQCtZvbli
|
||||
Args:
|
||||
a: 8
|
||||
b: 7
|
||||
|
||||
|
||||
Update from node tools:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: multiply
|
||||
|
||||
56.0
|
||||
|
||||
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: math_agent
|
||||
|
||||
56
|
||||
|
||||
|
||||
```
|
||||
|
||||
## 2. Create supervisor with `langgraph-supervisor`
|
||||
|
||||
To implement out multi-agent system, we will use [`create_supervisor`][langgraph_supervisor.supervisor.create_supervisor] from the prebuilt `langgraph-supervisor` library:
|
||||
|
||||
```python
|
||||
from langgraph_supervisor import create_supervisor
|
||||
from langchain.chat_models import init_chat_model
|
||||
|
||||
supervisor = create_supervisor(
|
||||
model=init_chat_model("openai:gpt-4.1"),
|
||||
agents=[research_agent, math_agent],
|
||||
prompt=(
|
||||
"You are a supervisor managing two agents:\n"
|
||||
"- a research agent. Assign research-related tasks to this agent\n"
|
||||
"- a math agent. Assign math-related tasks to this agent\n"
|
||||
"Assign work to one agent at a time, do not call agents in parallel.\n"
|
||||
"Do not do any work yourself."
|
||||
),
|
||||
add_handoff_back_messages=True,
|
||||
output_mode="full_history",
|
||||
).compile()
|
||||
```
|
||||
|
||||
```python
|
||||
from IPython.display import display, Image
|
||||
|
||||
display(Image(supervisor.get_graph().draw_mermaid_png()))
|
||||
```
|
||||
|
||||

|
||||
|
||||
**Note:** When you run this code, it will generate and display a visual representation of the supervisor graph showing the flow between the supervisor and worker agents.
|
||||
|
||||
Let's now run it with a query that requires both agents:
|
||||
|
||||
* research agent will look up the necessary GDP information
|
||||
* math agent will perform division to find the percentage of NY state GDP, as requested
|
||||
|
||||
```python
|
||||
for chunk in supervisor.stream(
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "find US and New York state GDP in 2024. what % of US GDP was New York state?",
|
||||
}
|
||||
]
|
||||
},
|
||||
):
|
||||
pretty_print_messages(chunk, last_message=True)
|
||||
|
||||
final_message_history = chunk["supervisor"]["messages"]
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_to_research_agent
|
||||
|
||||
Successfully transferred to research_agent
|
||||
|
||||
|
||||
Update from node research_agent:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_back_to_supervisor
|
||||
|
||||
Successfully transferred back to supervisor
|
||||
|
||||
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_to_math_agent
|
||||
|
||||
Successfully transferred to math_agent
|
||||
|
||||
|
||||
Update from node math_agent:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_back_to_supervisor
|
||||
|
||||
Successfully transferred back to supervisor
|
||||
|
||||
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: supervisor
|
||||
|
||||
In 2024, the US GDP was $29.18 trillion and New York State's GDP was $2.297 trillion. New York State accounted for approximately 7.87% of the total US GDP in 2024.
|
||||
|
||||
|
||||
```
|
||||
|
||||
## 3. Create supervisor from scratch
|
||||
|
||||
Let's now implement this same multi-agent system from scratch. We will need to:
|
||||
|
||||
1. [Set up how the supervisor communicates](#set-up-agent-communication) with individual agents
|
||||
2. [Create the supervisor agent](#create-supervisor-agent)
|
||||
3. Combine supervisor and worker agents into a [single multi-agent graph](#create-multi-agent-graph).
|
||||
|
||||
### Set up agent communication
|
||||
|
||||
We will need to define a way for the supervisor agent to communicate with the worker agents. A common way to implement this in multi-agent architectures is using **handoffs**, where one agent *hands off* control to another. Handoffs allow you to specify:
|
||||
|
||||
- **destination**: target agent to transfer to
|
||||
- **payload**: information to pass to that agent
|
||||
|
||||
We will implement handoffs via **handoff tools** and give these tools to the supervisor agent: when the supervisor calls these tools, it will hand off control to a worker agent, passing the full message history to that agent.
|
||||
|
||||
```python
|
||||
from typing import Annotated
|
||||
from langchain_core.tools import tool, InjectedToolCallId
|
||||
from langgraph.prebuilt import InjectedState
|
||||
from langgraph.graph import StateGraph, START, MessagesState
|
||||
from langgraph.types import Command
|
||||
|
||||
|
||||
def create_handoff_tool(*, agent_name: str, description: str | None = None):
|
||||
name = f"transfer_to_{agent_name}"
|
||||
description = description or f"Ask {agent_name} for help."
|
||||
|
||||
@tool(name, description=description)
|
||||
def handoff_tool(
|
||||
state: Annotated[MessagesState, InjectedState],
|
||||
tool_call_id: Annotated[str, InjectedToolCallId],
|
||||
) -> Command:
|
||||
tool_message = {
|
||||
"role": "tool",
|
||||
"content": f"Successfully transferred to {agent_name}",
|
||||
"name": name,
|
||||
"tool_call_id": tool_call_id,
|
||||
}
|
||||
# highlight-next-line
|
||||
return Command(
|
||||
# highlight-next-line
|
||||
goto=agent_name, # (1)!
|
||||
# highlight-next-line
|
||||
update={**state, "messages": state["messages"] + [tool_message]}, # (2)!
|
||||
# highlight-next-line
|
||||
graph=Command.PARENT, # (3)!
|
||||
)
|
||||
|
||||
return handoff_tool
|
||||
|
||||
|
||||
# Handoffs
|
||||
assign_to_research_agent = create_handoff_tool(
|
||||
agent_name="research_agent",
|
||||
description="Assign task to a researcher agent.",
|
||||
)
|
||||
|
||||
assign_to_math_agent = create_handoff_tool(
|
||||
agent_name="math_agent",
|
||||
description="Assign task to a math agent.",
|
||||
)
|
||||
```
|
||||
|
||||
1. Name of the agent or node to hand off to.
|
||||
2. Take the agent's messages and add them to the parent's state as part of the handoff. The next agent will see the parent state.
|
||||
3. Indicate to LangGraph that we need to navigate to agent node in a **parent** multi-agent graph.
|
||||
|
||||
### Create supervisor agent
|
||||
|
||||
Then, let's create the supervisor agent with the handoff tools we just defined. We will use the prebuilt [`create_react_agent`][langgraph.prebuilt.chat_agent_executor.create_react_agent]:
|
||||
|
||||
```python
|
||||
supervisor_agent = create_react_agent(
|
||||
model="openai:gpt-4.1",
|
||||
tools=[assign_to_research_agent, assign_to_math_agent],
|
||||
prompt=(
|
||||
"You are a supervisor managing two agents:\n"
|
||||
"- a research agent. Assign research-related tasks to this agent\n"
|
||||
"- a math agent. Assign math-related tasks to this agent\n"
|
||||
"Assign work to one agent at a time, do not call agents in parallel.\n"
|
||||
"Do not do any work yourself."
|
||||
),
|
||||
name="supervisor",
|
||||
)
|
||||
```
|
||||
|
||||
### Create multi-agent graph
|
||||
|
||||
Putting this all together, let's create a graph for our overall multi-agent system. We will add the supervisor and the individual agents as subgraph [nodes](../../concepts/low_level.md#nodes).
|
||||
|
||||
```python
|
||||
from langgraph.graph import END
|
||||
|
||||
# Define the multi-agent supervisor graph
|
||||
supervisor = (
|
||||
StateGraph(MessagesState)
|
||||
# NOTE: `destinations` is only needed for visualization and doesn't affect runtime behavior
|
||||
.add_node(supervisor_agent, destinations=("research_agent", "math_agent", END))
|
||||
.add_node(research_agent)
|
||||
.add_node(math_agent)
|
||||
.add_edge(START, "supervisor")
|
||||
# always return back to the supervisor
|
||||
.add_edge("research_agent", "supervisor")
|
||||
.add_edge("math_agent", "supervisor")
|
||||
.compile()
|
||||
)
|
||||
```
|
||||
|
||||
Notice that we've added explicit [edges](../../concepts/low_level.md#edges) from worker agents back to the supervisor — this means that they are guaranteed to return control back to the supervisor. If you want the agents to respond directly to the user (i.e., turn the system into a router, you can remove these edges).
|
||||
|
||||
```python
|
||||
from IPython.display import display, Image
|
||||
|
||||
display(Image(supervisor.get_graph().draw_mermaid_png()))
|
||||
```
|
||||
|
||||

|
||||
|
||||
**Note:** When you run this code, it will generate and display a visual representation of the multi-agent supervisor graph showing the flow between the supervisor and worker agents.
|
||||
|
||||
With the multi-agent graph created, let's now run it!
|
||||
|
||||
```python
|
||||
for chunk in supervisor.stream(
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "find US and New York state GDP in 2024. what % of US GDP was New York state?",
|
||||
}
|
||||
]
|
||||
},
|
||||
):
|
||||
pretty_print_messages(chunk, last_message=True)
|
||||
|
||||
final_message_history = chunk["supervisor"]["messages"]
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_to_research_agent
|
||||
|
||||
Successfully transferred to research_agent
|
||||
|
||||
|
||||
Update from node research_agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: research_agent
|
||||
|
||||
- US GDP in 2024 is projected to be about $28.18 trillion USD (Statista; CBO projection).
|
||||
- New York State's nominal GDP for 2024 is estimated at approximately $2.16 trillion USD (various economic reports).
|
||||
- New York State's share of US GDP in 2024 is roughly 7.7%.
|
||||
|
||||
Sources:
|
||||
- https://www.statista.com/statistics/216985/forecast-of-us-gross-domestic-product/
|
||||
- https://nyassembly.gov/Reports/WAM/2025economic_revenue/2025_report.pdf?v=1740533306
|
||||
|
||||
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_to_math_agent
|
||||
|
||||
Successfully transferred to math_agent
|
||||
|
||||
|
||||
Update from node math_agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: math_agent
|
||||
|
||||
US GDP in 2024: $28.18 trillion
|
||||
New York State GDP in 2024: $2.16 trillion
|
||||
Percentage of US GDP from New York State: 7.67%
|
||||
|
||||
|
||||
Update from node supervisor:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: supervisor
|
||||
|
||||
Here are your results:
|
||||
|
||||
- 2024 US GDP (projected): $28.18 trillion USD
|
||||
- 2024 New York State GDP (estimated): $2.16 trillion USD
|
||||
- New York State's share of US GDP: approximately 7.7%
|
||||
|
||||
If you need the calculation steps or sources, let me know!
|
||||
|
||||
|
||||
```
|
||||
|
||||
Let's examine the full resulting message history:
|
||||
|
||||
```python
|
||||
for message in final_message_history:
|
||||
message.pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================ Human Message ==================================
|
||||
|
||||
find US and New York state GDP in 2024. what % of US GDP was New York state?
|
||||
================================== Ai Message ===================================
|
||||
Name: supervisor
|
||||
Tool Calls:
|
||||
transfer_to_research_agent (call_KlGgvF5ahlAbjX8d2kHFjsC3)
|
||||
Call ID: call_KlGgvF5ahlAbjX8d2kHFjsC3
|
||||
Args:
|
||||
================================= Tool Message ==================================
|
||||
Name: transfer_to_research_agent
|
||||
|
||||
Successfully transferred to research_agent
|
||||
================================== Ai Message ===================================
|
||||
Name: research_agent
|
||||
Tool Calls:
|
||||
tavily_search (call_ZOaTVUA6DKrOjWQldLhtrsO2)
|
||||
Call ID: call_ZOaTVUA6DKrOjWQldLhtrsO2
|
||||
Args:
|
||||
query: US GDP 2024 estimate or actual
|
||||
search_depth: advanced
|
||||
tavily_search (call_QsRAasxW9K03lTlqjuhNLFbZ)
|
||||
Call ID: call_QsRAasxW9K03lTlqjuhNLFbZ
|
||||
Args:
|
||||
query: New York state GDP 2024 estimate or actual
|
||||
search_depth: advanced
|
||||
================================= Tool Message ==================================
|
||||
Name: tavily_search
|
||||
|
||||
{"query": "US GDP 2024 estimate or actual", "follow_up_questions": null, "answer": null, "images": [], "results": [{"url": "https://www.advisorperspectives.com/dshort/updates/2025/05/29/gdp-gross-domestic-product-q1-2025-second-estimate", "title": "Q1 GDP Second Estimate: Real GDP at -0.2%, Higher Than Expected", "content": "> Real gross domestic product (GDP) decreased at an annual rate of 0.2 percent in the first quarter of 2025 (January, February, and March), according to the second estimate released by the U.S. Bureau of Economic Analysis. In the fourth quarter of 2024, real GDP increased 2.4 percent. The decrease in real GDP in the first quarter primarily reflected an increase in imports, which are a subtraction in the calculation of GDP, and a decrease in government spending. These movements were partly [...] by [Harry Mamaysky](https://www.advisor```
|
||||
```
|
||||
|
||||
!!! important
|
||||
You can see that the supervisor system appends **all** of the individual agent messages (i.e., their internal tool-calling loop) to the full message history. This means that on every supervisor turn, supervisor agent sees this full history. If you want more control over:
|
||||
|
||||
* **how inputs are passed to agents**: you can use LangGraph [`Send()`][langgraph.types.Send] primitive to directly send data to the worker agents during the handoff. See the [task delegation](#4-create-delegation-tasks) example below
|
||||
* **how agent outputs are added**: you can control how much of the agent's internal message history is added to the overall supervisor message history by wrapping the agent in a separate node function:
|
||||
|
||||
```python
|
||||
def call_research_agent(state):
|
||||
# return agent's final response,
|
||||
# excluding inner monologue
|
||||
response = research_agent.invoke(state)
|
||||
# highlight-next-line
|
||||
return {"messages": response["messages"][-1]}
|
||||
```
|
||||
|
||||
## 4. Create delegation tasks
|
||||
|
||||
So far the individual agents relied on **interpreting full message history** to determine their tasks. An alternative approach is to ask the supervisor to **formulate a task explicitly**. We can do so by adding a `task_description` parameter to the `handoff_tool` function.
|
||||
|
||||
```python
|
||||
from langgraph.types import Send
|
||||
|
||||
|
||||
def create_task_description_handoff_tool(
|
||||
*, agent_name: str, description: str | None = None
|
||||
):
|
||||
name = f"transfer_to_{agent_name}"
|
||||
description = description or f"Ask {agent_name} for help."
|
||||
|
||||
@tool(name, description=description)
|
||||
def handoff_tool(
|
||||
# this is populated by the supervisor LLM
|
||||
task_description: Annotated[
|
||||
str,
|
||||
"Description of what the next agent should do, including all of the relevant context.",
|
||||
],
|
||||
# these parameters are ignored by the LLM
|
||||
state: Annotated[MessagesState, InjectedState],
|
||||
) -> Command:
|
||||
task_description_message = {"role": "user", "content": task_description}
|
||||
agent_input = {**state, "messages": [task_description_message]}
|
||||
return Command(
|
||||
# highlight-next-line
|
||||
goto=[Send(agent_name, agent_input)],
|
||||
graph=Command.PARENT,
|
||||
)
|
||||
|
||||
return handoff_tool
|
||||
|
||||
|
||||
assign_to_research_agent_with_description = create_task_description_handoff_tool(
|
||||
agent_name="research_agent",
|
||||
description="Assign task to a researcher agent.",
|
||||
)
|
||||
|
||||
assign_to_math_agent_with_description = create_task_description_handoff_tool(
|
||||
agent_name="math_agent",
|
||||
description="Assign task to a math agent.",
|
||||
)
|
||||
|
||||
supervisor_agent_with_description = create_react_agent(
|
||||
model="openai:gpt-4.1",
|
||||
tools=[
|
||||
assign_to_research_agent_with_description,
|
||||
assign_to_math_agent_with_description,
|
||||
],
|
||||
prompt=(
|
||||
"You are a supervisor managing two agents:\n"
|
||||
"- a research agent. Assign research-related tasks to this assistant\n"
|
||||
"- a math agent. Assign math-related tasks to this assistant\n"
|
||||
"Assign work to one agent at a time, do not call agents in parallel.\n"
|
||||
"Do not do any work yourself."
|
||||
),
|
||||
name="supervisor",
|
||||
)
|
||||
|
||||
supervisor_with_description = (
|
||||
StateGraph(MessagesState)
|
||||
.add_node(
|
||||
supervisor_agent_with_description, destinations=("research_agent", "math_agent")
|
||||
)
|
||||
.add_node(research_agent)
|
||||
.add_node(math_agent)
|
||||
.add_edge(START, "supervisor")
|
||||
.add_edge("research_agent", "supervisor")
|
||||
.add_edge("math_agent", "supervisor")
|
||||
.compile()
|
||||
)
|
||||
```
|
||||
|
||||
!!! note
|
||||
We're using [`Send()`][langgraph.types.Send] primitive in the `handoff_tool`. This means that instead of receiving the full `supervisor` graph state as input, each worker agent only sees the contents of the `Send` payload. In this example, we're sending the task description as a single "human" message.
|
||||
|
||||
Let's now running it with the same input query:
|
||||
|
||||
```python
|
||||
for chunk in supervisor_with_description.stream(
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "find US and New York state GDP in 2024. what % of US GDP was New York state?",
|
||||
}
|
||||
]
|
||||
},
|
||||
subgraphs=True,
|
||||
):
|
||||
pretty_print_messages(chunk, last_message=True)
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from subgraph supervisor:
|
||||
|
||||
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: supervisor
|
||||
Tool Calls:
|
||||
transfer_to_research_agent (call_tk8q8py8qK6MQz6Kj6mijKua)
|
||||
Call ID: call_tk8q8py8qK6MQz6Kj6mijKua
|
||||
Args:
|
||||
task_description: Find the 2024 GDP (Gross Domestic Product) for both the United States and New York state, using the most up-to-date and reputable sources available. Provide both GDP values and cite the data sources.
|
||||
|
||||
|
||||
Update from subgraph research_agent:
|
||||
|
||||
|
||||
Update from node agent:
|
||||
|
||||
|
||||
================================== Ai Message ==================================
|
||||
Name: research_agent
|
||||
Tool Calls:
|
||||
tavily_search (call_KqvhSvOIhAvXNsT6BOwbPlRB)
|
||||
Call ID: call_KqvhSvOIhAvXNsT6BOwbPlRB
|
||||
Args:
|
||||
query: 2024 United States GDP value from a reputable source
|
||||
search_depth: advanced
|
||||
tavily_search (call_kbbAWBc9KwCWKHmM5v04H88t)
|
||||
Call ID: call_kbbAWBc9KwCWKHmM5v04H88t
|
||||
Args:
|
||||
query: 2024 New York state GDP value from a reputable source
|
||||
search_depth: advanced
|
||||
|
||||
|
||||
Update from subgraph research_agent:
|
||||
|
||||
|
||||
Update from node tools:
|
||||
|
||||
|
||||
================================= Tool Message ==================================
|
||||
Name: tavily_search
|
||||
|
||||
{"query": "2024 United States GDP value from a reputable source", "follow_up_questions": null, "answer": null, "images": [], "results": [{"url": "https://www.focus-economics.com/countries/united-states/", "title": "United States Economy Overview - Focus Economics", "content": "The United States' Macroeconomic Analysis:\n------------------------------------------\n\n**Nominal GDP of USD 29,185 billion in 2024.**\n\n**Nominal GDP of USD 29,179 billion in 2024.**\n\n**GDP per capita of USD 86,635 compared to the global average of USD 10,589.**\n\n**GDP per capita of USD 86,652 compared to the global average of USD 10,589.**\n\n**Average real GDP growth of 2.5% over the last decade.**\n\n**Average real GDP growth of ```
|
||||
```
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 25 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 14 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 14 KiB |
@@ -1,21 +0,0 @@
|
||||
# Examples
|
||||
|
||||
The pages in this section provide end-to-end examples for the following topics:
|
||||
|
||||
## General
|
||||
|
||||
- [Agentic RAG](./rag/langgraph_adaptive_rag.ipynb)
|
||||
- [Agent Supervisor](./multi_agent/agent_supervisor.ipynb)
|
||||
- [SQL agent](./sql-agent.ipynb)
|
||||
- [Graph runs in LangSmith](../how-tos/run-id-langsmith.ipynb)
|
||||
|
||||
## LangGraph Platform
|
||||
|
||||
- [Set up custom authentication](./auth/getting_started.md)
|
||||
- [Make conversations private](./auth/resource_auth.md)
|
||||
- [Connect an authentication provider](./auth/add_auth_server.md)
|
||||
- [Rebuild graph at runtime](../cloud/deployment/graph_rebuild.md)
|
||||
- [Use RemoteGraph](../how-tos/use-remote-graph.md)
|
||||
- [Deploy CrewAI, AutoGen, and other frameworks](../how-tos/autogen-langgraph-platform.ipynb)
|
||||
- [Integrate LangGraph into a React app](../cloud/how-tos/use_stream_react.md)
|
||||
- [Implement Generative User Interfaces with LangGraph](../cloud/how-tos/generative_ui_react.md)
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 80 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 102 KiB |
File diff suppressed because one or more lines are too long
@@ -0,0 +1,522 @@
|
||||
# Agentic RAG
|
||||
|
||||
In this tutorial we will build a [retrieval agent](https://python.langchain.com/docs/tutorials/qa_chat_history). Retrieval agents are useful when you want an LLM to make a decision about whether to retrieve context from a vectorstore or respond to the user directly.
|
||||
|
||||
By the end of the tutorial we will have done the following:
|
||||
|
||||
1. Fetch and preprocess documents that will be used for retrieval.
|
||||
2. Index those documents for semantic search and create a retriever tool for the agent.
|
||||
3. Build an agentic RAG system that can decide when to use the retriever tool.
|
||||
|
||||

|
||||
|
||||
## Setup
|
||||
|
||||
Let's download the required packages and set our API keys:
|
||||
|
||||
```python
|
||||
%%capture --no-stderr
|
||||
%pip install -U --quiet langgraph "langchain[openai]" langchain-community langchain-text-splitters
|
||||
```
|
||||
|
||||
```python
|
||||
import getpass
|
||||
import os
|
||||
|
||||
|
||||
def _set_env(key: str):
|
||||
if key not in os.environ:
|
||||
os.environ[key] = getpass.getpass(f"{key}:")
|
||||
|
||||
|
||||
_set_env("OPENAI_API_KEY")
|
||||
```
|
||||
|
||||
!!! tip
|
||||
Sign up for LangSmith to quickly spot issues and improve the performance of your LangGraph projects. [LangSmith](https://docs.smith.langchain.com) lets you use trace data to debug, test, and monitor your LLM apps built with LangGraph.
|
||||
|
||||
|
||||
## 1. Preprocess documents
|
||||
|
||||
1. Fetch documents to use in our RAG system. We will use three of the most recent pages from [Lilian Weng's excellent blog](https://lilianweng.github.io/). We'll start by fetching the content of the pages using `WebBaseLoader` utility:
|
||||
|
||||
```python
|
||||
from langchain_community.document_loaders import WebBaseLoader
|
||||
|
||||
urls = [
|
||||
"https://lilianweng.github.io/posts/2024-11-28-reward-hacking/",
|
||||
"https://lilianweng.github.io/posts/2024-07-07-hallucination/",
|
||||
"https://lilianweng.github.io/posts/2024-04-12-diffusion-video/",
|
||||
]
|
||||
|
||||
docs = [WebBaseLoader(url).load() for url in urls]
|
||||
```
|
||||
|
||||
```python
|
||||
docs[0][0].page_content.strip()[:1000]
|
||||
```
|
||||
|
||||
2. Split the fetched documents into smaller chunks for indexing into our vectorstore:
|
||||
|
||||
```python
|
||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||
|
||||
docs_list = [item for sublist in docs for item in sublist]
|
||||
|
||||
text_splitter = RecursiveCharacterTextSplitter.from_tiktoken_encoder(
|
||||
chunk_size=100, chunk_overlap=50
|
||||
)
|
||||
doc_splits = text_splitter.split_documents(docs_list)
|
||||
```
|
||||
|
||||
```python
|
||||
doc_splits[0].page_content.strip()
|
||||
```
|
||||
|
||||
## 2. Create a retriever tool
|
||||
|
||||
Now that we have our split documents, we can index them into a vector store that we'll use for semantic search.
|
||||
|
||||
1. Use an in-memory vector store and OpenAI embeddings:
|
||||
|
||||
```python
|
||||
from langchain_core.vectorstores import InMemoryVectorStore
|
||||
from langchain_openai import OpenAIEmbeddings
|
||||
|
||||
vectorstore = InMemoryVectorStore.from_documents(
|
||||
documents=doc_splits, embedding=OpenAIEmbeddings()
|
||||
)
|
||||
retriever = vectorstore.as_retriever()
|
||||
```
|
||||
|
||||
2. Create a retriever tool using LangChain's prebuilt `create_retriever_tool`:
|
||||
|
||||
```python
|
||||
from langchain.tools.retriever import create_retriever_tool
|
||||
|
||||
retriever_tool = create_retriever_tool(
|
||||
retriever,
|
||||
"retrieve_blog_posts",
|
||||
"Search and return information about Lilian Weng blog posts.",
|
||||
)
|
||||
```
|
||||
|
||||
3. Test the tool:
|
||||
|
||||
```python
|
||||
retriever_tool.invoke({"query": "types of reward hacking"})
|
||||
```
|
||||
|
||||
## 3. Generate query
|
||||
|
||||
Now we will start building components ([nodes](../../concepts/low_level.md#nodes) and [edges](../../concepts/low_level.md#edges)) for our agentic RAG graph. Note that the components will operate on the [`MessagesState`](../../concepts/low_level.md#messagesstate) — graph state that contains a `messages` key with a list of [chat messages](https://python.langchain.com/docs/concepts/messages/).
|
||||
|
||||
1. Build a `generate_query_or_respond` node. It will call an LLM to generate a response based on the current graph state (list of messages). Given the input messages, it will decide to retrieve using the retriever tool, or respond directly to the user. Note that we're giving the chat model access to the `retriever_tool` we created earlier via `.bind_tools`:
|
||||
|
||||
```python
|
||||
from langgraph.graph import MessagesState
|
||||
from langchain.chat_models import init_chat_model
|
||||
|
||||
response_model = init_chat_model("openai:gpt-4.1", temperature=0)
|
||||
|
||||
|
||||
def generate_query_or_respond(state: MessagesState):
|
||||
"""Call the model to generate a response based on the current state. Given
|
||||
the question, it will decide to retrieve using the retriever tool, or simply respond to the user.
|
||||
"""
|
||||
response = (
|
||||
response_model
|
||||
# highlight-next-line
|
||||
.bind_tools([retriever_tool]).invoke(state["messages"])
|
||||
)
|
||||
return {"messages": [response]}
|
||||
```
|
||||
|
||||
2. Try it on a random input:
|
||||
|
||||
```python
|
||||
input = {"messages": [{"role": "user", "content": "hello!"}]}
|
||||
generate_query_or_respond(input)["messages"][-1].pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================== Ai Message ==================================
|
||||
|
||||
Hello! How can I help you today?
|
||||
```
|
||||
|
||||
3. Ask a question that requires semantic search:
|
||||
|
||||
```python
|
||||
input = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
}
|
||||
]
|
||||
}
|
||||
generate_query_or_respond(input)["messages"][-1].pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
retrieve_blog_posts (call_tYQxgfIlnQUDMdtAhdbXNwIM)
|
||||
Call ID: call_tYQxgfIlnQUDMdtAhdbXNwIM
|
||||
Args:
|
||||
query: types of reward hacking
|
||||
```
|
||||
|
||||
## 4. Grade documents
|
||||
|
||||
1. Add a [conditional edge](../../concepts/low_level.md#conditional-edges) — `grade_documents` — to determine whether the retrieved documents are relevant to the question. We will use a model with a structured output schema `GradeDocuments` for document grading. The `grade_documents` function will return the name of the node to go to based on the grading decision (`generate_answer` or `rewrite_question`):
|
||||
|
||||
```python
|
||||
from pydantic import BaseModel, Field
|
||||
from typing import Literal
|
||||
|
||||
GRADE_PROMPT = (
|
||||
"You are a grader assessing relevance of a retrieved document to a user question. \n "
|
||||
"Here is the retrieved document: \n\n {context} \n\n"
|
||||
"Here is the user question: {question} \n"
|
||||
"If the document contains keyword(s) or semantic meaning related to the user question, grade it as relevant. \n"
|
||||
"Give a binary score 'yes' or 'no' score to indicate whether the document is relevant to the question."
|
||||
)
|
||||
|
||||
|
||||
# highlight-next-line
|
||||
class GradeDocuments(BaseModel):
|
||||
"""Grade documents using a binary score for relevance check."""
|
||||
|
||||
binary_score: str = Field(
|
||||
description="Relevance score: 'yes' if relevant, or 'no' if not relevant"
|
||||
)
|
||||
|
||||
|
||||
grader_model = init_chat_model("openai:gpt-4.1", temperature=0)
|
||||
|
||||
|
||||
def grade_documents(
|
||||
state: MessagesState,
|
||||
) -> Literal["generate_answer", "rewrite_question"]:
|
||||
"""Determine whether the retrieved documents are relevant to the question."""
|
||||
question = state["messages"][0].content
|
||||
context = state["messages"][-1].content
|
||||
|
||||
prompt = GRADE_PROMPT.format(question=question, context=context)
|
||||
response = (
|
||||
grader_model
|
||||
# highlight-next-line
|
||||
.with_structured_output(GradeDocuments).invoke(
|
||||
[{"role": "user", "content": prompt}]
|
||||
)
|
||||
)
|
||||
score = response.binary_score
|
||||
|
||||
if score == "yes":
|
||||
return "generate_answer"
|
||||
else:
|
||||
return "rewrite_question"
|
||||
```
|
||||
|
||||
2. Run this with irrelevant documents in the tool response:
|
||||
|
||||
```python
|
||||
from langchain_core.messages import convert_to_messages
|
||||
|
||||
input = {
|
||||
"messages": convert_to_messages(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "1",
|
||||
"name": "retrieve_blog_posts",
|
||||
"args": {"query": "types of reward hacking"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "content": "meow", "tool_call_id": "1"},
|
||||
]
|
||||
)
|
||||
}
|
||||
grade_documents(input)
|
||||
```
|
||||
|
||||
3. Confirm that the relevant documents are classified as such:
|
||||
|
||||
```python
|
||||
input = {
|
||||
"messages": convert_to_messages(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "1",
|
||||
"name": "retrieve_blog_posts",
|
||||
"args": {"query": "types of reward hacking"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "reward hacking can be categorized into two types: environment or goal misspecification, and reward tampering",
|
||||
"tool_call_id": "1",
|
||||
},
|
||||
]
|
||||
)
|
||||
}
|
||||
grade_documents(input)
|
||||
```
|
||||
|
||||
## 5. Rewrite question
|
||||
|
||||
1. Build the `rewrite_question` node. The retriever tool can return potentially irrelevant documents, which indicates a need to improve the original user question. To do so, we will call the `rewrite_question` node:
|
||||
|
||||
```python
|
||||
REWRITE_PROMPT = (
|
||||
"Look at the input and try to reason about the underlying semantic intent / meaning.\n"
|
||||
"Here is the initial question:"
|
||||
"\n ------- \n"
|
||||
"{question}"
|
||||
"\n ------- \n"
|
||||
"Formulate an improved question:"
|
||||
)
|
||||
|
||||
|
||||
def rewrite_question(state: MessagesState):
|
||||
"""Rewrite the original user question."""
|
||||
messages = state["messages"]
|
||||
question = messages[0].content
|
||||
prompt = REWRITE_PROMPT.format(question=question)
|
||||
response = response_model.invoke([{"role": "user", "content": prompt}])
|
||||
return {"messages": [{"role": "user", "content": response.content}]}
|
||||
```
|
||||
|
||||
2. Try it out:
|
||||
|
||||
```python
|
||||
input = {
|
||||
"messages": convert_to_messages(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "1",
|
||||
"name": "retrieve_blog_posts",
|
||||
"args": {"query": "types of reward hacking"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "content": "meow", "tool_call_id": "1"},
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
response = rewrite_question(input)
|
||||
print(response["messages"][-1]["content"])
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
What are the different types of reward hacking described by Lilian Weng, and how does she explain them?
|
||||
```
|
||||
|
||||
## 6. Generate an answer
|
||||
|
||||
1. Build `generate_answer` node: if we pass the grader checks, we can generate the final answer based on the original question and the retrieved context:
|
||||
|
||||
```python
|
||||
GENERATE_PROMPT = (
|
||||
"You are an assistant for question-answering tasks. "
|
||||
"Use the following pieces of retrieved context to answer the question. "
|
||||
"If you don't know the answer, just say that you don't know. "
|
||||
"Use three sentences maximum and keep the answer concise.\n"
|
||||
"Question: {question} \n"
|
||||
"Context: {context}"
|
||||
)
|
||||
|
||||
|
||||
def generate_answer(state: MessagesState):
|
||||
"""Generate an answer."""
|
||||
question = state["messages"][0].content
|
||||
context = state["messages"][-1].content
|
||||
prompt = GENERATE_PROMPT.format(question=question, context=context)
|
||||
response = response_model.invoke([{"role": "user", "content": prompt}])
|
||||
return {"messages": [response]}
|
||||
```
|
||||
|
||||
2. Try it:
|
||||
|
||||
```python
|
||||
input = {
|
||||
"messages": convert_to_messages(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "1",
|
||||
"name": "retrieve_blog_posts",
|
||||
"args": {"query": "types of reward hacking"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"content": "reward hacking can be categorized into two types: environment or goal misspecification, and reward tampering",
|
||||
"tool_call_id": "1",
|
||||
},
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
response = generate_answer(input)
|
||||
response["messages"][-1].pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================== Ai Message ==================================
|
||||
|
||||
Lilian Weng categorizes reward hacking into two types: environment or goal misspecification, and reward tampering. She considers reward hacking as a broad concept that includes both of these categories. Reward hacking occurs when an agent exploits flaws or ambiguities in the reward function to achieve high rewards without performing the intended behaviors.
|
||||
```
|
||||
|
||||
## 7. Assemble the graph
|
||||
|
||||
* Start with a `generate_query_or_respond` and determine if we need to call `retriever_tool`
|
||||
* Route to next step using `tools_condition`:
|
||||
* If `generate_query_or_respond` returned `tool_calls`, call `retriever_tool` to retrieve context
|
||||
* Otherwise, respond directly to the user
|
||||
* Grade retrieved document content for relevance to the question (`grade_documents`) and route to next step:
|
||||
* If not relevant, rewrite the question using `rewrite_question` and then call `generate_query_or_respond` again
|
||||
* If relevant, proceed to `generate_answer` and generate final response using the `ToolMessage` with the retrieved document context
|
||||
|
||||
```python
|
||||
from langgraph.graph import StateGraph, START, END
|
||||
from langgraph.prebuilt import ToolNode
|
||||
from langgraph.prebuilt import tools_condition
|
||||
|
||||
workflow = StateGraph(MessagesState)
|
||||
|
||||
# Define the nodes we will cycle between
|
||||
workflow.add_node(generate_query_or_respond)
|
||||
workflow.add_node("retrieve", ToolNode([retriever_tool]))
|
||||
workflow.add_node(rewrite_question)
|
||||
workflow.add_node(generate_answer)
|
||||
|
||||
workflow.add_edge(START, "generate_query_or_respond")
|
||||
|
||||
# Decide whether to retrieve
|
||||
workflow.add_conditional_edges(
|
||||
"generate_query_or_respond",
|
||||
# Assess LLM decision (call `retriever_tool` tool or respond to the user)
|
||||
tools_condition,
|
||||
{
|
||||
# Translate the condition outputs to nodes in our graph
|
||||
"tools": "retrieve",
|
||||
END: END,
|
||||
},
|
||||
)
|
||||
|
||||
# Edges taken after the `action` node is called.
|
||||
workflow.add_conditional_edges(
|
||||
"retrieve",
|
||||
# Assess agent decision
|
||||
grade_documents,
|
||||
)
|
||||
workflow.add_edge("generate_answer", END)
|
||||
workflow.add_edge("rewrite_question", "generate_query_or_respond")
|
||||
|
||||
# Compile
|
||||
graph = workflow.compile()
|
||||
```
|
||||
|
||||
Visualize the graph:
|
||||
|
||||
```python
|
||||
from IPython.display import Image, display
|
||||
|
||||
display(Image(graph.get_graph().draw_mermaid_png()))
|
||||
```
|
||||
|
||||

|
||||
|
||||
## 8. Run the agentic RAG
|
||||
|
||||
```python
|
||||
for chunk in graph.stream(
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "What does Lilian Weng say about types of reward hacking?",
|
||||
}
|
||||
]
|
||||
}
|
||||
):
|
||||
for node, update in chunk.items():
|
||||
print("Update from node", node)
|
||||
update["messages"][-1].pretty_print()
|
||||
print("\n\n")
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Update from node generate_query_or_respond
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
retrieve_blog_posts (call_NYu2vq4km9nNNEFqJwefWKu1)
|
||||
Call ID: call_NYu2vq4km9nNNEFqJwefWKu1
|
||||
Args:
|
||||
query: types of reward hacking
|
||||
|
||||
|
||||
|
||||
Update from node retrieve
|
||||
================================= Tool Message ==================================
|
||||
Name: retrieve_blog_posts
|
||||
|
||||
(Note: Some work defines reward tampering as a distinct category of misalignment behavior from reward hacking. But I consider reward hacking as a broader concept here.)
|
||||
At a high level, reward hacking can be categorized into two types: environment or goal misspecification, and reward tampering.
|
||||
|
||||
Why does Reward Hacking Exist?#
|
||||
|
||||
Pan et al. (2022) investigated reward hacking as a function of agent capabilities, including (1) model size, (2) action space resolution, (3) observation space noise, and (4) training time. They also proposed a taxonomy of three types of misspecified proxy rewards:
|
||||
|
||||
Let's Define Reward Hacking#
|
||||
Reward shaping in RL is challenging. Reward hacking occurs when an RL agent exploits flaws or ambiguities in the reward function to obtain high rewards without genuinely learning the intended behaviors or completing the task as designed. In recent years, several related concepts have been proposed, all referring to some form of reward hacking:
|
||||
|
||||
|
||||
|
||||
Update from node generate_answer
|
||||
================================== Ai Message ==================================
|
||||
|
||||
Lilian Weng categorizes reward hacking into two types: environment or goal misspecification, and reward tampering. She considers reward hacking as a broad concept that includes both of these categories. Reward hacking occurs when an agent exploits flaws or ambiguities in the reward function to achieve high rewards without performing the intended behaviors.
|
||||
```
|
||||
File diff suppressed because one or more lines are too long
Binary file not shown.
|
After Width: | Height: | Size: 20 KiB |
@@ -0,0 +1,549 @@
|
||||
# Build a SQL agent
|
||||
|
||||
In this tutorial, we will walk through how to build an agent that can answer questions about a SQL database.
|
||||
|
||||
At a high level, the agent will:
|
||||
|
||||
1. Fetch the available tables from the database
|
||||
2. Decide which tables are relevant to the question
|
||||
3. Fetch the schemas for the relevant tables
|
||||
4. Generate a query based on the question and information from the schemas
|
||||
5. Double-check the query for common mistakes using an LLM
|
||||
6. Execute the query and return the results
|
||||
7. Correct mistakes surfaced by the database engine until the query is successful
|
||||
8. Formulate a response based on the results
|
||||
|
||||
!!! warning "Security note"
|
||||
Building Q&A systems of SQL databases requires executing model-generated SQL queries. There are inherent risks in doing this. Make sure that your database connection permissions are always scoped as narrowly as possible for your agent's needs. This will mitigate though not eliminate the risks of building a model-driven system.
|
||||
|
||||
## 1. Setup
|
||||
|
||||
Let's first install some dependencies. This tutorial uses SQL database and tool abstractions from [langchain-community](https://python.langchain.com/docs/concepts/architecture/#langchain-community). We will also require a LangChain [chat model](https://python.langchain.com/docs/concepts/chat_models/).
|
||||
|
||||
```python
|
||||
%%capture --no-stderr
|
||||
%pip install -U langgraph langchain_community "langchain[openai]"
|
||||
```
|
||||
|
||||
!!! tip
|
||||
Sign up for LangSmith to quickly spot issues and improve the performance of your LangGraph projects. [LangSmith](https://docs.smith.langchain.com) lets you use trace data to debug, test, and monitor your LLM apps built with LangGraph.
|
||||
|
||||
### Select a LLM
|
||||
|
||||
First we [initialize our LLM](https://python.langchain.com/docs/how_to/chat_models_universal_init/). Any model supporting [tool-calling](https://python.langchain.com/docs/integrations/chat/#featured-providers) should work. We use OpenAI below.
|
||||
|
||||
```python
|
||||
from langchain.chat_models import init_chat_model
|
||||
|
||||
llm = init_chat_model("openai:gpt-4.1")
|
||||
```
|
||||
|
||||
### Configure the database
|
||||
|
||||
We will be creating a SQLite database for this tutorial. SQLite is a lightweight database that is easy to set up and use. We will be loading the `chinook` database, which is a sample database that represents a digital media store.
|
||||
Find more information about the database [here](https://www.sqlitetutorial.net/sqlite-sample-database/).
|
||||
|
||||
For convenience, we have hosted the database (`Chinook.db`) on a public GCS bucket.
|
||||
|
||||
```python
|
||||
import requests
|
||||
|
||||
url = "https://storage.googleapis.com/benchmarks-artifacts/chinook/Chinook.db"
|
||||
|
||||
response = requests.get(url)
|
||||
|
||||
if response.status_code == 200:
|
||||
# Open a local file in binary write mode
|
||||
with open("Chinook.db", "wb") as file:
|
||||
# Write the content of the response (the file) to the local file
|
||||
file.write(response.content)
|
||||
print("File downloaded and saved as Chinook.db")
|
||||
else:
|
||||
print(f"Failed to download the file. Status code: {response.status_code}")
|
||||
```
|
||||
|
||||
We will use a handy SQL database wrapper available in the `langchain_community` package to interact with the database. The wrapper provides a simple interface to execute SQL queries and fetch results:
|
||||
|
||||
```python
|
||||
from langchain_community.utilities import SQLDatabase
|
||||
|
||||
db = SQLDatabase.from_uri("sqlite:///Chinook.db")
|
||||
|
||||
print(f"Dialect: {db.dialect}")
|
||||
print(f"Available tables: {db.get_usable_table_names()}")
|
||||
print(f'Sample output: {db.run("SELECT * FROM Artist LIMIT 5;")}')
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
Dialect: sqlite
|
||||
Available tables: ['Album', 'Artist', 'Customer', 'Employee', 'Genre', 'Invoice', 'InvoiceLine', 'MediaType', 'Playlist', 'PlaylistTrack', 'Track']
|
||||
Sample output: [(1, 'AC/DC'), (2, 'Accept'), (3, 'Aerosmith'), (4, 'Alanis Morissette'), (5, 'Alice In Chains')]
|
||||
```
|
||||
|
||||
### Tools for database interactions
|
||||
|
||||
`langchain-community` implements some built-in tools for interacting with our `SQLDatabase`, including tools for listing tables, reading table schemas, and checking and running queries:
|
||||
|
||||
```python
|
||||
from langchain_community.agent_toolkits import SQLDatabaseToolkit
|
||||
|
||||
toolkit = SQLDatabaseToolkit(db=db, llm=llm)
|
||||
|
||||
tools = toolkit.get_tools()
|
||||
|
||||
for tool in tools:
|
||||
print(f"{tool.name}: {tool.description}\n")
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
sql_db_query: Input to this tool is a detailed and correct SQL query, output is a result from the database. If the query is not correct, an error message will be returned. If an error is returned, rewrite the query, check the query, and try again. If you encounter an issue with Unknown column 'xxxx' in 'field list', use sql_db_schema to query the correct table fields.
|
||||
|
||||
sql_db_schema: Input to this tool is a comma-separated list of tables, output is the schema and sample rows for those tables. Be sure that the tables actually exist by calling sql_db_list_tables first! Example Input: table1, table2, table3
|
||||
|
||||
sql_db_list_tables: Input is an empty string, output is a comma-separated list of tables in the database.
|
||||
|
||||
sql_db_query_checker: Use this tool to double check if your query is correct before executing it. Always use this tool before executing a query with sql_db_query!
|
||||
|
||||
```
|
||||
|
||||
## 2. Using a prebuilt agent
|
||||
|
||||
Given these tools, we can initialize a pre-built agent in a single line. To customize our agents behavior, we write a descriptive system prompt.
|
||||
|
||||
```python
|
||||
from langgraph.prebuilt import create_react_agent
|
||||
|
||||
system_prompt = """
|
||||
You are an agent designed to interact with a SQL database.
|
||||
Given an input question, create a syntactically correct {dialect} query to run,
|
||||
then look at the results of the query and return the answer. Unless the user
|
||||
specifies a specific number of examples they wish to obtain, always limit your
|
||||
query to at most {top_k} results.
|
||||
|
||||
You can order the results by a relevant column to return the most interesting
|
||||
examples in the database. Never query for all the columns from a specific table,
|
||||
only ask for the relevant columns given the question.
|
||||
|
||||
You MUST double check your query before executing it. If you get an error while
|
||||
executing a query, rewrite the query and try again.
|
||||
|
||||
DO NOT make any DML statements (INSERT, UPDATE, DELETE, DROP etc.) to the
|
||||
database.
|
||||
|
||||
To start you should ALWAYS look at the tables in the database to see what you
|
||||
can query. Do NOT skip this step.
|
||||
|
||||
Then you should query the schema of the most relevant tables.
|
||||
""".format(
|
||||
dialect=db.dialect,
|
||||
top_k=5,
|
||||
)
|
||||
|
||||
agent = create_react_agent(
|
||||
llm,
|
||||
tools,
|
||||
prompt=system_prompt,
|
||||
)
|
||||
```
|
||||
|
||||
!!! note
|
||||
This system prompt includes a number of instructions, such as always running specific tools before or after others. In the [next section](#3-customizing-the-agent), we will enforce these behaviors through the graph's structure, providing us a greater degree of control and allowing us to simplify the prompt.
|
||||
|
||||
Let's run this agent on a sample query and observe its behavior:
|
||||
|
||||
```python
|
||||
question = "Which genre on average has the longest tracks?"
|
||||
|
||||
for step in agent.stream(
|
||||
{"messages": [{"role": "user", "content": question}]},
|
||||
stream_mode="values",
|
||||
):
|
||||
step["messages"][-1].pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================ Human Message =================================
|
||||
|
||||
Which genre on average has the longest tracks?
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_list_tables (call_d8lCgywSroCgpVl558nmXKwA)
|
||||
Call ID: call_d8lCgywSroCgpVl558nmXKwA
|
||||
Args:
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_list_tables
|
||||
|
||||
Album, Artist, Customer, Employee, Genre, Invoice, InvoiceLine, MediaType, Playlist, PlaylistTrack, Track
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_schema (call_nNf6IIUcwMYLIkE0l6uWkZHe)
|
||||
Call ID: call_nNf6IIUcwMYLIkE0l6uWkZHe
|
||||
Args:
|
||||
table_names: Genre, Track
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_schema
|
||||
|
||||
|
||||
CREATE TABLE "Genre" (
|
||||
"GenreId" INTEGER NOT NULL,
|
||||
"Name" NVARCHAR(120),
|
||||
PRIMARY KEY ("GenreId")
|
||||
)
|
||||
|
||||
/*
|
||||
3 rows from Genre table:
|
||||
GenreId Name
|
||||
1 Rock
|
||||
2 Jazz
|
||||
3 Metal
|
||||
*/
|
||||
|
||||
|
||||
CREATE TABLE "Track" (
|
||||
"TrackId" INTEGER NOT NULL,
|
||||
"Name" NVARCHAR(200) NOT NULL,
|
||||
"AlbumId" INTEGER,
|
||||
"MediaTypeId" INTEGER NOT NULL,
|
||||
"GenreId" INTEGER,
|
||||
"Composer" NVARCHAR(220),
|
||||
"Milliseconds" INTEGER NOT NULL,
|
||||
"Bytes" INTEGER,
|
||||
"UnitPrice" NUMERIC(10, 2) NOT NULL,
|
||||
PRIMARY KEY ("TrackId"),
|
||||
FOREIGN KEY("MediaTypeId") REFERENCES "MediaType" ("MediaTypeId"),
|
||||
FOREIGN KEY("GenreId") REFERENCES "Genre" ("GenreId"),
|
||||
FOREIGN KEY("AlbumId") REFERENCES "Album" ("AlbumId")
|
||||
)
|
||||
|
||||
/*
|
||||
3 rows from Track table:
|
||||
TrackId Name AlbumId MediaTypeId GenreId Composer Milliseconds Bytes UnitPrice
|
||||
1 For Those About To Rock (We Salute You) 1 1 1 Angus Young, Malcolm Young, Brian Johnson 343719 11170334 0.99
|
||||
2 Balls to the Wall 2 2 1 None 342562 5510424 0.99
|
||||
3 Fast As a Shark 3 2 1 F. Baltes, S. Kaufman, U. Dirkscneider & W. Hoffman 230619 3990994 0.99
|
||||
*/
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_query_checker (call_urTRmtiGtTxkwHtscec7Fd2K)
|
||||
Call ID: call_urTRmtiGtTxkwHtscec7Fd2K
|
||||
Args:
|
||||
query: SELECT Genre.Name, AVG(Track.Milliseconds) AS AvgMilliseconds
|
||||
FROM Track
|
||||
JOIN Genre ON Track.GenreId = Genre.GenreId
|
||||
GROUP BY Genre.Name
|
||||
ORDER BY AvgMilliseconds DESC
|
||||
LIMIT 1;
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_query_checker
|
||||
|
||||
\`\`\`sql
|
||||
SELECT Genre.Name, AVG(Track.Milliseconds) AS AvgMilliseconds
|
||||
FROM Track
|
||||
JOIN Genre ON Track.GenreId = Genre.GenreId
|
||||
GROUP BY Genre.Name
|
||||
ORDER BY AvgMilliseconds DESC
|
||||
LIMIT 1;
|
||||
\`\`\`
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_query (call_RNMqyUEMv0rvy0UxSwrXY2AV)
|
||||
Call ID: call_RNMqyUEMv0rvy0UxSwrXY2AV
|
||||
Args:
|
||||
query: SELECT Genre.Name, AVG(Track.Milliseconds) AS AvgMilliseconds
|
||||
FROM Track
|
||||
JOIN Genre ON Track.GenreId = Genre.GenreId
|
||||
GROUP BY Genre.Name
|
||||
ORDER BY AvgMilliseconds DESC
|
||||
LIMIT 1;
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_query
|
||||
|
||||
[('Sci Fi & Fantasy', 2911783.0384615385)]
|
||||
================================== Ai Message ==================================
|
||||
|
||||
The genre with the longest average track length is "Sci Fi & Fantasy," with an average duration of about 2,911,783 milliseconds (approximately 48.5 minutes) per track.
|
||||
```
|
||||
|
||||
This worked well enough: the agent correctly listed the tables, obtained the schemas, wrote a query, checked the query, and ran it to inform its final response.
|
||||
|
||||
!!! tip
|
||||
You can inspect all aspects of the above run, including steps taken, tools invoked, what prompts were seen by the LLM, and more in the [LangSmith trace](https://smith.langchain.com/public/bd594960-73e3-474b-b6f2-db039d7c713a/r).
|
||||
|
||||
## 3. Customizing the agent
|
||||
|
||||
The prebuilt agent lets us get started quickly, but at each step the agent has access to the full set of tools. Above, we relied on the system prompt to constrain its behavior— for example, we instructed the agent to always start with the "list tables" tool, and to always run a query-checker tool before executing the query.
|
||||
|
||||
We can enforce a higher degree of control in LangGraph by customizing the agent. Below, we implement a simple ReAct-agent setup, with dedicated nodes for specific tool-calls. We will use the same [state](../../concepts/low_level.md#state) as the pre-built agent.
|
||||
|
||||
We construct dedicated nodes for the following steps:
|
||||
|
||||
- Listing DB tables
|
||||
- Calling the "get schema" tool
|
||||
- Generating a query
|
||||
- Checking the query
|
||||
|
||||
Putting these steps in dedicated nodes lets us (1) force tool-calls when needed, and (2) customize the prompts associated with each step.
|
||||
|
||||
```python
|
||||
from typing import Literal
|
||||
from langchain_core.messages import AIMessage
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.graph import END, START, MessagesState, StateGraph
|
||||
from langgraph.prebuilt import ToolNode
|
||||
|
||||
|
||||
get_schema_tool = next(tool for tool in tools if tool.name == "sql_db_schema")
|
||||
get_schema_node = ToolNode([get_schema_tool], name="get_schema")
|
||||
|
||||
run_query_tool = next(tool for tool in tools if tool.name == "sql_db_query")
|
||||
run_query_node = ToolNode([run_query_tool], name="run_query")
|
||||
|
||||
|
||||
# Example: create a predetermined tool call
|
||||
def list_tables(state: MessagesState):
|
||||
tool_call = {
|
||||
"name": "sql_db_list_tables",
|
||||
"args": {},
|
||||
"id": "abc123",
|
||||
"type": "tool_call",
|
||||
}
|
||||
tool_call_message = AIMessage(content="", tool_calls=[tool_call])
|
||||
|
||||
list_tables_tool = next(tool for tool in tools if tool.name == "sql_db_list_tables")
|
||||
tool_message = list_tables_tool.invoke(tool_call)
|
||||
response = AIMessage(f"Available tables: {tool_message.content}")
|
||||
|
||||
return {"messages": [tool_call_message, tool_message, response]}
|
||||
|
||||
|
||||
# Example: force a model to create a tool call
|
||||
def call_get_schema(state: MessagesState):
|
||||
# Note that LangChain enforces that all models accept `tool_choice="any"`
|
||||
# as well as `tool_choice=<string name of tool>`.
|
||||
llm_with_tools = llm.bind_tools([get_schema_tool], tool_choice="any")
|
||||
response = llm_with_tools.invoke(state["messages"])
|
||||
|
||||
return {"messages": [response]}
|
||||
|
||||
|
||||
generate_query_system_prompt = """
|
||||
You are an agent designed to interact with a SQL database.
|
||||
Given an input question, create a syntactically correct {dialect} query to run,
|
||||
then look at the results of the query and return the answer. Unless the user
|
||||
specifies a specific number of examples they wish to obtain, always limit your
|
||||
query to at most {top_k} results.
|
||||
|
||||
You can order the results by a relevant column to return the most interesting
|
||||
examples in the database. Never query for all the columns from a specific table,
|
||||
only ask for the relevant columns given the question.
|
||||
|
||||
DO NOT make any DML statements (INSERT, UPDATE, DELETE, DROP etc.) to the database.
|
||||
""".format(
|
||||
dialect=db.dialect,
|
||||
top_k=5,
|
||||
)
|
||||
|
||||
|
||||
def generate_query(state: MessagesState):
|
||||
system_message = {
|
||||
"role": "system",
|
||||
"content": generate_query_system_prompt,
|
||||
}
|
||||
# We do not force a tool call here, to allow the model to
|
||||
# respond naturally when it obtains the solution.
|
||||
llm_with_tools = llm.bind_tools([run_query_tool])
|
||||
response = llm_with_tools.invoke([system_message] + state["messages"])
|
||||
|
||||
return {"messages": [response]}
|
||||
|
||||
|
||||
check_query_system_prompt = """
|
||||
You are a SQL expert with a strong attention to detail.
|
||||
Double check the {dialect} query for common mistakes, including:
|
||||
- Using NOT IN with NULL values
|
||||
- Using UNION when UNION ALL should have been used
|
||||
- Using BETWEEN for exclusive ranges
|
||||
- Data type mismatch in predicates
|
||||
- Properly quoting identifiers
|
||||
- Using the correct number of arguments for functions
|
||||
- Casting to the correct data type
|
||||
- Using the proper columns for joins
|
||||
|
||||
If there are any of the above mistakes, rewrite the query. If there are no mistakes,
|
||||
just reproduce the original query.
|
||||
|
||||
You will call the appropriate tool to execute the query after running this check.
|
||||
""".format(dialect=db.dialect)
|
||||
|
||||
|
||||
def check_query(state: MessagesState):
|
||||
system_message = {
|
||||
"role": "system",
|
||||
"content": check_query_system_prompt,
|
||||
}
|
||||
|
||||
# Generate an artificial user message to check
|
||||
tool_call = state["messages"][-1].tool_calls[0]
|
||||
user_message = {"role": "user", "content": tool_call["args"]["query"]}
|
||||
llm_with_tools = llm.bind_tools([run_query_tool], tool_choice="any")
|
||||
response = llm_with_tools.invoke([system_message, user_message])
|
||||
response.id = state["messages"][-1].id
|
||||
|
||||
return {"messages": [response]}
|
||||
```
|
||||
|
||||
Finally, we assemble these steps into a workflow using the Graph API. We define a [conditional edge](../../concepts/low_level.md#conditional-edges) at the query generation step that will route to the query checker if a query is generated, or end if there are no tool calls present, such that the LLM has delivered a response to the query.
|
||||
|
||||
```python
|
||||
def should_continue(state: MessagesState) -> Literal[END, "check_query"]:
|
||||
messages = state["messages"]
|
||||
last_message = messages[-1]
|
||||
if not last_message.tool_calls:
|
||||
return END
|
||||
else:
|
||||
return "check_query"
|
||||
|
||||
|
||||
builder = StateGraph(MessagesState)
|
||||
builder.add_node(list_tables)
|
||||
builder.add_node(call_get_schema)
|
||||
builder.add_node(get_schema_node, "get_schema")
|
||||
builder.add_node(generate_query)
|
||||
builder.add_node(check_query)
|
||||
builder.add_node(run_query_node, "run_query")
|
||||
|
||||
builder.add_edge(START, "list_tables")
|
||||
builder.add_edge("list_tables", "call_get_schema")
|
||||
builder.add_edge("call_get_schema", "get_schema")
|
||||
builder.add_edge("get_schema", "generate_query")
|
||||
builder.add_conditional_edges(
|
||||
"generate_query",
|
||||
should_continue,
|
||||
)
|
||||
builder.add_edge("check_query", "run_query")
|
||||
builder.add_edge("run_query", "generate_query")
|
||||
|
||||
agent = builder.compile()
|
||||
```
|
||||
|
||||
We visualize the application below:
|
||||
|
||||
```python
|
||||
from IPython.display import Image, display
|
||||
from langchain_core.runnables.graph import CurveStyle, MermaidDrawMethod, NodeStyles
|
||||
|
||||
display(Image(agent.get_graph().draw_mermaid_png()))
|
||||
```
|
||||
|
||||

|
||||
|
||||
**Note:** When you run this code, it will generate and display a visual representation of the SQL agent graph showing the flow between the different nodes (list_tables → call_get_schema → get_schema → generate_query → check_query → run_query).
|
||||
|
||||
We can now invoke the graph exactly as before:
|
||||
|
||||
```python
|
||||
question = "Which genre on average has the longest tracks?"
|
||||
|
||||
for step in agent.stream(
|
||||
{"messages": [{"role": "user", "content": question}]},
|
||||
stream_mode="values",
|
||||
):
|
||||
step["messages"][-1].pretty_print()
|
||||
```
|
||||
|
||||
**Output:**
|
||||
```
|
||||
================================ Human Message =================================
|
||||
|
||||
Which genre on average has the longest tracks?
|
||||
================================== Ai Message ==================================
|
||||
|
||||
Available tables: Album, Artist, Customer, Employee, Genre, Invoice, InvoiceLine, MediaType, Playlist, PlaylistTrack, Track
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_schema (call_qxKtYiHgf93AiTDin9ez5wFp)
|
||||
Call ID: call_qxKtYiHgf93AiTDin9ez5wFp
|
||||
Args:
|
||||
table_names: Genre,Track
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_schema
|
||||
|
||||
|
||||
CREATE TABLE "Genre" (
|
||||
"GenreId" INTEGER NOT NULL,
|
||||
"Name" NVARCHAR(120),
|
||||
PRIMARY KEY ("GenreId")
|
||||
)
|
||||
|
||||
/*
|
||||
3 rows from Genre table:
|
||||
GenreId Name
|
||||
1 Rock
|
||||
2 Jazz
|
||||
3 Metal
|
||||
*/
|
||||
|
||||
|
||||
CREATE TABLE "Track" (
|
||||
"TrackId" INTEGER NOT NULL,
|
||||
"Name" NVARCHAR(200) NOT NULL,
|
||||
"AlbumId" INTEGER,
|
||||
"MediaTypeId" INTEGER NOT NULL,
|
||||
"GenreId" INTEGER,
|
||||
"Composer" NVARCHAR(220),
|
||||
"Milliseconds" INTEGER NOT NULL,
|
||||
"Bytes" INTEGER,
|
||||
"UnitPrice" NUMERIC(10, 2) NOT NULL,
|
||||
PRIMARY KEY ("TrackId"),
|
||||
FOREIGN KEY("MediaTypeId") REFERENCES "MediaType" ("MediaTypeId"),
|
||||
FOREIGN KEY("GenreId") REFERENCES "Genre" ("GenreId"),
|
||||
FOREIGN KEY("AlbumId") REFERENCES "Album" ("AlbumId")
|
||||
)
|
||||
|
||||
/*
|
||||
3 rows from Track table:
|
||||
TrackId Name AlbumId MediaTypeId GenreId Composer Milliseconds Bytes UnitPrice
|
||||
1 For Those About To Rock (We Salute You) 1 1 1 Angus Young, Malcolm Young, Brian Johnson 343719 11170334 0.99
|
||||
2 Balls to the Wall 2 2 1 None 342562 5510424 0.99
|
||||
3 Fast As a Shark 3 2 1 F. Baltes, S. Kaufman, U. Dirkscneider & W. Hoffman 230619 3990994 0.99
|
||||
*/
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_query (call_RPN3GABMfb6DTaFTLlwnZxVN)
|
||||
Call ID: call_RPN3GABMfb6DTaFTLlwnZxVN
|
||||
Args:
|
||||
query: SELECT Genre.Name, AVG(Track.Milliseconds) AS AvgTrackLength
|
||||
FROM Track
|
||||
JOIN Genre ON Track.GenreId = Genre.GenreId
|
||||
GROUP BY Genre.GenreId
|
||||
ORDER BY AvgTrackLength DESC
|
||||
LIMIT 1;
|
||||
================================== Ai Message ==================================
|
||||
Tool Calls:
|
||||
sql_db_query (call_PR4s8ymiF3ZQLaoZADXtdqcl)
|
||||
Call ID: call_PR4s8ymiF3ZQLaoZADXtdqcl
|
||||
Args:
|
||||
query: SELECT Genre.Name, AVG(Track.Milliseconds) AS AvgTrackLength
|
||||
FROM Track
|
||||
JOIN Genre ON Track.GenreId = Genre.GenreId
|
||||
GROUP BY Genre.GenreId
|
||||
ORDER BY AvgTrackLength DESC
|
||||
LIMIT 1;
|
||||
================================= Tool Message =================================
|
||||
Name: sql_db_query
|
||||
|
||||
[('Sci Fi & Fantasy', 2911783.0384615385)]
|
||||
================================== Ai Message ==================================
|
||||
|
||||
The genre with the longest tracks on average is "Sci Fi & Fantasy," with an average track length of approximately 2,911,783 milliseconds.
|
||||
```
|
||||
|
||||
!!! tip
|
||||
See [LangSmith trace](https://smith.langchain.com/public/94b8c9ac-12f7-4692-8706-836a1f30f1ea/r) for the above run.
|
||||
|
||||
## Next steps
|
||||
|
||||
Check out [this guide](https://docs.smith.langchain.com/evaluation/how_to_guides/langgraph) for evaluating LangGraph applications, including SQL agents like this one, using LangSmith.
|
||||
Reference in New Issue
Block a user