From 1832ed46eff3fece900b15a5965173e1e71acd48 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Fri, 6 Sep 2024 12:28:10 -0700 Subject: [PATCH] Assert input/output channels aren't managed --- libs/langgraph/langgraph/graph/state.py | 13 ++++-- libs/langgraph/tests/test_state.py | 55 +++++++++++++++++++++++++ 2 files changed, 65 insertions(+), 3 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index baec31cdd..df38e07e2 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -151,8 +151,8 @@ class StateGraph(Graph): self.input = input self.output = output self._add_schema(state_schema) - self._add_schema(input) - self._add_schema(output) + self._add_schema(input, allow_managed=False) + self._add_schema(output, allow_managed=False) self.config_schema = config_schema self.waiting_edges: set[tuple[tuple[str, ...], str]] = set() @@ -162,10 +162,17 @@ class StateGraph(Graph): (start, end) for starts, end in self.waiting_edges for start in starts } - def _add_schema(self, schema: Type[Any]) -> None: + def _add_schema(self, schema: Type[Any], /, allow_managed: bool = True) -> None: if schema not in self.schemas: _warn_invalid_state_schema(schema) channels, managed = _get_channels(schema) + if managed and not allow_managed: + names = ", ".join(managed) + schema_name = getattr(schema, "__name__", "") + raise ValueError( + f"Invalid managed channels detected in {schema_name}: {names}." + " Managed channels are not permited in Input/Output schema." + ) self.schemas[schema] = {**channels, **managed} for key, channel in channels.items(): if key in self.channels: diff --git a/libs/langgraph/tests/test_state.py b/libs/langgraph/tests/test_state.py index 0234c5e2f..670a2c2e8 100644 --- a/libs/langgraph/tests/test_state.py +++ b/libs/langgraph/tests/test_state.py @@ -10,6 +10,7 @@ from pydantic.v1 import BaseModel from typing_extensions import Annotated, NotRequired, Required, TypedDict from langgraph.graph.state import StateGraph, _warn_invalid_state_schema +from langgraph.managed.shared_value import SharedValue class State(BaseModel): @@ -180,3 +181,57 @@ def test_state_schema_default_values(kw_only_: bool): assert ( set(json_schema["properties"].keys()) == expected_required | expected_optional ) + + +def test_raises_invalid_managed(): + class BadInputState(TypedDict): + some_thing: str + some_input_channel: Annotated[str, SharedValue.on("assistant_id")] + + class InputState(TypedDict): + some_thing: str + some_input_channel: str + + class BadOutputState(TypedDict): + some_thing: str + some_output_channel: Annotated[str, SharedValue.on("assistant_id")] + + class OutputState(TypedDict): + some_thing: str + some_output_channel: str + + class State(TypedDict): + some_thing: str + some_channel: Annotated[str, SharedValue.on("assistant_id")] + + # All OK + StateGraph(State, input=InputState, output=OutputState) + StateGraph(State) + StateGraph(State, input=State, output=State) + StateGraph(State, input=InputState) + StateGraph(State, input=InputState) + + bad_input_examples = [ + (State, BadInputState, OutputState), + (State, BadInputState, BadOutputState), + (State, BadInputState, State), + (State, BadInputState, None), + ] + for _state, _inp, _outp in bad_input_examples: + with pytest.raises( + ValueError, + match="Invalid managed channels detected in BadInputState: some_input_channel. Managed channels are not permited in Input/Output schema.", + ): + StateGraph(_state, input=_inp, output=_outp) + bad_output_examples = [ + (State, InputState, BadOutputState), + (None, InputState, BadOutputState), + (None, State, BadOutputState), + (State, None, BadOutputState), + ] + for _state, _inp, _outp in bad_output_examples: + with pytest.raises( + ValueError, + match="Invalid managed channels detected in BadOutputState: some_output_channel. Managed channels are not permited in Input/Output schema.", + ): + StateGraph(_state, input=_inp, output=_outp)