diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 518c73670..2c9c729ac 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -65,7 +65,6 @@ class StateNodeSpec(NamedTuple): runnable: Runnable metadata: dict[str, Any] input: Type[Any] - output: Type[Any] class StateGraph(Graph): @@ -136,6 +135,7 @@ class StateGraph(Graph): if state_schema is None: if input is None or output is None: raise ValueError("Must provide state_schema or input and output") + state_schema = input else: if input is None: input = state_schema @@ -167,10 +167,12 @@ class StateGraph(Graph): for key, channel in channels.items(): if key in self.channels: if self.channels[key] != channel: - print(self.channels[key], channel) - raise ValueError( - f"Channel '{key}' already exists with a different type" - ) + if isinstance(channel, LastValue): + pass + else: + raise ValueError( + f"Channel '{key}' already exists with a different type" + ) else: self.channels[key] = channel for key, managed in managed.items(): @@ -193,7 +195,6 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, - output: Optional[Type[Any]] = None, ) -> None: """Adds a new node to the state graph. Will take the name of the function/runnable as the node name. @@ -217,7 +218,6 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, - output: Optional[Type[Any]] = None, ) -> None: """Adds a new node to the state graph. @@ -240,7 +240,6 @@ class StateGraph(Graph): *, metadata: Optional[dict[str, Any]] = None, input: Optional[Type[Any]] = None, - output: Optional[Type[Any]] = None, ) -> None: """Adds a new node to the state graph. @@ -249,6 +248,8 @@ class StateGraph(Graph): Args: node (Union[str, RunnableLike)]: The function or runnable this node will run. action (Optional[RunnableLike]): The action associated with the node. (default: None) + metadata (Optional[dict[str, Any]]): The metadata associated with the node. (default: None) + input (Optional[Type[Any]]): The input schema for the node. (default: the graph's input schema) Raises: ValueError: If the key is already being used as a state key. @@ -313,21 +314,14 @@ class StateGraph(Graph): input_hint = hints[list(hints.keys())[0]] if isinstance(input_hint, type) and get_type_hints(input_hint): input = input_hint - if output is None: - output_hint = hints.get("return", Any) - if isinstance(output_hint, type) and get_type_hints(output_hint): - output = output_hint except TypeError: pass if input is not None: self._add_schema(input) - if output is not None: - self._add_schema(output) self.nodes[node] = StateNodeSpec( coerce_to_runnable(action, name=node, trace=False), metadata, input=input or self.schema, - output=output or self.schema, ) def add_edge(self, start_key: Union[str, list[str]], end_key: str) -> None: @@ -482,24 +476,17 @@ class CompiledStateGraph(CompiledGraph): def attach_node(self, key: str, node: Optional[StateNodeSpec]) -> None: if key == START: - input_schema = self.builder.input + output_keys = [ + k + for k, v in self.builder.schemas[self.builder.input].items() + if not isinstance(v, Context) and not is_managed_value(v) + ] else: - input_schema = node.input if node else self.builder.schema - input_values = { - k: v if is_managed_value(v) else k - for k, v in self.builder.schemas[input_schema].items() - } - is_single_input = len(input_values) == 1 and "__root__" in input_values + output_keys = list(self.builder.channels) - output_keys = [ - k - for k, v in self.builder.schemas[ - node.output if node else self.builder.schema - ].items() - if not is_managed_value(v) - ] - - def _get_state_key(input: dict, config: RunnableConfig, *, key: str) -> Any: + def _get_state_key( + input: Union[None, dict, Any], config: RunnableConfig, *, key: str + ) -> Any: if input is None: return SKIP_WRITE elif isinstance(input, dict): @@ -540,6 +527,13 @@ class CompiledStateGraph(CompiledGraph): ], ) else: + input_schema = node.input if node else self.builder.schema + input_values = { + k: v if is_managed_value(v) else k + for k, v in self.builder.schemas[input_schema].items() + } + is_single_input = len(input_values) == 1 and "__root__" in input_values + self.channels[key] = EphemeralValue(Any, guard=False) self.nodes[key] = PregelNode( triggers=[], diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 7fdd55913..fc5cddb87 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -232,7 +232,7 @@ def test_checkpoint_errors() -> None: graph.invoke("", {"configurable": {"thread_id": "thread-1"}}) -def test_node_schemas() -> None: +def test_node_schemas_custom_output() -> None: from langchain_core.messages import HumanMessage class State(TypedDict): @@ -240,6 +240,9 @@ def test_node_schemas() -> None: bye: str messages: Annotated[list[str], add_messages] + class Output(TypedDict): + messages: list[str] + class StateForA(TypedDict): hello: str messages: Annotated[list[str], add_messages] @@ -254,14 +257,14 @@ def test_node_schemas() -> None: bye: str now: int - def node_b(state: StateForB) -> StateForB: + def node_b(state: StateForB): assert state == { "bye": "world", "now": None, } return { "now": 123, - "hello": "again", # ignored because not in output schema + "hello": "again", } class StateForC(TypedDict): @@ -270,11 +273,11 @@ def test_node_schemas() -> None: def node_c(state: StateForC) -> StateForC: assert state == { - "hello": "there", + "hello": "again", "now": 123, } - builder = StateGraph(State) + builder = StateGraph(State, output=Output) builder.add_node("a", node_a) builder.add_node("b", node_b) builder.add_node("c", node_c) @@ -284,8 +287,26 @@ def test_node_schemas() -> None: graph = builder.compile() assert graph.invoke({"hello": "there", "bye": "world", "messages": "hello"}) == { - "hello": "there", - "bye": "world", + "messages": [HumanMessage(content="hello", id=AnyStr())], + } + + builder = StateGraph(input=State, output=Output) + builder.add_node("a", node_a) + builder.add_node("b", node_b) + builder.add_node("c", node_c) + builder.add_edge(START, "a") + builder.add_edge("a", "b") + builder.add_edge("b", "c") + graph = builder.compile() + + assert graph.invoke( + { + "hello": "there", + "bye": "world", + "messages": "hello", + "now": 345, # ignored because not in input schema + } + ) == { "messages": [HumanMessage(content="hello", id=AnyStr())], } diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 47a9c0b4f..b7d84766c 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -389,7 +389,7 @@ async def test_cancel_graph_astream_events_v2( await checkpointer.__aexit__(None, None, None) -async def test_node_schemas() -> None: +async def test_node_schemas_custom_output() -> None: from langchain_core.messages import HumanMessage class State(TypedDict): @@ -397,11 +397,14 @@ async def test_node_schemas() -> None: bye: str messages: Annotated[list[str], add_messages] + class Output(TypedDict): + messages: list[str] + class StateForA(TypedDict): hello: str messages: Annotated[list[str], add_messages] - async def node_a(state: StateForA) -> State: + async def node_a(state: StateForA): assert state == { "hello": "there", "messages": [HumanMessage(content="hello", id=AnyStr())], @@ -411,27 +414,27 @@ async def test_node_schemas() -> None: bye: str now: int - async def node_b(state: StateForB) -> StateForB: + async def node_b(state: StateForB): assert state == { "bye": "world", "now": None, } return { "now": 123, - "hello": "again", # ignored because not in output schema + "hello": "again", } class StateForC(TypedDict): hello: str now: int - async def node_c(state: StateForC) -> StateForC: + async def node_c(state: StateForC): assert state == { - "hello": "there", + "hello": "again", "now": 123, } - builder = StateGraph(State) + builder = StateGraph(State, output=Output) builder.add_node("a", node_a) builder.add_node("b", node_b) builder.add_node("c", node_c) @@ -443,8 +446,26 @@ async def test_node_schemas() -> None: assert await graph.ainvoke( {"hello": "there", "bye": "world", "messages": "hello"} ) == { - "hello": "there", - "bye": "world", + "messages": [HumanMessage(content="hello", id=AnyStr())], + } + + builder = StateGraph(input=State, output=Output) + builder.add_node("a", node_a) + builder.add_node("b", node_b) + builder.add_node("c", node_c) + builder.add_edge(START, "a") + builder.add_edge("a", "b") + builder.add_edge("b", "c") + graph = builder.compile() + + assert await graph.ainvoke( + { + "hello": "there", + "bye": "world", + "messages": "hello", + "now": 345, # ignored because not in input schema + } + ) == { "messages": [HumanMessage(content="hello", id=AnyStr())], }