mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-27 12:04:58 +02:00
## 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 :)
145 lines
4.5 KiB
Python
145 lines
4.5 KiB
Python
"""Shared utility functions for async and sync clients."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import functools
|
|
import os
|
|
import re
|
|
from collections.abc import Mapping
|
|
from typing import Any, cast
|
|
|
|
import httpx
|
|
|
|
import langgraph_sdk
|
|
from langgraph_sdk.schema import RunCreateMetadata
|
|
|
|
RESERVED_HEADERS = ("x-api-key",)
|
|
|
|
NOT_PROVIDED = cast(None, object())
|
|
|
|
|
|
def _get_api_key(api_key: str | None = NOT_PROVIDED) -> str | None:
|
|
"""Get the API key from the environment.
|
|
Precedence:
|
|
1. explicit string argument
|
|
2. LANGGRAPH_API_KEY (if api_key not provided)
|
|
3. LANGSMITH_API_KEY (if api_key not provided)
|
|
4. LANGCHAIN_API_KEY (if api_key not provided)
|
|
|
|
Args:
|
|
api_key: The API key to use. Can be:
|
|
- A string: use this exact API key
|
|
- None: explicitly skip loading from environment
|
|
- NOT_PROVIDED (default): auto-load from environment variables
|
|
"""
|
|
if isinstance(api_key, str):
|
|
return api_key
|
|
if api_key is NOT_PROVIDED:
|
|
# api_key is not explicitly provided, try to load from environment
|
|
for prefix in ["LANGGRAPH", "LANGSMITH", "LANGCHAIN"]:
|
|
if env := os.getenv(f"{prefix}_API_KEY"):
|
|
return env.strip().strip('"').strip("'")
|
|
# api_key is explicitly None, don't load from environment
|
|
return None
|
|
|
|
|
|
def _get_headers(
|
|
api_key: str | None,
|
|
custom_headers: Mapping[str, str] | None,
|
|
) -> dict[str, str]:
|
|
"""Combine api_key and custom user-provided headers."""
|
|
custom_headers = custom_headers or {}
|
|
for header in RESERVED_HEADERS:
|
|
if header in custom_headers:
|
|
raise ValueError(f"Cannot set reserved header '{header}'")
|
|
|
|
headers = {
|
|
"User-Agent": f"langgraph-sdk-py/{langgraph_sdk.__version__}",
|
|
**custom_headers,
|
|
}
|
|
resolved_api_key = _get_api_key(api_key)
|
|
if resolved_api_key:
|
|
headers["x-api-key"] = resolved_api_key
|
|
|
|
return headers
|
|
|
|
|
|
def _orjson_default(obj: Any) -> Any:
|
|
is_class = isinstance(obj, type)
|
|
if hasattr(obj, "model_dump") and callable(obj.model_dump):
|
|
if is_class:
|
|
raise TypeError(
|
|
f"Cannot JSON-serialize type object: {obj!r}. Did you mean to pass an instance of the object instead?"
|
|
f"\nReceived type: {obj!r}"
|
|
)
|
|
return obj.model_dump()
|
|
elif hasattr(obj, "dict") and callable(obj.dict):
|
|
if is_class:
|
|
raise TypeError(
|
|
f"Cannot JSON-serialize type object: {obj!r}. Did you mean to pass an instance of the object instead?"
|
|
f"\nReceived type: {obj!r}"
|
|
)
|
|
return obj.dict()
|
|
elif isinstance(obj, (set, frozenset)):
|
|
return list(obj)
|
|
else:
|
|
raise TypeError(f"Object of type {type(obj)} is not JSON serializable")
|
|
|
|
|
|
# Compiled regex pattern for extracting run metadata from Content-Location header
|
|
_RUN_METADATA_PATTERN = re.compile(
|
|
r"(\/threads\/(?P<thread_id>.+))?\/runs\/(?P<run_id>.+)"
|
|
)
|
|
|
|
|
|
def _get_run_metadata_from_response(
|
|
response: httpx.Response,
|
|
) -> RunCreateMetadata | None:
|
|
"""Extract run metadata from the response headers."""
|
|
if (content_location := response.headers.get("Content-Location")) and (
|
|
match := _RUN_METADATA_PATTERN.search(content_location)
|
|
):
|
|
return RunCreateMetadata(
|
|
run_id=match.group("run_id"),
|
|
thread_id=match.group("thread_id") or None,
|
|
)
|
|
|
|
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}
|
|
|
|
|
|
_registered_transports: list[httpx.ASGITransport] = []
|
|
|
|
|
|
# Do not move; this is used in the server.
|
|
def configure_loopback_transports(app: Any) -> None:
|
|
for transport in _registered_transports:
|
|
transport.app = app
|
|
|
|
|
|
@functools.lru_cache(maxsize=1)
|
|
def get_asgi_transport() -> type[httpx.ASGITransport]:
|
|
try:
|
|
from langgraph_api import asgi_transport # type: ignore[unresolved-import]
|
|
|
|
return asgi_transport.ASGITransport
|
|
except ImportError:
|
|
# Older versions of the server
|
|
return httpx.ASGITransport
|