diff --git a/libs/langgraph/langgraph/pregel/main.py b/libs/langgraph/langgraph/pregel/main.py index 37bd069c2..eeea9083c 100644 --- a/libs/langgraph/langgraph/pregel/main.py +++ b/libs/langgraph/langgraph/pregel/main.py @@ -9,8 +9,8 @@ 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 functools import lru_cache, partial +from inspect import Parameter, isclass, signature from typing import Any, Callable, Generic, Union, cast, get_type_hints from uuid import UUID, uuid5 @@ -3232,6 +3232,15 @@ def _output( yield payload +@lru_cache(maxsize=128) +def _get_schema_fields(context_schema: type) -> set[str] | None: + # Must be dataclass right now. + sig = signature(context_schema) + if any(p.kind == Parameter.VAR_KEYWORD for p in sig.parameters.values()): + return None + return set(get_type_hints(context_schema).keys()) + + def _coerce_context( context_schema: type[ContextT] | None, context: Any ) -> ContextT | None: @@ -3253,10 +3262,17 @@ def _coerce_context( 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] + if isinstance(context, dict): + if is_typeddict(context_schema): + return cast(ContextT, context) + + if is_dataclass(context_schema): + fields = _get_schema_fields(context_schema) + if fields is None: + return context_schema(**context) + return context_schema(**{k: v for k, v in context.items() if k in fields}) + + if issubclass(context_schema, BaseModel): + return context_schema(**context) return cast(ContextT, context) diff --git a/libs/langgraph/tests/test_runtime.py b/libs/langgraph/tests/test_runtime.py index 0407b84d2..9761b1e7e 100644 --- a/libs/langgraph/tests/test_runtime.py +++ b/libs/langgraph/tests/test_runtime.py @@ -52,6 +52,14 @@ def test_context_runtime() -> None: {"message": "hello world"}, context=Context(api_key="sk_123456") ) assert result == {"message": "api key: sk_123456"} + result = compiled.invoke( + {"message": "hello world"}, context={"api_key": "sk_1234567"} + ) + assert result == {"message": "api key: sk_1234567"} + result = compiled.invoke( + {"message": "hello world"}, context={"api_key": "sk_1234568", "foo": "bar"} + ) + assert result == {"message": "api key: sk_1234568"} def test_override_runtime() -> None: @@ -134,7 +142,8 @@ def test_context_coercion_dataclass() -> None: # Test dict coercion with all fields result = compiled.invoke( - {"message": "test"}, context={"api_key": "sk_test", "timeout": 60} + {"message": "test"}, + context={"api_key": "sk_test", "timeout": 60, "extra_arg": "extra_value"}, ) assert result == {"message": "api_key: sk_test, timeout: 60"} @@ -149,6 +158,53 @@ def test_context_coercion_dataclass() -> None: assert result == {"message": "api_key: sk_test3, timeout: 90"} +def test_context_coercion_dataclass_custom_init() -> None: + """Test that dict context is coerced to dataclass.""" + + @dataclass(init=False) + class Context: + api_key: str + timeout: int = 30 + + def __init__(self, api_key: str, timeout: int = 30, **kwargs: Any): + self.api_key = api_key + self.timeout = timeout + setattr(self, "extra_arg", kwargs.get("extra_arg")) + + 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}, extra_arg: {runtime.context.extra_arg}" + } + + 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, "extra_arg": "extra_value"}, + ) + assert result == { + "message": "api_key: sk_test, timeout: 60, extra_arg: extra_value" + } + + # 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, extra_arg: None"} + + # 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, extra_arg: None"} + + def test_context_coercion_pydantic() -> None: """Test that dict context is coerced to Pydantic model.""" @@ -174,7 +230,12 @@ def test_context_coercion_pydantic() -> None: # Test dict coercion with all fields result = compiled.invoke( {"message": "test"}, - context={"api_key": "sk_test", "timeout": 60, "tags": ["prod", "v2"]}, + context={ + "api_key": "sk_test", + "timeout": 60, + "tags": ["prod", "v2"], + "extra_arg": "extra_value", + }, ) assert result == {"message": "api_key: sk_test, timeout: 60, tags: ['prod', 'v2']"} @@ -214,7 +275,8 @@ def test_context_coercion_typeddict() -> None: # Test dict passes through for TypedDict result = compiled.invoke( - {"message": "test"}, context={"api_key": "sk_test", "timeout": 60} + {"message": "test"}, + context={"api_key": "sk_test", "timeout": 60, "extra_arg": "extra_value"}, ) assert result == {"message": "api_key: sk_test, timeout: 60"} @@ -271,12 +333,6 @@ def test_context_coercion_errors() -> None: 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: