diff --git a/libs/langgraph/tests/test_state.py b/libs/langgraph/tests/test_state.py index 0234c5e2f..fcd28cb56 100644 --- a/libs/langgraph/tests/test_state.py +++ b/libs/langgraph/tests/test_state.py @@ -106,14 +106,22 @@ def test_state_schema_optional_values(total_: bool): val5: Annotated[Required[str], "foo"] val6: Annotated[NotRequired[str], "bar"] + class OutputState(SomeParentState, total=total_): # type: ignore + out_val1: str + out_val2: Optional[str] + out_val3: Required[str] + out_val4: NotRequired[dict] + out_val5: Annotated[Required[str], "foo"] + out_val6: Annotated[NotRequired[str], "bar"] + class State(InputState): # this would be ignored val4: dict - builder = StateGraph(State, input=InputState) + builder = StateGraph(State, input=InputState, output=OutputState) builder.add_node("n", lambda x: x) builder.add_edge("__start__", "n") graph = builder.compile() - model = graph.input_schema + model = graph.get_input_schema() json_schema = model.schema() if total_ is False: @@ -133,6 +141,23 @@ def test_state_schema_optional_values(total_: bool): set(json_schema["properties"].keys()) == expected_required | expected_optional ) + # Check output schema. Should be the same process + output_schema = graph.get_output_schema().schema() + if total_ is False: + expected_required = set() + expected_optional = {"out_val2", "out_val1"} + else: + expected_required = {"out_val1"} + expected_optional = {"out_val2"} + + expected_required |= {"val0a", "out_val3", "out_val5"} + expected_optional |= {"val0b", "out_val4", "out_val6"} + + assert set(output_schema.get("required", set())) == expected_required + assert ( + set(output_schema["properties"].keys()) == expected_required | expected_optional + ) + @pytest.mark.parametrize("kw_only_", [False, True]) def test_state_schema_default_values(kw_only_: bool): @@ -160,23 +185,24 @@ def test_state_schema_default_values(kw_only_: bool): builder.add_node("n", lambda x: x) builder.add_edge("__start__", "n") graph = builder.compile() - model = graph.input_schema - json_schema = model.schema() + for model in [graph.get_input_schema(), graph.get_output_schema()]: + json_schema = model.schema() - expected_required = {"val1", "val7"} - expected_optional = { - "val2", - "val3", - "val4", - "val5", - "val6", - "val8", - "val9", - "val10", - "val11", - } + expected_required = {"val1", "val7"} + expected_optional = { + "val2", + "val3", + "val4", + "val5", + "val6", + "val8", + "val9", + "val10", + "val11", + } - assert set(json_schema.get("required", set())) == expected_required - assert ( - set(json_schema["properties"].keys()) == expected_required | expected_optional - ) + assert set(json_schema.get("required", set())) == expected_required + assert ( + set(json_schema["properties"].keys()) + == expected_required | expected_optional + )