diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index adc326ea0..ac7d92bd0 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -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, @@ -3224,3 +3224,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) diff --git a/libs/langgraph/tests/test_runtime.py b/libs/langgraph/tests/test_runtime.py index 536853109..0407b84d2 100644 --- a/libs/langgraph/tests/test_runtime.py +++ b/libs/langgraph/tests/test_runtime.py @@ -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"} + )