Use START

This commit is contained in:
William Fu-Hinthorn
2024-08-16 15:43:24 -07:00
parent fc95028738
commit b5429b6342
19 changed files with 114 additions and 434 deletions
@@ -235,7 +235,7 @@
" # Call the chat bot\n",
" chat_bot_response = my_chat_bot(messages)\n",
" # Respond with an AI Message\n",
" return {\"messages\":[AIMessage(content=chat_bot_response[\"content\"])]}"
" return {\"messages\": [AIMessage(content=chat_bot_response[\"content\"])]}"
]
},
{
@@ -270,7 +270,7 @@
" # Call the simulated user\n",
" response = simulated_user.invoke({\"messages\": new_messages})\n",
" # This response is an AI message - we need to flip this to be a human message\n",
" return {\"messages\":[HumanMessage(content=response.content)]}"
" return {\"messages\": [HumanMessage(content=response.content)]}"
]
},
{
@@ -331,6 +331,7 @@
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
"\n",
"graph_builder = StateGraph(State)\n",
"graph_builder.add_node(\"user\", simulated_user_node)\n",
"graph_builder.add_node(\"chat_bot\", chat_bot_node)\n",
@@ -182,9 +182,11 @@
"from typing import Annotated\n",
"from typing_extensions import TypedDict\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
"\n",
"memory = MemorySaver()\n",
"workflow = StateGraph(State)\n",
"workflow.add_node(\"info\", chain)\n",
@@ -141,15 +141,18 @@
" print(\"----\")\n",
" return \"Sunny!\"\n",
"\n",
"model = ChatAnthropic(model_name=\"claude-3-5-sonnet-20240620\").bind_tools([weather_search])\n",
"\n",
"model = ChatAnthropic(model_name=\"claude-3-5-sonnet-20240620\").bind_tools(\n",
" [weather_search]\n",
")\n",
"\n",
"\n",
"class State(MessagesState):\n",
" \"\"\"Simple state.\"\"\"\n",
"\n",
"\n",
"def call_llm(state):\n",
" return {\n",
" \"messages\": [model.invoke(state['messages'])]\n",
" }\n",
" return {\"messages\": [model.invoke(state[\"messages\"])]}\n",
"\n",
"\n",
"def human_review_node(state):\n",
@@ -159,28 +162,30 @@
"def run_tool(state):\n",
" new_messages = []\n",
" tools = {\"weather_search\": weather_search}\n",
" tool_calls = state['messages'][-1].tool_calls\n",
" tool_calls = state[\"messages\"][-1].tool_calls\n",
" for tool_call in tool_calls:\n",
" tool = tools[tool_call['name']]\n",
" result = tool.invoke(tool_call['args'])\n",
" new_messages.append({\n",
" \"role\": \"tool\",\n",
" \"name\": tool_call['name'],\n",
" \"content\": result,\n",
" \"tool_call_id\": tool_call['id']\n",
" })\n",
" tool = tools[tool_call[\"name\"]]\n",
" result = tool.invoke(tool_call[\"args\"])\n",
" new_messages.append(\n",
" {\n",
" \"role\": \"tool\",\n",
" \"name\": tool_call[\"name\"],\n",
" \"content\": result,\n",
" \"tool_call_id\": tool_call[\"id\"],\n",
" }\n",
" )\n",
" return {\"messages\": new_messages}\n",
"\n",
"\n",
"def route_after_llm(state) -> Literal[END, \"human_review_node\"]:\n",
" if len(state['messages'][-1].tool_calls) == 0:\n",
" if len(state[\"messages\"][-1].tool_calls) == 0:\n",
" return END\n",
" else:\n",
" return \"human_review_node\"\n",
"\n",
"\n",
"def route_after_human(state) -> Literal[\"run_tool\", \"call_llm\"]:\n",
" if isinstance(state['messages'][-1], AIMessage):\n",
" if isinstance(state[\"messages\"][-1], AIMessage):\n",
" return \"run_tool\"\n",
" else:\n",
" return \"call_llm\"\n",
@@ -460,35 +465,35 @@
"print(\"Current State:\")\n",
"print(state.values)\n",
"print(\"\\nCurrent Tool Call ID:\")\n",
"current_content = state.values['messages'][-1].content\n",
"current_id = state.values['messages'][-1].id\n",
"tool_call_id = state.values['messages'][-1].tool_calls[0]['id']\n",
"current_content = state.values[\"messages\"][-1].content\n",
"current_id = state.values[\"messages\"][-1].id\n",
"tool_call_id = state.values[\"messages\"][-1].tool_calls[0][\"id\"]\n",
"print(tool_call_id)\n",
"\n",
"# We now need to construct a replacement tool call.\n",
"# We will change the argument to be `San Francisco, USA`\n",
"# Note that we could change any number of arguments or tool names - it just has to be a valid one\n",
"new_message = {\n",
" \"role\": \"assistant\", \n",
" \"role\": \"assistant\",\n",
" \"content\": current_content,\n",
" \"tool_calls\": [\n",
" {\n",
" \"id\": tool_call_id,\n",
" \"name\": \"weather_search\",\n",
" \"args\": {\"city\": \"San Francisco, USA\"}\n",
" \"args\": {\"city\": \"San Francisco, USA\"},\n",
" }\n",
" ],\n",
" # This is important - this needs to be the same as the message you replacing!\n",
" # Otherwise, it will show up as a separate message\n",
" \"id\": current_id\n",
" \"id\": current_id,\n",
"}\n",
"graph.update_state(\n",
" # This is the config which represents this thread\n",
" thread, \n",
" thread,\n",
" # This is the updated value we want to push\n",
" {\"messages\": [new_message]}, \n",
" {\"messages\": [new_message]},\n",
" # We push this update acting as our human_review_node\n",
" as_node=\"human_review_node\"\n",
" as_node=\"human_review_node\",\n",
")\n",
"\n",
"# Let's now continue executing from here\n",
@@ -595,26 +600,26 @@
"print(\"Current State:\")\n",
"print(state.values)\n",
"print(\"\\nCurrent Tool Call ID:\")\n",
"tool_call_id = state.values['messages'][-1].tool_calls[0]['id']\n",
"tool_call_id = state.values[\"messages\"][-1].tool_calls[0][\"id\"]\n",
"print(tool_call_id)\n",
"\n",
"# We now need to construct a replacement tool call.\n",
"# We will change the argument to be `San Francisco, USA`\n",
"# Note that we could change any number of arguments or tool names - it just has to be a valid one\n",
"new_message = {\n",
" \"role\": \"tool\", \n",
" \"role\": \"tool\",\n",
" # This is our natural language feedback\n",
" \"content\": \"User requested changes: pass in the country as well\",\n",
" \"name\": \"weather_search\",\n",
" \"tool_call_id\": tool_call_id\n",
" \"tool_call_id\": tool_call_id,\n",
"}\n",
"graph.update_state(\n",
" # This is the config which represents this thread\n",
" thread, \n",
" thread,\n",
" # This is the updated value we want to push\n",
" {\"messages\": [new_message]}, \n",
" {\"messages\": [new_message]},\n",
" # We push this update acting as our human_review_node\n",
" as_node=\"human_review_node\"\n",
" as_node=\"human_review_node\",\n",
")\n",
"\n",
"# Let's now continue executing from here\n",
+4
View File
@@ -33,15 +33,19 @@
"from langgraph.graph import StateGraph, START, END\n",
"from typing import TypedDict\n",
"\n",
"\n",
"class InputState(TypedDict):\n",
" question: str\n",
"\n",
"\n",
"class OutputState(TypedDict):\n",
" answer: str\n",
"\n",
"\n",
"def answer_node(state: InputState):\n",
" return {\"answer\": \"bye\"}\n",
"\n",
"\n",
"graph = StateGraph(input=InputState, output=OutputState)\n",
"graph.add_node(answer_node)\n",
"graph.add_edge(START, \"answer_node\")\n",
+7 -3
View File
@@ -526,7 +526,7 @@
" \"tasks\": tasks,\n",
" }\n",
" )\n",
" return {\"messages\":[scheduled_tasks]}"
" return {\"messages\": [scheduled_tasks]}"
]
},
{
@@ -653,7 +653,7 @@
" )\n",
" ]\n",
" else:\n",
" return {\"messages\":response + [AIMessage(content=decision.action.response)]}\n",
" return {\"messages\": response + [AIMessage(content=decision.action.response)]}\n",
"\n",
"\n",
"def select_recent_messages(state) -> dict:\n",
@@ -726,9 +726,11 @@
"from langgraph.graph.message import add_messages\n",
"from typing import Annotated\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
"\n",
"graph_builder = StateGraph(State)\n",
"\n",
"# 1. Define vertices\n",
@@ -794,7 +796,9 @@
}
],
"source": [
"for step in chain.stream({\"messages\":[HumanMessage(content=\"What's the GDP of New York?\")]}):\n",
"for step in chain.stream(\n",
" {\"messages\": [HumanMessage(content=\"What's the GDP of New York?\")]}\n",
"):\n",
" print(step)\n",
" print(\"---\")"
]
+3 -3
View File
@@ -328,9 +328,9 @@
" \"set more_information_needed False and populate a blank string for the query.\"\n",
" )\n",
" input_messages = [system] + state[\"messages\"]\n",
" response = llm.bind_tools(\n",
" [QueryForTools], tool_choice=True\n",
" ).invoke(input_messages)\n",
" response = llm.bind_tools([QueryForTools], tool_choice=True).invoke(\n",
" input_messages\n",
" )\n",
" query = response.tool_calls[0][\"args\"][\"query\"]\n",
" tool_documents = vector_store.similarity_search(query)\n",
" if hack_remove_tool_condition:\n",
@@ -329,6 +329,7 @@
"\n",
"tools = [get_context, cite_context_sources]\n",
"\n",
"\n",
"# Define the function that calls the model\n",
"def call_model(state, config):\n",
" messages = state[\"messages\"]\n",
+2 -2
View File
@@ -72,12 +72,12 @@
"# Node to retrieve documents\n",
"def retrieve_documents(state: QueryOutputState) -> DocumentOutputState:\n",
" # Replace this with real logic\n",
" return {\"docs\": [state['query']] * 2}\n",
" return {\"docs\": [state[\"query\"]] * 2}\n",
"\n",
"\n",
"# Node to generate answer\n",
"def generate(state: GenerateInputState) -> OverallState:\n",
" return {\"answer\": \"\\n\\n\".join(state['docs'] + [state['question']])}\n",
" return {\"answer\": \"\\n\\n\".join(state[\"docs\"] + [state[\"question\"]])}\n",
"\n",
"\n",
"graph = StateGraph(OverallState)\n",
+10 -4
View File
@@ -630,7 +630,7 @@
" upsert=True,\n",
" )\n",
" )\n",
" await self.db[\"checkpoint_writes\"].bulk_write(operations)\n"
" await self.db[\"checkpoint_writes\"].bulk_write(operations)"
]
},
{
@@ -685,7 +685,9 @@
"metadata": {},
"outputs": [],
"source": [
"with MongoDBSaver.from_conn_info(host=\"localhost\", port=27017, db_name=\"checkpoints\") as checkpointer:\n",
"with MongoDBSaver.from_conn_info(\n",
" host=\"localhost\", port=27017, db_name=\"checkpoints\"\n",
") as checkpointer:\n",
" graph = create_react_agent(model, tools=tools, checkpointer=checkpointer)\n",
" config = {\"configurable\": {\"thread_id\": \"1\"}}\n",
" res = graph.invoke({\"messages\": [(\"human\", \"what's the weather in sf\")]}, config)\n",
@@ -796,10 +798,14 @@
"metadata": {},
"outputs": [],
"source": [
"async with AsyncMongoDBSaver.from_conn_info(host=\"localhost\", port=27017, db_name=\"checkpoints\") as checkpointer:\n",
"async with AsyncMongoDBSaver.from_conn_info(\n",
" host=\"localhost\", port=27017, db_name=\"checkpoints\"\n",
") as checkpointer:\n",
" graph = create_react_agent(model, tools=tools, checkpointer=checkpointer)\n",
" config = {\"configurable\": {\"thread_id\": \"2\"}}\n",
" res = await graph.ainvoke({\"messages\": [(\"human\", \"what's the weather in nyc\")]}, config)\n",
" res = await graph.ainvoke(\n",
" {\"messages\": [(\"human\", \"what's the weather in nyc\")]}, config\n",
" )\n",
"\n",
" latest_checkpoint = await checkpointer.aget(config)\n",
" latest_checkpoint_tuple = await checkpointer.aget_tuple(config)\n",
+3 -3
View File
@@ -134,7 +134,7 @@
"source": [
"from psycopg.rows import dict_row\n",
"\n",
"connection_kwargs ={\n",
"connection_kwargs = {\n",
" \"autocommit\": True,\n",
" \"prepare_threshold\": 0,\n",
" \"row_factory\": dict_row,\n",
@@ -166,7 +166,7 @@
" # Example configuration\n",
" conninfo=DB_URI,\n",
" max_size=20,\n",
" kwargs=connection_kwargs\n",
" kwargs=connection_kwargs,\n",
")\n",
"\n",
"with pool.connection() as conn:\n",
@@ -394,7 +394,7 @@
" # Example configuration\n",
" conninfo=DB_URI,\n",
" max_size=20,\n",
" kwargs=connection_kwargs\n",
" kwargs=connection_kwargs,\n",
") as pool, pool.connection() as conn:\n",
" checkpointer = AsyncPostgresSaver(conn)\n",
"\n",
+9 -3
View File
@@ -530,7 +530,9 @@
"\n",
" @classmethod\n",
" @asynccontextmanager\n",
" async def from_conn_info(cls, *, host: str, port: int, db: int) -> AsyncIterator[\"AsyncRedisSaver\"]:\n",
" async def from_conn_info(\n",
" cls, *, host: str, port: int, db: int\n",
" ) -> AsyncIterator[\"AsyncRedisSaver\"]:\n",
" conn = None\n",
" try:\n",
" conn = AsyncRedis(host=host, port=port, db=db)\n",
@@ -887,10 +889,14 @@
"metadata": {},
"outputs": [],
"source": [
"async with AsyncRedisSaver.from_conn_info(host=\"localhost\", port=6379, db=0) as checkpointer:\n",
"async with AsyncRedisSaver.from_conn_info(\n",
" host=\"localhost\", port=6379, db=0\n",
") as checkpointer:\n",
" graph = create_react_agent(model, tools=tools, checkpointer=checkpointer)\n",
" config = {\"configurable\": {\"thread_id\": \"2\"}}\n",
" res = await graph.ainvoke({\"messages\": [(\"human\", \"what's the weather in nyc\")]}, config)\n",
" res = await graph.ainvoke(\n",
" {\"messages\": [(\"human\", \"what's the weather in nyc\")]}, config\n",
" )\n",
"\n",
" latest_checkpoint = await checkpointer.aget(config)\n",
" latest_checkpoint_tuple = await checkpointer.aget_tuple(config)\n",
+1 -1
View File
@@ -269,7 +269,7 @@
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
" \n",
"\n",
"async def generation_node(state: Sequence[BaseMessage]):\n",
" return await generate.ainvoke({\"messages\": state})\n",
"\n",
+1
View File
@@ -392,6 +392,7 @@
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
"\n",
"MAX_ITERATIONS = 5\n",
"builder = StateGraph(State)\n",
"builder.add_node(\"draft\", first_responder.respond)\n",
+3 -1
View File
@@ -68,7 +68,9 @@
" # It's completely optional, but useful if you have many functions with similar names\n",
" gen = RunnableGenerator(my_generator).with_config(\n",
" tags=[\"should_stream\"],\n",
" callbacks=config.get(\"callbacks\", []) # <-- Propagate callbacks (Python <= 3.10)\n",
" callbacks=config.get(\n",
" \"callbacks\", []\n",
" ), # <-- Propagate callbacks (Python <= 3.10)\n",
" )\n",
" async for message in gen.astream(state):\n",
" messages.append(message)\n",
@@ -30,10 +30,7 @@
"id": "47f79af8-58d8-4a48-8d9a-88823d88701f",
"metadata": {},
"outputs": [],
"source": [
"%%capture --no-stderr\n",
"%pip install -U langgraph openai"
]
"source": ["%%capture --no-stderr\n%pip install -U langgraph openai"]
},
{
"cell_type": "code",
@@ -49,18 +46,7 @@
]
}
],
"source": [
"import getpass\n",
"import os\n",
"\n",
"\n",
"def _set_env(var: str):\n",
" if not os.environ.get(var):\n",
" os.environ[var] = getpass.getpass(f\"{var}: \")\n",
"\n",
"\n",
"_set_env(\"OPENAI_API_KEY\")"
]
"source": ["import getpass\nimport os\n\n\ndef _set_env(var: str):\n if not os.environ.get(var):\n os.environ[var] = getpass.getpass(f\"{var}: \")\n\n\n_set_env(\"OPENAI_API_KEY\")"]
},
{
"cell_type": "markdown",
@@ -84,94 +70,7 @@
"id": "d59234f9-173e-469d-a725-c13e0979663e",
"metadata": {},
"outputs": [],
"source": [
"from openai import AsyncOpenAI\n",
"from langchain_core.language_models.chat_models import ChatGenerationChunk\n",
"from langchain_core.messages import AIMessageChunk\n",
"from langchain_core.runnables.config import (\n",
" ensure_config,\n",
" get_callback_manager_for_config,\n",
")\n",
"\n",
"openai_client = AsyncOpenAI()\n",
"# define tool schema for openai tool calling\n",
"\n",
"tool = {\n",
" \"type\": \"function\",\n",
" \"function\": {\n",
" \"name\": \"get_items\",\n",
" \"description\": \"Use this tool to look up which items are in the given place.\",\n",
" \"parameters\": {\n",
" \"type\": \"object\",\n",
" \"properties\": {\"place\": {\"type\": \"string\"}},\n",
" \"required\": [\"place\"],\n",
" },\n",
" },\n",
"}\n",
"\n",
"\n",
"async def call_model(state, config=None):\n",
" config = ensure_config(config | {\"tags\": [\"agent_llm\"]})\n",
" callback_manager = get_callback_manager_for_config(config)\n",
" messages = state[\"messages\"]\n",
"\n",
" llm_run_manager = callback_manager.on_chat_model_start({}, [messages])[0]\n",
" response = await openai_client.chat.completions.create(\n",
" messages=messages, model=\"gpt-3.5-turbo\", tools=[tool], stream=True\n",
" )\n",
"\n",
" response_content = \"\"\n",
" role = None\n",
"\n",
" tool_call_id = None\n",
" tool_call_function_name = None\n",
" tool_call_function_arguments = \"\"\n",
" async for chunk in response:\n",
" delta = chunk.choices[0].delta\n",
" if delta.role is not None:\n",
" role = delta.role\n",
"\n",
" if delta.content:\n",
" response_content += delta.content\n",
" llm_run_manager.on_llm_new_token(delta.content)\n",
"\n",
" if delta.tool_calls:\n",
" # note: for simplicity we're only handling a single tool call here\n",
" if delta.tool_calls[0].function.name is not None:\n",
" tool_call_function_name = delta.tool_calls[0].function.name\n",
" tool_call_id = delta.tool_calls[0].id\n",
"\n",
" # note: we're wrapping the tools calls in ChatGenerationChunk so that the events from .astream_events in the graph can render tool calls correctly\n",
" tool_call_chunk = ChatGenerationChunk(\n",
" message=AIMessageChunk(\n",
" content=\"\",\n",
" additional_kwargs={\"tool_calls\": [delta.tool_calls[0].dict()]},\n",
" )\n",
" )\n",
" llm_run_manager.on_llm_new_token(\"\", chunk=tool_call_chunk)\n",
" tool_call_function_arguments += delta.tool_calls[0].function.arguments\n",
"\n",
" if tool_call_function_name is not None:\n",
" tool_calls = [\n",
" {\n",
" \"id\": tool_call_id,\n",
" \"function\": {\n",
" \"name\": tool_call_function_name,\n",
" \"arguments\": tool_call_function_arguments,\n",
" },\n",
" \"type\": \"function\",\n",
" }\n",
" ]\n",
" else:\n",
" tool_calls = None\n",
"\n",
" response_message = {\n",
" \"role\": role,\n",
" \"content\": response_content,\n",
" \"tool_calls\": tool_calls,\n",
" }\n",
" return {\"messages\": [response_message]}"
]
"source": ["from openai import AsyncOpenAI\nfrom langchain_core.language_models.chat_models import ChatGenerationChunk\nfrom langchain_core.messages import AIMessageChunk\nfrom langchain_core.runnables.config import (\n ensure_config,\n get_callback_manager_for_config,\n)\n\nopenai_client = AsyncOpenAI()\n# define tool schema for openai tool calling\n\ntool = {\n \"type\": \"function\",\n \"function\": {\n \"name\": \"get_items\",\n \"description\": \"Use this tool to look up which items are in the given place.\",\n \"parameters\": {\n \"type\": \"object\",\n \"properties\": {\"place\": {\"type\": \"string\"}},\n \"required\": [\"place\"],\n },\n },\n}\n\n\nasync def call_model(state, config=None):\n config = ensure_config(config | {\"tags\": [\"agent_llm\"]})\n callback_manager = get_callback_manager_for_config(config)\n messages = state[\"messages\"]\n\n llm_run_manager = callback_manager.on_chat_model_start({}, [messages])[0]\n response = await openai_client.chat.completions.create(\n messages=messages, model=\"gpt-3.5-turbo\", tools=[tool], stream=True\n )\n\n response_content = \"\"\n role = None\n\n tool_call_id = None\n tool_call_function_name = None\n tool_call_function_arguments = \"\"\n async for chunk in response:\n delta = chunk.choices[0].delta\n if delta.role is not None:\n role = delta.role\n\n if delta.content:\n response_content += delta.content\n llm_run_manager.on_llm_new_token(delta.content)\n\n if delta.tool_calls:\n # note: for simplicity we're only handling a single tool call here\n if delta.tool_calls[0].function.name is not None:\n tool_call_function_name = delta.tool_calls[0].function.name\n tool_call_id = delta.tool_calls[0].id\n\n # note: we're wrapping the tools calls in ChatGenerationChunk so that the events from .astream_events in the graph can render tool calls correctly\n tool_call_chunk = ChatGenerationChunk(\n message=AIMessageChunk(\n content=\"\",\n additional_kwargs={\"tool_calls\": [delta.tool_calls[0].dict()]},\n )\n )\n llm_run_manager.on_llm_new_token(\"\", chunk=tool_call_chunk)\n tool_call_function_arguments += delta.tool_calls[0].function.arguments\n\n if tool_call_function_name is not None:\n tool_calls = [\n {\n \"id\": tool_call_id,\n \"function\": {\n \"name\": tool_call_function_name,\n \"arguments\": tool_call_function_arguments,\n },\n \"type\": \"function\",\n }\n ]\n else:\n tool_calls = None\n\n response_message = {\n \"role\": role,\n \"content\": response_content,\n \"tool_calls\": tool_calls,\n }\n return {\"messages\": [response_message]}"]
},
{
"cell_type": "markdown",
@@ -187,62 +86,7 @@
"id": "b90941d8-afe4-42ec-9262-9c3b87c3b1ec",
"metadata": {},
"outputs": [],
"source": [
"import json\n",
"from langchain_core.callbacks import adispatch_custom_event\n",
"\n",
"\n",
"async def get_items(place: str) -> str:\n",
" \"\"\"Use this tool to look up which items are in the given place.\"\"\"\n",
"\n",
" # this can be replaced with any actual streaming logic that you might have\n",
" def stream(place: str):\n",
" if \"bed\" in place: # For under the bed\n",
" yield from [\"socks\", \"shoes\", \"dust bunnies\"]\n",
" elif \"shelf\" in place: # For 'shelf'\n",
" yield from [\"books\", \"penciles\", \"pictures\"]\n",
" else: # if the agent decides to ask about a different place\n",
" yield \"cat snacks\"\n",
"\n",
" tokens = []\n",
" for token in stream(place):\n",
" await adispatch_custom_event(\n",
" # this will allow you to filter events by name\n",
" \"tool_call_token_stream\",\n",
" {\n",
" \"function_name\": \"get_items\",\n",
" \"arguments\": {\"place\": place},\n",
" \"tool_output_token\": token,\n",
" },\n",
" # this will allow you to filter events by tags\n",
" config={\"tags\": [\"tool_call\"]},\n",
" )\n",
" tokens.append(token)\n",
"\n",
" return \", \".join(tokens)\n",
"\n",
"\n",
"# define mapping to look up functions when running tools\n",
"function_name_to_function = {\"get_items\": get_items}\n",
"\n",
"\n",
"async def call_tools(state):\n",
" messages = state[\"messages\"]\n",
"\n",
" tool_call = messages[-1][\"tool_calls\"][0]\n",
" function_name = tool_call[\"function\"][\"name\"]\n",
" function_arguments = tool_call[\"function\"][\"arguments\"]\n",
" arguments = json.loads(function_arguments)\n",
"\n",
" function_response = await function_name_to_function[function_name](**arguments)\n",
" tool_message = {\n",
" \"tool_call_id\": tool_call[\"id\"],\n",
" \"role\": \"tool\",\n",
" \"name\": function_name,\n",
" \"content\": function_response,\n",
" }\n",
" return {\"messages\": [tool_message]}"
]
"source": ["import json\nfrom langchain_core.callbacks import adispatch_custom_event\n\n\nasync def get_items(place: str) -> str:\n \"\"\"Use this tool to look up which items are in the given place.\"\"\"\n\n # this can be replaced with any actual streaming logic that you might have\n def stream(place: str):\n if \"bed\" in place: # For under the bed\n yield from [\"socks\", \"shoes\", \"dust bunnies\"]\n elif \"shelf\" in place: # For 'shelf'\n yield from [\"books\", \"penciles\", \"pictures\"]\n else: # if the agent decides to ask about a different place\n yield \"cat snacks\"\n\n tokens = []\n for token in stream(place):\n await adispatch_custom_event(\n # this will allow you to filter events by name\n \"tool_call_token_stream\",\n {\n \"function_name\": \"get_items\",\n \"arguments\": {\"place\": place},\n \"tool_output_token\": token,\n },\n # this will allow you to filter events by tags\n config={\"tags\": [\"tool_call\"]},\n )\n tokens.append(token)\n\n return \", \".join(tokens)\n\n\n# define mapping to look up functions when running tools\nfunction_name_to_function = {\"get_items\": get_items}\n\n\nasync def call_tools(state):\n messages = state[\"messages\"]\n\n tool_call = messages[-1][\"tool_calls\"][0]\n function_name = tool_call[\"function\"][\"name\"]\n function_arguments = tool_call[\"function\"][\"arguments\"]\n arguments = json.loads(function_arguments)\n\n function_response = await function_name_to_function[function_name](**arguments)\n tool_message = {\n \"tool_call_id\": tool_call[\"id\"],\n \"role\": \"tool\",\n \"name\": function_name,\n \"content\": function_response,\n }\n return {\"messages\": [tool_message]}"]
},
{
"cell_type": "markdown",
@@ -258,33 +102,7 @@
"id": "228260be-1f9a-4195-80e0-9604f8a5dba6",
"metadata": {},
"outputs": [],
"source": [
"import operator\n",
"from typing import Annotated, TypedDict, Literal\n",
"\n",
"from langgraph.graph import StateGraph, END\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list, operator.add]\n",
"\n",
"\n",
"def should_continue(state) -> Literal[\"tools\", END]:\n",
" messages = state[\"messages\"]\n",
" last_message = messages[-1]\n",
" if last_message[\"tool_calls\"]:\n",
" return \"tools\"\n",
" return END\n",
"\n",
"\n",
"workflow = StateGraph(State)\n",
"workflow.set_entry_point(\"model\")\n",
"workflow.add_node(\"model\", call_model) # i.e. our \"agent\"\n",
"workflow.add_node(\"tools\", call_tools)\n",
"workflow.add_conditional_edges(\"model\", should_continue)\n",
"workflow.add_edge(\"tools\", \"model\")\n",
"graph = workflow.compile()"
]
"source": ["import operator\nfrom typing import Annotated, TypedDict, Literal\n\nfrom langgraph.graph import StateGraph, END, START\n\n\nclass State(TypedDict):\n messages: Annotated[list, operator.add]\n\n\ndef should_continue(state) -> Literal[\"tools\", END]:\n messages = state[\"messages\"]\n last_message = messages[-1]\n if last_message[\"tool_calls\"]:\n return \"tools\"\n return END\n\n\nworkflow = StateGraph(State)\nworkflow.add_edge(START, \"model\")\nworkflow.add_node(\"model\", call_model) # i.e. our \"agent\"\nworkflow.add_node(\"tools\", call_tools)\nworkflow.add_conditional_edges(\"model\", should_continue)\nworkflow.add_edge(\"tools\", \"model\")\ngraph = workflow.compile()"]
},
{
"cell_type": "markdown",
@@ -318,14 +136,7 @@
]
}
],
"source": [
"async for event in graph.astream_events(\n",
" {\"messages\": [{\"role\": \"user\", \"content\": \"what's in the bedroom\"}]}, version=\"v2\"\n",
"):\n",
" tags = event.get(\"tags\", [])\n",
" if event[\"event\"] == \"on_custom_event\" and \"tool_call\" in tags:\n",
" print(\"Tool token\", event[\"data\"][\"tool_output_token\"])"
]
"source": ["async for event in graph.astream_events(\n {\"messages\": [{\"role\": \"user\", \"content\": \"what's in the bedroom\"}]}, version=\"v2\"\n):\n tags = event.get(\"tags\", [])\n if event[\"event\"] == \"on_custom_event\" and \"tool_call\" in tags:\n print(\"Tool token\", event[\"data\"][\"tool_output_token\"])"]
}
],
"metadata": {
@@ -30,10 +30,7 @@
"id": "47f79af8-58d8-4a48-8d9a-88823d88701f",
"metadata": {},
"outputs": [],
"source": [
"%%capture --no-stderr\n",
"%pip install -U langgraph openai"
]
"source": ["%%capture --no-stderr\n%pip install -U langgraph openai"]
},
{
"cell_type": "code",
@@ -49,18 +46,7 @@
]
}
],
"source": [
"import getpass\n",
"import os\n",
"\n",
"\n",
"def _set_env(var: str):\n",
" if not os.environ.get(var):\n",
" os.environ[var] = getpass.getpass(f\"{var}: \")\n",
"\n",
"\n",
"_set_env(\"OPENAI_API_KEY\")"
]
"source": ["import getpass\nimport os\n\n\ndef _set_env(var: str):\n if not os.environ.get(var):\n os.environ[var] = getpass.getpass(f\"{var}: \")\n\n\n_set_env(\"OPENAI_API_KEY\")"]
},
{
"cell_type": "markdown",
@@ -84,94 +70,7 @@
"id": "d59234f9-173e-469d-a725-c13e0979663e",
"metadata": {},
"outputs": [],
"source": [
"from openai import AsyncOpenAI\n",
"from langchain_core.language_models.chat_models import ChatGenerationChunk\n",
"from langchain_core.messages import AIMessageChunk\n",
"from langchain_core.runnables.config import (\n",
" ensure_config,\n",
" get_callback_manager_for_config,\n",
")\n",
"\n",
"openai_client = AsyncOpenAI()\n",
"# define tool schema for openai tool calling\n",
"\n",
"tool = {\n",
" \"type\": \"function\",\n",
" \"function\": {\n",
" \"name\": \"get_items\",\n",
" \"description\": \"Use this tool to look up which items are in the given place.\",\n",
" \"parameters\": {\n",
" \"type\": \"object\",\n",
" \"properties\": {\"place\": {\"type\": \"string\"}},\n",
" \"required\": [\"place\"],\n",
" },\n",
" },\n",
"}\n",
"\n",
"\n",
"async def call_model(state, config=None):\n",
" config = ensure_config(config | {\"tags\": [\"agent_llm\"]})\n",
" callback_manager = get_callback_manager_for_config(config)\n",
" messages = state[\"messages\"]\n",
"\n",
" llm_run_manager = callback_manager.on_chat_model_start({}, [messages])[0]\n",
" response = await openai_client.chat.completions.create(\n",
" messages=messages, model=\"gpt-3.5-turbo\", tools=[tool], stream=True\n",
" )\n",
"\n",
" response_content = \"\"\n",
" role = None\n",
"\n",
" tool_call_id = None\n",
" tool_call_function_name = None\n",
" tool_call_function_arguments = \"\"\n",
" async for chunk in response:\n",
" delta = chunk.choices[0].delta\n",
" if delta.role is not None:\n",
" role = delta.role\n",
"\n",
" if delta.content:\n",
" response_content += delta.content\n",
" llm_run_manager.on_llm_new_token(delta.content)\n",
"\n",
" if delta.tool_calls:\n",
" # note: for simplicity we're only handling a single tool call here\n",
" if delta.tool_calls[0].function.name is not None:\n",
" tool_call_function_name = delta.tool_calls[0].function.name\n",
" tool_call_id = delta.tool_calls[0].id\n",
"\n",
" # note: we're wrapping the tools calls in ChatGenerationChunk so that the events from .astream_events in the graph can render tool calls correctly\n",
" tool_call_chunk = ChatGenerationChunk(\n",
" message=AIMessageChunk(\n",
" content=\"\",\n",
" additional_kwargs={\"tool_calls\": [delta.tool_calls[0].dict()]},\n",
" )\n",
" )\n",
" llm_run_manager.on_llm_new_token(\"\", chunk=tool_call_chunk)\n",
" tool_call_function_arguments += delta.tool_calls[0].function.arguments\n",
"\n",
" if tool_call_function_name is not None:\n",
" tool_calls = [\n",
" {\n",
" \"id\": tool_call_id,\n",
" \"function\": {\n",
" \"name\": tool_call_function_name,\n",
" \"arguments\": tool_call_function_arguments,\n",
" },\n",
" \"type\": \"function\",\n",
" }\n",
" ]\n",
" else:\n",
" tool_calls = None\n",
"\n",
" response_message = {\n",
" \"role\": role,\n",
" \"content\": response_content,\n",
" \"tool_calls\": tool_calls,\n",
" }\n",
" return {\"messages\": [response_message]}"
]
"source": ["from openai import AsyncOpenAI\nfrom langchain_core.language_models.chat_models import ChatGenerationChunk\nfrom langchain_core.messages import AIMessageChunk\nfrom langchain_core.runnables.config import (\n ensure_config,\n get_callback_manager_for_config,\n)\n\nopenai_client = AsyncOpenAI()\n# define tool schema for openai tool calling\n\ntool = {\n \"type\": \"function\",\n \"function\": {\n \"name\": \"get_items\",\n \"description\": \"Use this tool to look up which items are in the given place.\",\n \"parameters\": {\n \"type\": \"object\",\n \"properties\": {\"place\": {\"type\": \"string\"}},\n \"required\": [\"place\"],\n },\n },\n}\n\n\nasync def call_model(state, config=None):\n config = ensure_config(config | {\"tags\": [\"agent_llm\"]})\n callback_manager = get_callback_manager_for_config(config)\n messages = state[\"messages\"]\n\n llm_run_manager = callback_manager.on_chat_model_start({}, [messages])[0]\n response = await openai_client.chat.completions.create(\n messages=messages, model=\"gpt-3.5-turbo\", tools=[tool], stream=True\n )\n\n response_content = \"\"\n role = None\n\n tool_call_id = None\n tool_call_function_name = None\n tool_call_function_arguments = \"\"\n async for chunk in response:\n delta = chunk.choices[0].delta\n if delta.role is not None:\n role = delta.role\n\n if delta.content:\n response_content += delta.content\n llm_run_manager.on_llm_new_token(delta.content)\n\n if delta.tool_calls:\n # note: for simplicity we're only handling a single tool call here\n if delta.tool_calls[0].function.name is not None:\n tool_call_function_name = delta.tool_calls[0].function.name\n tool_call_id = delta.tool_calls[0].id\n\n # note: we're wrapping the tools calls in ChatGenerationChunk so that the events from .astream_events in the graph can render tool calls correctly\n tool_call_chunk = ChatGenerationChunk(\n message=AIMessageChunk(\n content=\"\",\n additional_kwargs={\"tool_calls\": [delta.tool_calls[0].dict()]},\n )\n )\n llm_run_manager.on_llm_new_token(\"\", chunk=tool_call_chunk)\n tool_call_function_arguments += delta.tool_calls[0].function.arguments\n\n if tool_call_function_name is not None:\n tool_calls = [\n {\n \"id\": tool_call_id,\n \"function\": {\n \"name\": tool_call_function_name,\n \"arguments\": tool_call_function_arguments,\n },\n \"type\": \"function\",\n }\n ]\n else:\n tool_calls = None\n\n response_message = {\n \"role\": role,\n \"content\": response_content,\n \"tool_calls\": tool_calls,\n }\n return {\"messages\": [response_message]}"]
},
{
"cell_type": "markdown",
@@ -187,41 +86,7 @@
"id": "b756ea32",
"metadata": {},
"outputs": [],
"source": [
"import json\n",
"\n",
"\n",
"async def get_items(place: str) -> str:\n",
" \"\"\"Use this tool to look up which items are in the given place.\"\"\"\n",
" if \"bed\" in place: # For under the bed\n",
" return \"socks, shoes and dust bunnies\"\n",
" if \"shelf\" in place: # For 'shelf'\n",
" return \"books, penciles and pictures\"\n",
" else: # if the agent decides to ask about a different place\n",
" return \"cat snacks\"\n",
"\n",
"\n",
"# define mapping to look up functions when running tools\n",
"function_name_to_function = {\"get_items\": get_items}\n",
"\n",
"\n",
"async def call_tools(state):\n",
" messages = state[\"messages\"]\n",
"\n",
" tool_call = messages[-1][\"tool_calls\"][0]\n",
" function_name = tool_call[\"function\"][\"name\"]\n",
" function_arguments = tool_call[\"function\"][\"arguments\"]\n",
" arguments = json.loads(function_arguments)\n",
"\n",
" function_response = await function_name_to_function[function_name](**arguments)\n",
" tool_message = {\n",
" \"tool_call_id\": tool_call[\"id\"],\n",
" \"role\": \"tool\",\n",
" \"name\": function_name,\n",
" \"content\": function_response,\n",
" }\n",
" return {\"messages\": [tool_message]}"
]
"source": ["import json\n\n\nasync def get_items(place: str) -> str:\n \"\"\"Use this tool to look up which items are in the given place.\"\"\"\n if \"bed\" in place: # For under the bed\n return \"socks, shoes and dust bunnies\"\n if \"shelf\" in place: # For 'shelf'\n return \"books, penciles and pictures\"\n else: # if the agent decides to ask about a different place\n return \"cat snacks\"\n\n\n# define mapping to look up functions when running tools\nfunction_name_to_function = {\"get_items\": get_items}\n\n\nasync def call_tools(state):\n messages = state[\"messages\"]\n\n tool_call = messages[-1][\"tool_calls\"][0]\n function_name = tool_call[\"function\"][\"name\"]\n function_arguments = tool_call[\"function\"][\"arguments\"]\n arguments = json.loads(function_arguments)\n\n function_response = await function_name_to_function[function_name](**arguments)\n tool_message = {\n \"tool_call_id\": tool_call[\"id\"],\n \"role\": \"tool\",\n \"name\": function_name,\n \"content\": function_response,\n }\n return {\"messages\": [tool_message]}"]
},
{
"cell_type": "markdown",
@@ -237,33 +102,7 @@
"id": "228260be-1f9a-4195-80e0-9604f8a5dba6",
"metadata": {},
"outputs": [],
"source": [
"import operator\n",
"from typing import Annotated, TypedDict, Literal\n",
"\n",
"from langgraph.graph import StateGraph, END\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list, operator.add]\n",
"\n",
"\n",
"def should_continue(state) -> Literal[\"tools\", END]:\n",
" messages = state[\"messages\"]\n",
" last_message = messages[-1]\n",
" if last_message[\"tool_calls\"]:\n",
" return \"tools\"\n",
" return END\n",
"\n",
"\n",
"workflow = StateGraph(State)\n",
"workflow.set_entry_point(\"model\")\n",
"workflow.add_node(\"model\", call_model) # i.e. our \"agent\"\n",
"workflow.add_node(\"tools\", call_tools)\n",
"workflow.add_conditional_edges(\"model\", should_continue)\n",
"workflow.add_edge(\"tools\", \"model\")\n",
"graph = workflow.compile()"
]
"source": ["import operator\nfrom typing import Annotated, TypedDict, Literal\n\nfrom langgraph.graph import StateGraph, END, START\n\n\nclass State(TypedDict):\n messages: Annotated[list, operator.add]\n\n\ndef should_continue(state) -> Literal[\"tools\", END]:\n messages = state[\"messages\"]\n last_message = messages[-1]\n if last_message[\"tool_calls\"]:\n return \"tools\"\n return END\n\n\nworkflow = StateGraph(State)\nworkflow.add_edge(START, \"model\")\nworkflow.add_node(\"model\", call_model) # i.e. our \"agent\"\nworkflow.add_node(\"tools\", call_tools)\nworkflow.add_conditional_edges(\"model\", should_continue)\nworkflow.add_edge(\"tools\", \"model\")\ngraph = workflow.compile()"]
},
{
"cell_type": "markdown",
@@ -328,14 +167,7 @@
]
}
],
"source": [
"async for event in graph.astream_events(\n",
" {\"messages\": [{\"role\": \"user\", \"content\": \"what's in the bedroom\"}]}, version=\"v2\"\n",
"):\n",
" tags = event.get(\"tags\", [])\n",
" if event[\"event\"] == \"on_chat_model_stream\" and \"agent_llm\" in tags:\n",
" print(\"LLM token\", event[\"data\"][\"chunk\"].dict())"
]
"source": ["async for event in graph.astream_events(\n {\"messages\": [{\"role\": \"user\", \"content\": \"what's in the bedroom\"}]}, version=\"v2\"\n):\n tags = event.get(\"tags\", [])\n if event[\"event\"] == \"on_chat_model_stream\" and \"agent_llm\" in tags:\n print(\"LLM token\", event[\"data\"][\"chunk\"].dict())"]
},
{
"cell_type": "code",
@@ -343,7 +175,7 @@
"id": "adb0f7bc-6e51-478e-bd32-8f72df072d6c",
"metadata": {},
"outputs": [],
"source": []
"source": [""]
}
],
"metadata": {
@@ -169,9 +169,7 @@
"from langchain_core.output_parsers import JsonOutputParser\n",
"\n",
"# JSON\n",
"llm = ChatOllama(model=\"llama3.1\", \n",
" format=\"json\", \n",
" temperature=0)\n",
"llm = ChatOllama(model=\"llama3.1\", format=\"json\", temperature=0)\n",
"\n",
"\n",
"prompt = PromptTemplate(\n",
@@ -210,6 +208,7 @@
"from IPython.display import Image, display\n",
"from langgraph.graph import START, END, StateGraph\n",
"\n",
"\n",
"class GraphState(TypedDict):\n",
" \"\"\"\n",
" Represents the state of our graph.\n",
@@ -356,7 +355,7 @@
"workflow.add_node(\"web_search\", web_search) # web search\n",
"\n",
"# Build graph\n",
"workflow.set_entry_point(\"retrieve\")\n",
"workflow.add_edge(START, retrieve)\n",
"workflow.add_edge(\"retrieve\", \"grade_documents\")\n",
"workflow.add_conditional_edges(\n",
" \"grade_documents\",\n",
@@ -381,21 +380,22 @@
"metadata": {},
"outputs": [],
"source": [
"import uuid \n",
"import uuid\n",
"\n",
"\n",
"def predict_custom_agent_answer(example: dict):\n",
" \n",
" config = {\"configurable\": {\"thread_id\": str(uuid.uuid4())}}\n",
" \n",
"\n",
" state_dict = custom_graph.invoke(\n",
" {\"question\": example[\"input\"], \"steps\": []}, config\n",
" )\n",
" \n",
"\n",
" return {\"response\": state_dict[\"generation\"], \"steps\": state_dict[\"steps\"]}\n",
"\n",
"\n",
"example = {\"input\": \"What are the types of agent memory?\"}\n",
"#response = predict_custom_agent_answer(example)\n",
"#response"
"# response = predict_custom_agent_answer(example)\n",
"# response"
]
},
{
@@ -544,6 +544,7 @@
" \"generate_answer\",\n",
"]\n",
"\n",
"\n",
"def check_trajectory_custom(root_run: Run, example: Example) -> dict:\n",
" \"\"\"\n",
" Check if all expected tools are called in exact order and without any additional tool calls.\n",
@@ -134,6 +134,7 @@
" for d in web_results\n",
" ]\n",
"\n",
"\n",
"# Tool list\n",
"tools = [retrieve_documents, web_search]"
]
@@ -152,9 +153,11 @@
"from langgraph.graph.message import AnyMessage, add_messages\n",
"from typing_extensions import TypedDict\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list[AnyMessage], add_messages]\n",
"\n",
"\n",
"class Assistant:\n",
" def __init__(self, runnable: Runnable):\n",
" \"\"\"\n",
@@ -291,6 +294,7 @@
"source": [
"import uuid\n",
"\n",
"\n",
"def predict_react_agent_answer(example: dict):\n",
" \"\"\"Use this for answer evaluation\"\"\"\n",
"\n",
+2 -2
View File
@@ -457,13 +457,13 @@
"source": [
"from langchain_core.runnables import RunnableLambda\n",
"\n",
"from langgraph.graph import END, StateGraph\n",
"from langgraph.graph import END, START, StateGraph\n",
"\n",
"graph_builder = StateGraph(AgentState)\n",
"\n",
"\n",
"graph_builder.add_node(\"agent\", agent)\n",
"graph_builder.set_entry_point(\"agent\")\n",
"graph_builder.add_edge(START, \"agent\")\n",
"\n",
"graph_builder.add_node(\"update_scratchpad\", update_scratchpad)\n",
"graph_builder.add_edge(\"update_scratchpad\", \"agent\")\n",