diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 1416d14d0..f576a98cd 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -20,9 +20,7 @@ from typing import ( from langchain_core.pydantic_v1 import BaseModel from langchain_core.runnables import Runnable, RunnableConfig from langchain_core.runnables.base import RunnableLike -from langchain_core.runnables.utils import ( - create_model, -) +from langchain_core.runnables.utils import create_model from langgraph.channels.base import BaseChannel from langgraph.channels.binop import BinaryOperatorAggregate @@ -33,14 +31,7 @@ from langgraph.channels.named_barrier_value import NamedBarrierValue from langgraph.checkpoint.base import BaseCheckpointSaver from langgraph.constants import NS_END, NS_SEP, TAG_HIDDEN from langgraph.errors import InvalidUpdateError -from langgraph.graph.graph import ( - END, - START, - Branch, - CompiledGraph, - Graph, - Send, -) +from langgraph.graph.graph import END, START, Branch, CompiledGraph, Graph, Send from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, @@ -53,7 +44,7 @@ from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore -from langgraph.utils import RunnableCallable, coerce_to_runnable +from langgraph.utils import RunnableCallable, coerce_to_runnable, get_field_default logger = logging.getLogger(__name__) @@ -498,7 +489,16 @@ class CompiledStateGraph(CompiledGraph): return create_model( # type: ignore[call-overload] self.get_name("Input"), **{ - k: (self.channels[k].UpdateType, None) + k: ( + self.channels[k].UpdateType, + ( + get_field_default( + k, + self.channels[k].UpdateType, + self.builder.input, + ) + ), + ) for k in self.builder.schemas[self.builder.input] if isinstance(self.channels[k], BaseChannel) }, diff --git a/libs/langgraph/langgraph/utils.py b/libs/langgraph/langgraph/utils.py index 64cb21128..cc3fd4760 100644 --- a/libs/langgraph/langgraph/utils.py +++ b/libs/langgraph/langgraph/utils.py @@ -4,7 +4,7 @@ import inspect import sys from contextvars import copy_context from functools import partial, wraps -from typing import Any, AsyncIterator, Awaitable, Callable, Optional +from typing import Any, AsyncIterator, Awaitable, Callable, Optional, Type, Union from langchain_core.runnables.base import ( Runnable, @@ -19,7 +19,14 @@ from langchain_core.runnables.config import ( var_child_runnable_config, ) from langchain_core.runnables.utils import accepts_config -from typing_extensions import TypeGuard +from typing_extensions import ( + Annotated, + NotRequired, + ReadOnly, + Required, + TypeGuard, + get_origin, +) try: from langchain_core.runnables.config import _set_config_context @@ -34,8 +41,6 @@ except ImportError: class StrEnum(str, enum.Enum): """A string enum.""" - pass - class RunnableCallable(Runnable): """A much simpler version of RunnableLambda that requires sync and async functions.""" @@ -183,3 +188,95 @@ def coerce_to_runnable(thing: RunnableLike, *, name: str, trace: bool) -> Runnab f"Expected a Runnable, callable or dict." f"Instead got an unsupported type: {type(thing)}" ) + + +def _is_optional_type(type_: Any) -> bool: + """Check if a type is Optional.""" + + if hasattr(type_, "__origin__") and hasattr(type_, "__args__"): + origin = get_origin(type_) + if origin is Optional: + return True + if origin is Union: + return any( + arg is type(None) or _is_optional_type(arg) for arg in type_.__args__ + ) + if origin is Annotated: + return _is_optional_type(type_.__args__[0]) + return origin is None + if hasattr(type_, "__bound__") and type_.__bound__ is not None: + return _is_optional_type(type_.__bound__) + return type_ is None + + +def _is_required_type(type_: Any) -> Optional[bool]: + """Check if an annotation is marked as Required/NotRequired. + + Returns: + - True if required + - False if not required + - None if not annotated with either + """ + origin = get_origin(type_) + if origin is Required: + return True + if origin is NotRequired: + return False + if origin is Annotated or getattr(origin, "__args__", None): + # See https://typing.readthedocs.io/en/latest/spec/typeddict.html#interaction-with-annotated + return _is_required_type(type_.__args__[0]) + return None + + +def _is_readonly_type(type_: Any) -> bool: + """Check if an annotation is marked as ReadOnly. + + Returns: + - True if is read only + - False if not read only + """ + + # See: https://typing.readthedocs.io/en/latest/spec/typeddict.html#typing-readonly-type-qualifier + origin = get_origin(type_) + if origin is Annotated: + return _is_readonly_type(type_.__args__[0]) + if origin is ReadOnly: + return True + return False + + +_DEFAULT_KEYS = frozenset() + + +def get_field_default(name: str, type_: Any, schema: Type[Any]) -> Any: + """Determine the default value for a field in a state schema. + + This is based on: + If TypedDict: + - Required/NotRequired + - total=False -> everything optional + - Type annotation (Optional/Union[None]) + """ + optional_keys = getattr(schema, "__optional_keys__", _DEFAULT_KEYS) + irq = _is_required_type(type_) + if name in optional_keys: + # Either total=False or explicit NotRequired. + # No type annotation trumps this. + if irq: + # Unless it's earlier versions of python & explicit Required + return ... + return None + if irq is not None: + if irq: + # Handle Required[] + # (we already handled NotRequired and total=False) + return ... + # Handle NotRequired[] for earlier versions of python + return None + # Note, we ignore ReadOnly attributes, + # as they don't make much sense. (we don't care if you mutate the state in your node) + # and mutating state in your node has no effect on our graph state. + # Base case is the annotation + if _is_optional_type(type_): + return None + return ... diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index 14cf25277..3c3dd49d0 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -366,7 +366,7 @@ ''' # --- # name: test_conditional_entrypoint_to_multiple_state_graph - '{"title": "LangGraphInput", "type": "object", "properties": {"locations": {"title": "Locations", "type": "array", "items": {"type": "string"}}, "results": {"title": "Results", "type": "array", "items": {"type": "string"}}}}' + '{"title": "LangGraphInput", "type": "object", "properties": {"locations": {"title": "Locations", "type": "array", "items": {"type": "string"}}, "results": {"title": "Results", "type": "array", "items": {"type": "string"}}}, "required": ["locations", "results"]}' # --- # name: test_conditional_entrypoint_to_multiple_state_graph.1 '{"title": "LangGraphOutput", "type": "object", "properties": {"locations": {"title": "Locations", "type": "array", "items": {"type": "string"}}, "results": {"title": "Results", "type": "array", "items": {"type": "string"}}}}' @@ -4855,7 +4855,7 @@ ''' # --- # name: test_prebuilt_tool_chat - '{"title": "LangGraphInput", "type": "object", "properties": {"messages": {"title": "Messages", "type": "array", "items": {"$ref": "#/definitions/BaseMessage"}}}, "definitions": {"BaseMessage": {"title": "BaseMessage", "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}}, "required": ["content", "type"]}}}' + '{"title": "LangGraphInput", "type": "object", "properties": {"messages": {"title": "Messages", "type": "array", "items": {"$ref": "#/definitions/BaseMessage"}}}, "required": ["messages"], "definitions": {"BaseMessage": {"title": "BaseMessage", "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}}, "required": ["content", "type"]}}}' # --- # name: test_prebuilt_tool_chat.1 '{"title": "LangGraphOutput", "type": "object", "properties": {"messages": {"title": "Messages", "type": "array", "items": {"$ref": "#/definitions/BaseMessage"}}}, "definitions": {"BaseMessage": {"title": "BaseMessage", "description": "Base abstract message class.\\n\\nMessages are the inputs and outputs of ChatModels.", "type": "object", "properties": {"content": {"title": "Content", "anyOf": [{"type": "string"}, {"type": "array", "items": {"anyOf": [{"type": "string"}, {"type": "object"}]}}]}, "additional_kwargs": {"title": "Additional Kwargs", "type": "object"}, "response_metadata": {"title": "Response Metadata", "type": "object"}, "type": {"title": "Type", "type": "string"}, "name": {"title": "Name", "type": "string"}, "id": {"title": "Id", "type": "string"}}, "required": ["content", "type"]}}}' @@ -5094,7 +5094,7 @@ '{"title": "LangGraphConfig", "type": "object", "properties": {"configurable": {"$ref": "#/definitions/Configurable"}}, "definitions": {"Configurable": {"title": "Configurable", "type": "object", "properties": {"tools": {"title": "Tools", "type": "array", "items": {"type": "string"}}}}}}' # --- # name: test_state_graph_w_config_inherited_state_keys.1 - '{"title": "LangGraphInput", "type": "object", "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"title": "Agent Outcome", "anyOf": [{"$ref": "#/definitions/AgentAction"}, {"$ref": "#/definitions/AgentFinish"}]}, "intermediate_steps": {"title": "Intermediate Steps", "type": "array", "items": {"type": "array", "minItems": 2, "maxItems": 2, "items": [{"$ref": "#/definitions/AgentAction"}, {"type": "string"}]}}}, "definitions": {"AgentAction": {"title": "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.", "type": "object", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"title": "Tool Input", "anyOf": [{"type": "string"}, {"type": "object"}]}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentAction", "enum": ["AgentAction"], "type": "string"}}, "required": ["tool", "tool_input", "log"]}, "AgentFinish": {"title": "AgentFinish", "description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "type": "object", "properties": {"return_values": {"title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentFinish", "enum": ["AgentFinish"], "type": "string"}}, "required": ["return_values", "log"]}}}' + '{"title": "LangGraphInput", "type": "object", "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"title": "Agent Outcome", "anyOf": [{"$ref": "#/definitions/AgentAction"}, {"$ref": "#/definitions/AgentFinish"}]}, "intermediate_steps": {"title": "Intermediate Steps", "type": "array", "items": {"type": "array", "minItems": 2, "maxItems": 2, "items": [{"$ref": "#/definitions/AgentAction"}, {"type": "string"}]}}}, "required": ["input"], "definitions": {"AgentAction": {"title": "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.", "type": "object", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"title": "Tool Input", "anyOf": [{"type": "string"}, {"type": "object"}]}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentAction", "enum": ["AgentAction"], "type": "string"}}, "required": ["tool", "tool_input", "log"]}, "AgentFinish": {"title": "AgentFinish", "description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "type": "object", "properties": {"return_values": {"title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentFinish", "enum": ["AgentFinish"], "type": "string"}}, "required": ["return_values", "log"]}}}' # --- # name: test_state_graph_w_config_inherited_state_keys.2 '{"title": "LangGraphOutput", "type": "object", "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"title": "Agent Outcome", "anyOf": [{"$ref": "#/definitions/AgentAction"}, {"$ref": "#/definitions/AgentFinish"}]}, "intermediate_steps": {"title": "Intermediate Steps", "type": "array", "items": {"type": "array", "minItems": 2, "maxItems": 2, "items": [{"$ref": "#/definitions/AgentAction"}, {"type": "string"}]}}}, "definitions": {"AgentAction": {"title": "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.", "type": "object", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"title": "Tool Input", "anyOf": [{"type": "string"}, {"type": "object"}]}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentAction", "enum": ["AgentAction"], "type": "string"}}, "required": ["tool", "tool_input", "log"]}, "AgentFinish": {"title": "AgentFinish", "description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "type": "object", "properties": {"return_values": {"title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"title": "Type", "default": "AgentFinish", "enum": ["AgentFinish"], "type": "string"}}, "required": ["return_values", "log"]}}}' diff --git a/libs/langgraph/tests/test_state.py b/libs/langgraph/tests/test_state.py index 392a3f22a..12cd7f653 100644 --- a/libs/langgraph/tests/test_state.py +++ b/libs/langgraph/tests/test_state.py @@ -1,11 +1,11 @@ import warnings from typing import Annotated as Annotated2 -from typing import Any +from typing import Any, Optional import pytest from langchain_core.runnables import RunnableConfig from pydantic.v1 import BaseModel -from typing_extensions import Annotated, TypedDict +from typing_extensions import Annotated, NotRequired, Required, TypedDict from langgraph.graph.state import StateGraph, _warn_invalid_state_schema @@ -88,3 +88,45 @@ def test_state_schema_with_type_hint(): for i, c in enumerate(graph.stream(input_state, stream_mode="updates")): node_name = actions[i].__name__ assert c[node_name] == output_state + + +@pytest.mark.parametrize("total_", [True, False]) +def test_state_schema_optional_values(total_: bool): + class SomeParentState(TypedDict): + val0a: str + val0b: Optional[str] + + class InputState(SomeParentState, total=total_): # type: ignore + val1: str + val2: Optional[str] + val3: Required[str] + val4: NotRequired[dict] + val5: Annotated[Required[str], "foo"] + val6: Annotated[NotRequired[str], "bar"] + + class State(InputState): # this would be ignored + val4: dict + + builder = StateGraph(State, input=InputState) + builder.add_node("n", lambda x: x) + builder.add_edge("__start__", "n") + graph = builder.compile() + model = graph.input_schema + json_schema = model.schema() + + if total_ is False: + expected_required = set() + expected_optional = {"val2", "val1"} + else: + expected_required = {"val1"} + + expected_optional = {"val2"} + + # The others should always have precedence based on the required annotation + expected_required |= {"val0a", "val3", "val5"} + expected_optional |= {"val0b", "val4", "val6"} + + assert set(json_schema.get("required", set())) == expected_required + assert ( + set(json_schema["properties"].keys()) == expected_required | expected_optional + ) diff --git a/libs/langgraph/tests/test_utils.py b/libs/langgraph/tests/test_utils.py index 3a914c23b..bee8fcd9c 100644 --- a/libs/langgraph/tests/test_utils.py +++ b/libs/langgraph/tests/test_utils.py @@ -1,15 +1,32 @@ import functools import sys import uuid -from typing import TypedDict +from typing import ( + Any, + Callable, + Dict, + ForwardRef, + List, + Literal, + Optional, + TypedDict, + TypeVar, + Union, +) from unittest.mock import patch import langsmith import pytest +from typing_extensions import Annotated, NotRequired, Required from langgraph.graph import END, StateGraph from langgraph.graph.graph import CompiledGraph -from langgraph.utils import is_async_callable, is_async_generator +from langgraph.utils import ( + _is_optional_type, + get_field_default, + is_async_callable, + is_async_generator, +) pytestmark = pytest.mark.anyio @@ -121,3 +138,96 @@ async def test_runnable_callable_tracing_nested_async(rt_graph: CompiledGraph) - with langsmith.tracing_context(enabled=True): res = await rt_graph.ainvoke({"foo": 1}) assert isinstance(res["node_run_id"], uuid.UUID) + + +def test_is_optional_type(): + assert _is_optional_type(None) + assert not _is_optional_type(type(None)) + assert _is_optional_type(Optional[list]) + assert not _is_optional_type(int) + assert _is_optional_type(Optional[Literal[1, 2, 3]]) + assert not _is_optional_type(Literal[1, 2, 3]) + assert _is_optional_type(Optional[List[int]]) + assert _is_optional_type(Optional[Dict[str, int]]) + assert not _is_optional_type(List[Optional[int]]) + assert _is_optional_type(Union[Optional[str], Optional[int]]) + assert _is_optional_type( + Union[ + Union[Optional[str], Optional[int]], Union[Optional[float], Optional[dict]] + ] + ) + assert not _is_optional_type(Union[Union[str, int], Union[float, dict]]) + + assert _is_optional_type(Union[int, None]) + assert _is_optional_type(Union[str, None, int]) + assert _is_optional_type(Union[None, str, int]) + assert not _is_optional_type(Union[int, str]) + + assert not _is_optional_type(Any) # Do we actually want this? + assert _is_optional_type(Optional[Any]) + + class MyClass: + pass + + assert _is_optional_type(Optional[MyClass]) + assert not _is_optional_type(MyClass) + assert _is_optional_type(Optional[ForwardRef("MyClass")]) + assert not _is_optional_type(ForwardRef("MyClass")) + + assert _is_optional_type(Optional[Union[List[int], Dict[str, Optional[int]]]]) + assert not _is_optional_type(Union[List[int], Dict[str, Optional[int]]]) + + assert _is_optional_type(Optional[Callable[[int], str]]) + assert not _is_optional_type(Callable[[int], Optional[str]]) + + T = TypeVar("T") + assert _is_optional_type(Optional[T]) + assert not _is_optional_type(T) + + U = TypeVar("U", bound=Optional[T]) # type: ignore + assert _is_optional_type(U) + + +def test_is_required(): + class MyBaseTypedDict(TypedDict): + val_1: Required[Optional[str]] + val_2: Required[str] + val_3: NotRequired[str] + val_4: NotRequired[Optional[str]] + val_5: Annotated[NotRequired[int], "foo"] + val_6: NotRequired[Annotated[int, "foo"]] + val_7: Annotated[Required[int], "foo"] + val_8: Required[Annotated[int, "foo"]] + val_9: Optional[str] + val_10: str + + annos = MyBaseTypedDict.__annotations__ + assert get_field_default("val_1", annos["val_1"], MyBaseTypedDict) == ... + assert get_field_default("val_2", annos["val_2"], MyBaseTypedDict) == ... + assert get_field_default("val_3", annos["val_3"], MyBaseTypedDict) is None + assert get_field_default("val_4", annos["val_4"], MyBaseTypedDict) is None + # See https://peps.python.org/pep-0655/#interaction-with-annotated + assert get_field_default("val_5", annos["val_5"], MyBaseTypedDict) is None + assert get_field_default("val_6", annos["val_6"], MyBaseTypedDict) is None + assert get_field_default("val_7", annos["val_7"], MyBaseTypedDict) == ... + assert get_field_default("val_8", annos["val_8"], MyBaseTypedDict) == ... + assert get_field_default("val_9", annos["val_9"], MyBaseTypedDict) is None + assert get_field_default("val_10", annos["val_10"], MyBaseTypedDict) == ... + + class MyChildDict(MyBaseTypedDict): + val_11: int + val_11b: Optional[int] + val_11c: Union[int, None, str] + + class MyGrandChildDict(MyChildDict, total=False): + val_12: int + val_13: Required[str] + + cannos = MyChildDict.__annotations__ + gcannos = MyGrandChildDict.__annotations__ + assert get_field_default("val_11", cannos["val_11"], MyChildDict) == ... + assert get_field_default("val_11b", cannos["val_11b"], MyChildDict) is None + assert get_field_default("val_11c", cannos["val_11c"], MyChildDict) is None + assert get_field_default("val_12", gcannos["val_12"], MyGrandChildDict) is None + assert get_field_default("val_9", gcannos["val_9"], MyGrandChildDict) is None + assert get_field_default("val_13", gcannos["val_13"], MyGrandChildDict) == ...