mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
Merge branch 'main' into david/05-28/support-image-distro-config
This commit is contained in:
@@ -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"
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
@@ -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).
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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).
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 %}
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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, ""
|
||||
),
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)}. ")
|
||||
|
||||
@@ -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
@@ -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"]
|
||||
}
|
||||
|
||||
@@ -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"]}
|
||||
|
||||
@@ -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
|
||||
|
||||
Generated
+387
-358
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,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",
|
||||
|
||||
@@ -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
@@ -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,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);
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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<
|
||||
|
||||
Reference in New Issue
Block a user