Merge branch 'main' into david/05-28/support-image-distro-config

This commit is contained in:
Asamu David
2025-06-03 18:18:22 +01:00
committed by GitHub
69 changed files with 1442 additions and 9734 deletions
+11
View File
@@ -0,0 +1,11 @@
# Please see the documentation for all configuration options:
# https://docs.github.com/github/administering-a-repository/configuration-options-for-dependency-updates
# and
# https://docs.github.com/code-security/dependabot/dependabot-version-updates/configuration-options-for-the-dependabot.yml-file
version: 2
updates:
- package-ecosystem: "github-actions"
directory: "/"
schedule:
interval: "weekly"
+2 -2
View File
@@ -49,7 +49,7 @@ jobs:
- name: Get .mypy_cache to speed up mypy
if: steps.changed-files.outputs.all
uses: actions/cache@v3
uses: actions/cache@v4
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "2"
with:
@@ -75,7 +75,7 @@ jobs:
- name: Get .mypy_cache_test to speed up mypy
if: steps.changed-files.outputs.all
uses: actions/cache@v3
uses: actions/cache@v4
env:
SEGMENT_DOWNLOAD_TIMEOUT_MIN: "2"
with:
+4 -4
View File
@@ -166,9 +166,9 @@ jobs:
run:
working-directory: ${{ matrix.working-directory }}
steps:
- uses: actions/checkout@v3
- uses: actions/checkout@v4
- name: Setup Node.js (LTS)
uses: actions/setup-node@v3
uses: actions/setup-node@v4
with:
node-version: "20"
cache: "yarn"
@@ -192,9 +192,9 @@ jobs:
run:
working-directory: ${{ matrix.working-directory }}
steps:
- uses: actions/checkout@v3
- uses: actions/checkout@v4
- name: Setup Node.js (LTS)
uses: actions/setup-node@v3
uses: actions/setup-node@v4
with:
node-version: "20"
cache: "yarn"
+1 -1
View File
@@ -145,7 +145,7 @@ jobs:
- name: Configure GitHub Pages
if: github.ref == 'refs/heads/main'
uses: actions/configure-pages@v4
uses: actions/configure-pages@v5
- name: Upload Pages Artifact
# if: github.ref == 'refs/heads/main'
+1 -1
View File
@@ -22,7 +22,7 @@ jobs:
- uses: actions/checkout@v4
# JS Build
- name: Use Node.js
uses: actions/setup-node@v3
uses: actions/setup-node@v4
with:
node-version: "20"
cache: "yarn"
+1 -1
View File
@@ -109,7 +109,7 @@ Here are some high-level tips on writing a good how-to guide:
LangGraph's conceptual guides fall under the **Explanation** quadrant of Diataxis. They should cover LangChain terms and concepts
in a more abstract way than how-to guides or tutorials, and should be geared towards curious users interested in
gaining a deeper understanding of the framework. Try to avoid excessively large code examples. The goal here is to
impart perspective to the user rather than to finish a practical project. These guides should cover **why** things work they way they do.
impart perspective to the user rather than to finish a practical project. These guides should cover **why** things work the way they do.
To quote the Diataxis website:
+56 -4
View File
@@ -1,9 +1,16 @@
"""mkdocs hooks for adding custom logic to documentation pipeline.
Lifecycle events: https://www.mkdocs.org/dev-guide/plugins/#events
"""
import logging
import os
import posixpath
import re
from typing import Any, Dict
from bs4 import BeautifulSoup
from mkdocs.config.defaults import MkDocsConfig
from mkdocs.structure.files import Files, File
from mkdocs.structure.pages import Page
@@ -101,8 +108,7 @@ REDIRECT_MAP = {
"how-tos/deploy-self-hosted.md": "cloud/deployment/self_hosted_data_plane.md",
"concepts/self_hosted.md": "concepts/langgraph_self_hosted_data_plane.md",
# assistant redirects
"cloud/how-tos/assistant_versioning.md": "cloud/how-tos/configuration_cloud.md"
"cloud/how-tos/assistant_versioning.md": "cloud/how-tos/configuration_cloud.md",
}
@@ -292,7 +298,7 @@ Redirecting...
"""
def write_html(site_dir, old_path, new_path):
def _write_html(site_dir, old_path, new_path):
"""Write an HTML file in the site_dir with a meta redirect to the new page"""
# Determine all relevant paths
old_path_abs = os.path.join(site_dir, old_path)
@@ -308,6 +314,52 @@ def write_html(site_dir, old_path, new_path):
f.write(content)
def _inject_gtm(html: str) -> str:
"""Inject Google Tag Manager code into the HTML.
Code to inject Google Tag Manager noscript tag immediately after <body>.
This is done via hooks rather than via a template because the MkDocs material
theme does not seem to allow placing the code immediately after the <body> tag
without modifying the template files directly.
Args:
html: The HTML content to modify.
Returns:
The modified HTML content with GTM code injected.
"""
# Code was copied from Google Tag Manager setup instructions.
gtm_code = """
<!-- Google Tag Manager (noscript) -->
<noscript><iframe src="https://www.googletagmanager.com/ns.html?id=GTM-T35S4S46"
height="0" width="0" style="display:none;visibility:hidden"></iframe></noscript>
<!-- End Google Tag Manager (noscript) -->
"""
soup = BeautifulSoup(html, "html.parser")
body = soup.body
if body:
# Insert the GTM code as raw HTML at the top of <body>
body.insert(0, BeautifulSoup(gtm_code, "html.parser"))
return str(soup)
else:
return html # fallback if no <body> found
def on_post_page(output: str, page: Page, config: MkDocsConfig) -> str:
"""Inject Google Tag Manager noscript tag immediately after <body>.
Args:
output: The HTML output of the page.
page: The page instance.
config: The MkDocs configuration object.
Returns:
modified HTML output with GTM code injected.
"""
return _inject_gtm(output)
# Create HTML files for redirects after site dir has been built
def on_post_build(config):
use_directory_urls = config.get("use_directory_urls")
@@ -324,4 +376,4 @@ def on_post_build(config):
+ hash
+ suffix
)
write_html(config["site_dir"], old_html_path, new_html_path)
_write_html(config["site_dir"], old_html_path, new_html_path)
+1 -1
View File
@@ -38,7 +38,7 @@ client = MultiServerMCPClient(
"transport": "stdio",
},
"weather": {
# Ensure your start your weather server on port 8000
# Ensure you start your weather server on port 8000
"url": "http://localhost:8000/mcp",
"transport": "streamable_http",
}
+1 -1
View File
@@ -88,7 +88,7 @@ ny_response = agent.invoke(
When the agent is invoked the second time with the same `thread_id`, the original message history from the first conversation is automatically included, allowing the agent to infer that the user is asking specifically about the **weather** in New York.
!!! Note "LangGraph Platform providers a production-ready checkpointer"
!!! Note "LangGraph Platform provides a production-ready checkpointer"
If you're using [LangGraph Platform](./deployment.md), during deployment your checkpointer will be automatically configured to use a production-ready database.
+1 -1
View File
@@ -9,7 +9,7 @@ hide:
# Multi-agent
A single agent might struggle if it needs to specialize in multiple domains or manage many tools. To tackle this, you can break your agent into smaller, independent agents and composing them into a [multi-agent system](../concepts/multi_agent.md).
A single agent might struggle if it needs to specialize in multiple domains or manage many tools. To tackle this, you can break your agent into smaller, independent agents and compose them into a [multi-agent system](../concepts/multi_agent.md).
In multi-agent systems, agents need to communicate between each other. They do so via [handoffs](#handoffs) — a primitive that describes which agent to hand control to and the payload to send to that agent.
+1 -1
View File
@@ -16,4 +16,4 @@ Users can add an array of additional lines to add to the Dockerfile following th
}
```
This would install the system packages required to use Pillow if we were working with `jpeq` or `png` image formats.
This would install the system packages required to use Pillow if we were working with `jpeg` or `png` image formats.
+5 -5
View File
@@ -20,7 +20,7 @@ my-app/
|-- openai_agent.py # code for your graph
```
where the graph is defined in `openai_agent.py`.
where the graph is defined in `openai_agent.py`.
### No rebuild
@@ -28,11 +28,11 @@ In the standard LangGraph API configuration, the server uses the compiled graph
```python
from langchain_openai import ChatOpenAI
from langgraph.graph import END, START, MessageGraph
from langgraph.graph import END, START, StateGraph, MessagesState
model = ChatOpenAI(temperature=0)
graph_workflow = MessageGraph()
graph_workflow = StateGraph(MessagesState)
graph_workflow.add_node("agent", model)
graph_workflow.add_edge("agent", END)
@@ -61,7 +61,7 @@ To make your graph rebuild on each new run with custom configuration, you need t
from typing import Annotated
from typing_extensions import TypedDict
from langchain_openai import ChatOpenAI
from langgraph.graph import END, START, MessageGraph
from langgraph.graph import END, START
from langgraph.graph.state import StateGraph
from langgraph.graph.message import add_messages
from langgraph.prebuilt import ToolNode
@@ -144,4 +144,4 @@ Finally, you need to specify the path to your graph-making function (`make_graph
}
```
See more info on LangGraph API configuration file [here](../reference/cli.md#configuration-file)
See more info on LangGraph API configuration file [here](../reference/cli.md#configuration-file)
@@ -41,7 +41,8 @@ Before deploying, review the [conceptual guide for the Self-Hosted Control Plane
pullPolicy: IfNotPresent
tag: "aa9dff4"
1. In your `values.yaml` file, enable the `langgraphPlatform` option. Note that you must also have a valid ingress setup:
1. In your `langsmith_config.yaml` file, enable the `langgraphPlatform` option. Note that you must also have a valid ingress setup:
config:
langgraphPlatform:
enabled: true
+1 -1
View File
@@ -95,7 +95,7 @@ my-app/
## Define Graphs
Implement your graphs! Graphs can be defined in a single file or multiple files. Make note of the variable names of each [CompiledGraph][langgraph.graph.graph.CompiledGraph] to be included in the LangGraph application. The variable names will be used later when creating the [LangGraph configuration file](../reference/cli.md#configuration-file).
Implement your graphs! Graphs can be defined in a single file or multiple files. Make note of the variable names of each [CompiledStateGraph][langgraph.graph.state.CompiledStateGraph] to be included in the LangGraph application. The variable names will be used later when creating the [LangGraph configuration file](../reference/cli.md#configuration-file).
Example `agent.py` file, which shows how to import from other modules you define (code for the modules is not shown here, please see [this repository](https://github.com/langchain-ai/langgraph-example) to see their implementation):
@@ -108,7 +108,7 @@ my-app/
## Define Graphs
Implement your graphs! Graphs can be defined in a single file or multiple files. Make note of the variable names of each [CompiledGraph][langgraph.graph.graph.CompiledGraph] to be included in the LangGraph application. The variable names will be used later when creating the [LangGraph configuration file](../reference/cli.md#configuration-file).
Implement your graphs! Graphs can be defined in a single file or multiple files. Make note of the variable names of each [CompiledStateGraph][langgraph.graph.state.CompiledStateGraph] to be included in the LangGraph application. The variable names will be used later when creating the [LangGraph configuration file](../reference/cli.md#configuration-file).
Example `agent.py` file, which shows how to import from other modules you define (code for the modules is not shown here, please see [this repository](https://github.com/langchain-ai/langgraph-example-pyproject) to see their implementation):
+113 -4
View File
@@ -1,8 +1,8 @@
# How to integrate LangGraph into your React application
How to integrate LangGraph into your React application# How to integrate LangGraph into your React application
!!! info "Prerequisites"
!!! info "Prerequisites"
- [LangGraph Platform](../../concepts/langgraph_platform.md)
- [LangGraph Platform](../../concepts/langgraph_platform.md)
- [LangGraph Server](../../concepts/langgraph_server.md)
The `useStream()` React hook provides a seamless way to integrate LangGraph into your React applications. It handles all the complexities of streaming, state management, and branching logic, letting you focus on building great chat experiences.
@@ -113,6 +113,115 @@ export default function App() {
}
```
### Resume a stream after page refresh
The `useStream()` hook can automatically resume an ongoing run upon mounting by setting `reconnectOnMount: true`. This is useful for continuing a stream after a page refresh, ensuring no messages and events generated during the downtime are lost.
```tsx
const thread = useStream<{ messages: Message[] }>({
apiUrl: "http://localhost:2024",
assistantId: "agent",
reconnectOnMount: true,
});
```
By default the ID of the created run is stored in `window.sessionStorage`, which can be swapped by passing a custom storage in `reconnectOnMount` instead. The storage is used to persist the in-flight run ID for a thread (under `lg:stream:${threadId}` key).
```tsx
const thread = useStream<{ messages: Message[] }>({
apiUrl: "http://localhost:2024",
assistantId: "agent",
reconnectOnMount: () => window.localStorage,
});
```
You can also manually manage the resuming process by using the run callbacks to persist the run metadata and the `joinStream` function to resume the stream. Make sure to pass `streamResumable: true` when creating the run; otherwise some events might be lost.
````tsx
import type { Message } from "@langchain/langgraph-sdk";
import { useStream } from "@langchain/langgraph-sdk/react";
import { useCallback, useState, useEffect, useRef } from "react";
export default function App() {
const [threadId, onThreadId] = useSearchParam("threadId");
const thread = useStream<{ messages: Message[] }>({
apiUrl: "http://localhost:2024",
assistantId: "agent",
threadId,
onThreadId,
onCreated: (run) => {
window.sessionStorage.setItem(`resume:${run.thread_id}`, run.run_id);
},
onFinish: (_, run) => {
window.sessionStorage.removeItem(`resume:${run?.thread_id}`);
},
});
// Ensure that we only join the stream once per thread.
const joinedThreadId = useRef<string | null>(null);
useEffect(() => {
if (!threadId) return;
const resume = window.sessionStorage.getItem(`resume:${threadId}`);
if (resume && joinedThreadId.current !== threadId) {
thread.joinStream(resume);
joinedThreadId.current = threadId;
}
}, [threadId]);
return (
<form
onSubmit={(e) => {
e.preventDefault();
const form = e.target as HTMLFormElement;
const message = new FormData(form).get("message") as string;
thread.submit(
{ messages: [{ type: "human", content: message }] },
{ streamResumable: true }
);
}}
>
<div>
{thread.messages.map((message) => (
<div key={message.id}>{message.content as string}</div>
))}
</div>
<input type="text" name="message" />
<button type="submit">Send</button>
</form>
);
}
// Utility method to retrieve and persist data in URL as search param
function useSearchParam(key: string) {
const [value, setValue] = useState<string | null>(() => {
const params = new URLSearchParams(window.location.search);
return params.get(key) ?? null;
});
const update = useCallback(
(value: string | null) => {
setValue(value);
const url = new URL(window.location.href);
if (value == null) {
url.searchParams.delete(key);
} else {
url.searchParams.set(key, value);
}
window.history.pushState({}, "", url.toString());
},
[key]
);
return [value, update] as const;
}
```
### Thread Management
Keep track of conversations with built-in thread management. You can access the current thread ID and get notified when new threads are created:
@@ -127,7 +236,7 @@ const thread = useStream<{ messages: Message[] }>({
threadId: threadId,
onThreadId: setThreadId,
});
```
````
We recommend storing the `threadId` in your URL's query parameters to let users resume conversations after page refreshes.
+1 -1
View File
@@ -32,7 +32,7 @@ Below are examples of directory structures for Python and JavaScript application
│ ├── utils # utilities for your graph
│ │ ├── __init__.py
│ │ ├── tools.py # tools for your graph
│ │ ├── nodes.py # node functions for you graph
│ │ ├── nodes.py # node functions for your graph
│ │ └── state.py # state definition of your graph
│ ├── __init__.py
│ └── agent.py # code for constructing your graph
+1 -1
View File
@@ -26,4 +26,4 @@ Once you've created an assistant, subsequent edits to that assistant will create
## Learn more
* The LangGraph Cloud API provides several endpoints for creating and managing assistants their versions. See the [API reference](../cloud/reference/api/api_ref.html#tag/assistants) for more details.
* The LangGraph Cloud API provides several endpoints for creating and managing assistants and their versions. See the [API reference](../cloud/reference/api/api_ref.html#tag/assistants) for more details.
+1 -1
View File
@@ -59,7 +59,7 @@ For more information, please see:
!!! info "Important"
The Self-Hosted Control Plane deployment option is currently in beta stage and requires an [Enterprise](../concepts/plans.md) plan.
The [Self-Hosted Control Plane](./langgraph_self_hosted_control_plane.md) deployment option is a fully self-hosted model for deployment where you manage the [control plane](./langgraph_control_plane.md) and [data plane](./langgraph_data_plane.md) in your cloud. This option give you full control and responsibility of the control plane and data plane infrastructure.
The [Self-Hosted Control Plane](./langgraph_self_hosted_control_plane.md) deployment option is a fully self-hosted model for deployment where you manage the [control plane](./langgraph_control_plane.md) and [data plane](./langgraph_data_plane.md) in your cloud. This option gives you full control and responsibility of the control plane and data plane infrastructure.
Build a Docker image using the [LangGraph CLI](./langgraph_cli.md) and deploy your LangGraph Server from the [control plane UI](./langgraph_control_plane.md#control-plane-ui).
+2 -2
View File
@@ -59,8 +59,8 @@ Yes! You can use LangGraph with any LLMs. The main reason we use LLMs that suppo
Yes! LangGraph is totally ambivalent to what LLMs are used under the hood. The main reason we use closed LLMs in most of the tutorials is that they seamlessly support tool calling, while OSS LLMs often don't. But tool calling is not necessary (see [this section](#does-langgraph-work-with-llms-that-dont-support-tool-calling)) so you can totally use LangGraph with OSS LLMs.
## Can I use LangGraph Studio without logging to LangSmith
## Can I use LangGraph Studio without logging in to LangSmith
Yes! You can use the [development version of LangGraph Server](../tutorials/langgraph-platform/local-server.md) to run the backend locally.
This will connect to the studio frontend hosted as part of LangSmith.
If you set an environment variable of `LANGSMITH_TRACING=false` then no traces will be sent to LangSmith.
If you set an environment variable of `LANGSMITH_TRACING=false`, then no traces will be sent to LangSmith.
+3 -3
View File
@@ -9,7 +9,7 @@ The term "data plane" is used broadly to refer to [LangGraph Servers](./langgrap
## Server Infrastructure
In addition to the [LangGraph Server](./langgraph_server.md) itself, the following infrastructure for each server are also included in the broad definition of "data plane":
In addition to the [LangGraph Server](./langgraph_server.md) itself, the following infrastructure components for each server are also included in the broad definition of "data plane":
- Postgres
- Redis
@@ -44,7 +44,7 @@ All runs in a LangGraph Server are executed by a pool of background workers that
### Ephemeral metadata
Runs in a LangGraph Server may be retried for specific failures (currently only for transient Postgres errors encountered during the run). In order to limit the number of retries (currently limited to 3 attempts per run) we record the attempt number in a Redis string when is picked up. This contains no run-specific info other than its ID, and expires after a short delay.
Runs in a LangGraph Server may be retried for specific failures (currently only for transient Postgres errors encountered during the run). In order to limit the number of retries (currently limited to 3 attempts per run) we record the attempt number in a Redis string when it is picked up. This contains no run-specific info other than its ID, and expires after a short delay.
## Data Plane Features
@@ -62,7 +62,7 @@ For CPU utilization, the autoscaler targets 75% utilization. This means the auto
For number of pending runs, the autoscaler targets 10 pending runs. For example, if the current number of containers is 1, but the number of pending runs in 20, the autoscaler will scale up the deployment to 2 containers (20 pending runs / 2 containers = 10 pending runs per container).
Each metric is computed independently and the autoscaler will determine the scaling action based on the metric that results in the most number of containers.
Each metric is computed independently and the autoscaler will determine the scaling action based on the metric that results in the largest number of containers.
Scale down actions are delayed for 30 minutes before any action is taken. In other words, if the autoscaler decides to scale down a deployment, it will first wait for 30 minutes before scaling down. After 30 minutes, the metrics are recomputed and the deployment will scale down if the recomputed metrics result in a lower number of containers than the current number. Otherwise, the deployment remains scaled up. This "cool down" period ensures that deployments do not scale up and down too frequently.
+2 -2
View File
@@ -9,7 +9,7 @@ Develop, deploy, scale, and manage agents with **LangGraph Platform** — the pu
!!! tip "Get started with LangGraph Platform"
Check out the [LangGraph Platform quickstart](../tutorials/langgraph-platform/local-server.md) for instructions on how to use LangGraph Platform run a LangGraph application locally.
Check out the [LangGraph Platform quickstart](../tutorials/langgraph-platform/local-server.md) for instructions on how to use LangGraph Platform to run a LangGraph application locally.
## Why use LangGraph Platform?
@@ -33,4 +33,4 @@ LangGraph Platform makes it easy to get your agent running in production — wh
- **[LangGraph Studio](./langgraph_studio.md)**: Enables visualization, interaction, and debugging of agentic systems that implement the LangGraph Server API protocol. Studio also integrates with LangSmith to enable tracing, evaluation, and prompt engineering.
- **[Deployment](./deployment_options.md)**: There are four ways to deploy on LangGraph Platform: [Cloud Saas](../concepts/langgraph_cloud.md), [Self-Hosted Data Plane](../concepts/langgraph_self_hosted_data_plane.md), [Self-Hosted Control Plane](../concepts/langgraph_self_hosted_control_plane.md), and [Standalone Container](../concepts/langgraph_standalone_container.md).
- **[Deployment](./deployment_options.md)**: There are four ways to deploy on LangGraph Platform: [Cloud SaaS](../concepts/langgraph_cloud.md), [Self-Hosted Data Plane](../concepts/langgraph_self_hosted_data_plane.md), [Self-Hosted Control Plane](../concepts/langgraph_self_hosted_control_plane.md), and [Standalone Container](../concepts/langgraph_standalone_container.md).
@@ -9,10 +9,12 @@ There are two versions of the self-hosted deployment: [Self-Hosted Data Plane](.
- You use `langgraph-cli` and/or [LangGraph Studio](./langgraph_studio.md) app to test graph locally.
- You use `langgraph build` command to build image.
- You have a Self-Hosted LangSmith instance deployed.
- You are using Ingress for your LangSmith instance. All agents will be deployed as Kubernetes services behind this ingress.
## Self-Hosted Control Plane
The [Self-Hosted Control Plane](./langgraph_self_hosted_control_plane.md) deployment option is a fully self-hosted model for deployment where you manage the [control plane](./langgraph_control_plane.md) and [data plane](./langgraph_data_plane.md) in your cloud. This option give you full control and responsibility of the control plane and data plane infrastructure.
The [Self-Hosted Control Plane](./langgraph_self_hosted_control_plane.md) deployment option is a fully self-hosted model for deployment where you manage the [control plane](./langgraph_control_plane.md) and [data plane](./langgraph_data_plane.md) in your cloud. This option gives you full control and responsibility of the control plane and data plane infrastructure.
| | [Control plane](../concepts/langgraph_control_plane.md) | [Data plane](../concepts/langgraph_data_plane.md) |
|-------------------|-------------------|------------|
@@ -29,4 +31,4 @@ The [Self-Hosted Control Plane](./langgraph_self_hosted_control_plane.md) deploy
- **Kubernetes**: The Self-Hosted Control Plane deployment option supports deploying control plane and data plane infrastructure to any Kubernetes cluster.
!!! tip
If you would like to deploy to Kubernetes, you can use this [Helm chart](https://github.com/langchain-ai/helm/blob/main/charts/langgraph-cloud/README.md).
If you would like to enable this on your LangSmith instance, please follow the [Self-Hosted Control Plane deployment guide](../deployment/self_hosted_control_plane.md).
+1 -1
View File
@@ -26,7 +26,7 @@ Feature Differences:
|-------|------------|------------|
| [Cron Jobs](../cloud/concepts/cron_jobs.md) |❌|✅|
| [Custom Authentication](../concepts/auth.md) |❌|✅|
| [Deployment options](../concepts/deployment_options.md) | Standalone container | Cloud Saas, Self-Hosted Data Plane, Self-Hosted Control Plane, Standalone container
| [Deployment options](../concepts/deployment_options.md) | Standalone container | Cloud SaaS, Self-Hosted Data Plane, Self-Hosted Control Plane, Standalone container
## Application structure
+2 -2
View File
@@ -21,7 +21,7 @@ Key features of LangGraph Studio:
- Visualize your graph architecture
- [Run and interact with your agent](../cloud/how-tos/invoke_studio.md)
- [Manage assistants](../cloud/how-tos/studio/manage_assistants.md.md)
- [Manage assistants](../cloud/how-tos/studio/manage_assistants.md)
- [Manage threads](../cloud/how-tos/threads_studio.md)
- [Iterate on prompts](../cloud/how-tos/iterate_graph_studio.md)
- Manage [long term memory](memory.md)
@@ -33,7 +33,7 @@ Studio supports two modes:
### Graph mode
Graph mode exposes the full feature-set of Studio and is useful when you would like as many details about the execution of your agent, including the nodes traversed, intermediate states, and LangSmith integrations (such as adding to datasets an playground).
Graph mode exposes the full feature-set of Studio and is useful when you would like as many details about the execution of your agent, including the nodes traversed, intermediate states, and LangSmith integrations (such as adding to datasets and playground).
### Chat mode
+2 -2
View File
@@ -105,7 +105,7 @@ graph.invoke({"user_input":"My"})
There are two subtle and important points to note here:
1. We pass `state: InputState` as the input schema to `node_1`. But, we write out to `foo`, a channel in `OverallState`. How can we write out to a state channel that is not included in the input schema? This is because a node _can write to any state channel in the graph state._ The graph state is the union of of the state channels defined at initialization, which includes `OverallState` and the filters `InputState` and `OutputState`.
1. We pass `state: InputState` as the input schema to `node_1`. But, we write out to `foo`, a channel in `OverallState`. How can we write out to a state channel that is not included in the input schema? This is because a node _can write to any state channel in the graph state._ The graph state is the union of the state channels defined at initialization, which includes `OverallState` and the filters `InputState` and `OutputState`.
2. We initialize the graph with `StateGraph(OverallState,input=InputState,output=OutputState)`. So, how can we write to `PrivateState` in `node_2`? How does the graph gain access to this schema if it was not passed in the `StateGraph` initialization? We can do this because _nodes can also declare additional state channels_ as long as the state schema definition exists. In this case, the `PrivateState` schema is defined, so we can add `bar` as a new state channel in the graph and write to it.
@@ -167,7 +167,7 @@ In addition to keeping track of message IDs, the `add_messages` function will al
{"messages": [{"type": "human", "content": "message"}]}
```
Since the state updates are always deserialized into LangChain `Messages` when using `add_messages`, you should use dot notation to access message attributes, like `state["messages"][-1].content`. Below is an example of a graph that uses `add_messages` as it's reducer function.
Since the state updates are always deserialized into LangChain `Messages` when using `add_messages`, you should use dot notation to access message attributes, like `state["messages"][-1].content`. Below is an example of a graph that uses `add_messages` as its reducer function.
```python
from langchain_core.messages import AnyMessage
@@ -12,9 +12,9 @@
"\n",
"\n",
"1. **Run the graph** with initial inputs using `invoke` or `stream` APIs.\n",
"2. **Identify a checkpoint in an existing thread**: Use the [`get_state_history()`][langgraph.graph.graph.CompiledGraph.get_state_history] method to retrieve the execution history for a specific `thread_id` and locate the desired `checkpoint_id`. \n",
"2. **Identify a checkpoint in an existing thread**: Use the [`get_state_history()`][langgraph.graph.state.CompiledStateGraph.get_state_history] method to retrieve the execution history for a specific `thread_id` and locate the desired `checkpoint_id`. \n",
" Alternatively, set a [breakpoint](../../../concepts/breakpoints/) before the node(s) where you want execution to pause. You can then find the most recent checkpoint recorded up to that breakpoint.\n",
"3. **(Optional) modify the graph state**: Use the [`update_state`][langgraph.graph.graph.CompiledGraph.update_state] method to modify the graph’s state at the checkpoint and resume execution from alternative state.\n",
"3. **(Optional) modify the graph state**: Use the [`update_state`][langgraph.graph.state.CompiledStateGraph.update_state] method to modify the graph’s state at the checkpoint and resume execution from alternative state.\n",
"4. **Resume execution from the checkpoint**: Use the `invoke` or `stream` APIs with an input of `None` and a configuration containing the appropriate `thread_id` and `checkpoint_id`.\n",
"\n",
"## Example\n",
-35
View File
@@ -36,41 +36,6 @@
- aget_subgraphs
- with_config
::: langgraph.graph.graph.Graph
options:
show_if_no_docstring: true
show_root_heading: true
show_root_full_path: false
members:
- add_node
- add_edge
- add_conditional_edges
- compile
::: langgraph.graph.graph.CompiledGraph
options:
show_if_no_docstring: true
show_root_heading: true
show_root_full_path: false
members:
- stream
- astream
- invoke
- ainvoke
- get_state
- aget_state
- get_state_history
- aget_state_history
- update_state
- aupdate_state
- bulk_update_state
- abulk_update_state
- get_graph
- aget_graph
- get_subgraphs
- aget_subgraphs
- with_config
::: langgraph.graph.message
options:
members:
+3 -3
View File
@@ -89,7 +89,7 @@
},
{
"cell_type": "code",
"execution_count": 2,
"execution_count": null,
"id": "baf669a0-04ee-492d-80d8-8fcb658ed128",
"metadata": {},
"outputs": [],
@@ -313,8 +313,8 @@
"\n",
" builder.add_edge(\"finalizer\", END)\n",
"\n",
" # These functions let the step be used in a MessageGraph\n",
" # or a StateGraph with 'messages' as the key.\n",
" # These functions let the step be used in a\n",
" # StateGraph with 'messages' as the key.\n",
" def encode(x: Union[Sequence[AnyMessage], PromptValue]) -> dict:\n",
" \"\"\"Ensure the input is the correct format.\"\"\"\n",
" if isinstance(x, PromptValue):\n",
@@ -1,6 +1,6 @@
# Add tools
To handle queries you chatbot can't answer "from memory", integrate a web search tool. The chatbot can use this tool to find relevant information and provide better responses.
To handle queries that your chatbot can't answer "from memory", integrate a web search tool. The chatbot can use this tool to find relevant information and provide better responses.
!!! note
@@ -516,11 +516,13 @@
},
{
"cell_type": "code",
"execution_count": 12,
"execution_count": null,
"id": "b76b5ec3-0720-443d-85b1-c0e79659ca0a",
"metadata": {},
"outputs": [],
"source": [
"from pprint import pprint\n",
"\n",
"from langchain.schema import Document\n",
"\n",
"\n",
@@ -796,7 +798,7 @@
},
{
"cell_type": "code",
"execution_count": 14,
"execution_count": null,
"id": "29acc541-d726-4b75-84d1-a215845fe88a",
"metadata": {},
"outputs": [
@@ -823,8 +825,6 @@
}
],
"source": [
"from pprint import pprint\n",
"\n",
"# Run\n",
"inputs = {\n",
" \"question\": \"What player at the Bears expected to draft first in the 2024 NFL draft?\"\n",
+11 -2
View File
@@ -1,5 +1,16 @@
{% extends "base.html" %}
{% block analytics %}
<!-- Google Tag Manager -->
<script>(function(w,d,s,l,i){w[l]=w[l]||[];w[l].push({'gtm.start':
new Date().getTime(),event:'gtm.js'});var f=d.getElementsByTagName(s)[0],
j=d.createElement(s),dl=l!='dataLayer'?'&l='+l:'';j.async=true;j.src=
'https://www.googletagmanager.com/gtm.js?id='+i+dl;f.parentNode.insertBefore(j,f);
})(window,document,'script','dataLayer','GTM-T35S4S46');</script>
<!-- End Google Tag Manager -->
{% endblock %}
{% block extrahead %}
<meta name="algolia-site-verification" content="165B7E7C89E49946" />
<style>
@@ -185,7 +196,6 @@
</style>
{% endblock %}
{% block content %}
<div class="notebook-links">
{% if page.nb_url %}
@@ -209,7 +219,6 @@
{% endif %}
{% endblock %}
{% block announce %}
<strong>We are growing and hiring for multiple roles for LangChain, LangGraph and LangSmith. <a href="https://www.langchain.com/careers" target="_blank" rel="noopener noreferrer"> Join our team!</a></strong>
{% endblock %}
+1 -2
View File
@@ -44,9 +44,8 @@ stop-postgres:
docker compose -f tests/compose-postgres.yml down -v
start-dev-server:
LOG_LEVEL=warning uv run langgraph dev --config tests/example_app/langgraph.json --no-browser &
LOG_LEVEL=warning uv run langgraph dev --config tests/example_app/langgraph.json --no-browser & echo "$$!" > .devserver.pid
@echo "Dev server started."
@echo "Dev server PID: $$!" > .devserver.pid
stop-dev-server:
@if [ -f .devserver.pid ]; then \
+6 -2
View File
@@ -21,6 +21,8 @@ def fanout_to_subgraph() -> StateGraph:
class JokeOutput(TypedDict):
jokes: list[str]
class JokeState(JokeInput, JokeOutput): ...
async def bump(state: JokeOutput):
return {"jokes": [state["jokes"][0] + " a"]}
@@ -35,7 +37,7 @@ def fanout_to_subgraph() -> StateGraph:
return END if state["jokes"][0].endswith(" a" * 10) else "bump"
# subgraph
subgraph = StateGraph(input=JokeInput, output=JokeOutput)
subgraph = StateGraph(JokeState, input=JokeInput, output=JokeOutput)
subgraph.add_node("edit", edit)
subgraph.add_node("generate", generate)
subgraph.add_node("bump", bump)
@@ -69,6 +71,8 @@ def fanout_to_subgraph_sync() -> StateGraph:
class JokeOutput(TypedDict):
jokes: list[str]
class JokeState(JokeInput, JokeOutput): ...
def bump(state: JokeOutput):
return {"jokes": [state["jokes"][0] + " a"]}
@@ -83,7 +87,7 @@ def fanout_to_subgraph_sync() -> StateGraph:
return END if state["jokes"][0].endswith(" a" * 10) else "bump"
# subgraph
subgraph = StateGraph(input=JokeInput, output=JokeOutput)
subgraph = StateGraph(JokeState, input=JokeInput, output=JokeOutput)
subgraph.add_node("edit", edit)
subgraph.add_node("generate", generate)
subgraph.add_node("bump", bump)
@@ -1,84 +0,0 @@
import functools
import warnings
from typing import Any, Callable, TypeVar, Union, cast
class LangGraphDeprecationWarning(DeprecationWarning):
pass
F = TypeVar("F", bound=Callable[..., Any])
C = TypeVar("C", bound=type[Any])
def deprecated(
since: str, alternative: str, *, removal: str = "", example: str = ""
) -> Callable[[F], F]:
def decorator(obj: Union[F, C]) -> Union[F, C]:
removal_str = removal if removal else "a future version"
message = (
f"{obj.__name__} is deprecated as of version {since} and will be"
f" removed in {removal_str}. Use {alternative} instead.{example}"
)
if isinstance(obj, type):
original_init = obj.__init__ # type: ignore[misc]
@functools.wraps(original_init)
def new_init(self, *args: Any, **kwargs: Any) -> None: # type: ignore[no-untyped-def]
warnings.warn(message, LangGraphDeprecationWarning, stacklevel=2)
original_init(self, *args, **kwargs)
obj.__init__ = new_init # type: ignore[misc]
docstring = (
f"**Deprecated**: This class is deprecated as of version {since}. "
f"Use `{alternative}` instead."
)
if obj.__doc__:
docstring = docstring + f"\n\n{obj.__doc__}"
obj.__doc__ = docstring
return cast(C, obj)
elif callable(obj):
@functools.wraps(obj)
def wrapper(*args: Any, **kwargs: Any) -> Any:
warnings.warn(message, LangGraphDeprecationWarning, stacklevel=2)
return obj(*args, **kwargs)
docstring = (
f"**Deprecated**: This function is deprecated as of version {since}. "
f"Use `{alternative}` instead."
)
if obj.__doc__:
docstring = docstring + f"\n\n{obj.__doc__}"
wrapper.__doc__ = docstring
return cast(F, wrapper)
else:
raise TypeError(
f"Can only add deprecation decorator to classes or callables, got '{type(obj)}' instead."
)
return decorator
def deprecated_parameter(
arg_name: str, since: str, alternative: str, *, removal: str
) -> Callable[[F], F]:
def decorator(func: F) -> F:
@functools.wraps(func)
def wrapper(*args, **kwargs): # type: ignore[no-untyped-def]
if arg_name in kwargs:
warnings.warn(
f"Parameter '{arg_name}' in function '{func.__name__}' is "
f"deprecated as of version {since} and will be removed in version {removal}. "
f"Use '{alternative}' parameter instead.",
category=LangGraphDeprecationWarning,
stacklevel=2,
)
return func(*args, **kwargs)
return cast(F, wrapper)
return decorator
@@ -1,15 +1,14 @@
from langgraph.channels.any_value import AnyValue
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue
from langgraph.channels.last_value import LastValue, LastValueAfterFinish
from langgraph.channels.topic import Topic
from langgraph.channels.untracked_value import UntrackedValue
__all__ = [
"LastValue",
"LastValueAfterFinish",
"Topic",
"BinaryOperatorAggregate",
"UntrackedValue",
"EphemeralValue",
"AnyValue",
]
@@ -1,206 +0,0 @@
from collections.abc import Sequence, Set
from typing import Any, Generic, NamedTuple, Optional, Union
from typing_extensions import Self
from langgraph.channels.base import BaseChannel, Value
from langgraph.constants import MISSING
from langgraph.errors import EmptyChannelError, InvalidUpdateError
class WaitForNames(NamedTuple):
names: Set[Any]
class DynamicBarrierValue(
Generic[Value], BaseChannel[Value, Union[Value, WaitForNames], Set[Value]]
):
"""A channel that switches between two states
- in the "priming" state it can't be read from.
- if it receives a WaitForNames update, it switches to the "waiting" state.
- in the "waiting" state it collects named values until all are received.
- once all named values are received, it can be read once, and it switches
back to the "priming" state.
"""
__slots__ = ("names", "seen")
names: Optional[Set[Value]]
seen: set[Value]
def __init__(self, typ: type[Value]) -> None:
super().__init__(typ)
self.names = None
self.seen = set()
def __eq__(self, value: object) -> bool:
return isinstance(value, DynamicBarrierValue) and value.names == self.names
@property
def ValueType(self) -> type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> type[Value]:
"""The type of the update received by the channel."""
return self.typ
def copy(self) -> Self:
"""Return a copy of the channel."""
empty = self.__class__(self.typ)
empty.key = self.key
empty.names = self.names
empty.seen = self.seen.copy()
return empty
def checkpoint(self) -> tuple[Optional[Set[Value]], set[Value]]:
return (self.names, self.seen)
def from_checkpoint(
self, checkpoint: tuple[Optional[Set[Value]], set[Value]]
) -> Self:
empty = self.__class__(self.typ)
empty.key = self.key
if checkpoint is not MISSING:
names, seen = checkpoint
empty.names = names if names is not None else None
empty.seen = seen
return empty
def update(self, values: Sequence[Union[Value, WaitForNames]]) -> bool:
if wait_for_names := [v for v in values if isinstance(v, WaitForNames)]:
if len(wait_for_names) > 1:
raise InvalidUpdateError(
f"At key '{self.key}': Received multiple WaitForNames updates in the same step."
)
self.names = wait_for_names[0].names
return True
elif self.names is not None:
updated = False
for value in values:
assert not isinstance(value, WaitForNames)
if value in self.names and value not in self.seen:
self.seen.add(value)
updated = True
return updated
def get(self) -> Value:
if self.seen != self.names:
raise EmptyChannelError()
return None
def is_available(self) -> bool:
return self.seen == self.names
def consume(self) -> bool:
if self.seen == self.names:
self.seen = set()
self.names = None
return True
return False
class DynamicBarrierValueAfterFinish(
Generic[Value], BaseChannel[Value, Union[Value, WaitForNames], Set[Value]]
):
"""A channel that switches between two states
- in the "priming" state it can't be read from.
- if it receives a WaitForNames update, it switches to the "waiting" state.
- in the "waiting" state it collects named values until all are received.
- once all named values are received, and the finished flag is set, it can be read once, and it switches
back to the "priming" state.
"""
__slots__ = ("names", "seen", "finished")
names: Optional[Set[Value]]
seen: set[Value]
finished: bool
def __init__(self, typ: type[Value]) -> None:
super().__init__(typ)
self.names = None
self.seen = set()
self.finished = False
def __eq__(self, value: object) -> bool:
return (
isinstance(value, DynamicBarrierValueAfterFinish)
and value.names == self.names
)
@property
def ValueType(self) -> type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> type[Value]:
"""The type of the update received by the channel."""
return self.typ
def copy(self) -> Self:
"""Return a copy of the channel."""
empty = self.__class__(self.typ)
empty.key = self.key
empty.names = self.names
empty.seen = self.seen.copy()
empty.finished = self.finished
return empty
def checkpoint(self) -> tuple[Optional[Set[Value]], set[Value], bool]:
return (self.names, self.seen, self.finished)
def from_checkpoint(
self, checkpoint: tuple[Optional[Set[Value]], set[Value], bool]
) -> Self:
empty = self.__class__(self.typ)
empty.key = self.key
if checkpoint is not MISSING:
names, seen, finished = checkpoint
empty.names = names if names is not None else None
empty.seen = seen
empty.finished = finished
return empty
def update(self, values: Sequence[Union[Value, WaitForNames]]) -> bool:
if wait_for_names := [v for v in values if isinstance(v, WaitForNames)]:
if len(wait_for_names) > 1:
raise InvalidUpdateError(
f"At key '{self.key}': Received multiple WaitForNames updates in the same step."
)
self.names = wait_for_names[0].names
return True
elif self.names is not None:
updated = False
for value in values:
assert not isinstance(value, WaitForNames)
if value in self.names and value not in self.seen:
self.seen.add(value)
updated = True
return updated
def get(self) -> Value:
if not self.finished and self.seen != self.names:
raise EmptyChannelError()
return None
def is_available(self) -> bool:
return self.seen == self.names and self.finished
def consume(self) -> bool:
if self.finished and self.seen == self.names:
self.seen = set()
self.names = None
return True
return False
def finish(self) -> bool:
if not self.finished and self.seen == self.names:
self.finished = True
return True
else:
return False
@@ -1,66 +0,0 @@
from collections.abc import Sequence
from typing import Generic
from typing_extensions import Self
from langgraph.channels.base import BaseChannel, Value
from langgraph.constants import MISSING
from langgraph.errors import EmptyChannelError, InvalidUpdateError
class UntrackedValue(Generic[Value], BaseChannel[Value, Value, Value]):
"""Stores the last value received, never checkpointed."""
__slots__ = ("value", "guard")
def __init__(self, typ: type[Value], guard: bool = True) -> None:
super().__init__(typ)
self.guard = guard
self.value = MISSING
def __eq__(self, value: object) -> bool:
return isinstance(value, UntrackedValue) and value.guard == self.guard
@property
def ValueType(self) -> type[Value]:
"""The type of the value stored in the channel."""
return self.typ
@property
def UpdateType(self) -> type[Value]:
"""The type of the update received by the channel."""
return self.typ
def copy(self) -> Self:
"""Return a copy of the channel."""
empty = self.__class__(self.typ, self.guard)
empty.key = self.key
empty.value = self.value
return empty
def checkpoint(self) -> Value:
return MISSING
def from_checkpoint(self, checkpoint: Value) -> Self:
empty = self.__class__(self.typ, self.guard)
empty.key = self.key
return empty
def update(self, values: Sequence[Value]) -> bool:
if len(values) == 0:
return False
if len(values) != 1 and self.guard:
raise InvalidUpdateError(
f"At key '{self.key}': UntrackedValue(guard=True) can receive only one value per step. Use guard=False if you want to store any one of multiple values."
)
self.value = values[-1]
return True
def get(self) -> Value:
if self.value is MISSING:
raise EmptyChannelError()
return self.value
def is_available(self) -> bool:
return self.value is not MISSING
+2 -4
View File
@@ -1,13 +1,11 @@
from langgraph.graph.graph import END, START, Graph
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
from langgraph.constants import END, START
from langgraph.graph.message import MessagesState, add_messages
from langgraph.graph.state import StateGraph
__all__ = [
"END",
"START",
"Graph",
"StateGraph",
"MessageGraph",
"add_messages",
"MessagesState",
]
+1 -3
View File
@@ -87,7 +87,6 @@ def _get_branch_path_input_schema(
class Branch(NamedTuple):
path: Runnable[Any, Union[Hashable, list[Hashable]]]
ends: Optional[dict[Hashable, str]]
then: Optional[str] = None
input_schema: Optional[type[Any]] = None
@classmethod
@@ -95,7 +94,6 @@ class Branch(NamedTuple):
cls,
path: Runnable[Any, Union[Hashable, list[Hashable]]],
path_map: Optional[Union[dict[Hashable, str], list[str]]],
then: Optional[str] = None,
infer_schema: bool = False,
) -> "Branch":
# coerce path_map to a dictionary
@@ -123,7 +121,7 @@ class Branch(NamedTuple):
# infer input schema
input_schema = _get_branch_path_input_schema(path) if infer_schema else None
# create branch
return cls(path=path, ends=path_map_, then=then, input_schema=input_schema)
return cls(path=path, ends=path_map_, input_schema=input_schema)
def run(
self,
-445
View File
@@ -1,445 +0,0 @@
import logging
from collections import defaultdict
from collections.abc import Awaitable, Hashable, Sequence
from typing import (
Any,
Callable,
NamedTuple,
Optional,
Union,
cast,
overload,
)
from langchain_core.runnables import Runnable
from typing_extensions import Self
from langgraph.cache.base import BaseCache
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.constants import (
EMPTY_SEQ,
END,
NS_END,
NS_SEP,
START,
TAG_HIDDEN,
Send,
)
from langgraph.graph.branch import Branch
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.pregel.read import PregelNode
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
from langgraph.store.base import BaseStore
from langgraph.types import All, Checkpointer
from langgraph.utils.runnable import RunnableLike, coerce_to_runnable
logger = logging.getLogger(__name__)
class NodeSpec(NamedTuple):
runnable: Runnable
metadata: Optional[dict[str, Any]] = None
ends: Optional[Union[tuple[str, ...], dict[str, str]]] = EMPTY_SEQ
class Graph:
def __init__(self) -> None:
self.nodes: dict[str, NodeSpec] = {}
self.edges = set[tuple[str, str]]()
self.branches: defaultdict[str, dict[str, Branch]] = defaultdict(dict)
self.support_multiple_edges = False
self.compiled = False
@property
def _all_edges(self) -> set[tuple[str, str]]:
return self.edges
@overload
def add_node(
self,
node: RunnableLike,
*,
metadata: Optional[dict[str, Any]] = None,
) -> Self: ...
@overload
def add_node(
self,
node: str,
action: RunnableLike,
*,
metadata: Optional[dict[str, Any]] = None,
) -> Self: ...
def add_node(
self,
node: Union[str, RunnableLike],
action: Optional[RunnableLike] = None,
*,
metadata: Optional[dict[str, Any]] = None,
) -> Self:
"""Add a new node to the graph.
Args:
node: The function or runnable this node will run.
If a string is provided, it will be used as the node name, and action will be used as the function or runnable.
action: The action associated with the node. (default: None)
Will be used as the node function or runnable if `node` is a string (node name).
metadata: The metadata associated with the node. (default: None)
"""
if isinstance(node, str):
for character in (NS_SEP, NS_END):
if character in node:
raise ValueError(
f"'{character}' is a reserved character and is not allowed in the node names."
)
if self.compiled:
logger.warning(
"Adding a node to a graph that has already been compiled. This will "
"not be reflected in the compiled graph."
)
if not isinstance(node, str):
action = node
node = getattr(action, "name", getattr(action, "__name__"))
if node is None:
raise ValueError(
"Node name must be provided if action is not a function"
)
if action is None:
raise RuntimeError(
"Expected a function or Runnable action in add_node. Received None."
)
if node in self.nodes:
raise ValueError(f"Node `{node}` already present.")
if node == END or node == START:
raise ValueError(f"Node `{node}` is reserved.")
self.nodes[cast(str, node)] = NodeSpec(
coerce_to_runnable(action, name=cast(str, node), trace=False), metadata
)
return self
def add_edge(self, start_key: str, end_key: str) -> Self:
"""Add a directed edge from the start node to the end node.
Args:
start_key: The key of the start node of the edge.
end_key: The key of the end node of the edge.
"""
if self.compiled:
logger.warning(
"Adding an edge to a graph that has already been compiled. This will "
"not be reflected in the compiled graph."
)
if start_key == END:
raise ValueError("END cannot be a start node")
if end_key == START:
raise ValueError("START cannot be an end node")
# run this validation only for non-StateGraph graphs
if not hasattr(self, "channels") and start_key in set(
start for start, _ in self.edges
):
raise ValueError(
f"Already found path for node '{start_key}'.\n"
"For multiple edges, use StateGraph with an Annotated state key."
)
self.edges.add((start_key, end_key))
return self
def add_conditional_edges(
self,
source: str,
path: Union[
Callable[..., Union[Hashable, list[Hashable]]],
Callable[..., Awaitable[Union[Hashable, list[Hashable]]]],
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
then: Optional[str] = None,
) -> Self:
"""Add a conditional edge from the starting node to any number of destination nodes.
Args:
source: The starting node. This conditional edge will run when
exiting this node.
path: The callable that determines the next
node or nodes. If not specifying `path_map` it should return one or
more nodes. If it returns END, the graph will stop execution.
path_map: Optional mapping of paths to node
names. If omitted the paths returned by `path` should be node names.
then: The name of a node to execute after the nodes
selected by `path`.
Returns:
Self: The instance of the graph, allowing for method chaining.
Note: Without typehints on the `path` function's return value (e.g., `-> Literal["foo", "__end__"]:`)
or a path_map, the graph visualization assumes the edge could transition to any node in the graph.
""" # noqa: E501
if self.compiled:
logger.warning(
"Adding an edge to a graph that has already been compiled. This will "
"not be reflected in the compiled graph."
)
# find a name for the condition
path = coerce_to_runnable(path, name=None, trace=True)
name = path.name or "condition"
# validate the condition
if name in self.branches[source]:
raise ValueError(
f"Branch with name `{path.name}` already exists for node `{source}`"
)
# save it
self.branches[source][name] = Branch.from_path(path, path_map, then, False)
return self
def set_entry_point(self, key: str) -> Self:
"""Specifies the first node to be called in the graph.
Equivalent to calling `add_edge(START, key)`.
Parameters:
key (str): The key of the node to set as the entry point.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_edge(START, key)
def set_conditional_entry_point(
self,
path: Union[
Callable[..., Union[Hashable, list[Hashable]]],
Callable[..., Awaitable[Union[Hashable, list[Hashable]]]],
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
then: Optional[str] = None,
) -> Self:
"""Sets a conditional entry point in the graph.
Args:
path: The callable that determines the next
node or nodes. If not specifying `path_map` it should return one or
more nodes. If it returns END, the graph will stop execution.
path_map: Optional mapping of paths to node
names. If omitted the paths returned by `path` should be node names.
then: The name of a node to execute after the nodes
selected by `path`.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_conditional_edges(START, path, path_map, then)
def set_finish_point(self, key: str) -> Self:
"""Marks a node as a finish point of the graph.
If the graph reaches this node, it will cease execution.
Parameters:
key (str): The key of the node to set as the finish point.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_edge(key, END)
def validate(self, interrupt: Optional[Sequence[str]] = None) -> Self:
# assemble sources
all_sources = {src for src, _ in self._all_edges}
for start, branches in self.branches.items():
all_sources.add(start)
for cond, branch in branches.items():
if branch.then is not None:
if branch.ends is not None:
for end in branch.ends.values():
if end != END:
all_sources.add(end)
else:
for node in self.nodes:
if node != start and node != branch.then:
all_sources.add(node)
for name, spec in self.nodes.items():
if spec.ends:
all_sources.add(name)
# validate sources
for source in all_sources:
if source not in self.nodes and source != START:
raise ValueError(f"Found edge starting at unknown node '{source}'")
if START not in all_sources:
raise ValueError(
"Graph must have an entrypoint: add at least one edge from START to another node"
)
# assemble targets
all_targets = {end for _, end in self._all_edges}
for start, branches in self.branches.items():
for cond, branch in branches.items():
if branch.then is not None:
all_targets.add(branch.then)
if branch.ends is not None:
for end in branch.ends.values():
if end not in self.nodes and end != END:
raise ValueError(
f"At '{start}' node, '{cond}' branch found unknown target '{end}'"
)
all_targets.add(end)
else:
all_targets.add(END)
for node in self.nodes:
if node != start and node != branch.then:
all_targets.add(node)
for name, spec in self.nodes.items():
if spec.ends:
all_targets.update(spec.ends)
for target in all_targets:
if target not in self.nodes and target != END:
raise ValueError(f"Found edge ending at unknown node `{target}`")
# validate interrupts
if interrupt:
for node in interrupt:
if node not in self.nodes:
raise ValueError(f"Interrupt node `{node}` not found")
self.compiled = True
return self
def compile(
self,
checkpointer: Checkpointer = None,
interrupt_before: Optional[Union[All, list[str]]] = None,
interrupt_after: Optional[Union[All, list[str]]] = None,
debug: bool = False,
name: Optional[str] = None,
*,
cache: Optional[BaseCache] = None,
store: Optional[BaseStore] = None,
) -> "CompiledGraph":
"""Compiles the graph into a `CompiledGraph` object.
The compiled graph implements the `Runnable` interface and can be invoked,
streamed, batched, and run asynchronously.
Args:
checkpointer: A checkpoint saver object or flag.
If provided, this Checkpointer serves as a fully versioned "short-term memory" for the graph,
allowing it to be paused, resumed, and replayed from any point.
If None, it may inherit the parent graph's checkpointer when used as a subgraph.
If False, it will not use or inherit any checkpointer.
interrupt_before: An optional list of node names to interrupt before.
interrupt_after: An optional list of node names to interrupt after.
debug: A flag indicating whether to enable debug mode.
name: The name to use for the compiled graph.
Returns:
CompiledGraph: The compiled graph.
"""
# assign default values
interrupt_before = interrupt_before or []
interrupt_after = interrupt_after or []
# validate the graph
self.validate(
interrupt=(
(interrupt_before if interrupt_before != "*" else []) + interrupt_after
if interrupt_after != "*"
else []
)
)
# create empty compiled graph
compiled = CompiledGraph(
builder=self,
nodes={},
channels={START: EphemeralValue(Any), END: EphemeralValue(Any)},
input_channels=START,
output_channels=END,
stream_mode="values",
stream_channels=[],
checkpointer=checkpointer,
interrupt_before_nodes=interrupt_before,
interrupt_after_nodes=interrupt_after,
auto_validate=False,
debug=debug,
name=name or "LangGraph",
cache=cache,
store=store,
)
# attach nodes, edges, and branches
for key, node in self.nodes.items():
compiled.attach_node(key, node)
for start, end in self.edges:
compiled.attach_edge(start, end)
for start, branches in self.branches.items():
for name, branch in branches.items():
compiled.attach_branch(start, name, branch)
# validate the compiled graph
return compiled.validate()
class CompiledGraph(Pregel):
builder: Graph
def __init__(self, *, builder: Graph, **kwargs: Any) -> None:
super().__init__(**kwargs)
self.builder = builder
def attach_node(self, key: str, node: NodeSpec) -> None:
self.channels[key] = EphemeralValue(Any)
self.nodes[key] = (
PregelNode(channels=[], triggers=[], metadata=node.metadata)
| node.runnable
| ChannelWrite([ChannelWriteEntry(key)])
)
cast(list[str], self.stream_channels).append(key)
def attach_edge(self, start: str, end: str) -> None:
if end == END:
# publish to end channel
self.nodes[start].writers.append(ChannelWrite([ChannelWriteEntry(END)]))
else:
# subscribe to start channel
self.nodes[end].triggers.append(start)
cast(list[str], self.nodes[end].channels).append(start)
def attach_branch(self, start: str, name: str, branch: Branch) -> None:
def get_writes(
packets: Sequence[Union[str, Send]], static: bool = False
) -> Sequence[Union[ChannelWriteEntry, Send]]:
return [
(
ChannelWriteEntry(f"branch:{start}:{name}:{p}" if p != END else END)
if not isinstance(p, Send)
else p
)
for p in packets
]
# add hidden start node
if start == START and start not in self.nodes:
self.nodes[start] = (
NodeBuilder().subscribe_only(START).meta(TAG_HIDDEN).build()
)
# attach branch writer
self.nodes[start] |= branch.run(get_writes)
# attach branch readers
ends = branch.ends.values() if branch.ends else [node for node in self.nodes]
for end in ends:
if end != END:
channel_name = f"branch:{start}:{name}:{end}"
self.channels[channel_name] = EphemeralValue(Any)
self.nodes[end].triggers.append(channel_name)
cast(list[str], self.nodes[end].channels).append(channel_name)
-52
View File
@@ -24,7 +24,6 @@ from langchain_core.messages import (
from typing_extensions import TypedDict
from langgraph.constants import CONF, CONFIG_KEY_SEND
from langgraph.graph.state import StateGraph
Messages = Union[list[MessageLikeRepresentation], MessageLikeRepresentation]
@@ -226,57 +225,6 @@ def add_messages(
return merged
class MessageGraph(StateGraph):
"""A StateGraph where every node receives a list of messages as input and returns one or more messages as output.
MessageGraph is a subclass of StateGraph whose entire state is a single, append-only* list of messages.
Each node in a MessageGraph takes a list of messages as input and returns zero or more
messages as output. The `add_messages` function is used to merge the output messages from each node
into the existing list of messages in the graph's state.
Examples:
```pycon
>>> from langgraph.graph.message import MessageGraph
...
>>> builder = MessageGraph()
>>> builder.add_node("chatbot", lambda state: [("assistant", "Hello!")])
>>> builder.set_entry_point("chatbot")
>>> builder.set_finish_point("chatbot")
>>> builder.compile().invoke([("user", "Hi there.")])
[HumanMessage(content="Hi there.", id='...'), AIMessage(content="Hello!", id='...')]
```
```pycon
>>> from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
>>> from langgraph.graph.message import MessageGraph
...
>>> builder = MessageGraph()
>>> builder.add_node(
... "chatbot",
... lambda state: [
... AIMessage(
... content="Hello!",
... tool_calls=[{"name": "search", "id": "123", "args": {"query": "X"}}],
... )
... ],
... )
>>> builder.add_node(
... "search", lambda state: [ToolMessage(content="Searching...", tool_call_id="123")]
... )
>>> builder.set_entry_point("chatbot")
>>> builder.add_edge("chatbot", "search")
>>> builder.set_finish_point("search")
>>> builder.compile().invoke([HumanMessage(content="Hi there. Can you search for X?")])
{'messages': [HumanMessage(content="Hi there. Can you search for X?", id='b8b7d8f4-7f4d-4f4d-9c1d-f8b8d8f4d9c1'),
AIMessage(content="Hello!", id='f4d9c1d8-8d8f-4d9c-b8b7-d8f4f4d9c1d8'),
ToolMessage(content="Searching...", id='d8f4f4d9-c1d8-4f4d-b8b7-d8f4f4d9c1d8', tool_call_id="123")]}
```
"""
def __init__(self) -> None:
super().__init__(Annotated[list[AnyMessage], add_messages]) # type: ignore[arg-type]
class MessagesState(TypedDict):
messages: Annotated[list[AnyMessage], add_messages]
+151 -81
View File
@@ -25,15 +25,9 @@ from langchain_core.runnables import Runnable, RunnableConfig
from pydantic import BaseModel
from typing_extensions import Self
from langgraph._api.deprecation import LangGraphDeprecationWarning
from langgraph.cache.base import BaseCache
from langgraph.channels.base import BaseChannel
from langgraph.channels.binop import BinaryOperatorAggregate
from langgraph.channels.dynamic_barrier_value import (
DynamicBarrierValue,
DynamicBarrierValueAfterFinish,
WaitForNames,
)
from langgraph.channels.ephemeral_value import EphemeralValue
from langgraph.channels.last_value import LastValue, LastValueAfterFinish
from langgraph.channels.named_barrier_value import (
@@ -43,10 +37,12 @@ from langgraph.channels.named_barrier_value import (
from langgraph.checkpoint.base import Checkpoint
from langgraph.constants import (
EMPTY_SEQ,
END,
INTERRUPT,
MISSING,
NS_END,
NS_SEP,
START,
TAG_HIDDEN,
TASKS,
)
@@ -57,17 +53,11 @@ from langgraph.errors import (
create_error_message,
)
from langgraph.graph.branch import Branch
from langgraph.graph.graph import (
END,
START,
CompiledGraph,
Graph,
Send,
)
from langgraph.managed.base import (
ManagedValueSpec,
is_managed_value,
)
from langgraph.pregel import Pregel
from langgraph.pregel.read import ChannelRead, PregelNode
from langgraph.pregel.write import (
ChannelWrite,
@@ -75,8 +65,12 @@ from langgraph.pregel.write import (
ChannelWriteTupleEntry,
)
from langgraph.store.base import BaseStore
from langgraph.types import All, CachePolicy, Checkpointer, Command, RetryPolicy
from langgraph.utils.fields import get_field_default, get_update_as_tuples
from langgraph.types import All, CachePolicy, Checkpointer, Command, RetryPolicy, Send
from langgraph.utils.fields import (
get_cached_annotated_keys,
get_field_default,
get_update_as_tuples,
)
from langgraph.utils.pydantic import create_model
from langgraph.utils.runnable import RunnableLike, coerce_to_runnable
@@ -114,7 +108,7 @@ class StateNodeSpec(NamedTuple):
defer: bool = False
class StateGraph(Graph):
class StateGraph:
"""A graph whose nodes communicate by reading and writing to a shared state.
The signature of each node is State -> Partial<State>.
@@ -166,39 +160,34 @@ class StateGraph(Graph):
```
"""
nodes: dict[str, StateNodeSpec] # type: ignore[assignment]
edges: set[tuple[str, str]]
nodes: dict[str, StateNodeSpec]
branches: defaultdict[str, dict[str, Branch]]
channels: dict[str, BaseChannel]
managed: dict[str, ManagedValueSpec]
schemas: dict[type[Any], dict[str, Union[BaseChannel, ManagedValueSpec]]]
def __init__(
self,
state_schema: Optional[type[Any]] = None,
state_schema: type[Any],
config_schema: Optional[type[Any]] = None,
*,
input: Optional[type[Any]] = None,
output: Optional[type[Any]] = None,
) -> None:
super().__init__()
if state_schema is None:
if input is None or output is None:
raise ValueError("Must provide state_schema or input and output")
state_schema = input
warnings.warn(
"Initializing StateGraph without state_schema is deprecated. "
"Please pass in an explicit state_schema instead of just an input and output schema.",
LangGraphDeprecationWarning,
stacklevel=2,
)
else:
if input is None:
input = state_schema
if output is None:
output = state_schema
if input is None:
input = state_schema
if output is None:
output = state_schema
self.nodes = {}
self.edges = set[tuple[str, str]]()
self.branches = defaultdict(dict)
self.support_multiple_edges = False
self.compiled = False
self.schemas = {}
self.channels = {}
self.managed = {}
self.type_hints: dict[type[Any], dict[str, Any]] = {}
self.schema = state_schema
self.input = input
self.output = output
@@ -226,7 +215,6 @@ class StateGraph(Graph):
" Managed channels are not permitted in Input/Output schema."
)
self.schemas[schema] = {**channels, **managed}
self.type_hints[schema] = type_hints
for key, channel in channels.items():
if key in self.channels:
if self.channels[key] != channel:
@@ -373,7 +361,7 @@ class StateGraph(Graph):
raise ValueError(f"Node `{node}` is reserved.")
for character in (NS_SEP, NS_END):
if character in cast(str, node):
if character in node:
raise ValueError(
f"'{character}' is a reserved character and is not allowed in the node names."
)
@@ -428,8 +416,8 @@ class StateGraph(Graph):
if input is not None:
self._add_schema(input)
self.nodes[cast(str, node)] = StateNodeSpec(
coerce_to_runnable(action, name=cast(str, node), trace=False),
self.nodes[node] = StateNodeSpec(
coerce_to_runnable(action, name=node, trace=False),
metadata,
input=input or self.schema,
retry_policy=retry,
@@ -456,14 +444,30 @@ class StateGraph(Graph):
Returns:
Self: The instance of the state graph, allowing for method chaining.
"""
if isinstance(start_key, str):
return super().add_edge(start_key, end_key)
if self.compiled:
logger.warning(
"Adding an edge to a graph that has already been compiled. This will "
"not be reflected in the compiled graph."
)
if isinstance(start_key, str):
if start_key == END:
raise ValueError("END cannot be a start node")
if end_key == START:
raise ValueError("START cannot be an end node")
# run this validation only for non-StateGraph graphs
if not hasattr(self, "channels") and start_key in set(
start for start, _ in self.edges
):
raise ValueError(
f"Already found path for node '{start_key}'.\n"
"For multiple edges, use StateGraph with an Annotated state key."
)
self.edges.add((start_key, end_key))
return self
for start in start_key:
if start == END:
raise ValueError("END cannot be a start node")
@@ -486,7 +490,6 @@ class StateGraph(Graph):
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
then: Optional[str] = None,
) -> Self:
"""Add a conditional edge from the starting node to any number of destination nodes.
@@ -498,8 +501,6 @@ class StateGraph(Graph):
more nodes. If it returns END, the graph will stop execution.
path_map: Optional mapping of paths to node
names. If omitted the paths returned by `path` should be node names.
then: The name of a node to execute after the nodes
selected by `path`.
Returns:
Self: The instance of the graph, allowing for method chaining.
@@ -523,7 +524,7 @@ class StateGraph(Graph):
f"Branch with name `{path.name}` already exists for node `{source}`"
)
# save it
self.branches[source][name] = Branch.from_path(path, path_map, then, True)
self.branches[source][name] = Branch.from_path(path, path_map, True)
if schema := self.branches[source][name].input_schema:
self._add_schema(schema)
return self
@@ -570,6 +571,104 @@ class StateGraph(Graph):
return self
def set_entry_point(self, key: str) -> Self:
"""Specifies the first node to be called in the graph.
Equivalent to calling `add_edge(START, key)`.
Parameters:
key (str): The key of the node to set as the entry point.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_edge(START, key)
def set_conditional_entry_point(
self,
path: Union[
Callable[..., Union[Hashable, list[Hashable]]],
Callable[..., Awaitable[Union[Hashable, list[Hashable]]]],
Runnable[Any, Union[Hashable, list[Hashable]]],
],
path_map: Optional[Union[dict[Hashable, str], list[str]]] = None,
) -> Self:
"""Sets a conditional entry point in the graph.
Args:
path: The callable that determines the next
node or nodes. If not specifying `path_map` it should return one or
more nodes. If it returns END, the graph will stop execution.
path_map: Optional mapping of paths to node
names. If omitted the paths returned by `path` should be node names.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_conditional_edges(START, path, path_map)
def set_finish_point(self, key: str) -> Self:
"""Marks a node as a finish point of the graph.
If the graph reaches this node, it will cease execution.
Parameters:
key (str): The key of the node to set as the finish point.
Returns:
Self: The instance of the graph, allowing for method chaining.
"""
return self.add_edge(key, END)
def validate(self, interrupt: Optional[Sequence[str]] = None) -> Self:
# assemble sources
all_sources = {src for src, _ in self._all_edges}
for start, branches in self.branches.items():
all_sources.add(start)
for name, spec in self.nodes.items():
if spec.ends:
all_sources.add(name)
# validate sources
for source in all_sources:
if source not in self.nodes and source != START:
raise ValueError(f"Found edge starting at unknown node '{source}'")
if START not in all_sources:
raise ValueError(
"Graph must have an entrypoint: add at least one edge from START to another node"
)
# assemble targets
all_targets = {end for _, end in self._all_edges}
for start, branches in self.branches.items():
for cond, branch in branches.items():
if branch.ends is not None:
for end in branch.ends.values():
if end not in self.nodes and end != END:
raise ValueError(
f"At '{start}' node, '{cond}' branch found unknown target '{end}'"
)
all_targets.add(end)
else:
all_targets.add(END)
for node in self.nodes:
if node != start:
all_targets.add(node)
for name, spec in self.nodes.items():
if spec.ends:
all_targets.update(spec.ends)
for target in all_targets:
if target not in self.nodes and target != END:
raise ValueError(f"Found edge ending at unknown node `{target}`")
# validate interrupts
if interrupt:
for node in interrupt:
if node not in self.nodes:
raise ValueError(f"Interrupt node `{node}` not found")
self.compiled = True
return self
def compile(
self,
checkpointer: Checkpointer = None,
@@ -680,17 +779,19 @@ class StateGraph(Graph):
return compiled.validate()
class CompiledStateGraph(CompiledGraph):
class CompiledStateGraph(Pregel):
builder: StateGraph
schema_to_mapper: dict[type[Any], Optional[Callable[[Any], Any]]]
def __init__(
self,
*,
builder: StateGraph,
schema_to_mapper: dict[type[Any], Optional[Callable[[Any], Any]]],
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self.builder = builder
self.schema_to_mapper = schema_to_mapper
def get_input_schema(
@@ -754,7 +855,7 @@ class CompiledStateGraph(CompiledGraph):
else:
updates.extend(_get_updates(i) or ())
return updates
elif (t := type(input)) and get_type_hints(t):
elif (t := type(input)) and get_cached_annotated_keys(t):
return get_update_as_tuples(input, output_keys)
else:
msg = create_error_message(
@@ -794,7 +895,6 @@ class CompiledStateGraph(CompiledGraph):
mapper = _pick_mapper(
list(input_values),
input_schema,
self.builder.type_hints[input_schema],
)
self.schema_to_mapper[input_schema] = mapper
@@ -865,17 +965,6 @@ class CompiledStateGraph(CompiledGraph):
]
if not writes:
return []
if branch.then and branch.then != END:
writes.append(
ChannelWriteEntry(
f"branch:{start}:{name}::then",
WaitForNames(
frozenset(
p.node if isinstance(p, Send) else p for p in packets
)
),
)
)
return writes
if with_reader:
@@ -890,7 +979,7 @@ class CompiledStateGraph(CompiledGraph):
if schema in self.schema_to_mapper:
mapper = self.schema_to_mapper[schema]
else:
mapper = _pick_mapper(channels, schema, self.builder.type_hints[schema])
mapper = _pick_mapper(channels, schema)
self.schema_to_mapper[schema] = mapper
# create reader
reader: Optional[Callable[[RunnableConfig], Any]] = partial(
@@ -906,25 +995,6 @@ class CompiledStateGraph(CompiledGraph):
# attach branch publisher
self.nodes[start].writers.append(branch.run(get_writes, reader))
# attach then subscriber
if branch.then and branch.then != END:
ends = (
branch.ends.values()
if branch.ends
else [node for node in self.builder.nodes if node != branch.then]
)
channel_name = f"branch:{start}:{name}::then"
if self.builder.nodes[branch.then].defer:
self.channels[channel_name] = DynamicBarrierValueAfterFinish(str)
else:
self.channels[channel_name] = DynamicBarrierValue(str)
self.nodes[branch.then].triggers.append(channel_name)
for end in ends:
if end != END:
self.nodes[end].writers.append(
ChannelWrite((ChannelWriteEntry(channel_name, end),))
)
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
"""Migrate a checkpoint to new channel layout."""
@@ -1031,7 +1101,7 @@ class CompiledStateGraph(CompiledGraph):
def _pick_mapper(
state_keys: Sequence[str], schema: type[Any], type_hints: Optional[dict[str, Any]]
state_keys: Sequence[str], schema: type[Any]
) -> Optional[Callable[[Any], Any]]:
if state_keys == ["__root__"]:
return None
+4 -8
View File
@@ -2314,7 +2314,7 @@ class Pregel(PregelProtocol):
output_keys: The keys to stream, defaults to all non-context channels.
interrupt_before: Nodes to interrupt before, defaults to all nodes in the graph.
interrupt_after: Nodes to interrupt after, defaults to all nodes in the graph.
checkpoint_during: Whether to checkpoint intermediate steps, defaults to True. If False, only the final checkpoint is saved.
checkpoint_during: Whether to checkpoint intermediate steps, defaults to False. If False, only the final checkpoint is saved.
debug: Whether to print debug information during execution, defaults to False.
subgraphs: Whether to stream events from inside subgraphs, defaults to False.
If True, the events will be emitted as tuples `(namespace, data)`,
@@ -2421,7 +2421,7 @@ class Pregel(PregelProtocol):
debug=debug,
checkpoint_during=checkpoint_during
if checkpoint_during is not None
else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, True),
else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, False),
trigger_to_nodes=self.trigger_to_nodes,
migrate_checkpoint=self._migrate_checkpoint,
retry_policy=self.retry_policy,
@@ -2535,7 +2535,7 @@ class Pregel(PregelProtocol):
output_keys: The keys to stream, defaults to all non-context channels.
interrupt_before: Nodes to interrupt before, defaults to all nodes in the graph.
interrupt_after: Nodes to interrupt after, defaults to all nodes in the graph.
checkpoint_during: Whether to checkpoint intermediate steps, defaults to True. If False, only the final checkpoint is saved.
checkpoint_during: Whether to checkpoint intermediate steps, defaults to False. If False, only the final checkpoint is saved.
debug: Whether to print debug information during execution, defaults to False.
subgraphs: Whether to stream events from inside subgraphs, defaults to False.
If True, the events will be emitted as tuples `(namespace, data)`,
@@ -2664,7 +2664,7 @@ class Pregel(PregelProtocol):
debug=debug,
checkpoint_during=checkpoint_during
if checkpoint_during is not None
else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, True),
else config[CONF].get(CONFIG_KEY_CHECKPOINT_DURING, False),
trigger_to_nodes=self.trigger_to_nodes,
migrate_checkpoint=self._migrate_checkpoint,
retry_policy=self.retry_policy,
@@ -2743,7 +2743,6 @@ class Pregel(PregelProtocol):
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
checkpoint_during: bool | None = None,
debug: bool | None = None,
**kwargs: Any,
) -> dict[str, Any] | Any:
@@ -2776,7 +2775,6 @@ class Pregel(PregelProtocol):
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
checkpoint_during=checkpoint_during,
debug=debug,
**kwargs,
):
@@ -2811,7 +2809,6 @@ class Pregel(PregelProtocol):
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
checkpoint_during: bool | None = None,
debug: bool | None = None,
**kwargs: Any,
) -> dict[str, Any] | Any:
@@ -2845,7 +2842,6 @@ class Pregel(PregelProtocol):
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
checkpoint_during=checkpoint_during,
debug=debug,
**kwargs,
):
+17 -2
View File
@@ -144,6 +144,19 @@ def map_debug_task_results(
}
def rm_pregel_keys(config: Optional[RunnableConfig]) -> Optional[RunnableConfig]:
"""Remove pregel-specific keys from the config."""
if config is None:
return config
return {
"configurable": {
k: v
for k, v in config.get("configurable", {}).items()
if not k.startswith("__pregel_")
}
}
def map_debug_checkpoint(
step: int,
config: RunnableConfig,
@@ -183,8 +196,10 @@ def map_debug_checkpoint(
"timestamp": checkpoint["ts"],
"step": step,
"payload": {
"config": patch_checkpoint_map(config, metadata),
"parent_config": patch_checkpoint_map(parent_config, metadata),
"config": rm_pregel_keys(patch_checkpoint_map(config, metadata)),
"parent_config": rm_pregel_keys(
patch_checkpoint_map(parent_config, metadata)
),
"values": read_channels(channels, stream_channels),
"metadata": metadata,
"next": [t.name for t in tasks],
+7 -2
View File
@@ -564,7 +564,13 @@ class PregelLoop:
"debug",
map_debug_checkpoint,
self.step - 1, # printing checkpoint for previous step
self.checkpoint_config,
{
**self.checkpoint_config,
CONF: {
**self.checkpoint_config[CONF],
CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],
},
},
self.channels,
self.stream_keys,
self.checkpoint_metadata,
@@ -819,7 +825,6 @@ class PregelLoop:
**self.checkpoint_config,
CONF: {
**self.checkpoint_config[CONF],
# this is guaranteed to be set by code above
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
CONFIG_KEY_CHECKPOINT_NS, ""
),
+2 -2
View File
@@ -41,7 +41,7 @@ def run_with_retry(
except ParentCommand as exc:
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
cmd = exc.args[0]
if cmd.graph == ns:
if cmd.graph in (ns, task.name):
# this command is for the current graph, handle it
for w in task.writers:
w.invoke(cmd, config)
@@ -137,7 +137,7 @@ async def arun_with_retry(
except ParentCommand as exc:
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
cmd = exc.args[0]
if cmd.graph == ns:
if cmd.graph in (ns, task.name):
# this command is for the current graph, handle it
for w in task.writers:
w.invoke(cmd, config)
+3 -4
View File
@@ -14,7 +14,6 @@ from typing import (
TypeVar,
Union,
cast,
get_type_hints,
)
from langchain_core.runnables import Runnable, RunnableConfig
@@ -23,7 +22,7 @@ from xxhash import xxh3_128_hexdigest
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
from langgraph.utils.cache import default_cache_key
from langgraph.utils.fields import get_update_as_tuples
from langgraph.utils.fields import get_cached_annotated_keys, get_update_as_tuples
if TYPE_CHECKING:
from langgraph.pregel.protocol import PregelProtocol
@@ -350,8 +349,8 @@ class Command(Generic[N], ToolOutputMixin):
for t in self.update
):
return self.update
elif hints := get_type_hints(type(self.update)):
return get_update_as_tuples(self.update, tuple(hints.keys()))
elif keys := get_cached_annotated_keys(type(self.update)):
return get_update_as_tuples(self.update, keys)
elif self.update is not None:
return [("__root__", self.update)]
else:
+23
View File
@@ -1,4 +1,6 @@
import dataclasses
import types
import weakref
from collections.abc import Generator, Sequence
from typing import Annotated, Any, Optional, Union, get_type_hints
@@ -178,3 +180,24 @@ def get_update_as_tuples(input: Any, keys: Sequence[str]) -> list[tuple[str, Any
or (keep is not None and k in keep)
)
]
ANNOTATED_KEYS_CACHE: weakref.WeakKeyDictionary[type[Any], tuple[str, ...]] = (
weakref.WeakKeyDictionary()
)
def get_cached_annotated_keys(obj: type[Any]) -> tuple[str, ...]:
"""Return cached annotated keys for a Python class."""
if obj in ANNOTATED_KEYS_CACHE:
return ANNOTATED_KEYS_CACHE[obj]
if isinstance(obj, type):
keys: list[str] = []
for base in reversed(obj.__mro__):
ann = base.__dict__.get("__annotations__")
if ann is None or isinstance(ann, types.GetSetDescriptorType):
continue
keys.extend(ann.keys())
return ANNOTATED_KEYS_CACHE.setdefault(obj, tuple(keys))
else:
raise TypeError(f"Expected a type, got {type(obj)}. ")
+1 -1
View File
@@ -46,7 +46,7 @@ def get_fields(
return model.model_fields
if hasattr(model, "__fields__"):
return model.__fields__ # type: ignore[return-value]
return model.__fields__
msg = f"Expected a Pydantic model. Got {type(model)}"
raise TypeError(msg)
@@ -1,44 +1,4 @@
# serializer version: 1
# name: test_branch_then[memory]
'''
graph TD;
__start__ --> prepare;
prepare -.-> finish;
prepare -.-> tool_two_fast;
prepare -.-> tool_two_slow;
tool_two_fast --> finish;
tool_two_slow --> finish;
finish --> __end__;
'''
# ---
# name: test_branch_then[memory].1
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
prepare(prepare)
tool_two_slow(tool_two_slow)
tool_two_fast(tool_two_fast)
finish(finish)
__end__([<p>__end__</p>]):::last
__start__ --> prepare;
prepare -.-> finish;
prepare -.-> tool_two_fast;
prepare -.-> tool_two_slow;
tool_two_fast --> finish;
tool_two_slow --> finish;
finish --> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_conditional_graph[memory]
'''
{
@@ -422,28 +382,6 @@
'''
# ---
# name: test_start_branch_then[memory-in_memory]
'''
---
config:
flowchart:
curve: linear
---
graph TD;
__start__([<p>__start__</p>]):::first
tool_two_slow(tool_two_slow)
tool_two_fast(tool_two_fast)
__end__([<p>__end__</p>]):::last
__start__ -.-> tool_two_fast;
__start__ -.-> tool_two_slow;
tool_two_fast --> __end__;
tool_two_slow --> __end__;
classDef default fill:#f2f0ff,line-height:1.2
classDef first fill-opacity:0
classDef last fill:#bfb6fc
'''
# ---
# name: test_weather_subgraph[memory]
'''
---
@@ -541,21 +541,6 @@
'''
# ---
# name: test_in_one_fan_out_state_graph_then_defer_node[memory-True]
'''
graph TD;
__start__ --> rewrite_query;
analyzer_one --> qa;
analyzer_one --> retriever_one;
retriever_one -.-> qa;
retriever_two --> qa;
rewrite_query -.-> analyzer_one;
rewrite_query -.-> qa;
rewrite_query -.-> retriever_two;
qa --> __end__;
'''
# ---
# name: test_in_one_fan_out_state_graph_waiting_edge[memory]
'''
graph TD;
@@ -1541,7 +1541,9 @@ def test_latest_checkpoint_state_graph(
app = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "1"}}
assert [*app.stream({"query": "what is weather in sf"}, config)] == [
assert [
*app.stream({"query": "what is weather in sf"}, config, checkpoint_during=True)
] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
@@ -1557,7 +1559,7 @@ def test_latest_checkpoint_state_graph(
},
]
assert [*app.stream(Command(resume=""), config)] == [
assert [*app.stream(Command(resume=""), config, checkpoint_during=True)] == [
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
@@ -1582,7 +1584,10 @@ async def test_latest_checkpoint_state_graph_async(
config = {"configurable": {"thread_id": "1"}}
assert [
c async for c in app.astream({"query": "what is weather in sf"}, config)
c
async for c in app.astream(
{"query": "what is weather in sf"}, config, checkpoint_during=True
)
] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
@@ -1599,7 +1604,9 @@ async def test_latest_checkpoint_state_graph_async(
},
]
assert [c async for c in app.astream(Command(resume=""), config)] == [
assert [
c async for c in app.astream(Command(resume=""), config, checkpoint_during=True)
] == [
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+136 -514
View File
@@ -45,8 +45,8 @@ from langgraph.config import get_stream_writer
from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, START
from langgraph.errors import InvalidUpdateError
from langgraph.func import entrypoint, task
from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
from langgraph.graph import END, StateGraph
from langgraph.graph.message import MessagesState, add_messages
from langgraph.prebuilt.tool_node import ToolNode
from langgraph.pregel import (
GraphRecursionError,
@@ -83,81 +83,6 @@ logger = logging.getLogger(__name__)
def test_graph_validation() -> None:
def logic(inp: str) -> str:
return ""
workflow = Graph()
workflow.add_node("agent", logic)
workflow.set_entry_point("agent")
workflow.set_finish_point("agent")
assert workflow.compile(), "valid graph"
# Accept a dead-end
workflow = Graph()
workflow.add_node("agent", logic)
workflow.set_entry_point("agent")
workflow.compile()
workflow = Graph()
workflow.add_node("agent", logic)
workflow.set_finish_point("agent")
with pytest.raises(ValueError, match="must have an entrypoint"):
workflow.compile()
workflow = Graph()
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
workflow.set_entry_point("agent")
workflow.add_conditional_edges("agent", logic, {"continue": "tools", "exit": END})
workflow.add_edge("tools", "agent")
assert workflow.compile(), "valid graph"
workflow = Graph()
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
workflow.set_entry_point("tools")
workflow.add_conditional_edges("agent", logic, {"continue": "tools", "exit": END})
workflow.add_edge("tools", "agent")
assert workflow.compile(), "valid graph"
workflow = Graph()
workflow.set_entry_point("tools")
workflow.add_conditional_edges("agent", logic, {"continue": "tools", "exit": END})
workflow.add_edge("tools", "agent")
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
assert workflow.compile(), "valid graph"
workflow = Graph()
workflow.set_entry_point("tools")
workflow.add_conditional_edges(
"agent", logic, {"continue": "tools", "exit": END, "hmm": "extra"}
)
workflow.add_edge("tools", "agent")
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
with pytest.raises(ValueError, match="unknown"): # extra is not defined
workflow.compile()
workflow = Graph()
workflow.set_entry_point("agent")
workflow.add_conditional_edges("agent", logic, {"continue": "tools", "exit": END})
workflow.add_edge("tools", "extra")
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
with pytest.raises(ValueError, match="unknown"): # extra is not defined
workflow.compile()
workflow = Graph()
workflow.add_node("agent", logic)
workflow.add_node("tools", logic)
workflow.add_node("extra", logic)
workflow.set_entry_point("agent")
workflow.add_conditional_edges("agent", logic)
workflow.add_edge("tools", "agent")
# Accept, even though extra is dead-end
workflow.compile()
class State(TypedDict):
hello: str
@@ -261,7 +186,9 @@ def test_checkpoint_errors() -> None:
builder.add_edge(START, "parallel")
graph = builder.compile(checkpointer=FaultyPutWritesCheckpointer())
with pytest.raises(ValueError, match="Faulty put_writes"):
graph.invoke("", {"configurable": {"thread_id": "thread-1"}})
graph.invoke(
"", {"configurable": {"thread_id": "thread-1"}}, checkpoint_during=True
)
def test_config_json_schema() -> None:
@@ -488,11 +415,6 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
input_channels="input",
output_channels="output",
)
graph = Graph()
graph.add_node("add_one", add_one)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one")
gapp = graph.compile()
assert app.input_schema.model_json_schema() == {
"title": "LangGraphInput",
@@ -514,21 +436,6 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
assert app.invoke(2, output_keys=["output"]) == {"output": 3}
assert repr(app), "does not raise recursion error"
assert gapp.invoke(2, debug=True) == 3
@pytest.mark.parametrize(
"falsy_value",
[None, False, 0, "", [], {}, set(), frozenset(), 0.0, 0j],
)
def test_invoke_single_process_in_out_falsy_values(falsy_value: Any) -> None:
graph = Graph()
graph.add_node("return_falsy_const", lambda *args, **kwargs: falsy_value)
graph.set_entry_point("return_falsy_const")
graph.set_finish_point("return_falsy_const")
gapp = graph.compile()
assert gapp.invoke(1) == falsy_value
def test_invoke_single_process_in_write_kwargs(mocker: MockerFixture) -> None:
add_one = mocker.Mock(side_effect=lambda x: x + 1)
@@ -642,29 +549,6 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
with pytest.raises(GraphRecursionError):
app.invoke(2, {"recursion_limit": 1}, debug=1)
graph = Graph()
graph.add_node("add_one", add_one)
graph.add_node("add_one_more", add_one)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one_more")
graph.add_edge("add_one", "add_one_more")
gapp = graph.compile()
assert gapp.invoke(2) == 4
for step, values in enumerate(gapp.stream(2, debug=1), start=1):
if step == 1:
assert values == {
"add_one": 3,
}
elif step == 2:
assert values == {
"add_one_more": 4,
}
else:
assert 0, f"{step}:{values}"
assert step == 2
def test_run_from_checkpoint_id_retains_previous_writes(
sync_checkpointer: BaseCheckpointSaver,
@@ -706,7 +590,7 @@ def test_run_from_checkpoint_id_retains_previous_writes(
thread_id = uuid.uuid4()
thread1 = {"configurable": {"thread_id": str(thread_id)}}
result = graph.invoke({"myval": 1}, thread1)
result = graph.invoke({"myval": 1}, thread1, checkpoint_during=True)
assert result["myval"] == 4
history = [c for c in graph.get_state_history(thread1)]
@@ -771,16 +655,6 @@ def test_batch_two_processes_in_out() -> None:
{"output": 7},
]
graph = Graph()
graph.add_node("add_one", add_one_with_delay)
graph.add_node("add_one_more", add_one_with_delay)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one_more")
graph.add_edge("add_one", "add_one_more")
gapp = graph.compile()
assert gapp.batch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7]
def test_invoke_many_processes_in_out(mocker: MockerFixture) -> None:
test_size = 100
@@ -1568,7 +1442,7 @@ def test_invoke_checkpoint_three(
thread_1 = {"configurable": {"thread_id": "1"}}
# total starts out as 0, so output is 0+2=2
assert app.invoke(2, thread_1, debug=1) == 2
assert app.invoke(2, thread_1, checkpoint_during=True) == 2
state = app.get_state(thread_1)
assert state is not None
assert state.values.get("total") == 2
@@ -1578,7 +1452,7 @@ def test_invoke_checkpoint_three(
== sync_checkpointer.get(thread_1)["id"]
)
# total is now 2, so output is 2+3=5
assert app.invoke(3, thread_1) == 5
assert app.invoke(3, thread_1, checkpoint_during=True) == 5
state = app.get_state(thread_1)
assert state is not None
assert state.values.get("total") == 7
@@ -1588,7 +1462,7 @@ def test_invoke_checkpoint_three(
)
# total is now 2+5=7, so output would be 7+4=11, but raises ValueError
with pytest.raises(ValueError):
app.invoke(4, thread_1)
app.invoke(4, thread_1, checkpoint_during=True)
# checkpoint is updated with new input
state = app.get_state(thread_1)
assert state is not None
@@ -1596,7 +1470,7 @@ def test_invoke_checkpoint_three(
assert state.next == ("one",)
"""we checkpoint inputs and it failed on "one", so the next node is one"""
# we can recover from error by sending new inputs
assert app.invoke(2, thread_1) == 9
assert app.invoke(2, thread_1, checkpoint_during=True) == 9
state = app.get_state(thread_1)
assert state is not None
assert state.values.get("total") == 16, "total is now 7+9=16"
@@ -1604,8 +1478,8 @@ def test_invoke_checkpoint_three(
thread_2 = {"configurable": {"thread_id": "2"}}
# on a new thread, total starts out as 0, so output is 0+5=5
assert app.invoke(5, thread_2, debug=True) == 5
state = app.get_state({"configurable": {"thread_id": "1"}})
assert app.invoke(5, thread_2) == 5
state = app.get_state(thread_1)
assert state is not None
assert state.values.get("total") == 16
assert state.next == (), "checkpoint of other thread not touched"
@@ -1825,50 +1699,6 @@ def test_invoke_two_processes_no_in(mocker: MockerFixture) -> None:
Pregel(nodes={"one": one, "two": two})
def test_conditional_entrypoint_graph(snapshot: SnapshotAssertion) -> None:
def left(data: str) -> str:
return data + "->left"
def right(data: str) -> str:
return data + "->right"
def should_start(data: str) -> str:
# Logic to decide where to start
if len(data) > 10:
return "go-right"
else:
return "go-left"
# Define a new graph
workflow = Graph()
workflow.add_node("left", left)
workflow.add_node("right", right)
workflow.set_conditional_entry_point(
should_start, {"go-left": "left", "go-right": "right"}
)
workflow.add_conditional_edges("left", lambda data: END, {END: END})
workflow.add_edge("right", END)
app = workflow.compile()
assert json.dumps(app.get_input_schema().model_json_schema()) == snapshot
assert json.dumps(app.get_output_schema().model_json_schema()) == snapshot
assert json.dumps(app.get_graph().to_json(), indent=2) == snapshot
assert app.get_graph().draw_mermaid(with_styles=False) == snapshot
assert (
app.invoke("what is weather in sf", debug=True)
== "what is weather in sf->right"
)
assert [*app.stream("what is weather in sf")] == [
{"right": "what is weather in sf->right"},
]
def test_conditional_entrypoint_to_multiple_state_graph(
snapshot: SnapshotAssertion,
) -> None:
@@ -2507,275 +2337,6 @@ def test_in_one_fan_out_state_graph_defer_node(
]
@pytest.mark.parametrize("with_path_map", (True, False))
def test_in_one_fan_out_state_graph_then_defer_node(
snapshot: SnapshotAssertion,
sync_checkpointer: BaseCheckpointSaver,
with_path_map: bool,
) -> None:
def sorted_add(
x: list[str], y: Union[list[str], list[tuple[str, str]]]
) -> list[str]:
if isinstance(y[0], tuple):
for rem, _ in y:
x.remove(rem)
y = [t[1] for t in y]
return sorted(operator.add(x, y))
class State(TypedDict, total=False):
query: str
answer: str
docs: Annotated[list[str], sorted_add]
workflow = StateGraph(State)
@workflow.add_node
def rewrite_query(data: State) -> State:
return {"query": f"query: {data['query']}"}
def analyzer_one(data: State) -> State:
return {"query": f"analyzed: {data['query']}"}
def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]}
def retriever_two(data: State) -> State:
time.sleep(0.1) # to ensure stream order
return {"docs": ["doc3", "doc4"]}
def qa(data: State) -> State:
return {"answer": ",".join(data["docs"])}
workflow.add_node(analyzer_one)
workflow.add_node(retriever_one)
workflow.add_node(retriever_two)
workflow.add_node(qa, defer=True)
workflow.set_entry_point("rewrite_query")
workflow.add_conditional_edges(
"rewrite_query",
lambda _: ["analyzer_one", "retriever_two"],
["analyzer_one", "retriever_two"] if with_path_map else None,
then="qa",
)
workflow.add_edge("analyzer_one", "retriever_one")
app = workflow.compile()
if isinstance(sync_checkpointer, InMemorySaver) and with_path_map:
assert app.get_graph().draw_mermaid(with_styles=False) == snapshot
assert app.invoke({"query": "what is weather in sf"}) == {
"query": "analyzed: query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"],
"answer": "doc1,doc2,doc3,doc4",
}
assert [*app.stream({"query": "what is weather in sf"})] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
assert [*app.stream({"query": "what is weather in sf"}, stream_mode="debug")] == [
{
"type": "task",
"timestamp": AnyStr(),
"step": 1,
"payload": {
"id": AnyStr(),
"name": "rewrite_query",
"input": {"query": "what is weather in sf", "docs": []},
"triggers": ("branch:to:rewrite_query",),
},
},
{
"type": "task_result",
"timestamp": AnyStr(),
"step": 1,
"payload": {
"id": AnyStr(),
"name": "rewrite_query",
"error": None,
"result": [("query", "query: what is weather in sf")],
"interrupts": [],
},
},
{
"type": "task",
"timestamp": AnyStr(),
"step": 2,
"payload": {
"id": AnyStr(),
"name": "analyzer_one",
"input": {
"query": "query: what is weather in sf",
"docs": [],
},
"triggers": ("branch:to:analyzer_one",),
},
},
{
"type": "task",
"timestamp": AnyStr(),
"step": 2,
"payload": {
"id": AnyStr(),
"name": "retriever_two",
"input": {"query": "query: what is weather in sf", "docs": []},
"triggers": ("branch:to:retriever_two",),
},
},
{
"type": "task_result",
"timestamp": AnyStr(),
"step": 2,
"payload": {
"id": AnyStr(),
"name": "analyzer_one",
"error": None,
"result": [("query", "analyzed: query: what is weather in sf")],
"interrupts": [],
},
},
{
"type": "task_result",
"timestamp": AnyStr(),
"step": 2,
"payload": {
"id": AnyStr(),
"name": "retriever_two",
"error": None,
"result": [("docs", ["doc3", "doc4"])],
"interrupts": [],
},
},
{
"type": "task",
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": AnyStr(),
"name": "retriever_one",
"input": {
"query": "analyzed: query: what is weather in sf",
"docs": ["doc3", "doc4"],
},
"triggers": ("branch:to:retriever_one",),
},
},
{
"type": "task_result",
"timestamp": AnyStr(),
"step": 3,
"payload": {
"id": AnyStr(),
"name": "retriever_one",
"error": None,
"result": [("docs", ["doc1", "doc2"])],
"interrupts": [],
},
},
{
"type": "task",
"timestamp": AnyStr(),
"step": 4,
"payload": {
"id": AnyStr(),
"name": "qa",
"input": {
"query": "analyzed: query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"],
},
"triggers": ("branch:rewrite_query:condition::then", "branch:to:qa"),
},
},
{
"type": "task_result",
"timestamp": AnyStr(),
"step": 4,
"payload": {
"id": AnyStr(),
"name": "qa",
"error": None,
"result": [("answer", "doc1,doc2,doc3,doc4")],
"interrupts": [],
},
},
]
app_w_interrupt = workflow.compile(
checkpointer=sync_checkpointer,
interrupt_after=["analyzer_one"],
)
config = {"configurable": {"thread_id": "1"}}
assert [
c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config)
] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"__interrupt__": ()},
]
assert [c for c in app_w_interrupt.stream(None, config)] == [
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
app_w_interrupt = workflow.compile(
checkpointer=sync_checkpointer,
interrupt_before=["qa"],
)
config = {"configurable": {"thread_id": "2"}}
assert [
c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config)
] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
app_w_interrupt.update_state(config, {"docs": ["doc5"]})
expected_parent_config = list(app_w_interrupt.checkpointer.list(config, limit=2))[
-1
].config
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
"query": "analyzed: query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4", "doc5"],
},
tasks=(PregelTask(AnyStr(), "qa", (PULL, "qa")),),
next=("qa",),
config={
"configurable": {
"thread_id": "2",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
created_at=AnyStr(),
metadata={
"parents": {},
"source": "update",
"step": 4,
"thread_id": "2",
},
parent_config=expected_parent_config,
interrupts=(),
)
assert [c for c in app_w_interrupt.stream(None, config, debug=1)] == [
{"qa": {"answer": "doc1,doc2,doc3,doc4,doc5"}},
]
def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
snapshot: SnapshotAssertion, sync_checkpointer: BaseCheckpointSaver
) -> None:
@@ -3559,7 +3120,6 @@ def test_nested_graph_xray(snapshot: SnapshotAssertion) -> None:
tool_two_graph.set_conditional_entry_point(
lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast",
["tool_two_slow", "tool_two_fast"],
then=END,
)
tool_two = tool_two_graph.compile()
@@ -3568,7 +3128,7 @@ def test_nested_graph_xray(snapshot: SnapshotAssertion) -> None:
graph.add_node("tool_two", tool_two)
graph.add_node("tool_three", logic)
graph.set_conditional_entry_point(
lambda s: "tool_one", ["tool_one", "tool_two", "tool_three"], then=END
lambda s: "tool_one", ["tool_one", "tool_two", "tool_three"]
)
app = graph.compile()
@@ -4381,9 +3941,14 @@ def test_checkpoint_metadata(sync_checkpointer: BaseCheckpointSaver) -> None:
def test_remove_message_via_state_update(
sync_checkpointer: BaseCheckpointSaver,
) -> None:
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
from langchain_core.messages import (
AIMessage,
AnyMessage,
HumanMessage,
RemoveMessage,
)
workflow = MessageGraph()
workflow = StateGraph(Annotated[list[AnyMessage], add_messages])
workflow.add_node(
"chatbot",
lambda state: [
@@ -4414,9 +3979,14 @@ def test_remove_message_via_state_update(
def test_remove_message_from_node():
from langchain_core.messages import AIMessage, HumanMessage, RemoveMessage
from langchain_core.messages import (
AIMessage,
AnyMessage,
HumanMessage,
RemoveMessage,
)
workflow = MessageGraph()
workflow = StateGraph(Annotated[list[AnyMessage], add_messages])
workflow.add_node(
"chatbot",
lambda state: [
@@ -4805,7 +4375,7 @@ def test_debug_retry(sync_checkpointer: BaseCheckpointSaver):
graph = builder.compile(checkpointer=sync_checkpointer)
config = {"configurable": {"thread_id": "1"}}
graph.invoke({"messages": []}, config=config)
graph.invoke({"messages": []}, config=config, checkpoint_during=True)
# re-run step: 1
target_config = next(
@@ -4815,7 +4385,11 @@ def test_debug_retry(sync_checkpointer: BaseCheckpointSaver):
)
update_config = graph.update_state(target_config, values=None)
events = [*graph.stream(None, config=update_config, stream_mode="debug")]
events = [
*graph.stream(
None, config=update_config, stream_mode="debug", checkpoint_during=True
)
]
checkpoint_events = list(
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
@@ -4845,7 +4419,9 @@ def test_debug_retry(sync_checkpointer: BaseCheckpointSaver):
assert stream_parent_conf == history_parent_conf
def test_debug_subgraphs(sync_checkpointer: BaseCheckpointSaver):
def test_debug_subgraphs(
sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
):
class State(TypedDict):
messages: Annotated[list[str], operator.add]
@@ -4878,12 +4454,15 @@ def test_debug_subgraphs(sync_checkpointer: BaseCheckpointSaver):
{"messages": []},
config=config,
stream_mode="debug",
checkpoint_during=checkpoint_during,
)
]
checkpoint_events = list(
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
)
if not checkpoint_during:
checkpoint_events = checkpoint_events[:1]
checkpoint_history = list(graph.get_state_history(config))
assert len(checkpoint_events) == len(checkpoint_history)
@@ -4912,7 +4491,9 @@ def test_debug_subgraphs(sync_checkpointer: BaseCheckpointSaver):
assert stream_task.get("state") == history_task.state
def test_debug_nested_subgraphs(sync_checkpointer: BaseCheckpointSaver):
def test_debug_nested_subgraphs(
sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
):
from collections import defaultdict
class State(TypedDict):
@@ -4955,6 +4536,7 @@ def test_debug_nested_subgraphs(sync_checkpointer: BaseCheckpointSaver):
config=config,
stream_mode="debug",
subgraphs=True,
checkpoint_during=checkpoint_during,
)
]
@@ -4994,6 +4576,9 @@ def test_debug_nested_subgraphs(sync_checkpointer: BaseCheckpointSaver):
for checkpoint_events, checkpoint_history in zip(
stream_ns.values(), history_ns.values()
):
if not checkpoint_during:
checkpoint_events = checkpoint_events[-1:]
assert len(checkpoint_events) == len(checkpoint_history)
for stream, history in zip(checkpoint_events, checkpoint_history):
assert stream["values"] == history.values
assert stream["next"] == list(history.next)
@@ -5162,7 +4747,10 @@ def test_runnable_passthrough_node_graph() -> None:
assert graph.get_graph(xray=True).to_json() == graph.get_graph(xray=False).to_json()
def test_parent_command(sync_checkpointer: BaseCheckpointSaver) -> None:
@pytest.mark.parametrize("subgraph_persist", [True, False])
def test_parent_command(
sync_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
from langchain_core.messages import BaseMessage
from langchain_core.tools import tool
@@ -5174,7 +4762,7 @@ def test_parent_command(sync_checkpointer: BaseCheckpointSaver) -> None:
subgraph_builder = StateGraph(MessagesState)
subgraph_builder.add_node("tool", get_user_name)
subgraph_builder.add_edge(START, "tool")
subgraph = subgraph_builder.compile()
subgraph = subgraph_builder.compile(checkpointer=subgraph_persist)
class CustomParentState(TypedDict):
messages: Annotated[list[BaseMessage], add_messages]
@@ -5220,15 +4808,7 @@ def test_parent_command(sync_checkpointer: BaseCheckpointSaver) -> None:
"parents": {},
},
created_at=AnyStr(),
parent_config=(
{
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
}
),
parent_config=None,
tasks=(),
interrupts=(),
)
@@ -5788,7 +5368,9 @@ def test_concurrent_execution_thread_safety():
assert result["counter"] == 1
def test_checkpoint_recovery(sync_checkpointer: BaseCheckpointSaver):
def test_checkpoint_recovery(
sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
):
"""Test recovery from checkpoints after failures."""
class State(TypedDict):
@@ -5815,7 +5397,11 @@ def test_checkpoint_recovery(sync_checkpointer: BaseCheckpointSaver):
# First attempt should fail
with pytest.raises(RuntimeError):
graph.invoke({"steps": ["start"], "attempt": 1}, config)
graph.invoke(
{"steps": ["start"], "attempt": 1},
config,
checkpoint_during=checkpoint_during,
)
# Verify checkpoint state
state = graph.get_state(config)
@@ -5825,12 +5411,17 @@ def test_checkpoint_recovery(sync_checkpointer: BaseCheckpointSaver):
assert "RuntimeError('Simulated failure')" in state.tasks[0].error
# Retry with updated attempt count
result = graph.invoke({"steps": [], "attempt": 2}, config)
result = graph.invoke(
{"steps": [], "attempt": 2}, config, checkpoint_during=checkpoint_during
)
assert result == {"steps": ["start", "node1", "node2"], "attempt": 2}
# Verify checkpoint history shows both attempts
history = list(graph.get_state_history(config))
assert len(history) == 6 # Initial + failed attempt + successful attempt
if checkpoint_during:
assert len(history) == 6 # Initial + failed attempt + successful attempt
else:
assert len(history) == 2 # error + success
# Verify the error was recorded in checkpoint
failed_checkpoint = next(c for c in history if c.tasks and c.tasks[0].error)
@@ -5893,9 +5484,7 @@ def test_multiple_updates() -> None:
]
def test_falsy_return_from_task(
sync_checkpointer: BaseCheckpointSaver, snapshot: SnapshotAssertion
):
def test_falsy_return_from_task(sync_checkpointer: BaseCheckpointSaver):
"""Test with a falsy return from a task."""
@task
@@ -5915,15 +5504,11 @@ def test_falsy_return_from_task(
{
"payload": {
"config": {
"callbacks": None,
"configurable": {
"checkpoint_id": AnyStr(),
"checkpoint_ns": "",
"thread_id": AnyStr(),
},
"metadata": {},
"recursion_limit": 25,
"tags": [],
},
"metadata": {
"parents": {},
@@ -6014,7 +5599,6 @@ def test_falsy_return_from_task(
"type": "task_result",
},
]
print(type(configurable["configurable"]["thread_id"]))
assert [
c
for c in graph.stream(Command(resume="123"), configurable, stream_mode="debug")
@@ -6022,15 +5606,11 @@ def test_falsy_return_from_task(
{
"payload": {
"config": {
"callbacks": None,
"configurable": {
"checkpoint_id": AnyStr(),
"checkpoint_ns": "",
"thread_id": AnyStr(),
},
"metadata": {},
"recursion_limit": 25,
"tags": [],
},
"metadata": {
"parents": {},
@@ -6112,15 +5692,11 @@ def test_falsy_return_from_task(
{
"payload": {
"config": {
"callbacks": None,
"configurable": {
"checkpoint_id": AnyStr(),
"checkpoint_ns": "",
"thread_id": AnyStr(),
},
"metadata": {},
"recursion_limit": 25,
"tags": [],
},
"metadata": {
"parents": {},
@@ -6128,17 +5704,7 @@ def test_falsy_return_from_task(
"step": 0,
},
"next": [],
"parent_config": {
"callbacks": None,
"configurable": {
"checkpoint_id": AnyStr(),
"checkpoint_ns": "",
"thread_id": AnyStr(),
},
"metadata": {},
"recursion_limit": 25,
"tags": [],
},
"parent_config": None,
"tasks": [],
"values": None,
},
@@ -8046,7 +7612,9 @@ def test_pregel_node_copy() -> None:
graph.nodes["agent"].copy({})
def test_update_as_input(sync_checkpointer: BaseCheckpointSaver) -> None:
def test_update_as_input(
sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
class State(TypedDict):
foo: str
@@ -8065,13 +7633,17 @@ def test_update_as_input(sync_checkpointer: BaseCheckpointSaver) -> None:
.compile(checkpointer=sync_checkpointer)
)
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
"foo": "tool"
}
assert graph.invoke(
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {"foo": "tool"}
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
"foo": "tool"
}
assert graph.invoke(
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {"foo": "tool"}
def map_snapshot(i: StateSnapshot) -> dict:
return {
@@ -8109,11 +7681,14 @@ def test_update_as_input(sync_checkpointer: BaseCheckpointSaver) -> None:
for s in graph.get_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
if checkpoint_during:
assert new_history == history
else:
assert [new_history[0], new_history[4]] == history
def test_batch_update_as_input(
sync_checkpointer: BaseCheckpointSaver,
sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
class State(TypedDict):
foo: str
@@ -8145,7 +7720,11 @@ def test_batch_update_as_input(
.compile(checkpointer=sync_checkpointer)
)
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
assert graph.invoke(
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {
"foo": "map",
"tasks": [0, 1, 2],
}
@@ -8198,7 +7777,10 @@ def test_batch_update_as_input(
for s in graph.get_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
if checkpoint_during:
assert new_history == history
else:
assert new_history[:1] == history
def test_migration_graph(snapshot: SnapshotAssertion) -> None:
@@ -8336,3 +7918,43 @@ def test_imp_exception(
{"my_task": 2},
{"my_workflow": "done"},
]
@pytest.mark.parametrize("subgraph_persist", [True, False])
def test_parent_command_goto(
sync_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
class State(TypedDict):
dialog_state: Annotated[list[str], operator.add]
def node_a_child(state):
return {"dialog_state": ["a_child_state"]}
def node_b_child(state):
return Command(
graph=Command.PARENT,
goto="node_b_parent",
update={"dialog_state": ["b_child_state"]},
)
sub_builder = StateGraph(State)
sub_builder.add_node(node_a_child)
sub_builder.add_node(node_b_child)
sub_builder.add_edge(START, "node_a_child")
sub_builder.add_edge("node_a_child", "node_b_child")
sub_graph = sub_builder.compile(checkpointer=subgraph_persist)
def node_b_parent(state):
return {"dialog_state": ["node_b_parent"]}
main_builder = StateGraph(State)
main_builder.add_node(node_b_parent)
main_builder.add_edge(START, "subgraph_node")
main_builder.add_node("subgraph_node", sub_graph, destinations=("node_b_parent",))
main_graph = main_builder.compile(sync_checkpointer, name="parent")
config = {"configurable": {"thread_id": 1}}
assert main_graph.invoke(input={"dialog_state": ["init_state"]}, config=config) == {
"dialog_state": ["init_state", "b_child_state", "node_b_parent"]
}
+168 -299
View File
@@ -43,7 +43,7 @@ from langgraph.checkpoint.memory import InMemorySaver
from langgraph.constants import CONFIG_KEY_NODE_FINISHED, ERROR, PULL, PUSH, START
from langgraph.errors import InvalidUpdateError, NodeInterrupt
from langgraph.func import entrypoint, task
from langgraph.graph import END, Graph, StateGraph
from langgraph.graph import END, StateGraph
from langgraph.graph.message import MessagesState, add_messages
from langgraph.prebuilt.tool_node import ToolNode
from langgraph.pregel import GraphRecursionError, NodeBuilder, Pregel, StateSnapshot
@@ -154,13 +154,20 @@ async def test_checkpoint_errors() -> None:
builder.add_edge(START, "parallel")
graph = builder.compile(checkpointer=FaultyPutWritesCheckpointer())
with pytest.raises(ValueError, match="Faulty put_writes"):
await graph.ainvoke("", {"configurable": {"thread_id": "thread-1"}})
await graph.ainvoke(
"", {"configurable": {"thread_id": "thread-1"}}, checkpoint_during=True
)
with pytest.raises(ValueError, match="Faulty put_writes"):
async for _ in graph.astream("", {"configurable": {"thread_id": "thread-2"}}):
async for _ in graph.astream(
"", {"configurable": {"thread_id": "thread-2"}}, checkpoint_during=True
):
pass
with pytest.raises(ValueError, match="Faulty put_writes"):
async for _ in graph.astream_events(
"", {"configurable": {"thread_id": "thread-3"}}, version="v2"
"",
{"configurable": {"thread_id": "thread-3"}},
version="v2",
checkpoint_during=True,
):
pass
@@ -255,7 +262,10 @@ async def test_checkpoint_put_after_cancellation() -> None:
finally:
logs.append("awhile.end")
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.set_entry_point("agent")
builder.set_finish_point("agent")
@@ -264,14 +274,13 @@ async def test_checkpoint_put_after_cancellation() -> None:
thread1 = {"configurable": {"thread_id": "1"}}
# start the task
t = asyncio.create_task(graph.ainvoke(1, thread1))
t = asyncio.create_task(graph.ainvoke({"hello": "world"}, thread1))
# cancel after 0.2 seconds
await asyncio.sleep(0.2)
t.cancel()
# check logs before cancellation is handled
assert sorted(logs) == [
"awhile.start",
"checkpoint.aput.start",
], "Cancelled before checkpoint put started"
# wait for task to finish
try:
@@ -319,7 +328,10 @@ async def test_checkpoint_put_after_cancellation_stream_anext() -> None:
finally:
logs.append("awhile.end")
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.set_entry_point("agent")
builder.set_finish_point("agent")
@@ -328,7 +340,7 @@ async def test_checkpoint_put_after_cancellation_stream_anext() -> None:
thread1 = {"configurable": {"thread_id": "1"}}
# start the task
s = graph.astream(1, thread1)
s = graph.astream({"hello": "world"}, thread1)
t = asyncio.create_task(s.__anext__())
# cancel after 0.2 seconds
await asyncio.sleep(0.2)
@@ -336,7 +348,6 @@ async def test_checkpoint_put_after_cancellation_stream_anext() -> None:
# check logs before cancellation is handled
assert sorted(logs) == [
"awhile.start",
"checkpoint.aput.start",
], "Cancelled before checkpoint put started"
# wait for task to finish
try:
@@ -384,7 +395,10 @@ async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
finally:
logs.append("awhile.end")
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.set_entry_point("agent")
builder.set_finish_point("agent")
@@ -393,7 +407,9 @@ async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
thread1 = {"configurable": {"thread_id": "1"}}
# start the task
s = graph.astream_events(1, thread1, version="v2", include_names=["LangGraph"])
s = graph.astream_events(
{"hello": "world"}, thread1, version="v2", include_names=["LangGraph"]
)
# skip first event (happens right away)
await s.__anext__()
# start the task for 2nd event
@@ -403,7 +419,6 @@ async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
t.cancel()
# check logs before cancellation is handled
assert logs == [
"checkpoint.aput.start",
"awhile.start",
], "Cancelled before checkpoint put started"
# wait for task to finish
@@ -412,9 +427,9 @@ async def test_checkpoint_put_after_cancellation_stream_events_anext() -> None:
except asyncio.CancelledError:
# check logs after cancellation is handled
assert logs == [
"checkpoint.aput.start",
"awhile.start",
"awhile.end",
"checkpoint.aput.start",
"checkpoint.aput.end",
], "Checkpoint put is not cancelled"
else:
@@ -432,7 +447,10 @@ async def test_node_cancellation_on_external_cancel() -> None:
inner_task_cancelled = True
raise
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.set_entry_point("agent")
builder.set_finish_point("agent")
@@ -440,7 +458,7 @@ async def test_node_cancellation_on_external_cancel() -> None:
graph = builder.compile()
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(graph.ainvoke(1), 0.5)
await asyncio.wait_for(graph.ainvoke({"hello": "world"}), 0.5)
assert inner_task_cancelled
@@ -459,16 +477,19 @@ async def test_node_cancellation_on_other_node_exception() -> None:
async def iambad(input: Any) -> None:
raise ValueError("I am bad")
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.add_node("bad", iambad)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"])
graph = builder.compile()
with pytest.raises(ValueError, match="I am bad"):
# This will raise ValueError, not TimeoutError
await asyncio.wait_for(graph.ainvoke(1), 0.5)
await asyncio.wait_for(graph.ainvoke({"hello": "world"}), 0.5)
assert inner_task_cancelled
@@ -480,16 +501,19 @@ async def test_node_cancellation_on_other_node_exception_two() -> None:
async def iambad(input: Any) -> None:
raise ValueError("I am bad")
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.add_node("bad", iambad)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"])
graph = builder.compile()
with pytest.raises(ValueError, match="I am bad"):
# This will raise ValueError, not CancelledError
await graph.ainvoke(1)
await graph.ainvoke({"hello": "world"})
@NEEDS_CONTEXTVARS
@@ -590,12 +614,6 @@ async def test_dynamic_interrupt(async_checkpointer: BaseCheckpointSaver) -> Non
"step": 0,
"thread_id": "1",
},
{
"parents": {},
"source": "input",
"step": -1,
"thread_id": "1",
},
]
tup = await tool_two.checkpointer.aget_tuple(thread1)
assert await tool_two.aget_state(thread1) == StateSnapshot(
@@ -623,9 +641,7 @@ async def test_dynamic_interrupt(async_checkpointer: BaseCheckpointSaver) -> Non
"step": 0,
"thread_id": "1",
},
parent_config=(
[c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config
),
parent_config=None,
interrupts=(
Interrupt(
value="Just because...",
@@ -771,12 +787,6 @@ async def test_dynamic_interrupt_subgraph(
"step": 0,
"thread_id": "1",
},
{
"parents": {},
"source": "input",
"step": -1,
"thread_id": "1",
},
]
tup = await tool_two.checkpointer.aget_tuple(thread1)
assert await tool_two.aget_state(thread1) == StateSnapshot(
@@ -810,11 +820,7 @@ async def test_dynamic_interrupt_subgraph(
"step": 0,
"thread_id": "1",
},
parent_config=(
[c async for c in tool_two.checkpointer.alist(thread1root, limit=2)][
-1
].config
),
parent_config=None,
interrupts=(
Interrupt(
value="Just because...",
@@ -959,12 +965,6 @@ async def test_copy_checkpoint(async_checkpointer: BaseCheckpointSaver) -> None:
"step": 0,
"thread_id": "1",
},
{
"parents": {},
"source": "input",
"step": -1,
"thread_id": "1",
},
]
tup = await tool_two.checkpointer.aget_tuple(thread1)
@@ -1002,9 +1002,7 @@ async def test_copy_checkpoint(async_checkpointer: BaseCheckpointSaver) -> None:
"step": 0,
"thread_id": "1",
},
parent_config=(
[c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config
),
parent_config=None,
interrupts=(
Interrupt(
value="Just because...",
@@ -1044,9 +1042,7 @@ async def test_copy_checkpoint(async_checkpointer: BaseCheckpointSaver) -> None:
"thread_id": "1",
},
parent_config=(
[c async for c in tool_two.checkpointer.alist(thread1, limit=2)][
-1
].parent_config
[c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config
),
interrupts=(),
)
@@ -1080,7 +1076,7 @@ async def test_node_not_cancelled_on_other_node_interrupted(
builder = StateGraph(State)
builder.add_node("agent", awhile)
builder.add_node("bad", iambad)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END)
builder.set_conditional_entry_point(lambda _: ["agent", "bad"])
graph = builder.compile(checkpointer=async_checkpointer)
thread = {"configurable": {"thread_id": "1"}}
@@ -1137,18 +1133,21 @@ async def test_step_timeout_on_stream_hang(stream_hang_s: float) -> None:
async def alittlewhile(input: Any) -> None:
await asyncio.sleep(0.6)
return "1"
return {"hello": "1"}
builder = Graph()
class State(TypedDict):
hello: str
builder = StateGraph(State)
builder.add_node(awhile)
builder.add_node(alittlewhile)
builder.set_conditional_entry_point(lambda _: ["awhile", "alittlewhile"], then=END)
builder.set_conditional_entry_point(lambda _: ["awhile", "alittlewhile"])
graph = builder.compile()
graph.step_timeout = 1
with pytest.raises(asyncio.TimeoutError):
async for chunk in graph.astream(1, stream_mode="updates"):
assert chunk == {"alittlewhile": {"alittlewhile": "1"}}
async for chunk in graph.astream({"hello": "world"}, stream_mode="updates"):
assert chunk == {"alittlewhile": {"hello": "1"}}
await asyncio.sleep(stream_hang_s)
assert inner_task_cancelled
@@ -1406,11 +1405,6 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
input_channels="input",
output_channels="output",
)
graph = Graph()
graph.add_node("add_one", add_one)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one")
gapp = graph.compile()
assert app.input_schema.model_json_schema() == {
"title": "LangGraphInput",
@@ -1423,21 +1417,6 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
assert await app.ainvoke(2) == 3
assert await app.ainvoke(2, output_keys=["output"]) == {"output": 3}
assert await gapp.ainvoke(2) == 3
@pytest.mark.parametrize(
"falsy_value",
[None, False, 0, "", [], {}, set(), frozenset(), 0.0, 0j],
)
async def test_invoke_single_process_in_out_falsy_values(falsy_value: Any) -> None:
graph = Graph()
graph.add_node("return_falsy_const", lambda *args, **kwargs: falsy_value)
graph.set_entry_point("return_falsy_const")
graph.set_finish_point("return_falsy_const")
gapp = graph.compile()
assert falsy_value == await gapp.ainvoke(1)
async def test_invoke_single_process_in_write_kwargs(mocker: MockerFixture) -> None:
add_one = mocker.Mock(side_effect=lambda x: x + 1)
@@ -1567,29 +1546,6 @@ async def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
}
assert step == 2
graph = Graph()
graph.add_node("add_one", add_one)
graph.add_node("add_one_more", add_one)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one_more")
graph.add_edge("add_one", "add_one_more")
gapp = graph.compile()
assert await gapp.ainvoke(2) == 4
step = 0
async for values in gapp.astream(2):
step += 1
if step == 1:
assert values == {
"add_one": 3,
}
elif step == 2:
assert values == {
"add_one_more": 4,
}
assert step == 2
async def test_batch_two_processes_in_out() -> None:
async def add_one_with_delay(inp: int) -> int:
@@ -1619,16 +1575,6 @@ async def test_batch_two_processes_in_out() -> None:
{"output": 7},
]
graph = Graph()
graph.add_node("add_one", add_one_with_delay)
graph.add_node("add_one_more", add_one_with_delay)
graph.set_entry_point("add_one")
graph.set_finish_point("add_one_more")
graph.add_edge("add_one", "add_one_more")
gapp = graph.compile()
assert await gapp.abatch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7]
async def test_invoke_many_processes_in_out(mocker: MockerFixture) -> None:
test_size = 100
@@ -2104,7 +2050,7 @@ async def test_run_from_checkpoint_id_retains_previous_writes(
thread_id = uuid.uuid4()
thread1 = {"configurable": {"thread_id": str(thread_id)}}
result = await graph.ainvoke({"myval": 1}, thread1)
result = await graph.ainvoke({"myval": 1}, thread1, checkpoint_during=True)
assert result["myval"] == 4
history = [c async for c in graph.aget_state_history(thread1)]
@@ -3061,15 +3007,7 @@ async def test_send_react_interrupt(async_checkpointer: BaseCheckpointSaver) ->
"thread_id": "2",
},
created_at=AnyStr(),
parent_config=(
{
"configurable": {
"thread_id": "2",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
}
),
parent_config=None,
tasks=(
PregelTask(
id=AnyStr(),
@@ -3193,15 +3131,7 @@ async def test_send_react_interrupt(async_checkpointer: BaseCheckpointSaver) ->
"thread_id": "3",
},
created_at=AnyStr(),
parent_config=(
{
"configurable": {
"thread_id": "3",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
}
),
parent_config=None,
tasks=(
PregelTask(
id=AnyStr(),
@@ -3466,15 +3396,7 @@ async def test_send_react_interrupt_control(
"thread_id": "2",
},
created_at=AnyStr(),
parent_config=(
{
"configurable": {
"thread_id": "2",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
}
),
parent_config=None,
tasks=(
PregelTask(
id=AnyStr(),
@@ -3710,7 +3632,7 @@ async def test_invoke_checkpoint_three(
thread_1 = {"configurable": {"thread_id": "1"}}
# total starts out as 0, so output is 0+2=2
assert await app.ainvoke(2, thread_1) == 2
assert await app.ainvoke(2, thread_1, checkpoint_during=True) == 2
state = await app.aget_state(thread_1)
assert state is not None
assert state.values.get("total") == 2
@@ -3719,7 +3641,7 @@ async def test_invoke_checkpoint_three(
== (await async_checkpointer.aget(thread_1))["id"]
)
# total is now 2, so output is 2+3=5
assert await app.ainvoke(3, thread_1) == 5
assert await app.ainvoke(3, thread_1, checkpoint_during=True) == 5
state = await app.aget_state(thread_1)
assert state is not None
assert state.values.get("total") == 7
@@ -3729,7 +3651,7 @@ async def test_invoke_checkpoint_three(
)
# total is now 2+5=7, so output would be 7+4=11, but raises ValueError
with pytest.raises(ValueError):
await app.ainvoke(4, thread_1)
await app.ainvoke(4, thread_1, checkpoint_during=True)
# checkpoint is not updated
state = await app.aget_state(thread_1)
assert state is not None
@@ -3737,7 +3659,7 @@ async def test_invoke_checkpoint_three(
assert state.next == ("one",)
"""we checkpoint inputs and it failed on "one", so the next node is one"""
# we can recover from error by sending new inputs
assert await app.ainvoke(2, thread_1) == 9
assert await app.ainvoke(2, thread_1, checkpoint_during=True) == 9
state = await app.aget_state(thread_1)
assert state is not None
assert state.values.get("total") == 16, "total is now 7+9=16"
@@ -3746,7 +3668,7 @@ async def test_invoke_checkpoint_three(
thread_2 = {"configurable": {"thread_id": "2"}}
# on a new thread, total starts out as 0, so output is 0+5=5
assert await app.ainvoke(5, thread_2) == 5
state = await app.aget_state({"configurable": {"thread_id": "1"}})
state = await app.aget_state(thread_1)
assert state is not None
assert state.values.get("total") == 16
assert state.next == ()
@@ -3893,42 +3815,6 @@ async def test_invoke_two_processes_no_out(mocker: MockerFixture) -> None:
assert await app.ainvoke(2) is None
async def test_conditional_entrypoint_graph() -> None:
async def left(data: str) -> str:
return data + "->left"
async def right(data: str) -> str:
return data + "->right"
def should_start(data: str) -> str:
# Logic to decide where to start
if len(data) > 10:
return "go-right"
else:
return "go-left"
# Define a new graph
workflow = Graph()
workflow.add_node("left", left)
workflow.add_node("right", right)
workflow.set_conditional_entry_point(
should_start, {"go-left": "left", "go-right": "right"}
)
workflow.add_conditional_edges("left", lambda data: END)
workflow.add_edge("right", END)
app = workflow.compile()
assert await app.ainvoke("what is weather in sf") == "what is weather in sf->right"
assert [c async for c in app.astream("what is weather in sf")] == [
{"right": "what is weather in sf->right"},
]
async def test_conditional_entrypoint_graph_state() -> None:
class AgentState(TypedDict, total=False):
input: str
@@ -4156,97 +4042,6 @@ async def test_in_one_fan_out_state_graph_defer_node(
]
@pytest.mark.parametrize("with_path_map", (True, False))
async def test_in_one_fan_out_state_graph_then_defer_node(
async_checkpointer: BaseCheckpointSaver, with_path_map: bool
) -> None:
def sorted_add(
x: list[str], y: Union[list[str], list[tuple[str, str]]]
) -> list[str]:
if isinstance(y[0], tuple):
for rem, _ in y:
x.remove(rem)
y = [t[1] for t in y]
return sorted(operator.add(x, y))
class State(TypedDict, total=False):
query: str
answer: str
docs: Annotated[list[str], sorted_add]
async def rewrite_query(data: State) -> State:
return {"query": f"query: {data['query']}"}
async def analyzer_one(data: State) -> State:
return {"query": f"analyzed: {data['query']}"}
async def retriever_one(data: State) -> State:
return {"docs": ["doc1", "doc2"]}
async def retriever_two(data: State) -> State:
await asyncio.sleep(0.1)
return {"docs": ["doc3", "doc4"]}
async def qa(data: State) -> State:
return {"answer": ",".join(data["docs"])}
workflow = StateGraph(State)
workflow.add_node("rewrite_query", rewrite_query)
workflow.add_node("analyzer_one", analyzer_one)
workflow.add_node("retriever_one", retriever_one)
workflow.add_node("retriever_two", retriever_two)
workflow.add_node("qa", qa, defer=True)
workflow.set_entry_point("rewrite_query")
workflow.add_conditional_edges(
"rewrite_query",
lambda _: ["analyzer_one", "retriever_two"],
["analyzer_one", "retriever_two"] if with_path_map else None,
then="qa",
)
workflow.add_edge("analyzer_one", "retriever_one")
app = workflow.compile()
assert await app.ainvoke({"query": "what is weather in sf"}, debug=True) == {
"query": "analyzed: query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"],
"answer": "doc1,doc2,doc3,doc4",
}
assert [c async for c in app.astream({"query": "what is weather in sf"})] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
app_w_interrupt = workflow.compile(
checkpointer=async_checkpointer,
interrupt_after=["retriever_one"],
)
config = {"configurable": {"thread_id": "1"}}
assert [
c
async for c in app_w_interrupt.astream(
{"query": "what is weather in sf"}, config
)
] == [
{"rewrite_query": {"query": "query: what is weather in sf"}},
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c async for c in app_w_interrupt.astream(None, config)] == [
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
]
async def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
async_checkpointer: BaseCheckpointSaver,
) -> None:
@@ -6005,7 +5800,7 @@ async def test_debug_retry(async_checkpointer: BaseCheckpointSaver):
graph = builder.compile(checkpointer=async_checkpointer)
config = {"configurable": {"thread_id": "1"}}
await graph.ainvoke({"messages": []}, config=config)
await graph.ainvoke({"messages": []}, config=config, checkpoint_during=True)
# re-run step: 1
async for c in async_checkpointer.alist(config):
@@ -6017,7 +5812,10 @@ async def test_debug_retry(async_checkpointer: BaseCheckpointSaver):
update_config = await graph.aupdate_state(target_config, values=None)
events = [
c async for c in graph.astream(None, config=update_config, stream_mode="debug")
c
async for c in graph.astream(
None, config=update_config, stream_mode="debug", checkpoint_during=True
)
]
checkpoint_events = list(
@@ -6048,7 +5846,9 @@ async def test_debug_retry(async_checkpointer: BaseCheckpointSaver):
assert stream_parent_conf == history_parent_conf
async def test_debug_subgraphs(async_checkpointer: BaseCheckpointSaver):
async def test_debug_subgraphs(
async_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
):
class State(TypedDict):
messages: Annotated[list[str], operator.add]
@@ -6082,12 +5882,15 @@ async def test_debug_subgraphs(async_checkpointer: BaseCheckpointSaver):
{"messages": []},
config=config,
stream_mode="debug",
checkpoint_during=checkpoint_during,
)
]
checkpoint_events = list(
reversed([e["payload"] for e in events if e["type"] == "checkpoint"])
)
if not checkpoint_during:
checkpoint_events = checkpoint_events[:1]
checkpoint_history = [c async for c in graph.aget_state_history(config)]
assert len(checkpoint_events) == len(checkpoint_history)
@@ -6114,7 +5917,9 @@ async def test_debug_subgraphs(async_checkpointer: BaseCheckpointSaver):
assert stream_task.get("state") == history_task.state
async def test_debug_nested_subgraphs(async_checkpointer: BaseCheckpointSaver):
async def test_debug_nested_subgraphs(
async_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
from collections import defaultdict
class State(TypedDict):
@@ -6158,6 +5963,7 @@ async def test_debug_nested_subgraphs(async_checkpointer: BaseCheckpointSaver):
config=config,
stream_mode="debug",
subgraphs=True,
checkpoint_during=checkpoint_during,
)
]
@@ -6202,6 +6008,9 @@ async def test_debug_nested_subgraphs(async_checkpointer: BaseCheckpointSaver):
for checkpoint_events, checkpoint_history in zip(
stream_ns.values(), history_ns.values()
):
if not checkpoint_during:
checkpoint_events = checkpoint_events[-1:]
assert len(checkpoint_events) == len(checkpoint_history)
for stream, history in zip(checkpoint_events, checkpoint_history):
assert stream["values"] == history.values
assert stream["next"] == list(history.next)
@@ -6221,7 +6030,10 @@ async def test_debug_nested_subgraphs(async_checkpointer: BaseCheckpointSaver):
assert stream_task.get("state") == history_task.state
async def test_parent_command(async_checkpointer: BaseCheckpointSaver) -> None:
@pytest.mark.parametrize("subgraph_persist", [True, False])
async def test_parent_command(
async_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
from langchain_core.messages import BaseMessage
from langchain_core.tools import tool
@@ -6233,7 +6045,7 @@ async def test_parent_command(async_checkpointer: BaseCheckpointSaver) -> None:
subgraph_builder = StateGraph(MessagesState)
subgraph_builder.add_node("tool", get_user_name)
subgraph_builder.add_edge(START, "tool")
subgraph = subgraph_builder.compile()
subgraph = subgraph_builder.compile(checkpointer=subgraph_persist)
class CustomParentState(TypedDict):
messages: Annotated[list[BaseMessage], add_messages]
@@ -6282,15 +6094,7 @@ async def test_parent_command(async_checkpointer: BaseCheckpointSaver) -> None:
"parents": {},
},
created_at=AnyStr(),
parent_config=(
{
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
}
),
parent_config=None,
tasks=(),
interrupts=(),
)
@@ -6772,7 +6576,7 @@ async def test_concurrent_execution():
async def test_checkpoint_recovery_async(
async_checkpointer: BaseCheckpointSaver,
async_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
"""Test recovery from checkpoints after failures with async nodes."""
@@ -6802,7 +6606,11 @@ async def test_checkpoint_recovery_async(
# First attempt should fail
with pytest.raises(RuntimeError):
await graph.ainvoke({"steps": ["start"], "attempt": 1}, config)
await graph.ainvoke(
{"steps": ["start"], "attempt": 1},
config,
checkpoint_during=checkpoint_during,
)
# Verify checkpoint state
state = await graph.aget_state(config)
@@ -6811,12 +6619,17 @@ async def test_checkpoint_recovery_async(
assert state.next == ("node1",) # Should retry failed node
# Retry with updated attempt count
result = await graph.ainvoke({"steps": [], "attempt": 2}, config)
result = await graph.ainvoke(
{"steps": [], "attempt": 2}, config, checkpoint_during=checkpoint_during
)
assert result == {"steps": ["start", "node1", "node2"], "attempt": 2}
# Verify checkpoint history shows both attempts
history = [c async for c in graph.aget_state_history(config)]
assert len(history) == 6 # Initial + failed attempt + successful attempt
if checkpoint_during:
assert len(history) == 6 # Initial + failed attempt + successful attempt
else:
assert len(history) == 2 # error + success
# Verify the error was recorded in checkpoint
failed_checkpoint = next(c for c in history if c.tasks and c.tasks[0].error)
@@ -8291,7 +8104,9 @@ async def test_bulk_state_updates(async_checkpointer: BaseCheckpointSaver) -> No
)
async def test_update_as_input(async_checkpointer: BaseCheckpointSaver) -> None:
async def test_update_as_input(
async_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
class State(TypedDict):
foo: str
@@ -8311,11 +8126,15 @@ async def test_update_as_input(async_checkpointer: BaseCheckpointSaver) -> None:
)
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {"foo": "tool"}
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {"foo": "tool"}
def map_snapshot(i: StateSnapshot) -> dict:
@@ -8354,10 +8173,15 @@ async def test_update_as_input(async_checkpointer: BaseCheckpointSaver) -> None:
async for s in graph.aget_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
if checkpoint_during:
assert new_history == history
else:
assert [new_history[0], new_history[4]] == history
async def test_batch_update_as_input(async_checkpointer: BaseCheckpointSaver) -> None:
async def test_batch_update_as_input(
async_checkpointer: BaseCheckpointSaver, checkpoint_during: bool
) -> None:
class State(TypedDict):
foo: str
tasks: Annotated[list[int], operator.add]
@@ -8389,7 +8213,9 @@ async def test_batch_update_as_input(async_checkpointer: BaseCheckpointSaver) ->
)
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
{"foo": "input"},
{"configurable": {"thread_id": "1"}},
checkpoint_during=checkpoint_during,
) == {"foo": "map", "tasks": [0, 1, 2]}
def map_snapshot(i: StateSnapshot) -> dict:
@@ -8440,7 +8266,10 @@ async def test_batch_update_as_input(async_checkpointer: BaseCheckpointSaver) ->
async for s in graph.aget_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
if checkpoint_during:
assert new_history == history
else:
assert new_history[:1] == history
async def test_draw_invalid():
@@ -8824,3 +8653,43 @@ async def test_imp_exception(
"parent_ids": [],
},
]
@pytest.mark.parametrize("subgraph_persist", [True, False])
async def test_parent_command_goto(
async_checkpointer: BaseCheckpointSaver, subgraph_persist: bool
) -> None:
class State(TypedDict):
dialog_state: Annotated[list[str], operator.add]
async def node_a_child(state):
return {"dialog_state": ["a_child_state"]}
async def node_b_child(state):
return Command(
graph=Command.PARENT,
goto="node_b_parent",
update={"dialog_state": ["b_child_state"]},
)
sub_builder = StateGraph(State)
sub_builder.add_node(node_a_child)
sub_builder.add_node(node_b_child)
sub_builder.add_edge(START, "node_a_child")
sub_builder.add_edge("node_a_child", "node_b_child")
sub_graph = sub_builder.compile(checkpointer=subgraph_persist)
async def node_b_parent(state):
return {"dialog_state": ["node_b_parent"]}
main_builder = StateGraph(State)
main_builder.add_node(node_b_parent)
main_builder.add_edge(START, "subgraph_node")
main_builder.add_node("subgraph_node", sub_graph, destinations=("node_b_parent",))
main_graph = main_builder.compile(async_checkpointer, name="parent")
config = {"configurable": {"thread_id": 1}}
assert await main_graph.ainvoke(
input={"dialog_state": ["init_state"]}, config=config
) == {"dialog_state": ["init_state", "b_child_state", "node_b_parent"]}
+6 -4
View File
@@ -18,7 +18,7 @@ import pytest
from typing_extensions import NotRequired, Required, TypedDict
from langgraph.graph import END, StateGraph
from langgraph.graph.graph import CompiledGraph
from langgraph.graph.state import CompiledStateGraph
from langgraph.utils.config import _is_not_empty
from langgraph.utils.fields import (
_is_optional_type,
@@ -103,7 +103,7 @@ def test_is_generator() -> None:
@pytest.fixture
def rt_graph() -> CompiledGraph:
def rt_graph() -> CompiledStateGraph:
class State(TypedDict):
foo: int
node_run_id: int
@@ -120,7 +120,7 @@ def rt_graph() -> CompiledGraph:
return graph.compile()
def test_runnable_callable_tracing_nested(rt_graph: CompiledGraph) -> None:
def test_runnable_callable_tracing_nested(rt_graph: CompiledStateGraph) -> None:
with patch("langsmith.client.Client", spec=langsmith.Client) as mock_client:
with patch("langchain_core.tracers.langchain.get_client") as mock_get_client:
mock_get_client.return_value = mock_client
@@ -133,7 +133,9 @@ def test_runnable_callable_tracing_nested(rt_graph: CompiledGraph) -> None:
sys.version_info < (3, 11),
reason="Python 3.11+ is required for async contextvars support",
)
async def test_runnable_callable_tracing_nested_async(rt_graph: CompiledGraph) -> None:
async def test_runnable_callable_tracing_nested_async(
rt_graph: CompiledStateGraph,
) -> None:
with patch("langsmith.client.Client", spec=langsmith.Client) as mock_client:
with patch("langchain_core.tracers.langchain.get_client") as mock_get_client:
mock_get_client.return_value = mock_client
+387 -358
View File
File diff suppressed because it is too large Load Diff
@@ -36,8 +36,8 @@ from typing_extensions import Annotated, TypedDict
from langgraph.errors import ErrorCode, create_error_message
from langgraph.graph import END, StateGraph
from langgraph.graph.graph import CompiledGraph
from langgraph.graph.message import add_messages
from langgraph.graph.state import CompiledStateGraph
from langgraph.managed import IsLastStep, RemainingSteps
from langgraph.prebuilt.tool_node import ToolNode
from langgraph.store.base import BaseStore
@@ -257,7 +257,7 @@ def create_react_agent(
debug: bool = False,
version: Literal["v1", "v2"] = "v2",
name: Optional[str] = None,
) -> CompiledGraph:
) -> CompiledStateGraph:
"""Creates an agent graph that calls tools in a loop until a stopping condition is met.
For more details on using `create_react_agent`, visit [Agents](https://langchain-ai.github.io/langgraph/agents/overview/) documentation.
@@ -629,7 +629,7 @@ def tools_condition(
Args:
state: The state to check for
tool calls. Must have a list of messages (MessageGraph) or have the
tool calls. Must have a list of messages or have the
"messages" key (StateGraph).
Returns:
@@ -2,7 +2,7 @@
in a langchain graph. It applies a pydantic schema to tool_calls in the models' outputs,
and returns a ToolMessage with the validated content. If the schema is not valid, it
returns a ToolMessage with the error message. The ValidationNode can be used in a
StateGraph with a "messages" key or in a MessageGraph. If multiple tool calls are
StateGraph with a "messages" key. If multiple tool calls are
requested, they will be run in parallel.
"""
@@ -49,7 +49,7 @@ def _default_format_error(
class ValidationNode(RunnableCallable):
"""A node that validates all tools requests from the last AIMessage.
It can be used either in StateGraph with a "messages" key or in MessageGraph.
It can be used in StateGraph with a "messages" key.
!!! note
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "@langchain/langgraph-sdk",
"version": "0.0.78",
"version": "0.0.82",
"description": "Client library for interacting with the LangGraph API",
"type": "module",
"packageManager": "yarn@1.22.19",
+24 -3
View File
@@ -86,12 +86,18 @@ function getRunMetadataFromResponse(
};
}
export type RequestHook = (
url: URL,
init: RequestInit,
) => Promise<RequestInit> | RequestInit;
export interface ClientConfig {
apiUrl?: string;
apiKey?: string;
callerOptions?: AsyncCallerParams;
timeoutMs?: number;
defaultHeaders?: Record<string, string | null | undefined>;
onRequest?: RequestHook;
}
class BaseClient {
@@ -103,6 +109,8 @@ class BaseClient {
protected defaultHeaders: Record<string, string | null | undefined>;
protected onRequest?: RequestHook;
constructor(config?: ClientConfig) {
const callerOptions = {
maxRetries: 4,
@@ -136,6 +144,7 @@ class BaseClient {
// Regex to remove trailing slash, if present
this.apiUrl = config?.apiUrl?.replace(/\/$/, "") || defaultApiUrl;
this.defaultHeaders = config?.defaultHeaders || {};
this.onRequest = config?.onRequest;
const apiKey = getApiKey(config?.apiKey);
if (apiKey) {
this.defaultHeaders["X-Api-Key"] = apiKey;
@@ -230,9 +239,14 @@ class BaseClient {
withResponse?: boolean;
},
): Promise<T | [T, Response]> {
const response = await this.asyncCaller.fetch(
...this.prepareFetchOptions(path, options),
);
const [url, init] = this.prepareFetchOptions(path, options);
let finalInit = init;
if (this.onRequest) {
finalInit = await this.onRequest(url, init);
}
const response = await this.asyncCaller.fetch(url, finalInit);
const body = (() => {
if (response.status === 202 || response.status === 204) {
@@ -965,6 +979,7 @@ export class RunsClient<
after_seconds: payload?.afterSeconds,
if_not_exists: payload?.ifNotExists,
checkpoint_during: payload?.checkpointDuring,
langsmith_tracer: payload?._langsmithTracer,
};
const [run, response] = await this.fetch<Run>(`/threads/${threadId}/runs`, {
@@ -1045,6 +1060,12 @@ export class RunsClient<
after_seconds: payload?.afterSeconds,
if_not_exists: payload?.ifNotExists,
checkpoint_during: payload?.checkpointDuring,
langsmith_tracer: payload?._langsmithTracer
? {
project_name: payload?._langsmithTracer?.projectName,
example_id: payload?._langsmithTracer?.exampleId,
}
: undefined,
};
const endpoint =
threadId == null ? `/runs/wait` : `/threads/${threadId}/runs/wait`;
+37 -32
View File
@@ -1,51 +1,56 @@
export { Client } from "./client.js";
export { Client, getApiKey } from "./client.js";
export type { ClientConfig, RequestHook } from "./client.js";
export type {
AssistantBase,
Assistant,
AssistantVersion,
AssistantBase,
AssistantGraph,
AssistantVersion,
Checkpoint,
Config,
Cron,
CronCreateForThreadResponse,
CronCreateResponse,
DefaultValues,
GraphSchema,
Interrupt,
Item,
ListNamespaceResponse,
Metadata,
Run,
Thread,
ThreadTask,
ThreadState,
ThreadStatus,
Cron,
Checkpoint,
Interrupt,
ListNamespaceResponse,
Item,
SearchItem,
SearchItemsResponse,
CronCreateResponse,
CronCreateForThreadResponse,
Thread,
ThreadState,
ThreadStatus,
ThreadTask,
} from "./schema.js";
export { overrideFetchImplementation } from "./singletons/fetch.js";
export type { OnConflictBehavior, Command } from "./types.js";
export type { StreamMode } from "./types.stream.js";
export type {
ValuesStreamEvent,
Command,
OnConflictBehavior,
RunsInvokePayload,
} from "./types.js";
export type {
AIMessage,
FunctionMessage,
HumanMessage,
Message,
RemoveMessage,
SystemMessage,
ToolMessage,
} from "./types.messages.js";
export type {
CustomStreamEvent,
DebugStreamEvent,
ErrorStreamEvent,
EventsStreamEvent,
FeedbackStreamEvent,
MessagesStreamEvent,
MessagesTupleStreamEvent,
MetadataStreamEvent,
StreamMode,
UpdatesStreamEvent,
CustomStreamEvent,
MessagesStreamEvent,
DebugStreamEvent,
EventsStreamEvent,
ErrorStreamEvent,
FeedbackStreamEvent,
ValuesStreamEvent,
} from "./types.stream.js";
export type {
Message,
HumanMessage,
AIMessage,
ToolMessage,
SystemMessage,
FunctionMessage,
RemoveMessage,
} from "./types.messages.js";
+7 -4
View File
@@ -7,7 +7,7 @@ export interface UIMessage<
id: string;
name: TName;
props: TProps;
metadata: {
metadata?: {
merge?: boolean;
run_id?: string;
name?: string;
@@ -51,9 +51,12 @@ export function uiMessageReducer(
const index = state.findIndex((ui) => ui.id === event.id);
if (index !== -1) {
newState[index] = event.metadata.merge
? { ...event, props: { ...state[index].props, ...event.props } }
: event;
newState[index] =
typeof event.metadata === "object" &&
event.metadata != null &&
event.metadata.merge
? { ...event, props: { ...state[index].props, ...event.props } }
: event;
} else {
newState.push(event);
}
+30 -11
View File
@@ -405,6 +405,11 @@ type GetCustomEventType<Bag extends BagTemplate> = Bag extends {
? Bag["CustomEventType"]
: unknown;
interface RunCallbackMeta {
run_id: string;
thread_id: string;
}
export interface UseStreamOptions<
StateType extends Record<string, unknown> = Record<string, unknown>,
Bag extends BagTemplate = BagTemplate,
@@ -450,17 +455,20 @@ export interface UseStreamOptions<
/**
* Callback that is called when an error occurs.
*/
onError?: (error: unknown) => void;
onError?: (error: unknown, run: RunCallbackMeta | undefined) => void;
/**
* Callback that is called when the stream is finished.
*/
onFinish?: (state: ThreadState<StateType>) => void;
onFinish?: (
state: ThreadState<StateType>,
run: RunCallbackMeta | undefined,
) => void;
/**
* Callback that is called when a new stream is created.
*/
onCreated?: (run: { run_id: string; thread_id: string }) => void;
onCreated?: (run: RunCallbackMeta) => void;
/**
* Callback that is called when an update event is received.
@@ -846,8 +854,12 @@ export function useStream<
action: (signal: AbortSignal) => Promise<{
onSuccess: () => Promise<ThreadState<StateType>[]>;
stream: AsyncGenerator<EventStreamEvent>;
getCallbackMeta: () => { thread_id: string; run_id: string } | undefined;
}>,
) {
let getCallbackMeta:
| (() => { thread_id: string; run_id: string } | undefined)
| undefined;
try {
setIsLoading(true);
setStreamError(undefined);
@@ -856,6 +868,7 @@ export function useStream<
abortRef.current = new AbortController();
const run = await action(abortRef.current.signal);
getCallbackMeta = run.getCallbackMeta;
let streamError: StreamError | undefined;
for await (const { event, data } of run.stream) {
@@ -916,7 +929,7 @@ export function useStream<
if (streamError != null) throw streamError;
const lastHead = result.at(0);
if (lastHead) onFinish?.(lastHead);
if (lastHead) onFinish?.(lastHead, getCallbackMeta?.());
} catch (error) {
if (
!(
@@ -926,7 +939,7 @@ export function useStream<
) {
console.error(error);
setStreamError(error);
onError?.(error);
onError?.(error, getCallbackMeta?.());
}
} finally {
setIsLoading(false);
@@ -953,6 +966,7 @@ export function useStream<
return history.mutate(threadId);
},
stream,
getCallbackMeta: () => ({ thread_id: threadId, run_id: runId }),
};
});
};
@@ -1004,6 +1018,9 @@ export function useStream<
// @ts-expect-error
if (checkpoint != null) delete checkpoint.thread_id;
let rejoinKey: `lg:stream:${string}` | undefined;
let callbackMeta: RunCallbackMeta | undefined;
const streamResumable =
submitOptions?.streamResumable ?? !!runMetadataStorage;
const stream = client.runs.stream(usableThreadId, assistantId, {
input: values as Record<string, unknown>,
@@ -1017,29 +1034,31 @@ export function useStream<
onCompletion: submitOptions?.onCompletion,
onDisconnect:
submitOptions?.onDisconnect ??
(runMetadataStorage ? "continue" : "cancel"),
(streamResumable ? "continue" : "cancel"),
signal,
checkpoint,
streamMode,
streamSubgraphs: submitOptions?.streamSubgraphs,
streamResumable: submitOptions?.streamResumable ?? !!runMetadataStorage,
streamResumable,
onRunCreated(params) {
const runParams = {
callbackMeta = {
run_id: params.run_id,
thread_id: params.thread_id ?? usableThreadId,
};
if (runMetadataStorage) {
rejoinKey = `lg:stream:${runParams.thread_id}`;
runMetadataStorage.setItem(rejoinKey, runParams.run_id);
rejoinKey = `lg:stream:${callbackMeta.thread_id}`;
runMetadataStorage.setItem(rejoinKey, callbackMeta.run_id);
}
onCreated?.(runParams);
onCreated?.(callbackMeta);
},
}) as AsyncGenerator<EventStreamEvent>;
return {
stream,
getCallbackMeta: () => callbackMeta,
onSuccess: () => {
if (rejoinKey) runMetadataStorage?.removeItem(rejoinKey);
return history.mutate(usableThreadId);
+8 -1
View File
@@ -1,3 +1,4 @@
import { LangChainTracer } from "@langchain/core/tracers/tracer_langchain";
import { Checkpoint, Config, Metadata } from "./schema.js";
import { StreamMode } from "./types.stream.js";
@@ -41,7 +42,7 @@ export interface Command {
goto?: Send | Send[] | string | string[];
}
interface RunsInvokePayload {
export interface RunsInvokePayload {
/**
* Input to the run. Pass `null` to resume from the current state of the thread.
*/
@@ -140,6 +141,12 @@ interface RunsInvokePayload {
* Callback when a run is created.
*/
onRunCreated?: (params: { run_id: string; thread_id?: string }) => void;
/**
* @internal
* For LangSmith tracing purposes only. Not part of the public API.
*/
_langsmithTracer?: LangChainTracer;
}
export interface RunsStreamPayload<