Update test

This commit is contained in:
William Fu-Hinthorn
2024-09-06 12:09:37 -07:00
parent fd36ed794c
commit 093dab1035
+46 -20
View File
@@ -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
)