diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index 01b16aa69..49bd34e37 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -321,7 +321,11 @@ class CompiledStateGraph(CompiledGraph): triggers=[START], channels=[START], writers=[ - ChannelWrite(state_write_entries, tags=[TAG_HIDDEN]), + ChannelWrite( + state_write_entries, + tags=[TAG_HIDDEN], + require_at_least_one_of=state_keys, + ), ], ) else: @@ -345,6 +349,7 @@ class CompiledStateGraph(CompiledGraph): ChannelWrite( [ChannelWriteEntry(key, key)] + state_write_entries, tags=[TAG_HIDDEN], + require_at_least_one_of=state_keys, ), ], ).pipe(node) diff --git a/langgraph/pregel/read.py b/langgraph/pregel/read.py index 2e90a505b..a6041d5da 100644 --- a/langgraph/pregel/read.py +++ b/langgraph/pregel/read.py @@ -125,6 +125,7 @@ class PregelNode(RunnableBindingBase): writers[-2] = ChannelWrite( writes=writers[-2].writes + writers[-1].writes, tags=writers[-2].config["tags"] if writers[-2].config else None, + require_at_least_one_of=writers[-2].require_at_least_one_of, ) writers.pop() return writers diff --git a/langgraph/pregel/write.py b/langgraph/pregel/write.py index 5dd2ceca4..82610c322 100644 --- a/langgraph/pregel/write.py +++ b/langgraph/pregel/write.py @@ -43,15 +43,21 @@ class ChannelWrite(RunnableCallable): - runnable to map input, or None to use the input, or any other value to use instead - whether to skip writing if the mapped value is None """ + require_at_least_one_of: Optional[Sequence[str]] + """ + If defined, at least one of these channels must be written to. + """ def __init__( self, writes: Sequence[Union[ChannelWriteEntry, Packet]], *, tags: Optional[list[str]] = None, + require_at_least_one_of: Optional[Sequence[str]] = None, ): super().__init__(func=self._write, afunc=self._awrite, name=None, tags=tags) self.writes = writes + self.require_at_least_one_of = require_at_least_one_of def __repr_args__(self) -> Any: return [("writes", self.writes)] @@ -100,7 +106,11 @@ class ChannelWrite(RunnableCallable): if not write.skip_none or val is not None ] # write packets and values - self.do_write(config, writes + values) + self.do_write( + config, + writes + values, + self.require_at_least_one_of if input is not None else None, + ) return input async def _awrite(self, input: Any, config: RunnableConfig) -> None: @@ -132,13 +142,27 @@ class ChannelWrite(RunnableCallable): if not write.skip_none or val is not None ] # write packets and values - self.do_write(config, writes + values) + self.do_write( + config, + writes + values, + self.require_at_least_one_of if input is not None else None, + ) return input @staticmethod - def do_write(config: RunnableConfig, values: List[Tuple[str, Any]]) -> None: + def do_write( + config: RunnableConfig, + values: List[Tuple[str, Any]], + require_at_least_one_of: Optional[Sequence[str]] = None, + ) -> None: + filtered = [(chan, val) for chan, val in values if val is not SKIP_WRITE] + if require_at_least_one_of is not None: + if not {chan for chan, _ in filtered} & set(require_at_least_one_of): + raise InvalidUpdateError( + f"Must write to at least one of {require_at_least_one_of}" + ) write: TYPE_SEND = config["configurable"][CONFIG_KEY_SEND] - write([(chan, val) for chan, val in values if val is not SKIP_WRITE]) + write(filtered) @staticmethod def is_writer(runnable: Runnable) -> bool: diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 3d18fcd6c..a70533a9f 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -141,6 +141,21 @@ def test_graph_validation() -> None: with pytest.raises(ValueError): # extra is dead-end workflow.compile() + class State(TypedDict): + hello: str + + def node_a(state: State) -> State: + # typo + return {"hel": "world"} + + builder = StateGraph(State) + builder.add_node("a", node_a) + builder.set_entry_point("a") + builder.set_finish_point("a") + graph = builder.compile() + with pytest.raises(InvalidUpdateError): + assert graph.invoke({"hello": "there"}) == {"hello": "world"} + def test_invoke_single_process_in_out(mocker: MockerFixture) -> None: add_one = mocker.Mock(side_effect=lambda x: x + 1)