feat(langgraph): new context api (replacing config['configurable'] and config_schema) (#5243)

This commit is contained in:
Sydney Runkle
2025-07-15 09:20:20 -04:00
committed by GitHub
parent e0bf4a7bc3
commit 0837263542
30 changed files with 845 additions and 386 deletions
+33 -21
View File
@@ -12,56 +12,66 @@ LangGraph provides **three** primary ways to supply context:
| Type | Description | Mutable? | Lifetime |
|------------------------------------------------------------------------------|-----------------------------------------------|----------|-------------------------|
| [**Config**](#config-static-context) | data passed at the start of a run | ❌ | per run |
| [**Runtime Context**](#runtime-context) | data passed at the start of a run | ❌ | per run |
| [**Short-term memory (State)**](#short-term-memory-mutable-context) | dynamic data that can change during execution | ✅ | per run or conversation |
| [**Long-term memory (Store)**](#long-term-memory-cross-conversation-context) | data that can be shared between conversations | ✅ | across conversations |
## Provide runtime context
### Config (static context)
### Runtime Context
Config is for immutable data like user metadata or API keys. Use
when you have values that don't change mid-run.
!!! note "`config['configurable']` -> `runtime.context`"
Specify configuration using a key called **"configurable"** which is reserved
for this purpose:
In LangGraph < v1.0, static runtime context was passed via the `config['configurable']` key, paired with a `config_schema` argument
to `StateGraph` or `Pregel`. This is now deprecated and will be removed in v2.0.
As of LangGraph v1.0, the Runtime object is recommended to access static context and runtime-specific information like the store and stream writer.
Runtime context is for immutable data like user metadata or API keys. Use this when you have values that don't change mid-run.
Specify static context via the `context` argument to `invoke` / `stream`, which is reserved for this purpose:
```python
@dataclass
class ContextSchema:
user_name: str
graph.invoke( # (1)!
{"messages": [{"role": "user", "content": "hi!"}]}, # (2)!
# highlight-next-line
config={"configurable": {"user_id": "user_123"}} # (3)!
context={"user_name": "John Smith"} # (3)!
)
```
1. This is the invocation of the agent or graph. The `invoke` method runs the underlying graph with the provided input.
2. This example uses messages as an input, which is common, but your application may use different input structures.
3. This is where you pass the configuration data. The `config` parameter allows you to provide additional context that the agent can use during its execution.
3. This is where you pass the runtime data. The `context` parameter allows you to provide additional dependencies that the agent can use during its execution.
=== "Agent prompt"
```python
from langchain_core.messages import AnyMessage
from langchain_core.runnables import RunnableConfig
from langgraph.runtime import get_runtime
from langgraph.prebuilt.chat_agent_executor import AgentState
from langgraph.prebuilt import create_react_agent
# highlight-next-line
def prompt(state: AgentState, config: RunnableConfig) -> list[AnyMessage]:
user_name = config["configurable"].get("user_name")
system_msg = f"You are a helpful assistant. Address the user as {user_name}."
def prompt(state: AgentState) -> list[AnyMessage]:
runtime = get_runtime(ContextSchema)
system_msg = f"You are a helpful assistant. Address the user as {runtime.context.user_name}."
return [{"role": "system", "content": system_msg}] + state["messages"]
agent = create_react_agent(
model="anthropic:claude-3-7-sonnet-latest",
tools=[get_weather],
prompt=prompt
prompt=prompt,
context_schema=ContextSchema
)
agent.invoke(
{"messages": [{"role": "user", "content": "what is the weather in sf"}]},
# highlight-next-line
config={"configurable": {"user_name": "John Smith"}}
context={"user_name": "John Smith"}
)
```
@@ -70,11 +80,11 @@ graph.invoke( # (1)!
=== "Workflow node"
```python
from langchain_core.runnables import RunnableConfig
from langgraph.runtime import Runtime
# highlight-next-line
def node(state: State, config: RunnableConfig):
user_name = config["configurable"].get("user_name")
def node(state: State, config: Runtime[ContextSchema]):
user_name = runtime.context.user_name
...
```
@@ -83,14 +93,16 @@ graph.invoke( # (1)!
=== "In a tool"
```python
from langchain_core.runnables import RunnableConfig
from langgraph.runtime import get_runtime
@tool
# highlight-next-line
def get_user_info(config: RunnableConfig) -> str:
def get_user_email() -> str:
"""Retrieve user information based on user ID."""
user_id = config["configurable"].get("user_id")
return "User is John Smith" if user_id == "user_123" else "Unknown user"
# simulate fetching user info from a database
runtime = get_runtime(ContextSchema)
email = get_user_email_from_db(runtime.context.user_name)
return email
```
See the [tool calling guide](../how-tos/tool-calling.md#configuration) for details.
+3 -3
View File
@@ -108,11 +108,11 @@ from langgraph.graph import StateGraph, END, START
from my_agent.utils.nodes import call_model, should_continue, tool_node # import nodes
from my_agent.utils.state import AgentState # import state
# Define the config
class GraphConfig(TypedDict):
# Define the runtime context
class GraphContext(TypedDict):
model_name: Literal["anthropic", "openai"]
workflow = StateGraph(AgentState, config_schema=GraphConfig)
workflow = StateGraph(AgentState, context_schema=GraphContext)
workflow.add_node("agent", call_model)
workflow.add_node("action", tool_node)
workflow.add_edge(START, "agent")
@@ -121,11 +121,11 @@ from langgraph.graph import StateGraph, END, START
from my_agent.utils.nodes import call_model, should_continue, tool_node # import nodes
from my_agent.utils.state import AgentState # import state
# Define the config
class GraphConfig(TypedDict):
# Define the runtime context
class GraphContext(TypedDict):
model_name: Literal["anthropic", "openai"]
workflow = StateGraph(AgentState, config_schema=GraphConfig)
workflow = StateGraph(AgentState, context_schema=GraphContext)
workflow.add_node("agent", call_model)
workflow.add_node("action", tool_node)
workflow.add_edge(START, "agent")
@@ -2,21 +2,20 @@
In this guide we will show how to create, configure, and manage an [assistant](../../concepts/assistants.md).
First, as a brief refresher on the concept of configurations, consider the following simple `call_model` node and configuration schema. Observe that this node tries to read and use the `model_name` as defined by the `config` object's `configurable`.
First, as a brief refresher on the concept of runtime context, consider the following simple `call_model` node and configuration schema. Observe that this node tries to read and use the `model_provider` as defined by the `Runtime` object's `context` property.
=== "Python"
```python
@dataclass
class ContextSchema:
llm_provider: str = "anthropic"
class ConfigSchema(TypedDict):
model_name: str
builder = StateGraph(AgentState, context_schema=ContextSchema)
builder = StateGraph(AgentState, config_schema=ConfigSchema)
def call_model(state, config):
def call_model(state, runtime: Runtime[ContextSchema]):
messages = state["messages"]
model_name = config.get('configurable', {}).get("model_name", "anthropic")
model = _get_model(model_name)
model = _get_model(runtime.context.llm_provider)
response = model.invoke(messages)
# We return a list, because this will get added to the existing list
return {"messages": [response]}
+15 -16
View File
@@ -459,33 +459,32 @@ LangGraph can easily handle migrations of graph definitions (nodes, edges, and s
- State keys that are renamed lose their saved state in existing threads
- State keys whose types change in incompatible ways could currently cause issues in threads with state from before the change -- if this is a blocker please reach out and we can prioritize a solution.
## Configuration
## Runtime Context
When creating a graph, you can also mark that certain parts of the graph are configurable. This is commonly done to enable easily switching between models or system prompts. This allows you to create a single "cognitive architecture" (the graph) but have multiple different instance of it.
You can optionally specify a `config_schema` when creating a graph.
When creating a graph, you can specify a `context_schema` for runtime context passed to nodes. This is useful for passing
information to nodes that is not part of the graph state. For example, you might want to pass dependencies such as model name or a database connection.
```python
class ConfigSchema(TypedDict):
llm: str
@dataclass
class ContextSchema:
llm_provider: str = "openai"
graph = StateGraph(State, config_schema=ConfigSchema)
graph = StateGraph(State, context_schema=ContextSchema)
```
You can then pass this configuration into the graph using the `configurable` config field.
You can then pass this context into the graph using the `context` parameter of the `invoke` method.
```python
config = {"configurable": {"llm": "anthropic"}}
graph.invoke(inputs, config=config)
graph.invoke(inputs, context={"llm_provider": "anthropic"})
```
You can then access and use this configuration inside a node or conditional edge:
You can then access and use this context inside a node or conditional edge:
```python
def node_a(state, config):
llm_type = config.get("configurable", {}).get("llm", "openai")
llm = get_llm(llm_type)
from langgraph.runtime import Runtime
def node_a(state: State, runtime: Runtime[ContextSchema]):
llm = get_llm(runtime.context.llm_provider)
...
```
@@ -496,7 +495,7 @@ See [this guide](../how-tos/graph-api.md#add-runtime-configuration) for a full b
The recursion limit sets the maximum number of [super-steps](#graphs) the graph can execute during a single execution. Once the limit is reached, LangGraph will raise `GraphRecursionError`. By default this value is set to 25 steps. The recursion limit can be set on any graph at runtime, and is passed to `.invoke`/`.stream` via the config dictionary. Importantly, `recursion_limit` is a standalone `config` key and should not be passed inside the `configurable` key as all other user-defined configuration. See the example below:
```python
graph.invoke(inputs, config={"recursion_limit": 5, "configurable":{"llm": "anthropic"}})
graph.invoke(inputs, config={"recursion_limit": 5}, context={"llm": "anthropic"})
```
Read [this how-to](https://langchain-ai.github.io/langgraph/how-tos/recursion-limit/) to learn more about how the recursion limit works.
+30 -30
View File
@@ -513,12 +513,12 @@ To add runtime configuration:
See below for a simple example:
```python
from langchain_core.runnables import RunnableConfig
from langgraph.graph import END, StateGraph, START
from langgraph.runtime import Runtime
from typing_extensions import TypedDict
# 1. Specify config schema
class ConfigSchema(TypedDict):
class ContextSchema(TypedDict):
my_runtime_value: str
# 2. Define a graph that accesses the config in a node
@@ -526,18 +526,18 @@ class State(TypedDict):
my_state_value: str
# highlight-next-line
def node(state: State, config: RunnableConfig):
def node(state: State, runtime: Runtime[ContextSchema]):
# highlight-next-line
if config["configurable"]["my_runtime_value"] == "a":
if runtime.context["my_runtime_value"] == "a":
return {"my_state_value": 1}
# highlight-next-line
elif config["configurable"]["my_runtime_value"] == "b":
elif runtime.context["my_runtime_value"] == "b":
return {"my_state_value": 2}
else:
raise ValueError("Unknown values.")
# highlight-next-line
builder = StateGraph(State, config_schema=ConfigSchema)
builder = StateGraph(State, context_schema=ContextSchema)
builder.add_node(node)
builder.add_edge(START, "node")
builder.add_edge("node", END)
@@ -546,9 +546,9 @@ graph = builder.compile()
# 3. Pass in configuration at runtime:
# highlight-next-line
print(graph.invoke({}, {"configurable": {"my_runtime_value": "a"}}))
print(graph.invoke({}, context={"my_runtime_value": "a"}))
# highlight-next-line
print(graph.invoke({}, {"configurable": {"my_runtime_value": "b"}}))
print(graph.invoke({}, context={"my_runtime_value": "b"}))
```
```
{'my_state_value': 1}
@@ -559,27 +559,28 @@ print(graph.invoke({}, {"configurable": {"my_runtime_value": "b"}}))
Below we demonstrate a practical example in which we configure what LLM to use at runtime. We will use both OpenAI and Anthropic models.
```python
from dataclasses import dataclass
from langchain.chat_models import init_chat_model
from langchain_core.runnables import RunnableConfig
from langgraph.graph import MessagesState
from langgraph.graph import END, StateGraph, START
from langgraph.graph import MessagesState, END, StateGraph, START
from langgraph.runtime import Runtime
from typing_extensions import TypedDict
class ConfigSchema(TypedDict):
model: str
@dataclass
class ContextSchema:
model_provider: str = "anthropic"
MODELS = {
"anthropic": init_chat_model("anthropic:claude-3-5-haiku-latest"),
"openai": init_chat_model("openai:gpt-4.1-mini"),
}
def call_model(state: MessagesState, config: RunnableConfig):
model = config["configurable"].get("model", "anthropic")
model = MODELS[model]
def call_model(state: MessagesState, runtime: Runtime[ContextSchema]):
model = MODELS[runtime.context.model_provider]
response = model.invoke(state["messages"])
return {"messages": [response]}
builder = StateGraph(MessagesState, config_schema=ConfigSchema)
builder = StateGraph(MessagesState, context_schema=ContextSchema)
builder.add_node("model", call_model)
builder.add_edge(START, "model")
builder.add_edge("model", END)
@@ -591,8 +592,7 @@ print(graph.invoke({}, {"configurable": {"my_runtime_value": "b"}}))
# With no configuration, uses default (Anthropic)
response_1 = graph.invoke({"messages": [input_message]})["messages"][-1]
# Or, can set OpenAI
config = {"configurable": {"model": "openai"}}
response_2 = graph.invoke({"messages": [input_message]}, config=config)["messages"][-1]
response_2 = graph.invoke({"messages": [input_message]}, context={"model_provider": "openai"})["messages"][-1]
print(response_1.response_metadata["model_name"])
print(response_2.response_metadata["model_name"])
@@ -606,32 +606,33 @@ print(graph.invoke({}, {"configurable": {"my_runtime_value": "b"}}))
Below we demonstrate a practical example in which we configure two parameters: the LLM and system message to use at runtime.
```python
from dataclasses import dataclass
from typing import Optional
from langchain.chat_models import init_chat_model
from langchain_core.messages import SystemMessage
from langchain_core.runnables import RunnableConfig
from langgraph.graph import END, MessagesState, StateGraph, START
from langgraph.runtime import Runtime
from typing_extensions import TypedDict
class ConfigSchema(TypedDict):
model: Optional[str]
system_message: Optional[str]
@dataclass
class ContextSchema:
model_provider: str = "anthropic"
system_message: str | None = None
MODELS = {
"anthropic": init_chat_model("anthropic:claude-3-5-haiku-latest"),
"openai": init_chat_model("openai:gpt-4.1-mini"),
}
def call_model(state: MessagesState, config: RunnableConfig):
model = config["configurable"].get("model", "anthropic")
model = MODELS[model]
def call_model(state: MessagesState, runtime: Runtime[ContextSchema]):
model = MODELS[runtime.context.model_provider]
messages = state["messages"]
if system_message := config["configurable"].get("system_message"):
if (system_message := runtime.context.system_message):
messages = [SystemMessage(system_message)] + messages
response = model.invoke(messages)
return {"messages": [response]}
builder = StateGraph(MessagesState, config_schema=ConfigSchema)
builder = StateGraph(MessagesState, context_schema=ContextSchema)
builder.add_node("model", call_model)
builder.add_edge(START, "model")
builder.add_edge("model", END)
@@ -640,8 +641,7 @@ print(graph.invoke({}, {"configurable": {"my_runtime_value": "b"}}))
# Usage
input_message = {"role": "user", "content": "hi"}
config = {"configurable": {"model": "openai", "system_message": "Respond in Italian."}}
response = graph.invoke({"messages": [input_message]}, config)
response = graph.invoke({"messages": [input_message]}, context={"model_provider": "openai", "system_message": "Respond in Italian."})
for message in response["messages"]:
message.pretty_print()
```
+26 -30
View File
@@ -280,8 +280,8 @@
"from typing import Optional, Dict, Any\n",
"from typing_extensions import Annotated, TypedDict\n",
"from langgraph.graph import StateGraph\n",
"from langgraph.types import Runtime\n",
"\n",
"from langchain_core.runnables import RunnableConfig\n",
"from langgraph.checkpoint.memory import MemorySaver\n",
"from langgraph.types import Send\n",
"\n",
@@ -307,22 +307,25 @@
" depth: Annotated[int, operator.add]\n",
"\n",
"\n",
"class Configuration(TypedDict, total=False):\n",
"class Context(TypedDict, total=False):\n",
" max_depth: int\n",
" threshold: float\n",
" k: int\n",
" beam_size: int\n",
"\n",
"class EnsuredContext(TypedDict):\n",
" max_depth: int\n",
" threshold: float\n",
" k: int\n",
" beam_size: int\n",
"\n",
"def _ensure_configurable(config: RunnableConfig) -> Configuration:\n",
"def _ensure_context(ctx: Context) -> EnsuredContext:\n",
" \"\"\"Get params that configure the search algorithm.\"\"\"\n",
" configurable = config.get(\"configurable\", {})\n",
" return {\n",
" **configurable,\n",
" \"max_depth\": configurable.get(\"max_depth\", 10),\n",
" \"threshold\": config.get(\"threshold\", 0.9),\n",
" \"k\": configurable.get(\"k\", 5),\n",
" \"beam_size\": configurable.get(\"beam_size\", 3),\n",
" \"max_depth\": ctx.get(\"max_depth\", 10),\n",
" \"threshold\": ctx.get(\"threshold\", 0.9),\n",
" \"k\": ctx.get(\"k\", 5),\n",
" \"beam_size\": ctx.get(\"beam_size\", 3)\n",
" }\n",
"\n",
"\n",
@@ -330,9 +333,9 @@
" seed: Optional[Candidate]\n",
"\n",
"\n",
"def expand(state: ExpansionState, *, config: RunnableConfig) -> Dict[str, List[str]]:\n",
"def expand(state: ExpansionState, *, runtime: Runtime[Context]) -> Dict[str, List[Candidate]]:\n",
" \"\"\"Generate the next state.\"\"\"\n",
" configurable = _ensure_configurable(config)\n",
" ctx = _ensure_context(runtime.context)\n",
" if not state.get(\"seed\"):\n",
" candidate_str = \"\"\n",
" else:\n",
@@ -342,9 +345,8 @@
" {\n",
" \"problem\": state[\"problem\"],\n",
" \"candidate\": candidate_str,\n",
" \"k\": configurable[\"k\"],\n",
" \"k\": ctx[\"k\"],\n",
" },\n",
" config=config,\n",
" )\n",
" except Exception:\n",
" return {\"candidates\": []}\n",
@@ -354,7 +356,7 @@
" return {\"candidates\": new_candidates}\n",
"\n",
"\n",
"def score(state: ToTState) -> Dict[str, List[float]]:\n",
"def score(state: ToTState) -> Dict[str, Any]:\n",
" \"\"\"Evaluate the candidate generations.\"\"\"\n",
" candidates = state[\"candidates\"]\n",
" scored = []\n",
@@ -364,10 +366,10 @@
"\n",
"\n",
"def prune(\n",
" state: ToTState, *, config: RunnableConfig\n",
") -> Dict[str, List[Dict[str, Any]]]:\n",
" state: ToTState, *, runtime: Runtime[Context]\n",
") -> Dict[str, Any]:\n",
" scored_candidates = state[\"scored_candidates\"]\n",
" beam_size = _ensure_configurable(config)[\"beam_size\"]\n",
" beam_size = _ensure_context(runtime.context)[\"beam_size\"]\n",
" organized = sorted(\n",
" scored_candidates, key=lambda candidate: candidate[1], reverse=True\n",
" )\n",
@@ -383,11 +385,11 @@
"\n",
"\n",
"def should_terminate(\n",
" state: ToTState, config: RunnableConfig\n",
" state: ToTState, runtime: Runtime[Context]\n",
") -> Union[Literal[\"__end__\"], Send]:\n",
" configurable = _ensure_configurable(config)\n",
" solved = state[\"candidates\"][0].score >= configurable[\"threshold\"]\n",
" if solved or state[\"depth\"] >= configurable[\"max_depth\"]:\n",
" ctx = _ensure_context(runtime.context)\n",
" solved = state[\"candidates\"][0].score >= ctx[\"threshold\"]\n",
" if solved or state[\"depth\"] >= ctx[\"max_depth\"]:\n",
" return \"__end__\"\n",
" return [\n",
" Send(\"expand\", {**state, \"somevalseed\": candidate})\n",
@@ -396,7 +398,7 @@
"\n",
"\n",
"# Create the graph\n",
"builder = StateGraph(state_schema=ToTState, config_schema=Configuration)\n",
"builder = StateGraph(state_schema=ToTState, context_schema=Context)\n",
"\n",
"# Add nodes\n",
"builder.add_node(expand)\n",
@@ -467,13 +469,7 @@
}
],
"source": [
"config = {\n",
" \"configurable\": {\n",
" \"thread_id\": \"test_1\",\n",
" \"depth\": 10,\n",
" }\n",
"}\n",
"for step in graph.stream({\"problem\": puzzles[42]}, config):\n",
"for step in graph.stream({\"problem\": puzzles[42]}, config={\"configurable\": {\"thread_id\": \"test_1\"}}, context={\"depth\": 10}):\n",
" print(step)"
]
},
@@ -491,7 +487,7 @@
}
],
"source": [
"final_state = graph.get_state(config)\n",
"final_state = graph.get_state({'configurable': {'thread_id': 'test_1'}})\n",
"winning_solution = final_state.values[\"candidates\"][0]\n",
"search_depth = final_state.values[\"depth\"]\n",
"if winning_solution[1] == 1:\n",
+2 -2
View File
@@ -49,12 +49,12 @@ def call_model(state, config):
tool_node = ToolNode(tools)
class ConfigSchema(TypedDict):
class ContextSchema(TypedDict):
model: Literal["anthropic", "openai"]
# Define a new graph
workflow = StateGraph(AgentState, config_schema=ConfigSchema)
workflow = StateGraph(AgentState, context_schema=ContextSchema)
# Define the two nodes we will cycle between
workflow.add_node("agent", call_model)
+87 -70
View File
@@ -48,11 +48,10 @@ from langgraph._internal._config import (
get_callback_manager_for_config,
patch_config,
)
from langgraph._internal._typing import UNSET
from langgraph.constants import (
CONF,
CONFIG_KEY_PREVIOUS,
CONFIG_KEY_STORE,
CONFIG_KEY_STREAM_WRITER,
CONFIG_KEY_RUNTIME,
)
from langgraph.store.base import BaseStore
from langgraph.types import StreamWriter
@@ -128,45 +127,52 @@ ANY_TYPE = object()
ASYNCIO_ACCEPTS_CONTEXT = sys.version_info >= (3, 11)
# List of keyword arguments that can be injected at runtime from the config object.
# List of keyword arguments that can be injected into nodes / tasks / tools at runtime.
# A named argument may appear multiple times if it appears with distinct types.
KWARGS_CONFIG_KEYS: tuple[tuple[str, tuple[Any, ...], str, Any], ...] = (
(
sys.intern("writer"),
"config",
(RunnableConfig, "RunnableConfig", inspect.Parameter.empty),
# for now, use config directly, eventually, will pop off of Runtime
"N/A",
inspect.Parameter.empty,
),
(
"writer",
(StreamWriter, "StreamWriter", inspect.Parameter.empty),
CONFIG_KEY_STREAM_WRITER,
"stream_writer",
lambda _: None,
),
(
# Covers store that is not optional (will raise an error if a store
# cannot be injected).
sys.intern("store"),
"store",
(
BaseStore,
"BaseStore",
inspect.Parameter.empty,
),
CONFIG_KEY_STORE,
"store",
inspect.Parameter.empty,
),
(
# Covers store that is optional. Will set to None if not found in config.
sys.intern("store"),
"store",
(
Optional[BaseStore],
# Best effort to catch some forward references.
# This will not work for cases like `"Union[None, BaseStore]"`,
# we'll need to re-write logic to use get_type_hints()
# to resolve forward references.
"Optional[BaseStore]",
),
CONFIG_KEY_STORE,
"store",
None,
),
(
sys.intern("previous"),
"previous",
(ANY_TYPE,),
CONFIG_KEY_PREVIOUS,
"previous",
inspect.Parameter.empty,
),
(
"runtime",
(ANY_TYPE,),
# we never hit this block, we just inject runtime directly
"N/A",
inspect.Parameter.empty,
),
)
@@ -174,7 +180,7 @@ KWARGS_CONFIG_KEYS: tuple[tuple[str, tuple[Any, ...], str, Any], ...] = (
config keys, default values and type annotations.
Used to configure keyword arguments that can be injected at runtime
from the config object as kwargs to `invoke`, `ainvoke`, `stream` and `astream`.
from the `Runtime` object as kwargs to `invoke`, `ainvoke`, `stream` and `astream`.
For a keyword to be injected from the config object, the function signature
must contain a kwarg with the same name and a matching type annotation.
@@ -182,8 +188,10 @@ must contain a kwarg with the same name and a matching type annotation.
Each tuple contains:
- the name of the kwarg in the function signature
- the type annotation(s) for the kwarg
- the config key to look for the value in
- the default value for the kwarg
- the `Runtime` attribute for fetching the value (N/A if not applicable)
This is fully internal and should be further refactored to use `get_type_hints`
to resolve forward references and optional types formatted like BaseStore | None.
"""
VALID_KINDS = (inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.KEYWORD_ONLY)
@@ -250,7 +258,6 @@ class RunnableCallable(Runnable):
trace: bool = True,
recurse: bool = True,
explode_args: bool = False,
func_accepts_config: bool | None = None,
**kwargs: Any,
) -> None:
self.name = name
@@ -277,31 +284,23 @@ class RunnableCallable(Runnable):
if func is None and afunc is None:
raise ValueError("At least one of func or afunc must be provided.")
if func_accepts_config is not None:
self.func_accepts_config = func_accepts_config
self.func_accepts: dict[str, tuple[str, Any]] = {}
else:
params = inspect.signature(cast(Callable, func or afunc)).parameters
self.func_accepts: dict[str, tuple[str, Any]] = {}
params = inspect.signature(cast(Callable, func or afunc)).parameters
self.func_accepts_config = "config" in params
# Mapping from kwarg name to (config key, default value) to be used.
# The default value is used if the config key is not found in the config.
self.func_accepts = {}
for kw, typ, runtime_key, default in KWARGS_CONFIG_KEYS:
p = params.get(kw)
for kw, typ, config_key, default in KWARGS_CONFIG_KEYS:
p = params.get(kw)
if p is None or p.kind not in VALID_KINDS:
# If parameter is not found or is not a valid kind, skip
continue
if p is None or p.kind not in VALID_KINDS:
# If parameter is not found or is not a valid kind, skip
continue
if typ != (ANY_TYPE,) and p.annotation not in typ:
# A specific type is required, but the function annotation does
# not match the expected type.
continue
if typ != (ANY_TYPE,) and p.annotation not in typ:
# A specific type is required, but the function annotation does
# not match the expected type.
continue
# If the kwarg is accepted by the function, store the default value
self.func_accepts[kw] = (config_key, default)
# If the kwarg is accepted by the function, store the key / runtime attribute to inject
self.func_accepts[kw] = (runtime_key, default)
def __repr__(self) -> str:
repr_args = {
@@ -328,25 +327,33 @@ class RunnableCallable(Runnable):
else:
args = (input,)
kwargs = {**self.kwargs, **kwargs}
if self.func_accepts_config:
kwargs["config"] = config
_conf = config[CONF]
for kw, (config_key, default_value) in self.func_accepts.items():
runtime = config[CONF].get(CONFIG_KEY_RUNTIME)
for kw, (runtime_key, default) in self.func_accepts.items():
# If the kwarg is already set, use the set value
if kw in kwargs:
continue
if (
# If the kwarg is requested, but isn't in the config AND has no
# default value, raise an error
config_key not in _conf and default_value is inspect.Parameter.empty
):
raise ValueError(
f"Missing required config key '{config_key}' for '{self.name}'."
)
kw_value: Any = UNSET
if kw == "config":
kw_value = config
elif runtime:
if kw == "runtime":
kw_value = runtime
else:
try:
kw_value = getattr(runtime, runtime_key)
except AttributeError:
pass
kwargs[kw] = _conf.get(config_key, default_value)
if kw_value is UNSET:
if default is inspect.Parameter.empty:
raise ValueError(
f"Missing required config key '{runtime_key}' for '{self.name}'."
)
kw_value = default
kwargs[kw] = kw_value
if self.trace:
callback_manager = get_callback_manager_for_config(config, self.tags)
@@ -392,23 +399,33 @@ class RunnableCallable(Runnable):
else:
args = (input,)
kwargs = {**self.kwargs, **kwargs}
if self.func_accepts_config:
kwargs["config"] = config
_conf = config[CONF]
for kw, (config_key, default_value) in self.func_accepts.items():
runtime = config[CONF].get(CONFIG_KEY_RUNTIME)
for kw, (runtime_key, default) in self.func_accepts.items():
# If the kwarg has already been set, use the set value
if kw in kwargs:
continue
if (
# If the kwarg is requested, but isn't in the config AND has no
# default value, raise an error
config_key not in _conf and default_value is inspect.Parameter.empty
):
raise ValueError(
f"Missing required config key '{config_key}' for '{self.name}'."
)
kwargs[kw] = _conf.get(config_key, default_value)
kw_value: Any = UNSET
if kw == "config":
kw_value = config
elif runtime:
if kw == "runtime":
kw_value = runtime
else:
try:
kw_value = getattr(runtime, runtime_key)
except AttributeError:
pass
if kw_value is UNSET:
if default is inspect.Parameter.empty:
raise ValueError(
f"Missing required config key '{runtime_key}' for '{self.name}'."
)
kw_value = default
kwargs[kw] = kw_value
if self.trace:
callback_manager = get_async_callback_manager_for_config(config, self.tags)
run_manager = await callback_manager.on_chain_start(
@@ -0,0 +1,39 @@
"""Internal utilities for the Runtime class."""
from __future__ import annotations
from dataclasses import replace
from typing import Any, cast
from typing_extensions import TypedDict, Unpack
from langgraph.runtime import Runtime
from langgraph.store.base import BaseStore
from langgraph.types import StreamWriter
class RuntimePatch(TypedDict, total=False):
"""Patch structure for the Runtime class."""
context: Any
store: BaseStore | None
stream_writer: StreamWriter
previous: Any
def patch_runtime(runtime: Runtime, **overrides: Unpack[RuntimePatch]) -> Runtime:
"""Patch the runtime with the given overrides, returning a new instance."""
return replace(runtime, **overrides)
def patch_runtime_non_null(
runtime: Runtime, **overrides: Unpack[RuntimePatch]
) -> Runtime:
"""Patch the runtime with the given overrides, returning a new instance.
Only patch fields with overrides that are not None.
"""
return replace(
runtime,
**cast(dict[str, Any], {k: v for k, v in overrides.items() if v is not None}),
)
+4 -5
View File
@@ -5,7 +5,7 @@ from typing import Any
from langchain_core.runnables import RunnableConfig
from langchain_core.runnables.config import var_child_runnable_config
from langgraph.constants import CONF, CONFIG_KEY_STORE, CONFIG_KEY_STREAM_WRITER
from langgraph.constants import CONF, CONFIG_KEY_RUNTIME
from langgraph.store.base import BaseStore
from langgraph.types import StreamWriter
@@ -114,8 +114,7 @@ def get_store() -> BaseStore:
3
```
"""
config = get_config()
return config[CONF][CONFIG_KEY_STORE]
return get_config()[CONF][CONFIG_KEY_RUNTIME].store
def get_stream_writer() -> StreamWriter:
@@ -181,5 +180,5 @@ def get_stream_writer() -> StreamWriter:
{'custom_data': 'Hello!'}
```
"""
config = get_config()
return config[CONF].get(CONFIG_KEY_STREAM_WRITER, _no_op_stream_writer)
runtime = get_config()[CONF][CONFIG_KEY_RUNTIME]
return runtime.stream_writer
+3 -8
View File
@@ -79,10 +79,6 @@ CONFIG_KEY_CHECKPOINTER = sys.intern("__pregel_checkpointer")
# holds a `BaseCheckpointSaver` passed from parent graph to child graphs
CONFIG_KEY_STREAM = sys.intern("__pregel_stream")
# holds a `StreamProtocol` passed from parent graph to child graphs
CONFIG_KEY_STREAM_WRITER = sys.intern("__pregel_stream_writer")
# holds a `StreamWriter` for stream_mode=custom
CONFIG_KEY_STORE = sys.intern("__pregel_store")
# holds a `BaseStore` made available to managed values
CONFIG_KEY_CACHE = sys.intern("__pregel_cache")
# holds a `BaseCache` made available to subgraphs
CONFIG_KEY_RESUMING = sys.intern("__pregel_resuming")
@@ -101,12 +97,12 @@ CONFIG_KEY_NODE_FINISHED = sys.intern("__pregel_node_finished")
# holds a callback to be called when a node is finished
CONFIG_KEY_SCRATCHPAD = sys.intern("__pregel_scratchpad")
# holds a mutable dict for temporary storage scoped to the current task
CONFIG_KEY_PREVIOUS = sys.intern("__pregel_previous")
# holds the previous return value from a stateful Pregel graph.
CONFIG_KEY_RUNNER_SUBMIT = sys.intern("__pregel_runner_submit")
# holds a function that receives tasks from runner, executes them and returns results
CONFIG_KEY_CHECKPOINT_DURING = sys.intern("__pregel_checkpoint_during")
# holds a boolean indicating whether to checkpoint during the run (or only at the end)
CONFIG_KEY_RUNTIME = sys.intern("__pregel_runtime")
# holds a `Runtime` instance with context, store, stream writer, etc.
# --- Other constants ---
PUSH = sys.intern("__pregel_push")
@@ -137,8 +133,7 @@ RESERVED = {
CONFIG_KEY_READ,
CONFIG_KEY_CHECKPOINTER,
CONFIG_KEY_STREAM,
CONFIG_KEY_STREAM_WRITER,
CONFIG_KEY_STORE,
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_RESUMING,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_CHECKPOINT_MAP,
+19 -9
View File
@@ -12,6 +12,7 @@ from typing import (
Callable,
Generic,
TypeVar,
cast,
get_args,
get_origin,
overload,
@@ -38,7 +39,8 @@ from langgraph.pregel._read import PregelNode
from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry
from langgraph.store.base import BaseStore
from langgraph.types import _DC_KWARGS, CachePolicy, RetryPolicy, StreamMode
from langgraph.warnings import LangGraphDeprecatedSinceV05
from langgraph.typing import ContextT
from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV10
__all__ = ("task", "entrypoint")
@@ -215,7 +217,7 @@ S = TypeVar("S")
# In this form, the `final` attribute should play nicely with IDE autocompletion,
# and type checking tools.
# In addition, we'll be able to surface this information in the API Reference.
class entrypoint:
class entrypoint(Generic[ContextT]):
"""Define a LangGraph workflow using the `entrypoint` decorator.
### Function signature
@@ -231,10 +233,9 @@ class entrypoint:
| Parameter | Description |
|------------------|----------------------------------------------------------------------------------------------------|
| **`store`** | An instance of [BaseStore][langgraph.store.base.BaseStore]. Useful for long-term memory. |
| **`writer`** | A [StreamWriter][langgraph.types.StreamWriter] instance for writing custom data to a stream. |
| **`config`** | A configuration object (aka RunnableConfig) that holds run-time configuration values. |
| **`previous`** | The previous return value for the given thread (available only when a checkpointer is provided). |
| **`runtime`** | A Runtime object that contains information about the current run, including context, store, writer | |
The entrypoint decorator can be applied to sync functions or async functions.
@@ -254,7 +255,7 @@ class entrypoint:
store: A generalized key-value store. Some implementations may support
semantic search capabilities through an optional `index` configuration.
cache: A cache to use for caching the results of the workflow.
config_schema: Specifies the schema for the configuration object that will be
context_schema: Specifies the schema for the context object that will be
passed to the workflow.
cache_policy: A cache policy to use for caching the results of the workflow.
retry_policy: A retry policy (or list of policies) to use for the workflow in case of a failure.
@@ -376,26 +377,35 @@ class entrypoint:
checkpointer: BaseCheckpointSaver | None = None,
store: BaseStore | None = None,
cache: BaseCache | None = None,
config_schema: type[Any] | None = None,
context_schema: type[ContextT] | None = None,
cache_policy: CachePolicy | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> None:
"""Initialize the entrypoint decorator."""
if (config_schema := kwargs.get("config_schema", UNSET)) is not UNSET:
warnings.warn(
"`config_schema` is deprecated and will be removed. Please use `context_schema` instead.",
category=LangGraphDeprecatedSinceV10,
stacklevel=2,
)
if context_schema is None:
context_schema = cast(type[ContextT], config_schema)
if (retry := kwargs.get("retry", UNSET)) is not UNSET:
warnings.warn(
"`retry` is deprecated and will be removed. Please use `retry_policy` instead.",
category=LangGraphDeprecatedSinceV05,
)
if retry_policy is None:
retry_policy = retry # type: ignore[assignment]
retry_policy = cast("RetryPolicy | Sequence[RetryPolicy]", retry)
self.checkpointer = checkpointer
self.store = store
self.cache = cache
self.cache_policy = cache_policy
self.retry_policy = retry_policy
self.config_schema = config_schema
self.context_schema = context_schema
@dataclass(**_DC_KWARGS)
class final(Generic[R, S]):
@@ -527,5 +537,5 @@ class entrypoint:
cache=self.cache,
cache_policy=self.cache_policy,
retry_policy=self.retry_policy or (),
config_type=self.config_schema,
context_schema=self.context_schema, # type: ignore[arg-type]
)
@@ -134,7 +134,6 @@ class BranchSpec(NamedTuple):
reader=reader,
name=None,
trace=False,
func_accepts_config=True,
),
list(
zip_longest(
+11 -3
View File
@@ -9,9 +9,10 @@ from langchain_core.runnables import Runnable, RunnableConfig
from typing_extensions import TypeAlias
from langgraph.constants import EMPTY_SEQ
from langgraph.runtime import Runtime
from langgraph.store.base import BaseStore
from langgraph.types import CachePolicy, RetryPolicy, StreamWriter
from langgraph.typing import NodeInputT, NodeInputT_contra
from langgraph.typing import ContextT, NodeInputT, NodeInputT_contra
_DC_SLOTS = {"slots": True} if sys.version_info >= (3, 10) else {}
@@ -61,6 +62,12 @@ class _NodeWithConfigWriterStore(Protocol[NodeInputT_contra]):
) -> Any: ...
class _NodeWithRuntime(Protocol[NodeInputT_contra, ContextT]):
def __call__(
self, state: NodeInputT_contra, *, runtime: Runtime[ContextT]
) -> Any: ...
# TODO: we probably don't want to explicitly support the config / store signatures once
# we move to adding a context arg. Maybe what we do is we add support for kwargs with param spec
# this is purely for typing purposes though, so can easily change in the coming weeks.
@@ -73,13 +80,14 @@ StateNode: TypeAlias = Union[
_NodeWithConfigWriter[NodeInputT],
_NodeWithConfigStore[NodeInputT],
_NodeWithConfigWriterStore[NodeInputT],
_NodeWithRuntime[NodeInputT, ContextT],
Runnable[NodeInputT, Any],
]
@dataclass(**_DC_SLOTS)
class StateNodeSpec(Generic[NodeInputT]):
runnable: StateNode[NodeInputT]
class StateNodeSpec(Generic[NodeInputT, ContextT]):
runnable: StateNode[NodeInputT, ContextT]
metadata: dict[str, Any] | None
input_schema: type[NodeInputT]
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None
+59 -39
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
import inspect
import logging
import sys
import typing
import warnings
from collections import defaultdict
@@ -83,8 +84,13 @@ from langgraph.types import (
RetryPolicy,
Send,
)
from langgraph.typing import InputT, NodeInputT, OutputT, StateT
from langgraph.warnings import LangGraphDeprecatedSinceV05
from langgraph.typing import ContextT, InputT, NodeInputT, OutputT, StateT
from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV10
if sys.version_info < (3, 10):
NoneType = type(None)
else:
from types import NoneType as NoneType
__all__ = ("StateGraph", "CompiledStateGraph")
@@ -105,14 +111,14 @@ def _warn_invalid_state_schema(schema: type[Any] | Any) -> None:
)
def _get_node_name(node: StateNode) -> str:
def _get_node_name(node: StateNode[Any, ContextT]) -> str:
try:
return getattr(node, "__name__", node.__class__.__name__)
except AttributeError:
raise TypeError(f"Unsupported node type: {type(node)}")
class StateGraph(Generic[StateT, InputT, OutputT]):
class StateGraph(Generic[StateT, ContextT, InputT, OutputT]):
"""A graph whose nodes communicate by reading and writing to a shared state.
The signature of each node is State -> Partial<State>.
@@ -122,8 +128,10 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
Args:
state_schema: The schema class that defines the state.
config_schema: The schema class that defines the configuration.
Use this to expose configurable parameters in your API.
context_schema: The schema class that defines the runtime context.
Use this to expose immutable context data to your nodes, like user_id, db_conn, etc.
input_schema: The schema class that defines the input to the graph.
output_schema: The schema class that defines the output from the graph.
Example:
```python
@@ -131,6 +139,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
from typing_extensions import Annotated, TypedDict
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import StateGraph
from langgraph.rumtime import Runtime
def reducer(a: list, b: int | None) -> list:
if b is not None:
@@ -140,13 +149,13 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
class State(TypedDict):
x: Annotated[list, reducer]
class ConfigSchema(TypedDict):
class Context(TypedDict):
r: float
graph = StateGraph(State, config_schema=ConfigSchema)
graph = StateGraph(state_schema=State, context_schema=Context)
def node(state: State, config: RunnableConfig) -> dict:
r = config["configurable"].get("r", 1.0)
def node(state: State, runtime: Runtime[Context]) -> dict:
r = runtie.context.get("r", 1.0)
x = state["x"][-1]
next_value = x * r * (1 - x)
return {"x": next_value}
@@ -156,16 +165,13 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
graph.set_finish_point("A")
compiled = graph.compile()
print(compiled.config_specs)
# [ConfigurableFieldSpec(id='r', annotation=<class 'float'>, name=None, description=None, default=None, is_shared=False, dependencies=None)]
step1 = compiled.invoke({"x": 0.5}, {"configurable": {"r": 3.0}})
step1 = compiled.invoke({"x": 0.5}, context={"r": 3.0})
# {'x': [0.5, 0.75]}
```
"""
edges: set[tuple[str, str]]
nodes: dict[str, StateNodeSpec]
nodes: dict[str, StateNodeSpec[Any, ContextT]]
branches: defaultdict[str, dict[str, BranchSpec]]
channels: dict[str, BaseChannel]
managed: dict[str, ManagedValueSpec]
@@ -174,18 +180,28 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
compiled: bool
state_schema: type[StateT]
context_schema: type[ContextT] | None
input_schema: type[InputT]
output_schema: type[OutputT]
def __init__(
self,
state_schema: type[StateT],
config_schema: type[Any] | None = None,
context_schema: type[ContextT] | None = None,
*,
input_schema: type[InputT] | None = None,
output_schema: type[OutputT] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> None:
if (config_schema := kwargs.get("config_schema", UNSET)) is not UNSET:
warnings.warn(
"`config_schema` is deprecated and will be removed. Please use `context_schema` instead.",
category=LangGraphDeprecatedSinceV10,
stacklevel=2,
)
if context_schema is None:
context_schema = cast(type[ContextT], config_schema)
if (input_ := kwargs.get("input", UNSET)) is not UNSET:
warnings.warn(
"`input` is deprecated and will be removed. Please use `input_schema` instead.",
@@ -193,7 +209,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
stacklevel=2,
)
if input_schema is None:
input_schema = cast(Union[type[InputT], None], input_)
input_schema = cast(type[InputT], input_)
if (output := kwargs.get("output", UNSET)) is not UNSET:
warnings.warn(
@@ -202,7 +218,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
stacklevel=2,
)
if output_schema is None:
output_schema = cast(Union[type[OutputT], None], output)
output_schema = cast(type[OutputT], output)
self.nodes = {}
self.edges = set()
@@ -216,7 +232,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
self.state_schema = state_schema
self.input_schema = cast(type[InputT], input_schema or state_schema)
self.output_schema = cast(type[OutputT], output_schema or state_schema)
self.config_schema = config_schema
self.context_schema = context_schema
self._add_schema(self.state_schema)
self._add_schema(self.input_schema, allow_managed=False)
@@ -263,7 +279,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
@overload
def add_node(
self,
node: StateNode[StateT],
node: StateNode[NodeInputT, ContextT],
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
@@ -281,7 +297,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
@overload
def add_node(
self,
node: StateNode[NodeInputT],
node: StateNode[NodeInputT, ContextT],
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
@@ -300,7 +316,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
def add_node(
self,
node: str,
action: StateNode[StateT],
action: StateNode[NodeInputT, ContextT],
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
@@ -316,8 +332,8 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
@overload
def add_node(
self,
node: str,
action: StateNode[NodeInputT],
node: str | StateNode[NodeInputT, ContextT],
action: StateNode[NodeInputT, ContextT] | None = None,
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
@@ -332,8 +348,8 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
def add_node(
self,
node: str | StateNode[StateT] | StateNode[NodeInputT],
action: StateNode[StateT] | StateNode[NodeInputT] | None = None,
node: str | StateNode[NodeInputT, ContextT],
action: StateNode[NodeInputT, ContextT] | None = None,
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
@@ -498,8 +514,8 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
ends = destinations
if input_schema is not None:
self.nodes[node] = StateNodeSpec[NodeInputT](
coerce_to_runnable(action, name=node, trace=False),
self.nodes[node] = StateNodeSpec[NodeInputT, ContextT](
coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]
metadata,
input_schema=input_schema,
retry_policy=retry_policy,
@@ -509,7 +525,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
)
elif inferred_input_schema is not None:
self.nodes[node] = StateNodeSpec(
coerce_to_runnable(action, name=node, trace=False),
coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]
metadata,
input_schema=inferred_input_schema,
retry_policy=retry_policy,
@@ -518,8 +534,8 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
defer=defer,
)
else:
self.nodes[node] = StateNodeSpec[StateT](
coerce_to_runnable(action, name=node, trace=False),
self.nodes[node] = StateNodeSpec[StateT, ContextT](
coerce_to_runnable(action, name=node, trace=False), # type: ignore[arg-type]
metadata,
input_schema=self.state_schema,
retry_policy=retry_policy,
@@ -636,7 +652,10 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
def add_sequence(
self,
nodes: Sequence[StateNode[StateT] | tuple[str, StateNode[StateT]]],
nodes: Sequence[
StateNode[NodeInputT, ContextT]
| tuple[str, StateNode[NodeInputT, ContextT]]
],
) -> Self:
"""Add a sequence of nodes that will be executed in the provided order.
@@ -782,7 +801,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
interrupt_after: All | list[str] | None = None,
debug: bool = False,
name: str | None = None,
) -> CompiledStateGraph[StateT, InputT, OutputT]:
) -> CompiledStateGraph[StateT, ContextT, InputT, OutputT]:
"""Compiles the state graph into a `CompiledStateGraph` object.
The compiled graph implements the `Runnable` interface and can be invoked,
@@ -834,10 +853,10 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
]
)
compiled = CompiledStateGraph[StateT, InputT, OutputT](
compiled = CompiledStateGraph[StateT, ContextT, InputT, OutputT](
builder=self,
schema_to_mapper={},
config_type=self.config_schema,
context_schema=self.context_schema,
nodes={},
channels={
**self.channels,
@@ -876,15 +895,16 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
class CompiledStateGraph(
Pregel[StateT, InputT, OutputT], Generic[StateT, InputT, OutputT]
Pregel[StateT, ContextT, InputT, OutputT],
Generic[StateT, ContextT, InputT, OutputT],
):
builder: StateGraph[StateT, InputT, OutputT]
builder: StateGraph[StateT, ContextT, InputT, OutputT]
schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None]
def __init__(
self,
*,
builder: StateGraph[StateT, InputT, OutputT],
builder: StateGraph[StateT, ContextT, InputT, OutputT],
schema_to_mapper: dict[type[Any], Callable[[Any], Any] | None],
**kwargs: Any,
) -> None:
@@ -912,7 +932,7 @@ class CompiledStateGraph(
name=self.get_name("Output"),
)
def attach_node(self, key: str, node: StateNodeSpec | None) -> None:
def attach_node(self, key: str, node: StateNodeSpec[Any, ContextT] | None) -> None:
if key == START:
output_keys = [
k
+20 -13
View File
@@ -26,6 +26,7 @@ from langchain_core.runnables.config import RunnableConfig
from xxhash import xxh3_128_hexdigest
from langgraph._internal._config import merge_configs, patch_config
from langgraph._internal._runtime import patch_runtime_non_null
from langgraph.channels.base import BaseChannel
from langgraph.channels.topic import Topic
from langgraph.checkpoint.base import (
@@ -42,12 +43,11 @@ from langgraph.constants import (
CONFIG_KEY_CHECKPOINT_MAP,
CONFIG_KEY_CHECKPOINT_NS,
CONFIG_KEY_CHECKPOINTER,
CONFIG_KEY_PREVIOUS,
CONFIG_KEY_READ,
CONFIG_KEY_RESUME_MAP,
CONFIG_KEY_RUNTIME,
CONFIG_KEY_SCRATCHPAD,
CONFIG_KEY_SEND,
CONFIG_KEY_STORE,
CONFIG_KEY_TASK_ID,
EMPTY_SEQ,
ERROR,
@@ -72,6 +72,7 @@ from langgraph.pregel._io import read_channels
from langgraph.pregel._log import logger
from langgraph.pregel._read import INPUT_CACHE_KEY_TYPE, PregelNode
from langgraph.pregel._scratchpad import PregelScratchpad
from langgraph.runtime import DEFAULT_RUNTIME
from langgraph.store.base import BaseStore
from langgraph.types import (
All,
@@ -583,6 +584,10 @@ def prepare_single_task(
step,
stop,
)
runtime = patch_runtime_non_null(
configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME),
store=store,
)
return PregelExecutableTask(
name,
call.input,
@@ -604,7 +609,6 @@ def prepare_single_task(
managed,
PregelTaskWrites(task_path, name, writes, triggers),
),
CONFIG_KEY_STORE: (store or configurable.get(CONFIG_KEY_STORE)),
CONFIG_KEY_CHECKPOINTER: (
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
),
@@ -615,6 +619,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: scratchpad,
CONFIG_KEY_RUNTIME: runtime,
},
),
triggers,
@@ -709,6 +714,11 @@ def prepare_single_task(
step,
stop,
)
runtime = patch_runtime_non_null(
configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME),
store=store,
previous=checkpoint["channel_values"].get(PREVIOUS, None),
)
return PregelExecutableTask(
packet.node,
packet.arg,
@@ -731,7 +741,6 @@ def prepare_single_task(
managed,
PregelTaskWrites(task_path, packet.node, writes, triggers),
),
CONFIG_KEY_STORE: (store or configurable.get(CONFIG_KEY_STORE)),
CONFIG_KEY_CHECKPOINTER: (
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
),
@@ -742,9 +751,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: scratchpad,
CONFIG_KEY_PREVIOUS: checkpoint["channel_values"].get(
PREVIOUS, None
),
CONFIG_KEY_RUNTIME: runtime,
},
),
triggers,
@@ -846,6 +853,11 @@ def prepare_single_task(
)
else:
cache_key = None
runtime = patch_runtime_non_null(
configurable.get(CONFIG_KEY_RUNTIME, DEFAULT_RUNTIME),
previous=checkpoint["channel_values"].get(PREVIOUS, None),
store=store,
)
return PregelExecutableTask(
name,
val,
@@ -877,9 +889,6 @@ def prepare_single_task(
triggers,
),
),
CONFIG_KEY_STORE: (
store or configurable.get(CONFIG_KEY_STORE)
),
CONFIG_KEY_CHECKPOINTER: (
checkpointer
or configurable.get(CONFIG_KEY_CHECKPOINTER)
@@ -891,9 +900,7 @@ def prepare_single_task(
CONFIG_KEY_CHECKPOINT_ID: None,
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
CONFIG_KEY_SCRATCHPAD: scratchpad,
CONFIG_KEY_PREVIOUS: checkpoint["channel_values"].get(
PREVIOUS, None
),
CONFIG_KEY_RUNTIME: runtime,
},
),
triggers,
-1
View File
@@ -46,7 +46,6 @@ class ChannelRead(RunnableCallable):
tags=tags,
name=None,
trace=False,
func_accepts_config=True,
)
self.fresh = fresh
self.mapper = mapper
@@ -64,7 +64,6 @@ class ChannelWrite(RunnableCallable):
name=None,
tags=tags,
trace=False,
func_accepts_config=True,
)
self.writes = cast(
list[Union[ChannelWriteEntry, ChannelWriteTupleEntry, Send]], writes
+134 -38
View File
@@ -4,10 +4,13 @@ import asyncio
import concurrent
import concurrent.futures
import queue
import warnings
import weakref
from collections import defaultdict, deque
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from dataclasses import is_dataclass
from functools import partial
from inspect import isclass
from typing import Any, Callable, Generic, Union, cast, get_type_hints
from uuid import UUID, uuid5
@@ -22,8 +25,8 @@ from langchain_core.runnables.config import (
get_callback_manager_for_config,
)
from langchain_core.runnables.graph import Graph
from pydantic import BaseModel
from typing_extensions import Self
from pydantic import BaseModel, TypeAdapter
from typing_extensions import Self, Unpack, deprecated, is_typeddict
from langgraph._internal._config import (
ensure_config,
@@ -44,6 +47,7 @@ from langgraph._internal._runnable import (
RunnableSeq,
coerce_to_runnable,
)
from langgraph._internal._typing import DeprecatedKwargs
from langgraph.cache.base import BaseCache
from langgraph.channels.base import BaseChannel
from langgraph.channels.topic import Topic
@@ -64,10 +68,9 @@ from langgraph.constants import (
CONFIG_KEY_NODE_FINISHED,
CONFIG_KEY_READ,
CONFIG_KEY_RUNNER_SUBMIT,
CONFIG_KEY_RUNTIME,
CONFIG_KEY_SEND,
CONFIG_KEY_STORE,
CONFIG_KEY_STREAM,
CONFIG_KEY_STREAM_WRITER,
CONFIG_KEY_TASK_ID,
CONFIG_KEY_THREAD_ID,
END,
@@ -113,6 +116,7 @@ from langgraph.pregel._validate import validate_graph, validate_keys
from langgraph.pregel._write import ChannelWrite, ChannelWriteEntry
from langgraph.pregel.debug import get_bolded_text, get_colored_text, tasks_w_writes
from langgraph.pregel.protocol import PregelProtocol, StreamChunk, StreamProtocol
from langgraph.runtime import Runtime
from langgraph.store.base import BaseStore
from langgraph.types import (
All,
@@ -125,7 +129,8 @@ from langgraph.types import (
StateUpdate,
StreamMode,
)
from langgraph.typing import InputT, OutputT, StateT
from langgraph.typing import ContextT, InputT, OutputT, StateT
from langgraph.warnings import LangGraphDeprecatedSinceV10
try:
from langchain_core.tracers._streaming import _StreamingCallbackHandler
@@ -299,7 +304,10 @@ class NodeBuilder:
)
class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, OutputT]):
class Pregel(
PregelProtocol[StateT, ContextT, InputT, OutputT],
Generic[StateT, ContextT, InputT, OutputT],
):
"""Pregel manages the runtime behavior for LangGraph applications.
## Overview
@@ -592,7 +600,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
"""Cache policy to use for all nodes. Can be overridden by individual nodes.
Defaults to None."""
config_type: type[Any] | None = None
context_schema: type[ContextT] | None = None
config: RunnableConfig | None = None
@@ -620,11 +628,22 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
cache: BaseCache | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] = (),
cache_policy: CachePolicy | None = None,
config_type: type[Any] | None = None,
context_schema: type[ContextT] | None = None,
config: RunnableConfig | None = None,
trigger_to_nodes: Mapping[str, Sequence[str]] | None = None,
name: str = "LangGraph",
**deprecated_kwargs: Unpack[DeprecatedKwargs],
) -> None:
if config_type := deprecated_kwargs.get("config_type"):
warnings.warn(
"`config_type` is deprecated and will be removed. Please use `context_schema` instead.",
category=LangGraphDeprecatedSinceV10,
stacklevel=2,
)
if context_schema is None:
context_schema = cast(type[ContextT], config_type)
self.nodes = {
k: v.build() if isinstance(v, NodeBuilder) else v for k, v in nodes.items()
}
@@ -651,7 +670,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
(retry_policy,) if isinstance(retry_policy, RetryPolicy) else retry_policy
)
self.cache_policy = cache_policy
self.config_type = config_type
self.context_schema = context_schema
self.config = config
self.trigger_to_nodes = trigger_to_nodes or {}
self.name = name
@@ -760,10 +779,23 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
self.trigger_to_nodes = _trigger_to_nodes(self.nodes)
return self
@deprecated(
"`config_schema` is deprecated. Use `get_context_jsonschema` for the relevant schema instead."
)
def config_schema(self, *, include: Sequence[str] | None = None) -> type[BaseModel]:
warnings.warn(
"`config_schema` is deprecated. Use `get_context_jsonschema` for the relevant schema instead.",
category=LangGraphDeprecatedSinceV10,
stacklevel=2,
)
include = include or []
fields = {
**({"configurable": (self.config_type, None)} if self.config_type else {}),
**(
{"configurable": (self.context_schema, None)}
if self.context_schema
else {}
),
**{
field_name: (field_type, None)
for field_name, field_type in get_type_hints(RunnableConfig).items()
@@ -772,12 +804,36 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
}
return create_model(self.get_name("Config"), field_definitions=fields)
@deprecated(
"`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead."
)
def get_config_jsonschema(
self, *, include: Sequence[str] | None = None
) -> dict[str, Any]:
schema = self.config_schema(include=include)
warnings.warn(
"`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead.",
category=LangGraphDeprecatedSinceV10,
stacklevel=2,
)
with warnings.catch_warnings():
warnings.filterwarnings("ignore", category=LangGraphDeprecatedSinceV10)
schema = self.config_schema(include=include)
return schema.model_json_schema()
def get_context_jsonschema(self) -> dict[str, Any] | None:
if (context_schema := self.context_schema) is None:
return None
if isclass(context_schema) and issubclass(context_schema, BaseModel):
return context_schema.model_json_schema()
elif is_typeddict(context_schema) or is_dataclass(context_schema):
return TypeAdapter(context_schema).json_schema()
else:
raise ValueError(
f"Invalid context schema type: {context_schema}. Must be a BaseModel, TypedDict or dataclass."
)
@property
def InputType(self) -> Any:
if isinstance(self.input_channels, str):
@@ -2327,8 +2383,8 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
"Checkpointer requires one or more of the following 'configurable' "
"keys: thread_id, checkpoint_ns, checkpoint_id"
)
if CONFIG_KEY_STORE in config.get(CONF, {}):
store: BaseStore | None = config[CONF][CONFIG_KEY_STORE]
if CONFIG_KEY_RUNTIME in config.get(CONF, {}):
store: BaseStore | None = config[CONF][CONFIG_KEY_RUNTIME].store
else:
store = self.store
if CONFIG_KEY_CACHE in config.get(CONF, {}):
@@ -2350,6 +2406,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
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,
@@ -2446,28 +2503,39 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
run_manager.inheritable_handlers.append(
StreamMessagesHandler(stream.put, subgraphs)
)
# set up custom stream mode
if "custom" in stream_modes:
config[CONF][CONFIG_KEY_STREAM_WRITER] = lambda c: stream.put(
(
tuple(
get_config()[CONF][CONFIG_KEY_CHECKPOINT_NS].split(NS_SEP)[
:-1
]
),
"custom",
c,
def stream_writer(c: Any) -> None:
stream.put(
(
tuple(
get_config()[CONF][CONFIG_KEY_CHECKPOINT_NS].split(
NS_SEP
)[:-1]
),
"custom",
c,
)
)
)
elif (
CONFIG_KEY_STREAM not in config[CONF]
and CONFIG_KEY_STREAM_WRITER in config[CONF]
):
# remove parent graph stream writer if subgraph streaming not requested
del config[CONF][CONFIG_KEY_STREAM_WRITER]
elif CONFIG_KEY_STREAM in config[CONF]:
stream_writer = config[CONF][CONFIG_KEY_RUNTIME].stream_writer
else:
def stream_writer(c: Any) -> None:
pass
# set checkpointing mode for subgraphs
if checkpoint_during is not None:
config[CONF][CONFIG_KEY_CHECKPOINT_DURING] = checkpoint_during
config[CONF][CONFIG_KEY_RUNTIME] = Runtime(
context=context,
store=store,
stream_writer=stream_writer,
previous=None,
)
with SyncPregelLoop(
input,
stream=StreamProtocol(stream.put, stream_modes),
@@ -2572,6 +2640,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
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,
@@ -2686,10 +2755,26 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
run_manager.inheritable_handlers.append(
StreamMessagesHandler(stream_put, subgraphs)
)
# set up custom stream mode
def stream_writer(c: Any) -> None:
aioloop.call_soon_threadsafe(
stream.put_nowait,
(
tuple(
get_config()[CONF][CONFIG_KEY_CHECKPOINT_NS].split(NS_SEP)[
:-1
]
),
"custom",
c,
),
)
if "custom" in stream_modes:
config[CONF][CONFIG_KEY_STREAM_WRITER] = (
lambda c: aioloop.call_soon_threadsafe(
def stream_writer(c: Any) -> None:
aioloop.call_soon_threadsafe(
stream.put_nowait,
(
tuple(
@@ -2701,16 +2786,23 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
c,
),
)
)
elif (
CONFIG_KEY_STREAM not in config[CONF]
and CONFIG_KEY_STREAM_WRITER in config[CONF]
):
# remove parent graph stream writer if subgraph streaming not requested
del config[CONF][CONFIG_KEY_STREAM_WRITER]
elif CONFIG_KEY_STREAM in config[CONF]:
stream_writer = config[CONF][CONFIG_KEY_RUNTIME].stream_writer
else:
def stream_writer(c: Any) -> None:
pass
# set checkpointing mode for subgraphs
if checkpoint_during is not None:
config[CONF][CONFIG_KEY_CHECKPOINT_DURING] = checkpoint_during
config[CONF][CONFIG_KEY_RUNTIME] = Runtime(
context=context,
store=store,
stream_writer=stream_writer,
previous=None,
)
async with AsyncPregelLoop(
input,
stream=StreamProtocol(stream.put_nowait, stream_modes),
@@ -2816,6 +2908,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
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,
@@ -2848,6 +2941,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
for chunk in self.stream(
input,
config,
context=context,
stream_mode=["updates", "values"]
if stream_mode == "values"
else stream_mode,
@@ -2891,6 +2985,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
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,
@@ -2924,6 +3019,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
async for chunk in self.astream(
input,
config,
context=context,
stream_mode=["updates", "values"]
if stream_mode == "values"
else stream_mode,
+6 -2
View File
@@ -9,12 +9,12 @@ 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.typing import InputT, OutputT, StateT
from langgraph.typing import ContextT, InputT, OutputT, StateT
__all__ = ("PregelProtocol", "StreamProtocol")
class PregelProtocol(Runnable[InputT, Any], Generic[StateT, InputT, OutputT]):
class PregelProtocol(Runnable[InputT, Any], Generic[StateT, ContextT, InputT, OutputT]):
@abstractmethod
def with_config(
self, config: RunnableConfig | None = None, **kwargs: Any
@@ -102,6 +102,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, InputT, OutputT]):
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,
@@ -114,6 +115,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, InputT, OutputT]):
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,
@@ -126,6 +128,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, InputT, OutputT]):
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,
) -> dict[str, Any] | Any: ...
@@ -136,6 +139,7 @@ class PregelProtocol(Runnable[InputT, Any], Generic[StateT, InputT, OutputT]):
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,
) -> dict[str, Any] | Any: ...
+55
View File
@@ -0,0 +1,55 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Generic, cast
from langgraph.config import get_config
from langgraph.constants import CONF, CONFIG_KEY_RUNTIME
from langgraph.store.base import BaseStore
from langgraph.types import _DC_KWARGS, StreamWriter
from langgraph.typing import ContextT
def _no_op_stream_writer(_: Any) -> None: ...
@dataclass(**_DC_KWARGS)
class Runtime(Generic[ContextT]):
"""Convenience class that bundles run-scoped context and graph configuration.
!!! version-added "Added in version 1.0.0."
"""
context: ContextT
"""Static context for the graph run, like user_id, db_conn, etc.
Can also be thought of as 'run dependencies'."""
store: BaseStore | None
"""Store for the graph run, enabling persistence and memory."""
stream_writer: StreamWriter
"""Function that writes to the custom stream."""
previous: Any | None
"""The previous return value for the given thread.
Only available with the functional API when a checkpointer is provided."""
DEFAULT_RUNTIME = Runtime(
context=None,
store=None,
stream_writer=_no_op_stream_writer,
previous=None,
)
def get_runtime(context_schema: type[ContextT] | None = None) -> Runtime[ContextT]:
"""Get the runtime for the current graph run."""
# TODO: in an ideal world, we would have a context manager for
# the runtime that's independent of the config. this will follow
# from the removal of the configurable packing
runtime = cast(Runtime[ContextT], get_config()[CONF].get(CONFIG_KEY_RUNTIME))
return runtime
+6 -8
View File
@@ -1,9 +1,9 @@
from __future__ import annotations
import dataclasses
import sys
from collections import deque
from collections.abc import Hashable, Sequence
from dataclasses import asdict, dataclass
from typing import (
TYPE_CHECKING,
Any,
@@ -122,7 +122,7 @@ class RetryPolicy(NamedTuple):
KeyFuncT = TypeVar("KeyFuncT", bound=Callable[..., Union[str, bytes]])
@dataclasses.dataclass(**_DC_KWARGS)
@dataclass(**_DC_KWARGS)
class CachePolicy(Generic[KeyFuncT]):
"""Configuration for caching nodes."""
@@ -138,7 +138,7 @@ _DEFAULT_INTERRUPT_ID = "placeholder-id"
@final
@dataclasses.dataclass(init=False, **_DC_SLOTS)
@dataclass(init=False, **_DC_SLOTS)
class Interrupt:
"""Information about an interrupt that occurred in a node.
@@ -218,7 +218,7 @@ class CacheKey(NamedTuple):
"""Time to live for the cache entry in seconds."""
@dataclasses.dataclass(**_T_DC_KWARGS)
@dataclass(**_T_DC_KWARGS)
class PregelExecutableTask:
name: str
input: Any
@@ -329,7 +329,7 @@ class Send:
N = TypeVar("N", bound=Hashable)
@dataclasses.dataclass(**_DC_KWARGS)
@dataclass(**_DC_KWARGS)
class Command(Generic[N], ToolOutputMixin):
"""One or more commands to update the graph's state and send messages to nodes.
@@ -362,9 +362,7 @@ class Command(Generic[N], ToolOutputMixin):
def __repr__(self) -> str:
# get all non-None values
contents = ", ".join(
f"{key}={value!r}"
for key, value in dataclasses.asdict(self).items()
if value
f"{key}={value!r}" for key, value in asdict(self).items() if value
)
return f"Command({contents})"
+17 -2
View File
@@ -12,6 +12,7 @@ __all__ = (
"StateT_contra",
"InputT",
"OutputT",
"ContextT",
)
StateT = TypeVar("StateT", bound=StateLike)
@@ -21,15 +22,29 @@ StateT_co = TypeVar("StateT_co", bound=StateLike, covariant=True)
StateT_contra = TypeVar("StateT_contra", bound=StateLike, contravariant=True)
ContextT = TypeVar("ContextT", bound=Union[StateLike, None], default=None)
"""Type variable used to represent graph run scoped context.
Defaults to `None`.
"""
ContextT_contra = TypeVar(
"ContextT_contra", bound=Union[StateLike, None], contravariant=True, default=None
)
InputT = TypeVar("InputT", bound=StateLike, default=StateT)
"""Type variable used to represent the input to a state graph.
Defaults to `StateT`.
"""
OutputT = TypeVar("OutputT", bound=Union[StateLike, None], default=StateT)
"""Type variable used to represent the output of a state graph."""
OutputT = TypeVar("OutputT", bound=StateLike, default=StateT)
"""Type variable used to represent the output of a state graph.
Defaults to `StateT`.
"""
NodeInputT = TypeVar("NodeInputT", bound=StateLike)
"""Type variable used to represent the input to a node."""
NodeInputT_contra = TypeVar("NodeInputT_contra", bound=StateLike, contravariant=True)
@@ -825,7 +825,7 @@
'''
# ---
# name: test_state_graph_w_config_inherited_state_keys
'{"$defs": {"Config": {"properties": {"tools": {"items": {"type": "string"}, "title": "Tools", "type": "array"}}, "title": "Config", "type": "object"}}, "properties": {"configurable": {"$ref": "#/$defs/Config", "default": null}}, "title": "LangGraphConfig", "type": "object"}'
'{"properties": {"tools": {"items": {"type": "string"}, "title": "Tools", "type": "array"}}, "title": "Context", "type": "object"}'
# ---
# name: test_state_graph_w_config_inherited_state_keys.1
'{"$defs": {"AgentAction": {"description": "Represents a request to execute an action by an agent.\\n\\nThe action consists of the name of the tool to execute and the input to pass\\nto the tool. The log is used to pass along extra information about the action.", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"anyOf": [{"type": "string"}, {"additionalProperties": true, "type": "object"}], "title": "Tool Input"}, "log": {"title": "Log", "type": "string"}, "type": {"const": "AgentAction", "default": "AgentAction", "title": "Type", "type": "string"}}, "required": ["tool", "tool_input", "log"], "title": "AgentAction", "type": "object"}, "AgentFinish": {"description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "properties": {"return_values": {"additionalProperties": true, "title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"const": "AgentFinish", "default": "AgentFinish", "title": "Type", "type": "string"}}, "required": ["return_values", "log"], "title": "AgentFinish", "type": "object"}}, "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"anyOf": [{"$ref": "#/$defs/AgentAction"}, {"$ref": "#/$defs/AgentFinish"}, {"type": "null"}], "title": "Agent Outcome"}, "intermediate_steps": {"items": {"maxItems": 2, "minItems": 2, "prefixItems": [{"$ref": "#/$defs/AgentAction"}, {"type": "string"}], "type": "array"}, "title": "Intermediate Steps", "type": "array"}}, "required": ["input", "agent_outcome"], "title": "AgentState", "type": "object"}'
+52 -1
View File
@@ -1,9 +1,12 @@
import pytest
from pytest_mock import MockerFixture
from typing_extensions import TypedDict
from langgraph.channels.last_value import LastValue
from langgraph.errors import NodeInterrupt
from langgraph.func import entrypoint, task
from langgraph.graph import StateGraph
from langgraph.pregel import NodeBuilder, Pregel
from langgraph.types import Interrupt, RetryPolicy
from langgraph.warnings import LangGraphDeprecatedSinceV05, LangGraphDeprecatedSinceV10
@@ -83,7 +86,7 @@ def test_constants_deprecation() -> None:
from langgraph.constants import Interrupt # noqa: F401
def test_pregel_deprecation() -> None:
def test_pregel_types_deprecation() -> None:
with pytest.warns(
LangGraphDeprecatedSinceV10,
match="Importing from langgraph.pregel.types is deprecated. Please use 'from langgraph.types import ...' instead.",
@@ -91,6 +94,54 @@ def test_pregel_deprecation() -> None:
from langgraph.pregel.types import StateSnapshot # noqa: F401
@pytest.mark.filterwarnings("ignore:`config_schema` is deprecated")
@pytest.mark.filterwarnings("ignore:`get_config_jsonschema` is deprecated")
def test_config_schema_deprecation() -> None:
with pytest.warns(
LangGraphDeprecatedSinceV10,
match="`config_schema` is deprecated and will be removed. Please use `context_schema` instead.",
):
builder = StateGraph(PlainState, config_schema=PlainState)
builder.add_node("test_node", lambda state: state)
builder.set_entry_point("test_node")
graph = builder.compile()
with pytest.warns(
LangGraphDeprecatedSinceV10,
match="`config_schema` is deprecated. Use `get_context_jsonschema` for the relevant schema instead.",
):
graph.config_schema()
with pytest.warns(
LangGraphDeprecatedSinceV10,
match="`get_config_jsonschema` is deprecated. Use `get_context_jsonschema` instead.",
):
graph.get_config_jsonschema()
def test_config_type_deprecation_pregel(mocker: MockerFixture) -> None:
add_one = mocker.Mock(side_effect=lambda x: x + 1)
chain = NodeBuilder().subscribe_only("input").do(add_one).write_to("output")
with pytest.warns(
LangGraphDeprecatedSinceV10,
match="`config_type` is deprecated and will be removed. Please use `context_schema` instead.",
):
Pregel(
nodes={
"one": chain,
},
channels={
"input": LastValue(int),
"output": LastValue(int),
},
input_channels="input",
output_channels="output",
config_type=PlainState,
)
def test_interrupt_attributes_deprecation() -> None:
interrupt = Interrupt(value="question", id="abc")
+25 -55
View File
@@ -7,7 +7,6 @@ import operator
import threading
import time
import uuid
import warnings
from collections import Counter, deque
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor
@@ -190,7 +189,7 @@ def test_checkpoint_errors() -> None:
)
def test_config_json_schema() -> None:
def test_context_json_schema() -> None:
"""Test that config json schema is generated properly."""
chain = NodeBuilder().subscribe_only("input").write_to("output")
@@ -210,37 +209,25 @@ def test_config_json_schema() -> None:
},
input_channels=["input", "ephemeral"],
output_channels="output",
config_type=Foo,
context_schema=Foo,
)
assert app.get_config_jsonschema() == {
"$defs": {
"Foo": {
"properties": {
"x": {
"title": "X",
"type": "integer",
},
"y": {
"default": "foo",
"title": "Y",
"type": "string",
},
},
"required": [
"x",
],
"title": "Foo",
"type": "object",
},
},
assert app.get_context_jsonschema() == {
"properties": {
"configurable": {
"$ref": "#/$defs/Foo",
"default": None,
"x": {
"title": "X",
"type": "integer",
},
"y": {
"default": "foo",
"title": "Y",
"type": "string",
},
},
"title": "LangGraphConfig",
"required": [
"x",
],
"title": "Foo",
"type": "object",
}
@@ -423,13 +410,7 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
"title": "LangGraphOutput",
"type": "integer",
}
with warnings.catch_warnings():
warnings.simplefilter("error") # raise warnings as errors
assert app.config_schema().model_json_schema() == {
"properties": {},
"title": "LangGraphConfig",
"type": "object",
}
assert app.get_context_jsonschema() is None
assert app.invoke(2) == 3
assert app.invoke(2, output_keys=["output"]) == {"output": 3}
@@ -1227,7 +1208,7 @@ def test_imp_task(
) -> None:
mapper_calls = 0
class Configurable(TypedDict):
class Context(TypedDict):
model: str
@task()
@@ -1237,7 +1218,7 @@ def test_imp_task(
time.sleep(input / 100)
return str(input) * 2
@entrypoint(checkpointer=sync_checkpointer, config_schema=Configurable)
@entrypoint(checkpointer=sync_checkpointer, context_schema=Context)
def graph(input: list[int]) -> list[str]:
futures = [mapper(i) for i in input]
mapped = [f.result() for f in futures]
@@ -1254,21 +1235,10 @@ def test_imp_task(
"items": {"type": "string"},
"title": "LangGraphOutput",
}
assert graph.get_config_jsonschema() == {
"$defs": {
"Configurable": {
"properties": {
"model": {"title": "Model", "type": "string"},
},
"required": ["model"],
"title": "Configurable",
"type": "object",
}
},
"properties": {
"configurable": {"$ref": "#/$defs/Configurable", "default": None}
},
"title": "LangGraphConfig",
assert graph.get_context_jsonschema() == {
"properties": {"model": {"title": "Model", "type": "string"}},
"required": ["model"],
"title": "Context",
"type": "object",
}
@@ -1770,7 +1740,7 @@ def test_state_graph_w_config_inherited_state_keys(snapshot: SnapshotAssertion)
"intermediate_steps",
}
class Config(TypedDict, total=False):
class Context(TypedDict, total=False):
tools: list[str]
# Assemble the tools
@@ -1827,7 +1797,7 @@ def test_state_graph_w_config_inherited_state_keys(snapshot: SnapshotAssertion)
return "continue"
# Define a new graph
builder = StateGraph(AgentState, Config)
builder = StateGraph(AgentState, Context)
builder.add_node("agent", agent)
builder.add_node("tools", execute_tools)
@@ -1842,7 +1812,7 @@ def test_state_graph_w_config_inherited_state_keys(snapshot: SnapshotAssertion)
app = builder.compile()
assert json.dumps(app.config_schema().model_json_schema()) == snapshot
assert json.dumps(app.get_context_jsonschema()) == snapshot
assert json.dumps(app.get_input_jsonschema()) == snapshot
assert json.dumps(app.get_output_jsonschema()) == snapshot
+122 -11
View File
@@ -5,6 +5,7 @@ from typing import Any, Optional
import pytest
from langgraph._internal._runnable import RunnableCallable
from langgraph.runtime import Runtime
from langgraph.store.base import BaseStore
from langgraph.types import StreamWriter
@@ -90,7 +91,22 @@ def test_runnable_callable_injectable_arguments() -> None:
assert store is None
return "success"
assert RunnableCallable(func_optional_store).invoke({"x": "1"}) == "success"
assert (
RunnableCallable(func_optional_store).invoke(
{"x": "1"},
config={
"configurable": {
"__pregel_runtime": Runtime(
store=None,
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
# Test BaseStore annotation
def func_required_store(inputs: Any, store: BaseStore) -> str:
@@ -108,7 +124,17 @@ def test_runnable_callable_injectable_arguments() -> None:
# Specify a value for store in the config
assert (
RunnableCallable(func_required_store).invoke(
{}, config={"configurable": {"__pregel_store": None}}
{},
config={
"configurable": {
"__pregel_runtime": Runtime(
store=None,
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -118,7 +144,16 @@ def test_runnable_callable_injectable_arguments() -> None:
RunnableCallable(func_optional_store).invoke(
{"x": "1"},
store=None,
config={"configurable": {"__pregel_store": "foobar"}},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar", # type: ignore[assignment]
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -134,7 +169,17 @@ def test_runnable_callable_injectable_arguments() -> None:
assert (
RunnableCallable(func_required_store_v2).invoke(
{}, config={"configurable": {"__pregel_store": "foobar"}}
{},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar", # type: ignore[assignment]
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -143,7 +188,16 @@ def test_runnable_callable_injectable_arguments() -> None:
# And manual override takes precedence.
{},
store="foobar",
config={"configurable": {"__pregel_store": "barbar"}},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar", # type: ignore[assignment]
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
@@ -192,7 +246,9 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
assert (
await RunnableCallable(
func=func_required_store, afunc=afunc_required_store
).ainvoke({})
).ainvoke(
{},
)
== "success"
)
@@ -200,7 +256,20 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
assert (
await RunnableCallable(
func=func_required_store, afunc=afunc_required_store
).ainvoke({}, store=None)
).ainvoke(
{},
store=None,
config={
"configurable": {
"__pregel_runtime": Runtime(
store=None,
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -208,7 +277,19 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
assert (
await RunnableCallable(
func=func_required_store, afunc=afunc_required_store
).ainvoke({}, config={"configurable": {"__pregel_store": None}})
).ainvoke(
{},
config={
"configurable": {
"__pregel_runtime": Runtime(
store=None,
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -219,7 +300,16 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
).ainvoke(
{"x": "1"},
store=None,
config={"configurable": {"__pregel_store": "foobar"}},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar",
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -244,7 +334,19 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
assert (
await RunnableCallable(
func=func_required_store_v2, afunc=afunc_required_store_v2
).ainvoke({}, config={"configurable": {"__pregel_store": "foobar"}})
).ainvoke(
{},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar",
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
@@ -255,7 +357,16 @@ async def test_runnable_callable_injectable_arguments_async() -> None:
# And manual override takes precedence.
{},
store="foobar",
config={"configurable": {"__pregel_store": "barbar"}},
config={
"configurable": {
"__pregel_runtime": Runtime(
store="foobar",
context=None,
stream_writer=lambda _: None,
previous=None,
)
}
},
)
== "success"
)
+47
View File
@@ -0,0 +1,47 @@
from dataclasses import dataclass
from typing import Any
from typing_extensions import TypedDict
from langgraph.graph import END, START, StateGraph
from langgraph.runtime import Runtime, get_runtime
@dataclass
class Context:
api_key: str
class State(TypedDict):
message: str
def test_injected_runtime() -> None:
def injected_runtime(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
return {"message": f"api key: {runtime.context.api_key}"}
graph = StateGraph(state_schema=State, context_schema=Context)
graph.add_node("injected_runtime", injected_runtime)
graph.add_edge(START, "injected_runtime")
graph.add_edge("injected_runtime", END)
compiled = graph.compile()
result = compiled.invoke(
{"message": "hello world"}, context=Context(api_key="sk_123456")
)
assert result == {"message": "api key: sk_123456"}
def test_context_runtime() -> None:
def context_runtime(state: State) -> dict[str, Any]:
runtime = get_runtime(Context)
return {"message": f"api key: {runtime.context.api_key}"}
graph = StateGraph(state_schema=State, context_schema=Context)
graph.add_node("context_runtime", context_runtime)
graph.add_edge(START, "context_runtime")
graph.add_edge("context_runtime", END)
compiled = graph.compile()
result = compiled.invoke(
{"message": "hello world"}, context=Context(api_key="sk_123456")
)
assert result == {"message": "api key: sk_123456"}
@@ -11,6 +11,7 @@ from typing import (
cast,
get_type_hints,
)
from warnings import warn
from langchain_core.language_models import (
BaseChatModel,
@@ -35,6 +36,7 @@ from pydantic import BaseModel
from typing_extensions import Annotated, TypedDict
from langgraph._internal._runnable import RunnableCallable, RunnableLike
from langgraph._internal._typing import UNSET
from langgraph.errors import ErrorCode, create_error_message
from langgraph.graph import END, StateGraph
from langgraph.graph.message import add_messages
@@ -43,6 +45,7 @@ from langgraph.managed import IsLastStep, RemainingSteps
from langgraph.prebuilt.tool_node import ToolNode
from langgraph.store.base import BaseStore
from langgraph.types import Checkpointer, Send
from langgraph.warnings import LangGraphDeprecatedSinceV10
StructuredResponse = Union[dict, BaseModel]
StructuredResponseSchema = Union[dict, type[BaseModel]]
@@ -252,7 +255,7 @@ def create_react_agent(
pre_model_hook: Optional[RunnableLike] = None,
post_model_hook: Optional[RunnableLike] = None,
state_schema: Optional[StateSchemaType] = None,
config_schema: Optional[Type[Any]] = None,
context_schema: Optional[Type[Any]] = None,
checkpointer: Optional[Checkpointer] = None,
store: Optional[BaseStore] = None,
interrupt_before: Optional[list[str]] = None,
@@ -260,6 +263,7 @@ def create_react_agent(
debug: bool = False,
version: Literal["v1", "v2"] = "v2",
name: Optional[str] = None,
**deprecated_kwargs: Any,
) -> CompiledStateGraph:
"""Creates an agent graph that calls tools in a loop until a stopping condition is met.
@@ -334,8 +338,7 @@ def create_react_agent(
state_schema: An optional state schema that defines graph state.
Must have `messages` and `remaining_steps` keys.
Defaults to `AgentState` that defines those two keys.
config_schema: An optional schema for configuration.
Use this to expose configurable parameters via agent.config_specs.
context_schema: An optional schema for runtime context.
checkpointer: An optional checkpoint saver object. This is used for persisting
the state of the graph (e.g., as chat memory) for a single thread (e.g., a single conversation).
store: An optional store object. This is used for persisting data
@@ -402,6 +405,15 @@ def create_react_agent(
print(chunk)
```
"""
if (config_schema := deprecated_kwargs.pop("config_schema", UNSET)) is not UNSET:
warn(
"`config_schema` is no longer supported. Use `context_schema` instead.",
category=LangGraphDeprecatedSinceV10,
)
if context_schema is not None:
context_schema = config_schema
if version not in ("v1", "v2"):
raise ValueError(
f"Invalid version {version}. Supported versions are 'v1' and 'v2'."
@@ -590,7 +602,7 @@ def create_react_agent(
if not tool_calling_enabled:
# Define a new graph
workflow = StateGraph(state_schema, config_schema=config_schema)
workflow = StateGraph(state_schema=state_schema, context_schema=context_schema)
workflow.add_node(
"agent",
RunnableCallable(call_model, acall_model),
@@ -657,7 +669,9 @@ def create_react_agent(
return [Send("tools", [tool_call]) for tool_call in tool_calls]
# Define a new graph
workflow = StateGraph(state_schema or AgentState, config_schema=config_schema)
workflow = StateGraph(
state_schema=state_schema or AgentState, context_schema=context_schema
)
# Define the two nodes we will cycle between
workflow.add_node(