mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-21 07:02:25 +02:00
Compare commits
195
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8e922b859c | ||
|
|
9a4c30135f | ||
|
|
bc6651f34c | ||
|
|
23bb5369b9 | ||
|
|
85b81371f7 | ||
|
|
54ab833b74 | ||
|
|
4ea936eaf4 | ||
|
|
9d81ec9ffd | ||
|
|
46b652a74c | ||
|
|
7013ca9a3f | ||
|
|
0c04aec664 | ||
|
|
1650c8508e | ||
|
|
e176b98fe7 | ||
|
|
eb1e1aa010 | ||
|
|
77c833e1e5 | ||
|
|
0ac29434a7 | ||
|
|
43f5a17416 | ||
|
|
7d0857f263 | ||
|
|
b82d70a66a | ||
|
|
5fb037171d | ||
|
|
d3bb2b9aa0 | ||
|
|
ea765b4134 | ||
|
|
66ff83dca9 | ||
|
|
254e398345 | ||
|
|
c7567ea219 | ||
|
|
8c0306c3f4 | ||
|
|
992b05a196 | ||
|
|
893a9646d3 | ||
|
|
eaa37a2ce9 | ||
|
|
daee8d88bb | ||
|
|
eaa18cc2dd | ||
|
|
b2d9a36308 | ||
|
|
03fc695d60 | ||
|
|
9994b09304 | ||
|
|
53f8558914 | ||
|
|
4bfcd84cee | ||
|
|
cd1d7be05f | ||
|
|
939a426a2e | ||
|
|
94c815f226 | ||
|
|
ef345aac5f | ||
|
|
e306258525 | ||
|
|
05cd317486 | ||
|
|
6dfed31a5e | ||
|
|
ee8374c4c0 | ||
|
|
2066894b5f | ||
|
|
6a2d20fd5b | ||
|
|
5b8b9f1067 | ||
|
|
c9cb8165d4 | ||
|
|
1f0348a5ca | ||
|
|
442ef0788e | ||
|
|
73f9ef0ef8 | ||
|
|
1f7a380548 | ||
|
|
8959f2aec5 | ||
|
|
0e7869eba4 | ||
|
|
d3f8478054 | ||
|
|
0cb1893475 | ||
|
|
e779c8e0b1 | ||
|
|
18b82cb8e2 | ||
|
|
972ab1a935 | ||
|
|
9cc2f37cca | ||
|
|
a2d7631f47 | ||
|
|
1e767c0653 | ||
|
|
d4c569cb7c | ||
|
|
a146df7f6a | ||
|
|
c52cc03e4b | ||
|
|
f206cfad8f | ||
|
|
d4c8b219c4 | ||
|
|
00855999d2 | ||
|
|
24bd0e1c1f | ||
|
|
3b59055192 | ||
|
|
67a16bec53 | ||
|
|
f97802eed5 | ||
|
|
bccb796ccc | ||
|
|
4068e9d135 | ||
|
|
a951334f7f | ||
|
|
75a727877f | ||
|
|
1cece3228c | ||
|
|
7ba48d75c9 | ||
|
|
7fb0628957 | ||
|
|
6edf29f043 | ||
|
|
c8a605cbc8 | ||
|
|
fbec207446 | ||
|
|
2223c82606 | ||
|
|
06f2eef74c | ||
|
|
62aa66cd4b | ||
|
|
8ffe9634b7 | ||
|
|
4b1d6d2aeb | ||
|
|
199ab46429 | ||
|
|
c758954519 | ||
|
|
5bfb3bb882 | ||
|
|
68a5c3f4c7 | ||
|
|
b7fb8e6afb | ||
|
|
c34c798763 | ||
|
|
792cd805a7 | ||
|
|
764929afd9 | ||
|
|
1e751a2256 | ||
|
|
e6726802f7 | ||
|
|
a541376d10 | ||
|
|
f9780330a6 | ||
|
|
24f7d7c439 | ||
|
|
9533d35a84 | ||
|
|
48de4a7234 | ||
|
|
67ac35c5bc | ||
|
|
12c7ddf3a3 | ||
|
|
d6caa3b00a | ||
|
|
3d48526c16 | ||
|
|
2f3bd69bf5 | ||
|
|
f10a0c6f32 | ||
|
|
2eba27b01f | ||
|
|
c0245a6ee9 | ||
|
|
9e82d23252 | ||
|
|
460c522902 | ||
|
|
298a19b573 | ||
|
|
d2ec46b927 | ||
|
|
22e4bf74fc | ||
|
|
96a0536ec4 | ||
|
|
8d33938173 | ||
|
|
1c73b1e45a | ||
|
|
465d5d648b | ||
|
|
b751e8bcee | ||
|
|
8ec3982056 | ||
|
|
f17f264a7a | ||
|
|
086443472f | ||
|
|
936e69404e | ||
|
|
b1a25abc73 | ||
|
|
0e70b8d94f | ||
|
|
e1aa1a4510 | ||
|
|
ae7dbd1fa5 | ||
|
|
3ec95153ce | ||
|
|
4836f8b18b | ||
|
|
e7fbdeeb13 | ||
|
|
ee650ab85f | ||
|
|
7a959f62cc | ||
|
|
82905297fd | ||
|
|
9b5549f759 | ||
|
|
fa96c0ac76 | ||
|
|
98b8ff904c | ||
|
|
951131c8ec | ||
|
|
8bcdba822e | ||
|
|
60fc49b448 | ||
|
|
1d21b4ba08 | ||
|
|
55ec0d3d2a | ||
|
|
c7dd7be030 | ||
|
|
47d38a3022 | ||
|
|
8e829f38af | ||
|
|
3f241d00a3 | ||
|
|
f0abf582dd | ||
|
|
2458f2d2e0 | ||
|
|
477a43dae0 | ||
|
|
fc8e6ec64f | ||
|
|
d6a457ef1d | ||
|
|
ce1077da40 | ||
|
|
969958695a | ||
|
|
dd16ae4ba5 | ||
|
|
e24e141253 | ||
|
|
eae1faa656 | ||
|
|
1976d6584c | ||
|
|
54e18445fc | ||
|
|
69dc29aaf9 | ||
|
|
aa5ff74845 | ||
|
|
6049aaa842 | ||
|
|
576aa1ca02 | ||
|
|
e28e97d5e0 | ||
|
|
59e7c63c93 | ||
|
|
be7dee1c3b | ||
|
|
0aafa04bac | ||
|
|
2e1adaa867 | ||
|
|
3f8b165592 | ||
|
|
9ed0fa196c | ||
|
|
0b9adc28c3 | ||
|
|
3b0255d1ef | ||
|
|
80c3ccba7b | ||
|
|
d4255a0645 | ||
|
|
80d61a2600 | ||
|
|
424f24720a | ||
|
|
2a71180c1d | ||
|
|
697f878e36 | ||
|
|
987b9da4ab | ||
|
|
4bff1df4b0 | ||
|
|
5104e31e35 | ||
|
|
7ae4739630 | ||
|
|
5db1949ae3 | ||
|
|
ddb29df667 | ||
|
|
fa467573d7 | ||
|
|
1a728a93c6 | ||
|
|
def69c59d2 | ||
|
|
aaa0cd6b51 | ||
|
|
a9c831c11b | ||
|
|
bad4d17c34 | ||
|
|
55219b23d8 | ||
|
|
8edbd39ad3 | ||
|
|
4b0fd834d8 | ||
|
|
0fd2748530 | ||
|
|
bc0a3419ed | ||
|
|
5cd47bac49 |
@@ -57,13 +57,13 @@ jobs:
|
|||||||
env:
|
env:
|
||||||
# these won't actually be used because of the VCR cassettes
|
# these won't actually be used because of the VCR cassettes
|
||||||
# but need to set them to avoid triggering getpass()
|
# but need to set them to avoid triggering getpass()
|
||||||
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
|
OPENAI_API_KEY: "very-secret-key"
|
||||||
ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }}
|
ANTHROPIC_API_KEY: "very-secret-key"
|
||||||
TAVILY_API_KEY: ${{ secrets.TAVILY_API_KEY }}
|
TAVILY_API_KEY: "very-secret-key"
|
||||||
LANGSMITH_API_KEY: ${{ secrets.LANGSMITH_API_KEY }}
|
LANGSMITH_API_KEY: "very-secret-key"
|
||||||
NOMIC_API_KEY: ${{ secrets.NOMIC_API_KEY }}
|
NOMIC_API_KEY: "very-secret-key"
|
||||||
COHERE_API_KEY: ${{ secrets.COHERE_API_KEY }}
|
COHERE_API_KEY: "very-secret-key"
|
||||||
FIREWORKS_API_KEY: ${{ secrets.FIREWORKS_API_KEY }}
|
FIREWORKS_API_KEY: "very-secret-key"
|
||||||
run: |
|
run: |
|
||||||
if [ "${{ github.event_name }}" = "workflow_dispatch" ] || [ "${{ github.event_name }}" = "schedule" ]; then
|
if [ "${{ github.event_name }}" = "workflow_dispatch" ] || [ "${{ github.event_name }}" = "schedule" ]; then
|
||||||
echo "Running all notebooks"
|
echo "Running all notebooks"
|
||||||
|
|||||||
@@ -1,29 +0,0 @@
|
|||||||
name: Check File Size
|
|
||||||
|
|
||||||
on:
|
|
||||||
push:
|
|
||||||
branches:
|
|
||||||
- main
|
|
||||||
pull_request:
|
|
||||||
branches:
|
|
||||||
- main
|
|
||||||
workflow_dispatch:
|
|
||||||
|
|
||||||
jobs:
|
|
||||||
file-size-check:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- uses: actions/checkout@v4
|
|
||||||
- name: Get changed files
|
|
||||||
id: changed-files
|
|
||||||
uses: tj-actions/changed-files@v44
|
|
||||||
- name: Filter by size
|
|
||||||
# TODO: roll back the web voyager hack
|
|
||||||
run: |
|
|
||||||
large_added_files=$(find ${{ steps.changed-files.outputs.added_files }} -maxdepth 0 -size +1M | grep -v "web_voyager" || true)
|
|
||||||
if [ -n "$large_added_files" ]; then
|
|
||||||
echo "Large files added: $large_added_files"
|
|
||||||
echo "# Large files added:" >> $GITHUB_STEP_SUMMARY
|
|
||||||
echo "$large_added_files" >> $GITHUB_STEP_SUMMARY
|
|
||||||
exit 1
|
|
||||||
fi
|
|
||||||
@@ -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>
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -30,6 +30,9 @@ packages:
|
|||||||
- name: "langgraph-bigtool"
|
- name: "langgraph-bigtool"
|
||||||
repo: "langchain-ai/langgraph-bigtool"
|
repo: "langchain-ai/langgraph-bigtool"
|
||||||
description: "Build LangGraph agents with large numbers of tools."
|
description: "Build LangGraph agents with large numbers of tools."
|
||||||
|
- name: "ai-data-science-team"
|
||||||
|
repo: "business-science/ai-data-science-team"
|
||||||
|
description: "An AI-powered data science team of agents to help you perform common data science tasks 10X faster."
|
||||||
- name: "langgraph-reflection"
|
- name: "langgraph-reflection"
|
||||||
repo: "langchain-ai/langgraph-reflection"
|
repo: "langchain-ai/langgraph-reflection"
|
||||||
description: "LangGraph agent that runs a reflection step."
|
description: "LangGraph agent that runs a reflection step."
|
||||||
|
|||||||
@@ -64,7 +64,7 @@ license = "MIT"
|
|||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
|
|
||||||
[tool.poetry.dependencies]
|
[tool.poetry.dependencies]
|
||||||
python = ">=3.9.0,<3.13"
|
python = ">=3.9"
|
||||||
langgraph = "^0.2.0"
|
langgraph = "^0.2.0"
|
||||||
langchain-fireworks = "^0.1.3"
|
langchain-fireworks = "^0.1.3"
|
||||||
|
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ Now, let's invoke our graph by interrupting before `ask_human` node:
|
|||||||
"messages": [
|
"messages": [
|
||||||
{
|
{
|
||||||
"role": "user",
|
"role": "user",
|
||||||
"content": "Use the search tool to ask the user where they are, then look up the weather there",
|
"content": "Ask the user where they are, then look up the weather there",
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
@@ -85,8 +85,7 @@ Now, let's invoke our graph by interrupting before `ask_human` node:
|
|||||||
messages: [
|
messages: [
|
||||||
{
|
{
|
||||||
role: "human",
|
role: "human",
|
||||||
content: "Use the search tool to ask the user where they are, then look up the weather there"
|
content: "Ask the user where they are, then look up the weather there" }
|
||||||
}
|
|
||||||
]
|
]
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -115,7 +114,7 @@ Now, let's invoke our graph by interrupting before `ask_human` node:
|
|||||||
--header 'Content-Type: application/json' \
|
--header 'Content-Type: application/json' \
|
||||||
--data "{
|
--data "{
|
||||||
\"assistant_id\": \"agent\",
|
\"assistant_id\": \"agent\",
|
||||||
\"input\": {\"messages\": [{\"role\": \"human\", \"content\": \"Use the search tool to ask the user where they are, then look up the weather there\"}]},
|
\"input\": {\"messages\": [{\"role\": \"human\", \"content\": \"Ask the user where they are, then look up the weather there\"}]},
|
||||||
\"interrupt_before\": [\"ask_human\"],
|
\"interrupt_before\": [\"ask_human\"],
|
||||||
\"stream_mode\": [
|
\"stream_mode\": [
|
||||||
\"updates\"
|
\"updates\"
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 39 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 93 KiB |
@@ -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**:
|
||||||
|
|
||||||
|
{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.
|
||||||
|
|
||||||
|
{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
|
|||||||
|
|
||||||
{width=1200}
|
{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).
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ As a result, there are many different types of [agent architectures](https://blo
|
|||||||
|
|
||||||
## Router
|
## Router
|
||||||
|
|
||||||
A router allows an LLM to select a single step from a specified set of options. This is an agent architecture that exhibits a relatively limited level of control because the LLM usually focuses on making a single decision and produces a specific output from limited set of pre-defined options. Routers typically employ a few different concepts to achieve this.
|
A router allows an LLM to select a single step from a specified set of options. This is an agent architecture that exhibits a relatively limited level of control because the LLM usually focuses on making a single decision and produces a specific output from a limited set of pre-defined options. Routers typically employ a few different concepts to achieve this.
|
||||||
|
|
||||||
### Structured Output
|
### Structured Output
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,8 @@
|
|||||||
|
---
|
||||||
|
search:
|
||||||
|
boost: 2
|
||||||
|
---
|
||||||
|
|
||||||
# LangGraph Platform
|
# LangGraph Platform
|
||||||
|
|
||||||
## Overview
|
## Overview
|
||||||
|
|||||||
@@ -89,7 +89,7 @@ def transfer_to_bob(state):
|
|||||||
)
|
)
|
||||||
```
|
```
|
||||||
|
|
||||||
This is a special case of updating the graph state from tools where in addition the state update, the control flow is included as well.
|
This is a special case of updating the graph state from tools where, in addition to the state update, the control flow is included as well.
|
||||||
|
|
||||||
!!! important
|
!!! important
|
||||||
|
|
||||||
@@ -235,7 +235,7 @@ supervisor = create_react_agent(model, tools)
|
|||||||
|
|
||||||
### Hierarchical
|
### Hierarchical
|
||||||
|
|
||||||
As you add more agents to your system, it might become too hard for the supervisor to manage all of them. The supervisor might start making poor decisions about which agent to call next, the context might become too complex for a single supervisor to keep track of. In other words, you end up with the same problems that motivated the multi-agent architecture in the first place.
|
As you add more agents to your system, it might become too hard for the supervisor to manage all of them. The supervisor might start making poor decisions about which agent to call next, or the context might become too complex for a single supervisor to keep track of. In other words, you end up with the same problems that motivated the multi-agent architecture in the first place.
|
||||||
|
|
||||||
To address this, you can design your system _hierarchically_. For example, you can create separate, specialized teams of agents managed by individual supervisors, and a top-level supervisor to manage the teams.
|
To address this, you can design your system _hierarchically_. For example, you can create separate, specialized teams of agents managed by individual supervisors, and a top-level supervisor to manage the teams.
|
||||||
|
|
||||||
@@ -339,9 +339,9 @@ builder.add_edge("agent_1", "agent_2")
|
|||||||
|
|
||||||
## Communication between agents
|
## Communication between agents
|
||||||
|
|
||||||
The most important thing when building multi-agent systems is figuring out how the agents communicate. There are few different considerations:
|
The most important thing when building multi-agent systems is figuring out how the agents communicate. There are a few different considerations:
|
||||||
|
|
||||||
- Do agents communicate via [**via graph state or via tool calls**](#graph-state-vs-tool-calls)?
|
- Do agents communicate [**via graph state or via tool calls**](#graph-state-vs-tool-calls)?
|
||||||
- What if two agents have [**different state schemas**](#different-state-schemas)?
|
- What if two agents have [**different state schemas**](#different-state-schemas)?
|
||||||
- How to communicate over a [**shared message list**](#shared-message-list)?
|
- How to communicate over a [**shared message list**](#shared-message-list)?
|
||||||
|
|
||||||
|
|||||||
@@ -232,7 +232,7 @@ from langgraph.store.memory import InMemoryStore
|
|||||||
in_memory_store = InMemoryStore()
|
in_memory_store = InMemoryStore()
|
||||||
```
|
```
|
||||||
|
|
||||||
Memories are namespaced by a `tuple`, which in this specific example will be `(<user_id>, "memories")`. The namespace can be any length and represent anything, does not have be user specific.
|
Memories are namespaced by a `tuple`, which in this specific example will be `(<user_id>, "memories")`. The namespace can be any length and represent anything, does not have to be user specific.
|
||||||
|
|
||||||
```python
|
```python
|
||||||
user_id = "1"
|
user_id = "1"
|
||||||
@@ -387,6 +387,9 @@ We can access the memories and use them in our model call.
|
|||||||
def call_model(state: MessagesState, config: RunnableConfig, *, store: BaseStore):
|
def call_model(state: MessagesState, config: RunnableConfig, *, store: BaseStore):
|
||||||
# Get the user id from the config
|
# Get the user id from the config
|
||||||
user_id = config["configurable"]["user_id"]
|
user_id = config["configurable"]["user_id"]
|
||||||
|
|
||||||
|
# Namespace the memory
|
||||||
|
namespace = (user_id, "memories")
|
||||||
|
|
||||||
# Search based on the most recent message
|
# Search based on the most recent message
|
||||||
memories = store.search(
|
memories = store.search(
|
||||||
|
|||||||
@@ -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."
|
||||||
|
|||||||
@@ -33,7 +33,7 @@
|
|||||||
" )\n",
|
" )\n",
|
||||||
"```\n",
|
"```\n",
|
||||||
"\n",
|
"\n",
|
||||||
"If you are using [subgraphs](#subgraphs), you might want to navigate from a node a subgraph to a different subgraph (i.e. a different node in the parent graph). To do so, you can specify `graph=Command.PARENT` in `Command`:\n",
|
"If you are using [subgraphs](#subgraphs), you might want to navigate from a node within a subgraph to a different subgraph (i.e. a different node in the parent graph). To do so, you can specify `graph=Command.PARENT` in `Command`:\n",
|
||||||
"\n",
|
"\n",
|
||||||
"```python\n",
|
"```python\n",
|
||||||
"def my_node(state: State) -> Command[Literal[\"my_other_node\"]]:\n",
|
"def my_node(state: State) -> Command[Literal[\"my_other_node\"]]:\n",
|
||||||
|
|||||||
@@ -397,7 +397,8 @@
|
|||||||
"# We define a fake node to ask the human\n",
|
"# We define a fake node to ask the human\n",
|
||||||
"def ask_human(state):\n",
|
"def ask_human(state):\n",
|
||||||
" tool_call_id = state[\"messages\"][-1].tool_calls[0][\"id\"]\n",
|
" tool_call_id = state[\"messages\"][-1].tool_calls[0][\"id\"]\n",
|
||||||
" location = interrupt(\"Please provide your location:\")\n",
|
" ask = AskHuman.model_validate(state[\"messages\"][-1].tool_calls[0][\"args\"])\n",
|
||||||
|
" location = interrupt(ask.question)\n",
|
||||||
" tool_message = [{\"tool_call_id\": tool_call_id, \"type\": \"tool\", \"content\": location}]\n",
|
" tool_message = [{\"tool_call_id\": tool_call_id, \"type\": \"tool\", \"content\": location}]\n",
|
||||||
" return {\"messages\": tool_message}\n",
|
" return {\"messages\": tool_message}\n",
|
||||||
"\n",
|
"\n",
|
||||||
@@ -491,7 +492,7 @@
|
|||||||
" \"messages\": [\n",
|
" \"messages\": [\n",
|
||||||
" (\n",
|
" (\n",
|
||||||
" \"user\",\n",
|
" \"user\",\n",
|
||||||
" \"Use the search tool to ask the user where they are, then look up the weather there\",\n",
|
" \"Ask the user where they are, then look up the weather there\",\n",
|
||||||
" )\n",
|
" )\n",
|
||||||
" ]\n",
|
" ]\n",
|
||||||
" },\n",
|
" },\n",
|
||||||
|
|||||||
@@ -99,7 +99,7 @@
|
|||||||
"from typing import Literal\n",
|
"from typing import Literal\n",
|
||||||
"\n",
|
"\n",
|
||||||
"from langchain_anthropic import ChatAnthropic\n",
|
"from langchain_anthropic import ChatAnthropic\n",
|
||||||
"from langchain_core.messages import SystemMessage, RemoveMessage\n",
|
"from langchain_core.messages import SystemMessage, RemoveMessage, HumanMessage\n",
|
||||||
"from langgraph.checkpoint.memory import MemorySaver\n",
|
"from langgraph.checkpoint.memory import MemorySaver\n",
|
||||||
"from langgraph.graph import MessagesState, StateGraph, START, END\n",
|
"from langgraph.graph import MessagesState, StateGraph, START, END\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
|||||||
@@ -7,7 +7,7 @@
|
|||||||
"source": [
|
"source": [
|
||||||
"# How to manage conversation history\n",
|
"# How to manage conversation history\n",
|
||||||
"\n",
|
"\n",
|
||||||
"One of the most common use cases for persistence is to use it to keep track of conversation history. This is great - it makes it easy to continue conversations. As conversations get longer and longer, however, this conversation history can build up and take up more and more of the context window. This can often be undesirable as it leads to more expensive and longer calls to the LLM, and potentially ones that error. In order to prevent this from happening, you need to probably manage the conversation history.\n",
|
"One of the most common use cases for persistence is to use it to keep track of conversation history. This is great - it makes it easy to continue conversations. As conversations get longer and longer, however, this conversation history can build up and take up more and more of the context window. This can often be undesirable as it leads to more expensive and longer calls to the LLM, and potentially ones that error. In order to prevent this from happening, you need to properly manage the conversation history.\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Note: this guide focuses on how to do this in LangGraph, where you can fully customize how this is done. If you want a more off-the-shelf solution, you can look into functionality provided in LangChain:\n",
|
"Note: this guide focuses on how to do this in LangGraph, where you can fully customize how this is done. If you want a more off-the-shelf solution, you can look into functionality provided in LangChain:\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
|||||||
@@ -38,7 +38,7 @@
|
|||||||
" </p>\n",
|
" </p>\n",
|
||||||
"</div> \n",
|
"</div> \n",
|
||||||
"\n",
|
"\n",
|
||||||
"The core technique the examples below is to **annotate** a parameter as \"injected\", meaning it will be injected by your program and should not be seen or populated by the LLM. Let the following codesnippet serve as a tl;dr:\n",
|
"The core technique in the examples below is to **annotate** a parameter as \"injected\", meaning it will be injected by your program and should not be seen or populated by the LLM. Let the following codesnippet serve as a tl;dr:\n",
|
||||||
"\n",
|
"\n",
|
||||||
"```python\n",
|
"```python\n",
|
||||||
"from typing import Annotated\n",
|
"from typing import Annotated\n",
|
||||||
|
|||||||
@@ -65,7 +65,7 @@
|
|||||||
"\n",
|
"\n",
|
||||||
"**Pros and Cons**\n",
|
"**Pros and Cons**\n",
|
||||||
"\n",
|
"\n",
|
||||||
"The benefit to this format is that you only need one LLM, and can save money and latency because of this. The downside to this option is that it isn't guaranteed that the single LLM will call the correct tool when you want it to. We can help the LLM by setting `tool_choice` to `any` when we use `bind_tools` which forces the LLM to select at least one tool at every turn, but this is far from a fool proof strategy. In addition, another downside is that the agent might call *multiple* tools, so we need to check for this explicitly in our routing function (or if we are using OpenAI we an set `parallell_tool_calling=False` to ensure only one tool is called at a time).\n",
|
"The benefit to this format is that you only need one LLM, and can save money and latency because of this. The downside to this option is that it isn't guaranteed that the single LLM will call the correct tool when you want it to. We can help the LLM by setting `tool_choice` to `any` when we use `bind_tools` which forces the LLM to select at least one tool at every turn, but this is far from a foolproof strategy. In addition, another downside is that the agent might call *multiple* tools, so we need to check for this explicitly in our routing function (or if we are using OpenAI we can set `parallell_tool_calling=False` to ensure only one tool is called at a time).\n",
|
||||||
"\n",
|
"\n",
|
||||||
"**Option 2**\n",
|
"**Option 2**\n",
|
||||||
"\n",
|
"\n",
|
||||||
|
|||||||
@@ -266,6 +266,235 @@
|
|||||||
" print(\"An exception was raised because bad_node sets `a` to an integer.\")\n",
|
" print(\"An exception was raised because bad_node sets `a` to an integer.\")\n",
|
||||||
" print(e)"
|
" print(e)"
|
||||||
]
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"id": "2270bc3c",
|
||||||
|
"metadata": {},
|
||||||
|
"source": [
|
||||||
|
"## Multiple Nodes\n",
|
||||||
|
"\n",
|
||||||
|
"Run-time validation will also work in a multi-node graph. In the example below `bad_node` updates `a` to an integer. \n",
|
||||||
|
"\n",
|
||||||
|
"Because run-time validation occurs on **inputs**, the validation error will occur when `ok_node` is called (not when `bad_node` returns an update to the state which is inconsistent with the schema)."
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"id": "d832cdcc",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"from langgraph.graph import StateGraph, START, END\n",
|
||||||
|
"from typing_extensions import TypedDict\n",
|
||||||
|
"\n",
|
||||||
|
"from pydantic import BaseModel\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# The overall state of the graph (this is the public state shared across nodes)\n",
|
||||||
|
"class OverallState(BaseModel):\n",
|
||||||
|
" a: str\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"def bad_node(state: OverallState):\n",
|
||||||
|
" return {\n",
|
||||||
|
" \"a\": 123 # Invalid\n",
|
||||||
|
" }\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"def ok_node(state: OverallState):\n",
|
||||||
|
" return {\"a\": \"goodbye\"}\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# Build the state graph\n",
|
||||||
|
"builder = StateGraph(OverallState)\n",
|
||||||
|
"builder.add_node(bad_node)\n",
|
||||||
|
"builder.add_node(ok_node)\n",
|
||||||
|
"builder.add_edge(START, \"bad_node\")\n",
|
||||||
|
"builder.add_edge(\"bad_node\", \"ok_node\")\n",
|
||||||
|
"builder.add_edge(\"ok_node\", END)\n",
|
||||||
|
"graph = builder.compile()\n",
|
||||||
|
"\n",
|
||||||
|
"# Test the graph with a valid input\n",
|
||||||
|
"try:\n",
|
||||||
|
" graph.invoke({\"a\": \"hello\"})\n",
|
||||||
|
"except Exception as e:\n",
|
||||||
|
" print(\"An exception was raised because bad_node sets `a` to an integer.\")\n",
|
||||||
|
" print(e)"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"id": "456b1f77",
|
||||||
|
"metadata": {},
|
||||||
|
"source": [
|
||||||
|
"## Advanced Pydantic Model Usage\n",
|
||||||
|
"\n",
|
||||||
|
"This section covers more advanced topics when using Pydantic models with LangGraph.\n",
|
||||||
|
"\n",
|
||||||
|
"### Serialization Behavior\n",
|
||||||
|
"\n",
|
||||||
|
"When using Pydantic models as state schemas, it's important to understand how serialization works, especially when:\n",
|
||||||
|
"- Passing Pydantic objects as inputs\n",
|
||||||
|
"- Receiving outputs from the graph\n",
|
||||||
|
"- Working with nested Pydantic models\n",
|
||||||
|
"\n",
|
||||||
|
"Let's see these behaviors in action:"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"id": "0e919cdc",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"from langgraph.graph import StateGraph, START, END\n",
|
||||||
|
"from pydantic import BaseModel\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"class NestedModel(BaseModel):\n",
|
||||||
|
" value: str\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"class ComplexState(BaseModel):\n",
|
||||||
|
" text: str\n",
|
||||||
|
" count: int\n",
|
||||||
|
" nested: NestedModel\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"def process_node(state: ComplexState):\n",
|
||||||
|
" # Node receives a validated Pydantic object\n",
|
||||||
|
" print(f\"Input state type: {type(state)}\")\n",
|
||||||
|
" print(f\"Nested type: {type(state.nested)}\")\n",
|
||||||
|
"\n",
|
||||||
|
" # Return a dictionary update\n",
|
||||||
|
" return {\"text\": state.text + \" processed\", \"count\": state.count + 1}\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# Build the graph\n",
|
||||||
|
"builder = StateGraph(ComplexState)\n",
|
||||||
|
"builder.add_node(\"process\", process_node)\n",
|
||||||
|
"builder.add_edge(START, \"process\")\n",
|
||||||
|
"builder.add_edge(\"process\", END)\n",
|
||||||
|
"graph = builder.compile()\n",
|
||||||
|
"\n",
|
||||||
|
"# Create a Pydantic instance for input\n",
|
||||||
|
"input_state = ComplexState(text=\"hello\", count=0, nested=NestedModel(value=\"test\"))\n",
|
||||||
|
"print(f\"Input object type: {type(input_state)}\")\n",
|
||||||
|
"\n",
|
||||||
|
"# Invoke graph with a Pydantic instance\n",
|
||||||
|
"result = graph.invoke(input_state)\n",
|
||||||
|
"print(f\"Output type: {type(result)}\")\n",
|
||||||
|
"print(f\"Output content: {result}\")\n",
|
||||||
|
"\n",
|
||||||
|
"# Convert back to Pydantic model if needed\n",
|
||||||
|
"output_model = ComplexState(**result)\n",
|
||||||
|
"print(f\"Converted back to Pydantic: {type(output_model)}\")"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"id": "f13f28ce",
|
||||||
|
"metadata": {},
|
||||||
|
"source": [
|
||||||
|
"### Runtime Type Coercion\n",
|
||||||
|
"\n",
|
||||||
|
"Pydantic performs runtime type coercion for certain data types. This can be helpful but also lead to unexpected behavior if you're not aware of it."
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"id": "faf59316",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"from langgraph.graph import StateGraph, START, END\n",
|
||||||
|
"from pydantic import BaseModel\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"class CoercionExample(BaseModel):\n",
|
||||||
|
" # Pydantic will coerce string numbers to integers\n",
|
||||||
|
" number: int\n",
|
||||||
|
" # Pydantic will parse string booleans to bool\n",
|
||||||
|
" flag: bool\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"def inspect_node(state: CoercionExample):\n",
|
||||||
|
" print(f\"number: {state.number} (type: {type(state.number)})\")\n",
|
||||||
|
" print(f\"flag: {state.flag} (type: {type(state.flag)})\")\n",
|
||||||
|
" return {}\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"builder = StateGraph(CoercionExample)\n",
|
||||||
|
"builder.add_node(\"inspect\", inspect_node)\n",
|
||||||
|
"builder.add_edge(START, \"inspect\")\n",
|
||||||
|
"builder.add_edge(\"inspect\", END)\n",
|
||||||
|
"graph = builder.compile()\n",
|
||||||
|
"\n",
|
||||||
|
"# Demonstrate coercion with string inputs that will be converted\n",
|
||||||
|
"result = graph.invoke({\"number\": \"42\", \"flag\": \"true\"})\n",
|
||||||
|
"\n",
|
||||||
|
"# This would fail with a validation error\n",
|
||||||
|
"try:\n",
|
||||||
|
" graph.invoke({\"number\": \"not-a-number\", \"flag\": \"true\"})\n",
|
||||||
|
"except Exception as e:\n",
|
||||||
|
" print(f\"\\nExpected validation error: {e}\")"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"id": "2844475b",
|
||||||
|
"metadata": {},
|
||||||
|
"source": [
|
||||||
|
"### Working with Message Models\n",
|
||||||
|
"\n",
|
||||||
|
"When working with LangChain message types in your state schema, there are important considerations for serialization. You should use `AnyMessage` (rather than `BaseMessage`) for proper serialization/deserialization when using message objects over the wire:"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"id": "bd0734b0",
|
||||||
|
"metadata": {},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"from langgraph.graph import StateGraph, START, END\n",
|
||||||
|
"from pydantic import BaseModel\n",
|
||||||
|
"from langchain_core.messages import HumanMessage, AIMessage, AnyMessage\n",
|
||||||
|
"from typing import List\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"class ChatState(BaseModel):\n",
|
||||||
|
" messages: List[AnyMessage]\n",
|
||||||
|
" context: str\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"def add_message(state: ChatState):\n",
|
||||||
|
" return {\"messages\": state.messages + [AIMessage(content=\"Hello there!\")]}\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"builder = StateGraph(ChatState)\n",
|
||||||
|
"builder.add_node(\"add_message\", add_message)\n",
|
||||||
|
"builder.add_edge(START, \"add_message\")\n",
|
||||||
|
"builder.add_edge(\"add_message\", END)\n",
|
||||||
|
"graph = builder.compile()\n",
|
||||||
|
"\n",
|
||||||
|
"# Create input with a message\n",
|
||||||
|
"initial_state = ChatState(\n",
|
||||||
|
" messages=[HumanMessage(content=\"Hi\")], context=\"Customer support chat\"\n",
|
||||||
|
")\n",
|
||||||
|
"\n",
|
||||||
|
"result = graph.invoke(initial_state)\n",
|
||||||
|
"print(f\"Output: {result}\")\n",
|
||||||
|
"\n",
|
||||||
|
"# Convert back to Pydantic model to see message types\n",
|
||||||
|
"output_model = ChatState(**result)\n",
|
||||||
|
"for i, msg in enumerate(output_model.messages):\n",
|
||||||
|
" print(f\"Message {i}: {type(msg).__name__} - {msg.content}\")"
|
||||||
|
]
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"metadata": {
|
"metadata": {
|
||||||
|
|||||||
@@ -210,7 +210,7 @@
|
|||||||
"id": "cbb06aea-6654-4245-91f8-af6e8f2b5377",
|
"id": "cbb06aea-6654-4245-91f8-af6e8f2b5377",
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"source": [
|
"source": [
|
||||||
"Let's now add personalization: we'll respond differently to the user based on the state values AFTER the state has been updated from the tool. To achieve this, let's define a function that will dynamically construct the system prompt based on the graph state. It will be called ever time the LLM is called and the function output will be passed to the LLM:"
|
"Let's now add personalization: we'll respond differently to the user based on the state values AFTER the state has been updated from the tool. To achieve this, let's define a function that will dynamically construct the system prompt based on the graph state. It will be called every time the LLM is called and the function output will be passed to the LLM:"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
+1
-1
@@ -20,7 +20,7 @@ title: Home
|
|||||||
</p>
|
</p>
|
||||||
|
|
||||||
<style>
|
<style>
|
||||||
h1 {
|
.md-content h1 {
|
||||||
display: none;
|
display: none;
|
||||||
}
|
}
|
||||||
</style>
|
</style>
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|
||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -79,12 +79,12 @@ CREATE INDEX CONCURRENTLY IF NOT EXISTS store_prefix_idx ON store USING btree (p
|
|||||||
"""
|
"""
|
||||||
-- Add expires_at column to store table
|
-- Add expires_at column to store table
|
||||||
ALTER TABLE store
|
ALTER TABLE store
|
||||||
ADD COLUMN expires_at TIMESTAMP WITH TIME ZONE,
|
ADD COLUMN IF NOT EXISTS expires_at TIMESTAMP WITH TIME ZONE,
|
||||||
ADD COLUMN ttl_minutes INT;
|
ADD COLUMN IF NOT EXISTS ttl_minutes INT;
|
||||||
""",
|
""",
|
||||||
"""
|
"""
|
||||||
-- Add indexes for efficient TTL sweeping
|
-- Add indexes for efficient TTL sweeping
|
||||||
CREATE INDEX idx_store_expires_at ON store (expires_at)
|
CREATE INDEX IF NOT EXISTS idx_store_expires_at ON store (expires_at)
|
||||||
WHERE expires_at IS NOT NULL;
|
WHERE expires_at IS NOT NULL;
|
||||||
""",
|
""",
|
||||||
]
|
]
|
||||||
|
|||||||
Generated
+2
-2
@@ -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"
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[tool.poetry]
|
[tool.poetry]
|
||||||
name = "langgraph-checkpoint-postgres"
|
name = "langgraph-checkpoint-postgres"
|
||||||
version = "2.0.17"
|
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"
|
||||||
|
|||||||
@@ -60,15 +60,17 @@ async def store(request) -> AsyncIterator[AsyncPostgresStore]:
|
|||||||
) as store:
|
) as store:
|
||||||
store.MIGRATIONS = [
|
store.MIGRATIONS = [
|
||||||
(
|
(
|
||||||
mig.replace(
|
mig.replace("ttl_minutes INT;", "ttl_minutes FLOAT;")
|
||||||
"ADD COLUMN ttl_minutes INT;", "ADD COLUMN ttl_minutes FLOAT;"
|
|
||||||
)
|
|
||||||
if isinstance(mig, str)
|
if isinstance(mig, str)
|
||||||
else mig
|
else mig
|
||||||
)
|
)
|
||||||
for mig in store.MIGRATIONS
|
for mig in store.MIGRATIONS
|
||||||
]
|
]
|
||||||
await store.setup()
|
await store.setup()
|
||||||
|
async with store._cursor() as cur:
|
||||||
|
# drop the migration index
|
||||||
|
await cur.execute("DROP TABLE IF EXISTS store_migrations")
|
||||||
|
await store.setup() # Will fail if migrations aren't idempotent
|
||||||
|
|
||||||
if request.param == "pipe":
|
if request.param == "pipe":
|
||||||
async with AsyncPostgresStore.from_conn_string(
|
async with AsyncPostgresStore.from_conn_string(
|
||||||
|
|||||||
@@ -52,9 +52,7 @@ def store(request) -> PostgresStore:
|
|||||||
with PostgresStore.from_conn_string(conn_string, ttl=ttl_config) as store:
|
with PostgresStore.from_conn_string(conn_string, ttl=ttl_config) as store:
|
||||||
store.MIGRATIONS = [
|
store.MIGRATIONS = [
|
||||||
(
|
(
|
||||||
mig.replace(
|
mig.replace("ttl_minutes INT;", "ttl_minutes FLOAT;")
|
||||||
"ADD COLUMN ttl_minutes INT;", "ADD COLUMN ttl_minutes FLOAT;"
|
|
||||||
)
|
|
||||||
if isinstance(mig, str)
|
if isinstance(mig, str)
|
||||||
else mig
|
else mig
|
||||||
)
|
)
|
||||||
@@ -415,6 +413,10 @@ def _create_vector_store(
|
|||||||
ttl={"default_ttl": 2, "refresh_on_read": True} if enable_ttl else None,
|
ttl={"default_ttl": 2, "refresh_on_read": True} if enable_ttl else None,
|
||||||
) as store:
|
) as store:
|
||||||
store.setup()
|
store.setup()
|
||||||
|
with store._cursor() as cur:
|
||||||
|
# drop the migration index
|
||||||
|
cur.execute("DROP TABLE IF EXISTS store_migrations")
|
||||||
|
store.setup() # Will fail if migrations aren't idempotent
|
||||||
yield store
|
yield store
|
||||||
finally:
|
finally:
|
||||||
with Connection.connect(admin_conn_string, autocommit=True) as conn:
|
with Connection.connect(admin_conn_string, autocommit=True) as conn:
|
||||||
|
|||||||
@@ -56,7 +56,10 @@ class SqliteSaver(BaseCheckpointSaver[str]):
|
|||||||
>>> builder.add_node("add_one", lambda x: x + 1)
|
>>> builder.add_node("add_one", lambda x: x + 1)
|
||||||
>>> builder.set_entry_point("add_one")
|
>>> builder.set_entry_point("add_one")
|
||||||
>>> builder.set_finish_point("add_one")
|
>>> builder.set_finish_point("add_one")
|
||||||
>>> conn = sqlite3.connect("checkpoints.sqlite")
|
>>> # Create a new SqliteSaver instance
|
||||||
|
>>> # Note: check_same_thread=False is OK as the implementation uses a lock
|
||||||
|
>>> # to ensure thread safety.
|
||||||
|
>>> conn = sqlite3.connect("checkpoints.sqlite", check_same_thread=False)
|
||||||
>>> memory = SqliteSaver(conn)
|
>>> memory = SqliteSaver(conn)
|
||||||
>>> graph = builder.compile(checkpointer=memory)
|
>>> graph = builder.compile(checkpointer=memory)
|
||||||
>>> config = {"configurable": {"thread_id": "1"}}
|
>>> config = {"configurable": {"thread_id": "1"}}
|
||||||
|
|||||||
@@ -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"]: (
|
||||||
|
|||||||
@@ -45,3 +45,18 @@ def maybe_add_typed_methods(serde: SerializerProtocol) -> SerializerProtocol:
|
|||||||
return SerializerCompat(serde)
|
return SerializerCompat(serde)
|
||||||
|
|
||||||
return serde
|
return serde
|
||||||
|
|
||||||
|
|
||||||
|
class CipherProtocol(Protocol):
|
||||||
|
"""Protocol for encryption and decryption of data.
|
||||||
|
- `encrypt`: Encrypt plaintext.
|
||||||
|
- `decrypt`: Decrypt ciphertext.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def encrypt(self, plaintext: bytes) -> tuple[str, bytes]:
|
||||||
|
"""Encrypt plaintext. Returns a tuple (cipher name, ciphertext)."""
|
||||||
|
...
|
||||||
|
|
||||||
|
def decrypt(self, ciphername: str, ciphertext: bytes) -> bytes:
|
||||||
|
"""Decrypt ciphertext. Returns the plaintext."""
|
||||||
|
...
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
import os
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langgraph.checkpoint.serde.base import CipherProtocol, SerializerProtocol
|
||||||
|
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
|
||||||
|
|
||||||
|
|
||||||
|
class EncryptedSerializer(SerializerProtocol):
|
||||||
|
"""Serializer that encrypts and decrypts data using an encryption protocol."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, cipher: CipherProtocol, serde: SerializerProtocol = JsonPlusSerializer()
|
||||||
|
) -> None:
|
||||||
|
self.cipher = cipher
|
||||||
|
self.serde = serde
|
||||||
|
|
||||||
|
def dumps(self, obj: Any) -> bytes:
|
||||||
|
return self.serde.dumps(obj)
|
||||||
|
|
||||||
|
def loads(self, data: bytes) -> Any:
|
||||||
|
return self.serde.loads(data)
|
||||||
|
|
||||||
|
def dumps_typed(self, obj: Any) -> tuple[str, bytes]:
|
||||||
|
"""Serialize an object to a tuple (type, bytes) and encrypt the bytes."""
|
||||||
|
# serialize data
|
||||||
|
typ, data = self.serde.dumps_typed(obj)
|
||||||
|
# encrypt data
|
||||||
|
ciphername, ciphertext = self.cipher.encrypt(data)
|
||||||
|
# add cipher name to type
|
||||||
|
return f"{typ}+{ciphername}", ciphertext
|
||||||
|
|
||||||
|
def loads_typed(self, data: tuple[str, bytes]) -> Any:
|
||||||
|
enc_cipher, ciphertext = data
|
||||||
|
# unencrypted data
|
||||||
|
if "+" not in enc_cipher:
|
||||||
|
return self.serde.loads_typed(data)
|
||||||
|
# extract cipher name
|
||||||
|
typ, ciphername = enc_cipher.split("+", 1)
|
||||||
|
# decrypt data
|
||||||
|
decrypted_data = self.cipher.decrypt(ciphername, ciphertext)
|
||||||
|
# deserialize data
|
||||||
|
return self.serde.loads_typed((typ, decrypted_data))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pycryptodome_aes(
|
||||||
|
cls, serde: SerializerProtocol = JsonPlusSerializer(), **kwargs: Any
|
||||||
|
) -> "EncryptedSerializer":
|
||||||
|
"""Create an EncryptedSerializer using AES encryption."""
|
||||||
|
try:
|
||||||
|
from Crypto.Cipher import AES # type: ignore
|
||||||
|
except ImportError:
|
||||||
|
raise ImportError(
|
||||||
|
"Pycryptodome is not installed. Please install it with `pip install pycryptodome`."
|
||||||
|
) from None
|
||||||
|
|
||||||
|
# check if AES key is provided
|
||||||
|
if "key" in kwargs:
|
||||||
|
key: bytes = kwargs.pop("key")
|
||||||
|
else:
|
||||||
|
key_str = os.getenv("LANGGRAPH_AES_KEY")
|
||||||
|
if key_str is None:
|
||||||
|
raise ValueError("LANGGRAPH_AES_KEY environment variable is not set.")
|
||||||
|
key = key_str.encode()
|
||||||
|
if len(key) not in (16, 24, 32):
|
||||||
|
raise ValueError("LANGGRAPH_AES_KEY must be 16, 24, or 32 bytes long.")
|
||||||
|
|
||||||
|
# set default mode to EAX if not provided
|
||||||
|
if kwargs.get("mode") is None:
|
||||||
|
kwargs["mode"] = AES.MODE_EAX
|
||||||
|
|
||||||
|
class PycryptodomeAesCipher(CipherProtocol):
|
||||||
|
def encrypt(self, plaintext: bytes) -> tuple[str, bytes]:
|
||||||
|
cipher = AES.new(key, **kwargs)
|
||||||
|
ciphertext, tag = cipher.encrypt_and_digest(plaintext)
|
||||||
|
return "aes", cipher.nonce + tag + ciphertext
|
||||||
|
|
||||||
|
def decrypt(self, ciphername: str, ciphertext: bytes) -> bytes:
|
||||||
|
assert ciphername == "aes", f"Unsupported cipher: {ciphername}"
|
||||||
|
nonce = ciphertext[:16]
|
||||||
|
tag = ciphertext[16:32]
|
||||||
|
actual_ciphertext = ciphertext[32:]
|
||||||
|
|
||||||
|
cipher = AES.new(key, **kwargs, nonce=nonce)
|
||||||
|
return cipher.decrypt_and_verify(actual_ciphertext, tag)
|
||||||
|
|
||||||
|
return cls(PycryptodomeAesCipher(), serde)
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
[tool.poetry]
|
[tool.poetry]
|
||||||
name = "langgraph-checkpoint"
|
name = "langgraph-checkpoint"
|
||||||
version = "2.0.20"
|
version = "2.0.21"
|
||||||
description = "Library with base interfaces for LangGraph checkpoint savers."
|
description = "Library with base interfaces for LangGraph checkpoint savers."
|
||||||
authors = []
|
authors = []
|
||||||
license = "MIT"
|
license = "MIT"
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Generated
+19
-19
@@ -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"
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -58,9 +58,11 @@ WORKERS ?= auto
|
|||||||
XDIST_ARGS := $(if $(WORKERS),-n $(WORKERS) --dist worksteal,)
|
XDIST_ARGS := $(if $(WORKERS),-n $(WORKERS) --dist worksteal,)
|
||||||
MAXFAIL ?=
|
MAXFAIL ?=
|
||||||
MAXFAIL_ARGS := $(if $(MAXFAIL),--maxfail $(MAXFAIL),)
|
MAXFAIL_ARGS := $(if $(MAXFAIL),--maxfail $(MAXFAIL),)
|
||||||
|
# Add an '-x' if xdist is enabled
|
||||||
|
XDIST_ARGS := $(if $(WORKERS),-x $(XDIST_ARGS),)
|
||||||
|
|
||||||
test_watch:
|
test_watch:
|
||||||
make start-postgres && poetry run ptw . -- --ff -vv -x $(XDIST_ARGS) $(MAXFAIL_ARGS) --snapshot-update --tb short $(TEST); \
|
make start-postgres && poetry run ptw . -- --ff -vv $(XDIST_ARGS) $(MAXFAIL_ARGS) --snapshot-update --tb short $(TEST); \
|
||||||
EXIT_CODE=$$?; \
|
EXIT_CODE=$$?; \
|
||||||
make stop-postgres; \
|
make stop-postgres; \
|
||||||
exit $$EXIT_CODE
|
exit $$EXIT_CODE
|
||||||
|
|||||||
@@ -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>
|
||||||
|
|||||||
@@ -6,9 +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.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
|
||||||
|
|
||||||
|
|
||||||
@@ -42,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",
|
||||||
@@ -203,12 +211,168 @@ benchmarks = (
|
|||||||
]
|
]
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
"sequential_20",
|
||||||
|
create_sequential(20).compile(),
|
||||||
|
create_sequential(20).compile(),
|
||||||
|
{"messages": []}, # Empty list of messages
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"sequential_50",
|
||||||
|
create_sequential(50).compile(),
|
||||||
|
create_sequential(50).compile(),
|
||||||
|
{"messages": []}, # Empty list of messages
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"sequential_100",
|
||||||
|
create_sequential(100).compile(),
|
||||||
|
create_sequential(100).compile(),
|
||||||
|
{"messages": []}, # Empty list of messages
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"sequential_200",
|
||||||
|
create_sequential(200).compile(),
|
||||||
|
create_sequential(200).compile(),
|
||||||
|
{"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)
|
||||||
|
|||||||
@@ -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())
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
"""Create a sequential no-op graph consisting of a few hundred nodes."""
|
||||||
|
|
||||||
|
from langgraph.graph import MessagesState, StateGraph
|
||||||
|
from langgraph.utils.runnable import RunnableCallable
|
||||||
|
|
||||||
|
|
||||||
|
def create_sequential(number_nodes) -> StateGraph:
|
||||||
|
"""Create a sequential no-op graph consisting of a few hundred nodes."""
|
||||||
|
builder = StateGraph(MessagesState)
|
||||||
|
|
||||||
|
def noop(state: MessagesState) -> None:
|
||||||
|
"""No-op function."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def anoop(state: MessagesState) -> None:
|
||||||
|
"""No-op function."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
prev_node = "__start__"
|
||||||
|
|
||||||
|
for i in range(number_nodes):
|
||||||
|
name = f"node_{i}"
|
||||||
|
builder.add_node(name, RunnableCallable(noop, anoop))
|
||||||
|
builder.add_edge(prev_node, name)
|
||||||
|
prev_node = name
|
||||||
|
|
||||||
|
builder.add_edge(prev_node, "__end__")
|
||||||
|
return builder
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
|
||||||
|
import uvloop
|
||||||
|
|
||||||
|
graph = create_sequential(2000).compile()
|
||||||
|
input = {"messages": []} # Empty list of messages
|
||||||
|
config = {"recursion_limit": 20000000000}
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
len([c async for c in graph.astream(input, config=config)])
|
||||||
|
|
||||||
|
uvloop.install()
|
||||||
|
start = time.time()
|
||||||
|
asyncio.run(run())
|
||||||
|
end = time.time()
|
||||||
|
print(f"Time taken: {end - start:.4f} seconds")
|
||||||
@@ -1,8 +1,9 @@
|
|||||||
from typing import Generic, Optional, Sequence, Type
|
from typing import Any, Generic, Optional, Sequence, Type
|
||||||
|
|
||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from langgraph.channels.base import BaseChannel, Value
|
from langgraph.channels.base import BaseChannel, Value
|
||||||
|
from langgraph.constants import MISSING
|
||||||
from langgraph.errors import EmptyChannelError
|
from langgraph.errors import EmptyChannelError
|
||||||
|
|
||||||
|
|
||||||
@@ -12,6 +13,10 @@ class AnyValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
|
|
||||||
__slots__ = ("typ", "value")
|
__slots__ = ("typ", "value")
|
||||||
|
|
||||||
|
def __init__(self, typ: Any, key: str = "") -> None:
|
||||||
|
super().__init__(typ, key)
|
||||||
|
self.value = MISSING
|
||||||
|
|
||||||
def __eq__(self, value: object) -> bool:
|
def __eq__(self, value: object) -> bool:
|
||||||
return isinstance(value, AnyValue)
|
return isinstance(value, AnyValue)
|
||||||
|
|
||||||
@@ -34,17 +39,19 @@ class AnyValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
|
|
||||||
def update(self, values: Sequence[Value]) -> bool:
|
def update(self, values: Sequence[Value]) -> bool:
|
||||||
if len(values) == 0:
|
if len(values) == 0:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
del self.value
|
|
||||||
return True
|
|
||||||
except AttributeError:
|
|
||||||
return False
|
return False
|
||||||
|
else:
|
||||||
|
self.value = MISSING
|
||||||
|
return True
|
||||||
|
|
||||||
self.value = values[-1]
|
self.value = values[-1]
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def get(self) -> Value:
|
def get(self) -> Value:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
return self.value
|
|
||||||
except AttributeError:
|
|
||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.value is not MISSING
|
||||||
|
|||||||
@@ -64,6 +64,17 @@ class BaseChannel(Generic[Value, Update, C], ABC):
|
|||||||
"""
|
"""
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
"""Return True if the channel is available (not empty), False otherwise.
|
||||||
|
Subclasses should override this method to provide a more efficient
|
||||||
|
implementation than calling get() and catching EmptyChannelError.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
self.get()
|
||||||
|
return True
|
||||||
|
except EmptyChannelError:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseChannel",
|
"BaseChannel",
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ from typing import (
|
|||||||
from typing_extensions import NotRequired, Required, Self
|
from typing_extensions import NotRequired, Required, Self
|
||||||
|
|
||||||
from langgraph.channels.base import BaseChannel, Value
|
from langgraph.channels.base import BaseChannel, Value
|
||||||
|
from langgraph.constants import MISSING
|
||||||
from langgraph.errors import EmptyChannelError
|
from langgraph.errors import EmptyChannelError
|
||||||
|
|
||||||
|
|
||||||
@@ -51,7 +52,7 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
try:
|
try:
|
||||||
self.value = typ()
|
self.value = typ()
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
self.value = MISSING
|
||||||
|
|
||||||
def __eq__(self, value: object) -> bool:
|
def __eq__(self, value: object) -> bool:
|
||||||
return isinstance(value, BinaryOperatorAggregate) and (
|
return isinstance(value, BinaryOperatorAggregate) and (
|
||||||
@@ -81,7 +82,7 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
def update(self, values: Sequence[Value]) -> bool:
|
def update(self, values: Sequence[Value]) -> bool:
|
||||||
if not values:
|
if not values:
|
||||||
return False
|
return False
|
||||||
if not hasattr(self, "value"):
|
if self.value is MISSING:
|
||||||
self.value = values[0]
|
self.value = values[0]
|
||||||
values = values[1:]
|
values = values[1:]
|
||||||
for value in values:
|
for value in values:
|
||||||
@@ -89,7 +90,9 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def get(self) -> Value:
|
def get(self) -> Value:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
return self.value
|
|
||||||
except AttributeError:
|
|
||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.value is not MISSING
|
||||||
|
|||||||
@@ -85,6 +85,9 @@ class DynamicBarrierValue(
|
|||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.seen == self.names
|
||||||
|
|
||||||
def consume(self) -> bool:
|
def consume(self) -> bool:
|
||||||
if self.seen == self.names:
|
if self.seen == self.names:
|
||||||
self.seen = set()
|
self.seen = set()
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from typing import Any, Generic, Optional, Sequence, Type
|
|||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from langgraph.channels.base import BaseChannel, Value
|
from langgraph.channels.base import BaseChannel, Value
|
||||||
|
from langgraph.constants import MISSING
|
||||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||||
|
|
||||||
|
|
||||||
@@ -14,6 +15,7 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
def __init__(self, typ: Any, guard: bool = True) -> None:
|
def __init__(self, typ: Any, guard: bool = True) -> None:
|
||||||
super().__init__(typ)
|
super().__init__(typ)
|
||||||
self.guard = guard
|
self.guard = guard
|
||||||
|
self.value = MISSING
|
||||||
|
|
||||||
def __eq__(self, value: object) -> bool:
|
def __eq__(self, value: object) -> bool:
|
||||||
return isinstance(value, EphemeralValue) and value.guard == self.guard
|
return isinstance(value, EphemeralValue) and value.guard == self.guard
|
||||||
@@ -37,10 +39,10 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
|
|
||||||
def update(self, values: Sequence[Value]) -> bool:
|
def update(self, values: Sequence[Value]) -> bool:
|
||||||
if len(values) == 0:
|
if len(values) == 0:
|
||||||
try:
|
if self.value is not MISSING:
|
||||||
del self.value
|
self.value = MISSING
|
||||||
return True
|
return True
|
||||||
except AttributeError:
|
else:
|
||||||
return False
|
return False
|
||||||
if len(values) != 1 and self.guard:
|
if len(values) != 1 and self.guard:
|
||||||
raise InvalidUpdateError(
|
raise InvalidUpdateError(
|
||||||
@@ -51,7 +53,9 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def get(self) -> Value:
|
def get(self) -> Value:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
return self.value
|
|
||||||
except AttributeError:
|
|
||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.value is not MISSING
|
||||||
|
|||||||
@@ -1,8 +1,9 @@
|
|||||||
from typing import Generic, Optional, Sequence, Type
|
from typing import Any, Generic, Optional, Sequence, Type
|
||||||
|
|
||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from langgraph.channels.base import BaseChannel, Value
|
from langgraph.channels.base import BaseChannel, Value
|
||||||
|
from langgraph.constants import MISSING
|
||||||
from langgraph.errors import (
|
from langgraph.errors import (
|
||||||
EmptyChannelError,
|
EmptyChannelError,
|
||||||
ErrorCode,
|
ErrorCode,
|
||||||
@@ -16,6 +17,10 @@ class LastValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
|
|
||||||
__slots__ = ("value",)
|
__slots__ = ("value",)
|
||||||
|
|
||||||
|
def __init__(self, typ: Any, key: str = "") -> None:
|
||||||
|
super().__init__(typ, key)
|
||||||
|
self.value = MISSING
|
||||||
|
|
||||||
def __eq__(self, value: object) -> bool:
|
def __eq__(self, value: object) -> bool:
|
||||||
return isinstance(value, LastValue)
|
return isinstance(value, LastValue)
|
||||||
|
|
||||||
@@ -50,7 +55,9 @@ class LastValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def get(self) -> Value:
|
def get(self) -> Value:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
return self.value
|
|
||||||
except AttributeError:
|
|
||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.value is not MISSING
|
||||||
|
|||||||
@@ -60,6 +60,9 @@ class NamedBarrierValue(Generic[Value], BaseChannel[Value, Value, set[Value]]):
|
|||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.seen == self.names
|
||||||
|
|
||||||
def consume(self) -> bool:
|
def consume(self) -> bool:
|
||||||
if self.seen == self.names:
|
if self.seen == self.names:
|
||||||
self.seen = set()
|
self.seen = set()
|
||||||
|
|||||||
@@ -75,3 +75,6 @@ class Topic(
|
|||||||
return list(self.values)
|
return list(self.values)
|
||||||
else:
|
else:
|
||||||
raise EmptyChannelError
|
raise EmptyChannelError
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return bool(self.values)
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from typing import Generic, Optional, Sequence, Type
|
|||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from langgraph.channels.base import BaseChannel, Value
|
from langgraph.channels.base import BaseChannel, Value
|
||||||
|
from langgraph.constants import MISSING
|
||||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||||
|
|
||||||
|
|
||||||
@@ -14,6 +15,7 @@ class UntrackedValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
def __init__(self, typ: Type[Value], guard: bool = True) -> None:
|
def __init__(self, typ: Type[Value], guard: bool = True) -> None:
|
||||||
super().__init__(typ)
|
super().__init__(typ)
|
||||||
self.guard = guard
|
self.guard = guard
|
||||||
|
self.value = MISSING
|
||||||
|
|
||||||
def __eq__(self, value: object) -> bool:
|
def __eq__(self, value: object) -> bool:
|
||||||
return isinstance(value, UntrackedValue) and value.guard == self.guard
|
return isinstance(value, UntrackedValue) and value.guard == self.guard
|
||||||
@@ -48,7 +50,9 @@ class UntrackedValue(Generic[Value], BaseChannel[Value, Value, Value]):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
def get(self) -> Value:
|
def get(self) -> Value:
|
||||||
try:
|
if self.value is MISSING:
|
||||||
return self.value
|
|
||||||
except AttributeError:
|
|
||||||
raise EmptyChannelError()
|
raise EmptyChannelError()
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
def is_available(self) -> bool:
|
||||||
|
return self.value is not MISSING
|
||||||
|
|||||||
@@ -138,6 +138,7 @@ class Branch(NamedTuple):
|
|||||||
reader=reader,
|
reader=reader,
|
||||||
name=None,
|
name=None,
|
||||||
trace=False,
|
trace=False,
|
||||||
|
func_accepts_config=True,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from typing import (
|
from typing import (
|
||||||
@@ -31,6 +32,7 @@ from langgraph.constants import (
|
|||||||
)
|
)
|
||||||
from langgraph.graph.branch import Branch
|
from langgraph.graph.branch import Branch
|
||||||
from langgraph.pregel import Channel, Pregel
|
from langgraph.pregel import Channel, Pregel
|
||||||
|
from langgraph.pregel.protocol import PregelProtocol
|
||||||
from langgraph.pregel.read import PregelNode
|
from langgraph.pregel.read import PregelNode
|
||||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||||
from langgraph.types import All, Checkpointer
|
from langgraph.types import All, Checkpointer
|
||||||
@@ -418,7 +420,38 @@ class CompiledGraph(Pregel):
|
|||||||
*,
|
*,
|
||||||
xray: Union[int, bool] = False,
|
xray: Union[int, bool] = False,
|
||||||
) -> DrawableGraph:
|
) -> DrawableGraph:
|
||||||
return self.get_graph(config, xray=xray)
|
"""Returns a drawable representation of the computation graph."""
|
||||||
|
from langgraph.pregel.remote import RemoteGraph
|
||||||
|
|
||||||
|
# gather subgraphs
|
||||||
|
if xray:
|
||||||
|
subpregels: dict[str, PregelProtocol] = {
|
||||||
|
k: v
|
||||||
|
async for k, v in self.aget_subgraphs()
|
||||||
|
if isinstance(v, (CompiledGraph, RemoteGraph))
|
||||||
|
}
|
||||||
|
subgraphs = {
|
||||||
|
k: v
|
||||||
|
for k, v in zip(
|
||||||
|
subpregels,
|
||||||
|
await asyncio.gather(
|
||||||
|
*(
|
||||||
|
p.aget_graph(
|
||||||
|
config,
|
||||||
|
xray=xray
|
||||||
|
if isinstance(xray, bool) or xray <= 0
|
||||||
|
else xray - 1,
|
||||||
|
)
|
||||||
|
for p in subpregels.values()
|
||||||
|
)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
subgraphs = {}
|
||||||
|
|
||||||
|
# draw the graph
|
||||||
|
return self._draw_graph(config, subgraphs=subgraphs)
|
||||||
|
|
||||||
def get_graph(
|
def get_graph(
|
||||||
self,
|
self,
|
||||||
@@ -427,17 +460,36 @@ class CompiledGraph(Pregel):
|
|||||||
xray: Union[int, bool] = False,
|
xray: Union[int, bool] = False,
|
||||||
) -> DrawableGraph:
|
) -> DrawableGraph:
|
||||||
"""Returns a drawable representation of the computation graph."""
|
"""Returns a drawable representation of the computation graph."""
|
||||||
|
from langgraph.pregel.remote import RemoteGraph
|
||||||
|
|
||||||
|
# gather subgraphs
|
||||||
|
if xray:
|
||||||
|
subgraphs = {
|
||||||
|
k: v.get_graph(
|
||||||
|
config,
|
||||||
|
xray=xray if isinstance(xray, bool) or xray <= 0 else xray - 1,
|
||||||
|
)
|
||||||
|
for k, v in self.get_subgraphs()
|
||||||
|
if isinstance(v, (CompiledGraph, RemoteGraph))
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
subgraphs = {}
|
||||||
|
|
||||||
|
# draw the graph
|
||||||
|
return self._draw_graph(config, subgraphs=subgraphs)
|
||||||
|
|
||||||
|
def _draw_graph(
|
||||||
|
self,
|
||||||
|
config: Optional[RunnableConfig] = None,
|
||||||
|
*,
|
||||||
|
subgraphs: dict[str, DrawableGraph] = {},
|
||||||
|
) -> DrawableGraph:
|
||||||
|
# create the graph
|
||||||
graph = DrawableGraph()
|
graph = DrawableGraph()
|
||||||
start_nodes: dict[str, DrawableNode] = {
|
start_nodes: dict[str, DrawableNode] = {
|
||||||
START: graph.add_node(self.get_input_schema(config), START)
|
START: graph.add_node(self.get_input_schema(config), START)
|
||||||
}
|
}
|
||||||
end_nodes: dict[str, DrawableNode] = {}
|
end_nodes: dict[str, DrawableNode] = {}
|
||||||
if xray:
|
|
||||||
subgraphs = {
|
|
||||||
k: v for k, v in self.get_subgraphs() if isinstance(v, CompiledGraph)
|
|
||||||
}
|
|
||||||
else:
|
|
||||||
subgraphs = {}
|
|
||||||
|
|
||||||
def add_edge(
|
def add_edge(
|
||||||
start: str,
|
start: str,
|
||||||
@@ -463,13 +515,8 @@ class CompiledGraph(Pregel):
|
|||||||
metadata["__interrupt"] = "before"
|
metadata["__interrupt"] = "before"
|
||||||
elif key in self.interrupt_after_nodes:
|
elif key in self.interrupt_after_nodes:
|
||||||
metadata["__interrupt"] = "after"
|
metadata["__interrupt"] = "after"
|
||||||
if xray and key in subgraphs:
|
if key in subgraphs:
|
||||||
subgraph = subgraphs[key].get_graph(
|
subgraph = subgraphs[key]
|
||||||
config=config,
|
|
||||||
xray=xray - 1
|
|
||||||
if isinstance(xray, int) and not isinstance(xray, bool) and xray > 0
|
|
||||||
else xray,
|
|
||||||
)
|
|
||||||
subgraph.trim_first_node()
|
subgraph.trim_first_node()
|
||||||
subgraph.trim_last_node()
|
subgraph.trim_last_node()
|
||||||
if len(subgraph.nodes) >= 1:
|
if len(subgraph.nodes) >= 1:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from typing import (
|
|||||||
Union,
|
Union,
|
||||||
get_args,
|
get_args,
|
||||||
get_origin,
|
get_origin,
|
||||||
|
get_type_hints,
|
||||||
)
|
)
|
||||||
|
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -18,34 +19,56 @@ from typing_extensions import Annotated
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class SchemaCoercionMapper:
|
_cache: weakref.WeakKeyDictionary[Type[Any], dict[int, "SchemaCoercionMapper"]] = (
|
||||||
_cache: weakref.WeakKeyDictionary[Type[Any], dict[int, "SchemaCoercionMapper"]] = (
|
weakref.WeakKeyDictionary()
|
||||||
weakref.WeakKeyDictionary()
|
)
|
||||||
)
|
|
||||||
|
|
||||||
def __new__(cls, schema: Type[Any], max_depth: int = 5) -> "SchemaCoercionMapper":
|
|
||||||
if schema not in cls._cache:
|
class SchemaCoercionMapper:
|
||||||
cls._cache[schema] = {}
|
def __new__(
|
||||||
if max_depth in cls._cache[schema]:
|
cls,
|
||||||
return cls._cache[schema][max_depth]
|
schema: Type[Any],
|
||||||
|
type_hints: Optional[dict[str, Any]] = None,
|
||||||
|
max_depth: int = 12,
|
||||||
|
) -> "SchemaCoercionMapper":
|
||||||
|
if schema not in _cache:
|
||||||
|
_cache[schema] = {}
|
||||||
|
if max_depth in _cache[schema]:
|
||||||
|
return _cache[schema][max_depth]
|
||||||
|
|
||||||
inst = super().__new__(cls)
|
inst = super().__new__(cls)
|
||||||
cls._cache[schema][max_depth] = inst
|
_cache[schema][max_depth] = inst
|
||||||
return inst
|
return inst
|
||||||
|
|
||||||
def __init__(self, schema: Type[Any], max_depth: int = 5):
|
def __init__(
|
||||||
|
self,
|
||||||
|
schema: Type[Any],
|
||||||
|
type_hints: Optional[dict[str, Any]] = None,
|
||||||
|
max_depth: int = 12,
|
||||||
|
):
|
||||||
if hasattr(self, "_inited"):
|
if hasattr(self, "_inited"):
|
||||||
return
|
return
|
||||||
self._inited = True
|
self._inited = True
|
||||||
self.schema = schema
|
self.schema = schema
|
||||||
|
self.type_hints = (
|
||||||
|
type_hints
|
||||||
|
if type_hints is not None
|
||||||
|
else get_type_hints(schema, localns={schema.__name__: schema})
|
||||||
|
)
|
||||||
self.max_depth = max_depth
|
self.max_depth = max_depth
|
||||||
if hasattr(schema, "model_fields") and hasattr(schema, "model_construct"):
|
|
||||||
self._fields = {n: f.annotation for n, f in schema.model_fields.items()}
|
if issubclass(schema, BaseModel):
|
||||||
self._construct = schema.model_construct
|
self._fields = {
|
||||||
elif hasattr(schema, "__fields__") and callable(
|
n: self.type_hints.get(n, f.annotation)
|
||||||
getattr(schema, "construct", None)
|
for n, f in schema.model_fields.items()
|
||||||
):
|
}
|
||||||
self._fields = {n: f.annotation for n, f in schema.__fields__.items()}
|
self._construct: Callable[..., Any] = schema.model_construct
|
||||||
|
|
||||||
|
elif issubclass(schema, BaseModelV1):
|
||||||
|
self._fields = {
|
||||||
|
n: self.type_hints.get(n, f.annotation)
|
||||||
|
for n, f in schema.__fields__.items()
|
||||||
|
}
|
||||||
self._construct = schema.construct
|
self._construct = schema.construct
|
||||||
else:
|
else:
|
||||||
raise TypeError("Schema is neither valid Pydantic v1 nor v2 model.")
|
raise TypeError("Schema is neither valid Pydantic v1 nor v2 model.")
|
||||||
@@ -62,18 +85,23 @@ class SchemaCoercionMapper:
|
|||||||
processed = {}
|
processed = {}
|
||||||
if self._field_coercers is None:
|
if self._field_coercers is None:
|
||||||
self._field_coercers = {
|
self._field_coercers = {
|
||||||
n: self._build_coercer(t) for n, t in self._fields.items()
|
n: self._build_coercer(t, depth - 1) for n, t in self._fields.items()
|
||||||
}
|
}
|
||||||
for k, v in input_data.items():
|
for k, v in input_data.items():
|
||||||
fn = self._field_coercers.get(k)
|
fn = self._field_coercers.get(k)
|
||||||
processed[k] = fn(v, depth - 1) if fn else v
|
processed[k] = fn(v, depth - 1) if fn else v
|
||||||
return self._construct(**processed)
|
return self._construct(**processed)
|
||||||
|
|
||||||
def _build_coercer(self, field_type: Any) -> Callable[[Any, Any], Any]:
|
def _build_coercer(
|
||||||
|
self, field_type: Any, depth: int, throw: bool = False
|
||||||
|
) -> Callable[[Any, Any], Any]:
|
||||||
|
if depth == 0:
|
||||||
|
return self._passthrough
|
||||||
origin = get_origin(field_type)
|
origin = get_origin(field_type)
|
||||||
|
|
||||||
if origin is Annotated:
|
if origin is Annotated:
|
||||||
real_type, *_ = get_args(field_type)
|
real_type, *_ = get_args(field_type)
|
||||||
sub = self._build_coercer(real_type)
|
sub = self._build_coercer(real_type, depth - 1)
|
||||||
return lambda v, d: sub(v, d)
|
return lambda v, d: sub(v, d)
|
||||||
if isclass(field_type):
|
if isclass(field_type):
|
||||||
is_class_ = True
|
is_class_ = True
|
||||||
@@ -84,39 +112,54 @@ class SchemaCoercionMapper:
|
|||||||
is_base_model = False
|
is_base_model = False
|
||||||
|
|
||||||
if is_base_model:
|
if is_base_model:
|
||||||
mapper = SchemaCoercionMapper(field_type, self.max_depth)
|
mapper = SchemaCoercionMapper(field_type, max_depth=depth - 1)
|
||||||
return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v
|
return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v
|
||||||
if is_class_ and issubclass(field_type, BaseModelV1):
|
if is_class_ and issubclass(field_type, BaseModelV1):
|
||||||
mapper = SchemaCoercionMapper(field_type, self.max_depth)
|
mapper = SchemaCoercionMapper(field_type, max_depth=depth - 1)
|
||||||
return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v
|
return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v
|
||||||
if origin is list or field_type is list:
|
if origin is list or field_type is list:
|
||||||
args = get_args(field_type)
|
args = get_args(field_type)
|
||||||
if len(args) != 1:
|
if len(args) != 1:
|
||||||
return lambda v, d: v
|
return lambda v, d: v
|
||||||
sub = self._build_coercer(args[0])
|
sub = self._build_coercer(args[0], depth - 1)
|
||||||
|
|
||||||
def list_coercer(v: Any, d: Any) -> Any:
|
def list_coercer(v: Any, d: Any) -> Any:
|
||||||
if not isinstance(v, (list, tuple)):
|
if not isinstance(v, (list, tuple)):
|
||||||
raise TypeError(f"Expected list, got {type(v).__name__}")
|
return v
|
||||||
return [sub(x, d - 1) for x in v]
|
return [sub(x, d - 1) for x in v]
|
||||||
|
|
||||||
return list_coercer
|
return list_coercer
|
||||||
|
if origin is set or field_type is set:
|
||||||
|
args = get_args(field_type)
|
||||||
|
if len(args) != 1:
|
||||||
|
return lambda v, d: v
|
||||||
|
sub = self._build_coercer(args[0], depth - 1)
|
||||||
|
|
||||||
|
def set_coercer(v: Any, d: Any) -> Any:
|
||||||
|
if not isinstance(v, (list, tuple, set)):
|
||||||
|
return v
|
||||||
|
return {sub(x, d - 1) for x in v}
|
||||||
|
|
||||||
|
return set_coercer
|
||||||
if origin is dict or field_type is dict:
|
if origin is dict or field_type is dict:
|
||||||
args = get_args(field_type)
|
args = get_args(field_type)
|
||||||
if len(args) != 2:
|
if len(args) != 2:
|
||||||
|
|
||||||
def plain_dict_coercer(v: Any, d: Any) -> Any:
|
def dict_coercer(v: Any, d: Any) -> Any:
|
||||||
if not isinstance(v, dict):
|
if not isinstance(v, dict):
|
||||||
raise TypeError(f"Expected dict, got {type(v).__name__}")
|
if throw:
|
||||||
|
raise TypeError("Expected dict, got %s" % type(v))
|
||||||
return v
|
return v
|
||||||
|
|
||||||
return plain_dict_coercer
|
return dict_coercer
|
||||||
k_sub = self._build_coercer(args[0])
|
k_sub = self._build_coercer(args[0], depth - 1)
|
||||||
v_sub = self._build_coercer(args[1])
|
v_sub = self._build_coercer(args[1], depth - 1)
|
||||||
|
|
||||||
def dict_coercer(v: Any, d: Any) -> Any:
|
def dict_coercer(v: Any, d: Any) -> Any:
|
||||||
if not isinstance(v, dict):
|
if not isinstance(v, dict):
|
||||||
raise TypeError(f"Expected dict, got {type(v).__name__}")
|
if throw:
|
||||||
|
raise TypeError("Expected dict, got %s" % type(v))
|
||||||
|
return v
|
||||||
return {k_sub(k, d - 1): v_sub(val, d - 1) for k, val in v.items()}
|
return {k_sub(k, d - 1): v_sub(val, d - 1) for k, val in v.items()}
|
||||||
|
|
||||||
return dict_coercer
|
return dict_coercer
|
||||||
@@ -125,11 +168,11 @@ class SchemaCoercionMapper:
|
|||||||
targs = get_args(field_type)
|
targs = get_args(field_type)
|
||||||
if not targs:
|
if not targs:
|
||||||
return lambda v, d: v
|
return lambda v, d: v
|
||||||
subs = [self._build_coercer(a) for a in targs]
|
subs = [self._build_coercer(a, depth - 1) for a in targs]
|
||||||
|
|
||||||
def tuple_coercer(v: Any, d: Any) -> Any:
|
def tuple_coercer(v: Any, d: Any) -> Any:
|
||||||
if not isinstance(v, (list, tuple)):
|
if not isinstance(v, (list, tuple)):
|
||||||
raise TypeError(f"Expected tuple-like, got {type(v).__name__}")
|
return v
|
||||||
out = []
|
out = []
|
||||||
for i, sp in enumerate(subs):
|
for i, sp in enumerate(subs):
|
||||||
out.append(sp(v[i] if i < len(v) else None, d - 1))
|
out.append(sp(v[i] if i < len(v) else None, d - 1))
|
||||||
@@ -139,11 +182,13 @@ class SchemaCoercionMapper:
|
|||||||
if origin is Union:
|
if origin is Union:
|
||||||
uargs = get_args(field_type)
|
uargs = get_args(field_type)
|
||||||
subs, none_in_union = [], False
|
subs, none_in_union = [], False
|
||||||
for arg in uargs:
|
for ix, arg in enumerate(uargs):
|
||||||
if arg is type(None):
|
if arg is type(None):
|
||||||
none_in_union = True
|
none_in_union = True
|
||||||
else:
|
else:
|
||||||
subs.append(self._build_coercer(arg))
|
subs.append(
|
||||||
|
self._build_coercer(arg, depth - 1, throw=ix < len(uargs) - 1)
|
||||||
|
)
|
||||||
|
|
||||||
def union_coercer(v: Any, d: Any) -> Any:
|
def union_coercer(v: Any, d: Any) -> Any:
|
||||||
if v is None and none_in_union:
|
if v is None and none_in_union:
|
||||||
@@ -152,11 +197,14 @@ class SchemaCoercionMapper:
|
|||||||
for sp in subs:
|
for sp in subs:
|
||||||
try:
|
try:
|
||||||
return sp(v, d - 1)
|
return sp(v, d - 1)
|
||||||
except Exception as e:
|
except TypeError as e:
|
||||||
err = e
|
err = e
|
||||||
if err:
|
if err:
|
||||||
raise err
|
raise err
|
||||||
return v
|
return v
|
||||||
|
|
||||||
return union_coercer
|
return union_coercer
|
||||||
return lambda v, d: v
|
return self._passthrough
|
||||||
|
|
||||||
|
def _passthrough(self, v: Any, d: Any) -> Any:
|
||||||
|
return v
|
||||||
|
|||||||
@@ -185,6 +185,7 @@ class StateGraph(Graph):
|
|||||||
self.schemas = {}
|
self.schemas = {}
|
||||||
self.channels = {}
|
self.channels = {}
|
||||||
self.managed = {}
|
self.managed = {}
|
||||||
|
self.type_hints: dict[Type[Any], dict[str, Any]] = {}
|
||||||
self.schema = state_schema
|
self.schema = state_schema
|
||||||
self.input = input
|
self.input = input
|
||||||
self.output = output
|
self.output = output
|
||||||
@@ -203,7 +204,7 @@ class StateGraph(Graph):
|
|||||||
def _add_schema(self, schema: Type[Any], /, allow_managed: bool = True) -> None:
|
def _add_schema(self, schema: Type[Any], /, allow_managed: bool = True) -> None:
|
||||||
if schema not in self.schemas:
|
if schema not in self.schemas:
|
||||||
_warn_invalid_state_schema(schema)
|
_warn_invalid_state_schema(schema)
|
||||||
channels, managed = _get_channels(schema)
|
channels, managed, type_hints = _get_channels(schema)
|
||||||
if managed and not allow_managed:
|
if managed and not allow_managed:
|
||||||
names = ", ".join(managed)
|
names = ", ".join(managed)
|
||||||
schema_name = getattr(schema, "__name__", "")
|
schema_name = getattr(schema, "__name__", "")
|
||||||
@@ -212,6 +213,7 @@ class StateGraph(Graph):
|
|||||||
" Managed channels are not permitted in Input/Output schema."
|
" Managed channels are not permitted in Input/Output schema."
|
||||||
)
|
)
|
||||||
self.schemas[schema] = {**channels, **managed}
|
self.schemas[schema] = {**channels, **managed}
|
||||||
|
self.type_hints[schema] = type_hints
|
||||||
for key, channel in channels.items():
|
for key, channel in channels.items():
|
||||||
if key in self.channels:
|
if key in self.channels:
|
||||||
if self.channels[key] != channel:
|
if self.channels[key] != channel:
|
||||||
@@ -416,7 +418,7 @@ class StateGraph(Graph):
|
|||||||
and (vals := get_args(rargs[0]))
|
and (vals := get_args(rargs[0]))
|
||||||
):
|
):
|
||||||
ends = vals
|
ends = vals
|
||||||
except (TypeError, StopIteration):
|
except (NameError, TypeError, StopIteration):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
if destinations is not None:
|
if destinations is not None:
|
||||||
@@ -821,13 +823,19 @@ class CompiledStateGraph(CompiledGraph):
|
|||||||
input_values = {k: k for k in self.builder.schemas[input_schema]}
|
input_values = {k: k for k in self.builder.schemas[input_schema]}
|
||||||
is_single_input = len(input_values) == 1 and "__root__" in input_values
|
is_single_input = len(input_values) == 1 and "__root__" in input_values
|
||||||
|
|
||||||
|
branch_channel = f"branch:to:{key}"
|
||||||
self.channels[key] = EphemeralValue(Any, guard=False)
|
self.channels[key] = EphemeralValue(Any, guard=False)
|
||||||
|
self.channels[branch_channel] = EphemeralValue(Any, guard=False)
|
||||||
self.nodes[key] = PregelNode(
|
self.nodes[key] = PregelNode(
|
||||||
triggers=[],
|
triggers=[branch_channel],
|
||||||
# read state keys and managed values
|
# read state keys and managed values
|
||||||
channels=(list(input_values) if is_single_input else input_values),
|
channels=(list(input_values) if is_single_input else input_values),
|
||||||
# coerce state dict to schema class (eg. pydantic model)
|
# coerce state dict to schema class (eg. pydantic model)
|
||||||
mapper=_pick_mapper(list(input_values), input_schema),
|
mapper=_pick_mapper(
|
||||||
|
list(input_values),
|
||||||
|
input_schema,
|
||||||
|
self.builder.type_hints[input_schema],
|
||||||
|
),
|
||||||
writers=[
|
writers=[
|
||||||
# publish to this channel and state keys
|
# publish to this channel and state keys
|
||||||
ChannelWrite(
|
ChannelWrite(
|
||||||
@@ -851,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
|
||||||
@@ -865,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(
|
||||||
@@ -878,7 +890,7 @@ class CompiledStateGraph(CompiledGraph):
|
|||||||
if filtered := [p for p in packets if p != END]:
|
if filtered := [p for p in packets if p != END]:
|
||||||
writes = [
|
writes = [
|
||||||
(
|
(
|
||||||
ChannelWriteEntry(f"branch:{start}:{name}:{p}", start)
|
ChannelWriteEntry(f"branch:to:{p}", start)
|
||||||
if not isinstance(p, Send)
|
if not isinstance(p, Send)
|
||||||
else p
|
else p
|
||||||
)
|
)
|
||||||
@@ -902,33 +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,
|
||||||
|
)
|
||||||
)
|
)
|
||||||
for end in ends:
|
|
||||||
if end != END:
|
|
||||||
channel_name = f"branch:{start}:{name}:{end}"
|
|
||||||
self.channels[channel_name] = EphemeralValue(Any, guard=False)
|
|
||||||
self.nodes[end].triggers.append(channel_name)
|
|
||||||
|
|
||||||
# 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]
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -942,12 +952,12 @@ def _get_state_reader(
|
|||||||
select=select[0] if select == ["__root__"] else select,
|
select=select[0] if select == ["__root__"] else select,
|
||||||
fresh=True,
|
fresh=True,
|
||||||
# coerce state dict to schema class (eg. pydantic model)
|
# coerce state dict to schema class (eg. pydantic model)
|
||||||
mapper=_pick_mapper(state_keys, schema),
|
mapper=_pick_mapper(state_keys, schema, builder.type_hints[schema]),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _pick_mapper(
|
def _pick_mapper(
|
||||||
state_keys: Sequence[str], schema: Type[Any]
|
state_keys: Sequence[str], schema: Type[Any], type_hints: Optional[dict[str, Any]]
|
||||||
) -> Optional[Callable[[Any], Any]]:
|
) -> Optional[Callable[[Any], Any]]:
|
||||||
if state_keys == ["__root__"]:
|
if state_keys == ["__root__"]:
|
||||||
return None
|
return None
|
||||||
@@ -955,7 +965,7 @@ def _pick_mapper(
|
|||||||
if issubclass(schema, dict):
|
if issubclass(schema, dict):
|
||||||
return None
|
return None
|
||||||
if issubclass(schema, (BaseModel, BaseModelV1)):
|
if issubclass(schema, (BaseModel, BaseModelV1)):
|
||||||
return SchemaCoercionMapper(schema)
|
return SchemaCoercionMapper(schema, type_hints)
|
||||||
return partial(_coerce_state, schema)
|
return partial(_coerce_state, schema)
|
||||||
|
|
||||||
|
|
||||||
@@ -1010,25 +1020,36 @@ 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)
|
||||||
|
|
||||||
|
|
||||||
def _get_channels(
|
def _get_channels(
|
||||||
schema: Type[dict],
|
schema: Type[dict],
|
||||||
) -> tuple[dict[str, BaseChannel], dict[str, ManagedValueSpec]]:
|
) -> tuple[dict[str, BaseChannel], dict[str, ManagedValueSpec], dict[str, Any]]:
|
||||||
if not hasattr(schema, "__annotations__"):
|
if not hasattr(schema, "__annotations__"):
|
||||||
return {"__root__": _get_channel("__root__", schema, allow_managed=False)}, {}
|
return (
|
||||||
|
{"__root__": _get_channel("__root__", schema, allow_managed=False)},
|
||||||
|
{},
|
||||||
|
{},
|
||||||
|
)
|
||||||
|
|
||||||
|
type_hints = get_type_hints(schema, include_extras=True)
|
||||||
all_keys = {
|
all_keys = {
|
||||||
name: _get_channel(name, typ)
|
name: _get_channel(name, typ)
|
||||||
for name, typ in get_type_hints(schema, include_extras=True).items()
|
for name, typ in type_hints.items()
|
||||||
if name != "__slots__"
|
if name != "__slots__"
|
||||||
}
|
}
|
||||||
return (
|
return (
|
||||||
{k: v for k, v in all_keys.items() if isinstance(v, BaseChannel)},
|
{k: v for k, v in all_keys.items() if isinstance(v, BaseChannel)},
|
||||||
{k: v for k, v in all_keys.items() if is_managed_value(v)},
|
{k: v for k, v in all_keys.items() if is_managed_value(v)},
|
||||||
|
type_hints,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,4 @@
|
|||||||
import functools
|
import binascii
|
||||||
import itertools
|
import itertools
|
||||||
import sys
|
import sys
|
||||||
from collections import defaultdict, deque
|
from collections import defaultdict, deque
|
||||||
@@ -19,7 +19,6 @@ from typing import (
|
|||||||
cast,
|
cast,
|
||||||
overload,
|
overload,
|
||||||
)
|
)
|
||||||
from uuid import UUID
|
|
||||||
|
|
||||||
from langchain_core.callbacks import Callbacks
|
from langchain_core.callbacks import Callbacks
|
||||||
from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager
|
from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager
|
||||||
@@ -28,6 +27,7 @@ from langchain_core.runnables.config import RunnableConfig
|
|||||||
from langgraph.channels.base import BaseChannel
|
from langgraph.channels.base import BaseChannel
|
||||||
from langgraph.checkpoint.base import (
|
from langgraph.checkpoint.base import (
|
||||||
BaseCheckpointSaver,
|
BaseCheckpointSaver,
|
||||||
|
ChannelVersions,
|
||||||
Checkpoint,
|
Checkpoint,
|
||||||
PendingWrite,
|
PendingWrite,
|
||||||
V,
|
V,
|
||||||
@@ -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,37 @@ 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("-", ""))
|
||||||
|
null_version = checkpoint_null_version(checkpoint)
|
||||||
tasks: list[Union[PregelTask, PregelExecutableTask]] = []
|
tasks: list[Union[PregelTask, PregelExecutableTask]] = []
|
||||||
# Consume pending_sends from previous step
|
# Consume pending_sends from previous step
|
||||||
for idx, _ in enumerate(checkpoint["pending_sends"]):
|
for idx, _ in enumerate(checkpoint["pending_sends"]):
|
||||||
@@ -381,6 +422,8 @@ def prepare_next_tasks(
|
|||||||
(PUSH, idx),
|
(PUSH, idx),
|
||||||
None,
|
None,
|
||||||
checkpoint=checkpoint,
|
checkpoint=checkpoint,
|
||||||
|
checkpoint_id_bytes=checkpoint_id_bytes,
|
||||||
|
checkpoint_null_version=null_version,
|
||||||
pending_writes=pending_writes,
|
pending_writes=pending_writes,
|
||||||
processes=processes,
|
processes=processes,
|
||||||
channels=channels,
|
channels=channels,
|
||||||
@@ -393,13 +436,36 @@ 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,
|
||||||
checkpoint=checkpoint,
|
checkpoint=checkpoint,
|
||||||
|
checkpoint_id_bytes=checkpoint_id_bytes,
|
||||||
|
checkpoint_null_version=null_version,
|
||||||
pending_writes=pending_writes,
|
pending_writes=pending_writes,
|
||||||
processes=processes,
|
processes=processes,
|
||||||
channels=channels,
|
channels=channels,
|
||||||
@@ -415,11 +481,16 @@ def prepare_next_tasks(
|
|||||||
return {t.id: t for t in tasks}
|
return {t.id: t for t in tasks}
|
||||||
|
|
||||||
|
|
||||||
|
PUSH_TRIGGER = (PUSH,)
|
||||||
|
|
||||||
|
|
||||||
def prepare_single_task(
|
def prepare_single_task(
|
||||||
task_path: tuple[Any, ...],
|
task_path: tuple[Any, ...],
|
||||||
task_id_checksum: Optional[str],
|
task_id_checksum: Optional[str],
|
||||||
*,
|
*,
|
||||||
checkpoint: Checkpoint,
|
checkpoint: Checkpoint,
|
||||||
|
checkpoint_id_bytes: bytes,
|
||||||
|
checkpoint_null_version: Optional[V],
|
||||||
pending_writes: list[PendingWrite],
|
pending_writes: list[PendingWrite],
|
||||||
processes: Mapping[str, PregelNode],
|
processes: Mapping[str, PregelNode],
|
||||||
channels: Mapping[str, BaseChannel],
|
channels: Mapping[str, BaseChannel],
|
||||||
@@ -433,7 +504,6 @@ def prepare_single_task(
|
|||||||
) -> Union[None, PregelTask, PregelExecutableTask]:
|
) -> Union[None, PregelTask, PregelExecutableTask]:
|
||||||
"""Prepares a single task for the next Pregel step, given a task path, which
|
"""Prepares a single task for the next Pregel step, given a task path, which
|
||||||
uniquely identifies a PUSH or PULL task within the graph."""
|
uniquely identifies a PUSH or PULL task within the graph."""
|
||||||
checkpoint_id = UUID(checkpoint["id"]).bytes
|
|
||||||
configurable = config.get(CONF, {})
|
configurable = config.get(CONF, {})
|
||||||
parent_ns = configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")
|
parent_ns = configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||||
|
|
||||||
@@ -446,10 +516,10 @@ def prepare_single_task(
|
|||||||
if name is None:
|
if name is None:
|
||||||
raise ValueError("`call` functions must have a `__name__` attribute")
|
raise ValueError("`call` functions must have a `__name__` attribute")
|
||||||
# create task id
|
# create task id
|
||||||
triggers = [PUSH]
|
triggers: Sequence[str] = PUSH_TRIGGER
|
||||||
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
||||||
task_id = _uuid5_str(
|
task_id = _uuid5_str(
|
||||||
checkpoint_id,
|
checkpoint_id_bytes,
|
||||||
checkpoint_ns,
|
checkpoint_ns,
|
||||||
str(step),
|
str(step),
|
||||||
name,
|
name,
|
||||||
@@ -507,6 +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[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||||
pending_writes,
|
pending_writes,
|
||||||
task_id,
|
task_id,
|
||||||
),
|
),
|
||||||
@@ -539,12 +610,12 @@ def prepare_single_task(
|
|||||||
)
|
)
|
||||||
return
|
return
|
||||||
# create task id
|
# create task id
|
||||||
triggers = [PUSH]
|
triggers = PUSH_TRIGGER
|
||||||
checkpoint_ns = (
|
checkpoint_ns = (
|
||||||
f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node
|
f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node
|
||||||
)
|
)
|
||||||
task_id = _uuid5_str(
|
task_id = _uuid5_str(
|
||||||
checkpoint_id,
|
checkpoint_id_bytes,
|
||||||
checkpoint_ns,
|
checkpoint_ns,
|
||||||
str(step),
|
str(step),
|
||||||
packet.node,
|
packet.node,
|
||||||
@@ -616,6 +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[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||||
pending_writes,
|
pending_writes,
|
||||||
task_id,
|
task_id,
|
||||||
),
|
),
|
||||||
@@ -640,21 +712,17 @@ def prepare_single_task(
|
|||||||
if name not in processes:
|
if name not in processes:
|
||||||
return
|
return
|
||||||
proc = processes[name]
|
proc = processes[name]
|
||||||
version_type = type(next(iter(checkpoint["channel_versions"].values()), None))
|
if checkpoint_null_version is None:
|
||||||
null_version = version_type() # type: ignore[misc]
|
|
||||||
if null_version is None:
|
|
||||||
return
|
return
|
||||||
seen = checkpoint["versions_seen"].get(name, {})
|
|
||||||
# If any of the channels read by this process were updated
|
# If any of the channels read by this process were updated
|
||||||
if triggers := sorted(
|
if _triggers(
|
||||||
chan
|
channels,
|
||||||
for chan in proc.triggers
|
checkpoint["channel_versions"],
|
||||||
if not isinstance(
|
checkpoint["versions_seen"].get(name),
|
||||||
read_channel(channels, chan, return_exception=True), EmptyChannelError
|
checkpoint_null_version,
|
||||||
)
|
proc,
|
||||||
and checkpoint["channel_versions"].get(chan, null_version) # type: ignore[operator]
|
|
||||||
> seen.get(chan, null_version)
|
|
||||||
):
|
):
|
||||||
|
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)
|
||||||
@@ -671,7 +739,7 @@ def prepare_single_task(
|
|||||||
# create task id
|
# create task id
|
||||||
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
||||||
task_id = _uuid5_str(
|
task_id = _uuid5_str(
|
||||||
checkpoint_id,
|
checkpoint_id_bytes,
|
||||||
checkpoint_ns,
|
checkpoint_ns,
|
||||||
str(step),
|
str(step),
|
||||||
name,
|
name,
|
||||||
@@ -714,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,
|
||||||
@@ -723,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,
|
||||||
),
|
),
|
||||||
@@ -741,6 +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[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||||
pending_writes,
|
pending_writes,
|
||||||
task_id,
|
task_id,
|
||||||
),
|
),
|
||||||
@@ -761,13 +833,59 @@ def prepare_single_task(
|
|||||||
return PregelTask(task_id, name, task_path[:3])
|
return PregelTask(task_id, name, task_path[:3])
|
||||||
|
|
||||||
|
|
||||||
|
def checkpoint_null_version(
|
||||||
|
checkpoint: Checkpoint,
|
||||||
|
) -> Optional[V]:
|
||||||
|
"""Get the null version for the checkpoint, if available."""
|
||||||
|
for version in checkpoint["channel_versions"].values():
|
||||||
|
return type(version)()
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _triggers(
|
||||||
|
channels: Mapping[str, BaseChannel],
|
||||||
|
versions: ChannelVersions,
|
||||||
|
seen: Optional[ChannelVersions],
|
||||||
|
null_version: V,
|
||||||
|
proc: PregelNode,
|
||||||
|
) -> Sequence[str]:
|
||||||
|
if seen is None:
|
||||||
|
for chan in proc.triggers:
|
||||||
|
if channels[chan].is_available():
|
||||||
|
return (chan,)
|
||||||
|
else:
|
||||||
|
for chan in proc.triggers:
|
||||||
|
if channels[chan].is_available() and versions.get( # type: ignore[operator]
|
||||||
|
chan, null_version
|
||||||
|
) > seen.get(chan, null_version):
|
||||||
|
return (chan,)
|
||||||
|
return EMPTY_SEQ
|
||||||
|
|
||||||
|
|
||||||
def _scratchpad(
|
def _scratchpad(
|
||||||
|
parent_scratchpad: Optional[PregelScratchpad],
|
||||||
pending_writes: list[PendingWrite],
|
pending_writes: list[PendingWrite],
|
||||||
task_id: str,
|
task_id: str,
|
||||||
) -> PregelScratchpad:
|
) -> PregelScratchpad:
|
||||||
|
# None cannot be used as a resume value, because it would be difficult to
|
||||||
|
# distinguish from missing when used over http
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def get_null_resume(consume: bool = False) -> Any:
|
||||||
|
if null_resume_write is None:
|
||||||
|
if parent_scratchpad is not None:
|
||||||
|
return parent_scratchpad.get_null_resume(consume)
|
||||||
|
return None
|
||||||
|
if consume:
|
||||||
|
try:
|
||||||
|
pending_writes.remove(null_resume_write)
|
||||||
|
return null_resume_write[2]
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
return null_resume_write[2]
|
||||||
|
|
||||||
# using itertools.count as an atomic counter (+= 1 is not thread-safe)
|
# using itertools.count as an atomic counter (+= 1 is not thread-safe)
|
||||||
return PregelScratchpad(
|
return PregelScratchpad(
|
||||||
# call
|
# call
|
||||||
@@ -777,10 +895,7 @@ def _scratchpad(
|
|||||||
resume=next(
|
resume=next(
|
||||||
(w[2] for w in pending_writes if w[0] == task_id and w[1] == RESUME), []
|
(w[2] for w in pending_writes if w[0] == task_id and w[1] == RESUME), []
|
||||||
),
|
),
|
||||||
null_resume=null_resume_write[2] if null_resume_write is not None else None,
|
get_null_resume=get_null_resume,
|
||||||
_consume_null_resume=functools.partial(pending_writes.remove, null_resume_write)
|
|
||||||
if null_resume_write is not None
|
|
||||||
else lambda: None,
|
|
||||||
# subgraph
|
# subgraph
|
||||||
subgraph_counter=itertools.count(0).__next__,
|
subgraph_counter=itertools.count(0).__next__,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -137,7 +137,12 @@ def map_debug_task_results(
|
|||||||
"result": [
|
"result": [
|
||||||
w for w in writes if w[0] in stream_channels_list or w[0] == RETURN
|
w for w in writes if w[0] in stream_channels_list or w[0] == RETURN
|
||||||
],
|
],
|
||||||
"interrupts": [asdict(w[1]) for w in writes if w[0] == INTERRUPT],
|
"interrupts": [
|
||||||
|
asdict(v)
|
||||||
|
for w in writes
|
||||||
|
if w[0] == INTERRUPT
|
||||||
|
for v in (w[1] if isinstance(w[1], Sequence) else [w[1]])
|
||||||
|
],
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -293,8 +298,9 @@ def tasks_w_writes(
|
|||||||
),
|
),
|
||||||
tuple(
|
tuple(
|
||||||
v
|
v
|
||||||
for tid, n, v in pending_writes
|
for tid, n, vv in pending_writes
|
||||||
if tid == task.id and n == INTERRUPT
|
if tid == task.id and n == INTERRUPT
|
||||||
|
for v in (vv if isinstance(vv, Sequence) else [vv])
|
||||||
),
|
),
|
||||||
states.get(task.id) if states else None,
|
states.get(task.id) if states else None,
|
||||||
(
|
(
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ from langgraph.constants import (
|
|||||||
NULL_TASK_ID,
|
NULL_TASK_ID,
|
||||||
RESUME,
|
RESUME,
|
||||||
RETURN,
|
RETURN,
|
||||||
SELF,
|
|
||||||
START,
|
START,
|
||||||
TAG_HIDDEN,
|
TAG_HIDDEN,
|
||||||
TASKS,
|
TASKS,
|
||||||
@@ -28,7 +27,7 @@ def is_task_id(task_id: str) -> bool:
|
|||||||
"""Check if a string is a valid task id."""
|
"""Check if a string is a valid task id."""
|
||||||
try:
|
try:
|
||||||
UUID(task_id)
|
UUID(task_id)
|
||||||
except ValueError:
|
except Exception:
|
||||||
return False
|
return False
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -38,14 +37,11 @@ def read_channel(
|
|||||||
chan: str,
|
chan: str,
|
||||||
*,
|
*,
|
||||||
catch: bool = True,
|
catch: bool = True,
|
||||||
return_exception: bool = False,
|
|
||||||
) -> Any:
|
) -> Any:
|
||||||
try:
|
try:
|
||||||
return channels[chan].get()
|
return channels[chan].get()
|
||||||
except EmptyChannelError as exc:
|
except EmptyChannelError:
|
||||||
if return_exception:
|
if catch:
|
||||||
return exc
|
|
||||||
elif catch:
|
|
||||||
return None
|
return None
|
||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
@@ -84,7 +80,7 @@ def map_command(
|
|||||||
if isinstance(send, Send):
|
if isinstance(send, Send):
|
||||||
yield (NULL_TASK_ID, TASKS, send)
|
yield (NULL_TASK_ID, TASKS, send)
|
||||||
elif isinstance(send, str):
|
elif isinstance(send, str):
|
||||||
yield (NULL_TASK_ID, f"branch:{START}:{SELF}:{send}", START)
|
yield (NULL_TASK_ID, f"branch:to:{send}", START)
|
||||||
else:
|
else:
|
||||||
raise TypeError(
|
raise TypeError(
|
||||||
f"In Command.goto, expected Send/str, got {type(send).__name__}"
|
f"In Command.goto, expected Send/str, got {type(send).__name__}"
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
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
|
||||||
@@ -79,6 +81,7 @@ from langgraph.pregel.algo import (
|
|||||||
GetNextVersion,
|
GetNextVersion,
|
||||||
PregelTaskWrites,
|
PregelTaskWrites,
|
||||||
apply_writes,
|
apply_writes,
|
||||||
|
checkpoint_null_version,
|
||||||
increment,
|
increment,
|
||||||
prepare_next_tasks,
|
prepare_next_tasks,
|
||||||
prepare_single_task,
|
prepare_single_task,
|
||||||
@@ -207,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,
|
||||||
@@ -230,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])
|
||||||
@@ -264,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)
|
||||||
@@ -347,12 +352,16 @@ class PregelLoop(LoopProtocol):
|
|||||||
):
|
):
|
||||||
self.to_interrupt.append(task)
|
self.to_interrupt.append(task)
|
||||||
return
|
return
|
||||||
|
checkpoint_id_bytes = binascii.unhexlify(self.checkpoint["id"].replace("-", ""))
|
||||||
|
null_version = checkpoint_null_version(self.checkpoint)
|
||||||
if pushed := cast(
|
if pushed := cast(
|
||||||
Optional[PregelExecutableTask],
|
Optional[PregelExecutableTask],
|
||||||
prepare_single_task(
|
prepare_single_task(
|
||||||
(PUSH, task.path, write_idx, task.id, call),
|
(PUSH, task.path, write_idx, task.id, call),
|
||||||
None,
|
None,
|
||||||
checkpoint=self.checkpoint,
|
checkpoint=self.checkpoint,
|
||||||
|
checkpoint_id_bytes=checkpoint_id_bytes,
|
||||||
|
checkpoint_null_version=null_version,
|
||||||
pending_writes=self.checkpoint_pending_writes,
|
pending_writes=self.checkpoint_pending_writes,
|
||||||
processes=self.nodes,
|
processes=self.nodes,
|
||||||
channels=self.channels,
|
channels=self.channels,
|
||||||
@@ -400,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"
|
||||||
@@ -421,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(),
|
||||||
@@ -487,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 = []
|
||||||
|
|
||||||
@@ -565,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)
|
||||||
@@ -586,16 +599,9 @@ class PregelLoop(LoopProtocol):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
# this can be set only when there are input_writes
|
||||||
|
updated_channels: Optional[set[str]] = None
|
||||||
|
|
||||||
# take resume value from parent
|
|
||||||
if scratchpad := cast(
|
|
||||||
Optional[PregelScratchpad], configurable.get(CONFIG_KEY_SCRATCHPAD)
|
|
||||||
):
|
|
||||||
if (
|
|
||||||
isinstance(scratchpad, PregelScratchpad)
|
|
||||||
and scratchpad.null_resume is not None
|
|
||||||
):
|
|
||||||
self.put_writes(NULL_TASK_ID, [(RESUME, scratchpad.null_resume)])
|
|
||||||
# map command to writes
|
# map command to writes
|
||||||
if isinstance(self.input, Command):
|
if isinstance(self.input, Command):
|
||||||
if self.input.resume is not None and not self.checkpointer:
|
if self.input.resume is not None and not self.checkpointer:
|
||||||
@@ -615,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, [])],
|
||||||
@@ -664,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,
|
||||||
[
|
[
|
||||||
@@ -694,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():
|
||||||
@@ -779,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(),
|
||||||
@@ -794,11 +801,14 @@ class PregelLoop(LoopProtocol):
|
|||||||
[w for t in self.tasks.values() for w in t.writes],
|
[w for t in self.tasks.values() for w in t.writes],
|
||||||
self.channels,
|
self.channels,
|
||||||
)
|
)
|
||||||
# emit INTERRUPT event
|
# emit INTERRUPT if exception is empty (otherwise emitted by put_writes)
|
||||||
self._emit(
|
if exc_value is not None and (not exc_value.args or not exc_value.args[0]):
|
||||||
"updates",
|
self._emit(
|
||||||
lambda: iter([{INTERRUPT: cast(GraphInterrupt, exc_value).args[0]}]),
|
"updates",
|
||||||
)
|
lambda: iter(
|
||||||
|
[{INTERRUPT: cast(GraphInterrupt, exc_value).args[0]}]
|
||||||
|
),
|
||||||
|
)
|
||||||
# save final output
|
# save final output
|
||||||
self.output = read_channels(self.channels, self.output_keys)
|
self.output = read_channels(self.channels, self.output_keys)
|
||||||
# suppress interrupt
|
# suppress interrupt
|
||||||
@@ -829,7 +839,25 @@ class PregelLoop(LoopProtocol):
|
|||||||
"tags", EMPTY_SEQ
|
"tags", EMPTY_SEQ
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
if writes[0][0] != ERROR and writes[0][0] != INTERRUPT:
|
if writes[0][0] == INTERRUPT:
|
||||||
|
self._emit(
|
||||||
|
"updates",
|
||||||
|
lambda: iter(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
INTERRUPT: tuple(
|
||||||
|
v
|
||||||
|
for w in writes
|
||||||
|
if w[0] == INTERRUPT
|
||||||
|
for v in (
|
||||||
|
w[1] if isinstance(w[1], Sequence) else (w[1],)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
}
|
||||||
|
]
|
||||||
|
),
|
||||||
|
)
|
||||||
|
elif writes[0][0] != ERROR:
|
||||||
self._emit(
|
self._emit(
|
||||||
"updates",
|
"updates",
|
||||||
map_output_updates,
|
map_output_updates,
|
||||||
@@ -865,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,
|
||||||
@@ -881,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:
|
||||||
@@ -1006,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,
|
||||||
@@ -1022,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:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import asyncio
|
|||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
import weakref
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from typing import (
|
from typing import (
|
||||||
Any,
|
Any,
|
||||||
@@ -25,12 +26,10 @@ from langgraph.constants import (
|
|||||||
CONF,
|
CONF,
|
||||||
CONFIG_KEY_CALL,
|
CONFIG_KEY_CALL,
|
||||||
CONFIG_KEY_SCRATCHPAD,
|
CONFIG_KEY_SCRATCHPAD,
|
||||||
CONFIG_KEY_SEND,
|
|
||||||
ERROR,
|
ERROR,
|
||||||
INTERRUPT,
|
INTERRUPT,
|
||||||
MISSING,
|
MISSING,
|
||||||
NO_WRITES,
|
NO_WRITES,
|
||||||
PUSH,
|
|
||||||
RESUME,
|
RESUME,
|
||||||
RETURN,
|
RETURN,
|
||||||
TAG_HIDDEN,
|
TAG_HIDDEN,
|
||||||
@@ -48,7 +47,9 @@ E = TypeVar("E", threading.Event, asyncio.Event)
|
|||||||
|
|
||||||
class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]):
|
class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]):
|
||||||
event: E
|
event: E
|
||||||
callback: Callable[[PregelExecutableTask, Optional[BaseException]], None]
|
callback: weakref.ref[
|
||||||
|
Callable[[PregelExecutableTask, Optional[BaseException]], None]
|
||||||
|
]
|
||||||
counter: int
|
counter: int
|
||||||
done: set[F]
|
done: set[F]
|
||||||
lock: threading.Lock
|
lock: threading.Lock
|
||||||
@@ -56,7 +57,9 @@ class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]):
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
event: E,
|
event: E,
|
||||||
callback: Callable[[PregelExecutableTask, Optional[BaseException]], None],
|
callback: weakref.ref[
|
||||||
|
Callable[[PregelExecutableTask, Optional[BaseException]], None]
|
||||||
|
],
|
||||||
future_type: Type[F],
|
future_type: Type[F],
|
||||||
# used for generic typing, newer py supports FutureDict[...](...)
|
# used for generic typing, newer py supports FutureDict[...](...)
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -85,7 +88,7 @@ class FuturesDict(Generic[F, E], dict[F, Optional[PregelExecutableTask]]):
|
|||||||
fut: F,
|
fut: F,
|
||||||
) -> None:
|
) -> None:
|
||||||
try:
|
try:
|
||||||
self.callback(task, _exception(fut))
|
self.callback()(task, _exception(fut)) # type: ignore[misc]
|
||||||
finally:
|
finally:
|
||||||
with self.lock:
|
with self.lock:
|
||||||
self.done.add(fut)
|
self.done.add(fut)
|
||||||
@@ -102,10 +105,13 @@ class PregelRunner:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
submit: Submit,
|
submit: weakref.ref[Submit],
|
||||||
put_writes: Callable[[str, Sequence[tuple[str, Any]]], None],
|
put_writes: weakref.ref[Callable[[str, Sequence[tuple[str, Any]]], None]],
|
||||||
schedule_task: Callable[
|
schedule_task: weakref.ref[
|
||||||
[PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask]
|
Callable[
|
||||||
|
[PregelExecutableTask, int, Optional[Call]],
|
||||||
|
Optional[PregelExecutableTask],
|
||||||
|
]
|
||||||
],
|
],
|
||||||
use_astream: bool = False,
|
use_astream: bool = False,
|
||||||
node_finished: Optional[Callable[[str], None]] = None,
|
node_finished: Optional[Callable[[str], None]] = None,
|
||||||
@@ -125,99 +131,9 @@ class PregelRunner:
|
|||||||
retry_policy: Optional[RetryPolicy] = None,
|
retry_policy: Optional[RetryPolicy] = None,
|
||||||
get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None,
|
get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None,
|
||||||
) -> Iterator[None]:
|
) -> Iterator[None]:
|
||||||
def writer(
|
|
||||||
task: PregelExecutableTask,
|
|
||||||
writes: Sequence[tuple[str, Any]],
|
|
||||||
*,
|
|
||||||
calls: Optional[Sequence[Call]] = None,
|
|
||||||
) -> Sequence[Optional[concurrent.futures.Future]]:
|
|
||||||
if all(w[0] != PUSH for w in writes):
|
|
||||||
return task.config[CONF][CONFIG_KEY_SEND](writes)
|
|
||||||
|
|
||||||
# schedule PUSH tasks, collect futures
|
|
||||||
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
|
|
||||||
rtn: dict[int, Optional[concurrent.futures.Future]] = {}
|
|
||||||
for idx, w in enumerate(writes):
|
|
||||||
# bail if not a PUSH write
|
|
||||||
if w[0] != PUSH:
|
|
||||||
continue
|
|
||||||
# schedule the next task, if the callback returns one
|
|
||||||
wcall = calls[idx] if calls else None
|
|
||||||
if next_task := self.schedule_task(
|
|
||||||
task, scratchpad.call_counter(), wcall
|
|
||||||
):
|
|
||||||
if fut := next(
|
|
||||||
(
|
|
||||||
f
|
|
||||||
for f, t in futures.items()
|
|
||||||
if t is not None and t == next_task.id
|
|
||||||
),
|
|
||||||
None,
|
|
||||||
):
|
|
||||||
# if the parent task was retried,
|
|
||||||
# the next task might already be running
|
|
||||||
rtn[idx] = fut
|
|
||||||
elif next_task.writes:
|
|
||||||
# if it already ran, return the result
|
|
||||||
fut = concurrent.futures.Future()
|
|
||||||
ret = next(
|
|
||||||
(v for c, v in next_task.writes if c == RETURN), MISSING
|
|
||||||
)
|
|
||||||
if ret is not MISSING:
|
|
||||||
fut.set_result(ret)
|
|
||||||
elif exc := next(
|
|
||||||
(v for c, v in next_task.writes if c == ERROR), None
|
|
||||||
):
|
|
||||||
fut.set_exception(
|
|
||||||
exc
|
|
||||||
if isinstance(exc, BaseException)
|
|
||||||
else Exception(exc)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
fut.set_result(None)
|
|
||||||
rtn[idx] = fut
|
|
||||||
else:
|
|
||||||
# schedule the next task
|
|
||||||
fut = self.submit(
|
|
||||||
run_with_retry,
|
|
||||||
next_task,
|
|
||||||
retry_policy,
|
|
||||||
configurable={
|
|
||||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
|
||||||
CONFIG_KEY_CALL: partial(call, next_task),
|
|
||||||
},
|
|
||||||
__reraise_on_exit__=reraise,
|
|
||||||
# starting a new task in the next tick ensures
|
|
||||||
# updates from this tick are committed/streamed first
|
|
||||||
__next_tick__=True,
|
|
||||||
)
|
|
||||||
futures[fut] = next_task
|
|
||||||
rtn[idx] = fut
|
|
||||||
return [rtn.get(i) for i in range(len(writes))]
|
|
||||||
|
|
||||||
def call(
|
|
||||||
task: PregelExecutableTask,
|
|
||||||
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
|
||||||
input: Any,
|
|
||||||
*,
|
|
||||||
retry: Optional[RetryPolicy] = None,
|
|
||||||
callbacks: Callbacks = None,
|
|
||||||
) -> concurrent.futures.Future[Any]:
|
|
||||||
if asyncio.iscoroutinefunction(func):
|
|
||||||
raise RuntimeError("In an sync context async tasks cannot be called")
|
|
||||||
(fut,) = writer(
|
|
||||||
task,
|
|
||||||
[(PUSH, None)],
|
|
||||||
calls=[Call(func, input, retry=retry, callbacks=callbacks)],
|
|
||||||
)
|
|
||||||
assert fut is not None, "writer did not return a future for call"
|
|
||||||
# return a chained future to ensure commit() callback is called
|
|
||||||
# before the returned future is resolved, to ensure stream order etc
|
|
||||||
return chain_future(fut, concurrent.futures.Future())
|
|
||||||
|
|
||||||
tasks = tuple(tasks)
|
tasks = tuple(tasks)
|
||||||
futures = FuturesDict(
|
futures = FuturesDict(
|
||||||
callback=self.commit,
|
callback=weakref.WeakMethod(self.commit),
|
||||||
event=threading.Event(),
|
event=threading.Event(),
|
||||||
future_type=concurrent.futures.Future,
|
future_type=concurrent.futures.Future,
|
||||||
)
|
)
|
||||||
@@ -231,8 +147,15 @@ class PregelRunner:
|
|||||||
t,
|
t,
|
||||||
retry_policy,
|
retry_policy,
|
||||||
configurable={
|
configurable={
|
||||||
CONFIG_KEY_SEND: partial(writer, t),
|
CONFIG_KEY_CALL: partial(
|
||||||
CONFIG_KEY_CALL: partial(call, t),
|
_call,
|
||||||
|
weakref.ref(t),
|
||||||
|
retry=retry_policy,
|
||||||
|
futures=weakref.ref(futures),
|
||||||
|
schedule_task=self.schedule_task,
|
||||||
|
submit=self.submit,
|
||||||
|
reraise=reraise,
|
||||||
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
self.commit(t, None)
|
self.commit(t, None)
|
||||||
@@ -255,13 +178,20 @@ class PregelRunner:
|
|||||||
# schedule tasks
|
# schedule tasks
|
||||||
for t in tasks:
|
for t in tasks:
|
||||||
if not t.writes:
|
if not t.writes:
|
||||||
fut = self.submit(
|
fut = self.submit()( # type: ignore[misc]
|
||||||
run_with_retry,
|
run_with_retry,
|
||||||
t,
|
t,
|
||||||
retry_policy,
|
retry_policy,
|
||||||
configurable={
|
configurable={
|
||||||
CONFIG_KEY_SEND: partial(writer, t),
|
CONFIG_KEY_CALL: partial(
|
||||||
CONFIG_KEY_CALL: partial(call, t),
|
_call,
|
||||||
|
weakref.ref(t),
|
||||||
|
retry=retry_policy,
|
||||||
|
futures=weakref.ref(futures),
|
||||||
|
schedule_task=self.schedule_task,
|
||||||
|
submit=self.submit,
|
||||||
|
reraise=reraise,
|
||||||
|
),
|
||||||
},
|
},
|
||||||
__reraise_on_exit__=reraise,
|
__reraise_on_exit__=reraise,
|
||||||
)
|
)
|
||||||
@@ -313,125 +243,10 @@ class PregelRunner:
|
|||||||
retry_policy: Optional[RetryPolicy] = None,
|
retry_policy: Optional[RetryPolicy] = None,
|
||||||
get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None,
|
get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None,
|
||||||
) -> AsyncIterator[None]:
|
) -> AsyncIterator[None]:
|
||||||
def writer(
|
|
||||||
task: PregelExecutableTask,
|
|
||||||
writes: Sequence[tuple[str, Any]],
|
|
||||||
*,
|
|
||||||
calls: Optional[Sequence[Call]] = None,
|
|
||||||
) -> Sequence[Optional[asyncio.Future]]:
|
|
||||||
if all(w[0] != PUSH for w in writes):
|
|
||||||
return task.config[CONF][CONFIG_KEY_SEND](writes)
|
|
||||||
|
|
||||||
# schedule PUSH tasks, collect futures
|
|
||||||
scratchpad: PregelScratchpad = task.config[CONF][CONFIG_KEY_SCRATCHPAD]
|
|
||||||
rtn: dict[int, Optional[asyncio.Future]] = {}
|
|
||||||
for idx, w in enumerate(writes):
|
|
||||||
# bail if not a PUSH write
|
|
||||||
if w[0] != PUSH:
|
|
||||||
continue
|
|
||||||
# schedule the next task, if the callback returns one
|
|
||||||
wcall = calls[idx] if calls is not None else None
|
|
||||||
if next_task := self.schedule_task(
|
|
||||||
task, scratchpad.call_counter(), wcall
|
|
||||||
):
|
|
||||||
# if the parent task was retried,
|
|
||||||
# the next task might already be running
|
|
||||||
if fut := next(
|
|
||||||
(
|
|
||||||
f
|
|
||||||
for f, t in futures.items()
|
|
||||||
if t is not None and t == next_task.id
|
|
||||||
),
|
|
||||||
None,
|
|
||||||
):
|
|
||||||
# if the parent task was retried,
|
|
||||||
# the next task might already be running
|
|
||||||
rtn[idx] = fut
|
|
||||||
elif next_task.writes:
|
|
||||||
# if it already ran, return the result
|
|
||||||
fut = asyncio.Future(loop=loop)
|
|
||||||
ret = next(
|
|
||||||
(v for c, v in next_task.writes if c == RETURN), MISSING
|
|
||||||
)
|
|
||||||
if ret is not MISSING:
|
|
||||||
fut.set_result(ret)
|
|
||||||
elif exc := next(
|
|
||||||
(v for c, v in next_task.writes if c == ERROR), None
|
|
||||||
):
|
|
||||||
fut.set_exception(
|
|
||||||
exc
|
|
||||||
if isinstance(exc, BaseException)
|
|
||||||
else Exception(exc)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
fut.set_result(None)
|
|
||||||
rtn[idx] = fut
|
|
||||||
else:
|
|
||||||
# schedule the next task
|
|
||||||
fut = cast(
|
|
||||||
asyncio.Future,
|
|
||||||
self.submit(
|
|
||||||
arun_with_retry,
|
|
||||||
next_task,
|
|
||||||
retry_policy,
|
|
||||||
stream=self.use_astream,
|
|
||||||
configurable={
|
|
||||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
|
||||||
CONFIG_KEY_CALL: partial(call, next_task),
|
|
||||||
},
|
|
||||||
__name__=t.name,
|
|
||||||
__cancel_on_exit__=True,
|
|
||||||
__reraise_on_exit__=reraise,
|
|
||||||
# starting a new task in the next tick ensures
|
|
||||||
# updates from this tick are committed/streamed first
|
|
||||||
__next_tick__=True,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
futures[fut] = next_task
|
|
||||||
rtn[idx] = fut
|
|
||||||
return [rtn.get(i) for i in range(len(writes))]
|
|
||||||
|
|
||||||
def call(
|
|
||||||
task: PregelExecutableTask,
|
|
||||||
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
|
||||||
input: Any,
|
|
||||||
*,
|
|
||||||
retry: Optional[RetryPolicy] = None,
|
|
||||||
callbacks: Callbacks = None,
|
|
||||||
) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]:
|
|
||||||
(fut,) = writer(
|
|
||||||
task,
|
|
||||||
[(PUSH, None)],
|
|
||||||
calls=[Call(func, input, retry=retry, callbacks=callbacks)],
|
|
||||||
)
|
|
||||||
assert fut is not None, "writer did not return a future for call"
|
|
||||||
# return a chained future to ensure commit() callback is called
|
|
||||||
# before the returned future is resolved, to ensure stream order etc
|
|
||||||
try:
|
|
||||||
in_async = asyncio.current_task() is not None
|
|
||||||
except RuntimeError:
|
|
||||||
in_async = False
|
|
||||||
# if in async context return an async future
|
|
||||||
# otherwise return a chained sync future
|
|
||||||
if in_async:
|
|
||||||
if isinstance(fut, asyncio.Task):
|
|
||||||
sfut: Union[asyncio.Future[Any], concurrent.futures.Future[Any]] = (
|
|
||||||
asyncio.Future(loop=loop)
|
|
||||||
)
|
|
||||||
loop.call_soon_threadsafe(chain_future, fut, sfut)
|
|
||||||
return sfut
|
|
||||||
else:
|
|
||||||
# already wrapped in a future
|
|
||||||
return fut
|
|
||||||
else:
|
|
||||||
sfut = concurrent.futures.Future()
|
|
||||||
loop.call_soon_threadsafe(chain_future, fut, sfut)
|
|
||||||
return sfut
|
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
tasks = tuple(tasks)
|
tasks = tuple(tasks)
|
||||||
futures = FuturesDict(
|
futures = FuturesDict(
|
||||||
callback=self.commit,
|
callback=weakref.WeakMethod(self.commit),
|
||||||
event=asyncio.Event(),
|
event=asyncio.Event(),
|
||||||
future_type=asyncio.Future,
|
future_type=asyncio.Future,
|
||||||
)
|
)
|
||||||
@@ -446,8 +261,17 @@ class PregelRunner:
|
|||||||
retry_policy,
|
retry_policy,
|
||||||
stream=self.use_astream,
|
stream=self.use_astream,
|
||||||
configurable={
|
configurable={
|
||||||
CONFIG_KEY_SEND: partial(writer, t),
|
CONFIG_KEY_CALL: partial(
|
||||||
CONFIG_KEY_CALL: partial(call, t),
|
_acall,
|
||||||
|
weakref.ref(t),
|
||||||
|
stream=self.use_astream,
|
||||||
|
retry=retry_policy,
|
||||||
|
futures=weakref.ref(futures),
|
||||||
|
schedule_task=self.schedule_task,
|
||||||
|
submit=self.submit,
|
||||||
|
reraise=reraise,
|
||||||
|
loop=loop,
|
||||||
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
self.commit(t, None)
|
self.commit(t, None)
|
||||||
@@ -472,14 +296,23 @@ class PregelRunner:
|
|||||||
if not t.writes:
|
if not t.writes:
|
||||||
fut = cast(
|
fut = cast(
|
||||||
asyncio.Future,
|
asyncio.Future,
|
||||||
self.submit(
|
self.submit()( # type: ignore[misc]
|
||||||
arun_with_retry,
|
arun_with_retry,
|
||||||
t,
|
t,
|
||||||
retry_policy,
|
retry_policy,
|
||||||
stream=self.use_astream,
|
stream=self.use_astream,
|
||||||
configurable={
|
configurable={
|
||||||
CONFIG_KEY_SEND: partial(writer, t),
|
CONFIG_KEY_CALL: partial(
|
||||||
CONFIG_KEY_CALL: partial(call, t),
|
_acall,
|
||||||
|
weakref.ref(t),
|
||||||
|
retry=retry_policy,
|
||||||
|
stream=self.use_astream,
|
||||||
|
futures=weakref.ref(futures),
|
||||||
|
schedule_task=self.schedule_task,
|
||||||
|
submit=self.submit,
|
||||||
|
reraise=reraise,
|
||||||
|
loop=loop,
|
||||||
|
),
|
||||||
},
|
},
|
||||||
__name__=t.name,
|
__name__=t.name,
|
||||||
__cancel_on_exit__=True,
|
__cancel_on_exit__=True,
|
||||||
@@ -539,19 +372,20 @@ class PregelRunner:
|
|||||||
# for cancelled tasks, also save error in task,
|
# for cancelled tasks, also save error in task,
|
||||||
# so loop can finish super-step
|
# so loop can finish super-step
|
||||||
task.writes.append((ERROR, exception))
|
task.writes.append((ERROR, exception))
|
||||||
self.put_writes(task.id, task.writes)
|
self.put_writes()(task.id, task.writes) # type: ignore[misc]
|
||||||
elif exception:
|
elif exception:
|
||||||
if isinstance(exception, GraphInterrupt):
|
if isinstance(exception, GraphInterrupt):
|
||||||
# save interrupt to checkpointer
|
# save interrupt to checkpointer
|
||||||
if interrupts := [(INTERRUPT, i) for i in exception.args[0]]:
|
if exception.args[0]:
|
||||||
|
writes = [(INTERRUPT, exception.args[0])]
|
||||||
if resumes := [w for w in task.writes if w[0] == RESUME]:
|
if resumes := [w for w in task.writes if w[0] == RESUME]:
|
||||||
interrupts.extend(resumes)
|
writes.extend(resumes)
|
||||||
self.put_writes(task.id, interrupts)
|
self.put_writes()(task.id, writes) # type: ignore[misc]
|
||||||
elif isinstance(exception, GraphBubbleUp):
|
elif isinstance(exception, GraphBubbleUp):
|
||||||
raise exception
|
raise exception
|
||||||
else:
|
else:
|
||||||
# save error to checkpointer
|
# save error to checkpointer
|
||||||
self.put_writes(task.id, [(ERROR, exception)])
|
self.put_writes()(task.id, [(ERROR, exception)]) # type: ignore[misc]
|
||||||
else:
|
else:
|
||||||
if self.node_finished and (
|
if self.node_finished and (
|
||||||
task.config is None or TAG_HIDDEN not in task.config.get("tags", [])
|
task.config is None or TAG_HIDDEN not in task.config.get("tags", [])
|
||||||
@@ -561,7 +395,7 @@ class PregelRunner:
|
|||||||
# add no writes marker
|
# add no writes marker
|
||||||
task.writes.append((NO_WRITES, None))
|
task.writes.append((NO_WRITES, None))
|
||||||
# save task writes to checkpointer
|
# save task writes to checkpointer
|
||||||
self.put_writes(task.id, task.writes)
|
self.put_writes()(task.id, task.writes) # type: ignore[misc]
|
||||||
|
|
||||||
|
|
||||||
def _should_stop_others(
|
def _should_stop_others(
|
||||||
@@ -608,6 +442,7 @@ def _panic_or_proceed(
|
|||||||
done.add(fut)
|
done.add(fut)
|
||||||
else:
|
else:
|
||||||
inflight.add(fut)
|
inflight.add(fut)
|
||||||
|
interrupts: list[GraphInterrupt] = []
|
||||||
while done:
|
while done:
|
||||||
# if any task failed
|
# if any task failed
|
||||||
if exc := _exception(done.pop()):
|
if exc := _exception(done.pop()):
|
||||||
@@ -616,7 +451,14 @@ def _panic_or_proceed(
|
|||||||
inflight.pop().cancel()
|
inflight.pop().cancel()
|
||||||
# raise the exception
|
# raise the exception
|
||||||
if panic:
|
if panic:
|
||||||
raise exc
|
if isinstance(exc, GraphInterrupt):
|
||||||
|
# collect interrupts
|
||||||
|
interrupts.append(exc)
|
||||||
|
else:
|
||||||
|
raise exc
|
||||||
|
# raise combined interrupts
|
||||||
|
if interrupts:
|
||||||
|
raise GraphInterrupt(tuple(i for exc in interrupts for i in exc.args[0]))
|
||||||
if inflight:
|
if inflight:
|
||||||
# if we got here means we timed out
|
# if we got here means we timed out
|
||||||
while inflight:
|
while inflight:
|
||||||
@@ -624,3 +466,193 @@ def _panic_or_proceed(
|
|||||||
inflight.pop().cancel()
|
inflight.pop().cancel()
|
||||||
# raise timeout error
|
# raise timeout error
|
||||||
raise timeout_exc_cls("Timed out")
|
raise timeout_exc_cls("Timed out")
|
||||||
|
|
||||||
|
|
||||||
|
def _call(
|
||||||
|
task: weakref.ref[PregelExecutableTask],
|
||||||
|
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
||||||
|
input: Any,
|
||||||
|
*,
|
||||||
|
retry: Optional[RetryPolicy] = None,
|
||||||
|
callbacks: Callbacks = None,
|
||||||
|
futures: weakref.ref[FuturesDict],
|
||||||
|
schedule_task: weakref.ref[
|
||||||
|
Callable[
|
||||||
|
[PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask]
|
||||||
|
]
|
||||||
|
],
|
||||||
|
submit: weakref.ref[Submit],
|
||||||
|
reraise: bool,
|
||||||
|
) -> concurrent.futures.Future[Any]:
|
||||||
|
if asyncio.iscoroutinefunction(func):
|
||||||
|
raise RuntimeError("In an sync context async tasks cannot be called")
|
||||||
|
|
||||||
|
fut: Optional[concurrent.futures.Future] = None
|
||||||
|
# schedule PUSH tasks, collect futures
|
||||||
|
scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr]
|
||||||
|
# schedule the next task, if the callback returns one
|
||||||
|
if next_task := schedule_task()( # type: ignore[misc]
|
||||||
|
task(), # type: ignore[arg-type]
|
||||||
|
scratchpad.call_counter(),
|
||||||
|
Call(func, input, retry=retry, callbacks=callbacks),
|
||||||
|
):
|
||||||
|
if fut := next(
|
||||||
|
(
|
||||||
|
f
|
||||||
|
for f, t in futures().items() # type: ignore[union-attr]
|
||||||
|
if t is not None and t == next_task.id
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
):
|
||||||
|
# if the parent task was retried,
|
||||||
|
# the next task might already be running
|
||||||
|
pass
|
||||||
|
elif next_task.writes:
|
||||||
|
# if it already ran, return the result
|
||||||
|
fut = concurrent.futures.Future()
|
||||||
|
ret = next((v for c, v in next_task.writes if c == RETURN), MISSING)
|
||||||
|
if ret is not MISSING:
|
||||||
|
fut.set_result(ret)
|
||||||
|
elif exc := next((v for c, v in next_task.writes if c == ERROR), None):
|
||||||
|
fut.set_exception(
|
||||||
|
exc if isinstance(exc, BaseException) else Exception(exc)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
fut.set_result(None)
|
||||||
|
else:
|
||||||
|
# schedule the next task
|
||||||
|
fut = submit()( # type: ignore[misc]
|
||||||
|
run_with_retry,
|
||||||
|
next_task,
|
||||||
|
retry,
|
||||||
|
configurable={
|
||||||
|
CONFIG_KEY_CALL: partial(
|
||||||
|
_call,
|
||||||
|
weakref.ref(next_task),
|
||||||
|
futures=futures,
|
||||||
|
retry=retry,
|
||||||
|
callbacks=callbacks,
|
||||||
|
schedule_task=schedule_task,
|
||||||
|
submit=submit,
|
||||||
|
reraise=reraise,
|
||||||
|
),
|
||||||
|
},
|
||||||
|
__reraise_on_exit__=reraise,
|
||||||
|
# starting a new task in the next tick ensures
|
||||||
|
# updates from this tick are committed/streamed first
|
||||||
|
__next_tick__=True,
|
||||||
|
)
|
||||||
|
futures()[fut] = next_task # type: ignore[index]
|
||||||
|
fut = cast(Union[asyncio.Future, concurrent.futures.Future], fut)
|
||||||
|
# return a chained future to ensure commit() callback is called
|
||||||
|
# before the returned future is resolved, to ensure stream order etc
|
||||||
|
return chain_future(fut, concurrent.futures.Future())
|
||||||
|
|
||||||
|
|
||||||
|
def _acall(
|
||||||
|
task: weakref.ref[PregelExecutableTask],
|
||||||
|
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
||||||
|
input: Any,
|
||||||
|
*,
|
||||||
|
retry: Optional[RetryPolicy] = None,
|
||||||
|
callbacks: Callbacks = None,
|
||||||
|
# injected dependencies
|
||||||
|
futures: weakref.ref[FuturesDict],
|
||||||
|
schedule_task: weakref.ref[
|
||||||
|
Callable[
|
||||||
|
[PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask]
|
||||||
|
]
|
||||||
|
],
|
||||||
|
submit: weakref.ref[Submit],
|
||||||
|
loop: asyncio.AbstractEventLoop,
|
||||||
|
reraise: bool = False,
|
||||||
|
stream: bool = False,
|
||||||
|
) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]:
|
||||||
|
fut: Optional[asyncio.Future] = None
|
||||||
|
# schedule PUSH tasks, collect futures
|
||||||
|
scratchpad: PregelScratchpad = task().config[CONF][CONFIG_KEY_SCRATCHPAD] # type: ignore[union-attr]
|
||||||
|
# schedule the next task, if the callback returns one
|
||||||
|
if next_task := schedule_task()( # type: ignore[misc]
|
||||||
|
task(), # type: ignore[arg-type]
|
||||||
|
scratchpad.call_counter(),
|
||||||
|
Call(func, input, retry=retry, callbacks=callbacks),
|
||||||
|
):
|
||||||
|
if fut := next(
|
||||||
|
(
|
||||||
|
f
|
||||||
|
for f, t in futures().items() # type: ignore[union-attr]
|
||||||
|
if t is not None and t == next_task.id
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
):
|
||||||
|
# if the parent task was retried,
|
||||||
|
# the next task might already be running
|
||||||
|
|
||||||
|
pass
|
||||||
|
elif next_task.writes:
|
||||||
|
# if it already ran, return the result
|
||||||
|
fut = asyncio.Future(loop=loop)
|
||||||
|
ret = next((v for c, v in next_task.writes if c == RETURN), MISSING)
|
||||||
|
if ret is not MISSING:
|
||||||
|
fut.set_result(ret)
|
||||||
|
elif exc := next((v for c, v in next_task.writes if c == ERROR), None):
|
||||||
|
fut.set_exception(
|
||||||
|
exc if isinstance(exc, BaseException) else Exception(exc)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
fut.set_result(None)
|
||||||
|
futures()[fut] = next_task # type: ignore[index]
|
||||||
|
else:
|
||||||
|
# schedule the next task
|
||||||
|
fut = cast(
|
||||||
|
asyncio.Future,
|
||||||
|
submit()( # type: ignore[misc]
|
||||||
|
arun_with_retry,
|
||||||
|
next_task,
|
||||||
|
retry,
|
||||||
|
stream=stream,
|
||||||
|
configurable={
|
||||||
|
CONFIG_KEY_CALL: partial(
|
||||||
|
_acall,
|
||||||
|
weakref.ref(next_task),
|
||||||
|
stream=stream,
|
||||||
|
futures=futures,
|
||||||
|
schedule_task=schedule_task,
|
||||||
|
submit=submit,
|
||||||
|
loop=loop,
|
||||||
|
reraise=reraise,
|
||||||
|
),
|
||||||
|
},
|
||||||
|
__name__=task().name, # type: ignore[union-attr]
|
||||||
|
__cancel_on_exit__=True,
|
||||||
|
__reraise_on_exit__=reraise,
|
||||||
|
# starting a new task in the next tick ensures
|
||||||
|
# updates from this tick are committed/streamed first
|
||||||
|
__next_tick__=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
futures()[fut] = next_task # type: ignore[index]
|
||||||
|
|
||||||
|
fut = cast(Union[asyncio.Future, concurrent.futures.Future], fut)
|
||||||
|
# return a chained future to ensure commit() callback is called
|
||||||
|
# before the returned future is resolved, to ensure stream order etc
|
||||||
|
try:
|
||||||
|
in_async = asyncio.current_task() is not None
|
||||||
|
except RuntimeError:
|
||||||
|
in_async = False
|
||||||
|
# if in async context return an async future
|
||||||
|
# otherwise return a chained sync future
|
||||||
|
if in_async:
|
||||||
|
if isinstance(fut, asyncio.Task):
|
||||||
|
sfut: Union[asyncio.Future[Any], concurrent.futures.Future[Any]] = (
|
||||||
|
asyncio.Future(loop=loop)
|
||||||
|
)
|
||||||
|
loop.call_soon_threadsafe(chain_future, fut, sfut)
|
||||||
|
return sfut
|
||||||
|
else:
|
||||||
|
# already wrapped in a future
|
||||||
|
return fut
|
||||||
|
else:
|
||||||
|
sfut = concurrent.futures.Future()
|
||||||
|
loop.call_soon_threadsafe(chain_future, fut, sfut)
|
||||||
|
return sfut
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -1,7 +1,10 @@
|
|||||||
from typing import Optional
|
import ast
|
||||||
|
import inspect
|
||||||
|
import textwrap
|
||||||
|
from typing import Any, Callable, Optional
|
||||||
|
|
||||||
from langchain_core.runnables import RunnableLambda, RunnableSequence
|
from langchain_core.runnables import RunnableLambda, RunnableSequence
|
||||||
from langchain_core.runnables.utils import get_function_nonlocals
|
from typing_extensions import override
|
||||||
|
|
||||||
from langgraph.checkpoint.base import ChannelVersions
|
from langgraph.checkpoint.base import ChannelVersions
|
||||||
from langgraph.pregel.protocol import PregelProtocol
|
from langgraph.pregel.protocol import PregelProtocol
|
||||||
@@ -55,3 +58,152 @@ def find_subgraph_pregel(candidate: Runnable) -> Optional[PregelProtocol]:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def get_function_nonlocals(func: Callable) -> list[Any]:
|
||||||
|
"""Get the nonlocal variables accessed by a function.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
func: The function to check.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List[Any]: The nonlocal variables accessed by the function.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
code = inspect.getsource(func)
|
||||||
|
tree = ast.parse(textwrap.dedent(code))
|
||||||
|
visitor = FunctionNonLocals()
|
||||||
|
visitor.visit(tree)
|
||||||
|
values: list[Any] = []
|
||||||
|
closure = (
|
||||||
|
inspect.getclosurevars(func.__wrapped__)
|
||||||
|
if hasattr(func, "__wrapped__") and callable(func.__wrapped__)
|
||||||
|
else inspect.getclosurevars(func)
|
||||||
|
)
|
||||||
|
candidates = {**closure.globals, **closure.nonlocals}
|
||||||
|
for k, v in candidates.items():
|
||||||
|
if k in visitor.nonlocals:
|
||||||
|
values.append(v)
|
||||||
|
for kk in visitor.nonlocals:
|
||||||
|
if "." in kk and kk.startswith(k):
|
||||||
|
vv = v
|
||||||
|
for part in kk.split(".")[1:]:
|
||||||
|
if vv is None:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
vv = getattr(vv, part)
|
||||||
|
except AttributeError:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
values.append(vv)
|
||||||
|
except (SyntaxError, TypeError, OSError, SystemError):
|
||||||
|
return []
|
||||||
|
|
||||||
|
return values
|
||||||
|
|
||||||
|
|
||||||
|
class FunctionNonLocals(ast.NodeVisitor):
|
||||||
|
"""Get the nonlocal variables accessed of a function."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.nonlocals: set[str] = set()
|
||||||
|
|
||||||
|
@override
|
||||||
|
def visit_FunctionDef(self, node: ast.FunctionDef) -> Any:
|
||||||
|
"""Visit a function definition.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node: The node to visit.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Any: The result of the visit.
|
||||||
|
"""
|
||||||
|
visitor = NonLocals()
|
||||||
|
visitor.visit(node)
|
||||||
|
self.nonlocals.update(visitor.loads - visitor.stores)
|
||||||
|
|
||||||
|
@override
|
||||||
|
def visit_AsyncFunctionDef(self, node: ast.AsyncFunctionDef) -> Any:
|
||||||
|
"""Visit an async function definition.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node: The node to visit.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Any: The result of the visit.
|
||||||
|
"""
|
||||||
|
visitor = NonLocals()
|
||||||
|
visitor.visit(node)
|
||||||
|
self.nonlocals.update(visitor.loads - visitor.stores)
|
||||||
|
|
||||||
|
@override
|
||||||
|
def visit_Lambda(self, node: ast.Lambda) -> Any:
|
||||||
|
"""Visit a lambda function.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node: The node to visit.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Any: The result of the visit.
|
||||||
|
"""
|
||||||
|
visitor = NonLocals()
|
||||||
|
visitor.visit(node)
|
||||||
|
self.nonlocals.update(visitor.loads - visitor.stores)
|
||||||
|
|
||||||
|
|
||||||
|
class NonLocals(ast.NodeVisitor):
|
||||||
|
"""Get nonlocal variables accessed."""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self.loads: set[str] = set()
|
||||||
|
self.stores: set[str] = set()
|
||||||
|
|
||||||
|
@override
|
||||||
|
def visit_Name(self, node: ast.Name) -> Any:
|
||||||
|
"""Visit a name node.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node: The node to visit.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Any: The result of the visit.
|
||||||
|
"""
|
||||||
|
if isinstance(node.ctx, ast.Load):
|
||||||
|
self.loads.add(node.id)
|
||||||
|
elif isinstance(node.ctx, ast.Store):
|
||||||
|
self.stores.add(node.id)
|
||||||
|
|
||||||
|
@override
|
||||||
|
def visit_Attribute(self, node: ast.Attribute) -> Any:
|
||||||
|
"""Visit an attribute node.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
node: The node to visit.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Any: The result of the visit.
|
||||||
|
"""
|
||||||
|
if isinstance(node.ctx, ast.Load):
|
||||||
|
parent = node.value
|
||||||
|
attr_expr = node.attr
|
||||||
|
while isinstance(parent, ast.Attribute):
|
||||||
|
attr_expr = parent.attr + "." + attr_expr
|
||||||
|
parent = parent.value
|
||||||
|
if isinstance(parent, ast.Name):
|
||||||
|
self.loads.add(parent.id + "." + attr_expr)
|
||||||
|
self.loads.discard(parent.id)
|
||||||
|
elif isinstance(parent, ast.Call):
|
||||||
|
if isinstance(parent.func, ast.Name):
|
||||||
|
self.loads.add(parent.func.id)
|
||||||
|
else:
|
||||||
|
parent = parent.func
|
||||||
|
attr_expr = ""
|
||||||
|
while isinstance(parent, ast.Attribute):
|
||||||
|
if attr_expr:
|
||||||
|
attr_expr = parent.attr + "." + attr_expr
|
||||||
|
else:
|
||||||
|
attr_expr = parent.attr
|
||||||
|
parent = parent.value
|
||||||
|
if isinstance(parent, ast.Name):
|
||||||
|
self.loads.add(parent.id + "." + attr_expr)
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -130,7 +130,12 @@ class Interrupt:
|
|||||||
value: Any
|
value: Any
|
||||||
resumable: bool = False
|
resumable: bool = False
|
||||||
ns: Optional[Sequence[str]] = None
|
ns: Optional[Sequence[str]] = None
|
||||||
when: Literal["during"] = "during"
|
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):
|
||||||
@@ -143,13 +148,20 @@ 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
|
||||||
writes: deque[tuple[str, Any]]
|
writes: deque[tuple[str, Any]]
|
||||||
config: RunnableConfig
|
config: RunnableConfig
|
||||||
triggers: list[str]
|
triggers: Sequence[str]
|
||||||
retry_policy: Optional[RetryPolicy]
|
retry_policy: Optional[RetryPolicy]
|
||||||
cache_policy: Optional[CachePolicy]
|
cache_policy: Optional[CachePolicy]
|
||||||
id: str
|
id: str
|
||||||
@@ -351,20 +363,11 @@ class PregelScratchpad:
|
|||||||
call_counter: Callable[[], int]
|
call_counter: Callable[[], int]
|
||||||
# interrupt
|
# interrupt
|
||||||
interrupt_counter: Callable[[], int]
|
interrupt_counter: Callable[[], int]
|
||||||
|
get_null_resume: Callable[[bool], Any]
|
||||||
resume: list[Any]
|
resume: list[Any]
|
||||||
null_resume: Optional[Any]
|
|
||||||
_consume_null_resume: Callable[[], None]
|
|
||||||
# subgraph
|
# subgraph
|
||||||
subgraph_counter: Callable[[], int]
|
subgraph_counter: Callable[[], int]
|
||||||
|
|
||||||
def consume_null_resume(self) -> Any:
|
|
||||||
if self.null_resume is not None:
|
|
||||||
value = self.null_resume
|
|
||||||
self._consume_null_resume()
|
|
||||||
self.null_resume = None
|
|
||||||
return value
|
|
||||||
raise ValueError("No null resume to consume")
|
|
||||||
|
|
||||||
|
|
||||||
def interrupt(value: Any) -> Any:
|
def interrupt(value: Any) -> Any:
|
||||||
"""Interrupt the graph with a resumable exception from within a node.
|
"""Interrupt the graph with a resumable exception from within a node.
|
||||||
@@ -480,9 +483,9 @@ def interrupt(value: Any) -> Any:
|
|||||||
if idx < len(scratchpad.resume):
|
if idx < len(scratchpad.resume):
|
||||||
return scratchpad.resume[idx]
|
return scratchpad.resume[idx]
|
||||||
# find current resume value
|
# find current resume value
|
||||||
if scratchpad.null_resume is not None:
|
v = scratchpad.get_null_resume(True)
|
||||||
|
if v is not None:
|
||||||
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
|
assert len(scratchpad.resume) == idx, (scratchpad.resume, idx)
|
||||||
v = scratchpad.consume_null_resume()
|
|
||||||
scratchpad.resume.append(v)
|
scratchpad.resume.append(v)
|
||||||
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
conf[CONFIG_KEY_SEND]([(RESUME, scratchpad.resume)])
|
||||||
return v
|
return v
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Generated
+51
-9
@@ -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"
|
||||||
@@ -2238,6 +2238,48 @@ files = [
|
|||||||
{file = "pycparser-2.22.tar.gz", hash = "sha256:491c8be9c040f5390f5bf44a5b07752bd07f56edf992381b05c701439eec10f6"},
|
{file = "pycparser-2.22.tar.gz", hash = "sha256:491c8be9c040f5390f5bf44a5b07752bd07f56edf992381b05c701439eec10f6"},
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "pycryptodome"
|
||||||
|
version = "3.21.0"
|
||||||
|
description = "Cryptographic library for Python"
|
||||||
|
optional = false
|
||||||
|
python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,>=2.7"
|
||||||
|
groups = ["dev"]
|
||||||
|
files = [
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-macosx_10_9_x86_64.whl", hash = "sha256:dad9bf36eda068e89059d1f07408e397856be9511d7113ea4b586642a429a4fd"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-manylinux2010_i686.whl", hash = "sha256:a1752eca64c60852f38bb29e2c86fca30d7672c024128ef5d70cc15868fa10f4"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-manylinux2010_x86_64.whl", hash = "sha256:3ba4cc304eac4d4d458f508d4955a88ba25026890e8abff9b60404f76a62c55e"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7cb087b8612c8a1a14cf37dd754685be9a8d9869bed2ffaaceb04850a8aeef7e"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-musllinux_1_1_aarch64.whl", hash = "sha256:26412b21df30b2861424a6c6d5b1d8ca8107612a4cfa4d0183e71c5d200fb34a"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-win32.whl", hash = "sha256:cc2269ab4bce40b027b49663d61d816903a4bd90ad88cb99ed561aadb3888dd3"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27m-win_amd64.whl", hash = "sha256:0fa0a05a6a697ccbf2a12cec3d6d2650b50881899b845fac6e87416f8cb7e87d"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27mu-manylinux2010_i686.whl", hash = "sha256:6cce52e196a5f1d6797ff7946cdff2038d3b5f0aba4a43cb6bf46b575fd1b5bb"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27mu-manylinux2010_x86_64.whl", hash = "sha256:a915597ffccabe902e7090e199a7bf7a381c5506a747d5e9d27ba55197a2c568"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27mu-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a4e74c522d630766b03a836c15bff77cb657c5fdf098abf8b1ada2aebc7d0819"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp27-cp27mu-musllinux_1_1_aarch64.whl", hash = "sha256:a3804675283f4764a02db05f5191eb8fec2bb6ca34d466167fc78a5f05bbe6b3"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-macosx_10_9_universal2.whl", hash = "sha256:2480ec2c72438430da9f601ebc12c518c093c13111a5c1644c82cdfc2e50b1e4"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-macosx_10_9_x86_64.whl", hash = "sha256:de18954104667f565e2fbb4783b56667f30fb49c4d79b346f52a29cb198d5b6b"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:2de4b7263a33947ff440412339cb72b28a5a4c769b5c1ca19e33dd6cd1dcec6e"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0714206d467fc911042d01ea3a1847c847bc10884cf674c82e12915cfe1649f8"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:7d85c1b613121ed3dbaa5a97369b3b757909531a959d229406a75b912dd51dd1"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-musllinux_1_1_aarch64.whl", hash = "sha256:8898a66425a57bcf15e25fc19c12490b87bd939800f39a03ea2de2aea5e3611a"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-musllinux_1_2_i686.whl", hash = "sha256:932c905b71a56474bff8a9c014030bc3c882cee696b448af920399f730a650c2"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:18caa8cfbc676eaaf28613637a89980ad2fd96e00c564135bf90bc3f0b34dd93"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-win32.whl", hash = "sha256:280b67d20e33bb63171d55b1067f61fbd932e0b1ad976b3a184303a3dad22764"},
|
||||||
|
{file = "pycryptodome-3.21.0-cp36-abi3-win_amd64.whl", hash = "sha256:b7aa25fc0baa5b1d95b7633af4f5f1838467f1815442b22487426f94e0d66c53"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp27-pypy_73-manylinux2010_x86_64.whl", hash = "sha256:2cb635b67011bc147c257e61ce864879ffe6d03342dc74b6045059dfbdedafca"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp27-pypy_73-win32.whl", hash = "sha256:4c26a2f0dc15f81ea3afa3b0c87b87e501f235d332b7f27e2225ecb80c0b1cdd"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp310-pypy310_pp73-macosx_10_15_x86_64.whl", hash = "sha256:d5ebe0763c982f069d3877832254f64974139f4f9655058452603ff559c482e8"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp310-pypy310_pp73-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:7ee86cbde706be13f2dec5a42b52b1c1d1cbb90c8e405c68d0755134735c8dc6"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp310-pypy310_pp73-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:0fd54003ec3ce4e0f16c484a10bc5d8b9bd77fa662a12b85779a2d2d85d67ee0"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp310-pypy310_pp73-win_amd64.whl", hash = "sha256:5dfafca172933506773482b0e18f0cd766fd3920bd03ec85a283df90d8a17bc6"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp39-pypy39_pp73-macosx_10_15_x86_64.whl", hash = "sha256:590ef0898a4b0a15485b05210b4a1c9de8806d3ad3d47f74ab1dc07c67a6827f"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp39-pypy39_pp73-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f35e442630bc4bc2e1878482d6f59ea22e280d7121d7adeaedba58c23ab6386b"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp39-pypy39_pp73-manylinux_2_5_i686.manylinux1_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:ff99f952db3db2fbe98a0b355175f93ec334ba3d01bbde25ad3a5a33abc02b58"},
|
||||||
|
{file = "pycryptodome-3.21.0-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:8acd7d34af70ee63f9a849f957558e49a98f8f1634f86a59d2be62bb8e93f71c"},
|
||||||
|
{file = "pycryptodome-3.21.0.tar.gz", hash = "sha256:f7787e0d469bdae763b876174cf2e6c0f7be79808af26b1da96f1a64bcf47297"},
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "pydantic"
|
name = "pydantic"
|
||||||
version = "2.9.2"
|
version = "2.9.2"
|
||||||
@@ -3509,4 +3551,4 @@ type = ["pytest-mypy"]
|
|||||||
[metadata]
|
[metadata]
|
||||||
lock-version = "2.1"
|
lock-version = "2.1"
|
||||||
python-versions = ">=3.9.0,<4.0"
|
python-versions = ">=3.9.0,<4.0"
|
||||||
content-hash = "eb85f0bcc0e8a715ef38afb58cf888f7c2ee8579ea6ed94900244365f24cddd9"
|
content-hash = "b8641a0b2d92bee0363602e69f99b23366b2035b7e17ff017708194e6fbd0ac5"
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
[tool.poetry]
|
[tool.poetry]
|
||||||
name = "langgraph"
|
name = "langgraph"
|
||||||
version = "0.3.10"
|
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"
|
||||||
@@ -37,6 +37,7 @@ uvloop = "0.21.0beta1"
|
|||||||
pyperf = "^2.7.0"
|
pyperf = "^2.7.0"
|
||||||
py-spy = "^0.3.14"
|
py-spy = "^0.3.14"
|
||||||
types-requests = "^2.32.0.20240914"
|
types-requests = "^2.32.0.20240914"
|
||||||
|
pycryptodome = "^3.21.0"
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
lint.select = [ "E", "F", "I", "TID251" ]
|
lint.select = [ "E", "F", "I", "TID251" ]
|
||||||
|
|||||||
File diff suppressed because one or more lines are too long
@@ -377,6 +377,19 @@
|
|||||||
|
|
||||||
'''
|
'''
|
||||||
# ---
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge[sqlite_aes]
|
||||||
|
'''
|
||||||
|
graph TD;
|
||||||
|
__start__ --> rewrite_query;
|
||||||
|
analyzer_one --> retriever_one;
|
||||||
|
qa --> __end__;
|
||||||
|
retriever_one --> qa;
|
||||||
|
retriever_two --> qa;
|
||||||
|
rewrite_query --> analyzer_one;
|
||||||
|
rewrite_query --> retriever_two;
|
||||||
|
|
||||||
|
'''
|
||||||
|
# ---
|
||||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[memory]
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[memory]
|
||||||
'''
|
'''
|
||||||
graph TD;
|
graph TD;
|
||||||
@@ -797,6 +810,76 @@
|
|||||||
'type': 'object',
|
'type': 'object',
|
||||||
})
|
})
|
||||||
# ---
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[sqlite_aes]
|
||||||
|
'''
|
||||||
|
graph TD;
|
||||||
|
__start__ --> rewrite_query;
|
||||||
|
analyzer_one --> retriever_one;
|
||||||
|
qa --> __end__;
|
||||||
|
retriever_one --> qa;
|
||||||
|
retriever_two --> qa;
|
||||||
|
rewrite_query --> analyzer_one;
|
||||||
|
rewrite_query -.-> retriever_two;
|
||||||
|
|
||||||
|
'''
|
||||||
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[sqlite_aes].1
|
||||||
|
dict({
|
||||||
|
'definitions': dict({
|
||||||
|
'InnerObject': dict({
|
||||||
|
'properties': dict({
|
||||||
|
'yo': dict({
|
||||||
|
'title': 'Yo',
|
||||||
|
'type': 'integer',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'yo',
|
||||||
|
]),
|
||||||
|
'title': 'InnerObject',
|
||||||
|
'type': 'object',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'properties': dict({
|
||||||
|
'inner': dict({
|
||||||
|
'$ref': '#/definitions/InnerObject',
|
||||||
|
}),
|
||||||
|
'query': dict({
|
||||||
|
'title': 'Query',
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'query',
|
||||||
|
'inner',
|
||||||
|
]),
|
||||||
|
'title': 'Input',
|
||||||
|
'type': 'object',
|
||||||
|
})
|
||||||
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1[sqlite_aes].2
|
||||||
|
dict({
|
||||||
|
'properties': dict({
|
||||||
|
'answer': dict({
|
||||||
|
'title': 'Answer',
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
'docs': dict({
|
||||||
|
'items': dict({
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
'title': 'Docs',
|
||||||
|
'type': 'array',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'answer',
|
||||||
|
'docs',
|
||||||
|
]),
|
||||||
|
'title': 'Output',
|
||||||
|
'type': 'object',
|
||||||
|
})
|
||||||
|
# ---
|
||||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[memory]
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[memory]
|
||||||
'''
|
'''
|
||||||
graph TD;
|
graph TD;
|
||||||
@@ -1217,6 +1300,76 @@
|
|||||||
'type': 'object',
|
'type': 'object',
|
||||||
})
|
})
|
||||||
# ---
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[sqlite_aes]
|
||||||
|
'''
|
||||||
|
graph TD;
|
||||||
|
__start__ --> rewrite_query;
|
||||||
|
analyzer_one --> retriever_one;
|
||||||
|
qa --> __end__;
|
||||||
|
retriever_one --> qa;
|
||||||
|
retriever_two --> qa;
|
||||||
|
rewrite_query --> analyzer_one;
|
||||||
|
rewrite_query -.-> retriever_two;
|
||||||
|
|
||||||
|
'''
|
||||||
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[sqlite_aes].1
|
||||||
|
dict({
|
||||||
|
'$defs': dict({
|
||||||
|
'InnerObject': dict({
|
||||||
|
'properties': dict({
|
||||||
|
'yo': dict({
|
||||||
|
'title': 'Yo',
|
||||||
|
'type': 'integer',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'yo',
|
||||||
|
]),
|
||||||
|
'title': 'InnerObject',
|
||||||
|
'type': 'object',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'properties': dict({
|
||||||
|
'inner': dict({
|
||||||
|
'$ref': '#/$defs/InnerObject',
|
||||||
|
}),
|
||||||
|
'query': dict({
|
||||||
|
'title': 'Query',
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'query',
|
||||||
|
'inner',
|
||||||
|
]),
|
||||||
|
'title': 'Input',
|
||||||
|
'type': 'object',
|
||||||
|
})
|
||||||
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2[sqlite_aes].2
|
||||||
|
dict({
|
||||||
|
'properties': dict({
|
||||||
|
'answer': dict({
|
||||||
|
'title': 'Answer',
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
'docs': dict({
|
||||||
|
'items': dict({
|
||||||
|
'type': 'string',
|
||||||
|
}),
|
||||||
|
'title': 'Docs',
|
||||||
|
'type': 'array',
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
'required': list([
|
||||||
|
'answer',
|
||||||
|
'docs',
|
||||||
|
]),
|
||||||
|
'title': 'Output',
|
||||||
|
'type': 'object',
|
||||||
|
})
|
||||||
|
# ---
|
||||||
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic_input[memory]
|
# name: test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic_input[memory]
|
||||||
'''
|
'''
|
||||||
graph TD;
|
graph TD;
|
||||||
@@ -1715,6 +1868,19 @@
|
|||||||
|
|
||||||
'''
|
'''
|
||||||
# ---
|
# ---
|
||||||
|
# name: test_in_one_fan_out_state_graph_waiting_edge_via_branch[sqlite_aes]
|
||||||
|
'''
|
||||||
|
graph TD;
|
||||||
|
__start__ --> rewrite_query;
|
||||||
|
analyzer_one --> retriever_one;
|
||||||
|
qa --> __end__;
|
||||||
|
retriever_one --> qa;
|
||||||
|
retriever_two --> qa;
|
||||||
|
rewrite_query --> analyzer_one;
|
||||||
|
rewrite_query -.-> retriever_two;
|
||||||
|
|
||||||
|
'''
|
||||||
|
# ---
|
||||||
# name: test_multiple_sinks_subgraphs
|
# name: test_multiple_sinks_subgraphs
|
||||||
'''
|
'''
|
||||||
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
%%{init: {'flowchart': {'curve': 'linear'}}}%%
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from langgraph.checkpoint.postgres.aio import (
|
|||||||
AsyncPostgresSaver,
|
AsyncPostgresSaver,
|
||||||
AsyncShallowPostgresSaver,
|
AsyncShallowPostgresSaver,
|
||||||
)
|
)
|
||||||
|
from langgraph.checkpoint.serde.encrypted import EncryptedSerializer
|
||||||
from langgraph.checkpoint.sqlite import SqliteSaver
|
from langgraph.checkpoint.sqlite import SqliteSaver
|
||||||
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
|
||||||
from langgraph.store.base import BaseStore
|
from langgraph.store.base import BaseStore
|
||||||
@@ -61,6 +62,15 @@ def checkpointer_sqlite():
|
|||||||
yield checkpointer
|
yield checkpointer
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(scope="function")
|
||||||
|
def checkpointer_sqlite_aes():
|
||||||
|
with SqliteSaver.from_conn_string(":memory:") as checkpointer:
|
||||||
|
checkpointer.serde = EncryptedSerializer.from_pycryptodome_aes(
|
||||||
|
key=b"1234567890123456"
|
||||||
|
)
|
||||||
|
yield checkpointer
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def _checkpointer_sqlite_aio():
|
async def _checkpointer_sqlite_aio():
|
||||||
async with AsyncSqliteSaver.from_conn_string(":memory:") as checkpointer:
|
async with AsyncSqliteSaver.from_conn_string(":memory:") as checkpointer:
|
||||||
@@ -437,6 +447,7 @@ REGULAR_CHECKPOINTERS_SYNC = [
|
|||||||
"postgres",
|
"postgres",
|
||||||
"postgres_pipe",
|
"postgres_pipe",
|
||||||
"postgres_pool",
|
"postgres_pool",
|
||||||
|
"sqlite_aes",
|
||||||
]
|
]
|
||||||
ALL_CHECKPOINTERS_SYNC = [
|
ALL_CHECKPOINTERS_SYNC = [
|
||||||
*REGULAR_CHECKPOINTERS_SYNC,
|
*REGULAR_CHECKPOINTERS_SYNC,
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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:"),
|
||||||
@@ -2500,7 +2500,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": (PULL, "tools"),
|
"langgraph_path": (PULL, "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -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:"),
|
||||||
@@ -2559,7 +2559,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 4,
|
"langgraph_step": 4,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": (PULL, "tools"),
|
"langgraph_path": (PULL, "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -2573,7 +2573,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 4,
|
"langgraph_step": 4,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": (PULL, "tools"),
|
"langgraph_path": (PULL, "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -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": ["retriever_one", "retriever_two"],
|
"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"),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -6706,7 +6715,7 @@ def test_branch_then(
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "tool_two_slow",
|
"name": "tool_two_slow",
|
||||||
"input": {"my_key": "value prepared", "market": "DE"},
|
"input": {"my_key": "value prepared", "market": "DE"},
|
||||||
"triggers": ["branch:prepare:condition:tool_two_slow"],
|
"triggers": ("branch:to:tool_two_slow",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -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,
|
||||||
@@ -10378,9 +10393,7 @@ def test_weather_subgraph(
|
|||||||
"langgraph_node": "weather_graph",
|
"langgraph_node": "weather_graph",
|
||||||
"langgraph_path": [PULL, "weather_graph"],
|
"langgraph_path": [PULL, "weather_graph"],
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_triggers": [
|
"langgraph_triggers": ["branch:to:weather_graph"],
|
||||||
"branch:router_node:route_after_prediction:weather_graph"
|
|
||||||
],
|
|
||||||
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
||||||
},
|
},
|
||||||
created_at=AnyStr(),
|
created_at=AnyStr(),
|
||||||
@@ -10492,9 +10505,7 @@ def test_weather_subgraph(
|
|||||||
"langgraph_node": "weather_graph",
|
"langgraph_node": "weather_graph",
|
||||||
"langgraph_path": [PULL, "weather_graph"],
|
"langgraph_path": [PULL, "weather_graph"],
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_triggers": [
|
"langgraph_triggers": ["branch:to:weather_graph"],
|
||||||
"branch:router_node:route_after_prediction:weather_graph"
|
|
||||||
],
|
|
||||||
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
||||||
},
|
},
|
||||||
created_at=AnyStr(),
|
created_at=AnyStr(),
|
||||||
|
|||||||
@@ -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:"),
|
||||||
@@ -2317,7 +2321,7 @@ async def test_prebuilt_tool_chat() -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": ("__pregel_pull", "tools"),
|
"langgraph_path": ("__pregel_pull", "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -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:"),
|
||||||
@@ -2376,7 +2384,7 @@ async def test_prebuilt_tool_chat() -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 4,
|
"langgraph_step": 4,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": ("__pregel_pull", "tools"),
|
"langgraph_path": ("__pregel_pull", "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -2390,7 +2398,7 @@ async def test_prebuilt_tool_chat() -> None:
|
|||||||
{
|
{
|
||||||
"langgraph_step": 4,
|
"langgraph_step": 4,
|
||||||
"langgraph_node": "tools",
|
"langgraph_node": "tools",
|
||||||
"langgraph_triggers": ["branch:agent:should_continue:tools"],
|
"langgraph_triggers": ("branch:to:tools",),
|
||||||
"langgraph_path": ("__pregel_pull", "tools"),
|
"langgraph_path": ("__pregel_pull", "tools"),
|
||||||
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
"langgraph_checkpoint_ns": AnyStr("tools:"),
|
||||||
},
|
},
|
||||||
@@ -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": ["retriever_one", "retriever_two"],
|
"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",
|
||||||
|
),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -4537,7 +4561,7 @@ async def test_branch_then(checkpointer_name: str) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "tool_two_slow",
|
"name": "tool_two_slow",
|
||||||
"input": {"my_key": "value prepared", "market": "DE"},
|
"input": {"my_key": "value prepared", "market": "DE"},
|
||||||
"triggers": ["branch:prepare:condition:tool_two_slow"],
|
"triggers": ("branch:to:tool_two_slow",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -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,
|
||||||
@@ -7231,9 +7279,7 @@ async def test_weather_subgraph(
|
|||||||
"langgraph_node": "weather_graph",
|
"langgraph_node": "weather_graph",
|
||||||
"langgraph_path": [PULL, "weather_graph"],
|
"langgraph_path": [PULL, "weather_graph"],
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_triggers": [
|
"langgraph_triggers": ["branch:to:weather_graph"],
|
||||||
"branch:router_node:route_after_prediction:weather_graph"
|
|
||||||
],
|
|
||||||
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
||||||
},
|
},
|
||||||
created_at=AnyStr(),
|
created_at=AnyStr(),
|
||||||
@@ -7347,9 +7393,7 @@ async def test_weather_subgraph(
|
|||||||
"langgraph_node": "weather_graph",
|
"langgraph_node": "weather_graph",
|
||||||
"langgraph_path": [PULL, "weather_graph"],
|
"langgraph_path": [PULL, "weather_graph"],
|
||||||
"langgraph_step": 2,
|
"langgraph_step": 2,
|
||||||
"langgraph_triggers": [
|
"langgraph_triggers": ["branch:to:weather_graph"],
|
||||||
"branch:router_node:route_after_prediction:weather_graph"
|
|
||||||
],
|
|
||||||
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
|
||||||
},
|
},
|
||||||
created_at=AnyStr(),
|
created_at=AnyStr(),
|
||||||
|
|||||||
@@ -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,
|
||||||
@@ -817,7 +820,7 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "one",
|
"name": "one",
|
||||||
"input": 2,
|
"input": 2,
|
||||||
"triggers": ["input"],
|
"triggers": ("input",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -828,7 +831,7 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "two",
|
"name": "two",
|
||||||
"input": [12],
|
"input": [12],
|
||||||
"triggers": ["inbox"],
|
"triggers": ("inbox",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -863,7 +866,7 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "two",
|
"name": "two",
|
||||||
"input": [3],
|
"input": [3],
|
||||||
"triggers": ["inbox"],
|
"triggers": ("inbox",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -3247,14 +3250,24 @@ def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
|
|||||||
|
|
||||||
assert [
|
assert [
|
||||||
c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config)
|
c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config)
|
||||||
] == [
|
] in (
|
||||||
{"rewrite_query": {"query": "query: what is weather in sf"}},
|
[
|
||||||
{"qa": {"answer": ""}},
|
{"rewrite_query": {"query": "query: what is weather in sf"}},
|
||||||
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
|
{"qa": {"answer": ""}},
|
||||||
{"retriever_two": {"docs": ["doc3", "doc4"]}},
|
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
|
||||||
{"retriever_one": {"docs": ["doc1", "doc2"]}},
|
{"retriever_two": {"docs": ["doc3", "doc4"]}},
|
||||||
{"__interrupt__": ()},
|
{"retriever_one": {"docs": ["doc1", "doc2"]}},
|
||||||
]
|
{"__interrupt__": ()},
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{"rewrite_query": {"query": "query: what is weather in sf"}},
|
||||||
|
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
|
||||||
|
{"qa": {"answer": ""}},
|
||||||
|
{"retriever_two": {"docs": ["doc3", "doc4"]}},
|
||||||
|
{"retriever_one": {"docs": ["doc1", "doc2"]}},
|
||||||
|
{"__interrupt__": ()},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
assert [c for c in app_w_interrupt.stream(None, config)] == [
|
||||||
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
|
{"qa": {"answer": "doc1,doc2,doc3,doc4"}},
|
||||||
@@ -4199,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"]
|
||||||
)
|
)
|
||||||
@@ -4621,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(
|
||||||
@@ -5969,9 +5929,7 @@ def test_falsy_return_from_task(
|
|||||||
"a": 5,
|
"a": 5,
|
||||||
},
|
},
|
||||||
"name": "graph",
|
"name": "graph",
|
||||||
"triggers": [
|
"triggers": ("__start__",),
|
||||||
"__start__",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
"step": 0,
|
"step": 0,
|
||||||
"timestamp": AnyStr(),
|
"timestamp": AnyStr(),
|
||||||
@@ -5985,9 +5943,7 @@ def test_falsy_return_from_task(
|
|||||||
{},
|
{},
|
||||||
),
|
),
|
||||||
"name": "falsy_task",
|
"name": "falsy_task",
|
||||||
"triggers": [
|
"triggers": ("__pregel_push",),
|
||||||
"__pregel_push",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
"step": 0,
|
"step": 0,
|
||||||
"timestamp": AnyStr(),
|
"timestamp": AnyStr(),
|
||||||
@@ -6094,9 +6050,7 @@ def test_falsy_return_from_task(
|
|||||||
"a": 5,
|
"a": 5,
|
||||||
},
|
},
|
||||||
"name": "graph",
|
"name": "graph",
|
||||||
"triggers": [
|
"triggers": ("__start__",),
|
||||||
"__start__",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
"step": 0,
|
"step": 0,
|
||||||
"timestamp": AnyStr(),
|
"timestamp": AnyStr(),
|
||||||
@@ -6110,9 +6064,7 @@ def test_falsy_return_from_task(
|
|||||||
{},
|
{},
|
||||||
),
|
),
|
||||||
"name": "falsy_task",
|
"name": "falsy_task",
|
||||||
"triggers": [
|
"triggers": ("__pregel_push",),
|
||||||
"__pregel_push",
|
|
||||||
],
|
|
||||||
},
|
},
|
||||||
"step": 0,
|
"step": 0,
|
||||||
"timestamp": AnyStr(),
|
"timestamp": AnyStr(),
|
||||||
@@ -6288,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)
|
||||||
@@ -6923,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:"),
|
||||||
@@ -7317,3 +7273,614 @@ def test_empty_invoke() -> None:
|
|||||||
"111": 111,
|
"111": 111,
|
||||||
"222": 222,
|
"222": 222,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||||
|
def test_parallel_interrupts(
|
||||||
|
request: pytest.FixtureRequest, checkpointer_name: str
|
||||||
|
) -> None:
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||||
|
|
||||||
|
# --- CHILD GRAPH ---
|
||||||
|
|
||||||
|
class ChildState(BaseModel):
|
||||||
|
prompt: str = Field(..., description="What is going to be asked to the user?")
|
||||||
|
human_input: Optional[str] = Field(None, description="What the human said")
|
||||||
|
human_inputs: Annotated[List[str], operator.add] = Field(
|
||||||
|
default_factory=list, description="All of my messages"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_human_input(state: ChildState):
|
||||||
|
human_input = interrupt(state.prompt)
|
||||||
|
|
||||||
|
return dict(
|
||||||
|
human_input=human_input, # update child state
|
||||||
|
human_inputs=[human_input], # update parent state
|
||||||
|
)
|
||||||
|
|
||||||
|
child_graph_builder = StateGraph(ChildState)
|
||||||
|
child_graph_builder.add_node("get_human_input", get_human_input)
|
||||||
|
child_graph_builder.add_edge(START, "get_human_input")
|
||||||
|
child_graph_builder.add_edge("get_human_input", END)
|
||||||
|
child_graph = child_graph_builder.compile()
|
||||||
|
|
||||||
|
# --- PARENT GRAPH ---
|
||||||
|
|
||||||
|
class ParentState(BaseModel):
|
||||||
|
prompts: List[str] = Field(
|
||||||
|
..., description="What is going to be asked to the user?"
|
||||||
|
)
|
||||||
|
human_inputs: Annotated[List[str], operator.add] = Field(
|
||||||
|
default_factory=list, description="All of my messages"
|
||||||
|
)
|
||||||
|
|
||||||
|
def assign_workers(state: ParentState):
|
||||||
|
return [
|
||||||
|
Send(
|
||||||
|
"child_graph",
|
||||||
|
dict(
|
||||||
|
prompt=prompt,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for prompt in state.prompts
|
||||||
|
]
|
||||||
|
|
||||||
|
def cleanup(state: ParentState):
|
||||||
|
assert len(state.human_inputs) == len(state.prompts)
|
||||||
|
|
||||||
|
parent_graph_builder = StateGraph(ParentState)
|
||||||
|
parent_graph_builder.add_node("child_graph", child_graph)
|
||||||
|
parent_graph_builder.add_node("cleanup", cleanup)
|
||||||
|
|
||||||
|
parent_graph_builder.add_conditional_edges(START, assign_workers, ["child_graph"])
|
||||||
|
parent_graph_builder.add_edge("child_graph", "cleanup")
|
||||||
|
parent_graph_builder.add_edge("cleanup", END)
|
||||||
|
|
||||||
|
parent_graph = parent_graph_builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
# --- CLIENT INVOCATION ---
|
||||||
|
|
||||||
|
thread_config = dict(
|
||||||
|
configurable=dict(
|
||||||
|
thread_id=str(uuid.uuid4()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
current_input = dict(
|
||||||
|
prompts=["a", "b"],
|
||||||
|
)
|
||||||
|
|
||||||
|
invokes = 0
|
||||||
|
events: dict[int, list[dict]] = {}
|
||||||
|
while invokes < 10:
|
||||||
|
# reset interrupt
|
||||||
|
invokes += 1
|
||||||
|
events[invokes] = []
|
||||||
|
current_interrupts: list[Interrupt] = []
|
||||||
|
|
||||||
|
# start / resume the graph
|
||||||
|
for event in parent_graph.stream(
|
||||||
|
input=current_input,
|
||||||
|
config=thread_config,
|
||||||
|
stream_mode="updates",
|
||||||
|
):
|
||||||
|
events[invokes].append(event)
|
||||||
|
# handle the interrupt
|
||||||
|
if "__interrupt__" in event:
|
||||||
|
current_interrupts.extend(event["__interrupt__"])
|
||||||
|
# assume that it breaks here, because it is an interrupt
|
||||||
|
|
||||||
|
# get human input and resume
|
||||||
|
if any(i.resumable for i in current_interrupts):
|
||||||
|
current_input = Command(resume=f"Resume #{invokes}")
|
||||||
|
|
||||||
|
# not more human input required, must be completed
|
||||||
|
else:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
assert False, "Detected infinite loop"
|
||||||
|
|
||||||
|
assert invokes == 3
|
||||||
|
assert len(events) == 3
|
||||||
|
|
||||||
|
assert events[1] == UnsortedSequence(
|
||||||
|
{
|
||||||
|
"__interrupt__": (
|
||||||
|
Interrupt(
|
||||||
|
value="a",
|
||||||
|
resumable=True,
|
||||||
|
ns=[
|
||||||
|
AnyStr("child_graph:"),
|
||||||
|
AnyStr("get_human_input:"),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"__interrupt__": (
|
||||||
|
Interrupt(
|
||||||
|
value="b",
|
||||||
|
resumable=True,
|
||||||
|
ns=[
|
||||||
|
AnyStr("child_graph:"),
|
||||||
|
AnyStr("get_human_input:"),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert events[2] in (
|
||||||
|
UnsortedSequence(
|
||||||
|
{
|
||||||
|
"__interrupt__": (
|
||||||
|
Interrupt(
|
||||||
|
value="a",
|
||||||
|
resumable=True,
|
||||||
|
ns=[
|
||||||
|
AnyStr("child_graph:"),
|
||||||
|
AnyStr("get_human_input:"),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
{"child_graph": {"human_inputs": ["Resume #1"]}},
|
||||||
|
),
|
||||||
|
UnsortedSequence(
|
||||||
|
{
|
||||||
|
"__interrupt__": (
|
||||||
|
Interrupt(
|
||||||
|
value="b",
|
||||||
|
resumable=True,
|
||||||
|
ns=[
|
||||||
|
AnyStr("child_graph:"),
|
||||||
|
AnyStr("get_human_input:"),
|
||||||
|
],
|
||||||
|
),
|
||||||
|
)
|
||||||
|
},
|
||||||
|
{"child_graph": {"human_inputs": ["Resume #1"]}},
|
||||||
|
),
|
||||||
|
)
|
||||||
|
assert events[3] == UnsortedSequence(
|
||||||
|
{
|
||||||
|
"child_graph": {"human_inputs": ["Resume #1"]},
|
||||||
|
"__metadata__": {"cached": True},
|
||||||
|
},
|
||||||
|
{"child_graph": {"human_inputs": ["Resume #2"]}},
|
||||||
|
{"cleanup": None},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||||
|
def test_parallel_interrupts_double(
|
||||||
|
request: pytest.FixtureRequest, checkpointer_name: str
|
||||||
|
) -> None:
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||||
|
|
||||||
|
# --- CHILD GRAPH ---
|
||||||
|
|
||||||
|
class ChildState(BaseModel):
|
||||||
|
prompt: str = Field(..., description="What is going to be asked to the user?")
|
||||||
|
human_input: Optional[str] = Field(None, description="What the human said")
|
||||||
|
human_inputs: Annotated[List[str], operator.add] = Field(
|
||||||
|
default_factory=list, description="All of my messages"
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_human_input(state: ChildState):
|
||||||
|
human_input = interrupt(state.prompt)
|
||||||
|
|
||||||
|
return dict(
|
||||||
|
human_inputs=[human_input], # update parent state
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_dolphin_input(state: ChildState):
|
||||||
|
human_input = interrupt(state.prompt)
|
||||||
|
|
||||||
|
return dict(
|
||||||
|
human_inputs=[human_input], # update parent state
|
||||||
|
)
|
||||||
|
|
||||||
|
child_graph_builder = StateGraph(ChildState)
|
||||||
|
child_graph_builder.add_node("get_human_input", get_human_input)
|
||||||
|
child_graph_builder.add_node("get_dolphin_input", get_dolphin_input)
|
||||||
|
child_graph_builder.add_edge(START, "get_human_input")
|
||||||
|
child_graph_builder.add_edge(START, "get_dolphin_input")
|
||||||
|
child_graph = child_graph_builder.compile()
|
||||||
|
|
||||||
|
# --- PARENT GRAPH ---
|
||||||
|
|
||||||
|
class ParentState(BaseModel):
|
||||||
|
prompts: List[str] = Field(
|
||||||
|
..., description="What is going to be asked to the user?"
|
||||||
|
)
|
||||||
|
human_inputs: Annotated[List[str], operator.add] = Field(
|
||||||
|
default_factory=list, description="All of my messages"
|
||||||
|
)
|
||||||
|
|
||||||
|
def assign_workers(state: ParentState):
|
||||||
|
return [
|
||||||
|
Send(
|
||||||
|
"child_graph",
|
||||||
|
dict(
|
||||||
|
prompt=prompt,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for prompt in state.prompts
|
||||||
|
]
|
||||||
|
|
||||||
|
def cleanup(state: ParentState):
|
||||||
|
assert len(state.human_inputs) == len(state.prompts) * 2
|
||||||
|
|
||||||
|
parent_graph_builder = StateGraph(ParentState)
|
||||||
|
parent_graph_builder.add_node("child_graph", child_graph)
|
||||||
|
parent_graph_builder.add_node("cleanup", cleanup)
|
||||||
|
|
||||||
|
parent_graph_builder.add_conditional_edges(START, assign_workers, ["child_graph"])
|
||||||
|
parent_graph_builder.add_edge("child_graph", "cleanup")
|
||||||
|
parent_graph_builder.add_edge("cleanup", END)
|
||||||
|
|
||||||
|
parent_graph = parent_graph_builder.compile(checkpointer=checkpointer)
|
||||||
|
|
||||||
|
# --- CLIENT INVOCATION ---
|
||||||
|
|
||||||
|
thread_config = dict(
|
||||||
|
configurable=dict(
|
||||||
|
thread_id=str(uuid.uuid4()),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
current_input = dict(
|
||||||
|
prompts=["a", "b"],
|
||||||
|
)
|
||||||
|
|
||||||
|
invokes = 0
|
||||||
|
events: dict[int, list[dict]] = {}
|
||||||
|
while invokes < 10:
|
||||||
|
# reset interrupt
|
||||||
|
invokes += 1
|
||||||
|
events[invokes] = []
|
||||||
|
current_interrupts: list[Interrupt] = []
|
||||||
|
|
||||||
|
# start / resume the graph
|
||||||
|
for event in parent_graph.stream(
|
||||||
|
input=current_input,
|
||||||
|
config=thread_config,
|
||||||
|
stream_mode="updates",
|
||||||
|
):
|
||||||
|
events[invokes].append(event)
|
||||||
|
# handle the interrupt
|
||||||
|
if "__interrupt__" in event:
|
||||||
|
current_interrupts.extend(event["__interrupt__"])
|
||||||
|
# assume that it breaks here, because it is an interrupt
|
||||||
|
|
||||||
|
# get human input and resume
|
||||||
|
if any(i.resumable for i in current_interrupts):
|
||||||
|
current_input = Command(resume=f"Resume #{invokes}")
|
||||||
|
|
||||||
|
# not more human input required, must be completed
|
||||||
|
else:
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
assert False, "Detected infinite loop"
|
||||||
|
|
||||||
|
assert invokes == 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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -27,11 +29,7 @@ from uuid import UUID
|
|||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
from langchain_core.language_models import GenericFakeChatModel
|
from langchain_core.language_models import GenericFakeChatModel
|
||||||
from langchain_core.runnables import (
|
from langchain_core.runnables import RunnableConfig, RunnableLambda, RunnablePassthrough
|
||||||
RunnableConfig,
|
|
||||||
RunnableLambda,
|
|
||||||
RunnablePassthrough,
|
|
||||||
)
|
|
||||||
from langchain_core.utils.aiter import aclosing
|
from langchain_core.utils.aiter import aclosing
|
||||||
from pytest_mock import MockerFixture
|
from pytest_mock import MockerFixture
|
||||||
from syrupy import SnapshotAssertion
|
from syrupy import SnapshotAssertion
|
||||||
@@ -56,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,
|
||||||
)
|
)
|
||||||
@@ -77,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,
|
||||||
@@ -938,10 +936,7 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
|
|||||||
async for c in tool_two.astream(
|
async for c in tool_two.astream(
|
||||||
{"my_key": "value ⛰️", "market": "DE"}, thread2
|
{"my_key": "value ⛰️", "market": "DE"}, thread2
|
||||||
)
|
)
|
||||||
] == [
|
] == UnsortedSequence(
|
||||||
{
|
|
||||||
"tool_one": {"my_key": " one"},
|
|
||||||
},
|
|
||||||
{
|
{
|
||||||
"__interrupt__": (
|
"__interrupt__": (
|
||||||
Interrupt(
|
Interrupt(
|
||||||
@@ -951,7 +946,10 @@ async def test_copy_checkpoint(checkpointer_name: str) -> None:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
},
|
},
|
||||||
]
|
{
|
||||||
|
"tool_one": {"my_key": " one"},
|
||||||
|
},
|
||||||
|
)
|
||||||
# resume with answer
|
# resume with answer
|
||||||
assert [
|
assert [
|
||||||
c async for c in tool_two.astream(Command(resume=" my answer"), thread2)
|
c async for c in tool_two.astream(Command(resume=" my answer"), thread2)
|
||||||
@@ -1672,7 +1670,7 @@ async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "one",
|
"name": "one",
|
||||||
"input": 2,
|
"input": 2,
|
||||||
"triggers": ["input"],
|
"triggers": ("input",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -1683,7 +1681,7 @@ async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "two",
|
"name": "two",
|
||||||
"input": [12],
|
"input": [12],
|
||||||
"triggers": ["inbox"],
|
"triggers": ("inbox",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -1718,7 +1716,7 @@ async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
|||||||
"id": AnyStr(),
|
"id": AnyStr(),
|
||||||
"name": "two",
|
"name": "two",
|
||||||
"input": [3],
|
"input": [3],
|
||||||
"triggers": ["inbox"],
|
"triggers": ("inbox",),
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -4524,6 +4522,7 @@ async def test_nested_pydantic_models(version: str) -> None:
|
|||||||
class NestedModel(BaseModel):
|
class NestedModel(BaseModel):
|
||||||
value: int
|
value: int
|
||||||
name: str
|
name: str
|
||||||
|
something: Optional[str] = None
|
||||||
|
|
||||||
# Forward reference model
|
# Forward reference model
|
||||||
class RecursiveModel(BaseModel):
|
class RecursiveModel(BaseModel):
|
||||||
@@ -4545,18 +4544,33 @@ 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):
|
||||||
|
x: int
|
||||||
|
my_enum: MyEnum
|
||||||
|
|
||||||
class State(BaseModel):
|
class State(BaseModel):
|
||||||
# Basic nested model tests
|
# Basic nested model tests
|
||||||
top_level: str
|
top_level: str
|
||||||
nested: NestedModel
|
nested: NestedModel
|
||||||
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_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]
|
||||||
]
|
]
|
||||||
|
list_nested_reversed: Annotated[
|
||||||
|
Union[list[dict[str, NestedModel]], NestedModel, dict, list],
|
||||||
|
lambda x, y: (x or []) + [y],
|
||||||
|
]
|
||||||
tuple_nested: tuple[str, NestedModel]
|
tuple_nested: tuple[str, NestedModel]
|
||||||
tuple_list_nested: list[tuple[int, NestedModel]]
|
tuple_list_nested: list[tuple[int, NestedModel]]
|
||||||
complex_tuple: tuple[str, dict[str, tuple[int, NestedModel]]]
|
complex_tuple: tuple[str, dict[str, tuple[int, NestedModel]]]
|
||||||
|
my_typed_dict: MyTypedDict
|
||||||
|
|
||||||
# Forward reference test
|
# Forward reference test
|
||||||
recursive: RecursiveModel
|
recursive: RecursiveModel
|
||||||
@@ -4572,8 +4586,12 @@ async def test_nested_pydantic_models(version: str) -> None:
|
|||||||
"top_level": "initial",
|
"top_level": "initial",
|
||||||
"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_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"],
|
||||||
"tuple_nested": ["tuple-key", {"value": 7, "name": "tuple-value"}],
|
"tuple_nested": ["tuple-key", {"value": 7, "name": "tuple-value"}],
|
||||||
"tuple_list_nested": [[1, {"value": 8, "name": "tuple-in-list"}]],
|
"tuple_list_nested": [[1, {"value": 8, "name": "tuple-in-list"}]],
|
||||||
"complex_tuple": [
|
"complex_tuple": [
|
||||||
@@ -5749,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"]
|
||||||
)
|
)
|
||||||
@@ -5882,10 +5900,12 @@ async def test_store_injected_async(checkpointer_name: str, store_name: str) ->
|
|||||||
):
|
):
|
||||||
assert isinstance(store, BaseStore)
|
assert isinstance(store, BaseStore)
|
||||||
await store.aput(
|
await store.aput(
|
||||||
namespace
|
(
|
||||||
if self.i is not None
|
namespace
|
||||||
and config["configurable"]["thread_id"] in (thread_1, thread_2)
|
if self.i is not None
|
||||||
else (f"foo_{self.i}", "bar"),
|
and config["configurable"]["thread_id"] in (thread_1, thread_2)
|
||||||
|
else (f"foo_{self.i}", "bar")
|
||||||
|
),
|
||||||
doc_id,
|
doc_id,
|
||||||
{
|
{
|
||||||
**doc,
|
**doc,
|
||||||
@@ -6997,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)
|
||||||
@@ -7571,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:"),
|
||||||
@@ -7827,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
|
||||||
|
|||||||
Generated
+6
-5
@@ -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,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,5 +1,7 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
import binascii
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
|
import weakref
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from contextlib import (
|
from contextlib import (
|
||||||
AbstractAsyncContextManager,
|
AbstractAsyncContextManager,
|
||||||
@@ -19,7 +21,7 @@ import langgraph.scheduler.kafka.serde as serde
|
|||||||
from langgraph.constants import CONFIG_KEY_DELEGATE, ERROR
|
from langgraph.constants import CONFIG_KEY_DELEGATE, ERROR
|
||||||
from langgraph.errors import CheckpointNotLatest, GraphDelegate, TaskNotFound
|
from langgraph.errors import CheckpointNotLatest, GraphDelegate, TaskNotFound
|
||||||
from langgraph.pregel import Pregel
|
from langgraph.pregel import Pregel
|
||||||
from langgraph.pregel.algo import prepare_single_task
|
from langgraph.pregel.algo import checkpoint_null_version, prepare_single_task
|
||||||
from langgraph.pregel.executor import (
|
from langgraph.pregel.executor import (
|
||||||
AsyncBackgroundExecutor,
|
AsyncBackgroundExecutor,
|
||||||
BackgroundExecutor,
|
BackgroundExecutor,
|
||||||
@@ -209,12 +211,17 @@ class AsyncKafkaExecutor(AbstractAsyncContextManager):
|
|||||||
for_execution=True,
|
for_execution=True,
|
||||||
checkpointer=self.graph.checkpointer,
|
checkpointer=self.graph.checkpointer,
|
||||||
store=self.graph.store,
|
store=self.graph.store,
|
||||||
|
checkpoint_id_bytes=binascii.unhexlify(
|
||||||
|
saved.checkpoint["id"].replace("-", "")
|
||||||
|
),
|
||||||
|
checkpoint_null_version=checkpoint_null_version(saved.checkpoint),
|
||||||
):
|
):
|
||||||
# execute task, saving writes
|
# execute task, saving writes
|
||||||
|
put_writes = partial(self._put_writes, submit, msg["config"])
|
||||||
runner = PregelRunner(
|
runner = PregelRunner(
|
||||||
submit=submit,
|
submit=weakref.ref(submit),
|
||||||
put_writes=partial(self._put_writes, submit, msg["config"]),
|
put_writes=weakref.ref(put_writes),
|
||||||
schedule_task=self._schedule_task,
|
schedule_task=weakref.WeakMethod(self._schedule_task),
|
||||||
)
|
)
|
||||||
async for _ in runner.atick([task], reraise=False):
|
async for _ in runner.atick([task], reraise=False):
|
||||||
pass
|
pass
|
||||||
@@ -421,12 +428,17 @@ class KafkaExecutor(AbstractContextManager):
|
|||||||
step=saved.metadata["step"] + 1,
|
step=saved.metadata["step"] + 1,
|
||||||
for_execution=True,
|
for_execution=True,
|
||||||
checkpointer=self.graph.checkpointer,
|
checkpointer=self.graph.checkpointer,
|
||||||
|
checkpoint_id_bytes=binascii.unhexlify(
|
||||||
|
saved.checkpoint["id"].replace("-", "")
|
||||||
|
),
|
||||||
|
checkpoint_null_version=checkpoint_null_version(saved.checkpoint),
|
||||||
):
|
):
|
||||||
# execute task, saving writes
|
# execute task, saving writes
|
||||||
|
put_writes = partial(self._put_writes, submit, msg["config"])
|
||||||
runner = PregelRunner(
|
runner = PregelRunner(
|
||||||
submit=submit,
|
submit=weakref.ref(submit),
|
||||||
put_writes=partial(self._put_writes, submit, msg["config"]),
|
put_writes=weakref.ref(put_writes),
|
||||||
schedule_task=self._schedule_task,
|
schedule_task=weakref.WeakMethod(self._schedule_task),
|
||||||
)
|
)
|
||||||
for _ in runner.tick([task], reraise=False):
|
for _ in runner.tick([task], reraise=False):
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -202,7 +202,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": None,
|
"checkpoint_id": None,
|
||||||
@@ -275,7 +275,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -378,7 +378,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -491,7 +491,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": None,
|
"checkpoint_id": None,
|
||||||
@@ -559,7 +559,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -683,7 +683,7 @@ async def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
|
|||||||
@@ -201,7 +201,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": None,
|
"checkpoint_id": None,
|
||||||
@@ -274,7 +274,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -377,7 +377,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -489,7 +489,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": None,
|
"checkpoint_id": None,
|
||||||
@@ -557,7 +557,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
@@ -681,7 +681,7 @@ def test_subgraph_w_interrupt(
|
|||||||
"subgraph_counter": None,
|
"subgraph_counter": None,
|
||||||
"call_counter": None,
|
"call_counter": None,
|
||||||
"interrupt_counter": None,
|
"interrupt_counter": None,
|
||||||
"null_resume": None,
|
"get_null_resume": None,
|
||||||
"resume": [],
|
"resume": [],
|
||||||
},
|
},
|
||||||
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
"checkpoint_id": c.config["configurable"]["checkpoint_id"],
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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,
|
||||||
|
})),
|
||||||
|
})),
|
||||||
},
|
},
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,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"
|
||||||
|
|||||||
Reference in New Issue
Block a user