Compare commits

..
106 Commits
Author SHA1 Message Date
Vadym BardaandGitHub 8e922b859c prebuilt: release 0.1.4 (#3977) 2025-03-21 12:05:15 -04:00
Vadym BardaandGitHub 9a4c30135f prebuilt: pass last message to structured response model in create_react_agent (#3976) 2025-03-21 11:58:27 -04:00
William FHandGitHub bc6651f34c Studio cli command (#3962) 2025-03-20 17:57:56 -07:00
William Fu-Hinthorn 23bb5369b9 Studio cli command 2025-03-20 17:51:08 -07:00
Nuno CamposandGitHub 85b81371f7 Use incremental storage in memory checkpointer (#3960)
- This makes our checkpoint benchmarks more closely resemble the
behavior of our prod checkpointers
- Also found and fixed a bug w multiple subgraphs in same node
accidentally sharing checkpoints
2025-03-20 16:33:47 -07:00
Nuno Campos 54ab833b74 Lint 2025-03-20 16:26:52 -07:00
Nuno Campos 4ea936eaf4 Lint 2025-03-20 16:24:57 -07:00
Nuno Campos 9d81ec9ffd Lint 2025-03-20 16:22:24 -07:00
Nuno Campos 46b652a74c Use incremental storage in memory checkpointer
- This makes our checkpoint benchmarks more closely resemble the behavior of our prod checkpointers
- Also found and fixed a bug w multiple subgraphs in same node accidentally sharing checkpoints
2025-03-20 16:11:36 -07:00
Vadym BardaandGitHub 7013ca9a3f docs: use hosted logo (#3959) 2025-03-20 18:17:22 -04:00
William FHandGitHub 0c04aec664 Include enum in check for pydantic state (#3955) 2025-03-20 12:03:22 -07:00
Eugene YurtsevandGitHub 1650c8508e benchmark: Add compilation only (#3932)
Add compilation benchmark alone
2025-03-20 14:51:31 -04:00
Eugene YurtsevandGitHub e176b98fe7 Add llms-txt resources (#3935) 2025-03-20 14:44:21 -04:00
William Fu-Hinthorn eb1e1aa010 Include enum in check for pydantic state 2025-03-20 10:29:16 -07:00
Nuno CamposandGitHub 77c833e1e5 Use fast path for prepare_next_tasks on input (#3931)
- When there are no values in checkpoint no need to run through all the
PULL candidates
- When there are input writes save updated_channels to use on the next
call to prepare_next_tasks
2025-03-20 08:46:48 -07:00
Nuno Campos 0ac29434a7 Lint 2025-03-20 08:40:05 -07:00
Nuno Campos 43f5a17416 Lint 2025-03-20 08:24:51 -07:00
Nuno Campos 7d0857f263 Lint 2025-03-20 08:24:31 -07:00
Nuno Campos b82d70a66a Lint 2025-03-20 08:22:16 -07:00
Nuno CamposandGitHub 5fb037171d Small perf improvements (#3949)
- RunnableCallable: Skip signature checks for internal callables where
we know the signatures ahead of time
- PregelNode: Avoid redoing subgraphs serarch when copying it
- CompiledStateGraph: Avoid copying PregelNode when attaching writers
2025-03-20 08:19:02 -07:00
Nuno Campos d3bb2b9aa0 Lint 2025-03-20 08:18:17 -07:00
Nuno Campos ea765b4134 More small perf improvements
- RunnableCallable: Skip signature checks for internal callables where we know the signatures ahead of time
- PregelNode: Avoid redoing subgraphs serarch when copying it
- CompiledStateGraph: Avoid copying PregelNode when attaching writers
2025-03-20 08:11:28 -07:00
William FHandGitHub 66ff83dca9 Lock (#3947) 2025-03-20 08:04:53 -07:00
William FHandGitHub 254e398345 Merge branch 'main' into wfh/reloack 2025-03-20 08:04:38 -07:00
Vadym BardaandGitHub c7567ea219 docs: improve search (#3948) 2025-03-20 11:02:59 -04:00
William Fu-Hinthorn 8c0306c3f4 Lock 2025-03-20 07:59:30 -07:00
William FHandGitHub 992b05a196 langgraph-checkpoint-postgres 2.0.19 (#3945) 2025-03-20 07:25:41 -07:00
William Fu-Hinthorn 893a9646d3 langgraph-checkpoint-postgres 2.0.19 2025-03-20 07:25:14 -07:00
William FHandGitHub eaa37a2ce9 Increase pg->checkpoint minbound (#3944) 2025-03-20 07:24:43 -07:00
William Fu-Hinthorn daee8d88bb Increase pg->checkpoint minbound 2025-03-20 07:24:15 -07:00
Nuno Campos eaa18cc2dd Use fast path for prepare_next_tasks on input
- When there are no values in checkpoint no need to run through all the PULL candidates
- When there are input writes save updated_channels to use on the next call to prepare_next_tasks
2025-03-19 18:14:04 -07:00
Nuno CamposandGitHub b2d9a36308 langgraph: incorporate information about previously updated channels to identify which tasks to execute next (#3916)
Leverage information about which channels were updated in the previous
step to determine which tasks should be triggered. This can result in
significant speed up in prepare_next_tasks in some situations.
2025-03-19 16:22:12 -07:00
William FHandGitHub 03fc695d60 Add refcount test (#3910) 2025-03-19 14:33:43 -07:00
William Fu-Hinthorn 9994b09304 merge 2025-03-19 14:27:19 -07:00
William FHandGitHub 53f8558914 Release 0.3.18 (#3925)
Includes:
- Explicit unsetting of runnable context var
- Weakref for PregelExecutableTask

both to reduce the chance of keeping a reference to an internal object
and preventing garbage collection
2025-03-19 14:11:46 -07:00
William Fu-Hinthorn 4bfcd84cee Cleanup ref count check 2025-03-19 14:10:01 -07:00
David DuongandGitHub cd1d7be05f feat(sdk): add bulk_update_state in SDK (#3923) 2025-03-19 22:09:53 +01:00
Really HimandGitHub 939a426a2e DOCS: Update state-model.ipynb to use "AnyMessage" (#3926)
## Description
The documentation for working with Pydantic and graph State recommends
to use `AnyMessage` when working with LangChain types, but the code
example uses `BaseMessage`.
2025-03-19 21:08:14 +00:00
William Fu-Hinthorn 94c815f226 Release 0.3.18
Includes:
- Explicit unsetting of runnable context var
- Weakref for PregelExecutableTask

both to reduce the chance of keeping a reference to an internal object and preventing
garbage collection
2025-03-19 14:04:01 -07:00
Tat Dat Duong ef345aac5f Fix typo 2025-03-19 22:02:55 +01:00
William FHandGitHub e306258525 Reference to PregelExecutableTask (#3924) 2025-03-19 14:00:43 -07:00
William Fu-Hinthorn 05cd317486 Update snapshots more 2025-03-19 13:54:10 -07:00
Tat Dat Duong 6dfed31a5e Update parameters 2025-03-19 21:39:35 +01:00
ThaparandGitHub ee8374c4c0 docs: Update bad link in libs/cli README.md (#3919)
Fixed reference hyperlink
2025-03-19 16:27:37 -04:00
William Fu-Hinthorn 2066894b5f Update snapshots 2025-03-19 13:25:02 -07:00
William Fu-Hinthorn 6a2d20fd5b Reference to PregelExecutableTask 2025-03-19 13:19:02 -07:00
Eugene Yurtsev 5b8b9f1067 Update doc-string 2025-03-19 16:10:50 -04:00
Eugene Yurtsev c9cb8165d4 x 2025-03-19 16:08:25 -04:00
Tat Dat Duong 1f0348a5ca Fix docstring 2025-03-19 21:07:07 +01:00
Tat Dat Duong 442ef0788e Add graph_id back 2025-03-19 21:05:43 +01:00
Eugene Yurtsev 73f9ef0ef8 add type 2025-03-19 16:00:49 -04:00
Tat Dat Duong 1f7a380548 Fix typo 2025-03-19 20:54:15 +01:00
Eugene Yurtsev 8959f2aec5 lint 2025-03-19 15:54:07 -04:00
Eugene Yurtsev 0e7869eba4 Merge branch 'main' into ey/optimize_triggers 2025-03-19 15:52:29 -04:00
William FHandGitHub d3f8478054 Unset config context after function end (#3922) 2025-03-19 12:51:39 -07:00
Tat Dat Duong 0cb1893475 Update for JS as well 2025-03-19 20:46:22 +01:00
Tat Dat Duong e779c8e0b1 Merge into create 2025-03-19 20:42:48 +01:00
Eugene Yurtsev 18b82cb8e2 x 2025-03-19 15:29:50 -04:00
Tat Dat Duong 972ab1a935 Revert docstring for update_state 2025-03-19 20:22:05 +01:00
William Fu-Hinthorn 9cc2f37cca Unset config context after function end 2025-03-19 12:14:29 -07:00
Tat Dat Duong a2d7631f47 Fix in async client 2025-03-19 20:10:42 +01:00
Tat Dat Duong 1e767c0653 feat(sdk): add bulk_update_state in SDK 2025-03-19 20:05:08 +01:00
David DuongandGitHub d4c569cb7c feat(sdk-js): add bulkUpdateState method (#3878) 2025-03-19 19:25:37 +01:00
David DuongandGitHub a146df7f6a release(langgraph): 0.3.17 (#3918) 2025-03-19 19:10:01 +01:00
Tat Dat Duong c52cc03e4b release(langgraph): 0.3.17 2025-03-19 19:02:51 +01:00
Nuno CamposandGitHub f206cfad8f Store all triggers in task (#3912)
- These are used to update seen version
2025-03-19 09:47:06 -07:00
Nuno Campos d4c8b219c4 Update tests 2025-03-19 09:40:34 -07:00
Tat Dat Duong 00855999d2 Bump to 0.0.59 2025-03-19 17:02:19 +01:00
Tat Dat Duong 24bd0e1c1f Fix formatting 2025-03-19 16:59:53 +01:00
David DuongandGitHub 3b59055192 feat(pregel): add bulk update state method (#3737)
This method is useful for recreating a thread from a list of checkpoint
writes. A new method is needed to clone a checkpoint that has been
created from multiple writes (functional API, map-reduce)

Port of https://github.com/langchain-ai/langgraphjs/pull/969 and
https://github.com/langchain-ai/langgraphjs/pull/1007
2025-03-19 16:59:07 +01:00
Eugene Yurtsev 67a16bec53 more typos 2025-03-19 11:48:31 -04:00
Eugene Yurtsev f97802eed5 x 2025-03-19 11:44:42 -04:00
Eugene Yurtsev bccb796ccc x 2025-03-19 11:38:46 -04:00
Eugene Yurtsev 4068e9d135 x 2025-03-19 11:27:56 -04:00
Eugene Yurtsev a951334f7f qxqx 2025-03-19 11:18:00 -04:00
lc-arjunandGitHub 75a727877f feat: add docs for prompt engineering (#3846) 2025-03-19 10:51:54 -04:00
Tat Dat Duong 1cece3228c Fix indent bug 2025-03-19 15:37:43 +01:00
Tat Dat Duong 7ba48d75c9 Another merge issue 2025-03-19 15:05:22 +01:00
Tat Dat Duong 7fb0628957 Remove duplicated test 2025-03-19 14:44:07 +01:00
Tat Dat Duong 6edf29f043 Fix rebase artifacts 2025-03-19 14:39:59 +01:00
Tat Dat Duong c8a605cbc8 Update PregelProtocol 2025-03-19 14:29:11 +01:00
Tat Dat Duong fbec207446 Apply formatting 2025-03-19 14:29:09 +01:00
Tat Dat Duong 2223c82606 Update to match JS 2025-03-19 14:28:55 +01:00
Tat Dat Duong 06f2eef74c Fix bug with stale task_id 2025-03-19 14:22:25 +01:00
Tat Dat Duong 62aa66cd4b Add better docstrings 2025-03-19 14:22:25 +01:00
Tat Dat Duong 8ffe9634b7 Clearer breakdown 2025-03-19 14:22:25 +01:00
Tat Dat Duong 4b1d6d2aeb Rename to StateUpdate 2025-03-19 14:22:25 +01:00
Tat Dat Duong 199ab46429 Avoid using checkpointer.list in async 2025-03-19 14:22:25 +01:00
Tat Dat Duong c758954519 Use awith_checkpointer instead 2025-03-19 14:22:25 +01:00
Tat Dat Duong 5bfb3bb882 Avoid running with shallow checkpointer 2025-03-19 14:22:24 +01:00
Tat Dat Duong 68a5c3f4c7 Implement batch events for RemotePregel 2025-03-19 14:22:04 +01:00
Tat Dat Duong b7fb8e6afb Fix types 2025-03-19 14:22:04 +01:00
Tat Dat Duong c34c798763 Add tests 2025-03-19 14:22:04 +01:00
Tat Dat Duong 792cd805a7 Fix tests 2025-03-19 14:21:44 +01:00
Tat Dat Duong 764929afd9 Fix lint issues 2025-03-19 14:21:43 +01:00
Tat Dat Duong 1e751a2256 Fix typo 2025-03-19 14:21:43 +01:00
Tat Dat Duong e6726802f7 feat(pregel): add bulk update state method
This method is useful for recreating a thread from a list of checkpoint writes. A new method is needed to clone a checkpoint that has been created from multiple writes (functional API, map-reduce)

Port of https://github.com/langchain-ai/langgraphjs/pull/969
2025-03-19 14:21:43 +01:00
Nuno Campos a541376d10 Store all triggers in task
- These are used to update seen version
2025-03-18 21:17:26 -07:00
William Fu-Hinthorn f9780330a6 Add refcount test 2025-03-18 20:13:40 -07:00
Nuno CamposandGitHub 24f7d7c439 Add pydantic state benchmark case (#3872) 2025-03-18 18:10:50 -07:00
Nuno Campos 2f3bd69bf5 Add pydantic benchmark 2025-03-18 16:12:34 -07:00
Tat Dat Duong b1a25abc73 Add command 2025-03-18 21:22:40 +01:00
Tat Dat Duong 4bff1df4b0 Bump to 0.0.58 2025-03-17 18:14:51 +01:00
Tat Dat Duong 5104e31e35 feat(sdk-js): add bulkUpdateState method 2025-03-17 18:14:33 +01:00
Tat Dat Duong 7ae4739630 Make options optional 2025-03-17 18:07:14 +01:00
Tat Dat Duong 1a728a93c6 feat(sdk-js): add bulkUpdateState method 2025-03-17 17:22:03 +01:00
50 changed files with 2705 additions and 782 deletions
+3 -3
View File
@@ -1,7 +1,7 @@
<picture class="github-only"> <picture class="github-only">
<source media="(prefers-color-scheme: light)" srcset="docs/docs/static/wordmark_dark.svg"> <source media="(prefers-color-scheme: light)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg">
<source media="(prefers-color-scheme: dark)" srcset="docs/docs/static/wordmark_light.svg"> <source media="(prefers-color-scheme: dark)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_light.svg">
<img alt="LangGraph Logo" src="docs/docs/static/wordmark_dark.svg" width="80%"> <img alt="LangGraph Logo" src="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg" width="80%">
</picture> </picture>
<div> <div>
+2
View File
@@ -14,6 +14,8 @@ To run the documentation server locally you can run:
make serve-docs make serve-docs
``` ```
This will start the documentation server on [http://127.0.0.1:8000/langgraph/](http://127.0.0.1:8000/langgraph/).
## Execute notebooks ## Execute notebooks
If you would like to automatically execute all of the notebooks, to mimic the "Run notebooks" GHA, you can run: If you would like to automatically execute all of the notebooks, to mimic the "Run notebooks" GHA, you can run:
Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 93 KiB

+128 -3
View File
@@ -1,6 +1,133 @@
# Prompt Engineering in LangGraph Studio # Prompt Engineering in LangGraph Studio
In LangGraph Studio you can iterate on the prompts used within your graph by utilizing the LangSmith Playground. To do so: ## Overview
A central aspect of agent development is prompt engineering. LangGraph Studio makes it easy to iterate on the prompts used within your graph directly within the UI.
## Setup
The first step is to define your [configuration](https://langchain-ai.github.io/langgraph/how-tos/configuration/) such that LangGraph Studio is aware of the prompts you want to iterate on and which nodes they are associated with.
### Reference
When defining your configuration, you can use special metadata keys to instruct LangGraph Studio how to handle different fields. Here's a reference for the available configuration options:
#### `langgraph_nodes`
- **Description**: Specifies which graph nodes a configuration field is associated with.
- **Value Type**: Array of strings, where each string is the name of a node in your graph.
- **Usage Context**: Include in the `json_schema_extra` dictionary for Pydantic models or the `metadata["json_schema_extra"]` dictionary for dataclasses.
- **Required**: No, but necessary if you want a field to be editable for specific nodes in the UI.
- **Example**:
```python
system_prompt: str = Field(
default="You are a helpful AI assistant.",
json_schema_extra={"langgraph_nodes": ["call_model", "other_node"]},
)
```
#### `langgraph_type`
- **Description**: Specifies the type of configuration field, which determines how it's handled in the UI.
- **Value Type**: String
- **Supported Values**:
- `"prompt"`: Indicates the field contains prompt text that should be treated specially in the UI.
- **Usage Context**: Include in the `json_schema_extra` dictionary for Pydantic models or the `metadata["json_schema_extra"]` dictionary for dataclasses.
- **Required**: No, but helpful for prompt fields to enable special handling.
- **Example**:
```python
system_prompt: str = Field(
default="You are a helpful AI assistant.",
json_schema_extra={
"langgraph_nodes": ["call_model"],
"langgraph_type": "prompt",
},
)
```
### Example
For example, if you have a node called `call_model` whose system prompt you want to iterate on, you can define a configuration like the following.
```python
## Using Pydantic
from pydantic import BaseModel, Field
from typing import Annotated, Literal
class Configuration(BaseModel):
"""The configuration for the agent."""
system_prompt: str = Field(
default="You are a helpful AI assistant.",
description="The system prompt to use for the agent's interactions. "
"This prompt sets the context and behavior for the agent.",
json_schema_extra={
"langgraph_nodes": ["call_model"],
"langgraph_type": "prompt",
},
)
model: Annotated[
Literal[
"anthropic/claude-3-7-sonnet-latest",
"anthropic/claude-3-5-haiku-latest",
"openai/o1",
"openai/gpt-4o-mini",
"openai/o1-mini",
"openai/o3-mini",
],
{"__template_metadata__": {"kind": "llm"}},
] = Field(
default="openai/gpt-4o-mini",
description="The name of the language model to use for the agent's main interactions. "
"Should be in the form: provider/model-name.",
json_schema_extra={"langgraph_nodes": ["call_model"]},
)
## Using Dataclasses
from dataclasses import dataclass, field
@dataclass(kw_only=True)
class Configuration:
"""The configuration for the agent."""
system_prompt: str = field(
default="You are a helpful AI assistant.",
metadata={
"description": "The system prompt to use for the agent's interactions. "
"This prompt sets the context and behavior for the agent.",
"json_schema_extra": {"langgraph_nodes": ["call_model"]},
},
)
model: Annotated[str, {"__template_metadata__": {"kind": "llm"}}] = field(
default="anthropic/claude-3-5-sonnet-20240620",
metadata={
"description": "The name of the language model to use for the agent's main interactions. "
"Should be in the form: provider/model-name.",
"json_schema_extra": {"langgraph_nodes": ["call_model"]},
},
)
```
## Iterating on prompts
### Node Configuration
With this set up, running your graph and viewing in LangGraph Studio will result in the graph rendering like such.
**Note the configuration icon in the top right corner of the `call_model` node**:
![Graph in Studio](../img/studio_graph_with_configuration.png){width=1200}
Clicking this icon will open a modal where you can edit the configuration for all of the fields associated with the `call_model` node. From here, you can save your changes and apply them to the graph. Note that these values reflect the currently active assistant, and saving will update the assistant with the new values.
![Configuration modal](../img/studio_node_configuration.png){width=1200}
### Playground
LangGraph Studio also supports prompt engineering through an integration with the LangSmith Playground. To do so:
1. Open an existing thread or create a new one. 1. Open an existing thread or create a new one.
2. Within the thread log, any nodes that have made an LLM call will have a "View LLM Runs" button. Clicking this will open a popover with the LLM runs for that node. 2. Within the thread log, any nodes that have made an LLM call will have a "View LLM Runs" button. Clicking this will open a popover with the LLM runs for that node.
@@ -8,8 +135,6 @@ In LangGraph Studio you can iterate on the prompts used within your graph by uti
![Playground in Studio](../img/studio_playground.png){width=1200} ![Playground in Studio](../img/studio_playground.png){width=1200}
From here you can edit the prompt, test different model configurations and re-run just this LLM call without having to re-run the entire graph. When you are happy with your changes, you can copy the updated prompt back into your graph. From here you can edit the prompt, test different model configurations and re-run just this LLM call without having to re-run the entire graph. When you are happy with your changes, you can copy the updated prompt back into your graph.
For more information on how to use the LangSmith Playground, see the [LangSmith Playground documentation](https://docs.smith.langchain.com/prompt_engineering/how_to_guides#playground). For more information on how to use the LangSmith Playground, see the [LangSmith Playground documentation](https://docs.smith.langchain.com/prompt_engineering/how_to_guides#playground).
+5
View File
@@ -1,3 +1,8 @@
---
search:
boost: 2
---
# LangGraph Platform # LangGraph Platform
## Overview ## Overview
@@ -1,3 +1,8 @@
---
search:
exclude: true
---
# Human-in-the-loop # Human-in-the-loop
!!! note "Use the `interrupt` function instead." !!! note "Use the `interrupt` function instead."
+2 -2
View File
@@ -463,12 +463,12 @@
"source": [ "source": [
"from langgraph.graph import StateGraph, START, END\n", "from langgraph.graph import StateGraph, START, END\n",
"from pydantic import BaseModel\n", "from pydantic import BaseModel\n",
"from langchain_core.messages import HumanMessage, AIMessage, BaseMessage\n", "from langchain_core.messages import HumanMessage, AIMessage, AnyMessage\n",
"from typing import List\n", "from typing import List\n",
"\n", "\n",
"\n", "\n",
"class ChatState(BaseModel):\n", "class ChatState(BaseModel):\n",
" messages: List[BaseMessage]\n", " messages: List[AnyMessage]\n",
" context: str\n", " context: str\n",
"\n", "\n",
"\n", "\n",
+1 -1
View File
@@ -20,7 +20,7 @@ title: Home
</p> </p>
<style> <style>
h1 { .md-content h1 {
display: none; display: none;
} }
</style> </style>
+36
View File
@@ -0,0 +1,36 @@
# LLMs-txt for LangGraph
## Overview
LangGraph provides documentation files in the [`llms.txt`](https://llmstxt.org/) format, specifically `llms.txt` and `llms-full.txt`. These files allow large language models (LLMs) and agents to access programming documentation and APIs, particularly useful within integrated development environments (IDEs).
| Language Version | llms.txt | llms-full.txt |
|------------------|------------------------------------------------------------------------------------------------------------|----------------------------------------------------------------------------------------------------------------------|
| LangGraph Python | [https://langchain-ai.github.io/langgraph/llms.txt](https://langchain-ai.github.io/langgraph/llms.txt) | [https://langchain-ai.github.io/langgraph/llms-full.txt](https://langchain-ai.github.io/langgraph/llms-full.txt) |
| LangGraph JS | [https://langchain-ai.github.io/langgraphjs/llms.txt](https://langchain-ai.github.io/langgraphjs/llms.txt) | [https://langchain-ai.github.io/langgraphjs/llms-full.txt](https://langchain-ai.github.io/langgraphjs/llms-full.txt) |
## Differences Between `llms.txt` and `llms-full.txt`
- **`llms.txt`** is an index file containing links with brief descriptions of the content. An LLM or agent must follow these links to access detailed information.
- **`llms-full.txt`** includes all the detailed content directly in a single file, eliminating the need for additional navigation.
A key consideration when using `llms-full.txt` is its size. For extensive documentation, this file may become too large to fit into an LLM's context window.
## Using `llms.txt` via an MCP Server
As of March 9, 2025, IDEs [do not yet have robust native support for `llms.txt`](https://x.com/jeremyphoward/status/1902109312216129905?t=1eHFv2vdNdAckajnug0_Vw&s=19). However, you can utilize `llms.txt` effectively through an MCP server.
We provide an MCP server specifically designed to serve documentation, called [`mcpdoc`](https://github.com/langchain-ai/mcpdoc). This setup is compatible with IDEs and platforms such as Cursor, Windsurf, Claude, and Claude Code. Instructions for using `mcpdoc` with these tools are available in the repository.
## Using `llms-full.txt`
The LangGraph `llms-full.txt` file typically contains several hundred thousand tokens, exceeding the context window limitations of most LLMs. To effectively use this file:
1. **With IDEs (e.g., Cursor, Windsurf)**:
- Add the `llms-full.txt` as custom documentation. The IDE will automatically chunk and index the content, implementing Retrieval-Augmented Generation (RAG).
2. **Without IDE support**:
- Use a chat model with a large context window.
- Implement a RAG strategy to manage and query the documentation efficiently.
+5
View File
@@ -1,3 +1,8 @@
---
search:
boost: 2
---
# Deployment # Deployment
Get started deploying your LangGraph applications locally or on the cloud with Get started deploying your LangGraph applications locally or on the cloud with
+2 -1
View File
@@ -54,7 +54,7 @@ theme:
code: "Roboto Mono" code: "Roboto Mono"
plugins: plugins:
- search: - search:
separator: '[\s\u200b\-_,:!=\[\]()"`/]+|\.(?!\d)|&[lg]t;|(?!\b)(?=[A-Z][a-z])' separator: '[\s\u200b\-,:!=\[\]()"`/]+|\.(?!\d)|&[lg]t;'
- autorefs - autorefs
- mkdocstrings: - mkdocstrings:
handlers: handlers:
@@ -361,6 +361,7 @@ nav:
# NOTE: prebuilt.md is auto-generated by `make build-prebuilt` # NOTE: prebuilt.md is auto-generated by `make build-prebuilt`
- Prebuilt Agents: prebuilt.md - Prebuilt Agents: prebuilt.md
- Companies using LangGraph: adopters.md - Companies using LangGraph: adopters.md
- LLMS-txt: llms-txt-overview.md
- FAQ: concepts/faq.md - FAQ: concepts/faq.md
- Troubleshooting: - Troubleshooting:
- Troubleshooting: troubleshooting/errors/index.md - Troubleshooting: troubleshooting/errors/index.md
+2 -2
View File
@@ -397,7 +397,7 @@ typing-extensions = ">=4.7"
[[package]] [[package]]
name = "langgraph-checkpoint" name = "langgraph-checkpoint"
version = "2.0.18" version = "2.0.21"
description = "Library with base interfaces for LangGraph checkpoint savers." description = "Library with base interfaces for LangGraph checkpoint savers."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1404,4 +1404,4 @@ cffi = ["cffi (>=1.11)"]
[metadata] [metadata]
lock-version = "2.0" lock-version = "2.0"
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
content-hash = "369bfffecb9489835b43b8255932e043176a11d2f639aad2d055ffd89263ca1e" content-hash = "4b0efdd115566f294fcd876334f9c3787aafc81f2689473759d88189a71d4635"
+2 -2
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph-checkpoint-postgres" name = "langgraph-checkpoint-postgres"
version = "2.0.18" version = "2.0.19"
description = "Library with a Postgres implementation of LangGraph checkpoint saver." description = "Library with a Postgres implementation of LangGraph checkpoint saver."
authors = [] authors = []
license = "MIT" license = "MIT"
@@ -10,7 +10,7 @@ packages = [{ include = "langgraph" }]
[tool.poetry.dependencies] [tool.poetry.dependencies]
python = "^3.9.0,<4.0" python = "^3.9.0,<4.0"
langgraph-checkpoint = "^2.0.15" langgraph-checkpoint = "^2.0.21"
orjson = ">=3.10.1" orjson = ">=3.10.1"
psycopg = "^3.2.0" psycopg = "^3.2.0"
psycopg-pool = "^3.2.0" psycopg-pool = "^3.2.0"
@@ -7,7 +7,7 @@ from collections import defaultdict
from collections.abc import AsyncIterator, Iterator, Sequence from collections.abc import AsyncIterator, Iterator, Sequence
from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack from contextlib import AbstractAsyncContextManager, AbstractContextManager, ExitStack
from types import TracebackType from types import TracebackType
from typing import Any, Optional from typing import Any, Optional, Union
from langchain_core.runnables import RunnableConfig from langchain_core.runnables import RunnableConfig
@@ -70,6 +70,12 @@ class InMemorySaver(
tuple[str, str, str], tuple[str, str, str],
dict[tuple[str, int], tuple[str, str, tuple[str, bytes], str]], dict[tuple[str, int], tuple[str, str, tuple[str, bytes], str]],
] ]
blobs: dict[
tuple[
str, str, str, Union[str, int, float]
], # thread id, checkpoint ns, channel, version
tuple[str, bytes],
]
def __init__( def __init__(
self, self,
@@ -80,6 +86,7 @@ class InMemorySaver(
super().__init__(serde=serde) super().__init__(serde=serde)
self.storage = factory(lambda: defaultdict(dict)) self.storage = factory(lambda: defaultdict(dict))
self.writes = factory(dict) self.writes = factory(dict)
self.blobs = factory()
self.stack = ExitStack() self.stack = ExitStack()
if factory is not defaultdict: if factory is not defaultdict:
self.stack.enter_context(self.storage) # type: ignore[arg-type] self.stack.enter_context(self.storage) # type: ignore[arg-type]
@@ -107,6 +114,18 @@ class InMemorySaver(
) -> Optional[bool]: ) -> Optional[bool]:
return self.stack.__exit__(__exc_type, __exc_value, __traceback) return self.stack.__exit__(__exc_type, __exc_value, __traceback)
def _load_blobs(
self, thread_id: str, checkpoint_ns: str, versions: ChannelVersions
) -> dict[str, Any]:
channel_values: dict[str, Any] = {}
for k, v in versions.items():
kk = (thread_id, checkpoint_ns, k, v)
if kk in self.blobs:
vv = self.blobs[kk]
if vv[0] != "empty":
channel_values[k] = self.serde.loads_typed(vv)
return channel_values
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
"""Get a checkpoint tuple from the in-memory storage. """Get a checkpoint tuple from the in-memory storage.
@@ -121,8 +140,8 @@ class InMemorySaver(
Returns: Returns:
Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found. Optional[CheckpointTuple]: The retrieved checkpoint tuple, or None if no matching checkpoint was found.
""" """
thread_id = config["configurable"]["thread_id"] thread_id: str = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"].get("checkpoint_ns", "") checkpoint_ns: str = config["configurable"].get("checkpoint_ns", "")
if checkpoint_id := get_checkpoint_id(config): if checkpoint_id := get_checkpoint_id(config):
if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id): if saved := self.storage[thread_id][checkpoint_ns].get(checkpoint_id):
checkpoint, metadata, parent_checkpoint_id = saved checkpoint, metadata, parent_checkpoint_id = saved
@@ -140,10 +159,14 @@ class InMemorySaver(
) )
else: else:
sends = [] sends = []
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
return CheckpointTuple( return CheckpointTuple(
config=config, config=config,
checkpoint={ checkpoint={
**self.serde.loads_typed(checkpoint), **checkpoint_,
"channel_values": self._load_blobs(
thread_id, checkpoint_ns, checkpoint_["channel_versions"]
),
"pending_sends": [self.serde.loads_typed(s[2]) for s in sends], "pending_sends": [self.serde.loads_typed(s[2]) for s in sends],
}, },
metadata=self.serde.loads_typed(metadata), metadata=self.serde.loads_typed(metadata),
@@ -180,6 +203,9 @@ class InMemorySaver(
) )
else: else:
sends = [] sends = []
checkpoint_ = self.serde.loads_typed(checkpoint)
return CheckpointTuple( return CheckpointTuple(
config={ config={
"configurable": { "configurable": {
@@ -189,7 +215,10 @@ class InMemorySaver(
} }
}, },
checkpoint={ checkpoint={
**self.serde.loads_typed(checkpoint), **checkpoint_,
"channel_values": self._load_blobs(
thread_id, checkpoint_ns, checkpoint_["channel_versions"]
),
"pending_sends": [self.serde.loads_typed(s[2]) for s in sends], "pending_sends": [self.serde.loads_typed(s[2]) for s in sends],
}, },
metadata=self.serde.loads_typed(metadata), metadata=self.serde.loads_typed(metadata),
@@ -297,6 +326,8 @@ class InMemorySaver(
else: else:
sends = [] sends = []
checkpoint_: Checkpoint = self.serde.loads_typed(checkpoint)
yield CheckpointTuple( yield CheckpointTuple(
config={ config={
"configurable": { "configurable": {
@@ -306,7 +337,12 @@ class InMemorySaver(
} }
}, },
checkpoint={ checkpoint={
**self.serde.loads_typed(checkpoint), **checkpoint_,
"channel_values": self._load_blobs(
thread_id,
checkpoint_ns,
checkpoint_["channel_versions"],
),
"pending_sends": [ "pending_sends": [
self.serde.loads_typed(s[2]) for s in sends self.serde.loads_typed(s[2]) for s in sends
], ],
@@ -353,6 +389,11 @@ class InMemorySaver(
c.pop("pending_sends") # type: ignore[misc] c.pop("pending_sends") # type: ignore[misc]
thread_id = config["configurable"]["thread_id"] thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"]["checkpoint_ns"] checkpoint_ns = config["configurable"]["checkpoint_ns"]
values: dict[str, Any] = c.pop("channel_values") # type: ignore[misc]
for k, v in new_versions.items():
self.blobs[(thread_id, checkpoint_ns, k, v)] = (
self.serde.dumps_typed(values[k]) if k in values else ("empty", b"")
)
self.storage[thread_id][checkpoint_ns].update( self.storage[thread_id][checkpoint_ns].update(
{ {
checkpoint["id"]: ( checkpoint["id"]: (
+39 -7
View File
@@ -68,7 +68,9 @@ class TestMemorySaver:
}, },
"metadata": {"run_id": "my_run_id"}, "metadata": {"run_id": "my_run_id"},
} }
self.memory_saver.put(config, self.chkpnt_2, self.metadata_2, {}) self.memory_saver.put(
config, self.chkpnt_2, self.metadata_2, self.chkpnt_2["channel_versions"]
)
checkpoint = self.memory_saver.get_tuple(config) checkpoint = self.memory_saver.get_tuple(config)
assert checkpoint is not None assert checkpoint is not None
assert checkpoint.metadata == { assert checkpoint.metadata == {
@@ -80,9 +82,24 @@ class TestMemorySaver:
async def test_search(self) -> None: async def test_search(self) -> None:
# set up test # set up test
# save checkpoints # save checkpoints
self.memory_saver.put(self.config_1, self.chkpnt_1, self.metadata_1, {}) self.memory_saver.put(
self.memory_saver.put(self.config_2, self.chkpnt_2, self.metadata_2, {}) self.config_1,
self.memory_saver.put(self.config_3, self.chkpnt_3, self.metadata_3, {}) self.chkpnt_1,
self.metadata_1,
self.chkpnt_1["channel_versions"],
)
self.memory_saver.put(
self.config_2,
self.chkpnt_2,
self.metadata_2,
self.chkpnt_2["channel_versions"],
)
self.memory_saver.put(
self.config_3,
self.chkpnt_3,
self.metadata_3,
self.chkpnt_3["channel_versions"],
)
# call method / assertions # call method / assertions
query_1 = {"source": "input"} # search by 1 key query_1 = {"source": "input"} # search by 1 key
@@ -129,9 +146,24 @@ class TestMemorySaver:
async def test_asearch(self) -> None: async def test_asearch(self) -> None:
# set up test # set up test
# save checkpoints # save checkpoints
self.memory_saver.put(self.config_1, self.chkpnt_1, self.metadata_1, {}) self.memory_saver.put(
self.memory_saver.put(self.config_2, self.chkpnt_2, self.metadata_2, {}) self.config_1,
self.memory_saver.put(self.config_3, self.chkpnt_3, self.metadata_3, {}) self.chkpnt_1,
self.metadata_1,
self.chkpnt_1["channel_versions"],
)
self.memory_saver.put(
self.config_2,
self.chkpnt_2,
self.metadata_2,
self.chkpnt_2["channel_versions"],
)
self.memory_saver.put(
self.config_3,
self.chkpnt_3,
self.metadata_3,
self.chkpnt_3["channel_versions"],
)
# call method / assertions # call method / assertions
query_1 = {"source": "input"} # search by 1 key query_1 = {"source": "input"} # search by 1 key
+1 -1
View File
@@ -79,7 +79,7 @@ The CLI uses a `langgraph.json` configuration file with these key settings:
} }
``` ```
See the [full documentation](https://langchain-ai.github.io/langgraph/docs/cloud/reference/cli.html) for detailed configuration options. See the [full documentation](https://langchain-ai.github.io/langgraph/cloud/reference/cli/) for detailed configuration options.
## Development ## Development
+8
View File
@@ -574,6 +574,12 @@ def dockerfile(save_path: str, config: pathlib.Path, add_docker_compose: bool) -
help="Wait for a debugger client to connect to the debug port before starting the server", help="Wait for a debugger client to connect to the debug port before starting the server",
default=False, default=False,
) )
@click.option(
"--studio_url",
type=str,
default=None,
help="URL of the LangGraph Studio instance to connect to. Defaults to https://smith.langchain.com",
)
@cli.command( @cli.command(
"dev", "dev",
help="🏃‍♀️‍➡️ Run LangGraph API server in development mode with hot reloading and debugging support", help="🏃‍♀️‍➡️ Run LangGraph API server in development mode with hot reloading and debugging support",
@@ -588,6 +594,7 @@ def dev(
no_browser: bool, no_browser: bool,
debug_port: Optional[int], debug_port: Optional[int],
wait_for_client: bool, wait_for_client: bool,
studio_url: Optional[str],
): ):
"""CLI entrypoint for running the LangGraph API server.""" """CLI entrypoint for running the LangGraph API server."""
try: try:
@@ -651,6 +658,7 @@ def dev(
wait_for_client=wait_for_client, wait_for_client=wait_for_client,
auth=config_json.get("auth"), auth=config_json.get("auth"),
http=config_json.get("http"), http=config_json.get("http"),
studio_url=studio_url,
) )
+19 -19
View File
@@ -535,42 +535,42 @@ langgraph-sdk = ">=0.1.42,<0.2.0"
[[package]] [[package]]
name = "langgraph-api" name = "langgraph-api"
version = "0.0.27" version = "0.0.32"
description = "" description = ""
optional = true optional = true
python-versions = "<4.0,>=3.11.0" python-versions = "<4.0,>=3.11.0"
files = [ files = [
{file = "langgraph_api-0.0.27-py3-none-any.whl", hash = "sha256:9b21742238b15b8db9c2d3fd760a670332c8897d0bcbbd9d82e43b6ac15a7937"}, {file = "langgraph_api-0.0.32-py3-none-any.whl", hash = "sha256:7990cedc65f784813aba867c5bde3fdfae3fa4588baef1aa346cbeac7c3aebf1"},
{file = "langgraph_api-0.0.27.tar.gz", hash = "sha256:c21eb2b7fe3b93998379f7b13ad7d23b3ef06ab821b008c6b12b954acfb587ec"}, {file = "langgraph_api-0.0.32.tar.gz", hash = "sha256:6f5b698ad8d136b73c2c53bcfa30670e9244a318b08b5e9cf00a707ea57c058c"},
] ]
[package.dependencies] [package.dependencies]
cryptography = ">=43.0.3,<44.0.0" cryptography = ">=43.0.3,<44.0.0"
httpx = ">=0.25.0" httpx = ">=0.25.0"
jsonschema-rs = ">=0.20.0,<0.21.0" jsonschema-rs = ">=0.20.0,<0.30"
langchain-core = ">=0.2.38,<0.4.0" langchain-core = ">=0.2.38,<0.4.0"
langgraph = ">=0.2.56,<0.4.0" langgraph = ">=0.2.56,<0.4.0"
langgraph-checkpoint = ">=2.0.15,<3.0" langgraph-checkpoint = ">=2.0.21,<3.0"
langgraph-sdk = ">=0.1.53,<0.2.0" langgraph-sdk = ">=0.1.58,<0.2.0"
langsmith = ">=0.1.63,<0.4.0" langsmith = ">=0.1.63,<0.4.0"
orjson = ">=3.9.7" orjson = ">=3.9.7"
pyjwt = ">=2.9.0,<3.0.0" pyjwt = ">=2.9.0,<3.0.0"
sse-starlette = ">=2.1.0,<2.2.0" sse-starlette = ">=2.1.0,<2.2.0"
starlette = ">=0.38.6" starlette = ">=0.38.6"
structlog = ">=23.1.0,<24.0.0" structlog = ">=24.1.0,<26"
tenacity = ">=8.0.0" tenacity = ">=8.0.0"
uvicorn = ">=0.26.0" uvicorn = ">=0.26.0"
watchfiles = ">=0.13" watchfiles = ">=0.13"
[[package]] [[package]]
name = "langgraph-checkpoint" name = "langgraph-checkpoint"
version = "2.0.16" version = "2.0.21"
description = "Library with base interfaces for LangGraph checkpoint savers." description = "Library with base interfaces for LangGraph checkpoint savers."
optional = true optional = true
python-versions = "<4.0.0,>=3.9.0" python-versions = "<4.0.0,>=3.9.0"
files = [ files = [
{file = "langgraph_checkpoint-2.0.16-py3-none-any.whl", hash = "sha256:dfab51076a6eddb5f9e146cfe1b977e3dd6419168b2afa23ff3f4e47973bf06f"}, {file = "langgraph_checkpoint-2.0.21-py3-none-any.whl", hash = "sha256:ca89c2090cd9729f83f9782226935dc5ff9fe7756c24936f484ccb0ce367f87b"},
{file = "langgraph_checkpoint-2.0.16.tar.gz", hash = "sha256:49ba8cfa12b2aae845ccc3b1fbd1d7a8d3a6c4a2e387ab3a92fca40dd3d4baa5"}, {file = "langgraph_checkpoint-2.0.21.tar.gz", hash = "sha256:52beeb6dc1bd8c487b8315466cab271093b65eb97f54a0942dfe105cd20b237f"},
] ]
[package.dependencies] [package.dependencies]
@@ -594,13 +594,13 @@ langgraph-checkpoint = ">=2.0.10,<3.0.0"
[[package]] [[package]]
name = "langgraph-sdk" name = "langgraph-sdk"
version = "0.1.53" version = "0.1.58"
description = "SDK for interacting with LangGraph API" description = "SDK for interacting with LangGraph API"
optional = true optional = true
python-versions = "<4.0.0,>=3.9.0" python-versions = "<4.0.0,>=3.9.0"
files = [ files = [
{file = "langgraph_sdk-0.1.53-py3-none-any.whl", hash = "sha256:4fab62caad73661ffe4c3ababedcd0d7bfaaba986bee4416b9c28948458a3af5"}, {file = "langgraph_sdk-0.1.58-py3-none-any.whl", hash = "sha256:65f88cf5582da0c316714dc475126fa03c5f74d72bc0b9221dd42649de8e23d4"},
{file = "langgraph_sdk-0.1.53.tar.gz", hash = "sha256:12906ed965905fa27e0c28d9fa07dc6fd89e6895ff321ff049fdf3965d057cc4"}, {file = "langgraph_sdk-0.1.58.tar.gz", hash = "sha256:ef8b0e4c08af8c7efd3919497879c87a3627806b51e4ba5e8b06e0717e3d44cd"},
] ]
[package.dependencies] [package.dependencies]
@@ -1357,18 +1357,18 @@ full = ["httpx (>=0.27.0,<0.29.0)", "itsdangerous", "jinja2", "python-multipart
[[package]] [[package]]
name = "structlog" name = "structlog"
version = "23.3.0" version = "25.2.0"
description = "Structured Logging for Python" description = "Structured Logging for Python"
optional = true optional = true
python-versions = ">=3.8" python-versions = ">=3.8"
files = [ files = [
{file = "structlog-23.3.0-py3-none-any.whl", hash = "sha256:d6922a88ceabef5b13b9eda9c4043624924f60edbb00397f4d193bd754cde60a"}, {file = "structlog-25.2.0-py3-none-any.whl", hash = "sha256:0fecea2e345d5d491b72f3db2e5fcd6393abfc8cd06a4851f21fcd4d1a99f437"},
{file = "structlog-23.3.0.tar.gz", hash = "sha256:24b42b914ac6bc4a4e6f716e82ac70d7fb1e8c3b1035a765591953bfc37101a5"}, {file = "structlog-25.2.0.tar.gz", hash = "sha256:d9f9776944207d1035b8b26072b9b140c63702fd7aa57c2f85d28ab701bd8e92"},
] ]
[package.extras] [package.extras]
dev = ["structlog[tests,typing]"] dev = ["freezegun (>=0.2.8)", "mypy (>=1.4)", "pretend", "pytest (>=6.0)", "pytest-asyncio (>=0.17)", "rich", "simplejson", "twisted"]
docs = ["furo", "myst-parser", "sphinx", "sphinx-notfound-page", "sphinxcontrib-mermaid", "sphinxext-opengraph", "twisted"] docs = ["cogapp", "furo", "myst-parser", "sphinx", "sphinx-notfound-page", "sphinxcontrib-mermaid", "sphinxext-opengraph", "twisted"]
tests = ["freezegun (>=0.2.8)", "pretend", "pytest (>=6.0)", "pytest-asyncio (>=0.17)", "simplejson"] tests = ["freezegun (>=0.2.8)", "pretend", "pytest (>=6.0)", "pytest-asyncio (>=0.17)", "simplejson"]
typing = ["mypy (>=1.4)", "rich", "twisted"] typing = ["mypy (>=1.4)", "rich", "twisted"]
@@ -1717,4 +1717,4 @@ inmem = ["langgraph-api", "python-dotenv"]
[metadata] [metadata]
lock-version = "2.0" lock-version = "2.0"
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
content-hash = "d0e2bdcb600ad031867413025fcc58bb162609209359d63ca99a77060cf8cbb4" content-hash = "f5aa4d66f9c0b98b8321a70a82387dc6e5f3a3a7ecedd87ac00d6415199038f9"
+2 -2
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph-cli" name = "langgraph-cli"
version = "0.1.77" version = "0.1.78"
description = "CLI for interacting with LangGraph API" description = "CLI for interacting with LangGraph API"
authors = [] authors = []
license = "MIT" license = "MIT"
@@ -14,7 +14,7 @@ langgraph = "langgraph_cli.cli:cli"
[tool.poetry.dependencies] [tool.poetry.dependencies]
python = "^3.9.0,<4.0" python = "^3.9.0,<4.0"
click = "^8.1.7" click = "^8.1.7"
langgraph-api = { version = ">=0.0.27,<0.1.0", optional = true, python = ">=3.11,<4.0" } langgraph-api = { version = ">=0.0.32,<0.1.0", optional = true, python = ">=3.11,<4.0" }
python-dotenv = { version = ">=0.8.0", optional = true } python-dotenv = { version = ">=0.8.0", optional = true }
[tool.poetry.group.dev.dependencies] [tool.poetry.group.dev.dependencies]
+3 -3
View File
@@ -1,7 +1,7 @@
<picture class="github-only"> <picture class="github-only">
<source media="(prefers-color-scheme: light)" srcset="docs/docs/static/wordmark_dark.svg"> <source media="(prefers-color-scheme: light)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg">
<source media="(prefers-color-scheme: dark)" srcset="docs/docs/static/wordmark_light.svg"> <source media="(prefers-color-scheme: dark)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_light.svg">
<img alt="LangGraph Logo" src="docs/docs/static/wordmark_dark.svg" width="80%"> <img alt="LangGraph Logo" src="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg" width="80%">
</picture> </picture>
<div> <div>
+139
View File
@@ -6,10 +6,12 @@ from pyperf._runner import Runner
from uvloop import new_event_loop from uvloop import new_event_loop
from bench.fanout_to_subgraph import fanout_to_subgraph, fanout_to_subgraph_sync from bench.fanout_to_subgraph import fanout_to_subgraph, fanout_to_subgraph_sync
from bench.pydantic_state import pydantic_state
from bench.react_agent import react_agent from bench.react_agent import react_agent
from bench.sequential import create_sequential from bench.sequential import create_sequential
from bench.wide_state import wide_state from bench.wide_state import wide_state
from langgraph.checkpoint.memory import MemorySaver from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import StateGraph
from langgraph.pregel import Pregel from langgraph.pregel import Pregel
@@ -43,6 +45,11 @@ def run(graph: Pregel, input: dict):
) )
def compile_graph(graph: StateGraph) -> None:
"""Compile the graph."""
graph.compile()
benchmarks = ( benchmarks = (
( (
"fanout_to_subgraph_10x", "fanout_to_subgraph_10x",
@@ -228,12 +235,144 @@ benchmarks = (
create_sequential(200).compile(), create_sequential(200).compile(),
{"messages": []}, # Empty list of messages {"messages": []}, # Empty list of messages
), ),
(
"pydantic_state_25x300",
pydantic_state(300).compile(checkpointer=None),
pydantic_state(300).compile(checkpointer=None),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(5)
}
for i in range(5)
}
]
},
),
(
"pydantic_state_25x300_checkpoint",
pydantic_state(300).compile(checkpointer=MemorySaver()),
pydantic_state(300).compile(checkpointer=MemorySaver()),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(5)
}
for i in range(5)
}
]
},
),
(
"pydantic_state_15x600",
pydantic_state(600).compile(checkpointer=None),
pydantic_state(600).compile(checkpointer=None),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(5)
}
for i in range(3)
}
]
},
),
(
"pydantic_state_15x600_checkpoint",
pydantic_state(600).compile(checkpointer=MemorySaver()),
pydantic_state(600).compile(checkpointer=MemorySaver()),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(5)
}
for i in range(3)
}
]
},
),
(
"pydantic_state_9x1200",
pydantic_state(1200).compile(checkpointer=None),
pydantic_state(1200).compile(checkpointer=None),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(3)
}
for i in range(3)
}
]
},
),
(
"pydantic_state_9x1200_checkpoint",
pydantic_state(1200).compile(checkpointer=MemorySaver()),
pydantic_state(1200).compile(checkpointer=MemorySaver()),
{
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(3)
}
for i in range(3)
}
]
},
),
) )
r = Runner() r = Runner()
# Full graph run time
for name, agraph, graph, input in benchmarks: for name, agraph, graph, input in benchmarks:
r.bench_async_func(name, arun, agraph, input, loop_factory=new_event_loop) r.bench_async_func(name, arun, agraph, input, loop_factory=new_event_loop)
if graph is not None: if graph is not None:
r.bench_func(name + "_sync", run, graph, input) r.bench_func(name + "_sync", run, graph, input)
# Graph compilation times
compilation_benchmarks = (
(
"sequential_1000",
create_sequential(1_000),
),
(
"sequential_10000",
create_sequential(10_000),
),
(
"pydantic_state_25x300",
pydantic_state(300),
),
(
"pydantic_state_15x600",
pydantic_state(600),
),
(
"pydantic_state_9x1200",
pydantic_state(1200),
),
(
"wide_state_15x600",
wide_state(600),
),
(
"wide_state_9x1200",
wide_state(1200),
),
)
for name, graph in compilation_benchmarks:
r.bench_func(name + "_compilation", compile_graph, graph)
+327
View File
@@ -0,0 +1,327 @@
import operator
from functools import partial
from random import choice
from typing import Annotated, Optional, Sequence
from pydantic import BaseModel, Field, field_validator
from langgraph.constants import END, START
from langgraph.graph.state import StateGraph
def pydantic_state(n: int) -> StateGraph:
class State(BaseModel):
messages: Annotated[list, operator.add] = Field(default_factory=list)
@field_validator("messages", mode="after")
@classmethod
def validate_messages(cls, v):
if not isinstance(v, list):
raise TypeError("messages must be a list")
for msg in v:
if not isinstance(msg, dict):
raise TypeError("messages must be a list of dicts")
if not all(isinstance(k, str) for k in msg.keys()):
raise TypeError("messages must be a list of dicts with str keys")
return v
trigger_events: Annotated[list, operator.add] = Field(default_factory=list)
"""The external events that are converted by the graph."""
@field_validator("trigger_events", mode="after")
@classmethod
def validate_trigger_events(cls, v):
if not isinstance(v, list):
raise TypeError("trigger_events must be a list")
for event in v:
if not isinstance(event, dict):
raise TypeError("trigger_events must be a list of dicts")
if not all(isinstance(k, str) for k in event.keys()):
raise TypeError(
"trigger_events must be a list of dicts with str keys"
)
return v
primary_issue_medium: Annotated[str, lambda x, y: y or x] = Field(
default="email"
)
"""The primary issue medium for the current conversation."""
@field_validator("primary_issue_medium", mode="after")
@classmethod
def validate_primary_issue_medium(cls, v):
if not isinstance(v, str):
raise TypeError("primary_issue_medium must be a string")
return v
autoresponse: Annotated[Optional[dict], lambda _, y: y] = Field(
default=None
) # Always overwrite
@field_validator("autoresponse", mode="after")
@classmethod
def validate_autoresponse(cls, v):
if v is not None and not isinstance(v, dict):
raise TypeError("autoresponse must be a dict or None")
return v
issue: Annotated[dict | None, lambda x, y: y if y else x] = Field(default=None)
@field_validator("issue", mode="after")
@classmethod
def validate_issue(cls, v):
if v is not None and not isinstance(v, dict):
raise TypeError("issue must be a dict or None")
return v
relevant_rules: Optional[list[dict]] = Field(default=None)
"""SOPs fetched from the rulebook that are relevant to the current conversation."""
@field_validator("relevant_rules", mode="after")
@classmethod
def validate_relevant_rules(cls, v):
if v is None:
return v
if not isinstance(v, list):
raise TypeError("relevant_rules must be a list or None")
for rule in v:
if not isinstance(rule, dict):
raise TypeError("relevant_rules must be a list of dicts")
if not all(isinstance(k, str) for k in rule.keys()):
raise TypeError(
"relevant_rules must be a list of dicts with str keys"
)
return v
memory_docs: Optional[list[dict]] = Field(default=None)
"""Memory docs fetched from the memory service that are relevant to the current conversation."""
@field_validator("memory_docs", mode="after")
@classmethod
def validate_memory_docs(cls, v):
if v is None:
return v
if not isinstance(v, list):
raise TypeError("memory_docs must be a list or None")
for doc in v:
if not isinstance(doc, dict):
raise TypeError("memory_docs must be a list of dicts")
if not all(isinstance(k, str) for k in doc.keys()):
raise TypeError("memory_docs must be a list of dicts with str keys")
return v
categorizations: Annotated[list[dict], operator.add] = Field(
default_factory=list
)
"""The issue categorizations auto-generated by the AI."""
@field_validator("categorizations", mode="after")
@classmethod
def validate_categorizations(cls, v):
if not isinstance(v, list):
raise TypeError("categorizations must be a list")
for categorization in v:
if not isinstance(categorization, dict):
raise TypeError("categorizations must be a list of dicts")
if not all(isinstance(k, str) for k in categorization.keys()):
raise TypeError(
"categorizations must be a list of dicts with str keys"
)
return v
responses: Annotated[list[dict], operator.add] = Field(default_factory=list)
"""The draft responses recommended by the AI."""
@field_validator("responses", mode="after")
@classmethod
def validate_responses(cls, v):
if not isinstance(v, list):
raise TypeError("responses must be a list")
for response in v:
if not isinstance(response, dict):
raise TypeError("responses must be a list of dicts")
if not all(isinstance(k, str) for k in response.keys()):
raise TypeError("responses must be a list of dicts with str keys")
return v
user_info: Annotated[Optional[dict], lambda x, y: y if y is not None else x] = (
Field(default=None)
)
"""The current user state (by email)."""
@field_validator("user_info", mode="after")
@classmethod
def validate_user_info(cls, v):
if v is not None and not isinstance(v, dict):
raise TypeError("user_info must be a dict or None")
return v
crm_info: Annotated[Optional[dict], lambda x, y: y if y is not None else x] = (
Field(default=None)
)
"""The CRM information for organization the current user is from."""
@field_validator("crm_info", mode="after")
@classmethod
def validate_crm_info(cls, v):
if v is not None and not isinstance(v, dict):
raise TypeError("crm_info must be a dict or None")
return v
email_thread_id: Annotated[
Optional[str], lambda x, y: y if y is not None else x
] = Field(default=None)
"""The current email thread ID."""
@field_validator("email_thread_id", mode="after")
@classmethod
def validate_email_thread_id(cls, v):
if v is not None and not isinstance(v, str):
raise TypeError("email_thread_id must be a string or None")
return v
slack_participants: Annotated[dict, operator.or_] = Field(default_factory=dict)
"""The growing list of current slack participants."""
@field_validator("slack_participants", mode="after")
@classmethod
def validate_slack_participants(cls, v):
if not isinstance(v, dict):
raise TypeError("slack_participants must be a dict")
for participant in v:
if not isinstance(participant, str):
raise TypeError("slack_participants must be a dict with str keys")
return v
bot_id: Optional[str] = Field(default=None)
"""The ID of the bot user in the slack channel."""
@field_validator("bot_id", mode="after")
@classmethod
def validate_bot_id(cls, v):
if v is not None and not isinstance(v, str):
raise TypeError("bot_id must be a string or None")
return v
notified_assignees: Annotated[dict, operator.or_] = Field(default_factory=dict)
@field_validator("notified_assignees", mode="after")
def validate_notified_assignees(cls, v):
if not isinstance(v, dict):
raise TypeError("notified_assignees must be a dict")
for assignee in v:
if not isinstance(assignee, str):
raise TypeError("notified_assignees must be a dict with str keys")
return v
list_fields = {
"messages",
"trigger_events",
"categorizations",
"responses",
"memory_docs",
"relevant_rules",
}
dict_fields = {
"user_info",
"crm_info",
"slack_participants",
"notified_assignees",
"autoresponse",
"issue",
}
def read_write(read: str, write: Sequence[str], input: State) -> dict:
val = getattr(input, read)
val = {val: val} if isinstance(val, str) else val
val_single = val[-1] if isinstance(val, list) else val
val_list = val if isinstance(val, list) else [val]
return {
k: val_list
if k in list_fields
else val_single
if k in dict_fields
else "".join(choice("abcdefghijklmnopqrstuvwxyz") for _ in range(n))
for k in write
}
builder = StateGraph(State)
builder.add_edge(START, "one")
builder.add_node(
"one",
partial(read_write, "messages", ["trigger_events", "primary_issue_medium"]),
)
builder.add_edge("one", "two")
builder.add_node(
"two",
partial(read_write, "trigger_events", ["autoresponse", "issue"]),
)
builder.add_edge("two", "three")
builder.add_edge("two", "four")
builder.add_node(
"three",
partial(read_write, "autoresponse", ["relevant_rules"]),
)
builder.add_node(
"four",
partial(
read_write,
"trigger_events",
["categorizations", "responses", "memory_docs"],
),
)
builder.add_node(
"five",
partial(
read_write,
"categorizations",
[
"user_info",
"crm_info",
"email_thread_id",
"slack_participants",
"bot_id",
"notified_assignees",
],
),
)
builder.add_edge(["three", "four"], "five")
builder.add_edge("five", "six")
builder.add_node(
"six",
partial(read_write, "responses", ["messages"]),
)
builder.add_conditional_edges(
"six", lambda state: END if len(state.messages) > n else "one"
)
return builder
if __name__ == "__main__":
import asyncio
import uvloop
from langgraph.checkpoint.memory import MemorySaver
graph = pydantic_state(1000).compile(checkpointer=MemorySaver())
input = {
"messages": [
{
str(i) * 10: {
str(j) * 10: ["hi?" * 10, True, 1, 6327816386138, None] * 5
for j in range(5)
}
for i in range(5)
}
]
}
config = {"configurable": {"thread_id": "1"}, "recursion_limit": 20000000000}
async def run():
async for c in graph.astream(input, config=config):
print(c.keys())
uvloop.install()
asyncio.run(run())
+1
View File
@@ -138,6 +138,7 @@ class Branch(NamedTuple):
reader=reader, reader=reader,
name=None, name=None,
trace=False, trace=False,
func_accepts_config=True,
) )
) )
+29 -17
View File
@@ -859,8 +859,10 @@ class CompiledStateGraph(CompiledGraph):
# subscribe to channel # subscribe to channel
self.nodes[end].triggers.append(channel_name) self.nodes[end].triggers.append(channel_name)
# publish to channel # publish to channel
self.nodes[START] |= ChannelWrite( self.nodes[START].writers.append(
[ChannelWriteEntry(channel_name, START)], tags=[TAG_HIDDEN] ChannelWrite(
[ChannelWriteEntry(channel_name, START)], tags=[TAG_HIDDEN]
)
) )
elif end != END: elif end != END:
# subscribe to start channel # subscribe to start channel
@@ -873,8 +875,10 @@ class CompiledStateGraph(CompiledGraph):
self.nodes[end].triggers.append(channel_name) self.nodes[end].triggers.append(channel_name)
# publish to channel # publish to channel
for start in starts: for start in starts:
self.nodes[start] |= ChannelWrite( self.nodes[start].writers.append(
[ChannelWriteEntry(channel_name, start)], tags=[TAG_HIDDEN] ChannelWrite(
[ChannelWriteEntry(channel_name, start)], tags=[TAG_HIDDEN]
)
) )
def attach_branch( def attach_branch(
@@ -910,28 +914,31 @@ class CompiledStateGraph(CompiledGraph):
if start in self.builder.nodes if start in self.builder.nodes
else self.builder.schema else self.builder.schema
) )
# attach branch publisher
self.nodes[start] |= branch.run(
branch_writer,
_get_state_reader(self.builder, schema) if with_reader else None,
)
# attach branch subscribers # attach branch publisher
ends = ( self.nodes[start].writers.append(
branch.ends.values() branch.run(
if branch.ends branch_writer,
else [node for node in self.builder.nodes if node != branch.then] _get_state_reader(self.builder, schema) if with_reader else None,
)
) )
# attach then subscriber # attach then subscriber
if branch.then and branch.then != END: 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" channel_name = f"branch:{start}:{name}::then"
self.channels[channel_name] = DynamicBarrierValue(str) self.channels[channel_name] = DynamicBarrierValue(str)
self.nodes[branch.then].triggers.append(channel_name) self.nodes[branch.then].triggers.append(channel_name)
for end in ends: for end in ends:
if end != END: if end != END:
self.nodes[end] |= ChannelWrite( self.nodes[end].writers.append(
[ChannelWriteEntry(channel_name, end)], tags=[TAG_HIDDEN] ChannelWrite(
[ChannelWriteEntry(channel_name, end)], tags=[TAG_HIDDEN]
)
) )
@@ -1013,7 +1020,12 @@ async def _acontrol_branch(value: Any) -> Sequence[Union[str, Send]]:
CONTROL_BRANCH_PATH = RunnableCallable( CONTROL_BRANCH_PATH = RunnableCallable(
_control_branch, _acontrol_branch, tags=[TAG_HIDDEN], trace=False, recurse=False _control_branch,
_acontrol_branch,
tags=[TAG_HIDDEN],
trace=False,
recurse=False,
func_accepts_config=False,
) )
CONTROL_BRANCH = Branch(CONTROL_BRANCH_PATH, None) CONTROL_BRANCH = Branch(CONTROL_BRANCH_PATH, None)
File diff suppressed because it is too large Load Diff
+80 -19
View File
@@ -233,10 +233,21 @@ def apply_writes(
channels: Mapping[str, BaseChannel], channels: Mapping[str, BaseChannel],
tasks: Iterable[WritesProtocol], tasks: Iterable[WritesProtocol],
get_next_version: Optional[GetNextVersion], get_next_version: Optional[GetNextVersion],
) -> dict[str, list[Any]]: ) -> tuple[dict[str, list[Any]], set[str]]:
"""Apply writes from a set of tasks (usually the tasks from a Pregel step) """Apply writes from a set of tasks (usually the tasks from a Pregel step)
to the checkpoint and channels, and return managed values writes to be applied to the checkpoint and channels, and return managed values writes to be applied
externally.""" externally.
Args:
checkpoint: The checkpoint to update.
channels: The channels to update.
tasks: The tasks to apply writes from.
get_next_version: Optional function to determine the next version of a channel.
Returns:
A tuple containing the managed values writes to be applied externally, and
the set of channels that were updated in this step.
"""
# sort tasks on path, to ensure deterministic order for update application # sort tasks on path, to ensure deterministic order for update application
# any path parts after the 3rd are ignored for sorting # any path parts after the 3rd are ignored for sorting
# (we use them for eg. task ids which aren't good for sorting) # (we use them for eg. task ids which aren't good for sorting)
@@ -312,15 +323,14 @@ def apply_writes(
# Channels that weren't updated in this step are notified of a new step # Channels that weren't updated in this step are notified of a new step
if bump_step: if bump_step:
for chan in channels: for chan in channels:
if chan not in updated_channels: if channels[chan].is_available() and chan not in updated_channels:
if channels[chan].update([]) and get_next_version is not None: if channels[chan].update(EMPTY_SEQ) and get_next_version is not None:
checkpoint["channel_versions"][chan] = get_next_version( checkpoint["channel_versions"][chan] = get_next_version(
max_version, max_version,
channels[chan], channels[chan],
) )
# Return managed values writes to be applied externally # Return managed values writes to be applied externally
return pending_writes_by_managed return pending_writes_by_managed, updated_channels
@overload @overload
@@ -337,6 +347,8 @@ def prepare_next_tasks(
store: Literal[None] = None, store: Literal[None] = None,
checkpointer: Literal[None] = None, checkpointer: Literal[None] = None,
manager: Literal[None] = None, manager: Literal[None] = None,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
updated_channels: Optional[set[str]] = None,
) -> dict[str, PregelTask]: ... ) -> dict[str, PregelTask]: ...
@@ -354,6 +366,8 @@ def prepare_next_tasks(
store: Optional[BaseStore], store: Optional[BaseStore],
checkpointer: Optional[BaseCheckpointSaver], checkpointer: Optional[BaseCheckpointSaver],
manager: Union[None, ParentRunManager, AsyncParentRunManager], manager: Union[None, ParentRunManager, AsyncParentRunManager],
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
updated_channels: Optional[set[str]] = None,
) -> dict[str, PregelExecutableTask]: ... ) -> dict[str, PregelExecutableTask]: ...
@@ -370,10 +384,35 @@ def prepare_next_tasks(
store: Optional[BaseStore] = None, store: Optional[BaseStore] = None,
checkpointer: Optional[BaseCheckpointSaver] = None, checkpointer: Optional[BaseCheckpointSaver] = None,
manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
updated_channels: Optional[set[str]] = None,
) -> Union[dict[str, PregelTask], dict[str, PregelExecutableTask]]: ) -> Union[dict[str, PregelTask], dict[str, PregelExecutableTask]]:
"""Prepare the set of tasks that will make up the next Pregel step. """Prepare the set of tasks that will make up the next Pregel step.
This is the union of all PUSH tasks (Sends) and PULL tasks (nodes triggered
by edges).""" Args:
checkpoint: The current checkpoint.
pending_writes: The list of pending writes.
processes: The mapping of process names to PregelNode instances.
channels: The mapping of channel names to BaseChannel instances.
managed: The mapping of managed value names to functions.
config: The runnable configuration.
step: The current step.
for_execution: Whether the tasks are being prepared for execution.
store: An instance of BaseStore to make it available for usage within tasks.
checkpointer: Checkpointer instance used for saving checkpoints.
manager: The parent run manager to use for the tasks.
trigger_to_nodes: Optional: Mapping of channel names to the set of nodes
that are can be triggered by that channel.
updated_channels: Optional. Set of channel names that have been updated during
the previous step. Using in conjunction with trigger_to_nodes to speed
up the process of determining which nodes should be triggered in the next
step.
Returns:
A dictionary of tasks to be executed. The keys are the task ids and the values
are the tasks themselves. This is the union of all PUSH tasks (Sends)
and PULL tasks (nodes triggered by edges).
"""
checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", "")) checkpoint_id_bytes = binascii.unhexlify(checkpoint["id"].replace("-", ""))
null_version = checkpoint_null_version(checkpoint) null_version = checkpoint_null_version(checkpoint)
tasks: list[Union[PregelTask, PregelExecutableTask]] = [] tasks: list[Union[PregelTask, PregelExecutableTask]] = []
@@ -397,9 +436,30 @@ def prepare_next_tasks(
manager=manager, manager=manager,
): ):
tasks.append(task) tasks.append(task)
# This section is an optimization that allows which nodes will be active
# during the next step.
# When there's information about:
# 1. Which channels were updated in the previous step
# 2. Which nodes are triggered by which channels
# Then we can determine which nodes should be triggered in the next step
# without having to cycle through all nodes.
if updated_channels and trigger_to_nodes:
triggered_nodes: set[str] = set()
# Get all nodes that have triggers associated with an updated channel
for channel in updated_channels:
if node_ids := trigger_to_nodes.get(channel):
triggered_nodes.update(node_ids)
# Sort the nodes to ensure deterministic order
candidate_nodes: Iterable[str] = sorted(triggered_nodes)
elif not checkpoint["channel_versions"]:
candidate_nodes = ()
else:
candidate_nodes = processes.keys()
# Check if any processes should be run in next step # Check if any processes should be run in next step
# If so, prepare the values to be passed to them # If so, prepare the values to be passed to them
for name in processes: for name in candidate_nodes:
if task := prepare_single_task( if task := prepare_single_task(
(PULL, name), (PULL, name),
None, None,
@@ -517,7 +577,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: _scratchpad( CONFIG_KEY_SCRATCHPAD: _scratchpad(
config, config[CONF].get(CONFIG_KEY_SCRATCHPAD),
pending_writes, pending_writes,
task_id, task_id,
), ),
@@ -627,7 +687,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: _scratchpad( CONFIG_KEY_SCRATCHPAD: _scratchpad(
config, config[CONF].get(CONFIG_KEY_SCRATCHPAD),
pending_writes, pending_writes,
task_id, task_id,
), ),
@@ -655,13 +715,14 @@ def prepare_single_task(
if checkpoint_null_version is None: if checkpoint_null_version is None:
return return
# If any of the channels read by this process were updated # If any of the channels read by this process were updated
if triggers := _triggers( if _triggers(
channels, channels,
checkpoint["channel_versions"], checkpoint["channel_versions"],
checkpoint["versions_seen"].get(name), checkpoint["versions_seen"].get(name),
checkpoint_null_version, checkpoint_null_version,
proc, proc,
): ):
triggers = tuple(sorted(proc.triggers))
try: try:
val = next( val = next(
_proc_input(proc, managed, channels, for_execution=for_execution) _proc_input(proc, managed, channels, for_execution=for_execution)
@@ -721,7 +782,7 @@ def prepare_single_task(
CONFIG_KEY_SEND: partial( CONFIG_KEY_SEND: partial(
local_write, local_write,
writes.extend, writes.extend,
processes.keys(), tuple(processes.keys()),
), ),
CONFIG_KEY_READ: partial( CONFIG_KEY_READ: partial(
local_read, local_read,
@@ -730,7 +791,10 @@ def prepare_single_task(
channels, channels,
managed, managed,
PregelTaskWrites( PregelTaskWrites(
task_path[:3], name, writes, triggers task_path[:3],
name,
writes,
triggers,
), ),
config, config,
), ),
@@ -748,7 +812,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None, CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns, CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: _scratchpad( CONFIG_KEY_SCRATCHPAD: _scratchpad(
config, config[CONF].get(CONFIG_KEY_SCRATCHPAD),
pending_writes, pending_writes,
task_id, task_id,
), ),
@@ -799,7 +863,7 @@ def _triggers(
def _scratchpad( def _scratchpad(
config: RunnableConfig, parent_scratchpad: Optional[PregelScratchpad],
pending_writes: list[PendingWrite], pending_writes: list[PendingWrite],
task_id: str, task_id: str,
) -> PregelScratchpad: ) -> PregelScratchpad:
@@ -808,9 +872,6 @@ def _scratchpad(
null_resume_write = next( null_resume_write = next(
(w for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME), None (w for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME), None
) )
parent_scratchpad: Optional[PregelScratchpad] = config[CONF].get(
CONFIG_KEY_SCRATCHPAD
)
def get_null_resume(consume: bool = False) -> Any: def get_null_resume(consume: bool = False) -> Any:
if null_resume_write is None: if null_resume_write is None:
+25 -11
View File
@@ -1,6 +1,7 @@
import asyncio import asyncio
import binascii import binascii
import concurrent.futures import concurrent.futures
import dataclasses
from collections import defaultdict, deque from collections import defaultdict, deque
from contextlib import AsyncExitStack, ExitStack from contextlib import AsyncExitStack, ExitStack
from inspect import signature from inspect import signature
@@ -209,6 +210,7 @@ class PregelLoop(LoopProtocol):
manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None,
input_model: Optional[Type[BaseModel]] = None, input_model: Optional[Type[BaseModel]] = None,
debug: bool = False, debug: bool = False,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
) -> None: ) -> None:
super().__init__( super().__init__(
step=0, step=0,
@@ -232,6 +234,7 @@ class PregelLoop(LoopProtocol):
CONFIG_KEY_CHECKPOINT_ID not in config[CONF] CONFIG_KEY_CHECKPOINT_ID not in config[CONF]
or CONFIG_KEY_DEDUPE_TASKS in config[CONF] or CONFIG_KEY_DEDUPE_TASKS in config[CONF]
) )
self.trigger_to_nodes = trigger_to_nodes
self.debug = debug self.debug = debug
if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]: if self.stream is not None and CONFIG_KEY_STREAM in config[CONF]:
self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM]) self.stream = DuplexStream(self.stream, config[CONF][CONFIG_KEY_STREAM])
@@ -266,13 +269,13 @@ class PregelLoop(LoopProtocol):
self.checkpoint_config = patch_configurable( self.checkpoint_config = patch_configurable(
self.config, self.config,
{ {
CONFIG_KEY_CHECKPOINT_ID: config[CONF][CONFIG_KEY_CHECKPOINT_MAP][ CONFIG_KEY_CHECKPOINT_ID: self.config[CONF][
self.config[CONF][CONFIG_KEY_CHECKPOINT_NS] CONFIG_KEY_CHECKPOINT_MAP
] ][self.config[CONF][CONFIG_KEY_CHECKPOINT_NS]]
}, },
) )
else: else:
self.checkpoint_config = config self.checkpoint_config = self.config
self.checkpoint_ns = ( self.checkpoint_ns = (
tuple(cast(str, self.config[CONF][CONFIG_KEY_CHECKPOINT_NS]).split(NS_SEP)) tuple(cast(str, self.config[CONF][CONFIG_KEY_CHECKPOINT_NS]).split(NS_SEP))
if self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS) if self.config[CONF].get(CONFIG_KEY_CHECKPOINT_NS)
@@ -406,8 +409,10 @@ class PregelLoop(LoopProtocol):
if self.status != "pending": if self.status != "pending":
raise RuntimeError("Cannot tick when status is no longer 'pending'") raise RuntimeError("Cannot tick when status is no longer 'pending'")
updated_channels: set[str] | None = None
if self.input not in (INPUT_DONE, INPUT_RESUMING, INPUT_SHOULD_VALIDATE): if self.input not in (INPUT_DONE, INPUT_RESUMING, INPUT_SHOULD_VALIDATE):
self._first(input_keys=input_keys) updated_channels = self._first(input_keys=input_keys)
elif self.to_interrupt: elif self.to_interrupt:
# if we need to interrupt, do so # if we need to interrupt, do so
self.status = "interrupt_before" self.status = "interrupt_before"
@@ -427,7 +432,7 @@ class PregelLoop(LoopProtocol):
), ),
) )
# all tasks have finished # all tasks have finished
mv_writes = apply_writes( mv_writes, updated_channels = apply_writes(
self.checkpoint, self.checkpoint,
self.channels, self.channels,
self.tasks.values(), self.tasks.values(),
@@ -493,6 +498,8 @@ class PregelLoop(LoopProtocol):
manager=self.manager, manager=self.manager,
store=self.store, store=self.store,
checkpointer=self.checkpointer, checkpointer=self.checkpointer,
trigger_to_nodes=self.trigger_to_nodes,
updated_channels=updated_channels,
) )
self.to_interrupt = [] self.to_interrupt = []
@@ -571,11 +578,11 @@ class PregelLoop(LoopProtocol):
self.checkpoint["versions_seen"].get(INTERRUPT, {}).values(), self.checkpoint["versions_seen"].get(INTERRUPT, {}).values(),
default=None, default=None,
): ):
self.tasks[tid] = task._replace(scheduled=True) self.tasks[tid] = dataclasses.replace(task, scheduled=True)
else: else:
task.writes.append((k, v)) task.writes.append((k, v))
def _first(self, *, input_keys: Union[str, Sequence[str]]) -> None: def _first(self, *, input_keys: Union[str, Sequence[str]]) -> Optional[set[str]]:
# resuming from previous checkpoint requires # resuming from previous checkpoint requires
# - finding a previous checkpoint # - finding a previous checkpoint
# - receiving None input (outer graph) or RESUMING flag (subgraph) # - receiving None input (outer graph) or RESUMING flag (subgraph)
@@ -592,6 +599,8 @@ class PregelLoop(LoopProtocol):
), ),
) )
) )
# this can be set only when there are input_writes
updated_channels: Optional[set[str]] = None
# map command to writes # map command to writes
if isinstance(self.input, Command): if isinstance(self.input, Command):
@@ -612,7 +621,7 @@ class PregelLoop(LoopProtocol):
if null_writes := [ if null_writes := [
w[1:] for w in self.checkpoint_pending_writes if w[0] == NULL_TASK_ID w[1:] for w in self.checkpoint_pending_writes if w[0] == NULL_TASK_ID
]: ]:
mv_writes = apply_writes( mv_writes, _ = apply_writes(
self.checkpoint, self.checkpoint,
self.channels, self.channels,
[PregelTaskWrites((), INPUT, null_writes, [])], [PregelTaskWrites((), INPUT, null_writes, [])],
@@ -661,7 +670,7 @@ class PregelLoop(LoopProtocol):
manager=None, manager=None,
) )
# apply input writes # apply input writes
mv_writes = apply_writes( mv_writes, updated_channels = apply_writes(
self.checkpoint, self.checkpoint,
self.channels, self.channels,
[ [
@@ -691,6 +700,7 @@ class PregelLoop(LoopProtocol):
self.config = patch_configurable( self.config = patch_configurable(
self.config, {CONFIG_KEY_RESUMING: is_resuming} self.config, {CONFIG_KEY_RESUMING: is_resuming}
) )
return updated_channels
def _put_checkpoint(self, metadata: CheckpointMetadata) -> None: def _put_checkpoint(self, metadata: CheckpointMetadata) -> None:
for k, v in self.config["metadata"].items(): for k, v in self.config["metadata"].items():
@@ -776,7 +786,7 @@ class PregelLoop(LoopProtocol):
and self.checkpoint_pending_writes and self.checkpoint_pending_writes
and any(task.writes for task in self.tasks.values()) and any(task.writes for task in self.tasks.values())
): ):
mv_writes = apply_writes( mv_writes, _ = apply_writes(
self.checkpoint, self.checkpoint,
self.channels, self.channels,
self.tasks.values(), self.tasks.values(),
@@ -883,6 +893,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
input_model: Optional[Type[BaseModel]] = None, input_model: Optional[Type[BaseModel]] = None,
debug: bool = False, debug: bool = False,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
) -> None: ) -> None:
super().__init__( super().__init__(
input, input,
@@ -899,6 +910,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
interrupt_before=interrupt_before, interrupt_before=interrupt_before,
manager=manager, manager=manager,
debug=debug, debug=debug,
trigger_to_nodes=trigger_to_nodes,
) )
self.stack = ExitStack() self.stack = ExitStack()
if checkpointer: if checkpointer:
@@ -1024,6 +1036,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
input_model: Optional[Type[BaseModel]] = None, input_model: Optional[Type[BaseModel]] = None,
debug: bool = False, debug: bool = False,
trigger_to_nodes: Optional[Mapping[str, Sequence[str]]] = None,
) -> None: ) -> None:
super().__init__( super().__init__(
input, input,
@@ -1040,6 +1053,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
interrupt_before=interrupt_before, interrupt_before=interrupt_before,
manager=manager, manager=manager,
debug=debug, debug=debug,
trigger_to_nodes=trigger_to_nodes,
) )
self.stack = AsyncExitStack() self.stack = AsyncExitStack()
if checkpointer: if checkpointer:
+15 -1
View File
@@ -12,7 +12,7 @@ from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.graph import Graph as DrawableGraph from langchain_core.runnables.graph import Graph as DrawableGraph
from typing_extensions import Self from typing_extensions import Self
from langgraph.pregel.types import All, StateSnapshot, StreamMode from langgraph.pregel.types import All, StateSnapshot, StateUpdate, StreamMode
class PregelProtocol( class PregelProtocol(
@@ -69,6 +69,20 @@ class PregelProtocol(
limit: Optional[int] = None, limit: Optional[int] = None,
) -> AsyncIterator[StateSnapshot]: ... ) -> AsyncIterator[StateSnapshot]: ...
@abstractmethod
def bulk_update_state(
self,
config: RunnableConfig,
updates: Sequence[Sequence[StateUpdate]],
) -> RunnableConfig: ...
@abstractmethod
async def abulk_update_state(
self,
config: RunnableConfig,
updates: Sequence[Sequence[StateUpdate]],
) -> RunnableConfig: ...
@abstractmethod @abstractmethod
def update_state( def update_state(
self, self,
+11 -3
View File
@@ -62,7 +62,13 @@ class ChannelRead(RunnableCallable):
mapper: Optional[Callable[[Any], Any]] = None, mapper: Optional[Callable[[Any], Any]] = None,
tags: Optional[list[str]] = None, tags: Optional[list[str]] = None,
) -> None: ) -> None:
super().__init__(func=self._read, afunc=self._aread, tags=tags, name=None) super().__init__(
func=self._read,
afunc=self._aread,
tags=tags,
name=None,
func_accepts_config=True,
)
self.fresh = fresh self.fresh = fresh
self.mapper = mapper self.mapper = mapper
self.channel = channel self.channel = channel
@@ -161,6 +167,7 @@ class PregelNode(Runnable):
metadata: Optional[Mapping[str, Any]] = None, metadata: Optional[Mapping[str, Any]] = None,
bound: Optional[Runnable[Any, Any]] = None, bound: Optional[Runnable[Any, Any]] = None,
retry_policy: Optional[RetryPolicy] = None, retry_policy: Optional[RetryPolicy] = None,
subgraphs: Optional[Sequence[PregelProtocol]] = None,
) -> None: ) -> None:
self.channels = channels self.channels = channels
self.triggers = list(triggers) self.triggers = list(triggers)
@@ -170,7 +177,9 @@ class PregelNode(Runnable):
self.retry_policy = retry_policy self.retry_policy = retry_policy
self.tags = tags self.tags = tags
self.metadata = metadata self.metadata = metadata
if self.bound is not DEFAULT_BOUND: if subgraphs is not None:
self.subgraphs = subgraphs
elif self.bound is not DEFAULT_BOUND:
try: try:
subgraph = find_subgraph_pregel(self.bound) subgraph = find_subgraph_pregel(self.bound)
except Exception: except Exception:
@@ -184,7 +193,6 @@ class PregelNode(Runnable):
def copy(self, update: dict[str, Any]) -> PregelNode: def copy(self, update: dict[str, Any]) -> PregelNode:
attrs = {**self.__dict__, **update} attrs = {**self.__dict__, **update}
attrs.pop("subgraphs")
return PregelNode(**attrs) return PregelNode(**attrs)
@cached_property @cached_property
+14
View File
@@ -457,6 +457,20 @@ class RemoteGraph(PregelProtocol):
for state in states: for state in states:
yield self._create_state_snapshot(state) yield self._create_state_snapshot(state)
def bulk_update_state(
self,
config: RunnableConfig,
updates: list[tuple[Optional[dict[str, Any]], Optional[str]]],
) -> RunnableConfig:
raise NotImplementedError
async def abulk_update_state(
self,
config: RunnableConfig,
updates: list[tuple[Optional[dict[str, Any]], Optional[str]]],
) -> RunnableConfig:
raise NotImplementedError
def update_state( def update_state(
self, self,
config: RunnableConfig, config: RunnableConfig,
+13 -13
View File
@@ -149,7 +149,7 @@ class PregelRunner:
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_call, _call,
t, weakref.ref(t),
retry=retry_policy, retry=retry_policy,
futures=weakref.ref(futures), futures=weakref.ref(futures),
schedule_task=self.schedule_task, schedule_task=self.schedule_task,
@@ -185,7 +185,7 @@ class PregelRunner:
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_call, _call,
t, weakref.ref(t),
retry=retry_policy, retry=retry_policy,
futures=weakref.ref(futures), futures=weakref.ref(futures),
schedule_task=self.schedule_task, schedule_task=self.schedule_task,
@@ -263,7 +263,7 @@ class PregelRunner:
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_acall, _acall,
t, weakref.ref(t),
stream=self.use_astream, stream=self.use_astream,
retry=retry_policy, retry=retry_policy,
futures=weakref.ref(futures), futures=weakref.ref(futures),
@@ -304,7 +304,7 @@ class PregelRunner:
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_acall, _acall,
t, weakref.ref(t),
retry=retry_policy, retry=retry_policy,
stream=self.use_astream, stream=self.use_astream,
futures=weakref.ref(futures), futures=weakref.ref(futures),
@@ -469,7 +469,7 @@ def _panic_or_proceed(
def _call( def _call(
task: PregelExecutableTask, task: weakref.ref[PregelExecutableTask],
func: Callable[[Any], Union[Awaitable[Any], Any]], func: Callable[[Any], Union[Awaitable[Any], Any]],
input: Any, input: Any,
*, *,
@@ -489,10 +489,10 @@ def _call(
fut: Optional[concurrent.futures.Future] = None fut: Optional[concurrent.futures.Future] = None
# schedule PUSH tasks, collect futures # schedule PUSH tasks, collect futures
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD] scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr]
# schedule the next task, if the callback returns one # schedule the next task, if the callback returns one
if next_task := schedule_task()( # type: ignore[misc] if next_task := schedule_task()( # type: ignore[misc]
task, task(), # type: ignore[arg-type]
scratchpad.call_counter(), scratchpad.call_counter(),
Call(func, input, retry=retry, callbacks=callbacks), Call(func, input, retry=retry, callbacks=callbacks),
): ):
@@ -528,7 +528,7 @@ def _call(
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_call, _call,
next_task, weakref.ref(next_task),
futures=futures, futures=futures,
retry=retry, retry=retry,
callbacks=callbacks, callbacks=callbacks,
@@ -550,7 +550,7 @@ def _call(
def _acall( def _acall(
task: PregelExecutableTask, task: weakref.ref[PregelExecutableTask],
func: Callable[[Any], Union[Awaitable[Any], Any]], func: Callable[[Any], Union[Awaitable[Any], Any]],
input: Any, input: Any,
*, *,
@@ -570,10 +570,10 @@ def _acall(
) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]: ) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]:
fut: Optional[asyncio.Future] = None fut: Optional[asyncio.Future] = None
# schedule PUSH tasks, collect futures # schedule PUSH tasks, collect futures
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD] scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr]
# schedule the next task, if the callback returns one # schedule the next task, if the callback returns one
if next_task := schedule_task()( # type: ignore[misc] if next_task := schedule_task()( # type: ignore[misc]
task, task(), # type: ignore[arg-type]
scratchpad.call_counter(), scratchpad.call_counter(),
Call(func, input, retry=retry, callbacks=callbacks), Call(func, input, retry=retry, callbacks=callbacks),
): ):
@@ -614,7 +614,7 @@ def _acall(
configurable={ configurable={
CONFIG_KEY_CALL: partial( CONFIG_KEY_CALL: partial(
_acall, _acall,
next_task, weakref.ref(next_task),
stream=stream, stream=stream,
futures=futures, futures=futures,
schedule_task=schedule_task, schedule_task=schedule_task,
@@ -623,7 +623,7 @@ def _acall(
reraise=reraise, reraise=reraise,
), ),
}, },
__name__=task.name, __name__=task().name, # type: ignore[union-attr]
__cancel_on_exit__=True, __cancel_on_exit__=True,
__reraise_on_exit__=reraise, __reraise_on_exit__=reraise,
# starting a new task in the next tick ensures # starting a new task in the next tick ensures
+2
View File
@@ -7,6 +7,7 @@ from langgraph.types import (
PregelTask, PregelTask,
RetryPolicy, RetryPolicy,
StateSnapshot, StateSnapshot,
StateUpdate,
StreamMode, StreamMode,
StreamWriter, StreamWriter,
default_retry_on, default_retry_on,
@@ -14,6 +15,7 @@ from langgraph.types import (
__all__ = [ __all__ = [
"All", "All",
"StateUpdate",
"CachePolicy", "CachePolicy",
"PregelExecutableTask", "PregelExecutableTask",
"PregelTask", "PregelTask",
+7 -1
View File
@@ -57,7 +57,13 @@ class ChannelWrite(RunnableCallable):
tags: Optional[Sequence[str]] = None, tags: Optional[Sequence[str]] = None,
require_at_least_one_of: Optional[Sequence[str]] = None, # ignored require_at_least_one_of: Optional[Sequence[str]] = None, # ignored
): ):
super().__init__(func=self._write, afunc=self._awrite, name=None, tags=tags) super().__init__(
func=self._write,
afunc=self._awrite,
name=None,
tags=tags,
func_accepts_config=True,
)
self.writes = cast( self.writes = cast(
list[Union[ChannelWriteEntry, ChannelWriteTupleEntry, Send]], writes list[Union[ChannelWriteEntry, ChannelWriteTupleEntry, Send]], writes
) )
+13 -1
View File
@@ -133,6 +133,11 @@ class Interrupt:
when: Literal["during"] = dataclasses.field(default="during", repr=False) when: Literal["during"] = dataclasses.field(default="during", repr=False)
class StateUpdate(NamedTuple):
values: Optional[dict[str, Any]]
as_node: Optional[str] = None
class PregelTask(NamedTuple): class PregelTask(NamedTuple):
id: str id: str
name: str name: str
@@ -143,7 +148,14 @@ class PregelTask(NamedTuple):
result: Optional[Any] = None result: Optional[Any] = None
class PregelExecutableTask(NamedTuple): if sys.version_info > (3, 11):
_T_DC_KWARGS = {"weakref_slot": True, "slots": True, "frozen": True}
else:
_T_DC_KWARGS = {"frozen": True}
@dataclasses.dataclass(**_T_DC_KWARGS)
class PregelExecutableTask:
name: str name: str
input: Any input: Any
proc: Runnable proc: Runnable
+104 -44
View File
@@ -2,8 +2,8 @@ import asyncio
import enum import enum
import inspect import inspect
import sys import sys
from contextlib import AsyncExitStack from contextlib import AsyncExitStack, contextmanager
from contextvars import copy_context from contextvars import Context, Token, copy_context
from functools import partial, wraps from functools import partial, wraps
from typing import ( from typing import (
Any, Any,
@@ -11,6 +11,7 @@ from typing import (
Awaitable, Awaitable,
Callable, Callable,
Coroutine, Coroutine,
Generator,
Iterator, Iterator,
Optional, Optional,
Protocol, Protocol,
@@ -53,13 +54,69 @@ from langgraph.utils.config import (
patch_config, patch_config,
) )
try:
from langchain_core.runnables.config import _set_config_context def _set_config_context(
except ImportError: config: RunnableConfig,
# For forwards compatibility ) -> tuple[Token[Optional[RunnableConfig]], Optional[dict[str, Any]]]:
def _set_config_context(context: RunnableConfig) -> None: # type: ignore """Set the child Runnable config + tracing context.
"""Set the context for the current thread."""
var_child_runnable_config.set(context) Args:
config (RunnableConfig): The config to set.
"""
from langchain_core.tracers.langchain import LangChainTracer
config_token = var_child_runnable_config.set(config)
current_context = None
if (
(callbacks := config.get("callbacks"))
and (
parent_run_id := getattr(callbacks, "parent_run_id", None)
) # Is callback manager
and (
tracer := next(
(
handler
for handler in getattr(callbacks, "handlers", [])
if isinstance(handler, LangChainTracer)
),
None,
)
)
and (run := tracer.run_map.get(str(parent_run_id)))
):
from langsmith.run_helpers import _set_tracing_context, get_tracing_context
current_context = get_tracing_context()
_set_tracing_context({"parent": run})
return config_token, current_context
@contextmanager
def set_config_context(config: RunnableConfig) -> Generator[Context, None, None]:
"""Set the child Runnable config + tracing context.
Args:
config (RunnableConfig): The config to set.
"""
from langsmith.run_helpers import _set_tracing_context
ctx = copy_context()
config_token, _ = ctx.run(_set_config_context, config)
try:
yield ctx
finally:
ctx.run(var_child_runnable_config.reset, config_token)
ctx.run(
_set_tracing_context,
{
"parent": None,
"project_name": None,
"tags": None,
"metadata": None,
"enabled": None,
"client": None,
},
)
# Before Python 3.11 native StrEnum is not available # Before Python 3.11 native StrEnum is not available
@@ -194,6 +251,7 @@ class RunnableCallable(Runnable):
trace: bool = True, trace: bool = True,
recurse: bool = True, recurse: bool = True,
explode_args: bool = False, explode_args: bool = False,
func_accepts_config: Optional[bool] = None,
**kwargs: Any, **kwargs: Any,
) -> None: ) -> None:
self.name = name self.name = name
@@ -219,27 +277,32 @@ class RunnableCallable(Runnable):
# check signature # check signature
if func is None and afunc is None: if func is None and afunc is None:
raise ValueError("At least one of func or afunc must be provided.") raise ValueError("At least one of func or afunc must be provided.")
params = inspect.signature(cast(Callable, func or afunc)).parameters
self.func_accepts_config = "config" in params if func_accepts_config is not None:
# Mapping from kwarg name to (config key, default value) to be used. self.func_accepts_config = func_accepts_config
# The default value is used if the config key is not found in the config. self.func_accepts: dict[str, Tuple[str, Any]] = {}
self.func_accepts: dict[str, Tuple[str, Any]] = {} else:
params = inspect.signature(cast(Callable, func or afunc)).parameters
for kw, typ, config_key, default in KWARGS_CONFIG_KEYS: self.func_accepts_config = "config" in params
p = params.get(kw) # Mapping from kwarg name to (config key, default value) to be used.
# The default value is used if the config key is not found in the config.
self.func_accepts = {}
if p is None or p.kind not in VALID_KINDS: for kw, typ, config_key, default in KWARGS_CONFIG_KEYS:
# If parameter is not found or is not a valid kind, skip p = params.get(kw)
continue
if typ != (ANY_TYPE,) and p.annotation not in typ: if p is None or p.kind not in VALID_KINDS:
# A specific type is required, but the function annotation does # If parameter is not found or is not a valid kind, skip
# not match the expected type. continue
continue
# If the kwarg is accepted by the function, store the default value if typ != (ANY_TYPE,) and p.annotation not in typ:
self.func_accepts[kw] = (config_key, default) # A specific type is required, but the function annotation does
# not match the expected type.
continue
# If the kwarg is accepted by the function, store the default value
self.func_accepts[kw] = (config_key, default)
def __repr__(self) -> str: def __repr__(self) -> str:
repr_args = { repr_args = {
@@ -286,7 +349,6 @@ class RunnableCallable(Runnable):
kwargs[kw] = _conf.get(config_key, default_value) kwargs[kw] = _conf.get(config_key, default_value)
context = copy_context()
if self.trace: if self.trace:
callback_manager = get_callback_manager_for_config(config, self.tags) callback_manager = get_callback_manager_for_config(config, self.tags)
run_manager = callback_manager.on_chain_start( run_manager = callback_manager.on_chain_start(
@@ -297,17 +359,16 @@ class RunnableCallable(Runnable):
) )
try: try:
child_config = patch_config(config, callbacks=run_manager.get_child()) child_config = patch_config(config, callbacks=run_manager.get_child())
context = copy_context() with set_config_context(child_config) as context:
context.run(_set_config_context, child_config) ret = context.run(self.func, *args, **kwargs)
ret = context.run(self.func, *args, **kwargs)
except BaseException as e: except BaseException as e:
run_manager.on_chain_error(e) run_manager.on_chain_error(e)
raise raise
else: else:
run_manager.on_chain_end(ret) run_manager.on_chain_end(ret)
else: else:
context.run(_set_config_context, config) with set_config_context(config) as context:
ret = context.run(self.func, *args, **kwargs) ret = context.run(self.func, *args, **kwargs)
if isinstance(ret, Runnable) and self.recurse: if isinstance(ret, Runnable) and self.recurse:
return ret.invoke(input, config) return ret.invoke(input, config)
return ret return ret
@@ -342,7 +403,6 @@ class RunnableCallable(Runnable):
f"Missing required config key '{config_key}' for '{self.name}'." f"Missing required config key '{config_key}' for '{self.name}'."
) )
kwargs[kw] = _conf.get(config_key, default_value) kwargs[kw] = _conf.get(config_key, default_value)
context = copy_context()
if self.trace: if self.trace:
callback_manager = get_async_callback_manager_for_config(config, self.tags) callback_manager = get_async_callback_manager_for_config(config, self.tags)
run_manager = await callback_manager.on_chain_start( run_manager = await callback_manager.on_chain_start(
@@ -353,24 +413,24 @@ class RunnableCallable(Runnable):
) )
try: try:
child_config = patch_config(config, callbacks=run_manager.get_child()) child_config = patch_config(config, callbacks=run_manager.get_child())
context.run(_set_config_context, child_config) with set_config_context(child_config) as context:
coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs)) coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs))
if ASYNCIO_ACCEPTS_CONTEXT: if ASYNCIO_ACCEPTS_CONTEXT:
ret = await asyncio.create_task(coro, context=context) ret = await asyncio.create_task(coro, context=context)
else: else:
ret = await coro ret = await coro
except BaseException as e: except BaseException as e:
await run_manager.on_chain_error(e) await run_manager.on_chain_error(e)
raise raise
else: else:
await run_manager.on_chain_end(ret) await run_manager.on_chain_end(ret)
else: else:
context.run(_set_config_context, config) with set_config_context(config) as context:
if ASYNCIO_ACCEPTS_CONTEXT: if ASYNCIO_ACCEPTS_CONTEXT:
coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs)) coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs))
ret = await asyncio.create_task(coro, context=context) ret = await asyncio.create_task(coro, context=context)
else: else:
ret = await self.afunc(*args, **kwargs) ret = await self.afunc(*args, **kwargs)
if isinstance(ret, Runnable) and self.recurse: if isinstance(ret, Runnable) and self.recurse:
return await ret.ainvoke(input, config) return await ret.ainvoke(input, config)
return ret return ret
+9 -9
View File
@@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 2.0.0 and should not be changed by hand. # This file is automatically @generated by Poetry 2.0.1 and should not be changed by hand.
[[package]] [[package]]
name = "aiosqlite" name = "aiosqlite"
@@ -1324,14 +1324,14 @@ files = [
[[package]] [[package]]
name = "langchain-core" name = "langchain-core"
version = "0.3.44" version = "0.3.46"
description = "Building applications with LLMs through composability" description = "Building applications with LLMs through composability"
optional = false optional = false
python-versions = "<4.0,>=3.9" python-versions = "<4.0,>=3.9"
groups = ["main", "dev"] groups = ["main", "dev"]
files = [ files = [
{file = "langchain_core-0.3.44-py3-none-any.whl", hash = "sha256:d989ce8bd62f1d07765acd575e6ec1254aec0cf7775aaea39fe4af8102377459"}, {file = "langchain_core-0.3.46-py3-none-any.whl", hash = "sha256:28b5689fc347975ea520b5364ab4aee5567e661553bbee5e97cabf4596c28ce0"},
{file = "langchain_core-0.3.44.tar.gz", hash = "sha256:7c0a01e78360f007cbca448178fe7e032404068e6431dbe8ce905f84febbdfa5"}, {file = "langchain_core-0.3.46.tar.gz", hash = "sha256:5fca010eeb0a427be5aa8a8525e2112995dde790c584cef165be7c5e0ee1c2b5"},
] ]
[package.dependencies] [package.dependencies]
@@ -1348,7 +1348,7 @@ typing-extensions = ">=4.7"
[[package]] [[package]]
name = "langgraph-checkpoint" name = "langgraph-checkpoint"
version = "2.0.18" version = "2.0.21"
description = "Library with base interfaces for LangGraph checkpoint savers." description = "Library with base interfaces for LangGraph checkpoint savers."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1366,7 +1366,7 @@ url = "../checkpoint"
[[package]] [[package]]
name = "langgraph-checkpoint-postgres" name = "langgraph-checkpoint-postgres"
version = "2.0.16" version = "2.0.19"
description = "Library with a Postgres implementation of LangGraph checkpoint saver." description = "Library with a Postgres implementation of LangGraph checkpoint saver."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1375,7 +1375,7 @@ files = []
develop = true develop = true
[package.dependencies] [package.dependencies]
langgraph-checkpoint = "^2.0.15" langgraph-checkpoint = "^2.0.21"
orjson = ">=3.10.1" orjson = ">=3.10.1"
psycopg = "^3.2.0" psycopg = "^3.2.0"
psycopg-pool = "^3.2.0" psycopg-pool = "^3.2.0"
@@ -1404,7 +1404,7 @@ url = "../checkpoint-sqlite"
[[package]] [[package]]
name = "langgraph-prebuilt" name = "langgraph-prebuilt"
version = "0.1.2" version = "0.1.4"
description = "Library with high-level APIs for creating and executing LangGraph agents and tools." description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -1422,7 +1422,7 @@ url = "../prebuilt"
[[package]] [[package]]
name = "langgraph-sdk" name = "langgraph-sdk"
version = "0.1.55" version = "0.1.58"
description = "SDK for interacting with LangGraph API" description = "SDK for interacting with LangGraph API"
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph" name = "langgraph"
version = "0.3.16" version = "0.3.18"
description = "Building stateful, multi-actor applications with LLMs" description = "Building stateful, multi-actor applications with LLMs"
authors = [] authors = []
license = "MIT" license = "MIT"
+2 -59
View File
@@ -1,4 +1,3 @@
import asyncio
import os import os
import tempfile import tempfile
from collections import defaultdict from collections import defaultdict
@@ -13,7 +12,6 @@ from langgraph.checkpoint.base import (
CheckpointMetadata, CheckpointMetadata,
CheckpointTuple, CheckpointTuple,
SerializerProtocol, SerializerProtocol,
copy_checkpoint,
) )
from langgraph.checkpoint.memory import InMemorySaver, PersistentDict from langgraph.checkpoint.memory import InMemorySaver, PersistentDict
@@ -63,69 +61,14 @@ class MemorySaverAssertImmutable(InMemorySaver):
self.storage_for_copies[thread_id][checkpoint_ns][saved["id"]] self.storage_for_copies[thread_id][checkpoint_ns][saved["id"]]
) )
== saved == saved
) ), config["configurable"]["checkpoint_ns"]
self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = ( self.storage_for_copies[thread_id][checkpoint_ns][checkpoint["id"]] = (
self.serde.dumps_typed(copy_checkpoint(checkpoint)) self.serde.dumps_typed(checkpoint)
) )
# call super to write checkpoint # call super to write checkpoint
return super().put(config, checkpoint, metadata, new_versions) return super().put(config, checkpoint, metadata, new_versions)
class MemorySaverAssertCheckpointMetadata(InMemorySaver):
"""This custom checkpointer is for verifying that a run's configurable
fields are merged with the previous checkpoint config for each step in
the run. This is the desired behavior. Because the checkpointer's (a)put()
method is called for each step, the implementation of this checkpointer
should produce a side effect that can be asserted.
"""
def put(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> None:
"""The implementation of put() merges config["configurable"] (a run's
configurable fields) with the metadata field. The state of the
checkpoint metadata can be asserted to confirm that the run's
configurable fields were merged with the previous checkpoint config.
"""
configurable = config["configurable"].copy()
# remove checkpoint_id to make testing simpler
checkpoint_id = configurable.pop("checkpoint_id", None)
thread_id = config["configurable"]["thread_id"]
checkpoint_ns = config["configurable"]["checkpoint_ns"]
self.storage[thread_id][checkpoint_ns].update(
{
checkpoint["id"]: (
self.serde.dumps_typed(checkpoint),
# merge configurable fields and metadata
self.serde.dumps_typed({**configurable, **metadata}),
checkpoint_id,
)
}
)
return {
"configurable": {
"thread_id": config["configurable"]["thread_id"],
"checkpoint_id": checkpoint["id"],
}
}
async def aput(
self,
config: RunnableConfig,
checkpoint: Checkpoint,
metadata: CheckpointMetadata,
new_versions: ChannelVersions,
) -> RunnableConfig:
return await asyncio.get_running_loop().run_in_executor(
None, self.put, config, checkpoint, metadata, new_versions
)
class MemorySaverNoPending(InMemorySaver): class MemorySaverNoPending(InMemorySaver):
def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]: def get_tuple(self, config: RunnableConfig) -> Optional[CheckpointTuple]:
result = super().get_tuple(config) result = super().get_tuple(config)
+38 -23
View File
@@ -2483,7 +2483,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
{ {
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("start:agent",), "langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
"langgraph_path": (PULL, "agent"), "langgraph_path": (PULL, "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -2542,7 +2542,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
{ {
"langgraph_step": 3, "langgraph_step": 3,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("tools",), "langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
"langgraph_path": (PULL, "agent"), "langgraph_path": (PULL, "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -2585,7 +2585,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
{ {
"langgraph_step": 5, "langgraph_step": 5,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("tools",), "langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
"langgraph_path": (PULL, "agent"), "langgraph_path": (PULL, "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -5501,7 +5501,10 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "rewrite_query", "name": "rewrite_query",
"input": {"query": "what is weather in sf", "docs": []}, "input": {"query": "what is weather in sf", "docs": []},
"triggers": ("start:rewrite_query",), "triggers": (
"branch:to:rewrite_query",
"start:rewrite_query",
),
}, },
}, },
), ),
@@ -5532,7 +5535,10 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "retriever_one", "name": "retriever_one",
"input": {"query": "query: what is weather in sf", "docs": []}, "input": {"query": "query: what is weather in sf", "docs": []},
"triggers": ("rewrite_query",), "triggers": (
"branch:to:retriever_one",
"rewrite_query",
),
}, },
}, },
), ),
@@ -5546,7 +5552,10 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "retriever_two", "name": "retriever_two",
"input": {"query": "query: what is weather in sf", "docs": []}, "input": {"query": "query: what is weather in sf", "docs": []},
"triggers": ("rewrite_query",), "triggers": (
"branch:to:retriever_two",
"rewrite_query",
),
}, },
}, },
), ),
@@ -5608,7 +5617,7 @@ def test_in_one_fan_out_out_one_graph_state() -> None:
"query": "query: what is weather in sf", "query": "query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"], "docs": ["doc1", "doc2", "doc3", "doc4"],
}, },
"triggers": (AnyStr("retriever_"),), "triggers": ("branch:to:qa", "retriever_one", "retriever_two"),
}, },
}, },
), ),
@@ -6634,7 +6643,7 @@ def test_branch_then(
"id": AnyStr(), "id": AnyStr(),
"name": "prepare", "name": "prepare",
"input": {"my_key": "value", "market": "DE"}, "input": {"my_key": "value", "market": "DE"},
"triggers": ("start:prepare",), "triggers": ("branch:to:prepare", "start:prepare"),
}, },
}, },
{ {
@@ -6773,7 +6782,10 @@ def test_branch_then(
"id": AnyStr(), "id": AnyStr(),
"name": "finish", "name": "finish",
"input": {"my_key": "value prepared slow", "market": "DE"}, "input": {"my_key": "value prepared slow", "market": "DE"},
"triggers": ("branch:prepare:condition::then",), "triggers": (
"branch:prepare:condition::then",
"branch:to:finish",
),
}, },
}, },
{ {
@@ -7783,7 +7795,7 @@ def test_nested_graph_state(
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -7978,7 +7990,7 @@ def test_nested_graph_state(
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -8021,7 +8033,7 @@ def test_nested_graph_state(
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -8070,7 +8082,7 @@ def test_nested_graph_state(
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -8504,7 +8516,7 @@ def test_doubly_nested_graph_state(
"langgraph_node": "child_1", "langgraph_node": "child_1",
"langgraph_path": [PULL, AnyStr("child_1")], "langgraph_path": [PULL, AnyStr("child_1")],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=( parent_config=(
@@ -8588,7 +8600,10 @@ def test_doubly_nested_graph_state(
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
AnyStr("start:child_1"),
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=( parent_config=(
@@ -8635,7 +8650,7 @@ def test_doubly_nested_graph_state(
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -8931,7 +8946,7 @@ def test_doubly_nested_graph_state(
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -8970,7 +8985,7 @@ def test_doubly_nested_graph_state(
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -9022,7 +9037,7 @@ def test_doubly_nested_graph_state(
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -9076,7 +9091,7 @@ def test_doubly_nested_graph_state(
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -9131,7 +9146,7 @@ def test_doubly_nested_graph_state(
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -9193,7 +9208,7 @@ def test_doubly_nested_graph_state(
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -9255,7 +9270,7 @@ def test_doubly_nested_graph_state(
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=None, parent_config=None,
+72 -24
View File
@@ -2300,7 +2300,11 @@ async def test_prebuilt_tool_chat() -> None:
{ {
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("start:agent",), "langgraph_triggers": (
"branch:to:agent",
"start:agent",
"tools",
),
"langgraph_path": ("__pregel_pull", "agent"), "langgraph_path": ("__pregel_pull", "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -2359,7 +2363,11 @@ async def test_prebuilt_tool_chat() -> None:
{ {
"langgraph_step": 3, "langgraph_step": 3,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("tools",), "langgraph_triggers": (
"branch:to:agent",
"start:agent",
"tools",
),
"langgraph_path": ("__pregel_pull", "agent"), "langgraph_path": ("__pregel_pull", "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -2402,7 +2410,11 @@ async def test_prebuilt_tool_chat() -> None:
{ {
"langgraph_step": 5, "langgraph_step": 5,
"langgraph_node": "agent", "langgraph_node": "agent",
"langgraph_triggers": ("tools",), "langgraph_triggers": (
"branch:to:agent",
"start:agent",
"tools",
),
"langgraph_path": ("__pregel_pull", "agent"), "langgraph_path": ("__pregel_pull", "agent"),
"langgraph_checkpoint_ns": AnyStr("agent:"), "langgraph_checkpoint_ns": AnyStr("agent:"),
"checkpoint_ns": AnyStr("agent:"), "checkpoint_ns": AnyStr("agent:"),
@@ -3883,7 +3895,10 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "rewrite_query", "name": "rewrite_query",
"input": {"query": "what is weather in sf", "docs": []}, "input": {"query": "what is weather in sf", "docs": []},
"triggers": ("start:rewrite_query",), "triggers": (
"branch:to:rewrite_query",
"start:rewrite_query",
),
}, },
}, },
), ),
@@ -3914,7 +3929,10 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "retriever_one", "name": "retriever_one",
"input": {"query": "query: what is weather in sf", "docs": []}, "input": {"query": "query: what is weather in sf", "docs": []},
"triggers": ("rewrite_query",), "triggers": (
"branch:to:retriever_one",
"rewrite_query",
),
}, },
}, },
), ),
@@ -3928,7 +3946,10 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "retriever_two", "name": "retriever_two",
"input": {"query": "query: what is weather in sf", "docs": []}, "input": {"query": "query: what is weather in sf", "docs": []},
"triggers": ("rewrite_query",), "triggers": (
"branch:to:retriever_two",
"rewrite_query",
),
}, },
}, },
), ),
@@ -3990,7 +4011,7 @@ async def test_in_one_fan_out_out_one_graph_state() -> None:
"query": "query: what is weather in sf", "query": "query: what is weather in sf",
"docs": ["doc1", "doc2", "doc3", "doc4"], "docs": ["doc1", "doc2", "doc3", "doc4"],
}, },
"triggers": (AnyStr("retriever_"),), "triggers": ("branch:to:qa", "retriever_one", "retriever_two"),
}, },
}, },
), ),
@@ -4465,7 +4486,10 @@ async def test_branch_then(checkpointer_name: str) -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "prepare", "name": "prepare",
"input": {"my_key": "value", "market": "DE"}, "input": {"my_key": "value", "market": "DE"},
"triggers": ("start:prepare",), "triggers": (
"branch:to:prepare",
"start:prepare",
),
}, },
}, },
{ {
@@ -4609,7 +4633,10 @@ async def test_branch_then(checkpointer_name: str) -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "finish", "name": "finish",
"input": {"my_key": "value prepared slow", "market": "DE"}, "input": {"my_key": "value prepared slow", "market": "DE"},
"triggers": ("branch:prepare:condition::then",), "triggers": (
"branch:prepare:condition::then",
"branch:to:finish",
),
}, },
}, },
{ {
@@ -4778,7 +4805,10 @@ async def test_branch_then(checkpointer_name: str) -> None:
"id": AnyStr(), "id": AnyStr(),
"name": "prepare", "name": "prepare",
"input": {"my_key": "value", "market": "DE"}, "input": {"my_key": "value", "market": "DE"},
"triggers": ("start:prepare",), "triggers": (
"branch:to:prepare",
"start:prepare",
),
}, },
}, },
{ {
@@ -5333,7 +5363,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -5530,7 +5560,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -5573,7 +5603,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -5622,7 +5652,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "inner", "langgraph_node": "inner",
"langgraph_path": [PULL, "inner"], "langgraph_path": [PULL, "inner"],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": ["outer_1"], "langgraph_triggers": ["branch:to:inner", "outer_1"],
"langgraph_checkpoint_ns": AnyStr("inner:"), "langgraph_checkpoint_ns": AnyStr("inner:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -6060,7 +6090,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "child_1", "langgraph_node": "child_1",
"langgraph_path": [PULL, AnyStr("child_1")], "langgraph_path": [PULL, AnyStr("child_1")],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": ["branch:to:child_1", "start:child_1"],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=( parent_config=(
@@ -6146,7 +6176,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
"start:child_1",
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=( parent_config=(
@@ -6195,7 +6228,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": [
"branch:to:child",
AnyStr("parent_1"),
],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -6493,7 +6529,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -6532,7 +6568,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -6584,7 +6620,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
"langgraph_node": "child", "langgraph_node": "child",
"langgraph_path": [PULL, AnyStr("child")], "langgraph_path": [PULL, AnyStr("child")],
"langgraph_step": 2, "langgraph_step": 2,
"langgraph_triggers": [AnyStr("parent_1")], "langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
"langgraph_checkpoint_ns": AnyStr("child:"), "langgraph_checkpoint_ns": AnyStr("child:"),
}, },
created_at=AnyStr(), created_at=AnyStr(),
@@ -6642,7 +6678,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
AnyStr("start:child_1"),
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -6697,7 +6736,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
AnyStr("start:child_1"),
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -6759,7 +6801,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
AnyStr("start:child_1"),
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config={ parent_config={
@@ -6821,7 +6866,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
AnyStr("child_1"), AnyStr("child_1"),
], ],
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_triggers": [AnyStr("start:child_1")], "langgraph_triggers": [
"branch:to:child_1",
AnyStr("start:child_1"),
],
}, },
created_at=AnyStr(), created_at=AnyStr(),
parent_config=None, parent_config=None,
+328 -57
View File
@@ -1,5 +1,6 @@
import enum import enum
import functools import functools
import gc
import json import json
import logging import logging
import operator import operator
@@ -62,13 +63,16 @@ from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessageGraph, MessagesState, add_messages from langgraph.graph.message import MessageGraph, MessagesState, add_messages
from langgraph.prebuilt.tool_node import ToolNode from langgraph.prebuilt.tool_node import ToolNode
from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot
from langgraph.pregel.loop import SyncPregelLoop
from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.retry import RetryPolicy
from langgraph.pregel.runner import PregelRunner
from langgraph.store.base import BaseStore from langgraph.store.base import BaseStore
from langgraph.types import ( from langgraph.types import (
Command, Command,
Interrupt, Interrupt,
PregelTask, PregelTask,
Send, Send,
StateUpdate,
StreamWriter, StreamWriter,
interrupt, interrupt,
) )
@@ -80,7 +84,6 @@ from tests.conftest import (
REGULAR_CHECKPOINTERS_SYNC, REGULAR_CHECKPOINTERS_SYNC,
SHOULD_CHECK_SNAPSHOTS, SHOULD_CHECK_SNAPSHOTS,
) )
from tests.memory_assert import MemorySaverAssertCheckpointMetadata
from tests.messages import ( from tests.messages import (
_AnyIdAIMessage, _AnyIdAIMessage,
_AnyIdAIMessageChunk, _AnyIdAIMessageChunk,
@@ -4209,11 +4212,11 @@ def test_checkpoint_metadata() -> None:
workflow.add_edge("tools", "agent") workflow.add_edge("tools", "agent")
# graph w/o interrupt # graph w/o interrupt
checkpointer_1 = MemorySaverAssertCheckpointMetadata() checkpointer_1 = InMemorySaver()
app = workflow.compile(checkpointer=checkpointer_1) app = workflow.compile(checkpointer=checkpointer_1)
# graph w/ interrupt # graph w/ interrupt
checkpointer_2 = MemorySaverAssertCheckpointMetadata() checkpointer_2 = InMemorySaver()
app_w_interrupt = workflow.compile( app_w_interrupt = workflow.compile(
checkpointer=checkpointer_2, interrupt_before=["tools"] checkpointer=checkpointer_2, interrupt_before=["tools"]
) )
@@ -4631,59 +4634,6 @@ def test_multiple_sinks_subgraphs(snapshot: SnapshotAssertion) -> None:
assert app.get_graph(xray=True).draw_mermaid() == snapshot assert app.get_graph(xray=True).draw_mermaid() == snapshot
def test_subgraph_retries():
class State(TypedDict):
count: int
class ChildState(State):
some_list: Annotated[list, operator.add]
called_times = 0
class RandomError(ValueError):
"""This will be retried on."""
def parent_node(state: State):
return {"count": state["count"] + 1}
def child_node_a(state: ChildState):
nonlocal called_times
# We want it to retry only on node_b
# NOT re-compute the whole graph.
assert not called_times
called_times += 1
return {"some_list": ["val"]}
def child_node_b(state: ChildState):
raise RandomError("First attempt fails")
child = StateGraph(ChildState)
child.add_node(child_node_a)
child.add_node(child_node_b)
child.add_edge("__start__", "child_node_a")
child.add_edge("child_node_a", "child_node_b")
parent = StateGraph(State)
parent.add_node("parent_node", parent_node)
parent.add_node(
"child_graph",
child.compile(),
retry=RetryPolicy(
max_attempts=3,
retry_on=(RandomError,),
backoff_factor=0.0001,
initial_interval=0.0001,
),
)
parent.add_edge("parent_node", "child_graph")
parent.set_entry_point("parent_node")
checkpointer = InMemorySaver()
app = parent.compile(checkpointer=checkpointer)
with pytest.raises(RandomError):
app.invoke({"count": 0}, {"configurable": {"thread_id": "foo"}})
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
@pytest.mark.parametrize("store_name", ALL_STORES_SYNC) @pytest.mark.parametrize("store_name", ALL_STORES_SYNC)
def test_store_injected( def test_store_injected(
@@ -6290,6 +6240,7 @@ def test_double_interrupt_subgraph(
def invoke_sub_agent(state: AgentState): def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state) return subgraph.invoke(state)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
parent_agent = ( parent_agent = (
StateGraph(AgentState) StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent) .add_node("invoke_sub_agent", invoke_sub_agent)
@@ -6925,7 +6876,10 @@ def test_tags_stream_mode_messages() -> None:
{ {
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_node": "call_model", "langgraph_node": "call_model",
"langgraph_triggers": ("start:call_model",), "langgraph_triggers": (
"branch:to:call_model",
"start:call_model",
),
"langgraph_path": ("__pregel_pull", "call_model"), "langgraph_path": ("__pregel_pull", "call_model"),
"langgraph_checkpoint_ns": AnyStr("call_model:"), "langgraph_checkpoint_ns": AnyStr("call_model:"),
"checkpoint_ns": AnyStr("call_model:"), "checkpoint_ns": AnyStr("call_model:"),
@@ -7613,3 +7567,320 @@ def test_parallel_interrupts_double(
assert invokes == 5 assert invokes == 5
assert len(events) == 5 assert len(events) == 5
def test_pregel_loop_refcount():
gc.collect()
try:
gc.disable()
class State(TypedDict):
messages: Annotated[list, add_messages]
graph_builder = StateGraph(State)
def chatbot(state: State):
return {"messages": [("ai", "HIYA")]}
graph_builder.add_node("chatbot", chatbot)
graph_builder.set_entry_point("chatbot")
graph_builder.set_finish_point("chatbot")
graph = graph_builder.compile()
for _ in range(5):
graph.invoke({"messages": [{"role": "user", "content": "hi"}]})
assert (
len(
[obj for obj in gc.get_objects() if isinstance(obj, SyncPregelLoop)]
)
== 0
)
assert (
len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)])
== 0
)
finally:
gc.enable()
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC)
def test_bulk_state_updates(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class State(TypedDict):
foo: str
baz: str
def node_a(state: State) -> State:
return {"foo": "bar"}
def node_b(state: State) -> State:
return {"baz": "qux"}
graph = (
StateGraph(State)
.add_node("node_a", node_a)
.add_node("node_b", node_b)
.add_edge(START, "node_a")
.add_edge("node_a", "node_b")
.compile(checkpointer=checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First update with node_a
graph.bulk_update_state(
config,
[
[
StateUpdate(values={"foo": "bar"}, as_node="node_a"),
]
],
)
# Then bulk update with both nodes
graph.bulk_update_state(
config,
[
[
StateUpdate(values={"foo": "updated"}, as_node="node_a"),
StateUpdate(values={"baz": "new"}, as_node="node_b"),
]
],
)
state = graph.get_state(config)
assert state.values == {"foo": "updated", "baz": "new"}
# Check if there are only two checkpoints
checkpoints = list(checkpointer.list(config))
assert len(checkpoints) == 2
assert checkpoints[0].metadata["writes"] == {
"node_a": {"foo": "updated"},
"node_b": {"baz": "new"},
}
assert checkpoints[1].metadata["writes"] == {"node_a": {"foo": "bar"}}
# perform multiple steps at the same time
config = {"configurable": {"thread_id": "2"}}
graph.bulk_update_state(
config,
[
[
StateUpdate(values={"foo": "bar"}, as_node="node_a"),
],
[
StateUpdate(values={"foo": "updated"}, as_node="node_a"),
StateUpdate(values={"baz": "new"}, as_node="node_b"),
],
],
)
state = graph.get_state(config)
assert state.values == {"foo": "updated", "baz": "new"}
checkpoints = list(checkpointer.list(config))
assert len(checkpoints) == 2
assert checkpoints[0].metadata["writes"] == {
"node_a": {"foo": "updated"},
"node_b": {"baz": "new"},
}
assert checkpoints[1].metadata["writes"] == {"node_a": {"foo": "bar"}}
# Should raise error if updating without as_node
with pytest.raises(InvalidUpdateError):
graph.bulk_update_state(
config,
[
[
StateUpdate(values={"foo": "error"}, as_node=None),
StateUpdate(values={"bar": "error"}, as_node=None),
]
],
)
# Should raise if no updates are provided
with pytest.raises(ValueError, match="No supersteps provided"):
graph.bulk_update_state(config, [])
# Should raise if no updates are provided
with pytest.raises(ValueError, match="No updates provided"):
graph.bulk_update_state(config, [[], []])
# Should raise if __end__ or __copy__ update is applied in bulk
with pytest.raises(InvalidUpdateError):
graph.bulk_update_state(
config,
[
[
StateUpdate(values=None, as_node="__end__"),
StateUpdate(values=None, as_node="__copy__"),
],
],
)
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC)
def test_update_as_input(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class State(TypedDict):
foo: str
def agent(state: State) -> State:
return {"foo": "agent"}
def tool(state: State) -> State:
return {"foo": "tool"}
graph = (
StateGraph(State)
.add_node("agent", agent)
.add_node("tool", tool)
.add_edge(START, "agent")
.add_edge("agent", "tool")
.compile(checkpointer=checkpointer)
)
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
"foo": "tool"
}
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
"foo": "tool"
}
def map_snapshot(i: StateSnapshot) -> dict:
return {
"values": i.values,
"next": i.next,
"step": i.metadata.get("step"),
}
history = [
map_snapshot(s)
for s in graph.get_state_history({"configurable": {"thread_id": "1"}})
]
graph.bulk_update_state(
{"configurable": {"thread_id": "2"}},
[
# First turn
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent"}, "agent")],
[StateUpdate({"foo": "tool"}, "tool")],
# Second turn
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent"}, "agent")],
[StateUpdate({"foo": "tool"}, "tool")],
],
)
state = graph.get_state({"configurable": {"thread_id": "2"}})
assert state.values == {"foo": "tool"}
new_history = [
map_snapshot(s)
for s in graph.get_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_SYNC)
def test_batch_update_as_input(
request: pytest.FixtureRequest, checkpointer_name: str
) -> None:
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
class State(TypedDict):
foo: str
tasks: Annotated[list[int], operator.add]
def agent(state: State) -> State:
return {"foo": "agent"}
def map(state: State) -> Command["task"]:
return Command(
goto=[
Send("task", {"index": 0}),
Send("task", {"index": 1}),
Send("task", {"index": 2}),
],
update={"foo": "map"},
)
def task(state: dict) -> State:
return {"tasks": [state["index"]]}
graph = (
StateGraph(State)
.add_node("agent", agent)
.add_node("map", map)
.add_node("task", task)
.add_edge(START, "agent")
.add_edge("agent", "map")
.compile(checkpointer=checkpointer)
)
assert graph.invoke({"foo": "input"}, {"configurable": {"thread_id": "1"}}) == {
"foo": "map",
"tasks": [0, 1, 2],
}
def map_snapshot(i: StateSnapshot) -> dict:
return {
"values": i.values,
"next": i.next,
"step": i.metadata.get("step"),
"tasks": [t.name for t in i.tasks],
}
history = [
map_snapshot(s)
for s in graph.get_state_history({"configurable": {"thread_id": "1"}})
]
graph.bulk_update_state(
{"configurable": {"thread_id": "2"}},
[
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent", "tasks": []}, "agent")],
[
StateUpdate(
Command(
goto=[
Send("task", {"index": 0}),
Send("task", {"index": 1}),
Send("task", {"index": 2}),
],
update={"foo": "map"},
),
"map",
)
],
[
StateUpdate({"tasks": [0]}, "task"),
StateUpdate({"tasks": [1]}, "task"),
StateUpdate({"tasks": [2]}, "task"),
],
],
)
state = graph.get_state({"configurable": {"thread_id": "2"}})
assert state.values == {"foo": "map", "tasks": [0, 1, 2]}
new_history = [
map_snapshot(s)
for s in graph.get_state_history({"configurable": {"thread_id": "2"}})
]
assert new_history == history
+348 -8
View File
@@ -1,5 +1,7 @@
import asyncio import asyncio
import enum
import functools import functools
import gc
import logging import logging
import operator import operator
import random import random
@@ -52,13 +54,16 @@ from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.message import MessagesState, add_messages from langgraph.graph.message import MessagesState, add_messages
from langgraph.prebuilt.tool_node import ToolNode from langgraph.prebuilt.tool_node import ToolNode
from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot from langgraph.pregel import Channel, GraphRecursionError, Pregel, StateSnapshot
from langgraph.pregel.loop import AsyncPregelLoop
from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.retry import RetryPolicy
from langgraph.pregel.runner import PregelRunner
from langgraph.store.base import BaseStore from langgraph.store.base import BaseStore
from langgraph.types import ( from langgraph.types import (
Command, Command,
Interrupt, Interrupt,
PregelTask, PregelTask,
Send, Send,
StateUpdate,
StreamWriter, StreamWriter,
interrupt, interrupt,
) )
@@ -73,10 +78,7 @@ from tests.conftest import (
awith_store, awith_store,
) )
from tests.fake_tracer import FakeTracer from tests.fake_tracer import FakeTracer
from tests.memory_assert import ( from tests.memory_assert import MemorySaverNoPending
MemorySaverAssertCheckpointMetadata,
MemorySaverNoPending,
)
from tests.messages import ( from tests.messages import (
_AnyIdAIMessage, _AnyIdAIMessage,
_AnyIdAIMessageChunk, _AnyIdAIMessageChunk,
@@ -4542,8 +4544,13 @@ async def test_nested_pydantic_models(version: str) -> None:
name: str name: str
friends: list[str] = Field(default_factory=list) # IDs of friends friends: list[str] = Field(default_factory=list) # IDs of friends
class MyEnum(enum.Enum):
A = 1
B = 2
class MyTypedDict(TypedDict): class MyTypedDict(TypedDict):
x: int x: int
my_enum: MyEnum
class State(BaseModel): class State(BaseModel):
# Basic nested model tests # Basic nested model tests
@@ -4552,6 +4559,7 @@ async def test_nested_pydantic_models(version: str) -> None:
optional_nested: Optional[NestedModel] = None optional_nested: Optional[NestedModel] = None
dict_nested: dict[str, NestedModel] dict_nested: dict[str, NestedModel]
my_set: set[int] my_set: set[int]
my_enum: MyEnum
list_nested: Annotated[ list_nested: Annotated[
Union[dict, list[dict[str, NestedModel]]], lambda x, y: (x or []) + [y] Union[dict, list[dict[str, NestedModel]]], lambda x, y: (x or []) + [y]
] ]
@@ -4579,7 +4587,8 @@ async def test_nested_pydantic_models(version: str) -> None:
"nested": {"value": 42, "name": "test"}, "nested": {"value": 42, "name": "test"},
"optional_nested": {"value": 10, "name": "optional"}, "optional_nested": {"value": 10, "name": "optional"},
"my_set": [1, 2, 7], "my_set": [1, 2, 7],
"my_typed_dict": {"x": 1}, "my_enum": MyEnum.B,
"my_typed_dict": {"x": 1, "my_enum": MyEnum.A},
"dict_nested": {"a": {"value": 5, "name": "a"}}, "dict_nested": {"a": {"value": 5, "name": "a"}},
"list_nested": [{"a": {"value": 6, "name": "b"}}], "list_nested": [{"a": {"value": 6, "name": "b"}}],
"list_nested_reversed": ["foo", "bar"], "list_nested_reversed": ["foo", "bar"],
@@ -5758,11 +5767,11 @@ async def test_checkpoint_metadata() -> None:
workflow.add_edge("tools", "agent") workflow.add_edge("tools", "agent")
# graph w/o interrupt # graph w/o interrupt
checkpointer_1 = MemorySaverAssertCheckpointMetadata() checkpointer_1 = InMemorySaver()
app = workflow.compile(checkpointer=checkpointer_1) app = workflow.compile(checkpointer=checkpointer_1)
# graph w/ interrupt # graph w/ interrupt
checkpointer_2 = MemorySaverAssertCheckpointMetadata() checkpointer_2 = InMemorySaver()
app_w_interrupt = workflow.compile( app_w_interrupt = workflow.compile(
checkpointer=checkpointer_2, interrupt_before=["tools"] checkpointer=checkpointer_2, interrupt_before=["tools"]
) )
@@ -7008,6 +7017,8 @@ async def test_double_interrupt_subgraph(checkpointer_name: str) -> None:
def invoke_sub_agent(state: AgentState): def invoke_sub_agent(state: AgentState):
return subgraph.invoke(state) return subgraph.invoke(state)
thread = {"configurable": {"thread_id": str(uuid.uuid4())}}
parent_agent = ( parent_agent = (
StateGraph(AgentState) StateGraph(AgentState)
.add_node("invoke_sub_agent", invoke_sub_agent) .add_node("invoke_sub_agent", invoke_sub_agent)
@@ -7582,7 +7593,10 @@ async def test_tags_stream_mode_messages() -> None:
{ {
"langgraph_step": 1, "langgraph_step": 1,
"langgraph_node": "call_model", "langgraph_node": "call_model",
"langgraph_triggers": ("start:call_model",), "langgraph_triggers": (
"branch:to:call_model",
"start:call_model",
),
"langgraph_path": ("__pregel_pull", "call_model"), "langgraph_path": ("__pregel_pull", "call_model"),
"langgraph_checkpoint_ns": AnyStr("call_model:"), "langgraph_checkpoint_ns": AnyStr("call_model:"),
"checkpoint_ns": AnyStr("call_model:"), "checkpoint_ns": AnyStr("call_model:"),
@@ -7838,3 +7852,329 @@ async def test_handles_multiple_interrupts_from_tasks() -> None:
assert len(result) == 2 assert len(result) == 2
assert result[0] == "Added James!" assert result[0] == "Added James!"
assert result[1] == "Added Will!" assert result[1] == "Added Will!"
async def test_pregel_loop_refcount():
gc.collect()
try:
gc.disable()
class State(TypedDict):
messages: Annotated[list, add_messages]
graph_builder = StateGraph(State)
async def chatbot(state: State):
return {"messages": [("ai", "HIYA")]}
graph_builder.add_node("chatbot", chatbot)
graph_builder.set_entry_point("chatbot")
graph_builder.set_finish_point("chatbot")
graph = graph_builder.compile()
for _ in range(5):
await graph.ainvoke({"messages": [{"role": "user", "content": "hi"}]})
assert (
len(
[
obj
for obj in gc.get_objects()
if isinstance(obj, AsyncPregelLoop)
]
)
== 0
)
assert (
len([obj for obj in gc.get_objects() if isinstance(obj, PregelRunner)])
== 0
)
finally:
gc.enable()
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
async def test_bulk_state_updates(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer:
class State(TypedDict):
foo: str
baz: str
def node_a(state: State) -> State:
return {"foo": "bar"}
def node_b(state: State) -> State:
return {"baz": "qux"}
graph = (
StateGraph(State)
.add_node("node_a", node_a)
.add_node("node_b", node_b)
.add_edge(START, "node_a")
.add_edge("node_a", "node_b")
.compile(checkpointer=checkpointer)
)
config = {"configurable": {"thread_id": "1"}}
# First update with node_a
await graph.abulk_update_state(
config,
[
[
StateUpdate({"foo": "bar"}, "node_a"),
]
],
)
# Then bulk update with both nodes
await graph.abulk_update_state(
config,
[
[
StateUpdate({"foo": "updated"}, "node_a"),
StateUpdate({"baz": "new"}, "node_b"),
]
],
)
state = await graph.aget_state(config)
assert state.values == {"foo": "updated", "baz": "new"}
# Check if there are only two checkpoints
checkpoints = [
c async for c in checkpointer.alist({"configurable": {"thread_id": "1"}})
]
assert len(checkpoints) == 2
assert checkpoints[0].metadata["writes"] == {
"node_a": {"foo": "updated"},
"node_b": {"baz": "new"},
}
assert checkpoints[1].metadata["writes"] == {"node_a": {"foo": "bar"}}
# perform multiple steps at the same time
config = {"configurable": {"thread_id": "2"}}
await graph.abulk_update_state(
config,
[
[
StateUpdate({"foo": "bar"}, "node_a"),
],
[
StateUpdate({"foo": "updated"}, "node_a"),
StateUpdate({"baz": "new"}, "node_b"),
],
],
)
state = await graph.aget_state(config)
assert state.values == {"foo": "updated", "baz": "new"}
checkpoints = [
c async for c in checkpointer.alist({"configurable": {"thread_id": "1"}})
]
assert len(checkpoints) == 2
assert checkpoints[0].metadata["writes"] == {
"node_a": {"foo": "updated"},
"node_b": {"baz": "new"},
}
assert checkpoints[1].metadata["writes"] == {"node_a": {"foo": "bar"}}
# Should raise error if updating without as_node
with pytest.raises(InvalidUpdateError):
await graph.abulk_update_state(
config,
[
[
StateUpdate(values={"foo": "error"}, as_node=None),
StateUpdate(values={"bar": "error"}, as_node=None),
]
],
)
# Should raise if no updates are provided
with pytest.raises(ValueError, match="No supersteps provided"):
await graph.abulk_update_state(config, [])
# Should raise if no updates are provided
with pytest.raises(ValueError, match="No updates provided"):
await graph.abulk_update_state(config, [[], []])
# Should raise if __end__ or __copy__ update is applied in bulk
with pytest.raises(InvalidUpdateError):
await graph.abulk_update_state(
config,
[
[
StateUpdate(values=None, as_node="__end__"),
StateUpdate(values=None, as_node="__copy__"),
],
],
)
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
async def test_update_as_input(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer:
class State(TypedDict):
foo: str
def agent(state: State) -> State:
return {"foo": "agent"}
def tool(state: State) -> State:
return {"foo": "tool"}
graph = (
StateGraph(State)
.add_node("agent", agent)
.add_node("tool", tool)
.add_edge(START, "agent")
.add_edge("agent", "tool")
.compile(checkpointer=checkpointer)
)
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
) == {"foo": "tool"}
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
) == {"foo": "tool"}
def map_snapshot(i: StateSnapshot) -> dict:
return {
"values": i.values,
"next": i.next,
"step": i.metadata.get("step"),
}
history = [
map_snapshot(s)
async for s in graph.aget_state_history(
{"configurable": {"thread_id": "1"}}
)
]
await graph.abulk_update_state(
{"configurable": {"thread_id": "2"}},
[
# First turn
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent"}, "agent")],
[StateUpdate({"foo": "tool"}, "tool")],
# Second turn
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent"}, "agent")],
[StateUpdate({"foo": "tool"}, "tool")],
],
)
state = await graph.aget_state({"configurable": {"thread_id": "2"}})
assert state.values == {"foo": "tool"}
new_history = [
map_snapshot(s)
async for s in graph.aget_state_history(
{"configurable": {"thread_id": "2"}}
)
]
assert new_history == history
@pytest.mark.parametrize("checkpointer_name", REGULAR_CHECKPOINTERS_ASYNC)
async def test_batch_update_as_input(checkpointer_name: str) -> None:
async with awith_checkpointer(checkpointer_name) as checkpointer:
class State(TypedDict):
foo: str
tasks: Annotated[list[int], operator.add]
def agent(state: State) -> State:
return {"foo": "agent"}
def map(state: State) -> Command["task"]:
return Command(
goto=[
Send("task", {"index": 0}),
Send("task", {"index": 1}),
Send("task", {"index": 2}),
],
update={"foo": "map"},
)
def task(state: dict) -> State:
return {"tasks": [state["index"]]}
graph = (
StateGraph(State)
.add_node("agent", agent)
.add_node("map", map)
.add_node("task", task)
.add_edge(START, "agent")
.add_edge("agent", "map")
.compile(checkpointer=checkpointer)
)
assert await graph.ainvoke(
{"foo": "input"}, {"configurable": {"thread_id": "1"}}
) == {"foo": "map", "tasks": [0, 1, 2]}
def map_snapshot(i: StateSnapshot) -> dict:
return {
"values": i.values,
"next": i.next,
"step": i.metadata.get("step"),
"tasks": [t.name for t in i.tasks],
}
history = [
map_snapshot(s)
async for s in graph.aget_state_history(
{"configurable": {"thread_id": "1"}}
)
]
await graph.abulk_update_state(
{"configurable": {"thread_id": "2"}},
[
[StateUpdate({"foo": "input"}, "__input__")],
[StateUpdate({"foo": "input"}, "__start__")],
[StateUpdate({"foo": "agent", "tasks": []}, "agent")],
[
StateUpdate(
Command(
goto=[
Send("task", {"index": 0}),
Send("task", {"index": 1}),
Send("task", {"index": 2}),
],
update={"foo": "map"},
),
"map",
)
],
[
StateUpdate({"tasks": [0]}, "task"),
StateUpdate({"tasks": [1]}, "task"),
StateUpdate({"tasks": [2]}, "task"),
],
],
)
state = await graph.aget_state({"configurable": {"thread_id": "2"}})
assert state.values == {"foo": "map", "tasks": [0, 1, 2]}
new_history = [
map_snapshot(s)
async for s in graph.aget_state_history(
{"configurable": {"thread_id": "2"}}
)
]
assert new_history == history
@@ -719,9 +719,7 @@ def create_react_agent(
def generate_structured_response( def generate_structured_response(
state: StateSchema, config: RunnableConfig state: StateSchema, config: RunnableConfig
) -> StateSchema: ) -> StateSchema:
# NOTE: we exclude the last message because there is enough information messages = _get_state_value(state, "messages")
# for the LLM to generate the structured response
messages = _get_state_value(state, "messages")[:-1]
structured_response_schema = response_format structured_response_schema = response_format
if isinstance(response_format, tuple): if isinstance(response_format, tuple):
system_prompt, structured_response_schema = response_format system_prompt, structured_response_schema = response_format
@@ -736,9 +734,7 @@ def create_react_agent(
async def agenerate_structured_response( async def agenerate_structured_response(
state: StateSchema, config: RunnableConfig state: StateSchema, config: RunnableConfig
) -> StateSchema: ) -> StateSchema:
# NOTE: we exclude the last message because there is enough information messages = _get_state_value(state, "messages")
# for the LLM to generate the structured response
messages = _get_state_value(state, "messages")[:-1]
structured_response_schema = response_format structured_response_schema = response_format
if isinstance(response_format, tuple): if isinstance(response_format, tuple):
system_prompt, structured_response_schema = response_format system_prompt, structured_response_schema = response_format
+6 -5
View File
@@ -435,7 +435,7 @@ typing-extensions = ">=4.7"
[[package]] [[package]]
name = "langgraph" name = "langgraph"
version = "0.3.0" version = "0.3.18"
description = "Building stateful, multi-actor applications with LLMs" description = "Building stateful, multi-actor applications with LLMs"
optional = false optional = false
python-versions = ">=3.9.0,<4.0" python-versions = ">=3.9.0,<4.0"
@@ -446,6 +446,7 @@ develop = true
[package.dependencies] [package.dependencies]
langchain-core = ">=0.1,<0.4" langchain-core = ">=0.1,<0.4"
langgraph-checkpoint = "^2.0.10" langgraph-checkpoint = "^2.0.10"
langgraph-prebuilt = ">=0.1.1,<0.2"
langgraph-sdk = "^0.1.42" langgraph-sdk = "^0.1.42"
[package.source] [package.source]
@@ -454,7 +455,7 @@ url = "../langgraph"
[[package]] [[package]]
name = "langgraph-checkpoint" name = "langgraph-checkpoint"
version = "2.0.16" version = "2.0.21"
description = "Library with base interfaces for LangGraph checkpoint savers." description = "Library with base interfaces for LangGraph checkpoint savers."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -472,7 +473,7 @@ url = "../checkpoint"
[[package]] [[package]]
name = "langgraph-checkpoint-postgres" name = "langgraph-checkpoint-postgres"
version = "2.0.15" version = "2.0.19"
description = "Library with a Postgres implementation of LangGraph checkpoint saver." description = "Library with a Postgres implementation of LangGraph checkpoint saver."
optional = false optional = false
python-versions = "^3.9.0,<4.0" python-versions = "^3.9.0,<4.0"
@@ -481,7 +482,7 @@ files = []
develop = true develop = true
[package.dependencies] [package.dependencies]
langgraph-checkpoint = "^2.0.15" langgraph-checkpoint = "^2.0.21"
orjson = ">=3.10.1" orjson = ">=3.10.1"
psycopg = "^3.2.0" psycopg = "^3.2.0"
psycopg-pool = "^3.2.0" psycopg-pool = "^3.2.0"
@@ -492,7 +493,7 @@ url = "../checkpoint-postgres"
[[package]] [[package]]
name = "langgraph-checkpoint-sqlite" name = "langgraph-checkpoint-sqlite"
version = "2.0.5" version = "2.0.6"
description = "Library with a SQLite implementation of LangGraph checkpoint saver." description = "Library with a SQLite implementation of LangGraph checkpoint saver."
optional = false optional = false
python-versions = "^3.9.0" python-versions = "^3.9.0"
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph-prebuilt" name = "langgraph-prebuilt"
version = "0.1.3" version = "0.1.4"
description = "Library with high-level APIs for creating and executing LangGraph agents and tools." description = "Library with high-level APIs for creating and executing LangGraph agents and tools."
authors = [] authors = []
license = "MIT" license = "MIT"
+1 -1
View File
@@ -1,6 +1,6 @@
{ {
"name": "@langchain/langgraph-sdk", "name": "@langchain/langgraph-sdk",
"version": "0.0.57", "version": "0.0.60",
"description": "Client library for interacting with the LangGraph API", "description": "Client library for interacting with the LangGraph API",
"type": "module", "type": "module",
"packageManager": "yarn@1.22.19", "packageManager": "yarn@1.22.19",
+34 -1
View File
@@ -30,6 +30,7 @@ import type {
StreamEvent, StreamEvent,
CronsCreatePayload, CronsCreatePayload,
OnConflictBehavior, OnConflictBehavior,
Command,
} from "./types.js"; } from "./types.js";
import { mergeSignals } from "./utils/signals.js"; import { mergeSignals } from "./utils/signals.js";
import { getEnvironmentVariable } from "./utils/env.js"; import { getEnvironmentVariable } from "./utils/env.js";
@@ -481,15 +482,47 @@ export class ThreadsClient<
* Metadata for the thread. * Metadata for the thread.
*/ */
metadata?: Metadata; metadata?: Metadata;
/**
* ID of the thread to create.
*
* If not provided, a random UUID will be generated.
*/
threadId?: string; threadId?: string;
/**
* How to handle duplicate creation.
*
* @default "raise"
*/
ifExists?: OnConflictBehavior; ifExists?: OnConflictBehavior;
/**
* Graph ID to associate with the thread.
*/
graphId?: string;
/**
* Apply a list of supersteps when creating a thread, each containing a sequence of updates.
*
* Used for copying a thread between deployments.
*/
supersteps?: Array<{
updates: Array<{ values: unknown; command?: Command; asNode: string }>;
}>;
}): Promise<Thread<TStateType>> { }): Promise<Thread<TStateType>> {
return this.fetch<Thread<TStateType>>(`/threads`, { return this.fetch<Thread<TStateType>>(`/threads`, {
method: "POST", method: "POST",
json: { json: {
metadata: payload?.metadata, metadata: {
...payload?.metadata,
graph_id: payload?.graphId,
},
thread_id: payload?.threadId, thread_id: payload?.threadId,
if_exists: payload?.ifExists, if_exists: payload?.ifExists,
supersteps: payload?.supersteps?.map((s) => ({
updates: s.updates.map((u) => ({
values: u.values,
command: u.command,
as_node: u.asNode,
})),
})),
}, },
}); });
} }
+51 -5
View File
@@ -839,6 +839,8 @@ class ThreadsClient:
metadata: Json = None, metadata: Json = None,
thread_id: Optional[str] = None, thread_id: Optional[str] = None,
if_exists: Optional[OnConflictBehavior] = None, if_exists: Optional[OnConflictBehavior] = None,
supersteps: Optional[Sequence[dict[str, Sequence[dict[str, Any]]]]] = None,
graph_id: Optional[str] = None,
) -> Thread: ) -> Thread:
"""Create a new thread. """Create a new thread.
@@ -848,6 +850,9 @@ class ThreadsClient:
If None, ID will be a randomly generated UUID. If None, ID will be a randomly generated UUID.
if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood. if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood.
Must be either 'raise' (raise error if duplicate), or 'do_nothing' (return existing thread). Must be either 'raise' (raise error if duplicate), or 'do_nothing' (return existing thread).
supersteps: Apply a list of supersteps when creating a thread, each containing a sequence of updates.
Each update has `values` or `command` and `as_node`. Used for copying a thread between deployments.
graph_id: Optional graph ID to associate with the thread.
Returns: Returns:
Thread: The created thread. Thread: The created thread.
@@ -863,10 +868,28 @@ class ThreadsClient:
payload: Dict[str, Any] = {} payload: Dict[str, Any] = {}
if thread_id: if thread_id:
payload["thread_id"] = thread_id payload["thread_id"] = thread_id
if metadata: if metadata or graph_id:
payload["metadata"] = metadata payload["metadata"] = {
**(metadata or {}),
**({"graph_id": graph_id} if graph_id else {}),
}
if if_exists: if if_exists:
payload["if_exists"] = if_exists payload["if_exists"] = if_exists
if supersteps:
payload["supersteps"] = [
{
"updates": [
{
"values": u["values"],
"command": u.get("command"),
"as_node": u["as_node"],
}
for u in s["updates"]
]
}
for s in supersteps
]
return await self.http.post("/threads", json=payload) return await self.http.post("/threads", json=payload)
async def update(self, thread_id: str, *, metadata: dict[str, Any]) -> Thread: async def update(self, thread_id: str, *, metadata: dict[str, Any]) -> Thread:
@@ -3036,6 +3059,8 @@ class SyncThreadsClient:
metadata: Json = None, metadata: Json = None,
thread_id: Optional[str] = None, thread_id: Optional[str] = None,
if_exists: Optional[OnConflictBehavior] = None, if_exists: Optional[OnConflictBehavior] = None,
supersteps: Optional[Sequence[dict[str, Sequence[dict[str, Any]]]]] = None,
graph_id: Optional[str] = None,
) -> Thread: ) -> Thread:
"""Create a new thread. """Create a new thread.
@@ -3045,6 +3070,9 @@ class SyncThreadsClient:
If None, ID will be a randomly generated UUID. If None, ID will be a randomly generated UUID.
if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood. if_exists: How to handle duplicate creation. Defaults to 'raise' under the hood.
Must be either 'raise' (raise error if duplicate), or 'do_nothing' (return existing thread). Must be either 'raise' (raise error if duplicate), or 'do_nothing' (return existing thread).
supersteps: Apply a list of supersteps when creating a thread, each containing a sequence of updates.
Each update has `values` or `command` and `as_node`. Used for copying a thread between deployments.
graph_id: Optional graph ID to associate with the thread.
Returns: Returns:
Thread: The created thread. Thread: The created thread.
@@ -3060,10 +3088,28 @@ class SyncThreadsClient:
payload: Dict[str, Any] = {} payload: Dict[str, Any] = {}
if thread_id: if thread_id:
payload["thread_id"] = thread_id payload["thread_id"] = thread_id
if metadata: if metadata or graph_id:
payload["metadata"] = metadata payload["metadata"] = {
**(metadata or {}),
**({"graph_id": graph_id} if graph_id else {}),
}
if if_exists: if if_exists:
payload["if_exists"] = if_exists payload["if_exists"] = if_exists
if supersteps:
payload["supersteps"] = [
{
"updates": [
{
"values": u["values"],
"command": u.get("command"),
"as_node": u["as_node"],
}
for u in s["updates"]
]
}
for s in supersteps
]
return self.http.post("/threads", json=payload) return self.http.post("/threads", json=payload)
def update(self, thread_id: str, *, metadata: dict[str, Any]) -> Thread: def update(self, thread_id: str, *, metadata: dict[str, Any]) -> Thread:
@@ -3307,7 +3353,7 @@ class SyncThreadsClient:
Example Usage: Example Usage:
response = client.threads.update_state( response = await client.threads.update_state(
thread_id="my_thread_id", thread_id="my_thread_id",
values={"messages":[{"role": "user", "content": "hello!"}]}, values={"messages":[{"role": "user", "content": "hello!"}]},
as_node="my_node", as_node="my_node",
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "langgraph-sdk" name = "langgraph-sdk"
version = "0.1.57" version = "0.1.58"
description = "SDK for interacting with LangGraph API" description = "SDK for interacting with LangGraph API"
authors = [] authors = []
license = "MIT" license = "MIT"