mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
43 KiB
43 KiB
In [1]:
import re
import uuid
from langchain_core.tools import StructuredTool
def create_tool(company: str) -> dict:
"""Create schema for a placeholder tool."""
formatted_company = re.sub(r"[^\w\s]", "", company).replace(" ", "_")
def company_tool(year: int) -> str:
return f"{company} had revenues of $100 in {year}."
return StructuredTool.from_function(
company_tool,
name=formatted_company,
description=f"Information about {company}",
)
s_and_p_500_companies = [ # Abbreviated list for demonstration purposes
"3M",
"A.O. Smith",
"Abbott",
"Accenture",
"Advanced Micro Devices",
"Yum! Brands",
"Zebra Technologies",
"Zimmer Biomet",
"Zoetis",
]
tool_registry = {
str(uuid.uuid4()): create_tool(company) for company in s_and_p_500_companies
}In [2]:
from langchain_core.documents import Document
from langchain_core.vectorstores import InMemoryVectorStore, VectorStore
from langchain_openai import OpenAIEmbeddings
tool_documents = [
Document(
page_content=tool.description,
id=id,
metadata={"tool_name": tool.name},
)
for id, tool in tool_registry.items()
]
vector_store = InMemoryVectorStore(embedding=OpenAIEmbeddings())
document_ids = vector_store.add_documents(tool_documents)In [3]:
from typing import Annotated
from langchain_openai import ChatOpenAI
from typing_extensions import TypedDict
from langgraph.graph import StateGraph, START, END
from langgraph.graph.message import add_messages
from langgraph.prebuilt import ToolNode, tools_condition
class State(TypedDict):
messages: Annotated[list, add_messages]
selected_tools: list[str]
graph_builder = StateGraph(State)
tools = list(tool_registry.values())
llm = ChatOpenAI()
def agent(state: State):
selected_tools = [tool_registry[id] for id in state["selected_tools"]]
llm_with_tools = llm.bind_tools(selected_tools)
return {"messages": [llm_with_tools.invoke(state["messages"])]}
def select_tools(state: State):
last_user_message = state["messages"][-1]
query = last_user_message.content
tool_documents = vector_store.similarity_search(query)
return {"selected_tools": [document.id for document in tool_documents]}
graph_builder.add_node("agent", agent)
graph_builder.add_node("select_tools", select_tools)
tool_node = ToolNode(tools=tools)
graph_builder.add_node("tools", tool_node)
graph_builder.add_conditional_edges(
"agent",
tools_condition,
)
graph_builder.add_edge("tools", "agent")
graph_builder.add_edge("select_tools", "agent")
graph_builder.add_edge(START, "select_tools")
graph = graph_builder.compile()In [4]:
from IPython.display import Image, display
try:
display(Image(graph.get_graph().draw_mermaid_png()))
except Exception:
# This requires some extra dependencies and is optional
passIn [5]:
user_input = "Can you give me some information about AMD in 2022?"
result = graph.invoke({"messages": [("user", user_input)]})In [6]:
print(result["selected_tools"])['3b7d1528-6007-4473-a92f-b9b3341c3bfe', '8d77b753-c58a-41bf-9649-ad1a7326bc27', '514a6fc3-03d1-4e73-b410-c39309ad7b2f', '83c5cc8f-5111-46ed-874a-e0b883265ff6']
In [7]:
for message in result["messages"]:
message.pretty_print()================================[1m Human Message [0m================================= Can you give me some information about AMD in 2022? ==================================[1m Ai Message [0m================================== Tool Calls: Advanced_Micro_Devices (call_Htbv7Imx4BwSsYWhZvSSs6yW) Call ID: call_Htbv7Imx4BwSsYWhZvSSs6yW Args: year: 2022 =================================[1m Tool Message [0m================================= Name: Advanced_Micro_Devices Advanced Micro Devices had revenues of $100 in 2022. ==================================[1m Ai Message [0m================================== In 2022, Advanced Micro Devices had revenues of $100.
In [8]:
from langchain_core.messages import HumanMessage, SystemMessage, ToolMessage
from langchain_core.pydantic_v1 import BaseModel, Field
class QueryForTools(BaseModel):
"""Generate a query for additional tools."""
query: str = Field(..., description="Query for additional tools.")
def select_tools(state: State):
last_message = state["messages"][-1]
hack_remove_tool_condition = False
if isinstance(last_message, HumanMessage):
query = last_message.content
hack_remove_tool_condition = True
else:
assert isinstance(last_message, ToolMessage)
system = SystemMessage(
"Given this conversation, generate a query for additional tools. "
"The query should be a short string containing what type of information "
"is needed. If no further information is needed, "
"set more_information_needed False and populate a blank string for the query."
)
input_messages = [system] + state["messages"]
response = llm.bind_tools([QueryForTools], tool_choice=True).invoke(
input_messages
)
query = response.tool_calls[0]["args"]["query"]
tool_documents = vector_store.similarity_search(query)
if hack_remove_tool_condition:
# Remove needed tool
selected_tools = [
document.id
for document in tool_documents
if document.metadata["tool_name"] != "Advanced_Micro_Devices"
]
else:
selected_tools = [document.id for document in tool_documents]
return {"selected_tools": selected_tools}
graph_builder = StateGraph(State)
graph_builder.add_node("agent", agent)
graph_builder.add_node("select_tools", select_tools)
tool_node = ToolNode(tools=tools)
graph_builder.add_node("tools", tool_node)
graph_builder.add_conditional_edges(
"agent",
tools_condition,
)
graph_builder.add_edge("tools", "select_tools")
graph_builder.add_edge("select_tools", "agent")
graph_builder.add_edge(START, "select_tools")
graph = graph_builder.compile()In [9]:
from IPython.display import Image, display
try:
display(Image(graph.get_graph().draw_mermaid_png()))
except Exception:
# This requires some extra dependencies and is optional
passIn [18]:
user_input = "Can you give me some information about AMD in 2022?"
result = graph.invoke({"messages": [("user", user_input)]})In [19]:
for message in result["messages"]:
message.pretty_print()================================[1m Human Message [0m================================= Can you give me some information about AMD in 2022? ==================================[1m Ai Message [0m================================== Tool Calls: Accenture (call_L82JRUyIFilhzeTmPnNbPeVD) Call ID: call_L82JRUyIFilhzeTmPnNbPeVD Args: year: 2022 =================================[1m Tool Message [0m================================= Name: Accenture Accenture had revenues of $100 in 2022. ==================================[1m Ai Message [0m================================== Tool Calls: Advanced_Micro_Devices (call_k3zR9zS98gjiejmNgq6aVsXL) Call ID: call_k3zR9zS98gjiejmNgq6aVsXL Args: year: 2022 =================================[1m Tool Message [0m================================= Name: Advanced_Micro_Devices Advanced Micro Devices had revenues of $100 in 2022. ==================================[1m Ai Message [0m================================== In 2022, Advanced Micro Devices (AMD) had revenues of $100.