diff --git a/docs/docs/tutorials/functional_api/functional_api_test.ipynb b/docs/docs/tutorials/functional_api/functional_api_test.ipynb index ecb437454..d2f7fba73 100644 --- a/docs/docs/tutorials/functional_api/functional_api_test.ipynb +++ b/docs/docs/tutorials/functional_api/functional_api_test.ipynb @@ -11,7 +11,7 @@ }, { "cell_type": "code", - "execution_count": null, + "execution_count": 1, "metadata": {}, "outputs": [], "source": [ @@ -28,7 +28,7 @@ }, { "cell_type": "code", - "execution_count": 9, + "execution_count": 57, "metadata": {}, "outputs": [], "source": [ @@ -44,20 +44,16 @@ "### Vanilla Agent\n", "\n", "* No orchestration framework \n", - "* Optionally, use LangGraph to bind tools and specify tools " + "* Use LangChain to bind tools and specify tools " ] }, { "cell_type": "code", - "execution_count": 10, + "execution_count": 90, "metadata": {}, "outputs": [], "source": [ "from langchain_core.tools import tool\n", - "from langchain_openai import ChatOpenAI\n", - "\n", - "# LLM\n", - "llm = ChatOpenAI(model=\"gpt-4o\")\n", "\n", "# Define tools\n", "@tool\n", @@ -100,7 +96,7 @@ }, { "cell_type": "code", - "execution_count": 11, + "execution_count": 63, "metadata": {}, "outputs": [ { @@ -111,9 +107,11 @@ "\n", "Add 3 and 4.\n", "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add those numbers using the `add` function.\", 'type': 'text'}, {'id': 'toolu_01Lr23DbYsJvGSuzwQdXqDvH', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", "Tool Calls:\n", - " add (call_N1trAdi9h9vK0IW3vsvCRaaA)\n", - " Call ID: call_N1trAdi9h9vK0IW3vsvCRaaA\n", + " add (toolu_01Lr23DbYsJvGSuzwQdXqDvH)\n", + " Call ID: toolu_01Lr23DbYsJvGSuzwQdXqDvH\n", " Args:\n", " a: 3\n", " b: 4\n", @@ -181,46 +179,42 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "### Agent with short-term memory\n", + "### Agent with short-term memory (within a thread)\n", "\n", "* LangGraph persistence layer \n", - "* `@entrypoint` decorator indicates the start of a workflow. " + "\n", + "`@entrypoint` \n", + "* Decorator indicates the start of a workflow / agent \n", + "* Produces a Pregel object, an abstraction for managing a few things\n", + "* Execution -- Syncronous (invoke), Async (ainvoke), streaming (stream)\n", + "* State -- Checkpointing, Human in the loop (interrupt)\n", + "\n", + "[Optional: `@entrypoint.final`](https://langchain-ai.github.io/langgraph/concepts/functional_api/#entrypointfinal)\n", + "* Can be used to specify what to return vs what to checkpoint \n", + "\n", + "`@task`\n", + "* Results from tasks are saved as checkpoints\n", + "* Important for caching results (time-consuming operations)\n", + "* Support streaming updates from tasks\n", + "* Support tracing\n", + "\n", + "Calling a task -- \n", + "* When you call a task, it returns immediately with a future object.\n", + "* A future is a placeholder for a result that will be available later.\n", + "* `.result()` marks where in the code you actually need the task's result." ] }, { "cell_type": "code", - "execution_count": 12, + "execution_count": 65, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "================================\u001b[1m Human Message \u001b[0m=================================\n", - "\n", - "Add 3 and 4.\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "Tool Calls:\n", - " add (call_wVN0NiCHHueHRdaDGFEt1hFh)\n", - " Call ID: call_wVN0NiCHHueHRdaDGFEt1hFh\n", - " Args:\n", - " a: 3\n", - " b: 4\n", - "=================================\u001b[1m Tool Message \u001b[0m=================================\n", - "Name: add\n", - "\n", - "7\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "The result of adding 3 and 4 is 7.\n" - ] - } - ], + "outputs": [], "source": [ "import uuid\n", - "from langgraph.func import entrypoint # New \n", + "from langgraph.func import entrypoint, task # New \n", "from langgraph.checkpoint.memory import MemorySaver # New \n", "\n", + "@task # New\n", "def call_llm(messages: list[BaseMessage]):\n", " \"\"\"LLM decides whether to call a tool or not\"\"\"\n", " return llm_with_tools.invoke(\n", @@ -232,13 +226,15 @@ " + messages\n", " )\n", "\n", + "@task # New\n", "def call_tool(tool_call: ToolCall):\n", " \"\"\"Performs the tool call\"\"\"\n", "\n", " tool = tools_by_name[tool_call[\"name\"]]\n", " return tool.invoke(tool_call)\n", "\n", - "@entrypoint(checkpointer=MemorySaver()) # New \n", + "checkpointer = MemorySaver()\n", + "@entrypoint(checkpointer=checkpointer) # New \n", "def agent(messages: list[BaseMessage], previous: list[BaseMessage]): # New \n", " \"\"\" Tool calling agent \"\"\"\n", "\n", @@ -247,7 +243,7 @@ " messages = add_messages(previous, messages)\n", " \n", " # Call the LLM\n", - " llm_response = call_llm(messages)\n", + " llm_response = call_llm(messages).result()\n", "\n", " while True:\n", " if not llm_response.tool_calls:\n", @@ -255,29 +251,20 @@ "\n", " # Execute tools\n", " tool_results = [\n", - " call_tool(tool_call) for tool_call in llm_response.tool_calls\n", + " call_tool(tool_call).result() for tool_call in llm_response.tool_calls\n", " ]\n", " messages = add_messages(messages, [llm_response, *tool_results])\n", - " llm_response = call_llm(messages)\n", + " llm_response = call_llm(messages).result()\n", "\n", " messages = add_messages(messages, llm_response)\n", - " return messages\n", "\n", - "# Thread ID\n", - "thread_id = str(uuid.uuid4())\n", - "\n", - "# Config\n", - "config = {\"configurable\": {\"thread_id\": thread_id}}\n", - "\n", - "# Run with checkpointer to persist state in memory\n", - "messages = agent.invoke([HumanMessage(content=\"Add 3 and 4.\")], config)\n", - "for m in messages:\n", - " m.pretty_print()" + " # Return LLM response and save the full message history\n", + " return messages" ] }, { "cell_type": "code", - "execution_count": 57, + "execution_count": 66, "metadata": {}, "outputs": [ { @@ -288,9 +275,11 @@ "\n", "Add 3 and 4.\n", "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_013z1syyD8SvfDKdysUSwyeq', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", "Tool Calls:\n", - " add (call_FisH9R1uplO9Nyx7fzWDY6uE)\n", - " Call ID: call_FisH9R1uplO9Nyx7fzWDY6uE\n", + " add (toolu_013z1syyD8SvfDKdysUSwyeq)\n", + " Call ID: toolu_013z1syyD8SvfDKdysUSwyeq\n", " Args:\n", " a: 3\n", " b: 4\n", @@ -305,15 +294,21 @@ } ], "source": [ - "# Checkpoint state\n", - "agent_state = agent.get_state(config)\n", - "for m in agent_state.values:\n", + "# Thread ID\n", + "thread_id = str(uuid.uuid4())\n", + "\n", + "# Config\n", + "config = {\"configurable\": {\"thread_id\": thread_id}}\n", + "\n", + "# Run with checkpointer to persist state in memory\n", + "messages = agent.invoke([HumanMessage(content=\"Add 3 and 4.\")], config)\n", + "for m in messages:\n", " m.pretty_print()" ] }, { "cell_type": "code", - "execution_count": 58, + "execution_count": 67, "metadata": {}, "outputs": [ { @@ -324,9 +319,49 @@ "\n", "Add 3 and 4.\n", "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_013z1syyD8SvfDKdysUSwyeq', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", "Tool Calls:\n", - " add (call_FisH9R1uplO9Nyx7fzWDY6uE)\n", - " Call ID: call_FisH9R1uplO9Nyx7fzWDY6uE\n", + " add (toolu_013z1syyD8SvfDKdysUSwyeq)\n", + " Call ID: toolu_013z1syyD8SvfDKdysUSwyeq\n", + " Args:\n", + " a: 3\n", + " b: 4\n", + "=================================\u001b[1m Tool Message \u001b[0m=================================\n", + "Name: add\n", + "\n", + "7\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The sum of 3 and 4 is 7.\n" + ] + } + ], + "source": [ + "# Get the last checkpoint, which contains the full message history\n", + "agent_state = agent.get_state(config)\n", + "for m in agent_state.values:\n", + " m.pretty_print()" + ] + }, + { + "cell_type": "code", + "execution_count": 68, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================\u001b[1m Human Message \u001b[0m=================================\n", + "\n", + "Add 3 and 4.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_013z1syyD8SvfDKdysUSwyeq', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", + "Tool Calls:\n", + " add (toolu_013z1syyD8SvfDKdysUSwyeq)\n", + " Call ID: toolu_013z1syyD8SvfDKdysUSwyeq\n", " Args:\n", " a: 3\n", " b: 4\n", @@ -341,9 +376,11 @@ "\n", "Take the result and multiply it by 2.\n", "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll multiply the previous result (7) by 2 using the `multiply` function.\", 'type': 'text'}, {'id': 'toolu_01RhScbTXybkGxYG6RA1tQpM', 'input': {'a': 7, 'b': 2}, 'name': 'multiply', 'type': 'tool_use'}]\n", "Tool Calls:\n", - " multiply (call_Q7TTDpimEZf0QC6klEx2qTRJ)\n", - " Call ID: call_Q7TTDpimEZf0QC6klEx2qTRJ\n", + " multiply (toolu_01RhScbTXybkGxYG6RA1tQpM)\n", + " Call ID: toolu_01RhScbTXybkGxYG6RA1tQpM\n", " Args:\n", " a: 7\n", " b: 2\n", @@ -364,124 +401,29 @@ " m.pretty_print()" ] }, - { - "cell_type": "code", - "execution_count": 60, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "The result of multiplying 42 by 3 is 126.\n" - ] - } - ], - "source": [ - "# Continue with the same thread\n", - "for item in agent.stream([HumanMessage(content=\"Take the result and multiply it by 3.\")], config, stream_mode=\"values\"):\n", - " item[-1].pretty_print()" - ] - }, - { - "cell_type": "code", - "execution_count": 41, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "================================\u001b[1m Human Message \u001b[0m=================================\n", - "\n", - "Add 3 and 4.\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "Tool Calls:\n", - " add (call_5Ff1Bj5S4TYSYAoW1TpSFeyA)\n", - " Call ID: call_5Ff1Bj5S4TYSYAoW1TpSFeyA\n", - " Args:\n", - " a: 3\n", - " b: 4\n", - "=================================\u001b[1m Tool Message \u001b[0m=================================\n", - "Name: add\n", - "\n", - "7\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "The sum of 3 and 4 is 7.\n", - "================================\u001b[1m Human Message \u001b[0m=================================\n", - "\n", - "Take the result and multiply it by 2.\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "Tool Calls:\n", - " multiply (call_kk8W7CO9hAZgYwLVaISHGvKs)\n", - " Call ID: call_kk8W7CO9hAZgYwLVaISHGvKs\n", - " Args:\n", - " a: 7\n", - " b: 2\n", - "=================================\u001b[1m Tool Message \u001b[0m=================================\n", - "Name: multiply\n", - "\n", - "14\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "The result of multiplying 7 by 2 is 14.\n", - "================================\u001b[1m Human Message \u001b[0m=================================\n", - "\n", - "Take the result and multiply it by 3.\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "Tool Calls:\n", - " multiply (call_38ckbEtvlXDKAZLWMoTULVQs)\n", - " Call ID: call_38ckbEtvlXDKAZLWMoTULVQs\n", - " Args:\n", - " a: 14\n", - " b: 3\n", - "=================================\u001b[1m Tool Message \u001b[0m=================================\n", - "Name: multiply\n", - "\n", - "42\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "Multiplying 14 by 3 gives you 42.\n" - ] - } - ], - "source": [ - "# Checkpoint state\n", - "agent_state = agent.get_state(config)\n", - "for m in agent_state.values:\n", - " m.pretty_print()" - ] - }, { "cell_type": "markdown", "metadata": {}, "source": [ "### Agent with HITL\n", "\n", - "* Add interrupt to the workflow to allow for HITL" + "* Add interrupt to the workflow to allow for HITL.\n", + "* Re-executes from the start of the entrypoint.\n", + "* Output of each `@task` is cached / saved as a checkpoint." ] }, { "cell_type": "code", - "execution_count": 13, + "execution_count": 72, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "{'tool_call': {'name': 'add', 'args': {'a': 3, 'b': 4}, 'id': 'call_Np3qpF1w2n6VHgIEXNvw7duZ', 'type': 'tool_call'}, 'action': 'Please approve/reject the tool call'}\n" - ] - } - ], + "outputs": [], "source": [ "from langgraph.types import interrupt\n", "\n", + "@task\n", "def call_llm(messages: list[BaseMessage]):\n", " \"\"\"LLM decides whether to call a tool or not\"\"\"\n", + " print(\"Calling LLM or using cached LLM output!\")\n", " return llm_with_tools.invoke(\n", " [\n", " SystemMessage(\n", @@ -491,9 +433,10 @@ " + messages\n", " )\n", "\n", + "@task\n", "def call_tool(tool_call: ToolCall):\n", " \"\"\"Performs the tool call\"\"\"\n", - "\n", + " print(\"Calling tool or using cached tool output!\")\n", " # Interrupt the workflow to get a review from a human.\n", " is_approved = interrupt({ # New \n", " # Any json-serializable payload provided to interrupt as argument.\n", @@ -520,7 +463,7 @@ " messages = add_messages(previous, messages)\n", " \n", " # Call the LLM\n", - " llm_response = call_llm(messages)\n", + " llm_response = call_llm(messages).result()\n", "\n", " while True:\n", " if not llm_response.tool_calls:\n", @@ -528,14 +471,31 @@ "\n", " # Execute tools\n", " tool_results = [\n", - " call_tool(tool_call) for tool_call in llm_response.tool_calls\n", + " call_tool(tool_call).result() for tool_call in llm_response.tool_calls\n", " ]\n", " messages = add_messages(messages, [llm_response, *tool_results])\n", - " llm_response = call_llm(messages)\n", + " llm_response = call_llm(messages).result()\n", "\n", " messages = add_messages(messages, llm_response)\n", - " return messages\n", - "\n", + " return messages" + ] + }, + { + "cell_type": "code", + "execution_count": 76, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Calling LLM or using cached LLM output!\n", + "Calling tool or using cached tool output!\n", + "{'tool_call': {'name': 'add', 'args': {'a': 3, 'b': 4}, 'id': 'toolu_017uJdvp4SfEyPLVWamxBjGT', 'type': 'tool_call'}, 'action': 'Please approve/reject the tool call'}\n" + ] + } + ], + "source": [ "# Thread ID\n", "thread_id = str(uuid.uuid4())\n", "\n", @@ -544,18 +504,21 @@ "\n", "# Run until the interrupt \n", "for item in agent.stream([HumanMessage(content=\"Add 3 and 4.\")], config, stream_mode=\"updates\"):\n", - " print(item['__interrupt__'][0].value)" + " if '__interrupt__' in item:\n", + " print(item['__interrupt__'][0].value)" ] }, { "cell_type": "code", - "execution_count": 83, + "execution_count": 77, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ + "Calling tool or using cached tool output!\n", + "Calling LLM or using cached LLM output!\n", "==================================\u001b[1m Ai Message \u001b[0m==================================\n", "\n", "The sum of 3 and 4 is 7.\n" @@ -565,14 +528,270 @@ "source": [ "from langgraph.types import Command\n", "for item in agent.stream(Command(resume=True), config, stream_mode=\"updates\"):\n", - " item['agent'][-1].pretty_print()" + " if 'agent' in item:\n", + " item['agent'][-1].pretty_print()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ - "### Agent with HITL and Long-term memory\n", + "### Time travel\n", + "\n", + "* It can be useful to add time travel to the workflow.\n" + ] + }, + { + "cell_type": "code", + "execution_count": 108, + "metadata": {}, + "outputs": [], + "source": [ + "@task\n", + "def call_llm(messages: list[BaseMessage]):\n", + " \"\"\"LLM decides whether to call a tool or not\"\"\"\n", + " return llm_with_tools.invoke(\n", + " [\n", + " SystemMessage(\n", + " content=\"You are a helpful assistant tasked with performing arithmetic on a set of inputs.\"\n", + " )\n", + " ]\n", + " + messages\n", + " )\n", + "\n", + "@task\n", + "def call_tool(tool_call: ToolCall):\n", + " \"\"\"Performs the tool call\"\"\"\n", + " tool = tools_by_name[tool_call[\"name\"]]\n", + " return tool.invoke(tool_call)\n", + "\n", + "checkpointer = MemorySaver()\n", + "@entrypoint(checkpointer=checkpointer) \n", + "def agent(messages: list[BaseMessage], previous: list[BaseMessage]): # New \n", + " \"\"\" Tool calling agent \"\"\"\n", + " # Add previous messages from short-term memory to the current messages\n", + " if previous is not None:\n", + " messages = add_messages(previous, messages)\n", + " \n", + " # Call the LLM\n", + " llm_response = call_llm(messages).result()\n", + "\n", + " while True:\n", + " if not llm_response.tool_calls:\n", + " break\n", + "\n", + " # Execute tools\n", + " tool_results = [\n", + " call_tool(tool_call).result() for tool_call in llm_response.tool_calls\n", + " ]\n", + " messages = add_messages(messages, [llm_response, *tool_results])\n", + " llm_response = call_llm(messages).result()\n", + "\n", + " messages = add_messages(messages, llm_response)\n", + " return messages" + ] + }, + { + "cell_type": "code", + "execution_count": 109, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================\u001b[1m Human Message \u001b[0m=================================\n", + "\n", + "Add 3 and 4.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_015EBCekAyhHGVExDSCbSppi', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", + "Tool Calls:\n", + " add (toolu_015EBCekAyhHGVExDSCbSppi)\n", + " Call ID: toolu_015EBCekAyhHGVExDSCbSppi\n", + " Args:\n", + " a: 3\n", + " b: 4\n", + "=================================\u001b[1m Tool Message \u001b[0m=================================\n", + "Name: add\n", + "\n", + "7\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The sum of 3 and 4 is 7.\n" + ] + } + ], + "source": [ + "# Thread ID\n", + "thread_id = str(uuid.uuid4())\n", + "\n", + "# Config\n", + "config = {\"configurable\": {\"thread_id\": thread_id}}\n", + "\n", + "# Run with checkpointer to persist state in memory\n", + "messages = agent.invoke([HumanMessage(content=\"Add 3 and 4.\")], config)\n", + "for m in messages:\n", + " m.pretty_print()\n" + ] + }, + { + "cell_type": "code", + "execution_count": 110, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The result of multiplying 7 by 2 is 14.\n" + ] + } + ], + "source": [ + "# Second turn\n", + "for item in agent.stream([HumanMessage(content=\"Take the result and multiply it by 2.\")], config, stream_mode=\"updates\"):\n", + " if 'agent' in item:\n", + " item['agent'][-1].pretty_print()" + ] + }, + { + "cell_type": "code", + "execution_count": 111, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "{'configurable': {'thread_id': 'ac13180b-31f0-4ff5-8d71-0354addc0a59',\n", + " 'checkpoint_ns': '',\n", + " 'checkpoint_id': '1efddb89-50bb-6c16-8001-4b2a0f971b4e'}}" + ] + }, + "execution_count": 111, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "# Fork and do alternative second turn\n", + "to_fork_from = list(agent.get_state_history(config))[1].config\n", + "to_fork_from" + ] + }, + { + "cell_type": "code", + "execution_count": 112, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================\u001b[1m Human Message \u001b[0m=================================\n", + "\n", + "Add 3 and 4.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_015EBCekAyhHGVExDSCbSppi', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", + "Tool Calls:\n", + " add (toolu_015EBCekAyhHGVExDSCbSppi)\n", + " Call ID: toolu_015EBCekAyhHGVExDSCbSppi\n", + " Args:\n", + " a: 3\n", + " b: 4\n", + "=================================\u001b[1m Tool Message \u001b[0m=================================\n", + "Name: add\n", + "\n", + "7\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The sum of 3 and 4 is 7.\n" + ] + } + ], + "source": [ + "# Get the last checkpoint, which contains the full message history\n", + "agent_state = agent.get_state(to_fork_from)\n", + "for m in agent_state.values:\n", + " m.pretty_print()" + ] + }, + { + "cell_type": "code", + "execution_count": 113, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The result of multiplying 7 by 2 is 14.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The result of multiplying 7 by 2 is 14.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The result of multiplying 7 by 2 is 14.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The result of multiplying 7 by 2 is 14.\n" + ] + } + ], + "source": [ + "# Re-run the workflow from the fork\n", + "for chunk in agent.stream([HumanMessage(content=\"Take the result and multiply it by 3.\")], to_fork_from, stream_mode=\"updates\"):\n", + " if 'agent' in item:\n", + " item['agent'][-1].pretty_print()" + ] + }, + { + "cell_type": "code", + "execution_count": 114, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "================================\u001b[1m Human Message \u001b[0m=================================\n", + "\n", + "Add 3 and 4.\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "[{'text': \"I'll help you add 3 and 4 using the `add` function.\", 'type': 'text'}, {'id': 'toolu_015EBCekAyhHGVExDSCbSppi', 'input': {'a': 3, 'b': 4}, 'name': 'add', 'type': 'tool_use'}]\n", + "Tool Calls:\n", + " add (toolu_015EBCekAyhHGVExDSCbSppi)\n", + " Call ID: toolu_015EBCekAyhHGVExDSCbSppi\n", + " Args:\n", + " a: 3\n", + " b: 4\n", + "=================================\u001b[1m Tool Message \u001b[0m=================================\n", + "Name: add\n", + "\n", + "7\n", + "==================================\u001b[1m Ai Message \u001b[0m==================================\n", + "\n", + "The sum of 3 and 4 is 7.\n" + ] + } + ], + "source": [ + "agent_state = agent.get_state(to_fork_from)\n", + "for m in agent_state.values:\n", + " m.pretty_print()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Agent with HITL and Long-term memory (across threads)\n", "\n", "* Add interrupt to the workflow to allow for HITL\n", "* Add tool for [long-term memory](https://langchain-ai.github.io/langgraph/concepts/memory/#long-term-memory)" @@ -580,7 +799,7 @@ }, { "cell_type": "code", - "execution_count": 164, + "execution_count": 78, "metadata": {}, "outputs": [], "source": [ @@ -628,21 +847,14 @@ }, { "cell_type": "code", - "execution_count": 166, + "execution_count": 83, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "{'tool_call': {'name': 'upsert_memory', 'args': {'content': \"User's name is Lance and they live in San Francisco.\"}, 'id': 'call_4wMyPHYypNRscBdylzW6x3UD', 'type': 'tool_call'}, 'action': 'Please approve/reject the tool call'}\n" - ] - } - ], + "outputs": [], "source": [ "from langgraph.store.memory import InMemoryStore # New \n", "from langchain_core.messages import ToolMessage\n", "\n", + "@task\n", "def call_llm(messages: list[BaseMessage]):\n", " \"\"\"LLM decides whether to call a tool or not\"\"\"\n", " return llm_with_memory_tool.invoke( # New \n", @@ -654,6 +866,7 @@ " + messages\n", " )\n", "\n", + "@task\n", "def call_tool(tool_call: ToolCall, store: BaseStore):\n", "\n", " # Interrupt the workflow to get a review from a human.\n", @@ -680,153 +893,8 @@ " else: \n", " return \"Tool call rejected\"\n", "\n", - "@entrypoint(checkpointer=MemorySaver(), store=InMemoryStore()) \n", - "def agent(messages: list[BaseMessage], previous: list[BaseMessage], store: BaseStore): \n", - " \"\"\" Tool calling agent \"\"\"\n", - "\n", - " # Add previous messages from short-term memory to the current messages\n", - " if previous is not None:\n", - " messages = add_messages(previous, messages)\n", - " \n", - " # New \n", - " # Retrieve the most recent memories for context\n", - " memories = store.search( \n", - " (\"memories\"),\n", - " limit=10,\n", - " )\n", - "\n", - " # New\n", - " # Format memories for inclusion in the prompt\n", - " formatted = \"\\n\".join(f\"[{mem.key}]: {mem.value} (similarity: {mem.score})\" for mem in memories)\n", - " if formatted:\n", - " formatted = f\"\"\"\n", - "\n", - "{formatted}\n", - "\"\"\"\n", - "\n", - " # New\n", - " # Call the LLM\n", - " llm_response = call_llm([SystemMessage(content=f\"Here is some context for you about the user: {formatted}\"), *messages])\n", - "\n", - " while True:\n", - " if not llm_response.tool_calls:\n", - " break\n", - "\n", - " # Execute tools\n", - " tool_results = [\n", - " call_tool(tool_call, store) for tool_call in llm_response.tool_calls\n", - " ]\n", - " messages = add_messages(messages, [llm_response, *tool_results])\n", - " llm_response = call_llm(messages)\n", - "\n", - " messages = add_messages(messages, llm_response)\n", - " return messages\n", - "\n", - "# Thread ID\n", - "thread_id = str(uuid.uuid4())\n", - "\n", - "# Config\n", - "config = {\"configurable\": {\"thread_id\": thread_id}}\n", - "\n", - "# Run until the interrupt \n", - "for item in agent.stream([HumanMessage(content=\"Hi my name is Lance and I live in San Francisco.\")], config, stream_mode=\"updates\"):\n", - " if '__interrupt__' in item:\n", - " print(item['__interrupt__'][0].value)" - ] - }, - { - "cell_type": "code", - "execution_count": 167, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Tool call approved, Memory Added!\n", - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "Nice to meet you, Lance! How can I assist you today?\n" - ] - } - ], - "source": [ - "for item in agent.stream(Command(resume=True), config, stream_mode=\"updates\"):\n", - " item['agent'][-1].pretty_print()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "TODO: Clarify problem w/ *not* using `@task` in the above case!\n", - "\n", - "Seems it still runs once. " - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Adding tasks\n", - "\n", - "* TODO: Why?\n" - ] - }, - { - "cell_type": "code", - "execution_count": 162, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "{'tool_call': {'name': 'upsert_memory', 'args': {'content': \"User's name is Isaac and he lives in Palo Alto.\"}, 'id': 'call_OI9WYghIIgaw6WAlq35KftbA', 'type': 'tool_call'}, 'action': 'Please approve/reject the tool call'}\n" - ] - } - ], - "source": [ - "from langgraph.func import task # New \n", - "\n", - "@task\n", - "def call_llm(messages: list[BaseMessage]):\n", - " \"\"\"LLM decides whether to call a tool or not\"\"\"\n", - " return llm_with_memory_tool.invoke( # New \n", - " [\n", - " SystemMessage(\n", - " content=\"You are a helpful assistant tasked with storing memories.\" # New \n", - " )\n", - " ]\n", - " + messages\n", - " )\n", - "\n", - "@task\n", - "def call_tool(tool_call: ToolCall, store: BaseStore):\n", - "\n", - " # Interrupt the workflow to get a review from a human.\n", - " is_approved = interrupt({ # New \n", - " # Any json-serializable payload provided to interrupt as argument.\n", - " # It will be surfaced on the client side as an Interrupt when streaming data\n", - " # from the workflow.\n", - " \"tool_call\": tool_call, # The tool call we want reviewed.\n", - " # We can add any additional information that we need.\n", - " # For example, introduce a key called \"action\" with some instructions.\n", - " \"action\": \"Please approve/reject the tool call\",\n", - " })\n", - " \n", - " if is_approved:\n", - "\n", - " tool = tools_by_name[tool_call[\"name\"]]\n", - " tool.invoke({**tool_call[\"args\"], \"store\": store})\n", - "\n", - " # Tool message provides confirmation to the model that the actions it took were completed\n", - " results = ToolMessage(content=tool_call[\"args\"][\"content\"], tool_call_id=tool_call[\"id\"])\n", - " return results\n", - " else: \n", - " return \"Tool call rejected\"\n", - "\n", - "@entrypoint(checkpointer=MemorySaver(), store=InMemoryStore()) \n", + "in_memory_store = InMemoryStore()\n", + "@entrypoint(checkpointer=MemorySaver(), store=in_memory_store) \n", "def agent(messages: list[BaseMessage], previous: list[BaseMessage], store: BaseStore): \n", " \"\"\" Tool calling agent \"\"\"\n", "\n", @@ -866,8 +934,23 @@ " llm_response = call_llm(messages).result()\n", "\n", " messages = add_messages(messages, llm_response)\n", - " return messages\n", - "\n", + " return messages" + ] + }, + { + "cell_type": "code", + "execution_count": 84, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "{'tool_call': {'name': 'upsert_memory', 'args': {'content': \"User's name is Lance.\"}, 'id': 'toolu_01Ppw9sRsBnkpmfUPmeHZVyR', 'type': 'tool_call'}, 'action': 'Please approve/reject the tool call'}\n" + ] + } + ], + "source": [ "# Thread ID\n", "thread_id = str(uuid.uuid4())\n", "\n", @@ -875,39 +958,28 @@ "config = {\"configurable\": {\"thread_id\": thread_id}}\n", "\n", "# Run until the interrupt \n", - "for item in agent.stream([HumanMessage(content=\"Hi my name is Isaac and I live in Palo Alto.\")], config, stream_mode=\"updates\"):\n", + "for item in agent.stream([HumanMessage(content=\"Hi my name is Lance.\")], config, stream_mode=\"updates\"):\n", " if '__interrupt__' in item:\n", " print(item['__interrupt__'][0].value)" ] }, { "cell_type": "code", - "execution_count": 163, + "execution_count": 86, "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "==================================\u001b[1m Ai Message \u001b[0m==================================\n", - "\n", - "Hello Isaac! I've noted that you live in Palo Alto. How can I assist you today?\n", - "None\n" - ] - } - ], + "outputs": [], "source": [ "for item in agent.stream(Command(resume=True), config, stream_mode=\"updates\"):\n", " if 'agent' in item:\n", - " print(item['agent'][-1].pretty_print())" + " item['agent'][-1].pretty_print()" ] }, { - "cell_type": "markdown", + "cell_type": "code", + "execution_count": 87, "metadata": {}, - "source": [ - "### Adding Time Travel\n" - ] + "outputs": [], + "source": [] }, { "cell_type": "code",