mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-24 08:32:24 +02:00
Update test
This commit is contained in:
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user