mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
Merge branch 'main' into open-swe/e61ca3a1-8559-4dcc-a859-476b55a12775
This commit is contained in:
@@ -2571,7 +2571,7 @@ class Pregel(
|
||||
config[CONF][CONFIG_KEY_DURABILITY] = durability_
|
||||
|
||||
runtime = Runtime(
|
||||
context=context,
|
||||
context=_coerce_context(self.context_schema, context),
|
||||
store=store,
|
||||
stream_writer=stream_writer,
|
||||
previous=None,
|
||||
@@ -2866,7 +2866,7 @@ class Pregel(
|
||||
config[CONF][CONFIG_KEY_DURABILITY] = durability_
|
||||
|
||||
runtime = Runtime(
|
||||
context=context,
|
||||
context=_coerce_context(self.context_schema, context),
|
||||
store=store,
|
||||
stream_writer=stream_writer,
|
||||
previous=None,
|
||||
@@ -3235,3 +3235,33 @@ def _output(
|
||||
yield (ns, payload)
|
||||
else:
|
||||
yield payload
|
||||
|
||||
|
||||
def _coerce_context(
|
||||
context_schema: type[ContextT] | None, context: Any
|
||||
) -> ContextT | None:
|
||||
"""Coerce context input to the appropriate schema type.
|
||||
|
||||
If context is a dict and context_schema is a dataclass or pydantic model, we coerce.
|
||||
Else, we return the context as-is.
|
||||
|
||||
Args:
|
||||
context_schema: The schema type to coerce to (BaseModel, dataclass, or TypedDict)
|
||||
context: The context value to coerce
|
||||
|
||||
Returns:
|
||||
The coerced context value or None if context is None
|
||||
"""
|
||||
if context is None:
|
||||
return None
|
||||
|
||||
if context_schema is None:
|
||||
return context
|
||||
|
||||
schema_is_class = issubclass(context_schema, BaseModel) or is_dataclass(
|
||||
context_schema
|
||||
)
|
||||
if isinstance(context, dict) and schema_is_class:
|
||||
return context_schema(**context) # type: ignore[misc]
|
||||
|
||||
return cast(ContextT, context)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
@@ -106,3 +108,284 @@ def test_runtime_propogated_to_subgraph() -> None:
|
||||
context = Context(username="Alice")
|
||||
result = graph.invoke({}, context=context)
|
||||
assert result == {"subgraph": "Alice!", "main": "Alice!"}
|
||||
|
||||
|
||||
def test_context_coercion_dataclass() -> None:
|
||||
"""Test that dict context is coerced to dataclass."""
|
||||
|
||||
@dataclass
|
||||
class Context:
|
||||
api_key: str
|
||||
timeout: int = 30
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {
|
||||
"message": f"api_key: {runtime.context.api_key}, timeout: {runtime.context.timeout}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test dict coercion with all fields
|
||||
result = compiled.invoke(
|
||||
{"message": "test"}, context={"api_key": "sk_test", "timeout": 60}
|
||||
)
|
||||
assert result == {"message": "api_key: sk_test, timeout: 60"}
|
||||
|
||||
# Test dict coercion with default field
|
||||
result = compiled.invoke({"message": "test"}, context={"api_key": "sk_test2"})
|
||||
assert result == {"message": "api_key: sk_test2, timeout: 30"}
|
||||
|
||||
# Test with actual dataclass instance (should still work)
|
||||
result = compiled.invoke(
|
||||
{"message": "test"}, context=Context(api_key="sk_test3", timeout=90)
|
||||
)
|
||||
assert result == {"message": "api_key: sk_test3, timeout: 90"}
|
||||
|
||||
|
||||
def test_context_coercion_pydantic() -> None:
|
||||
"""Test that dict context is coerced to Pydantic model."""
|
||||
|
||||
class Context(BaseModel):
|
||||
api_key: str
|
||||
timeout: int = 30
|
||||
tags: list[str] = []
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {
|
||||
"message": f"api_key: {runtime.context.api_key}, timeout: {runtime.context.timeout}, tags: {runtime.context.tags}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test dict coercion with all fields
|
||||
result = compiled.invoke(
|
||||
{"message": "test"},
|
||||
context={"api_key": "sk_test", "timeout": 60, "tags": ["prod", "v2"]},
|
||||
)
|
||||
assert result == {"message": "api_key: sk_test, timeout: 60, tags: ['prod', 'v2']"}
|
||||
|
||||
# Test dict coercion with defaults
|
||||
result = compiled.invoke({"message": "test"}, context={"api_key": "sk_test2"})
|
||||
assert result == {"message": "api_key: sk_test2, timeout: 30, tags: []"}
|
||||
|
||||
# Test with actual Pydantic instance (should still work)
|
||||
result = compiled.invoke(
|
||||
{"message": "test"},
|
||||
context=Context(api_key="sk_test3", timeout=90, tags=["test"]),
|
||||
)
|
||||
assert result == {"message": "api_key: sk_test3, timeout: 90, tags: ['test']"}
|
||||
|
||||
|
||||
def test_context_coercion_typeddict() -> None:
|
||||
"""Test that dict context with TypedDict schema passes through as-is."""
|
||||
|
||||
class Context(TypedDict):
|
||||
api_key: str
|
||||
timeout: int
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
# TypedDict context is just a dict at runtime
|
||||
return {
|
||||
"message": f"api_key: {runtime.context['api_key']}, timeout: {runtime.context['timeout']}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test dict passes through for TypedDict
|
||||
result = compiled.invoke(
|
||||
{"message": "test"}, context={"api_key": "sk_test", "timeout": 60}
|
||||
)
|
||||
assert result == {"message": "api_key: sk_test, timeout: 60"}
|
||||
|
||||
|
||||
def test_context_coercion_none() -> None:
|
||||
"""Test that None context is handled properly."""
|
||||
|
||||
@dataclass
|
||||
class Context:
|
||||
api_key: str
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_without_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
# Should be None when no context provided
|
||||
return {"message": f"context is None: {runtime.context is None}"}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_without_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test with None context
|
||||
result = compiled.invoke({"message": "test"}, context=None)
|
||||
assert result == {"message": "context is None: True"}
|
||||
|
||||
# Test without context parameter (defaults to None)
|
||||
result = compiled.invoke({"message": "test"})
|
||||
assert result == {"message": "context is None: True"}
|
||||
|
||||
|
||||
def test_context_coercion_errors() -> None:
|
||||
"""Test error handling for invalid context."""
|
||||
|
||||
@dataclass
|
||||
class Context:
|
||||
api_key: str # Required field
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {"message": "should not reach here"}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test missing required field
|
||||
with pytest.raises(TypeError):
|
||||
compiled.invoke({"message": "test"}, context={"timeout": 60})
|
||||
|
||||
# Test invalid dict keys
|
||||
with pytest.raises(TypeError):
|
||||
compiled.invoke(
|
||||
{"message": "test"}, context={"api_key": "test", "invalid_field": "value"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_context_coercion_async() -> None:
|
||||
"""Test context coercion with async methods."""
|
||||
|
||||
@dataclass
|
||||
class Context:
|
||||
api_key: str
|
||||
async_mode: bool = True
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
async def async_node(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {
|
||||
"message": f"async api_key: {runtime.context.api_key}, async_mode: {runtime.context.async_mode}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", async_node)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test dict coercion with ainvoke
|
||||
result = await compiled.ainvoke(
|
||||
{"message": "test"}, context={"api_key": "sk_async", "async_mode": False}
|
||||
)
|
||||
assert result == {"message": "async api_key: sk_async, async_mode: False"}
|
||||
|
||||
# Test dict coercion with astream
|
||||
chunks = []
|
||||
async for chunk in compiled.astream(
|
||||
{"message": "test"}, context={"api_key": "sk_stream"}
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
# Find the chunk with our node output
|
||||
node_output = None
|
||||
for chunk in chunks:
|
||||
if "node" in chunk:
|
||||
node_output = chunk["node"]
|
||||
break
|
||||
|
||||
assert node_output == {"message": "async api_key: sk_stream, async_mode: True"}
|
||||
|
||||
|
||||
def test_context_coercion_stream() -> None:
|
||||
"""Test context coercion with sync stream method."""
|
||||
|
||||
@dataclass
|
||||
class Context:
|
||||
api_key: str
|
||||
stream_mode: str = "default"
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {
|
||||
"message": f"stream api_key: {runtime.context.api_key}, mode: {runtime.context.stream_mode}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
compiled = graph.compile()
|
||||
|
||||
# Test dict coercion with stream
|
||||
chunks = []
|
||||
for chunk in compiled.stream(
|
||||
{"message": "test"}, context={"api_key": "sk_stream", "stream_mode": "fast"}
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
# Find the chunk with our node output
|
||||
node_output = None
|
||||
for chunk in chunks:
|
||||
if "node" in chunk:
|
||||
node_output = chunk["node"]
|
||||
break
|
||||
|
||||
assert node_output == {"message": "stream api_key: sk_stream, mode: fast"}
|
||||
|
||||
|
||||
def test_context_coercion_pydantic_validation_errors() -> None:
|
||||
"""Test that Pydantic validation errors are raised."""
|
||||
|
||||
class Context(BaseModel):
|
||||
api_key: str
|
||||
timeout: int
|
||||
|
||||
class State(TypedDict):
|
||||
message: str
|
||||
|
||||
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
|
||||
return {
|
||||
"message": f"api_key: {runtime.context.api_key}, timeout: {runtime.context.timeout}"
|
||||
}
|
||||
|
||||
graph = StateGraph(state_schema=State, context_schema=Context)
|
||||
graph.add_node("node", node_with_context)
|
||||
graph.add_edge(START, "node")
|
||||
graph.add_edge("node", END)
|
||||
|
||||
compiled = graph.compile()
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
compiled.invoke(
|
||||
{"message": "test"}, context={"api_key": "sk_test", "timeout": "not_an_int"}
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user