mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-22 07:32:25 +02:00
24 KiB
24 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)
for id, tool in tool_registry.items()
]
vector_store = InMemoryVectorStore(embedding=OpenAIEmbeddings())
upsert_response = vector_store.upsert(tool_documents)
assert not upsert_response["failed"]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"])['3a42d402-4f6f-44ce-9ff4-d53d50f5f24a', '42aafb13-52e4-402b-80aa-db5be2e0e40f', '275bcddc-a918-4688-938f-9d2d6f232757', 'abc00b37-e0fa-470f-98d7-f055cc29a11f']
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_oXCnqxGmzQaP4OTEWvPghiKR) Call ID: call_oXCnqxGmzQaP4OTEWvPghiKR 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.