mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-29 03:09:45 +02:00
Merge branch 'main' into ey/optimize_triggers
This commit is contained in:
@@ -14,6 +14,8 @@ To run the documentation server locally you can run:
|
||||
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
|
||||
|
||||
If you would like to automatically execute all of the notebooks, to mimic the "Run notebooks" GHA, you can run:
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 39 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 93 KiB |
@@ -1,6 +1,133 @@
|
||||
# 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.
|
||||
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}
|
||||
|
||||
|
||||
|
||||
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).
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -570,7 +570,7 @@ def prepare_single_task(
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: _scratchpad(
|
||||
config,
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
),
|
||||
@@ -680,7 +680,7 @@ def prepare_single_task(
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: _scratchpad(
|
||||
config,
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
),
|
||||
@@ -708,13 +708,14 @@ def prepare_single_task(
|
||||
if checkpoint_null_version is None:
|
||||
return
|
||||
# If any of the channels read by this process were updated
|
||||
if triggers := _triggers(
|
||||
if _triggers(
|
||||
channels,
|
||||
checkpoint["channel_versions"],
|
||||
checkpoint["versions_seen"].get(name),
|
||||
checkpoint_null_version,
|
||||
proc,
|
||||
):
|
||||
triggers = tuple(sorted(proc.triggers))
|
||||
try:
|
||||
val = next(
|
||||
_proc_input(proc, managed, channels, for_execution=for_execution)
|
||||
@@ -774,7 +775,7 @@ def prepare_single_task(
|
||||
CONFIG_KEY_SEND: partial(
|
||||
local_write,
|
||||
writes.extend,
|
||||
processes.keys(),
|
||||
tuple(processes.keys()),
|
||||
),
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
@@ -783,7 +784,10 @@ def prepare_single_task(
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(
|
||||
task_path[:3], name, writes, triggers
|
||||
task_path[:3],
|
||||
name,
|
||||
writes,
|
||||
triggers,
|
||||
),
|
||||
config,
|
||||
),
|
||||
@@ -801,7 +805,7 @@ def prepare_single_task(
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_SCRATCHPAD: _scratchpad(
|
||||
config,
|
||||
config[CONF].get(CONFIG_KEY_SCRATCHPAD),
|
||||
pending_writes,
|
||||
task_id,
|
||||
),
|
||||
@@ -852,7 +856,7 @@ def _triggers(
|
||||
|
||||
|
||||
def _scratchpad(
|
||||
config: RunnableConfig,
|
||||
parent_scratchpad: Optional[PregelScratchpad],
|
||||
pending_writes: list[PendingWrite],
|
||||
task_id: str,
|
||||
) -> PregelScratchpad:
|
||||
@@ -861,9 +865,6 @@ def _scratchpad(
|
||||
null_resume_write = next(
|
||||
(w for w in pending_writes if w[0] == NULL_TASK_ID and w[1] == RESUME), None
|
||||
)
|
||||
parent_scratchpad: Optional[PregelScratchpad] = config[CONF].get(
|
||||
CONFIG_KEY_SCRATCHPAD
|
||||
)
|
||||
|
||||
def get_null_resume(consume: bool = False) -> Any:
|
||||
if null_resume_write is None:
|
||||
|
||||
@@ -12,7 +12,7 @@ from langchain_core.runnables import Runnable, RunnableConfig
|
||||
from langchain_core.runnables.graph import Graph as DrawableGraph
|
||||
from typing_extensions import Self
|
||||
|
||||
from langgraph.pregel.types import All, StateSnapshot, StreamMode
|
||||
from langgraph.pregel.types import All, StateSnapshot, StateUpdate, StreamMode
|
||||
|
||||
|
||||
class PregelProtocol(
|
||||
@@ -69,6 +69,20 @@ class PregelProtocol(
|
||||
limit: Optional[int] = None,
|
||||
) -> 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
|
||||
def update_state(
|
||||
self,
|
||||
|
||||
@@ -457,6 +457,20 @@ class RemoteGraph(PregelProtocol):
|
||||
for state in states:
|
||||
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(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
|
||||
@@ -7,6 +7,7 @@ from langgraph.types import (
|
||||
PregelTask,
|
||||
RetryPolicy,
|
||||
StateSnapshot,
|
||||
StateUpdate,
|
||||
StreamMode,
|
||||
StreamWriter,
|
||||
default_retry_on,
|
||||
@@ -14,6 +15,7 @@ from langgraph.types import (
|
||||
|
||||
__all__ = [
|
||||
"All",
|
||||
"StateUpdate",
|
||||
"CachePolicy",
|
||||
"PregelExecutableTask",
|
||||
"PregelTask",
|
||||
|
||||
@@ -133,6 +133,11 @@ class Interrupt:
|
||||
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):
|
||||
id: str
|
||||
name: str
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "langgraph"
|
||||
version = "0.3.16"
|
||||
version = "0.3.17"
|
||||
description = "Building stateful, multi-actor applications with LLMs"
|
||||
authors = []
|
||||
license = "MIT"
|
||||
|
||||
@@ -2483,7 +2483,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 1,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("start:agent",),
|
||||
"langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
@@ -2542,7 +2542,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 3,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("tools",),
|
||||
"langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
@@ -2585,7 +2585,7 @@ def test_prebuilt_tool_chat(snapshot: SnapshotAssertion) -> None:
|
||||
{
|
||||
"langgraph_step": 5,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("tools",),
|
||||
"langgraph_triggers": ("branch:to:agent", "start:agent", "tools"),
|
||||
"langgraph_path": (PULL, "agent"),
|
||||
"langgraph_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(),
|
||||
"name": "rewrite_query",
|
||||
"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(),
|
||||
"name": "retriever_one",
|
||||
"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(),
|
||||
"name": "retriever_two",
|
||||
"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",
|
||||
"docs": ["doc1", "doc2", "doc3", "doc4"],
|
||||
},
|
||||
"triggers": (AnyStr("retriever_"),),
|
||||
"triggers": ("branch:to:qa", "retriever_one", "retriever_two"),
|
||||
},
|
||||
},
|
||||
),
|
||||
@@ -6634,7 +6643,7 @@ def test_branch_then(
|
||||
"id": AnyStr(),
|
||||
"name": "prepare",
|
||||
"input": {"my_key": "value", "market": "DE"},
|
||||
"triggers": ("start:prepare",),
|
||||
"triggers": ("branch:to:prepare", "start:prepare"),
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -6773,7 +6782,10 @@ def test_branch_then(
|
||||
"id": AnyStr(),
|
||||
"name": "finish",
|
||||
"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_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -7978,7 +7990,7 @@ def test_nested_graph_state(
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -8021,7 +8033,7 @@ def test_nested_graph_state(
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -8070,7 +8082,7 @@ def test_nested_graph_state(
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -8504,7 +8516,7 @@ def test_doubly_nested_graph_state(
|
||||
"langgraph_node": "child_1",
|
||||
"langgraph_path": [PULL, AnyStr("child_1")],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=(
|
||||
@@ -8588,7 +8600,10 @@ def test_doubly_nested_graph_state(
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
AnyStr("start:child_1"),
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=(
|
||||
@@ -8635,7 +8650,7 @@ def test_doubly_nested_graph_state(
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -8931,7 +8946,7 @@ def test_doubly_nested_graph_state(
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -8970,7 +8985,7 @@ def test_doubly_nested_graph_state(
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -9022,7 +9037,7 @@ def test_doubly_nested_graph_state(
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -9076,7 +9091,7 @@ def test_doubly_nested_graph_state(
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -9131,7 +9146,7 @@ def test_doubly_nested_graph_state(
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -9193,7 +9208,7 @@ def test_doubly_nested_graph_state(
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -9255,7 +9270,7 @@ def test_doubly_nested_graph_state(
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", AnyStr("start:child_1")],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
|
||||
@@ -2300,7 +2300,11 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 1,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("start:agent",),
|
||||
"langgraph_triggers": (
|
||||
"branch:to:agent",
|
||||
"start:agent",
|
||||
"tools",
|
||||
),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
@@ -2359,7 +2363,11 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 3,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("tools",),
|
||||
"langgraph_triggers": (
|
||||
"branch:to:agent",
|
||||
"start:agent",
|
||||
"tools",
|
||||
),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_checkpoint_ns": AnyStr("agent:"),
|
||||
"checkpoint_ns": AnyStr("agent:"),
|
||||
@@ -2402,7 +2410,11 @@ async def test_prebuilt_tool_chat() -> None:
|
||||
{
|
||||
"langgraph_step": 5,
|
||||
"langgraph_node": "agent",
|
||||
"langgraph_triggers": ("tools",),
|
||||
"langgraph_triggers": (
|
||||
"branch:to:agent",
|
||||
"start:agent",
|
||||
"tools",
|
||||
),
|
||||
"langgraph_path": ("__pregel_pull", "agent"),
|
||||
"langgraph_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(),
|
||||
"name": "rewrite_query",
|
||||
"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(),
|
||||
"name": "retriever_one",
|
||||
"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(),
|
||||
"name": "retriever_two",
|
||||
"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",
|
||||
"docs": ["doc1", "doc2", "doc3", "doc4"],
|
||||
},
|
||||
"triggers": (AnyStr("retriever_"),),
|
||||
"triggers": ("branch:to:qa", "retriever_one", "retriever_two"),
|
||||
},
|
||||
},
|
||||
),
|
||||
@@ -4465,7 +4486,10 @@ async def test_branch_then(checkpointer_name: str) -> None:
|
||||
"id": AnyStr(),
|
||||
"name": "prepare",
|
||||
"input": {"my_key": "value", "market": "DE"},
|
||||
"triggers": ("start:prepare",),
|
||||
"triggers": (
|
||||
"branch:to:prepare",
|
||||
"start:prepare",
|
||||
),
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -4609,7 +4633,10 @@ async def test_branch_then(checkpointer_name: str) -> None:
|
||||
"id": AnyStr(),
|
||||
"name": "finish",
|
||||
"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(),
|
||||
"name": "prepare",
|
||||
"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_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -5530,7 +5560,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -5573,7 +5603,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -5622,7 +5652,7 @@ async def test_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "inner",
|
||||
"langgraph_path": [PULL, "inner"],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": ["outer_1"],
|
||||
"langgraph_triggers": ["branch:to:inner", "outer_1"],
|
||||
"langgraph_checkpoint_ns": AnyStr("inner:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -6060,7 +6090,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "child_1",
|
||||
"langgraph_path": [PULL, AnyStr("child_1")],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": ["branch:to:child_1", "start:child_1"],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=(
|
||||
@@ -6146,7 +6176,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
"start:child_1",
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=(
|
||||
@@ -6195,7 +6228,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child",
|
||||
AnyStr("parent_1"),
|
||||
],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -6493,7 +6529,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -6532,7 +6568,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -6584,7 +6620,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
"langgraph_node": "child",
|
||||
"langgraph_path": [PULL, AnyStr("child")],
|
||||
"langgraph_step": 2,
|
||||
"langgraph_triggers": [AnyStr("parent_1")],
|
||||
"langgraph_triggers": ["branch:to:child", AnyStr("parent_1")],
|
||||
"langgraph_checkpoint_ns": AnyStr("child:"),
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
@@ -6642,7 +6678,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
AnyStr("start:child_1"),
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -6697,7 +6736,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
AnyStr("start:child_1"),
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -6759,7 +6801,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
AnyStr("start:child_1"),
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config={
|
||||
@@ -6821,7 +6866,10 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
|
||||
AnyStr("child_1"),
|
||||
],
|
||||
"langgraph_step": 1,
|
||||
"langgraph_triggers": [AnyStr("start:child_1")],
|
||||
"langgraph_triggers": [
|
||||
"branch:to:child_1",
|
||||
AnyStr("start:child_1"),
|
||||
],
|
||||
},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
|
||||
@@ -69,6 +69,7 @@ from langgraph.types import (
|
||||
Interrupt,
|
||||
PregelTask,
|
||||
Send,
|
||||
StateUpdate,
|
||||
StreamWriter,
|
||||
interrupt,
|
||||
)
|
||||
@@ -6925,7 +6926,10 @@ def test_tags_stream_mode_messages() -> None:
|
||||
{
|
||||
"langgraph_step": 1,
|
||||
"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_checkpoint_ns": AnyStr("call_model:"),
|
||||
"checkpoint_ns": AnyStr("call_model:"),
|
||||
@@ -7613,3 +7617,286 @@ def test_parallel_interrupts_double(
|
||||
|
||||
assert invokes == 5
|
||||
assert len(events) == 5
|
||||
|
||||
|
||||
@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
|
||||
|
||||
@@ -59,6 +59,7 @@ from langgraph.types import (
|
||||
Interrupt,
|
||||
PregelTask,
|
||||
Send,
|
||||
StateUpdate,
|
||||
StreamWriter,
|
||||
interrupt,
|
||||
)
|
||||
@@ -7582,7 +7583,10 @@ async def test_tags_stream_mode_messages() -> None:
|
||||
{
|
||||
"langgraph_step": 1,
|
||||
"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_checkpoint_ns": AnyStr("call_model:"),
|
||||
"checkpoint_ns": AnyStr("call_model:"),
|
||||
@@ -7838,3 +7842,291 @@ async def test_handles_multiple_interrupts_from_tasks() -> None:
|
||||
assert len(result) == 2
|
||||
assert result[0] == "Added James!"
|
||||
assert result[1] == "Added Will!"
|
||||
|
||||
|
||||
@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
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
{
|
||||
"name": "@langchain/langgraph-sdk",
|
||||
"version": "0.0.57",
|
||||
"version": "0.0.59",
|
||||
"description": "Client library for interacting with the LangGraph API",
|
||||
"type": "module",
|
||||
"packageManager": "yarn@1.22.19",
|
||||
|
||||
@@ -30,6 +30,7 @@ import type {
|
||||
StreamEvent,
|
||||
CronsCreatePayload,
|
||||
OnConflictBehavior,
|
||||
Command,
|
||||
} from "./types.js";
|
||||
import { mergeSignals } from "./utils/signals.js";
|
||||
import { getEnvironmentVariable } from "./utils/env.js";
|
||||
@@ -638,6 +639,40 @@ export class ThreadsClient<
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Create a new thread from a batch states.
|
||||
*/
|
||||
async bulkUpdateState(
|
||||
supersteps: Array<{
|
||||
updates: Array<{ values: unknown; command?: Command; asNode: string }>;
|
||||
}>,
|
||||
options?: {
|
||||
graphId?: string;
|
||||
threadId?: string;
|
||||
metadata?: Metadata;
|
||||
ifExists?: OnConflictBehavior;
|
||||
},
|
||||
): Promise<Thread<TStateType>> {
|
||||
return this.fetch<Thread<TStateType>>("/threads/state/batch", {
|
||||
method: "POST",
|
||||
json: {
|
||||
supersteps: supersteps.map((s) => ({
|
||||
updates: s.updates.map((u) => ({
|
||||
values: u.values,
|
||||
command: u.command,
|
||||
as_node: u.asNode,
|
||||
})),
|
||||
})),
|
||||
thread_id: options?.threadId,
|
||||
metadata: {
|
||||
...options?.metadata,
|
||||
graph_id: options?.graphId,
|
||||
},
|
||||
if_exists: options?.ifExists,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* Patch the metadata of a thread.
|
||||
*
|
||||
|
||||
Reference in New Issue
Block a user